Repository navigation
[PyTorch] Add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK to warn on grouped MLP fusion fallback - #3631
[PyTorch] Add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK to warn on grouped MLP fusion fallback#3631ghadiaravi13 wants to merge 4 commits into
Conversation
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
With NVIDIA#3584, TE fuses GroupedLinear + activation + GroupedLinear into the CuTeDSL grouped MLP ops whenever possible and silently falls back to unfused ops otherwise. A fallback can cost a large share of end-to-end performance and memory, and its cause, for example a cuDNN frontend and cutlass-dsl mismatch that makes the kernels fail to import, is invisible. When NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK=1, warn once per fused op, activation, and reason whenever a grouped MLP pattern is left unfused: CUDA unavailable, unsupported GPU, cuDNN frontend too old or kernels failing to import, or an unsupported recipe, activation, or dimensions. Nothing changes when the variable is unset. Co-authored-by: Claude <noreply@anthropic.com> Signed-off-by: Ravi Ghadia <rghadia@nvidia.com>
10d8b64 to
5974f7b
Compare
timmoon10
left a comment
There was a problem hiding this comment.
This is not a bad idea. It's not fully airtight (you won't get any warning about SiTUGLU on Rubin, or QGeGLU on Rubin with CFE 1.29) and you might get some false positives (maybe SwiGLU is broken, but you don't care since your model uses SReLU), but for an opt-in diagnostic that's not too critical.
| try: | ||
| return get_pkg_version("nvidia-cudnn-frontend") | ||
| except PackageNotFoundError: | ||
| return "not installed" |
There was a problem hiding this comment.
Nit:
| return "not installed" | |
| return "cuDNN Frontend is not installed" |
|
|
||
| def _grouped_mlp_fallback_warnings_enabled() -> bool: |
There was a problem hiding this comment.
Memoizing will help us reduce how often we query the environment:
| def _grouped_mlp_fallback_warnings_enabled() -> bool: | |
| @functools.cache | |
| def _grouped_mlp_fallback_warnings_enabled() -> bool: |
| if not (recipe.mxfp8() or recipe.nvfp4()): | ||
| return f"requires an MXFP8 or NVFP4 recipe, got {type(recipe).__name__}" | ||
| # NVFP4 graph-safe grouped quantize currently requires RHT. | ||
| if recipe.nvfp4() and recipe.disable_rht: | ||
| return "NVFP4 requires RHT, but the recipe sets disable_rht=True" | ||
| # The fused MXFP8 backward reinterprets grad-output storage as E4M3. It | ||
| # cannot consume E5M2 gradients from Format.HYBRID. NVFP4 has separate formats. | ||
| if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: | ||
| return f"MXFP8 requires E4M3 gradients, but the recipe uses {recipe.fp8_format}" |
There was a problem hiding this comment.
#3325 added support for custom recipes to enable this fusion.
| if not (recipe.mxfp8() or recipe.nvfp4()): | |
| return f"requires an MXFP8 or NVFP4 recipe, got {type(recipe).__name__}" | |
| # NVFP4 graph-safe grouped quantize currently requires RHT. | |
| if recipe.nvfp4() and recipe.disable_rht: | |
| return "NVFP4 requires RHT, but the recipe sets disable_rht=True" | |
| # The fused MXFP8 backward reinterprets grad-output storage as E4M3. It | |
| # cannot consume E5M2 gradients from Format.HYBRID. NVFP4 has separate formats. | |
| if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: | |
| return f"MXFP8 requires E4M3 gradients, but the recipe uses {recipe.fp8_format}" | |
| if recipe.mxfp8(): | |
| if get_fp8_torch_dtype(recipe, fprop_tensor=True) != torch.float8_e4m3fn: | |
| # MXFP8 forward kernels only support E4M3 | |
| return f"MXFP8 requires E4M3 activations, but the recipe uses {recipe.fp8_format}" | |
| if get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: | |
| # MXFP8 backward kernels only support E4M3 | |
| return f"MXFP8 requires E4M3 gradients, but the recipe uses {recipe.fp8_format}" | |
| return None | |
| if recipe.nvfp4(): | |
| if recipe.disable_rht: | |
| # NVFP4 graph-safe grouped quantize currently requires RHT | |
| return "NVFP4 requires RHT, but the recipe sets disable_rht=True" | |
| return None | |
| if recipe.custom(): | |
| if not getattr(recipe, "enable_cutedsl_fused_grouped_mlp", False): | |
| return f"Custom recipe requires setting enable_cutedsl_fused_grouped_mlp=True" | |
| return None | |
| return f"Unsupported recipe ({type(recipe).__name__})" |
| monkeypatch.setattr( | ||
| fused_op_cls, "_unsupported_reason", classmethod(lambda cls: "broken cutlass-dsl") | ||
| ) |
There was a problem hiding this comment.
Shouldn't you only set this if recipe="unsupported_system"?
| monkeypatch.setattr( | |
| fused_op_cls, "_unsupported_reason", classmethod(lambda cls: "broken cutlass-dsl") | |
| ) | |
| if reason == "unsupported_system": | |
| monkeypatch.setattr( | |
| fused_op_cls, "_unsupported_reason", classmethod(lambda cls: "broken cutlass-dsl") | |
| ) |
| if reason == "unsupported_device": | ||
| monkeypatch.setattr(torch.cuda, "is_available", lambda: True) | ||
| monkeypatch.setattr(grouped_mlp_module, "get_device_compute_capability", lambda: (9, 0)) |
There was a problem hiding this comment.
This monkeypatching may work for now, but I'm wary of establishing it as a contract. If you run on Hopper, I could imagine experiencing errors even if we lie about being on Blackwell.
zhongbozhu
left a comment
There was a problem hiding this comment.
why do we just warn the fallback instead of giving us the option to crash the training if we cannot trigger the fusion?
|
@zhongbozhu I actually thought about that, and realized that a run without the fusion might still be useful. So I ended up proceeding with a warning instead of an error. Feel free to add suggestions in case I might be missing some broader implications. |
Description
Depends on #3584. This branch is stacked on #3584; only the last commit (
feat(pytorch): add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK) is new here. I'll rebase ontomainonce #3584 merges.With #3584, TE fuses
GroupedLinear+ activation +GroupedLinearinto the CuTeDSL grouped MLP ops whenever possible and falls back to unfused ops otherwise. That fallback is silent.A silent fallback is expensive to diagnose. We hit one where
nvidia-cudnn-frontend==1.26.0was installed alongsidenvidia-cutlass-dsl==4.8.0. cutlass-dsl 4.8 renamedReduxKindtoReductionKind(cuDNN frontend 1.28.0 handles both names), so the cuDNN frontend grouped GEMM kernels failed to import and every grouped MLP ran unfused. The only symptom was a large end-to-end slowdown and higher memory use, which took profiling and diffing container images to trace back.This PR adds an opt-in
NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK=1, in the spirit ofNVTE_WARN_IF_CUTEDSL_BACKEND_NOT_CHOSEN. When it is set, TE warns whenever a grouped MLP pattern is left unfused and says why. For the case above:Fusion decisions are identical with or without the variable, and nothing is printed when it is unset.
Type of change
Changes
GroupedLinear+ activation +GroupedLinearpattern that the fused op handles. Covered reasons:Format.HYBRID), or SReLU with UE5M3 scales;_GroupedMLP_CuTeGEMMBase._unsupported_reason();is_supported()is now_unsupported_reason() is Noneand stayslru_cached.GroupedMLP_CuTeGEMMUnaryoverrides_unsupported_reason()for its cuDNN frontend >= 1.24.0 requirement._is_grouped_mlp_fusion_candidate()into_grouped_mlp_recipe_unsupported_reason()so the warning can state the reason.docs/envvars.rst, and add a device-agnostic test covering system, recipe, dimension, and device fallbacks, with no warning when the variable is unset or no pattern is present.Checklist: