DSv4.1 mHC projection/RMS backward
DSv4.1 mHC projection/RMS backward
This experimental SM100 API computes the input and projection-weight gradients of the projection/RMS stage in DSv4.1 mHC. It expects the upstream projection and RMS gradients after the gate and normalization derivatives:
The two matrix products convert operands to TF32 and accumulate in FP32.
The RMS contribution uses FP32 arithmetic. dX is stored in BF16 and dW in
FP32. Callers must explicitly pass allow_tf32=True; this API does not promise
full FP32 matrix-product precision. The reduction of partial dW is in a fixed
order, without atomics. Outputs are overwritten, not accumulated.
The supported training profile has 4 residual streams, hidden size 5120 and 4096 tokens per rank. This is one backward stage, not the complete mHC forward or backward. Sinkhorn, gate derivatives and parameter-gradient accumulation remain the caller’s responsibility.
Requirements and tensor contract
Use a Blackwell SM100 GPU, PyTorch, cuda-tile>=1.5, a compatible system
tileiras compiler and the Frontend cutile,triton optional dependencies:
Install the PyTorch CUDA build matching your environment separately.
All tensors use compact row-major strides, with stride (1, 1) for the column
vectors. The last 8 columns of grad_proj are ignored, including if nonzero.
r must contain the positive RMS values from forward. All operands and
workspace must be on one device, disjoint and 16-byte aligned. Other shapes,
strides, devices and precisions are rejected. No conversion or repacking is
performed.
Allocating wrapper
The wrapper allocates outputs and workspace on the requested stream. Callers must establish producer/consumer stream dependencies and keep tensors alive until execution completes.
Prepared class API
compile() compiles and loads both kernels without executing tensor work.
execute() enqueues exactly two kernels and allocates no tensor storage.
Workspace is 15,728,640 bytes (15 MiB). Use separate scratch and outputs for
overlapping executions. Prepared execution supports CUDA Graph capture; keep
its buffers alive for all replays. current_stream accepts a PyTorch stream,
a CUDA stream handle, or None for the current stream of the tensors’ device.
Source attribution
The projection kernel derives from Megatron-LM’s
_ct_fused_grad_x_weight_kernel
under BSD-3-Clause. This implementation partitions the token reduction across
independent blocks, writes partial weight gradients, then reduces them in a
second kernel. The source notice is retained in _kernels.py.