Skip to content

TE Gemma tutorial attempt#2 - #1839

Merged
cyanguwa merged 85 commits into
NVIDIA:mainfrom
sudhakarsingh27:te_gemma_tutorial_base
Sep 17, 2025
Merged

TE Gemma tutorial attempt#2#1839
cyanguwa merged 85 commits into
NVIDIA:mainfrom
sudhakarsingh27:te_gemma_tutorial_base

Conversation

@sudhakarsingh27

@sudhakarsingh27sudhakarsingh27 commented Jun 2, 2025

Copy link
Copy Markdown
Member

Description

TLDR;
Adds a tutorial to showcase how to use Transformer Engine to accelerate generation with HuggingFace Gemma model.

Specifically, it showcases following features:

  1. How to use TE's TransformerLayer layer in place of HF's GemmaDecoderLayer in Gemma models. Monkey-patching and mapping correct module names.
  2. How to use non-paged and paged KV cache from TE
  3. How to speedup up generation with CUDA Graphs and fp8_model_init that keeps model params in FP8 precision.

Attempt#1 @ #829

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)

@sudhakarsingh27
sudhakarsingh27force-pushed the te_gemma_tutorial_base branch from 03729bc to 2a514cfCompareJune 2, 2025 21:10
@sudhakarsingh27
sudhakarsingh27force-pushed the te_gemma_tutorial_base branch from 2a514cf to 4757bfaCompareJune 2, 2025 21:19
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27
sudhakarsingh27force-pushed the te_gemma_tutorial_base branch 3 times, most recently from 5d7538e to 93960fdCompareJune 16, 2025 22:09
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
pre-commit-ciBotand others added 20 commits June 17, 2025 22:27
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
…ransformerEngine into te_gemma_tutorial_base_test
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
…mma_tutorial_base
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
def _get_weight_quantizers(self) -> List[Quantizer]:
"""Get the weight quantizers of the module."""
if not self.fp8:
if not self.fp8 and not self.fp8_calibration:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Do we need to do the same for group_linear?

# Copy the pre-step seqlens to the device in CUDA Graphs safe manner.
self.pre_step_seqlens[: len(pre_step_seqlens_temp)].copy_(
pre_step_seqlens_temp, non_blocking=False
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I missed this the first time, but why are we copying pre_step_seqlens_temp to CPU and then to GPU?

I think I wanted to say this while we were working on the tutorial, but we could extend our pre_step API to support cases where users provide a GPU tensor for step_dict too. But let's do that in a future PR.

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

self.pre_step_seqlens[: len(pre_step_seqlens_temp)].copy_(
pre_step_seqlens_temp, non_blocking=False
)

This is for CUDA graphs to work correctly since pre_step_seqlens could have different lengths. Since I'm referencing pre_step_seqlens_temp twice in this statement, I created a temporary CPU tensor. I could try checking if I could do that directly with a list but that'd mean creating a temporary list

Comment threaddocs/examples/te_gemma/te_gemma.py

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Could you read through this file again? I feel it's not as polished as te_gemm.py. There are some mistakes that can be easily spotted.

Also, could you please run the CI again after all changes are done? Just wanted to make sure the fp8_calibration, MHA changes are not breaking anything.

Nitpicks:

  • Animation 1 -> Figure 1, are we calling it animation in other tutorials?
  • Building on this foundation, the current objective is -> Building on that foundation, this tutorial is
  • ushered an era -> ushered in an era?
  • tutorial 2, etc. -> tutorial 2.
  • for both those cases -> for both use cases
  • CUDA Graph is not just used to avoid CPU launch overhead. It helps remove any other CPU overhead too - it only replays the GPU activities.
  • "Weight calibration involves calculating FP8 scaling factors from higher precision forward passes." -> Does it not require backward passes?
  • This eliminates the need to cast from higher precision to BF16, -> This eliminates the need to cast from higher precision to BF16 every time,
  • other artefacts used in the following tutorial. -> other artifacts used in this tutorial.
  • CUDA Graphs to function -> CUDA Graphs to function.
  • Blue blobs in the above figure are GPU kernels and whitespace b/w those -> Blue blobs in the top figure are GPU kernels and whitespace between those
  • The second big blob seems to be in different proportions between no-CG and CG, or is it just me.
  • It is highly recommended to familiarize oneself with the tutorial on FP8 precision to understand the necessity of scaling. -> Didn't we say this already? I feel we're repeating ourselves quite a bit in the fp8_autocast and fp8_model_init sections.
  • during the forwards in higher precision -> Is it just the forward that we need to run?
  • This approach is beneficial during training: one can perform one cast for both backward and forward passes, leading to speedups. However, performing a single cast for each forward pass introduces too much overhead to achieve a speedup. -> This paragraph doesn't quite make sense.
  • assert type(linear_fp8.weight.data) is te.float8_tensor.Float8Tensor -> Is .data a Float8Tensor, or .weight? Also, should we use Float8TensorBase?
  • Using less memory during generation (by storing weights in FP8 precision using fp8_model_init) -> Didn't we talk about this in the tutorial? Why "not extensively talked about"?
  • playing around -> playing around with?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

  • Could you read through this file again? I feel it's not as polished as te_gemm.py. There are some mistakes that can be easily spotted.

  • Also, could you please run the CI again after all changes are done? Just wanted to make sure the fp8_calibration, MHA changes are not breaking anything.

Nitpicks:

  • Animation 1 -> Figure 1, are we calling it animation in other tutorials?
    • I’m not sure if other tutorials have animations. Llama didn’t.
  • Building on this foundation, the current objective is -> Building on that foundation, this tutorial is
  • ushered an era -> ushered in an era?
  • tutorial 2, etc. -> tutorial 2.
    • Seems to be grammatically correct but improved grammar of the context around it
  • for both those cases -> for both use cases
  • CUDA Graph is not just used to avoid CPU launch overhead. It helps remove any other CPU overhead too - it only replays the GPU activities.
    • Updated the text
  • "Weight calibration involves calculating FP8 scaling factors from higher precision forward passes." -> Does it not require backward passes?
    • Backward pass is not needed to calibrate weights and inputs.
  • This eliminates the need to cast from higher precision to BF16, -> This eliminates the need to cast from higher precision to BF16 every time,
  • other artefacts used in the following tutorial. -> other artifacts used in this tutorial.
  • CUDA Graphs to function -> CUDA Graphs to function.
  • Blue blobs in the above figure are GPU kernels and whitespace b/w those -> Blue blobs in the top figure are GPU kernels and whitespace between those
  • The second big blob seems to be in different proportions between no-CG and CG, or is it just me.
    • You are correct, I just noticed that for first generated tokens, the CUDA graphs’ attn kernel is slightly longer than non CUDA graphs case. After that it becomes the same
  • It is highly recommended to familiarize oneself with the tutorial on FP8 precision to understand the necessity of scaling. -> Didn't we say this already? I feel we're repeating ourselves quite a bit in the fp8_autocast and fp8_model_init sections.
  • during the forwards in higher precision -> Is it just the forward that we need to run?
    • yes, only the fwd pass is needed
  • This approach is beneficial during training: one can perform one cast for both backward and forward passes, leading to speedups. However, performing a single cast for each forward pass introduces too much overhead to achieve a speedup. -> This paragraph doesn't quite make sense.
    • Updated paragraph
  • assert type(linear_fp8.weight.data) is te.float8_tensor.Float8Tensor -> Is .data a Float8Tensor, or .weight? Also, should we use Float8TensorBase?
    • This is just a small dirty check. We use FakeTensor for FP8 weights since PyT doesn’t have a notion of FP8 datatype, so weight is still in BF16 but weight.data is actual Float8Tensor.
    • Updated the check now to make it more intuitive
  • Using less memory during generation (by storing weights in FP8 precision using fp8_model_init) -> Didn't we talk about this in the tutorial? Why "not extensively talked about"?
    • We aren't talking about the memory usage of the models in the tutorial. I've now mentioned in there how fp8_model_init results in less memory usage.
  • playing around -> playing around with?

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Comment threaddocs/examples/te_gemma/te_gemma.py
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27

Copy link
Copy Markdown
MemberAuthor

/te-ci pytorch L0

Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
@sudhakarsingh27

Copy link
Copy Markdown
MemberAuthor

/te-ci pytorch L0

@sudhakarsingh27

Copy link
Copy Markdown
MemberAuthor

/te-ci pytorch L0

@cyanguwa
cyanguwa merged commit 7042d7a into NVIDIA:mainSep 17, 2025
23 checks passed
vthumbe1503 pushed a commit to vthumbe1503/TransformerEngine that referenced this pull request Sep 19, 2025
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Lower precision gated-act to accelerate FP8 current-scaling. (#2153)
* Applying the original precision as Norm outputs' and activation compuations.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding knob to control norm output precision.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Removing the knob and applying lower-precision norm with current-scaling only.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fix the error when quantizer==None
Signed-off-by: Ming Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
[PyTorch] Support activation CPU offloading in fusible ops (#2158)
* Add CPU offloading logic to ops. Fix test to compute dgrad.
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Make sure grads are contiguous in op backwards
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Add op-based MLP to CPU offloading tests
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Handle different weight cache behavior on Hopper/Blackwell
Add MXFP8 to CPU offload tests.
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Remove MXFP8 test
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
---------
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Do not use normalization forward + amax fusion if cuDNN backend is requested (#2174)
* Do not use norm fwd + amax fusion if cudnn backend is requested
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Read envirornment vairable directly to avoid include error
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
---------
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Fix unjoined comm stream in UB communicator (#2160)
Signed-off-by: djns99 <40156487+djns99@users.noreply.github.com>
FP8 Output Quantization for GEMM (#2123)
* Test working as I think it should work
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* revert accidental change
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Restrict the number of cases for unfused quantization, some fp8->fp8 cases are handled by cublas
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
fix merge conflict
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
bug: missed a } in the code
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Add cuBLASMp-backed GEMM-like API to TE common (#1824)
* Pick up cuBLASMp during build
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Change lib order to fix link error
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Context creation, incomplete...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Test fixure
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* A sanity AgGemm test, failing...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix axes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Take care of uneven distribution
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use MPI to get position of local matrices
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Refactor
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Refactor & fixes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Gemm-RS
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Gemm-AR, not working...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fixes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Setting all-reduce epilogue for gemm-ar
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use supported shapes for GEMM-AR
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak tolerance
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* First shot at fp8
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use TensorHolder in tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More test configs
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Support comm_sm_count
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Parametrize dtypes for A, B and D separately
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak scaling
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Amax ptr
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Flags parity with cublas_gemm, saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Cleanup
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Bias tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix bias test
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Aux, saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* aux_ld
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* A fix
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use test::Tensor
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Set scale inv
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove unsupported test configs
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Replace libcal with NCCL
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Add NVTX markers to API functions
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak GemmAr tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More test config
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix merge fallout
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove MPI dependency, comment API, add algo parameter
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix nvshmem dependency
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix nvshmem build
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Excluse CommGemm tests from L0_cppunittest
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Add cpp_distributed sh file for CI
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Adapt tp TensorAllocator
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Skip GemmAr test on unsupported HW
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Oversibscribe is needed on some clusters
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix incomplete libcal removal
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Move CI tests to L1
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Rename context to include NVTE prefix
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove leftover code
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* NVTE_WITH_CUBLASMP off by default
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More detailed NVTE_CHECK diag
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Comment API
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Include stdbool header for legacy C compilers
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove now unused argument
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Abstract away cuBLASMp algo behind our own enum
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More detailed shape diag messages
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update transformer_engine/common/include/transformer_engine/comm_gemm.h
Co-authored-by: Przemyslaw Tredak <ptrendx@gmail.com>
Signed-off-by: Vladimir Cherepanov <56651474+mk-61@users.noreply.github.com>
* Add license
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
---------
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Vladimir Cherepanov <56651474+mk-61@users.noreply.github.com>
Co-authored-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Przemyslaw Tredak <ptrendx@gmail.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
FP8 AllGather in FP8 GroupedGEMM + Fix Stream Usage Issue. (#2086)
* FP8 AllGather in FP8 GroupedGEMM
1. Support current scaling FP8 quantation with a given amax.
2. Support FP8 AG in fwd and BF16 RS in bwd.
3. The workflow is AR-max -> FP8 Quant -> FP8 AG -> FP8 GroupedGEMM.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Slightly refactor
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding documents of new args.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding unit-tests.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding license.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Move unit-tests to L1.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Move quantizaer store/reset into FP8 only.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding all layout support for Blackwell+
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopt the feedback from code-review.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fixed the wrong stream used by d2d in groupedGEMM FFI.
Signed-off-by: Ming Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Co-authored-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[JAX] Delay MeshResource validation until first usage (#2124)
Delay MeshResource validation until first usage
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Co-authored-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[JAX] Decouple Recipe and ScalingMode (#1728)
* Decouple recipe and scaling mode
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Expose global QuantizeConfig instance as a getter
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Format and lint
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Merge branch 'main' into dev/jberchtold/jax-scaling-mode-and-recipe-decoupling
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Rename UsageType to TensorSource
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Update test_layer.py
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
---------
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: jberchtold-nvidia <158520091+jberchtold-nvidia@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[JAX] `dot_1_output` sharding constraint + use AXIS_IS_UNSHARDED (#2128)
* add dot_1_output sharding constraint + use AXIS_IS_UNSHARDED
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
---------
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[JAX] Add amax input to DBiasQuantizePrimitive and FFI (#2118)
* add amax input to DBiasQuantizePrimitive and FFI
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* make sure amax is init with zero
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
* fix sharding rule
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
---------
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Further relax constraints to cuDNN 9.13 for disabling fused attn for kv caching (#2121)
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Temporarily remove comm_gemm tests (#2133)
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[PyTorch] Disable determinism for sm100 (#2130)
* disable determinism for sm100+ and cudnn<9.14
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* fix remaining CI failures
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* revert some changes
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* revert more changes
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* remove sm100 from determinism table
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
---------
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[PyTorch] ONNX export of FP8 Current Scaling (#2068)
* Compute amax in normalization forward in current scaling in untuned kernels
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* code drop
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* apply tims suggestions
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
---------
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Co-authored-by: Jan Bielak <jbielak@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[PyTorch][MOE] Tentative Fix For Replacing from_blob with empty for experts receiving zero tokens (#2134)
use torch empty for empty shape instead of from_blob
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
build: pull cached wheels (#2127)
* build: pull cached wheels
Signed-off-by: oliver könig <okoenig@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update setup.py
Signed-off-by: oliver könig <okoenig@nvidia.com>
---------
Signed-off-by: oliver könig <okoenig@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
feat: Add support for multiple quantization modes in the UB communicators (#2043)
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[Common] Add checks to CUDA kernel launch and CUDA API calls (#2074)
* add checks to cuda kernel launch and cuda API calls
Signed-off-by: Xin Yao <xiny@nvidia.com>
* Remove exceptions from destructors
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* fix weired dispatch in ln/rmsnorm
Signed-off-by: Xin Yao <xiny@nvidia.com>
---------
Signed-off-by: Xin Yao <xiny@nvidia.com>
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[PyTorch] Support bf16+fp8 cudagraph (#2098)
* support bf16+fp8 model
Signed-off-by: Robin Zhang <robinz@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* update
Signed-off-by: Robin Zhang <robinz@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* update
Signed-off-by: Robin Zhang <robinz@nvidia.com>
---------
Signed-off-by: Robin Zhang <robinz@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Dropout with 8-bit RNG (#2014)
* Add dropout kernel with 8-bit RNG
Co-authored-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix license
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Avoid ambiguous types
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Do not enforce dropout prob is representable in 8 bits
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Expand error message
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix small statistical bug from using less-equal instead of less-than
Refactor kernel implementations and add comments. Interpret masks as bytes rather than 16-bit uints.
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Fix linter warning
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Remove unnecessary helper function in PyTorch extensions
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Create GPU reload buffers on main stream (#2131)
* Create GPU relaod buffers on main stream
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fixed typo
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
* Fixed typo
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
---------
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
Co-authored-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
Co-authored-by: Paweł Gadziński <62263673+pggPL@users.noreply.github.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
mxfp8 unfused quant support, refined unit test, remove unecessary quantization code
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
missed a quant code removal
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
minor bug fix
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Add cuBLASMp-backed GEMM-like API to TE common (#1824)
* Pick up cuBLASMp during build
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Change lib order to fix link error
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Context creation, incomplete...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Test fixure
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* A sanity AgGemm test, failing...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix axes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Take care of uneven distribution
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use MPI to get position of local matrices
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Refactor
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Refactor & fixes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Gemm-RS
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Gemm-AR, not working...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fixes
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Setting all-reduce epilogue for gemm-ar
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use supported shapes for GEMM-AR
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak tolerance
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* First shot at fp8
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use TensorHolder in tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More test configs
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Support comm_sm_count
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Parametrize dtypes for A, B and D separately
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak scaling
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Amax ptr
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Flags parity with cublas_gemm, saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Cleanup
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Bias tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix bias test
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Aux, saving...
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* aux_ld
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* A fix
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Use test::Tensor
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Set scale inv
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove unsupported test configs
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Replace libcal with NCCL
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Add NVTX markers to API functions
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Tweak GemmAr tests
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More test config
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix merge fallout
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove MPI dependency, comment API, add algo parameter
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix nvshmem dependency
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix nvshmem build
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Excluse CommGemm tests from L0_cppunittest
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Add cpp_distributed sh file for CI
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Adapt tp TensorAllocator
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Skip GemmAr test on unsupported HW
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Oversibscribe is needed on some clusters
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Fix incomplete libcal removal
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Move CI tests to L1
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Rename context to include NVTE prefix
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove leftover code
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* NVTE_WITH_CUBLASMP off by default
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More detailed NVTE_CHECK diag
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Comment API
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Include stdbool header for legacy C compilers
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Remove now unused argument
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* Abstract away cuBLASMp algo behind our own enum
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* More detailed shape diag messages
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update transformer_engine/common/include/transformer_engine/comm_gemm.h
Co-authored-by: Przemyslaw Tredak <ptrendx@gmail.com>
Signed-off-by: Vladimir Cherepanov <56651474+mk-61@users.noreply.github.com>
* Add license
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
---------
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Vladimir Cherepanov <56651474+mk-61@users.noreply.github.com>
Co-authored-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Przemyslaw Tredak <ptrendx@gmail.com>
FP8 AllGather in FP8 GroupedGEMM + Fix Stream Usage Issue. (#2086)
* FP8 AllGather in FP8 GroupedGEMM
1. Support current scaling FP8 quantation with a given amax.
2. Support FP8 AG in fwd and BF16 RS in bwd.
3. The workflow is AR-max -> FP8 Quant -> FP8 AG -> FP8 GroupedGEMM.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Slightly refactor
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding documents of new args.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding unit-tests.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding license.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Move unit-tests to L1.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Move quantizaer store/reset into FP8 only.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adding all layout support for Blackwell+
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Adopt the feedback from code-review.
Signed-off-by: Ming Huang <mingh@nvidia.com>
* Fixed the wrong stream used by d2d in groupedGEMM FFI.
Signed-off-by: Ming Huang <mingh@nvidia.com>
---------
Signed-off-by: Ming Huang <mingh@nvidia.com>
Co-authored-by: Phuong Nguyen <phuonguyen@nvidia.com>
[JAX] Delay MeshResource validation until first usage (#2124)
Delay MeshResource validation until first usage
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Co-authored-by: Phuong Nguyen <phuonguyen@nvidia.com>
[JAX] Decouple Recipe and ScalingMode (#1728)
* Decouple recipe and scaling mode
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Expose global QuantizeConfig instance as a getter
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Format and lint
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Merge branch 'main' into dev/jberchtold/jax-scaling-mode-and-recipe-decoupling
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Rename UsageType to TensorSource
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
* Update test_layer.py
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
---------
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: jberchtold-nvidia <158520091+jberchtold-nvidia@users.noreply.github.com>
[JAX] `dot_1_output` sharding constraint + use AXIS_IS_UNSHARDED (#2128)
* add dot_1_output sharding constraint + use AXIS_IS_UNSHARDED
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
---------
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
[JAX] Add amax input to DBiasQuantizePrimitive and FFI (#2118)
* add amax input to DBiasQuantizePrimitive and FFI
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* make sure amax is init with zero
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
* fix sharding rule
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
---------
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Further relax constraints to cuDNN 9.13 for disabling fused attn for kv caching (#2121)
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
Temporarily remove comm_gemm tests (#2133)
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
[PyTorch] Disable determinism for sm100 (#2130)
* disable determinism for sm100+ and cudnn<9.14
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* fix remaining CI failures
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* revert some changes
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* revert more changes
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* remove sm100 from determinism table
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
---------
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
[PyTorch] ONNX export of FP8 Current Scaling (#2068)
* Compute amax in normalization forward in current scaling in untuned kernels
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* code drop
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
* apply tims suggestions
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
---------
Signed-off-by: Jan Bielak <jbielak@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Co-authored-by: Jan Bielak <jbielak@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
[PyTorch][MOE] Tentative Fix For Replacing from_blob with empty for experts receiving zero tokens (#2134)
use torch empty for empty shape instead of from_blob
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
build: pull cached wheels (#2127)
* build: pull cached wheels
Signed-off-by: oliver könig <okoenig@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update setup.py
Signed-off-by: oliver könig <okoenig@nvidia.com>
---------
Signed-off-by: oliver könig <okoenig@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
feat: Add support for multiple quantization modes in the UB communicators (#2043)
[Common] Add checks to CUDA kernel launch and CUDA API calls (#2074)
* add checks to cuda kernel launch and cuda API calls
Signed-off-by: Xin Yao <xiny@nvidia.com>
* Remove exceptions from destructors
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* fix weired dispatch in ln/rmsnorm
Signed-off-by: Xin Yao <xiny@nvidia.com>
---------
Signed-off-by: Xin Yao <xiny@nvidia.com>
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
[PyTorch] Support bf16+fp8 cudagraph (#2098)
* support bf16+fp8 model
Signed-off-by: Robin Zhang <robinz@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* update
Signed-off-by: Robin Zhang <robinz@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* update
Signed-off-by: Robin Zhang <robinz@nvidia.com>
---------
Signed-off-by: Robin Zhang <robinz@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Dropout with 8-bit RNG (#2014)
* Add dropout kernel with 8-bit RNG
Co-authored-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix license
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Avoid ambiguous types
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Do not enforce dropout prob is representable in 8 bits
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Expand error message
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix small statistical bug from using less-equal instead of less-than
Refactor kernel implementations and add comments. Interpret masks as bytes rather than 16-bit uints.
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* Fix linter warning
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Remove unnecessary helper function in PyTorch extensions
Signed-off-by: Tim Moon <tmoon@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Create GPU reload buffers on main stream (#2131)
* Create GPU relaod buffers on main stream
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fixed typo
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
* Fixed typo
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
---------
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
Signed-off-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
Co-authored-by: Selvaraj Anandaraj <selvaraja@login-ptyche01.ptyche.clusters.nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Selvaraj Anandaraj <selvaraja@login-preos01.a51.clusters.nvidia.com>
Co-authored-by: Paweł Gadziński <62263673+pggPL@users.noreply.github.com>
minor code cleanup
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
minor cosmetics
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Address review comment
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
minor comment update
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Fix CI failures for UB overlap changes (#2149)
Signed-off-by: djns99 <40156487+djns99@users.noreply.github.com>
minor bug: quantizer should not be none for unfused quantization
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[JAX] Fix failing fused attn tests for dropout=0.1 and bias for sm100 (#2135)
* Fix failing tests for dropout=0.1 and bias for fused attn for blackwell
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix the skip message
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* Assert in fused attn bwd pass for sm100
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
Add check for sm100
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Add support to get all devs in the process for jax
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Code clean up
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* Make get_all_device_compute_capability more pythonic, thereby avoiding unnecessary type conversion
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* Represent attn bias using enum instead of string
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
---------
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
fix linting error
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[PyTorch][CUDA Graph] Fix FP8 Weight Quantization Cache under CUDA Graph (#2119)
* add noop to comp amax
Signed-off-by: zhongboz <zhongboz@nvidia.com>
* fix for fp8 blockwise recipe
Signed-off-by: zhongboz <zhongboz@nvidia.com>
* resolve comments
Signed-off-by: zhongboz <zhongboz@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
address review comments
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* Update test_multi_process_distributed_grouped_gemm.py
change accidentally added while merging
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
* Update dense.py
change accidentally added while merging
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* address review comments
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* address revie comments
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Bug solved: delayed scaling quantization with mxfp8 inputs didnt work
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix the unit test error
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* just to trigger ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* address review comments: quantization inside gemm and outside both should exactly match for fp32 accumulation
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
* fix merge conflict
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
address review comments: quantization inside gemm and outside both should exactly match for fp32 accumulation
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
TE Gemma tutorial attempt#2 (#1839)
* add tutorial files and other local changes
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* remove extraneous code for easy debu
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* make cuda graphs work with non-paged and paged attention
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* perf imp for kv cache ops
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* add code for calibration
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* optimize kv_cache reindex and copy kernels
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* changes to make quantizers work with fp8_calibration
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* avoid reindexing from python side
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* rename variable from previous commit
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* minor fix
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* minor fix
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
* use quantizer only if needed
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* functionality of the tutorial tested and perf checked
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* remove files and update headers/licenses
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* update header/license
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* update tutorial for review
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* make weights downloadable on the fly; remove extra print statements
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix lint and update comments
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* add comma back, typo
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* sequence_start_positions should be None for training
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* add paged attention numberes and update requirements.txt file
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* more fixes
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* make tutorial work on blackwell
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* remove gemma FT tutorial for now
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* fixing the headings placement and rewording attention -> kv caching
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* fixes from comments
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* fix the images
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* misc fixes
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* add more comments to te_gemma.py and cleanup utils.py
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* add more information about the hierarchy of the classes used in the tutorial
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* add better cuda graphs picture
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* addd updated cuda graphs pictures
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* add illustrated cuda graphs
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* fix
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* small fixes in documentation
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* add torch.no_grad() to force reduced memory usage
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* some fixes from recent comments
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* more fixes from remaining comments
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* add te_rope_emb to class desc
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
* fix tutorial wording; add calibration fix to grouped_linear.py
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
---------
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Fix memory overhead of linear layer when all gather from sequence parallel (#2125)
* fix memory overhead of all gather from sequence parallel
Signed-off-by: Yuzhong Wang <yuzhongw@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update transformer_engine/pytorch/tensor/_internal/float8_blockwise_tensor_base.py
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
* quick fix the errors when for UB buffers
Signed-off-by: Yuzhong Wang <yuzhongw@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Update transformer_engine/pytorch/module/linear.py
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
* Avoid deallocating FP8 scale-invs since they are reused
Signed-off-by: Tim Moon <tmoon@nvidia.com>
---------
Signed-off-by: Yuzhong Wang <yuzhongw@nvidia.com>
Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Signed-off-by: Tim Moon <tmoon@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com>
Co-authored-by: Tim Moon <tmoon@nvidia.com>
Fix incorrect TP rank calculation when using data parallel (#2179)
Signed-off-by: djns99 <40156487+djns99@users.noreply.github.com>
[Pytorch] Add Cutlass Grouped GEMM Support for fine-grained MoE Model (#2045)
* feat: add cutlass group gemm support
Signed-off-by: Min Yang <min.yang@shopee.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* refactor: refactor multi tensor gemm interface
Signed-off-by: Min Yang <min.yang@shopee.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* refactor: refactor nvte_multi_stream_cublas_gemm func and add license info
Signed-off-by: Min Yang <min.yang@shopee.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* feat: add unit test for cutlass group gemm
Signed-off-by: Min Yang <min.yang@shopee.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* feat: add cutlass support type protect
Signed-off-by: Min Yang <min.yang@shopee.com>
* add tests and fix lint
Signed-off-by: Xin Yao <xiny@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* feat: fix unit tests error
Signed-off-by: Min Yang <min.yang@shopee.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* feat: refactor host workspace malloc
Signed-off-by: Min Yang <min.yang@shopee.com>
* update cutlass
Signed-off-by: Xin Yao <xiny@nvidia.com>
* update cutlass
Signed-off-by: Xin Yao <xiny@nvidia.com>
* further relex threshold and add a env var to warn fall back
Signed-off-by: Xin Yao <xiny@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: Min Yang <min.yang@shopee.com>
Signed-off-by: Xin Yao <xiny@nvidia.com>
Signed-off-by: alan yang <89962857+cassiewilliam@users.noreply.github.com>
Co-authored-by: Min Yang <min.yang@shopee.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Xin Yao <xiny@nvidia.com>
Co-authored-by: Phuong Nguyen <phuonguyen@nvidia.com>
[PyTorch] Support FA3 for MLA and with CP (#1907)
feature(FA3,MLA,CP):
1. Update FA3 to commit-id 3ba6f82 (tag 2.8.0.post2 with compile error fixed), PR-1604 support hdimQK != hdimV backward
2. Update get_attention_backend method because FA3 support MLA now
3. Add CP MLA support for FA3
4. Add unit tests for FA3 MLA CP
5. Update attention doc
Signed-off-by: zhujian <zhujian.whu.cs@gmail.com>
Fix cuDNN version checks when getting backend and for sm89 kv cache (#2185)
* Fix cudnn version checks for kv cache for sm89. Add cudnn version check in preparation for 9.14 when getting backend
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Minor fix for cuDNN version condition check
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
@sudhakarsingh27sudhakarsingh27 mentioned this pull request Feb 9, 2026
13 tasks
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants

@sudhakarsingh27@pggPL@cyanguwa