Skip to content

perf(logp): fused zero-workspace backward for the generic CUDA fused logp (#174) - #495

Open
fusheng-ji wants to merge 2 commits into
RL-Align:mainfrom
fusheng-ji:fix/174-chunked-logp-backward
Open

fusheng-ji wants to merge 2 commits into
RL-Align:mainfrom
fusheng-ji:fix/174-chunked-logp-backward

Conversation

@fusheng-ji

@fusheng-ji fusheng-ji commented Oct 7, 2026 •

Copy link
Copy Markdown

Summary

Addresses #174 (long-sequence logp memory safety and workspace sizing) for the generic CUDA fused logp op (FusedLogpGenericOp / _FusedLogpAutograd in rl_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:

  1. Bounded row-chunked FP32 backward (fused_logp_backward_chunked). The VJP is computed over row chunks of at most FUSED_LOGP_BWD_CHUNK_ELEMS = 2^27 FP32 elements, so the transient workspace is ~1 GiB regardless of N. The output is bitwise identical to the previous implementation for any chunk size. This path is kept as the fallback for _C builds that do not have the new kernel.
  2. Fused CUDA backward kernel (_C.fused_logp_backward). It computes dlogits = 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.
    • Its max / sum-exp passes visit each thread's elements in the same order as fused_logp_forward_kernel, with the same 256-thread block and the same blockReduceMax / blockReduceSum. The softmax used in the backward is therefore bitwise consistent with the forward logp.
    • Loads are batched 8 at a time (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.
    • It is deterministic and batch-invariant, with one block per row and a fixed reduction tree.
    • _FusedLogpAutograd.backward uses 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_sm90 path, linear_logp, and the PyTorch / Triton / Ascend logp ops.

Behavior changes

  • Out-of-range targets. The forward kernel already returns a constant 0 for targets outside [0, V), such as the ignore index -100. The old backward indexed probs[rows, labels] directly. Python wrapped -100 into a real vocab column, giving a nonzero gradient inconsistent with the forward, and targets >= V raised an indexing error. Those rows now get an exactly zero gradient.
  • Gradient numerics for valid targets. With the fused kernel, valid-target gradients are no longer bit-for-bit equal to the old torch.softmax-based VJP, because the softmax now comes from the forward kernel's own reductions.
    • Measured max abs difference against the chunked/old path: ~1e-8 for BF16, ~8e-6 for FP16, ~2e-6 for FP32.
    • The FP32 difference comes from the extension being built with --use_fast_math, so expf is the fast intrinsic. The forward kernel has the same property.
    • Against an FP64 reference, BF16 and FP16 error is identical to the old path.
    • On builds without the kernel, the chunked fallback is bitwise identical to the old path.

Results

Single B200, V = 151936, BF16 logits, median of 5. All variants call the same _C.fused_logp forward 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

N rows Before Chunked fallback Fused kernel
8192 11.59 GiB / 14.5 ms 3.32 GiB / 15.0 ms 2.32 GiB / 7.0 ms
16384 23.18 GiB / 28.7 ms 5.64 GiB / 29.8 ms 4.64 GiB / 13.7 ms
32768 46.37 GiB / 57.1 ms 10.27 GiB / 59.3 ms 9.27 GiB / 27.2 ms

Backward only

N rows Before Chunked fallback Fused kernel
8192 10.6 ms 11.2 ms 3.2 ms
16384 21.2 ms 22.3 ms 6.2 ms
32768 42.2 ms 44.4 ms 12.2 ms

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:

chunk_elems rows/chunk transient time
2^22 27 0.03 GiB 114.8 ms
2^24 110 0.13 GiB 54.8 ms
2^26 441 0.50 GiB 48.8 ms
2^27 883 1.00 GiB 44.4 ms
2^28 1766 2.00 GiB 43.3 ms

Tests

New tests/test_fused_logp_chunked_backward.py (46 tests).

Chunked fallback:

  • bitwise equal to the previous VJP for FP32/BF16 at chunk sizes of 1/3/7/64 rows (uneven tails included), on CPU and CUDA
  • agreement with the autograd log_softmax + gather gradient
  • zero rows for out-of-range targets (-100, V, V+5) while other rows stay bitwise unchanged

Fused kernel:

  • agreement with an FP64 reference for FP32/BF16/FP16 at V = 1, 255, 1031 and 151936
  • bitwise run-to-run determinism, and batch invariance across row subsets
  • out-of-range targets and an -inf-masked vocab tail: zero gradient, all finite
  • input validation (grad dtype, target dtype, length mismatch) and empty input

Dispatch:

  • FusedLogpGenericOp autograd uses the fused kernel
  • a backend without fused_logp_backward falls back to the chunked path

Workspace at 8k / 16k / 32k rows × 151936 vocab:

  • fused: zero bytes beyond the output
  • chunked: within the bound
  • skipped when free device memory is insufficient
pytest -q tests/test_fused_logp_chunked_backward.py
# 46 passed

pytest -q tests/test_fused_logp_chunked_backward.py tests/test_logp.py \
  tests/test_deterministic_logp.py tests/test_batch_invariant_logp.py \
  tests/test_op_accuracy.py tests/test_rl_kernel_loss_step.py tests/test_logprob_comparison.py \
  tests/test_gradient_invariance.py tests/test_ws1_gtest_gpu.py \
  tests/test_ws1_chain_integration.py tests/test_tolerance_contract.py
# 327 passed, 34 skipped

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:

  • ROCm. fused_logp_backward is registered next to fused_logp and 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.
  • Hopper (SM90) and Ampere were not run. The kernel uses no architecture-specific features.
Benchmark script (forward + backward)
import torch
import rl_engine.kernels.ops.cuda.loss.logp as m
from rl_engine.kernels.ops.base import _C

def original_vjp(logits, labels, g, dt):
    probs = torch.softmax(logits.float(), dim=-1)
    rows = torch.arange(logits.size(0), device=logits.device)
    probs[rows, labels] -= 1.0
    return (-g.reshape(-1, 1).float() * probs).to(dt)

class ForwardOnly:  # a _C build without fused_logp_backward
    fused_logp = staticmethod(_C.fused_logp)

chunked = m.fused_logp_backward_chunked
V = 151936
for n in (8192, 16384, 32768):
    x0 = torch.randn(n, V, device="cuda", dtype=torch.bfloat16)
    y = torch.randint(0, V, (n,), device="cuda")
    for name, backend, vjp in (("before", ForwardOnly(), original_vjp),
                               ("chunked", ForwardOnly(), chunked),
                               ("fused", _C, chunked)):
        m.fused_logp_backward_chunked = vjp
        def step():
            x = x0.detach().requires_grad_(True)
            m._FusedLogpAutograd.apply(x, y, backend).float().sum().backward()
            return x.grad
        step()  # warm-up
        torch.cuda.synchronize(); torch.cuda.reset_peak_memory_stats()
        base = torch.cuda.memory_allocated()
        s, e = torch.cuda.Event(True), torch.cuda.Event(True); ts = []
        for _ in range(5):
            s.record(); g = step(); e.record(); torch.cuda.synchronize()
            ts.append(s.elapsed_time(e)); del g
        peak = (torch.cuda.max_memory_allocated() - base) / 2**30
        print(n, name, f"peak={peak:.2f}GiB", f"t={sorted(ts)[2]:.1f}ms")
    m.fused_logp_backward_chunked = chunked
    del x0; torch.cuda.empty_cache()

Relation to other work

#234 also edits the backward in rl_engine/kernels/ops/cuda/loss/logp.py and csrc/fused_logp_kernel.cu. It currently conflicts with main; whichever lands second will need a rebase.

Summary by CodeRabbit

  • New Features
    • Added fused CUDA support for log-probability gradient computation, with a chunked fallback that limits temporary memory use.
    • Invalid target IDs now produce zero gradients in both computation paths.
  • Bug Fixes
    • Improved handling of long sequences by avoiding a full-tensor softmax workspace during backward computation.

…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).
@coderabbitai

coderabbitai Bot commented Oct 7, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 10adb04d-5b52-4bd5-97f8-a43129d09fd8
📥 Commits

Reviewing files that changed from the base of the PR and between 43f150f and 607fc3f.

📒 Files selected for processing (5)
  • csrc/fused_logp_kernel.cu
  • csrc/ops.cpp
  • rl_engine/_C.pyi
  • rl_engine/kernels/ops/cuda/loss/logp.py
  • tests/test_fused_logp_chunked_backward.py

Included review availability: This review used your included allowance. Your plan provides up to 2 included reviews per hour; 1 remain after this review.


📝 Walkthrough

Walkthrough

Adds 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.

Changes

Fused log-probability backward

Layer / File(s) Summary
CUDA kernel and backend binding
csrc/fused_logp_kernel.cu, csrc/ops.cpp, rl_engine/_C.pyi
Adds and exports fused_logp_backward. The CUDA entry point validates inputs, handles empty batches, and launches a per-row kernel that computes the gradient for valid targets and zeros gradients for invalid targets.
Python dispatch, fallback, and validation
rl_engine/kernels/ops/cuda/loss/logp.py, tests/test_fused_logp_chunked_backward.py
The autograd path uses the fused operation when available and otherwise uses an FP32 row-chunked fallback. Tests compare both paths with references and check invalid targets, validation, and workspace bounds.

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
Loading

Merge Risk: ⚪ Minimal · up to 607fc

No actionable merge-blocking issue was established. Merge after the normal build and test checks for the supported hardware.

Security Architecture Review

Security architecture risk: 🔵 Low · up to 607fc

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

  • Medium · reliability · inferred: The newly exposed native backward checks that its tensors share a CUDA device but does not bind execution to that device. A direct caller whose active GPU differs from the tensors can pass validation while the kernel uses the active GPU's current stream, potentially causing invalid device access or incorrect ordering of tensor reads and writes. This weakens multi-GPU failure containment relative to the former PyTorch backward. Existing native forward shares the assumption, and an actual failure through normal autograd was not demonstrated.
Security review details

Security Blast Radius

  • inferred — The demonstrated reachability is an in-process Python extension call and the existing generic operator path. Exercising the device-context gap requires control of CUDA tensor placement and active-device context within that process. The inspected evidence does not establish remote reachability, privilege gain, or cross-tenant access.

Trust Boundaries and Controls

  • observed — The native boundary independently validates tensor placement, ranks, dtypes, and cardinality rather than relying solely on Python normalization. These controls constrain raw-pointer inputs, but same-device equality does not establish that the launch's active CUDA device matches those allocations.

Resilience and Maintainability Implications

  • inferred — Removing full-tensor backward intermediates reduces memory-exhaustion pressure for long sequences. It does not impose a total workload quota: logits and the returned gradient still scale with row count times vocabulary size, and fallback workspace must accommodate at least one row.

Hardening Proposals

  • proposed — Make the native backward's device and stream ownership explicit by guarding the logits device before preparation and launch. Validate direct calls with tensors on a non-current GPU and non-default streams, including input readiness and output lifetime.
🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 10.34% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 29 functions across 5 files. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly and concisely describes the main change: a fused, zero-workspace backward path for generic CUDA fused log-probability computation.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
  • Fix all pre-merge checks with AI
✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

@Flink-ddd Flink-ddd added component: kernels Tasks involving the development of CUDA and Triton underlying operators platform: cuda Specific optimizations or bugs in NVIDIA graphics cards (such as FlashInfer, TMA optimizations) type: performance Performance optimization tasks aimed at increasing throughput and reducing latency etc. labels Oct 8, 2026

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

Labels

component: kernels Tasks involving the development of CUDA and Triton underlying operators platform: cuda Specific optimizations or bugs in NVIDIA graphics cards (such as FlashInfer, TMA optimizations) type: performance Performance optimization tasks aimed at increasing throughput and reducing latency etc.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants