Skip to content

[Common] Fall back from gated TMA kernels when shared memory does not fit - #3600

Open
ravimajeti wants to merge 3 commits into
NVIDIA:mainfrom
ravimajeti:codex/issue-3299
Open

ravimajeti wants to merge 3 commits into
NVIDIA:mainfrom
ravimajeti:codex/issue-3299

Conversation

@ravimajeti

Copy link
Copy Markdown
Contributor

Description

On GPUs with less shared memory per block than SM 10.0, the delayed-scaling gated activation path
(SwiGLU, GeGLU, ReGLU, QGeGLU, SReGLU) crashes for FP32 inputs. On an RTX 5090 (SM 12.0, opt-in limit
101,376 B per block), cast_fp8_gated_kernel requests 131,232 B for FP32→FP32 forward and
114,848–164,000 B for any FP32-input backward. cudaFuncSetAttribute fails with
CUDA Error: invalid argument. Dispatch selected the TMA kernel on every CC ≥ 10.0 device whenever
cols % 32 == 0, without checking the kernel's shared memory request against the device.

This PR computes the kernel's shared memory requirement in one place, compares it with the device's
cached opt-in per-block limit, and uses the existing non-TMA gated kernels when it does not fit. This
applies to both forward and backward. Configurations that fit keep the TMA kernel, including BF16/FP16
on SM 12.0 and every configuration on SM 10.0. The check does not rely on an architecture list, so it
also covers SM 12.1 and any other device with a smaller shared memory budget.

Fixes #3299

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

  • common/util/cuda_runtime.{h,cpp}: add cuda::max_shared_memory_per_block_optin(), which queries
    cudaDevAttrMaxSharedMemoryPerBlockOptin once per device and caches it (same pattern as sm_count()).
  • common/cast/fp8/gated_fp8.cuh: add cast_gated_tma_dynamic_shmem_size(), the single source for the
    kernel's dynamic shared memory size, now also used by the launcher. Add
    cast_gated_tma_fits_device() (dynamic + static shared memory ≤ device limit).
  • common/cast/dispatch/gated.cuh: forward and backward NVTE_DELAYED_TENSOR_SCALING dispatch use the
    TMA kernel only if it also fits the device. Otherwise they use the existing cast_gated_fwd /
    cast_gated_bwd path. The device query only runs on CC ≥ 10.0.

Testing

All on 1× RTX 5090 (CC 12.0, driver 580.105.08), CUDA 12.9.2, PyTorch 2.11.0+cu128, cuDNN 9.19,
NVTE_CUDA_ARCHS=120. Unpatched vs patched at eb56ba22. After rebasing onto a16bce3a, I re-ran the
issue reproducer (6/6 pass), the C++ GLU suite (250/250), and the PyTorch swiglu numerics (48/48).

Check Unpatched Patched
Issue reproducer (LayerNormMLP swiglu, FP32/BF16/FP16 × fwd, fwd+bwd) FP32 fails all pass
test_operator --gtest_filter=*ActTestSuite*GLU* 25 failed (5 GLUs × FP32 input × 5 output types) 250/250
test_operator --gtest_filter=*Act* – 1015/1015
test_numerics.py -k "test_layernorm_mlp_accuracy and swiglu" (L0 env vars) 16 failed (FP32) 48/48
test_numerics.py -k layernorm_mlp 96 CUDA errors + 16 other 0 CUDA errors; same 16 other + 16 newly-running (see note)
test_sanity.py -k "layernorm_mlp and swiglu" – 1536 passed, 2304 skipped

Kernel selection (torch.profiler): FP32 uses gated_act_kernel / dgated_act_kernel. BF16/FP16 still
use cast_fp8_gated_kernel for forward and backward.

Note: FP32 relu (not gated, unaffected by this PR) fails test_layernorm_mlp_accuracy on this GPU with
gradient mismatches, identically before and after. With the fix, FP32 reglu now runs and shows the same
mismatch. Both appear to come from TF32 GEMMs vs. the FP32 PyTorch reference: with
NVIDIA_TF32_OVERRIDE=0, reglu passes 48/48 and relu 47/48. This is pre-existing and outside the
scope of this PR.

Performance (CUDA events, median of 3 × 200 iterations):

  • BF16/FP16 before vs after: same kernel, within ±0.16% for the activation and ≤0.4% for LayerNormMLP
    (run-to-run noise). Host time per call is unchanged (the limit is cached).
  • For reference, on SM 12.0 the non-TMA kernel is not slower than TMA for this op. At large
    bandwidth-bound shapes the two are within 0.6%; at small shapes the non-TMA kernel is 1.2–1.7× faster.
    So falling back costs nothing here. SM 10.0 was not measured.

No new tests: the existing C++ and PyTorch tests above already cover the failing configurations and fail
without this change on SM 12.0.

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
  • 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

🤖 Generated with Claude Code

… fit

The delayed-scaling gated activation dispatch selected the TMA kernel
(cast_fp8_gated_kernel) on every device with compute capability 10.0+,
without checking its dynamic shared memory request against the device
limit. On SM 12.0 (opt-in limit 99 KiB per block) FP32 configurations
request up to 160 KiB, so cudaFuncSetAttribute fails with
"invalid argument" in both forward and backward (issue NVIDIA#3299).

Compute the kernel's shared memory requirement in one helper shared by
dispatch and launch, compare it with the device's cached opt-in
per-block limit, and use the existing non-TMA kernels when it does not
fit. Configurations that fit (e.g. BF16/FP16 on SM 12.0, everything on
SM 10.0) keep using the TMA kernel.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Signed-off-by: Ravi M <ravitejamajeti@gmail.com>
@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Oct 1, 2026
@ravimajeti
ravimajeti marked this pull request as ready for review October 1, 2026 06:45
@greptile-apps

greptile-apps Bot commented Oct 1, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[High risk] Adds runtime checks for GPU shared memory constraints in kernel dispatch.

The reviewed changes appear safe to merge; no actionable new failure was established.

Summary

The PR adds a device shared-memory fit check before selecting gated TMA kernels and falls back to the existing non-TMA path when the kernel does not fit. It also centralizes the device-limit query used by swizzle.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  A[Gated delayed-scaling dispatch] --> B{Aligned and SM 10.0+?}
  B -- No --> F[Non-TMA gated kernel]
  B -- Yes --> C{Dynamic + compiled static memory fits?}
  C -- Yes --> T[TMA gated kernel]
  C -- No --> F
Loading

Reviews (2) · Last reviewed commit: "[Common] Read gated TMA kernel static sh..."

@denera denera self-assigned this Oct 2, 2026
@denera
denera self-requested a review October 2, 2026 16:58

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

Thanks for the PR!

Looks good to me overall pending the two minor changes I recommended below.

Comment thread transformer_engine/common/util/cuda_runtime.cpp
Comment thread transformer_engine/common/cast/fp8/gated_fp8.cuh Outdated
ravimajeti and others added 2 commits October 3, 2026 16:19
Signed-off-by: Ravi M <ravitejamajeti@gmail.com>
… kernel

The fit check hard-coded the static shared memory of cast_fp8_gated_kernel
as its mbarrier array (32 B) and missed the staging array declared in
reduce_max (64 B). Query the compiled kernel instance instead via
cudaFuncGetAttributes, exposed as cuda::static_shared_memory_size() and
cached per kernel and device, so __shared__ arrays in called device
functions are always included.

Also replace the file-local get_max_dynamic_smem() in swizzle.cu with
cuda::max_shared_memory_per_block_optin().

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Signed-off-by: Ravi M <ravitejamajeti@gmail.com>
@ravimajeti
ravimajeti requested a review from denera October 4, 2026 00:53
@ravimajeti

Copy link
Copy Markdown
Contributor Author

Hi @denera thanks a lot for the quick review and the pointers, i've addressed your comments :)

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.

[Bug] Transformer Engine SM120 (Blackwell) Compatibility Issue

2 participants