Repository navigation
perf(logp): fused zero-workspace backward for the generic CUDA fused logp (#174) - #495
fusheng-ji wants to merge 2 commits into
Conversation
…sequences The generic CUDA fused logp VJP materialized FP32 softmax and FP32 grad over the full [N, V] logits, about 10 bytes per logit of transient memory (46 GiB peak at 32k rows x 152k vocab). Compute it in row chunks of at most 2^27 FP32 elements instead; the VJP is row-local, so the result is bitwise identical to the unchunked path for any chunk size. Rows whose target is outside [0, V) already return a constant 0 from the forward kernel; give them a zero gradient instead of wrapping negative targets (e.g. -100) into a real vocab index or failing on targets >= V. Part of RL-Align#174 (stage 1: workspace sizing and boundary safety).
Add _C.fused_logp_backward, which writes dlogits = grad * (one_hot - softmax) directly in the logits dtype with one block per row. Its max / sum-exp passes visit each thread's elements in the same order as fused_logp_forward_kernel (same 256-thread block and reductions), so the softmax is bitwise consistent with the forward logp; loads are batched 8 at a time for memory-level parallelism without changing the accumulation order. The backward needs no workspace beyond the returned gradient, and is deterministic and batch-invariant. Out-of-range targets get a zero gradient, matching the forward's constant output. _FusedLogpAutograd uses it when the extension provides it and keeps the row-chunked FP32 path as the fallback. Forward + backward at 32768 x 151936 BF16 on B200: 57.1 ms / 46.4 GiB peak (original) -> 27.2 ms / 9.3 GiB peak, the peak being the gradient itself. Part of RL-Align#174 (stage 2).
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (5)
Included review availability: This review used your included allowance. Your plan provides up to 2 included reviews per hour; 1 remain after this review. 📝 WalkthroughWalkthroughAdds a CUDA backward kernel for fused log-probability and exposes it through the Python extension. The autograd path uses the kernel when available and otherwise computes gradients in bounded row chunks. Tests cover both paths, input validation, gradient values, and workspace limits. ChangesFused log-probability backward
Priority: ➖ Normal Estimated code review effort: 3 (Moderate) | ~25 minutes Change: Bug fix Sequence Diagram(s)sequenceDiagram
participant AutogradBridge
participant BackendExtension
participant CudaKernel
AutogradBridge->>BackendExtension: call fused_logp_backward when available
BackendExtension->>CudaKernel: launch backward kernel
CudaKernel-->>BackendExtension: return logits gradients
BackendExtension-->>AutogradBridge: return logits gradients
AutogradBridge->>AutogradBridge: otherwise compute gradients in row chunks
Merge Risk: ⚪ Minimal · up to No actionable merge-blocking issue was established. Merge after the normal build and test checks for the supported hardware. Security Architecture ReviewSecurity architecture risk: 🔵 Low · up to The change reduces temporary GPU memory use without establishing a new external trust boundary. One multi-GPU execution assumption remains: the new native backward does not ensure that the active GPU matches its input tensors. Exposure appears limited to callers within the training process. Retained concerns
Security review detailsSecurity Blast Radius
Trust Boundaries and Controls
Resilience and Maintainability Implications
Hardening Proposals
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
Summary
Addresses #174 (long-sequence logp memory safety and workspace sizing) for the generic CUDA fused logp op (
FusedLogpGenericOp/_FusedLogpAutogradinrl_engine/kernels/ops/cuda/loss/logp.py).The previous backward materialized FP32 softmax and an FP32 gradient over the full
[N, V]logits, about 10 bytes of transient memory per logit. At 32k response tokens × 152k vocab, forward + backward peaked at 46 GiB. This is where 8k–32k CoT training runs out of memory.The PR has two commits:
fused_logp_backward_chunked). The VJP is computed over row chunks of at mostFUSED_LOGP_BWD_CHUNK_ELEMS = 2^27FP32 elements, so the transient workspace is ~1 GiB regardless ofN. The output is bitwise identical to the previous implementation for any chunk size. This path is kept as the fallback for_Cbuilds that do not have the new kernel._C.fused_logp_backward). It computesdlogits = grad * (one_hot - softmax)with one block per row and writes it directly in the logits dtype, so the backward allocates nothing beyond the returned gradient.fused_logp_forward_kernel, with the same 256-thread block and the sameblockReduceMax/blockReduceSum. The softmax used in the backward is therefore bitwise consistent with the forward logp.kFusedLogpBackwardUnroll) for memory-level parallelism. Accumulation order does not change: the unrolled kernel is bitwise identical to the non-unrolled version on all checked cases._FusedLogpAutograd.backwarduses the kernel when the backend provides it and otherwise uses the chunked path.Both paths also fix target handling at the boundary; see the behavior changes below.
Not touched: the forward kernels, the SM90
fused_logp_sm90path,linear_logp, and the PyTorch / Triton / Ascend logp ops.Behavior changes
0for targets outside[0, V), such as the ignore index-100. The old backward indexedprobs[rows, labels]directly. Python wrapped-100into a real vocab column, giving a nonzero gradient inconsistent with the forward, and targets>= Vraised an indexing error. Those rows now get an exactly zero gradient.torch.softmax-based VJP, because the softmax now comes from the forward kernel's own reductions.--use_fast_math, soexpfis the fast intrinsic. The forward kernel has the same property.Results
Single B200, V = 151936, BF16 logits, median of 5. All variants call the same
_C.fused_logpforward through_FusedLogpAutograd; only the backward differs. Peak memory is measured from before the step and includes the returned BF16 gradient (2.3 / 4.6 / 9.3 GiB).Forward + backward
Backward only
The fused backward reads each row 3 times and writes it once. At 32k rows it runs about 2.4× above the HBM floor (~5 ms at ~8 TB/s). A single-pass online max/sum would remove one read, but it would break the bitwise consistency with the forward's two-pass reduction, so this PR does not do it.
The chunk size of the fallback was chosen from this sweep at N = 32768:
Tests
New
tests/test_fused_logp_chunked_backward.py(46 tests).Chunked fallback:
log_softmax+gathergradient-100,V,V+5) while other rows stay bitwise unchangedFused kernel:
-inf-masked vocab tail: zero gradient, all finiteDispatch:
FusedLogpGenericOpautograd uses the fused kernelfused_logp_backwardfalls back to the chunked pathWorkspace at 8k / 16k / 32k rows × 151936 vocab:
All 34 skips need hardware not present (SM90-only kernels, Hopper-only comparisons, Ascend), plus one CUDA-dispatch guard in
test_logp.py.Environment: 1× NVIDIA B200, CUDA 13.0, PyTorch 2.13.0+cu130, Python 3.12. The extension was built with
TORCH_CUDA_ARCH_LIST=10.0 python setup.py build_ext --inplace. Formatting was checked with black 24.4.2 and isort 5.13.2 (--line-length=100), as pinned in.pre-commit-config.yaml.Not verified:
fused_logp_backwardis registered next tofused_logpand is also built for ROCm. It only uses the existing HIP-compatible helpers (fused_logp_shfl_down_32,blockReduce*), but I have not compiled or run it on ROCm.Benchmark script (forward + backward)
Relation to other work
#234 also edits the backward in
rl_engine/kernels/ops/cuda/loss/logp.pyandcsrc/fused_logp_kernel.cu. It currently conflicts withmain; whichever lands second will need a rebase.Summary by CodeRabbit