Repository navigation
[Common] Fall back from gated TMA kernels when shared memory does not fit - #3600
Open
ravimajeti wants to merge 3 commits into
Open
ravimajeti wants to merge 3 commits into
ravimajeti wants to merge 3 commits into
Conversation
… 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>
ravimajeti
marked this pull request as ready for review
October 1, 2026 06:45
Contributor
|
denera
self-requested a review
October 2, 2026 16:58
denera
requested changes
Oct 2, 2026
denera
left a comment
Collaborator
There was a problem hiding this comment.
Thanks for the PR!
Looks good to me overall pending the two minor changes I recommended below.
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>
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_kernelrequests 131,232 B for FP32→FP32 forward and114,848–164,000 B for any FP32-input backward.
cudaFuncSetAttributefails withCUDA Error: invalid argument. Dispatch selected the TMA kernel on every CC ≥ 10.0 device whenevercols % 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
Changes
common/util/cuda_runtime.{h,cpp}: addcuda::max_shared_memory_per_block_optin(), which queriescudaDevAttrMaxSharedMemoryPerBlockOptinonce per device and caches it (same pattern assm_count()).common/cast/fp8/gated_fp8.cuh: addcast_gated_tma_dynamic_shmem_size(), the single source for thekernel'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 backwardNVTE_DELAYED_TENSOR_SCALINGdispatch use theTMA kernel only if it also fits the device. Otherwise they use the existing
cast_gated_fwd/cast_gated_bwdpath. 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 ateb56ba22. After rebasing ontoa16bce3a, I re-ran theissue reproducer (6/6 pass), the C++ GLU suite (250/250), and the PyTorch swiglu numerics (48/48).
LayerNormMLPswiglu, FP32/BF16/FP16 × fwd, fwd+bwd)test_operator --gtest_filter=*ActTestSuite*GLU*test_operator --gtest_filter=*Act*test_numerics.py -k "test_layernorm_mlp_accuracy and swiglu"(L0 env vars)test_numerics.py -k layernorm_mlptest_sanity.py -k "layernorm_mlp and swiglu"Kernel selection (torch.profiler): FP32 uses
gated_act_kernel/dgated_act_kernel. BF16/FP16 stilluse
cast_fp8_gated_kernelfor forward and backward.Note: FP32
relu(not gated, unaffected by this PR) failstest_layernorm_mlp_accuracyon this GPU withgradient mismatches, identically before and after. With the fix, FP32
reglunow runs and shows the samemismatch. Both appear to come from TF32 GEMMs vs. the FP32 PyTorch reference: with
NVIDIA_TF32_OVERRIDE=0,reglupasses 48/48 andrelu47/48. This is pre-existing and outside thescope of this PR.
Performance (CUDA events, median of 3 × 200 iterations):
(run-to-run noise). Host time per call is unchanged (the limit is cached).
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:
🤖 Generated with Claude Code