Skip to content

fix(pt): enable second-order autograd for tabulate descriptors - #5537

Merged
OutisLi merged 2 commits into
deepmodeling:masterfrom
njzjz:fix/pt-tabulate-second-order-autograd
Jun 20, 2026
Merged

OutisLi merged 2 commits into
deepmodeling:masterfrom
njzjz:fix/pt-tabulate-second-order-autograd

Conversation

@njzjz

@njzjz njzjz commented Jun 15, 2026 •

Copy link
Copy Markdown
Member

Summary

  • wrap PyTorch tabulate descriptor first-derivative kernels in autograd Functions
  • connect se_a, se_atten, se_t, se_r, and se_t_tebd backward paths to existing grad-grad kernels
  • add second-order backward regression tests for all affected tabulate descriptor ops

Fixes #4994.

Tests

  • pytest source/tests/pt/test_tabulate_fusion_se_a.py source/tests/pt/test_tabulate_fusion_se_atten.py source/tests/pt/test_tabulate_fusion_se_r.py source/tests/pt/test_tabulate_fusion_se_t.py source/tests/pt/test_tabulate_fusion_se_t_tebd.py -q
  • ruff check .
  • ruff format .
  • commit hook suite, including clang-format

Summary by CodeRabbit

  • New Features

    • Improved/extended second-order gradient support for tabulate fusion operations (embeddings and attention), improving correctness for higher-order differentiation workflows.
  • Tests

    • Added second-order backward tests across multiple tabulate fusion ops.
    • Introduced shared finite-difference-based test utilities to validate autograd second-order gradients (including numerical tolerance checks).

@coderabbitai

coderabbitai Bot commented Jun 15, 2026 •

Copy link
Copy Markdown
Contributor

Review Change Stack

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration

Configuration used: Repository UI

Review profile: CHILL

Plan: Pro

Run ID: eb113315-2e19-4076-a85d-7b5b576e1ee7

📥 Commits

Reviewing files that changed from the base of the PR and between de7ebac and b0b422d.

📒 Files selected for processing (6)
  • source/tests/pt/tabulate_test_utils.py
  • source/tests/pt/test_tabulate_fusion_se_a.py
  • source/tests/pt/test_tabulate_fusion_se_atten.py
  • source/tests/pt/test_tabulate_fusion_se_r.py
  • source/tests/pt/test_tabulate_fusion_se_t.py
  • source/tests/pt/test_tabulate_fusion_se_t_tebd.py

📝 Walkthrough

Walkthrough

Adds second-order autograd support to all five PyTorch TabulateFusion embedding paths (SeA, SeAtten, SeT, SeR, SeTTebd) by introducing GradOp/GradGradOp torch::autograd::Function wrappers. Each Op's backward_t now delegates via apply() instead of calling *GradForward kernels directly. Introduces a shared test utility for finite-difference-based second-order gradient validation. Five corresponding test_second_order_backward test methods are added.

Changes

Second-order autograd for all FusionSe* ops

Layer / File(s) Summary
SeA GradOp and GradGradOp introduction and wiring
source/op/pt/tabulate_multi_device.cc
Introduces TabulateFusionSeAGradOp (forward_t calls TabulateFusionSeAGradForward, backward_t calls TabulateFusionSeAGradGradForward) and TabulateFusionSeAGradGradOp (dtype-dispatched forward_t computing dz_dy_tensor). TabulateFusionSeAOp::backward_t is rewired to delegate to TabulateFusionSeAGradOp::apply(...).
SeAtten GradOp refactor and is_sorted persistence
source/op/pt/tabulate_multi_device.cc
Refactors TabulateFusionSeAttenGradOp: forward_t saves is_sorted into ctx->saved_data; backward_t retrieves it, applies zeros_like for undefined grad outputs, and calls TabulateFusionSeAGradGradForward. TabulateFusionSeAttenOp::forward_t adds ctx->saved_data["is_sorted"]; backward_t delegates to TabulateFusionSeAttenGradOp::apply(...).
SeT, SeR, and SeTTebd GradOp wrappers
source/op/pt/tabulate_multi_device.cc
Introduces TabulateFusionSeTGradOp, TabulateFusionSeRGradOp, and TabulateFusionSeTTebdGradOp, each with forward_t calling the matching *GradForward kernel and backward_t calling the matching *GradGradForward kernel. The three parent Ops' backward_t bodies are updated to delegate via apply(...).
Test utilities for second-order backward validation
source/tests/pt/tabulate_test_utils.py
New utility module defining _nonuniform_like (reproducible nonuniform tensor), _project_first_grads (weighted scalar projection of first-order gradients), and assert_second_order_backward_matches_finite_difference (validates second-order gradients against central finite-difference approximations using dtype-dependent tolerances).
Second-order backward tests for all five ops
source/tests/pt/test_tabulate_fusion_se_a.py, source/tests/pt/test_tabulate_fusion_se_atten.py, source/tests/pt/test_tabulate_fusion_se_r.py, source/tests/pt/test_tabulate_fusion_se_t.py, source/tests/pt/test_tabulate_fusion_se_t_tebd.py
Adds test_second_order_backward to each test class. Each test executes the corresponding forward op, extracts descriptor_tensor, and validates second-order backward behavior by invoking assert_second_order_backward_matches_finite_difference against finite-difference approximations.

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~60 minutes

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 4.44% which is insufficient. The required threshold is 80.00%. 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 accurately describes the main change: enabling second-order autograd for tabulate descriptors is the core objective of this PR.
Linked Issues check ✅ Passed All code changes directly address issue #4994 by implementing the missing autograd wrappers for se_a, se_atten, se_t, se_r, and se_t_tebd descriptors' second-order derivatives.
Out of Scope Changes check ✅ Passed All changes are in-scope: the main implementation refactors autograd Functions to enable second-order derivatives, and test additions validate this functionality across all affected descriptor types.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests

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 and usage tips.

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

🧹 Nitpick comments (1)
source/op/pt/tabulate_multi_device.cc (1)

607-608: ⚡ Quick win

Remove unused private member variable.

The device member is declared but never used anywhere in TabulateFusionSeAGradOp. This appears to be leftover from an earlier implementation.

 class TabulateFusionSeAGradOp
     : public torch::autograd::Function<TabulateFusionSeAGradOp> {
- private:
-  std::string device;
-
  public:
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@source/op/pt/tabulate_multi_device.cc` around lines 607 - 608, Remove the
unused private member variable `device` from the `TabulateFusionSeAGradOp` class
since it is declared but never referenced anywhere in the implementation. Simply
delete the line containing `std::string device;` from the private section.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Nitpick comments:
In `@source/op/pt/tabulate_multi_device.cc`:
- Around line 607-608: Remove the unused private member variable `device` from
the `TabulateFusionSeAGradOp` class since it is declared but never referenced
anywhere in the implementation. Simply delete the line containing `std::string
device;` from the private section.

ℹ️ Review info
⚙️ Run configuration

Configuration used: Repository UI

Review profile: CHILL

Plan: Pro

Run ID: 10e5cae3-359e-4d0a-b565-f7fc90453f06

📥 Commits

Reviewing files that changed from the base of the PR and between 87d8557 and de7ebac.

📒 Files selected for processing (6)
  • source/op/pt/tabulate_multi_device.cc
  • source/tests/pt/test_tabulate_fusion_se_a.py
  • source/tests/pt/test_tabulate_fusion_se_atten.py
  • source/tests/pt/test_tabulate_fusion_se_r.py
  • source/tests/pt/test_tabulate_fusion_se_t.py
  • source/tests/pt/test_tabulate_fusion_se_t_tebd.py

@codecov

codecov Bot commented Jun 15, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 87.64045% with 22 lines in your changes missing coverage. Please review.
✅ Project coverage is 82.18%. Comparing base (87d8557) to head (b0b422d).
⚠️ Report is 248 commits behind head on master.

Files with missing lines Patch % Lines
source/op/pt/tabulate_multi_device.cc 87.64% 18 Missing and 4 partials ⚠️
Additional details and impacted files
@@            Coverage Diff             @@
##           master    #5537      +/-   ##
==========================================
- Coverage   82.18%   82.18%   -0.01%     
==========================================
  Files         890      896       +6     
  Lines      101358   102845    +1487     
  Branches     4240     4363     +123     
==========================================
+ Hits        83301    84520    +1219     
- Misses      16756    16970     +214     
- Partials     1301     1355      +54     

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

@njzjz
njzjz requested review from OutisLi and wanghan-iapcm June 15, 2026 04:38
Comment thread source/tests/pt/test_tabulate_fusion_se_a.py Outdated
@njzjz
njzjz requested a review from wanghan-iapcm June 18, 2026 10:11
@OutisLi
OutisLi added this pull request to the merge queue Jun 20, 2026
Merged via the queue into deepmodeling:master with commit b1cd6dc Jun 20, 2026
70 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

PyTorch Backend: Missing autograd wrapper for se_t descriptor's second-order derivatives

3 participants