Repository navigation
[mxfp8 moe training] refactor autograd func - #4176
Conversation
🔗 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 FailureAs of commit 001dc09 with merge base 02105d4 ( NEW FAILURE - The following job has failed:
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
07c67ab to
ba0003f
Compare
| scale_calculation_mode, | ||
| ) | ||
| else: | ||
| return _compute_fwd_auto( |
There was a problem hiding this comment.
the _auto suffix is weird here, just remove it?
There was a problem hiding this comment.
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). |
There was a problem hiding this comment.
do the triton kernels match the native pytorch reference exactly? not blocking this PR, i'm just curious
There was a problem hiding this comment.
Yes, there are no new kernels used in this PR btw, this is still using the existing triton kernel, tested here:
ao/test/prototype/mx_formats/test_kernels.py
Lines 478 to 512 in 4611835
i am planning to integrate the new cutedsl kernel for this in follow ups
vkuzo
left a comment
There was a problem hiding this comment.
skimmed it and looks good, lmk if you need a proper review
stack-info: PR: #4176, branch: danielvegamyhre/stack/159
ba0003f to
001dc09
Compare
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_wgradall 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_bwdpytest test/prototype/moe_training/test_training.py -s