Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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" + '
Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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('^' + ".*" + ' Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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('^' + ".*" + ' Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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" + ' Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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('^' + ".*" + ' Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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('^' + ".*" + ' Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067
, '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); } })(); })(); Add TE/JAX high-level modules, unittests and examples by jeng1220 · Pull Request #54 · NVIDIA/TransformerEngine · GitHub
Skip to content

Add TE/JAX high-level modules, unittests and examples - #54

Merged
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax
Mar 9, 2023
Merged

Add TE/JAX high-level modules, unittests and examples#54
timmoon10 merged 21 commits into
NVIDIA:mainfrom
jeng1220:rjeng/add_te_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 17, 2023

Copy link
Copy Markdown
Contributor

Add TE/JAX high-level modules and API

  • transformer_engine/jax/init.py
  • transformer_engine/jax/module.py
  • transformer_engine/jax/softmax.py
  • transformer_engine/jax/transformer.py

Add TE/JAX unittests to cover high-level API

  • qa/L0_jax_unittest/test.sh
  • tests/jax/test_layer.py
  • tests/jax/test_mnist.py
  • tests/jax/test_sharding.py
  • tests/jax/utils.py

Add TE/JAX exampels

  • examples/jax/encoder/test_single_gpu_bf16_training.py
  • examples/jax/encoder/test_single_gpu_fp8_training.py

TE/JAX bugfix

  • tests/jax/test_helper.py
  • transformer_engine/jax/fp8.py
  • transformer_engine/jax/sharding.py

@jeng1220
jeng1220force-pushed the rjeng/add_te_jax branch 6 times, most recently from 969d306 to 1277484CompareJanuary 19, 2023 16:23
@jeng1220jeng1220 changed the title [WIP] add TE/Jax python modulesAdd TE/JAX high-level modules, unittests and examplesFeb 25, 2023
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I have updated this PR(#54). Please help to review it.
Many thanks

@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The internal pipeline is #7424183

@ksivamanksivaman mentioned this pull request Feb 27, 2023
@timmoon10
timmoon10 self-requested a review February 28, 2023 21:53
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Overall looks good. I see that Pylint is complaining about "Too few public methods", but I think it just doesn't understand that Flax modules are dataclasses. We can just suppress those warnings.

Comment threadtests/jax/test_sharding.py Outdated
Comment threadtransformer_engine/jax/softmax.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/module.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
jeng1220and others added 4 commits March 6, 2023 07:51
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@nouiz

nouiz commented Mar 6, 2023

Copy link
Copy Markdown
Collaborator

The README should be updated to talk about JAX (or point to a JAX subpage).
This can be in this PR or a follow up.

Comment threadtransformer_engine/jax/transformer.py Outdated
Comment threadtransformer_engine/jax/transformer.py Outdated
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 7, 2023

Copy link
Copy Markdown
ContributorAuthor

The README should be updated to talk about JAX (or point to a JAX subpage). This can be in this PR or a follow up.

@nouiz ,
I will update doc and REAME.md in next PR

…_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

@timmoon10timmoon10 left a comment

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.

Looks good. I'm happy with merging once we fix the failing tests.

Running with #86, I don't see any more Pylint errors if we add # pylint: disable=too-many-function-args:

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

k_kernel=self.kernel_init(k_key, k_shape, dtype)
v_kernel=self.kernel_init(v_key, v_shape, dtype)

Comment threadtransformer_engine/jax/__init__.py Outdated
jeng1220and others added 5 commits March 7, 2023 17:20
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Comment threadexamples/jax/encoder/test_single_gpu_bf16_training.py Outdated
Comment threadexamples/jax/encoder/test_single_gpu_fp8_training.py
mingxu1067and others added 3 commits March 8, 2023 08:27
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

jeng1220 commented Mar 8, 2023

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
I think all existing conversations are solved, and unittest failures are fixed in 383d170.
Please help to run CI again.
Thanks

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadtransformer_engine/jax/fp8.py Outdated
Comment threadtransformer_engine/jax/fp8.py
Comment threadtransformer_engine/jax/fp8.py Outdated
@timmoon10

Copy link
Copy Markdown
Member

Pylint passes for me when I run manually on this branch. The only thing remaining thing before merging is to iron out a possible correctness issue.

Comment threadtransformer_engine/jax/fp8.py Outdated
jeng1220and others added 3 commits March 9, 2023 08:28
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 self-requested a review March 9, 2023 01:23
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@ksivaman

Copy link
Copy Markdown
Member

/te-ci

@timmoon10
timmoon10 merged commit bc9d57a into NVIDIA:mainMar 9, 2023
@ksivamanksivaman mentioned this pull request Mar 9, 2023

@nouiznouiz left a comment

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.

In fact, can you add the make_jaxpr assert in all the tests?

features = [256, 256, 128]
for feature in features:
x = DenseGeneral(features=feature, transpose_batch_sequence=False,
dtype=jnp.bfloat16, use_bias=True)(x) \

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.

should the dtype be fp8? What trigger the fp8 computation here?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

See #108

idx = idx + 1
batch = {k: v[perm, ...] for k, v in train_ds.items()}
state, metrics, grads = train_step(state, variables, batch, use_fp8_dense)

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.

Can you add an assert based on make_jaxpr and fp8 here? I'm not sure it is being used by the test.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

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

@nouiz ,
See this: #108

cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add transformer module , unittests and examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update tests/jax/test_sharding.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/transformer.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint: disable=line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove pylint: disable=too-many-func-args
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Fix the wrong broadcasting dim to dropout masks when enable transpose_bs.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Enable 2xACC for WGRAD and DGRAD by default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename LayerNormMlpBlock as LayerNormMLP
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor to avoid line-too-long
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename amax_history_size to amax_history_len
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* align dropout mask to TE/PyTorch as default
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* enlarge atol for decoder unittests
Two decoder unittests can pass in old JAX container(e.g., 23.02)
but can't in latest container (devel).
1. The actual(-0.020264) and desired(-0.020386) are very close.
2. The TE kernels are not changed, the diff should come from
new codegen behavior of XLA.
Thus, it is a common floating-point accumulated error.
Enlarge atol to avoid unittest failures.
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Adding Amax History Support
1. hide amax update in custom_vjp
2. replace amax indexing with roll(using circular buffer)
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* move kernel_init to __post_init__
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor encoder examples
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update transformer_engine/jax/fp8.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove envvar regarding 2xACC
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* remove unused import
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
---------
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Ming-Xu Huang <mingh@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/add_te_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Header change] Remove the `cudnn_backend.h` dependency, since the correct header is already included in cudnn.h
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.

6 participants

@jeng1220@timmoon10@nouiz@ksivaman@ptrendx@mingxu1067