Skip to content

[mxfp8 moe training] refactor autograd func - #4176

Merged
danielvegamyhre merged 1 commit into
mainfrom
danielvegamyhre/stack/159
Mar 26, 2026
Merged

danielvegamyhre merged 1 commit into
mainfrom
danielvegamyhre/stack/159

Conversation

@danielvegamyhre

@danielvegamyhre danielvegamyhre commented Mar 25, 2026 •

Copy link
Copy Markdown
Contributor

Stacked PRs:


[mxfp8 moe training] refactor autograd func

Summary

The autograd func has become a bit messy with individual "emulated vs non-emulated" branches everywhere, every helper, etc.

This PR refactors the autograd func to have 2 clean code paths: auto vs emulated. The bifurcation happens early, and all code after that point just has one distinct path of kernels/functions, making it easier to reason about and modify.

_compute_fwd, _compute_dgrad, and _compute_wgrad all now have distinct "emulated" vs "auto" paths that are cleanly separated at the beginning of the path.

Tests

  • pytest test/prototype/moe_training/test_mxfp8_grouped_mm.py -k dq_fwd_bwd
  • pytest test/prototype/moe_training/test_training.py -s

@pytorch-bot

pytorch-bot Bot commented Mar 25, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4176

Note: Links to docs will display an error until the docs builds have been completed.

❌ 1 New Failure

As of commit 001dc09 with merge base 02105d4 (image):

NEW FAILURE - The following job has failed:

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@danielvegamyhre
danielvegamyhre force-pushed the danielvegamyhre/stack/159 branch from 07c67ab to ba0003f Compare March 25, 2026 23:21
@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 25, 2026
@danielvegamyhre danielvegamyhre added mx module: training quantize_ api training flow moe labels Mar 25, 2026
@danielvegamyhre

Copy link
Copy Markdown
Contributor Author

@drisspg @vkuzo i addressed the comments from #4157 but got in a crazy git state, recreated the PR than wrangle with it, apologies. ready for review.

scale_calculation_mode,
)
else:
return _compute_fwd_auto(

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.

the _auto suffix is weird here, just remove it?

@danielvegamyhre danielvegamyhre Mar 26, 2026 •

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.

how about _compute_{fwd/dgrad/wgrad}_sm100? that is really what i am trying to communicate in the code to the reader:

  • _compute_{fwd/dgrad/wgrad} = logical top level function, dispatcher based on kernel preference
  • _compute_{fwd/dgrad/wgard}_sm100 = compute with sm100 kernels
  • _compute_{fwd/dgrad/wgard}_emulated = compute with emulated logic

scale_calculation_mode: ScaleCalculationMode,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Extract qdata and scales from MXTensor or quantize using Triton kernels (AUTO path).

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.

do the triton kernels match the native pytorch reference exactly? not blocking this PR, i'm just curious

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.

Yes, there are no new kernels used in this PR btw, this is still using the existing triton kernel, tested here:

def test_triton_mxfp8_dim0_randn(M, K, scaling_mode):
x = torch.randn(M, K, dtype=torch.bfloat16, device="cuda")
x_mx_ref, x_s_ref = triton_to_mxfp8_dim0_reference(
x, block_size=32, scaling_mode=scaling_mode
)
x_mx_t, x_s_t = triton_to_mxfp8_dim0(
x,
inner_block_size=32,
scaling_mode=scaling_mode.value.lower(),
)
torch.testing.assert_close(x_mx_t, x_mx_ref, rtol=0, atol=0)
torch.testing.assert_close(x_s_t, x_s_ref, rtol=0, atol=0)
@pytest.mark.skipif(not has_triton(), reason="unsupported without triton")
@pytest.mark.skipif(
not is_sm_at_least_100() and not is_MI350(),
reason="mxfp8 requires CUDA capability 10.0 or greater or ROCm gfx950 or greater.",
)
@pytest.mark.parametrize(
"scaling_mode", (ScaleCalculationMode.FLOOR, ScaleCalculationMode.RCEIL)
)
def test_triton_mxfp8_dim0_zeros(scaling_mode):
x = torch.zeros(128, 256, dtype=torch.bfloat16, device="cuda")
x_mx_ref, x_s_ref = triton_to_mxfp8_dim0_reference(
x, block_size=32, scaling_mode=scaling_mode
)
x_mx_t, x_s_t = triton_to_mxfp8_dim0(
x,
inner_block_size=32,
scaling_mode=scaling_mode.value.lower(),
)
assert not x_mx_t.isnan().any(), "quantized tensor should not contain NaNs"
torch.testing.assert_close(x_mx_t, x_mx_ref, rtol=0, atol=0)
torch.testing.assert_close(x_s_t, x_s_ref, rtol=0, atol=0)

i am planning to integrate the new cutedsl kernel for this in follow ups

@vkuzo vkuzo left a comment

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.

skimmed it and looks good, lmk if you need a proper review

stack-info: PR: #4176, branch: danielvegamyhre/stack/159
@danielvegamyhre
danielvegamyhre marked this pull request as draft March 26, 2026 17:08
@danielvegamyhre
danielvegamyhre force-pushed the danielvegamyhre/stack/159 branch from ba0003f to 001dc09 Compare March 26, 2026 17:08
@danielvegamyhre
danielvegamyhre marked this pull request as ready for review March 26, 2026 17:09
@danielvegamyhre
danielvegamyhre merged commit efbcb0e into main Mar 26, 2026
22 of 23 checks passed
Priyjain-amd pushed a commit to Priyjain-amd/ao that referenced this pull request May 26, 2026
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 moe mx

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants