Repository navigation
[mxfp8 training] require torch nightly and cuda 12.8+ - #4141
Merged
Merged
Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4141
Note: Links to docs will display an error until the docs builds have been completed. ✅ You can merge normally! (1 Unrelated Failure)As of commit a7fe353 with merge base 34322b5 ( BROKEN TRUNK - The following job failed but were present on the merge base:👉 Rebase onto the `viable/strict` branch to avoid these failures
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
jcaip
approved these changes
Mar 23, 2026
vkuzo
reviewed
Mar 23, 2026
|
|
||
| if isinstance(config, MXFP8TrainingOpConfig): | ||
| # MXFP8 training requires torch nightly build for cuda 12.8+ | ||
| assert "dev" in torch.__version__ and "cu12" in torch.version.cuda, ( |
Contributor
There was a problem hiding this comment.
would this work on any nightly (even before the PT core fix)? Also, would this break on stable when the corresponding PyTorch stable version is released? Ideally this should:
- fail when on a version of PyTorch without the core fix
- pass on nightly, stable or built-from-source versions of PyTorch with the core fix
Contributor
Author
There was a problem hiding this comment.
i updated it to check if pytorch installation has certain attributes that indicate the necessary PRs were included
danielvegamyhre
force-pushed
the
assert321
branch
from
March 23, 2026 21:20
f61e775 to
a7fe353
Compare
S1ro1
pushed a commit
to PrimeIntellect-ai/prime-rl
that referenced
this pull request
Jul 7, 2026
torchao's _get_tensor_cls_for_config (added in pytorch/ao#4141, "require torch nightly and cuda 12.8+", Mar 2026) asserts two post-2.11.0 DTensor symbols (_ops.scaled_mm_single_dim_strategy, _dispatch.is_pinned_handler) exist, else raises "install the latest torch nightly". This blanket-blocks all MXFP8 training, but is over-broad for us: prime-rl's EP path pulls expert weights local (.to_local()) before the grouped GEMM, so the DTensor scaled-mm strategy is never hit, and current torch's non-single-dim scaled_mm_strategy covers the rest. Neutralize just that assert (MXFP8 branch only) before quantize_, so enable_grouped_gemm works on our pinned torch 2.11 without a nightly bump. Verified 8xB200 SFT (GLM-0.5B, ep=8, grouped_gemm + a2a): 8/8 steps, all ranks; numerics vs bf16 within mxfp8 range (out ~4.8%, grads ~6-7%) for both mxfp8_rceil and mxfp8_rceil_wgrad_with_hp recipes. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
S1ro1
pushed a commit
to PrimeIntellect-ai/prime-rl
that referenced
this pull request
Jul 16, 2026
* Add MXFP8 Linear Layer * Add MXFP8 Receipe * Add MXFP8 * Refractor, remove constants, cleanup * Remove import from local path * Reformat * Add MXFP8 expert parallel impl to support TorchAO's MXTensor in backward for all2all in MXFP8 * Build torchao from source for py3.12 MXFP8 CUDA kernel Pin torchao to the pytorch/ao git source (rev 02105d4) and build its CUDA kernels from source against the local torch cu128 / cpython-3.12. The PyPI wheel ships the MXFP8 kernel as a cpython-310-only .so that fails to load on prime-rl's Python 3.12, so torchao::mxfp8_quantize was never registered with a CUDA backend and MXFP8 died in backward. Building from source produces a cpython-312 extension (_C_mxfp8.cpython-312-*.so). - no-build-isolation-package += torchao (build against local torch) - extra-build-variables: CUDA_HOME=/usr/local/cuda-12.8, TORCH_CUDA_ARCH_LIST=10.0a (B200/sm_100a), USE_CPP=1 Enables the MXFP8 all-to-all and dense-linear paths (verified on 8xB200). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> * Relax torchao MXFP8 grouped-GEMM torch-nightly gate torchao's _get_tensor_cls_for_config (added in pytorch/ao#4141, "require torch nightly and cuda 12.8+", Mar 2026) asserts two post-2.11.0 DTensor symbols (_ops.scaled_mm_single_dim_strategy, _dispatch.is_pinned_handler) exist, else raises "install the latest torch nightly". This blanket-blocks all MXFP8 training, but is over-broad for us: prime-rl's EP path pulls expert weights local (.to_local()) before the grouped GEMM, so the DTensor scaled-mm strategy is never hit, and current torch's non-single-dim scaled_mm_strategy covers the rest. Neutralize just that assert (MXFP8 branch only) before quantize_, so enable_grouped_gemm works on our pinned torch 2.11 without a nightly bump. Verified 8xB200 SFT (GLM-0.5B, ep=8, grouped_gemm + a2a): 8/8 steps, all ranks; numerics vs bf16 within mxfp8 range (out ~4.8%, grads ~6-7%) for both mxfp8_rceil and mxfp8_rceil_wgrad_with_hp recipes. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> * Use prebuilt torchao wheel from v0.6.0 release (x86_64) Replace the torchao build-from-source git dependency with the prebuilt cp312 CUDA wheel hosted on the prime-rl v0.6.0 release, matching how deep-ep / deep-gemm / nixl-cu12 are shipped. Avoids requiring the CUDA toolchain at install time on x86_64. - torchao source -> release wheel URL, scoped to platform_machine == 'x86_64' (aarch64 continues to resolve torchao from the default index, untouched) - drop torchao from no-build-isolation-package + the extra-build-variables block - add torchao to the exclude-newer-package trusted-sources list Wheel is H100 (sm90a) + B200 (sm100a), built from pytorch/ao@02105d4, verified in a fresh py3.12 env (mxfp8 op registers, a2a imports). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> * Sort imports in expert_parallel.py (ruff I001) Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> * Fix comments * Fix alignment of 32 for grouped gemm backward * Apply Samis Suggestions * Add back default ignore patterns * Fix name bug * Reformat with ruff * Config verifier * Fix kernel redirect for > 32 token groups * Ruff reformat * chore: update uv.lock for torchao dependency Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * style: fix import sorting in MXFP8 files Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * Move config check to config code * Fix ruff * Move config check to trainer --------- Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary