Skip to content

[mxfp8 training] require torch nightly and cuda 12.8+ - #4141

Merged
danielvegamyhre merged 1 commit into
mainfrom
assert321
Mar 24, 2026
Merged

danielvegamyhre merged 1 commit into
mainfrom
assert321

Conversation

@danielvegamyhre

Copy link
Copy Markdown
Contributor

Summary

@danielvegamyhre danielvegamyhre added mx module: training quantize_ api training flow labels Mar 23, 2026
@pytorch-bot

pytorch-bot Bot commented Mar 23, 2026 •

Copy link
Copy Markdown

🔗 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 (image):

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.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Mar 23, 2026
@danielvegamyhre
danielvegamyhre requested a review from vkuzo March 23, 2026 17:56

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, (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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:

  1. fail when on a version of PyTorch without the core fix
  2. pass on nightly, stable or built-from-source versions of PyTorch with the core fix

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

i updated it to check if pytorch installation has certain attributes that indicate the necessary PRs were included

@danielvegamyhre
danielvegamyhre merged commit 808d7b6 into main Mar 24, 2026
22 of 23 checks passed
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. module: training quantize_ api training flow mx

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants