Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); Adding documents to TE/JAX by mingxu1067 · Pull Request #87 · NVIDIA/TransformerEngine · GitHub
Skip to content

Adding documents to TE/JAX - #87

Merged
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs
Mar 14, 2023
Merged

Adding documents to TE/JAX#87
ksivaman merged 12 commits into
NVIDIA:mainfrom
mingxu1067:mingh/te_docs

Conversation

@mingxu1067

Copy link
Copy Markdown
Collaborator

No description provided.

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

@ksivaman

Copy link
Copy Markdown
Member

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@ksivaman
ksivaman self-requested a review March 9, 2023 07:14
@ksivaman

Copy link
Copy Markdown
Member

Since #54 is merged now, we can remove the WIP tag too

@jeng1220

Copy link
Copy Markdown
Contributor

@mingxu1067@jeng1220 Looks like there are a lot of duplicate commits here, could you please rebase with main?

@mingxu1067 ,
Could you update your branch first? So other colleagues can read the change easier and earlier.

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman and @jeng1220, Rebased.

@mingxu1067
mingxu1067force-pushed the mingh/te_docs branch 2 times, most recently from e823240 to d449c12CompareMarch 9, 2023 08:44
Comment threadtransformer_engine/jax/__init__.py Outdated
Comment on lines 5 to 10

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 ,
Could you help to make import order to be ordered alphabetically?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The copyright statement needs to be updated

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

@ksivaman

Copy link
Copy Markdown
Member

Could we combine #88 with this PR? I think the scope of this PR is documentation and that change falls within this. I want to avoid burdening the ci with a bunch of small PRs especially when they can be combined. This would speed up the dev/review/merge process. @mingxu1067

@mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

@ksivaman, ok, merged #88 with this PR and close #88
.

@ksivaman
ksivaman requested a review from timmoon10March 10, 2023 18:32
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@mingxu1067mingxu1067 changed the title [WIP, DO NOT MERGE] Adding documents to TE/JAXAdding documents to TE/JAXMar 13, 2023
@zlsh80826

Copy link
Copy Markdown
Collaborator

/te-ci

Comment threaddocs/api/jax.rst Outdated
-------

.. autoapiclass:: transformer_engine.jax.LayerNorm(epsilon=1e-6, layernorm_type='layernorm', **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
:members: __call__
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

consistency

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to shard batch along.
if it is None, then disabling data parallelism.
The axis name in Mesh used to shard batches along.
If it is None, then disabling data parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingdataparallelism.
IfitisNone, thendataparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threadtransformer_engine/jax/sharding.py Outdated
axis name in Mesh used to split model tensor along.
if it is None, then disabling tensor parallelism.
The axis name in Mesh used to split the hidden dimensions along.
If it is None, then disabling tensor parallelism.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IfitisNone, thendisablingtensorparallelism.
IfitisNone, thentensorparallelismisdisabled.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesn't read too well. Maybe just some grammar fix

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed as the suggestion below

Comment on lines +194 to +195
The key in given RNGs via flax.linen.Module.apply that
to generate Dropout masks in the core attention.

@ksivamanksivamanMar 13, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
ThekeyingivenRNGsviaflax.linen.Module.applythat
togenerateDropoutmasksinthecoreattention.
ThekeyinthegivenRNGsviaflax.linen.Module.applythatis
usedtogenerateDropoutmasksinthecoreattention.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this what it means? @mingxu1067

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is. Change to the suggestion. THX

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

kernel_init: Initializer, default =
flax.linen.initializers.variance_scaling(1.0, 'fan_in', 'normal')
used for initializing weights of QKV and Output projection weights.
Used for initializing weights of QKV and Output projection weights.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingweightsofQKVandOutputprojectionweights.
UsedforinitializingtheQKVandOutputprojectionweights.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

If set to False, the layer will not learn additive biases.
bias_init: Initializer, default = flax.linen.initializers.zeros
used for initializing bias of QKVO projections, only works when :attr:`use_bias=True`.
Used for initializing bias of QKVO projections, it only works when :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
UsedforinitializingbiasofQKVOprojections, itonlyworkswhen :attr:`use_bias=True`.
UsedforinitializingbiasofQKVOprojections, onlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

use_bias: bool, default = False
indicate whether to enable bias shifting for QKVO projections.
if set to False, the layer will not learn additive biases.
Indicate whether to enable bias shifting for QKVO projections.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
IndicatewhethertoenablebiasshiftingforQKVOprojections.
IndicatewhetherornottoenablebiasshiftingforQKVOprojections.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/fp8.py Outdated
A helper to update Flax's Collection.

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/fp8.py Outdated
updated_scale_inv = 1/updated_scale

Collection = [dict, FrozenDict]
Collection = [dict, Flax's FrozenDict]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Collection= [dict, Flax'sFrozenDict]
Collection= [dict, flax.core.frozen_dict.FrozenDict]

Comment threadtransformer_engine/jax/module.py Outdated
----------
scale_factor : float, default = 1.0
scale the inputs along the last dimension before running softmax.
Scale the inputs along the last dimension before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this scaling only across the last dimension? The whole softmax input is scaled, right?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it is.

Comment threadtransformer_engine/jax/module.py Outdated

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
FC1andFC2. Itonlyworkswhen :attr:`use_bias=True`.
FC1andFC2. Itonlyusedwhen :attr:`use_bias=True`.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@mingxu1067 This pattern actually exists throughout, could you please go through all the cases here change "only works" -> "only used" so that it gives the correct picture?

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Comment threadtransformer_engine/jax/module.py Outdated
scale the inputs along the last dimension before running softmax.
softmax_type : SoftmaxType, default = 'layernorm'
indicate the type of softmax.
Scale the whole (inputs + bias) before running softmax.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Scalethewhole (inputs+bias) beforerunningsoftmax.
Scalarfortheinputtosoftmax.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed

Comment threaddocs/api/jax.rst Outdated
Jax
=======

Types

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
Types
Enums

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this gives a better idea of what these are supposed to be

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

Comment threaddocs/api/jax.rst
Comment on lines +26 to +42
.. autoapiclass:: transformer_engine.jax.DenseGeneral(features, layernorm_type='layernorm', use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormDenseGeneral(features, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.LayerNormMLP(intermediate_dim=2048, layernorm_type='layernorm', epsilon=1e-6, use_bias=False, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.RelativePositionBiases(num_buckets, max_distance, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.MultiHeadAttention(head_dim, num_heads, **kwargs)
:members: __call__

.. autoapiclass:: transformer_engine.jax.TransformerLayer(hidden_size=512, mlp_hidden_size=2048, num_attention_heads=8, **kwargs)
:members: __call__

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All of these modules are showing up in the Functions category below in the generated docs. @mingxu1067

@mingxu1067mingxu1067Mar 14, 2023

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remove categories as PyTorch to make the style consistent.

.. toctree::

pytorch
jax

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main README should also mention JAX.
Grep for pytorch in that document and add JAX at those places and add a JAX example too as there is a PyTorch example.

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think README.md need more refactor to show all FWs supported by TE. Diretort adding JAX to all Pytorch appearance might mess up reading. We can submit a split PR to refactor README.md.

Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@ksivaman
ksivaman merged commit ed1a311 into NVIDIA:mainMar 14, 2023
nzmora-nvidia pushed a commit to nzmora-nvidia/TransformerEngine that referenced this pull request Mar 16, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Mar 31, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* Updated TE/JAX docs
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding TE/JAX docs' rst files
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Set DType as pybind11::module_local() to avoid generic_type errors.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Updating license and exporting more modules
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopting autoapi and removing enum_tools.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix typo
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Make jax.rst be style consistent.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Fixing doc statements as the suggestion from code review.
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Update the description of Softmax
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
* Removed categories in catalog as PyTorch
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Signed-off-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants

@mingxu1067@zlsh80826@ksivaman@jeng1220@nouiz