RMSNorm + RHT + Amax (SM100)
RMSNorm + RHT + Amax (SM100)
This is an experimental API and subject to change.
Overview
RMSNorm + RHT + amax: A fused CUTE DSL kernel for NVIDIA Blackwell GPUs (SM100+) that applies RMS normalization, a block-diagonal Hadamard transform with fixed block size 16, and a per-CTA amax reduction.
This frontend integration exposes the kernel as a standard FE-OSS Python API with:
- a class API (
RmsNormRhtAmaxSm100) - a wrapper API (
rmsnorm_rht_amax_wrapper_sm100) - grouped-gemm-style regression coverage for compile/execute, wrapper use, and cache reuse
Shapes
-
Inputs
X: activation tensor, shape(M, N)W: RMSNorm scale tensor, shape(N,)
-
Outputs
O: fused RMSNorm + RHT output tensor, shape(M, N)Amax: per-CTA max-abs tensor, shape(M / rows_per_cta,)
rows_per_cta is the number of rows reduced into each amax element.
Equations
For each row m:
Then apply the fixed Hadamard transform blockwise over 16-wide chunks:
where H_16 is the 16 x 16 Hadamard matrix and b indexes each 16-element block in the hidden dimension.
For each CTA covering rows_per_cta rows:
over every element produced by that CTA.
API Usage
High-level wrapper
When no overrides are supplied, the wrapper uses the upstream-tuned thread table when available and an upstream-style rows_per_cta heuristic.
Class API
Parameters
Input and output tensors
x_tensor/sample_x- Shape:
(M, N) - Layout: row-major contiguous
- Dtype:
torch.bfloat16
- Shape:
w_tensor/sample_w- Shape:
(N,) - Layout: contiguous
- Dtype:
torch.bfloat16
- Shape:
o_tensor/sample_o- Shape:
(M, N) - Layout: row-major contiguous
- Dtype:
torch.bfloat16
- Shape:
amax_tensor/sample_amax- Shape:
(M / rows_per_cta,) - Dtype:
torch.float32
- Shape:
Common parameters
eps: float- RMSNorm epsilon. Default:
1e-5
- RMSNorm epsilon. Default:
num_threads: Optional[int]- Threads per CTA. If omitted, the API uses the upstream-tuned table when possible, otherwise a valid fallback search.
rows_per_cta: Optional[int]- Rows processed by each CTA. If omitted, the wrapper uses the upstream-style heuristic over
{2, 4, 8}.
- Rows processed by each CTA. If omitted, the wrapper uses the upstream-style heuristic over
- CUDA stream (
current_stream)
Wrapper return values
Returns a TupleDict with keys:
o_tensoramax_tensor
Tuple unpacking order is (o_tensor, amax_tensor).
Support surface and constraints
- Requires SM100+.
Nmust be divisible by16.Nmust be divisible by the resolvednum_threads.EPT = N / num_threadsmust be at least8and divisible by8.Mmust be divisible byrows_per_cta.- Inputs and output are currently bf16 only.
- The frontend integration matches the upstream RMSNorm kernel semantics; it does not expose full LayerNorm mean/bias behavior.
Verification
Focused correctness and cache coverage live in:
test/python/fe_api/norm/test_rmsnorm_rht_amax.py