Skip to content

perf(pt): fuse the mixing-stack per-layer backward chain (opt-in) - #6060

Open
long-yi-2019 wants to merge 2 commits into
deepmodeling:masterfrom
long-yi-2019:pr/stack-bwd-fused
Open

long-yi-2019 wants to merge 2 commits into
deepmodeling:masterfrom
long-yi-2019:pr/stack-bwd-fused

Conversation

@long-yi-2019

@long-yi-2019 long-yi-2019 commented Oct 8, 2026 •

Copy link
Copy Markdown

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 reading gz back. 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 via use_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_K so 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 the gz tile in registers with exactly the pointwise kernel's operator sequence, and accumulates the GEMM with exactly the backward GEMM kernel's tl.dot order, so results match the relay bit-for-bit.

The m = 0 block (which carries the full scalar chain) and the |m| = 1 stripes 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, and not 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):

build per-eval (ms) steady peak VRAM alloc (MB)
master 153.2 22740
this PR (flag on) 150.6 (-1.7%) 22742 (unchanged)

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_alpha on 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.
  • Full source/tests/pt/model/test_descriptor_sezm_triton.py (43 passed) and test_descriptor_sezm.py (61 passed) on the branch with the flag off (default).

Summary by CodeRabbit

  • New Features
    • Added an opt-in fused backward path for supported gated mixing-stack layers, which can speed up gradient computation during training.
    • Control the option with the DP_STACK_BWD_FUSED_INFER setting; it is disabled by default. Unsupported configurations and higher-order gradient calculations continue to use the existing backward path.

…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.
Copilot AI balanced review requested due to automatic review settings October 8, 2026 11:44

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@github-actions github-actions Bot added the Python label Oct 8, 2026
@coderabbitai

coderabbitai Bot commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

Review in Change Stack →

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: Repository UI
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: d26cab04-572d-4531-998f-c73f2ebf883d
📥 Commits

Reviewing files that changed from the base of the PR and between dbca0b1 and 4455e67.

📒 Files selected for processing (3)
  • deepmd/pt_expt/kernels/triton/sezm/so2_value_path.py
  • deepmd/pt_expt/kernels/utils.py
  • source/tests/pt/model/test_descriptor_sezm_triton.py

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


📝 Walkthrough

Walkthrough

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

Changes

Fused mixing-stack backward

Layer / File(s) Summary
Environment switch
deepmd/pt_expt/kernels/utils.py, source/tests/pt/model/test_descriptor_sezm_triton.py
Adds use_stack_bwd_fused(), which reads DP_STACK_BWD_FUSED_INFER. Tests cover accepted truthy values, false values, unrecognized values, and an unset variable.
Fused kernel and launch
deepmd/pt_expt/kernels/triton/sezm/so2_value_path.py
Adds a kernel that recomputes gate sigmoids and pre-activation gradients, contracts them with block weights, and writes gradient columns for m = 0 and `
Eligibility, dispatch, and validation
deepmd/pt_expt/kernels/triton/sezm/so2_value_path.py, source/tests/pt/model/test_descriptor_sezm_triton.py
Selects the fused path when the flag and eligibility conditions are met. Eligible calls bypass the pointwise-backward and block-GEMM relay. Tests check gradient equivalence, higher-order gradients, dispatch counts, and fallback cases.

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
Loading

Suggested reviewers: outisli

Merge Risk: ⚪ Minimal · up to 4455e

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)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 60.00% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 20 functions across 3 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: an opt-in performance fusion of the mixing-stack per-layer backward chain.
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 💡 1
🧪 Generate unit tests (beta)
  • Create a new PR
🛠️ Fix failing CI checks 💡
  • Commit to this branch
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

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.

❤️ Share

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

@njzjz-bot njzjz-bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants