Skip to content

[CuTeDSL] Lower whole-warp warp_reduction to a single redux.sync when op/type/target allow - #3592

Open
zkyue wants to merge 1 commit into
NVIDIA:mainfrom
zkyue:feat/cutedsl-warp-reduction-redux
Open

zkyue wants to merge 1 commit into
NVIDIA:mainfrom
zkyue:feat/cutedsl-warp-reduction-redux

Conversation

@zkyue

@zkyue zkyue commented Sep 7, 2026

Copy link
Copy Markdown
Contributor

Summary

cute.arch.warp_reduction always emitted the butterfly tree — log2(32) = 5 rounds of
shfl.sync.bfly + op — even though cute.arch.warp_redux_sync already wraps the
single-instruction PTX redux.sync, and the C++ side already does this dispatch
(include/cutlass/functional.h, CUTLASS_ARCH_CREDUX_ENABLED). Kernels that wanted the
fast instruction had to bypass warp_reduction and call warp_redux_sync by hand
(e.g. mixed_input_fmha_decode.py, gqa_decode_*.py).

This PR makes warp_reduction lower to one redux.sync when — and only when — that is
bit-for-bit equivalent to the shuffle tree, and keeps the tree otherwise. No public API
changes; warp_reduction_max / warp_reduction_sum pick 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.

condition detail
group threads_in_group == 32 — redux.sync has no sub-warp groups
type exactly Int32, Uint32 or Float32 (the only operand widths redux.sync accepts)
operator (int) operator.add, cutlass_dsl.max / min, operator.and_ / or_ / xor
operator (f32) cute.arch.fmax / fmin, or `functools.partial(fmax
target (int) arch >= sm_80 (PTX ISA 7.0)
target (f32) arch.is_family_of(sm_100f) and CUDA ≥ 12.9 — PTX ISA: ".f32 type requires sm_100a and is supported on sm_100f from PTX ISA 8.8". Plain sm_100 / sm_103 (no suffix) fall back.

Deliberately not mapped (kept on the shuffle tree):

  • partial(fmax, abs=True): NVVM fmax(abs=True) is the xorsign-abs form
    sign(a ^ b) * max(|a|, |b|), whereas redux.sync.abs reduces magnitudes only — different results.
  • ftz=True: no redux.sync counterpart.
  • arbitrary lambdas / custom callables (cannot be inspected), Int64, Int16, Boolean, …

Semantic equivalence of the mapped cases:

  • integer add wraps mod 2³² on both paths; cutlass_dsl.max/min are signedness-aware and
    warp_redux_sync promotes max/min → umax/umin for Uint32.
  • fmax/fmin default to NaN-quiet maximumNumber; redux.sync.max.f32 without .NaN
    likewise ignores NaN inputs (canonical NaN only if all inputs are NaN). nan=True ↔ .NaN.
    Both treat +0.0 > -0.0.
  • both paths require every lane of the warp to execute the reduction (shfl.sync with full
    mask vs redux.sync with 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 resolver
      returning (kind, nan) or None.
    • _warp_reduction_redux_supported(value_type) — arch / toolchain gate per PTX ISA.
    • warp_reduction consults both and calls warp_redux_sync, else the unchanged shuffle loop.
    • warp_reduction_max / warp_reduction_sum use 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)
    • 38 device cases (Int32 / Uint32 / Float32 × operators, sub-warp groups, NaN-quiet,
      NaN-propagating, all-NaN, ±0.0 ties, denormals, ftz, xorsign-abs, warp_reduction_max/sum)
      checked against per-warp NumPy references;
    • compile-only codegen tests: cute.compile(..., options="--gpu-arch <sm> --keep-ptx ...")
      for sm_90a / sm_100 / sm_100a / sm_100f / sm_120a asserting exactly one redux.sync and no
      shfl.sync where the target allows it, and the unchanged shuffle count otherwise — these run
      on any GPU;
    • GPU-free unit tests of the operator resolver and the arch predicate.

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 f32
fallback 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_reduction call (CUTE_DSL_KEEP=ptx):

case sm_100a sm_100 (no suffix)
f32 fmax / fmin / warp_reduction_max 1× redux.sync.max.f32 / .min.f32, 0 shfl 5× shfl.sync.bfly
f32 partial(fmax, nan=True) 1× redux.sync.max.NaN.f32 5× shfl
f32 add (warp_reduction_sum) 5× shfl (unchanged) 5× shfl
f32 partial(fmax, abs=True) 5× shfl (unchanged) 5× shfl
f32 fmax, threads_in_group=16 4× shfl (unchanged) 4× shfl
i32 add / max / xor, u32 max 1× redux.sync.add.s32 / .max.s32 / .xor.b32 / .max.u32 same
i32 add, threads_in_group=8 3× shfl (unchanged) same
i32 custom lambda 5× shfl (unchanged) same

Microbenchmark, 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):

op shuffle tree redux.sync speedup
f32 fmax 1.353 ms 0.273 ms 4.95×
i32 add 1.350 ms 0.628 ms 2.15×

Outputs of the two paths are identical (torch.equal).

Notes for reviewers

  • The resolver compares callables by identity against the module-level fmax/fmin/
    cutlass_dsl.max/min objects (the dsl_user_op-wrapped functions users get via
    cute.arch.*), and unwraps functools.partial only when it has no positional args.
  • New imports (operator, cutlass.base_dsl, BaseDSL, target_version) follow
    numeric_conversion.py, which already imports the same names from the same modules.

… 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.
@github-actions

github-actions Bot commented Oct 7, 2026

Copy link
Copy Markdown

This PR has been labeled inactive-30d due to no recent activity in the past 30 days. Please close this PR if it is no longer required. Otherwise, please respond with a comment indicating any updates. This PR will be labeled inactive-90d if there is no activity in the next 60 days.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant