Repository navigation
perf(pt): fuse the mixing-stack per-layer backward chain (opt-in) - #6060
long-yi-2019 wants to merge 2 commits into
Conversation
…D_FUSED_INFER) Each gated layer's backward runs as recompute -> pointwise -> GEMM, materializing the sigmoid surface and the pre-activation gradient to HBM and reading both straight back. With the flag on (and the supported regime: fp32, narrow channels, no higher-order graph), one kernel walks the chain with both surfaces in registers, matching the relay bit-for-bit; the m = 0 and |m| = 1 column segments launch separately to keep the lighter segment off the heavy one's register profile. The 2.15 GB standing scratch never allocates; on the N=4096 eval workload it does not coincide with the frame's peak.
for more information, see https://pre-commit.ci
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (3)
Included review availability: This review used your included allowance. Your plan provides up to 8 included reviews per hour; 7 remain after this review. 📝 WalkthroughWalkthroughThe mixing-stack backward adds an opt-in fused Triton path for eligible fp32 gradients. The path recomputes gate sigmoids and writes input gradients directly. Tests compare fused and unfused gradients and verify higher-order gradient behavior and fallback dispatch. ChangesFused mixing-stack backward
Priority: ⬇️ Low Estimated code review effort: 3 (Moderate) | ~25 minutes Change: Feature Sequence Diagram(s)sequenceDiagram
participant Traversal as Backward traversal
participant Launcher as _launch_stack_bwd_fused
participant Kernel as _stack_bwd_fused_kernel
participant Buffer as Spare gradient buffer
Traversal->>Launcher: Select fused path for an eligible layer
Launcher->>Kernel: Launch m = 0 and |m| = 1 blocks
Kernel->>Buffer: Write next-layer gradient columns
Suggested reviewers: Merge Risk: ⚪ Minimal · up to No merge-blocking issue is established for the opt-in fused backward path; complete the normal CUDA test checks before merging. 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches 💡 1🧪 Generate unit tests (beta)
🛠️ Fix failing CI checks 💡
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
njzjz-bot
left a comment
There was a problem hiding this comment.
Agent: dot
Reviewed the three-file delta from dbca0b1 to 4455e67, including the existing recompute/pointwise/GEMM implementations, launch configuration lookup, autograd routing, and new tests. One nonblocking regression-test issue is noted inline; source inspection did not establish a production kernel correctness defect.
Validation: eight isolated checks of the exact-source environment parser passed. Repository tests and CUDA/Triton numerical or compilation checks were not run: this executor has neither PyTorch nor Triton. The author-reported CUDA results were not independently reproduced. At review time all seven PR-triggered Actions runs reported action_required (the CUDA run has no jobs); pre-commit passed and Read the Docs reported failure. No CI runs were activated or rerun.
|
|
||
| def test_non_fp32_fall_back(self): | ||
| """The fused kernel assumes fp32 rounding; other dtypes keep the relay.""" | ||
| calls, out = self._record_launches(3, 32, dtype=torch.bfloat16) |
There was a problem hiding this comment.
Agent: dot
[P3] Enable fusion before asserting that unsupported inputs fall back
This test and test_wide_channels_fall_back never set DP_STACK_BWD_FUSED_INFER, and _record_launches does not set it either. In the normal default-off test run, fused short-circuits at use_stack_bwd_fused(), so these assertions pass without exercising the unsupported-dtype/channel dispatch contract. In particular, removing the fp32 guard would still leave this bf16 regression test green. Please scope both calls with mock.patch.dict(os.environ, {"DP_STACK_BWD_FUSED_INFER": "1"}) (or equivalent) so the tests actually request the feature they expect to reject.
Motivation
Each gated layer of the mixing-stack backward runs as a three-kernel relay: recompute the gate sigmoids, pointwise backward writing the pre-activation gradient
gz, then the transposed-weight GEMM readinggzback. At the SeZM N=4096 eval workload (E ~ 646k, Cf = 32, lmax = 3) the relay materializes two surfaces per stack and walks them through HBM every gated layer:sig(F, E, L*Cf)fp32, ~0.5 GB, written by the recompute kernel and read straight back by the pointwise kernel;gz(F, E, ROW), ~1.65 GB, written by the pointwise kernel and read straight back by the GEMM.Change
Opt-in via
DP_STACK_BWD_FUSED_INFER(default OFF, read when the backward traversal runs viause_stack_bwd_fused(), matching the repository's convention of construction/call-time reads for these switches). When enabled, and in the supported regime — fp32,Cf < GATE_BMM_MIN_FOCUS_DIM(no cuBLAS gate path),Cf == BLOCK_Kso one K tile is exactly one row group, no weight gradients, no second-order surfaces, no upstream cotangents, and no higher-order gradient graph being built — a single kernel per layer replaces the relay: each program recomputes its row group's gate sigmoid, forms thegztile in registers with exactly the pointwise kernel's operator sequence, and accumulates the GEMM with exactly the backward GEMM kernel'stl.dotorder, so results match the relay bit-for-bit.The
m = 0block (which carries the full scalar chain) and the|m| = 1stripes launch as two PART passes over disjoint column segments: a shared kernel would push the light segment onto the heavy one's register profile (BLOCK_M=64/BN=64 spills 226 slots; the shipped 32/64 config stays in budget).Outside the regime the dispatch falls back to the existing relay unchanged. The fused kernel registers no next-derivative formula;
create_graph=True(e.g. force-loss training) keeps the operator that carries the hand-derived second order — the official backward already routes those differentiations structurally, andnot torch.is_grad_enabled()at the dispatch is a second belt.Measurements
SeZM model, N=4096 atoms (~646k edges), RTX 5090, torch 2.12.1+cu130,
DP_COMPILE_INFER=1,DP_TRITON_INFER=3, throughput mode (16 evals, P50):Peak VRAM is unchanged at this workload: the eliminated 2.15 GB standing scratch does not coincide with the frame's peak allocation on this shape (the forward does); what it removes is the allocation itself and the per-layer HBM round-trip of both surfaces. Bit-identity of the fused kernel against the relay is pinned by
TestSeZMStackBwdFused(fp32,apply_alphaon and off); two-state compiled-eval accuracy at N=4096: energy rel diff 3.7e-10, max force diff 4.9e-6 with max |f| ~ 3.9.Tests
TestSeZMStackBwdFused(CUDA): fused-vs-relay bitwise cross-check; dispatch follows the env flag (recording stand-in); higher-order backward preserved with the flag on; wide-channel (Cf >= GATE_BMM_MIN_FOCUS_DIM) and non-fp32 fallbacks.TestSeZMStackBwdFlagParser(CPU): env parsing.source/tests/pt/model/test_descriptor_sezm_triton.py(43 passed) andtest_descriptor_sezm.py(61 passed) on the branch with the flag off (default).Summary by CodeRabbit
DP_STACK_BWD_FUSED_INFERsetting; it is disabled by default. Unsupported configurations and higher-order gradient calculations continue to use the existing backward path.