Skip to content

[PyTorch] Add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK to warn on grouped MLP fusion fallback - #3631

Draft
ghadiaravi13 wants to merge 4 commits into
NVIDIA:mainfrom
ghadiaravi13:rghadia/warn-cutedsl-grouped-mlp-fallback
Draft

ghadiaravi13 wants to merge 4 commits into
NVIDIA:mainfrom
ghadiaravi13:rghadia/warn-cutedsl-grouped-mlp-fallback

Conversation

@ghadiaravi13

@ghadiaravi13 ghadiaravi13 commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

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 onto main once #3584 merges.

With #3584, TE fuses GroupedLinear + activation + GroupedLinear into 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.0 was installed alongside nvidia-cutlass-dsl==4.8.0. cutlass-dsl 4.8 renamed ReduxKind to ReductionKind (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 of NVTE_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:

UserWarning: Not fusing GroupedLinear + ScaledSReLU + GroupedLinear into GroupedMLP_CuTeGEMMUnary (failed to import cuDNN frontend kernels: grouped_gemm_srelu_wrapper_sm100 requires optional dependencies. Install with 'pip install nvidia-cudnn-frontend[cutedsl]': cannot import name 'ReduxKind' from 'cutlass._mlir.dialects.nvvm'). Falling back to unfused ops, which can be significantly slower and use more memory.

Fusion decisions are identical with or without the variable, and nothing is printed when it is unset.

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

  • Warn at each grouped MLP fusion fallback, but only when the ops contain a GroupedLinear + activation + GroupedLinear pattern that the fused op handles. Covered reasons:
    • system: CUDA unavailable, unsupported compute capability, cuDNN frontend too old, or kernels failing to import;
    • recipe: not MXFP8/NVFP4, NVFP4 without RHT or with 4over6, MXFP8 with E5M2 gradients (Format.HYBRID), or SReLU with UE5M3 scales;
    • activation not supported on this GPU or cuDNN frontend;
    • invalid dimensions.
  • Each warning names the fused op, the activation, and the reason, and is emitted once per (op, activation, reason).
  • Add _GroupedMLP_CuTeGEMMBase._unsupported_reason(); is_supported() is now _unsupported_reason() is None and stays lru_cached. GroupedMLP_CuTeGEMMUnary overrides _unsupported_reason() for its cuDNN frontend >= 1.24.0 requirement.
  • Factor the recipe checks of _is_grouped_mlp_fusion_candidate() into _grouped_mlp_recipe_unsupported_reason() so the warning can state the reason.
  • Document the variable in 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:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings (adds an opt-in warning)
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes (new test checked against the patched code in isolation; full suite pending CI)

zhongbozhu and others added 3 commits October 2, 2026 16:57
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Signed-off-by: zhongboz <zhongboz@nvidia.com>
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Oct 5, 2026
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>
@ghadiaravi13
ghadiaravi13 force-pushed the rghadia/warn-cutedsl-grouped-mlp-fallback branch from 10d8b64 to 5974f7b Compare October 5, 2026 23:18
@ghadiaravi13 ghadiaravi13 changed the title fix(pytorch): warn when CuTeDSL fused grouped MLP is requested but unavailable [PyTorch] Add NVTE_CUTEDSL_FUSED_GROUPED_MLP_WARN_FALLBACK to warn on grouped MLP fusion fallback Oct 5, 2026

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

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"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Nit:

Suggested change
return "not installed"
return "cuDNN Frontend is not installed"

Comment on lines 1012 to +1013

def _grouped_mlp_fallback_warnings_enabled() -> bool:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Memoizing will help us reduce how often we query the environment:

Suggested change
def _grouped_mlp_fallback_warnings_enabled() -> bool:
@functools.cache
def _grouped_mlp_fallback_warnings_enabled() -> bool:

Comment on lines +1071 to +1079
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}"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

#3325 added support for custom recipes to enable this fusion.

Suggested change
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__})"

Comment on lines +191 to +193
monkeypatch.setattr(
fused_op_cls, "_unsupported_reason", classmethod(lambda cls: "broken cutlass-dsl")
)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Shouldn't you only set this if recipe="unsupported_system"?

Suggested change
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")
)

Comment on lines +194 to +196
if reason == "unsupported_device":
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(grouped_mlp_module, "get_device_compute_capability", lambda: (9, 0))

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

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 zhongbozhu left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

why do we just warn the fallback instead of giving us the option to crash the training if we cannot trigger the fusion?

@ghadiaravi13

Copy link
Copy Markdown
Contributor Author

@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.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants