Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman
, '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 building workflow for TE/Jax by jeng1220 · Pull Request #53 · NVIDIA/TransformerEngine · GitHub
Skip to content

add building workflow for TE/Jax - #53

Merged
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax
Feb 24, 2023
Merged

add building workflow for TE/Jax#53
timmoon10 merged 38 commits into
NVIDIA:mainfrom
jeng1220:rjeng/init_jax

Conversation

@jeng1220

@jeng1220jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
Contributor

This is the first commit for TE/Jax.

The build process for TE/PyTorch and TE/Jax components is different. In the building step, the TE/Jax only depends on CPP STL and CUDA, so we simply add needed building commands in CMakeList.txt. Unlike PyTorch, Jax has nothing like torch.utils.cpp_extension.BuildExtension and torch.utils.cpp_extension.CUDAExtension.

To ensure that users can build only TE/PyTorch, only TE/Jax, or both, I created FrameworkBuilder to manage framework-specific stuff. For example, if Jax users want to build TE/Jax, then PyTorch should not be required in their environment.

The user can set up an environment variable - FRAMEWORK to select framework, such as:

FRAMEWORK=jax pip install .
FRAMEWORK=pytorch pip install .
FRAMEWORK=all pip install .

The default value is all

The other Jax modules will be submitted after this PR is merged.

@jeng1220

jeng1220 commented Jan 11, 2023

Copy link
Copy Markdown
ContributorAuthor

Addition, TE/Jax need pybind11 to build Python-CPP binding, so we need

apt install pybind11-dev
pip install pybind11

Both of above packages are needed. Otherwise, CMake will return:

CMake Error at CMakeLists.txt:32 (find_package):
Could not find a package configuration file provided by "pybind11" with any
of the following names:
pybind11Config.cmake
pybind11-config.cmake

Comment threadsetup.py Outdated
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@timmoon10 ,
The CI triggers build failed:
https://github.com/NVIDIA/TransformerEngine/actions/runs/3901104182/jobs/6662679011#step:4:85

As I mentioned, TE/Jax needs pybind11
We have to install pybind11-dev in the CI container.

$ apt install pybind11-dev

@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.

Sorry for the delay. I think this is mostly fine. I could imagine some future build errors if, for example, the CMake flags for the PyTorch and JAX extensions are incompatible. That seems unlikely though, and we can refactor setup.py as any issues come up.

The biggest issues are merge conflicts. In particular, #51 removes the scale_inv arg from many of the TE functions and makes it the responsibility of the frameworks. This will likely also affect #54.

Comment threadtransformer_engine/CMakeLists.txt Outdated
Comment threadtransformer_engine/jax/csrc/utils.h Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
Comment threadtransformer_engine/jax/csrc/modules.cc Outdated
@jeng1220
jeng1220force-pushed the rjeng/init_jax branch 3 times, most recently from 8f8a4ac to 8e70ae0CompareJanuary 19, 2023 11:29
Comment threadsetup.py Outdated
@timmoon10

timmoon10 commented Jan 19, 2023

Copy link
Copy Markdown
Member

Once we iron out the style and build issues, I think this is okay to merge.

I'm not thrilled with how we're refactoring layer norm and RMSNorm to use the same helper function. I think it's a premature optimization. In principle they are independent operations that just happen to look similar. But in practice I don't expect them to undergo significant API changes in the future, and this is an internal implementation detail anyways.

@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.

I see the linter found some style issues.

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Comment threadsetup.py Outdated
Comment threadtransformer_engine/__init__.py Outdated
Comment threadtests/jax/test_custom_call_compute.py
Comment threadtests/jax/test_custom_call_compute.py Outdated
Comment threadtests/jax/test_custom_call_shape.py Outdated
Comment threadtests/jax/test_helper.py Outdated
Comment threadtests/jax/utils.py Outdated
Comment on lines +73 to +102
class MajorShardingType(Enum):
"""
The major sharding type to indicate sharding pattern.
`SINGLE` means single process training.
`DP` means data parallel traiing.
`TP` means tensor parallel traiing.
`DPTP` means data and tensor parallel traiing.
"""
SINGLE = 0
DP = 1
TP = 2
DPTP = 3


class ShardingType(Enum):
"""
The sharding type to indicate sharding pattern.
`SINGLE` means no sharding.
`DP` means sharding along data parallelism.
`TP_COL` means sharding along column-split tensor parallelism.
`TP_ROW` means sharding along row-split tensor parallelism.
`DP_TP_COL` means sharding along data and column-split tensor parallelism.
`DP_TP_ROW` means sharding along data and row-split tensor parallelism.
"""
SINGLE = (MajorShardingType.SINGLE, "single")
DP = (MajorShardingType.DP, "dp")
TP_COL = (MajorShardingType.TP, "tp_col")
TP_ROW = (MajorShardingType.TP, "tp_row")
DP_TP_COL = (MajorShardingType.DPTP, "dp_tp_col")
DP_TP_ROW = (MajorShardingType.DPTP, "dp_tp_row")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I have a feeling enums are the wrong data structure here. The number of enums grows exponentially with the number of sharding dimensions and I frequently see messy patterns like:

ifsharding_typein (ShardingType.TP_COL, ShardingType.DP_TP_COL):
tp_col_impl()
else:
default_impl()

It would feel better to keep track of separate bools for DP and TP. In fact, I wonder if it's better to split up TP and treat TP_COL and TP_ROW as completely orthogonal sharding axes. Treating the sharding axes as orthogonal would help with code encapsulation and make it easier to do things like add new sharding axes (e.g. over the sequence dim) without needing to touch everything.

Maybe not something to do right now if deadlines are urgent, but something to think about for a future refactor.

@jeng1220jeng1220Feb 23, 2023

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.

This is an action item in our schedule. But to refactor this part, probably will take more than 2 weeks. That will be too late to release TE/JAX.

Those sharding stuff is low-level implementation in TE/JAX. The TE/JAX users don't and shouldn't call them directly. In the next PR, we will provide high-level module to hide these stuff.

I fully agree this part should be better. We will continue refactoring after first TE/JAX release is done.

@timmoon10
timmoon10 self-requested a review February 22, 2023 20:05
@timmoon10

timmoon10 commented Feb 22, 2023

Copy link
Copy Markdown
Member

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

jeng1220and others added 2 commits February 23, 2023 19:52
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>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

The test failures are unrelated to this PR (see #67 (comment)), so I think this is ready after addressing @ksivaman's concerns. I don't have full confidence that I got everything and in the future I would try to keep each PR with a clearly defined scope. It's much faster and safer to review multiple simple PRs than a single massive one.

@timmoon10 ,
I understand it is difficult to review so much files at once.
The next PR will include less files. Hope it can be easier for code review.
I appreciate your help very much.

@timmoon10

Copy link
Copy Markdown
Member

/te-ci

Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
@jeng1220

Copy link
Copy Markdown
ContributorAuthor

@ksivaman and @timmoon10 ,
I have merged the main

@timmoon10

Copy link
Copy Markdown
Member

@ksivaman and I think this is ready to merge, but we'd like to discuss with @ptrendx tomorrow. Barring any issues, we'll merge it by California afternoon.

@timmoon10
timmoon10 merged commit a3ec6a5 into NVIDIA:mainFeb 24, 2023
@ksivamanksivaman mentioned this pull request Feb 27, 2023
cyanguwa pushed a commit to cyanguwa/TransformerEngine that referenced this pull request Apr 1, 2023
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* add building workflow for jax modules
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* replace bit_cast with reinterpret_cast
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add nvtx to cmake check list
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor rmsnorm fwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor layernorm_bwd
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* set pytorch as default in setup.py
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* rename extension from *.cc to *.cpp
cpplint cannot recognize *.cc file, so rename the extension
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* refactor style, to align TE/PyTorch
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add pybinding, unittest and qa
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix license
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* disable c-extension-no-member and no-name-in-module
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* add dataclass avoid pylint error
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update transformer_engine/__init__.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
fix typo
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* Update tests/jax/test_custom_call_shape.py
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* fix conflict due to PR62
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* fix c-extension-no-member and no-name-in-module
1. add transformer_engine_jax into extension-pkg-whitelist
2. convert pylintrc from CRLF to LF format
Signed-off-by: Ryan Jeng <rjeng@nvidia.com>
* Update setup.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Jeng Bai-Cheng <jeng1220@users.noreply.github.com>
* remove pylint:disable and refactor import order
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: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
@jeng1220
jeng1220 deleted the rjeng/init_jax branch April 24, 2023 01:45
zhiyu-deep pushed a commit to zhiyu-deep/TransformerEngine that referenced this pull request Sep 3, 2024
[Bug Fix] Added a padding mask like sub-graph in sdpa node when
kv-sequence length is not a multiple of 64 and padding mask is not
enabled. This allows graphs with kv- sequence length not a multiple of
64 to be executed on cudnn version 8.9.5 onwards. cudnn versions prior
to this now correctly return NOT_SUPPORTED as expected.
[Bug Fix] Fixed an issue where creation of graph object leads to
compilation error in some compilers.
[Bug Fix] cudnn frontend now correctly sets the stream to on the handle.
This affected only the python bindings.
[Internal change] Streamlined includes of cudnn graph API header files
into `cudnn_frontend.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.

4 participants

@jeng1220@timmoon10@ptrendx@ksivaman