Repository navigation
Conversation
… op/type/target allow warp_reduction always emitted the log2(32) = 5 round shfl.sync.bfly tree even though warp_redux_sync already wraps the single-instruction redux.sync. Add a structural resolver that maps recognised operators (operator.add, cutlass_dsl.max/min, operator.and_/or_/xor for Int32/Uint32; fmax/fmin and partial(fmax, nan=True) for Float32) to a redux.sync kind, gate it on the PTX ISA target rules (integer forms sm_80+, .f32 forms sm_100 family + CUDA 12.9), and fall back to the shuffle tree for everything else (custom lambdas, other types, sub-warp groups, fmax(abs=True)/ftz=True, older targets). warp_reduction_max/warp_reduction_sum now pass named module-level operators instead of lambdas so they take the fast path too. Verified on B200 (sm_100a, CUDA 13.3): all fast-path cases emit exactly one redux.sync and zero shfl.sync; fallback cases are unchanged; device results match the pre-change implementation bit-for-bit on the new test file.
|
This PR has been labeled |
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.
Summary
cute.arch.warp_reductionalways emitted the butterfly tree —log2(32) = 5rounds ofshfl.sync.bfly+op— even thoughcute.arch.warp_redux_syncalready wraps thesingle-instruction PTX
redux.sync, and the C++ side already does this dispatch(
include/cutlass/functional.h,CUTLASS_ARCH_CREDUX_ENABLED). Kernels that wanted thefast instruction had to bypass
warp_reductionand callwarp_redux_syncby hand(e.g.
mixed_input_fmha_decode.py,gqa_decode_*.py).This PR makes
warp_reductionlower to oneredux.syncwhen — and only when — that isbit-for-bit equivalent to the shuffle tree, and keeps the tree otherwise. No public API
changes;
warp_reduction_max/warp_reduction_sumpick up the fast path automatically.Dispatch rules
The fast path is taken iff all of the following hold; otherwise the existing shuffle
tree is emitted unchanged.
threads_in_group == 32—redux.synchas no sub-warp groupsInt32,Uint32orFloat32(the only operand widthsredux.syncaccepts)operator.add,cutlass_dsl.max/min,operator.and_/or_/xorcute.arch.fmax/fmin, or `functools.partial(fmaxarch >= sm_80(PTX ISA 7.0)arch.is_family_of(sm_100f)and CUDA ≥ 12.9 — PTX ISA: ".f32type requires sm_100a and is supported on sm_100f from PTX ISA 8.8". Plainsm_100/sm_103(no suffix) fall back.Deliberately not mapped (kept on the shuffle tree):
partial(fmax, abs=True): NVVMfmax(abs=True)is the xorsign-abs formsign(a ^ b) * max(|a|, |b|), whereasredux.sync.absreduces magnitudes only — different results.ftz=True: noredux.synccounterpart.Int64,Int16,Boolean, …Semantic equivalence of the mapped cases:
cutlass_dsl.max/minare signedness-aware andwarp_redux_syncpromotesmax/min→umax/uminforUint32.fmax/fmindefault to NaN-quietmaximumNumber;redux.sync.max.f32without.NaNlikewise ignores NaN inputs (canonical NaN only if all inputs are NaN).
nan=True↔.NaN.Both treat
+0.0 > -0.0.shfl.syncwith fullmask vs
redux.syncwith full membermask), so the convergence contract is unchanged.Changes
python/CuTeDSL/cutlass/cute/arch/nvvm_wrappers.py_warp_reduction_redux_kind(value_type, op, threads_in_group)— pure structural resolverreturning
(kind, nan)orNone._warp_reduction_redux_supported(value_type)— arch / toolchain gate per PTX ISA.warp_reductionconsults both and callswarp_redux_sync, else the unchanged shuffle loop.warp_reduction_max/warp_reduction_sumuse named module-level operators(
_warp_reduction_max_op,_warp_reduction_sum_op) with identical bodies instead of lambdas,so they are recognisable.
test/examples/CuTeDSL/test_warp_reduction.py(new)NaN-propagating, all-NaN, ±0.0 ties, denormals,
ftz, xorsign-abs,warp_reduction_max/sum)checked against per-warp NumPy references;
cute.compile(..., options="--gpu-arch <sm> --keep-ptx ...")for sm_90a / sm_100 / sm_100a / sm_100f / sm_120a asserting exactly one
redux.syncand noshfl.syncwhere the target allows it, and the unchanged shuffle count otherwise — these runon any GPU;
Verification (B200, sm_100a, CUDA 13.3, DSL runtime 4.6.2 + this patch)
Test file: 139 passed on the patched tree (also with
CUTE_DSL_ARCH=sm_100, i.e. forced f32fallback on the same GPU). The 38 device cases also pass bit-for-bit against the unpatched
install, i.e. the references describe pre-existing behaviour and the change is
semantics-preserving — including ±0.0 ties, all-NaN warps (canonical NaN on both paths), NaN
payload handling, and denormal inputs with and without
ftz.Emitted PTX per
warp_reductioncall (CUTE_DSL_KEEP=ptx):fmax/fmin/warp_reduction_maxredux.sync.max.f32/.min.f32, 0 shflshfl.sync.bflypartial(fmax, nan=True)redux.sync.max.NaN.f32warp_reduction_sum)partial(fmax, abs=True)fmax,threads_in_group=16redux.sync.add.s32/.max.s32/.xor.b32/.max.u32threads_in_group=8Microbenchmark, loop-carried
acc = warp_reduction(acc) + x× 4096 iterations,2368 blocks × 256 threads, median of 20 launches (shuffle path forced via an unrecognised lambda
with the same body):
redux.syncfmaxOutputs of the two paths are identical (
torch.equal).Notes for reviewers
fmax/fmin/cutlass_dsl.max/minobjects (thedsl_user_op-wrapped functions users get viacute.arch.*), and unwrapsfunctools.partialonly when it has no positional args.operator,cutlass.base_dsl,BaseDSL,target_version) follownumeric_conversion.py, which already imports the same names from the same modules.