feat(tf2): support DPA4 descriptor - #5749
Conversation
|
Note Reviews pausedIt looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughThis PR adds TF2 support for DPA4/SeZM descriptors and fitting networks, registers model and component aliases, adapts PT checkpoints, extends TensorFlow-backed array operations, normalizes force shapes, and adds backend and training tests. ChangesTF2 DPA4/SeZM integration
Estimated code review effort: 5 (Critical) | ~120 minutes Possibly related PRs
Suggested reviewers: 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
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 |
There was a problem hiding this comment.
Actionable comments posted: 2
🧹 Nitpick comments (1)
deepmd/tf2/common.py (1)
362-387: 🚀 Performance & Scalability | 🔵 Trivial | ⚡ Quick win
__getattribute__allocates a newseton every attribute access.This override intercepts every attribute read on the module, and for each read it calls
_tf2_array_variable_attr_names()/_tf2_array_variable_list_attr_names(), both of which materialize a freshset(...)from the class-level tuple. On hot descriptor/fitting paths this adds a per-access allocation. Consider a cheaper membership check against the raw tuple (or a cached frozenset) to avoid rebuilding the set on every access.♻️ Cheaper membership check
def __getattribute__(self, name: str) -> Any: if not name.startswith("_tf2_"): - array_attrs = object.__getattribute__( - self, - "_tf2_array_variable_attr_names", - )() - if name in array_attrs: + if name in object.__getattribute__( + self, "_tf2_array_variable_attrs", () + ):🤖 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 `@deepmd/tf2/common.py` around lines 362 - 387, The __getattribute__ override in common.py is doing extra work on every attribute read by calling _tf2_array_variable_attr_names() and _tf2_array_variable_list_attr_names(), which rebuild sets repeatedly. Update __getattribute__ to use a cheaper membership path for the array/list attribute names, such as checking the underlying tuple directly or reusing a cached frozenset, while keeping the existing storage-name lookup and to_tensorflow_array conversion behavior unchanged.
🤖 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.
Inline comments:
In `@deepmd/dpmodel/array_api.py`:
- Around line 285-296: The TensorFlow branch in array_api.py is overwriting
prefilled -inf entries because tf.maximum(x_tensor, reduced) replaces
empty-segment sentinels with dtype minimum values. Update the
unsorted-segment-max handling in this branch so empty segments remain -inf,
matching the behavior expected by segment.py and the other backends; use the
existing x_tensor, indices_tensor, values_tensor, and reduced flow, but avoid
applying a blanket maximum that changes sentinel slots.
In `@deepmd/tf2/train/trainer.py`:
- Around line 1490-1493: The shape-iteration guard in the trainer’s
rank-checking helper only handles TypeError, but tf.TensorShape(None) can also
raise ValueError during tracing. Update the try/except around iter(shape) to
return None for both exception types in the same helper path so unknown-rank
shapes are handled safely.
---
Nitpick comments:
In `@deepmd/tf2/common.py`:
- Around line 362-387: The __getattribute__ override in common.py is doing extra
work on every attribute read by calling _tf2_array_variable_attr_names() and
_tf2_array_variable_list_attr_names(), which rebuild sets repeatedly. Update
__getattribute__ to use a cheaper membership path for the array/list attribute
names, such as checking the underlying tuple directly or reusing a cached
frozenset, while keeping the existing storage-name lookup and
to_tensorflow_array conversion behavior unchanged.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: Repository UI
Review profile: CHILL
Plan: Pro
Run ID: 76f9add6-fdb8-4b7b-8908-ab6d1ed973d8
📒 Files selected for processing (11)
deepmd/dpmodel/array_api.pydeepmd/tf2/common.pydeepmd/tf2/descriptor/__init__.pydeepmd/tf2/descriptor/dpa4.pydeepmd/tf2/fitting/__init__.pydeepmd/tf2/fitting/dpa4_ener.pydeepmd/tf2/model/ener_model.pydeepmd/tf2/model/model.pydeepmd/tf2/train/trainer.pysource/tests/consistent/descriptor/test_dpa4.pysource/tests/consistent/fitting/test_dpa4_ener.py
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## master #5749 +/- ##
==========================================
- Coverage 79.47% 79.28% -0.20%
==========================================
Files 1072 1074 +2
Lines 125055 125638 +583
Branches 4536 4569 +33
==========================================
+ Hits 99385 99607 +222
- Misses 24044 24393 +349
- Partials 1626 1638 +12 ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
There was a problem hiding this comment.
🧹 Nitpick comments (1)
deepmd/tf2/descriptor/se_atten_v2.py (1)
17-21: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valueRemove the extra refresh call here.
TF2Modulealready refreshes_refresh_tf2_trackable_lists()fordeepmd.tf2.descriptor.se_atten_v2, so this second call is redundant and can be dropped.🤖 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 `@deepmd/tf2/descriptor/se_atten_v2.py` around lines 17 - 21, The DescrptSeAttenV2.deserialize method is calling _refresh_tf2_trackable_lists() twice because TF2Module already performs that refresh for deepmd.tf2.descriptor.se_atten_v2. Remove the explicit refresh call from DescrptSeAttenV2.deserialize and keep the deserialization flow limited to delegating to DescrptSeAttenV2DP.deserialize.__func__(cls, data) and returning the object.
🤖 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 `@deepmd/tf2/descriptor/se_atten_v2.py`:
- Around line 17-21: The DescrptSeAttenV2.deserialize method is calling
_refresh_tf2_trackable_lists() twice because TF2Module already performs that
refresh for deepmd.tf2.descriptor.se_atten_v2. Remove the explicit refresh call
from DescrptSeAttenV2.deserialize and keep the deserialization flow limited to
delegating to DescrptSeAttenV2DP.deserialize.__func__(cls, data) and returning
the object.
ℹ️ Review info
⚙️ Run configuration
Configuration used: Repository UI
Review profile: CHILL
Plan: Pro
Run ID: e4a0ef32-c262-4207-8de7-5e2559fd9048
📒 Files selected for processing (6)
deepmd/dpmodel/array_api.pydeepmd/tf2/common.pydeepmd/tf2/descriptor/se_atten_v2.pydeepmd/tf2/train/trainer.pysource/tests/consistent/descriptor/test_dpa4.pysource/tests/consistent/fitting/test_dpa4_ener.py
🚧 Files skipped from review as they are similar to previous changes (3)
- deepmd/tf2/train/trainer.py
- source/tests/consistent/descriptor/test_dpa4.py
- deepmd/tf2/common.py
|
Pushed follow-ups |
Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh
|
Pushed follow-up commit 4328329 for the remaining review feedback. Changes:
Validation:
Coding agent: Codex |
OutisLi
left a comment
There was a problem hiding this comment.
Requesting changes for the dynamic-shape DPA4 force-loss correctness issue described inline.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved graph-safety correctness finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved fitting-trainability correctness finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved frozen-descriptor tracking correctness finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved normalized-exclusion routing correctness finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved default-random-gamma training semantics finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved DPA4 property-fitting factory compatibility finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved DPA4 parameter-promotion resident-memory finding.
OutisLi
left a comment
There was a problem hiding this comment.
Adding the approved PT DPA4 .pt to TF2 conversion integration finding.
Make DPA4 dynamic-shape training graph-safe, preserve frozen state, and support the schema-approved factory and PT conversion paths. Coding-Agent: Codex Codex-Version: codex-cli 0.144.1 Model: gpt-5.6-sol Reasoning-Effort: xhigh
There was a problem hiding this comment.
Actionable comments posted: 1
🤖 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.
Inline comments:
In `@deepmd/dpmodel/descriptor/dpa4_nn/so2.py`:
- Around line 747-762: Update the shape validation around the visible
x_local/radial_feat checks and _project_radial() to use runtime shape values
rather than assuming concrete symbolic rank metadata. Explicitly reject inputs
with rank below 3 before indexing dimensions, compare dimensions through runtime
shape handling, and reshape radial_feat using its runtime batch dimension or -1.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: Repository UI
Review profile: CHILL
Plan: Pro
Run ID: 741870f5-a817-453c-80e8-1695eec6fccb
📒 Files selected for processing (14)
deepmd/dpmodel/descriptor/dpa4.pydeepmd/dpmodel/descriptor/dpa4_nn/so2.pydeepmd/dpmodel/fitting/dpa4_ener.pydeepmd/tf2/common.pydeepmd/tf2/descriptor/dpa4.pydeepmd/tf2/descriptor/se_atten_v2.pydeepmd/tf2/fitting/dpa4_ener.pydeepmd/tf2/model/base_model.pydeepmd/tf2/model/model.pydeepmd/tf2/train/trainer.pysource/tests/tf2/test_dpa4.pysource/tests/tf2/test_dpa4_conversion.pysource/tests/tf2/test_model_factory.pysource/tests/tf2/test_training.py
💤 Files with no reviewable changes (1)
- deepmd/tf2/fitting/dpa4_ener.py
🚧 Files skipped from review as they are similar to previous changes (1)
- deepmd/tf2/common.py
Address the outstanding requested-change review comments. Coding-Agent: Codex Codex-Version: codex-cli 0.144.6 Model: gpt-5.6-sol Reasoning-Effort: xhigh
Resolve the DPA4 graph-interface conflict on the current master API and port the empty-edge regression to call_graph. Coding-Agent: Codex Codex-Version: codex-cli 0.144.6 Model: gpt-5.6-sol Reasoning-Effort: xhigh
OutisLi
left a comment
There was a problem hiding this comment.
TF2 DPA4 must support random_gamma=True; rejecting the standard DPA4 configuration is blocking.
Coding-Agent: Codex Codex-Version: codex-cli 0.144.6 Model: gpt-5.6-sol Reasoning-Effort: xhigh
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 10 out of 10 changed files in this pull request and generated no new comments.
Suppressed comments (3)
deepmd/tf2/descriptor/dpa4.py:317
DescrptDPA4is introduced as a TF2 wrapper forDescrptDPA4DP, but this module does not register aregister_dpmodel_mapping(DescrptDPA4DP, ...)converter. That meanstry_convert_module()cannot convert a dpmodel DPA4 descriptor instance into its TF2 wrapper when encountered as a nested object during TF2 value conversion. Add an explicit mapping (consistent with other components in this file) so dpmodel → TF2 conversion works reliably.
@BaseDescriptor.register("SeZM")
@BaseDescriptor.register("sezm")
@BaseDescriptor.register("DPA4")
@BaseDescriptor.register("dpa4")
@tf2_module
class DescrptDPA4(DescrptDPA4DP):
def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
self._tf2_training_mode = False
_promote_trainable_tree(self)
@classmethod
def deserialize(cls, data: dict) -> "DescrptDPA4":
obj = super().deserialize(data)
return _promote_trainable_tree(obj)
deepmd/tf2/descriptor/se_atten_v2.py:19
- This
deserialize()override bypasses thetf2_module-added post-deserialization refresh logic (_refresh_tf2_trackable_lists) that was introduced to ensure lists containing convertedtf.Modules are properly re-wrapped as trackable containers. Ifse_atten_v2contains module lists, this can lead to missing checkpoint tracking after deserialize. Recommended fix: after constructing the object viaDescrptSeAttenV2DP.deserialize.__func__, explicitly callobj._refresh_tf2_trackable_lists()when available (mirroring the wrapper behavior indeepmd/tf2/common.py).
@BaseDescriptor.register("se_atten_v2")
class DescrptSeAttenV2(DescrptDPA1, DescrptSeAttenV2DP):
@classmethod
def deserialize(cls, data: dict) -> "DescrptSeAttenV2":
return DescrptSeAttenV2DP.deserialize.__func__(cls, data)
deepmd/dpmodel/array_api.py:324
- In the TF branch of
xp_maximum_at,tf.shape(x_tensor, out_type=tf.int64)[0]is recomputed multiple times and feeds several segment ops. Since this is likely on a hot path (descriptor graph execution), cachenum_segments = tf.shape(x_tensor, out_type=tf.int64)[0]once and reuse it forunsorted_segment_max/min/sum. Optionally, computesegment_counts/touchedfirst and gate theall_negative_infinitycorrection to touched segments to reduce extra work for largexwith sparse updates.
x_tensor = x.unwrap()
indices_tensor = tf.reshape(tf.cast(indices.unwrap(), tf.int64), (-1,))
values_tensor = values.unwrap()
reduced = tf.math.unsorted_segment_max(
values_tensor,
indices_tensor,
tf.shape(x_tensor, out_type=tf.int64)[0],
)
if values_tensor.dtype.is_floating:
# TensorFlow uses the lowest finite value as the identity of
# unsorted_segment_max. Restore the true maximum-at identity when
# every update for a touched segment element is negative infinity.
all_negative_infinity = (
tf.math.unsorted_segment_min(
tf.cast(
tf.math.is_inf(values_tensor) & (values_tensor < 0),
tf.int32,
),
indices_tensor,
tf.shape(x_tensor, out_type=tf.int64)[0],
)
> 0
)
reduced = tf.where(
all_negative_infinity,
tf.cast(float("-inf"), values_tensor.dtype),
reduced,
)
segment_counts = tf.math.unsorted_segment_sum(
tf.ones_like(indices_tensor, dtype=tf.int32),
indices_tensor,
tf.shape(x_tensor, out_type=tf.int64)[0],
)
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 10 out of 10 changed files in this pull request and generated no new comments.
Suppressed comments (3)
deepmd/tf2/descriptor/dpa4.py:318
DescrptDPA4.deserialize()calls_promote_trainable_tree()aftersuper().deserialize(), butsuper().deserialize()already instantiatescls(**config), which runsDescrptDPA4.__init__()and promotes the tree once. The second promotion re-createstf.Variableobjects unnecessarily (and may change trackable identity).
@classmethod
def deserialize(cls, data: dict) -> "DescrptDPA4":
obj = super().deserialize(data)
return _promote_trainable_tree(obj)
deepmd/tf2/descriptor/dpa4.py:257
SO2Linear.deserialize()promotesweight_ma second time even thoughsuper().deserialize()constructs the object viacls(**config), which already runs__init__()and promotesweight_m. Re-promoting here re-creates a fresh set oftf.Variableobjects (extra allocations and can disrupt trackable identity).
This issue also appears on line 314 of the same file.
@classmethod
def deserialize(cls, data: dict) -> "SO2Linear":
obj = super().deserialize(data)
_promote_parameter_lists(obj, ("weight_m",), trainable=bool(obj.trainable))
return obj
deepmd/tf2/train/trainer.py:1332
- PR description says trainer changes are excluded ("descriptor-only"), but this PR modifies the TF2 trainer to pass a
trainingflag and mutate descriptor state via_set_model_training_mode(). Either update the PR description/scope or split the trainer change into a separate PR so the stated scope matches the actual changes.
def _call_model(
self,
task_key: str,
input_dict: dict[str, Any],
*,
label_dict: dict[str, Any] | None = None,
do_virial: bool = True,
training: bool,
) -> dict[str, Any]:
model = self.models[task_key]
self._set_model_training_mode(model, training)
_trainer._call_model now requires the keyword-only training argument; the atomic-virial-disabled regression was written before that signature change and failed with TypeError. Pass training=False (eval path). Coding-Agent: opencode opencode-Version: 1.18.9 Model: ustc/deepseek-v4-flash Reasoning-Effort: max
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 11 out of 11 changed files in this pull request and generated no new comments.
Suppressed comments (1)
deepmd/tf2/descriptor/se_atten_v2.py:19
- The new DescrptSeAttenV2.deserialize override bypasses the tf2_module deserialize wrapper’s post-deserialization
_refresh_tf2_trackable_lists()call (see deepmd/tf2/common.py:432-446). That refresh is what rebuilds list containers so nested tf.Module items are tracked correctly after deserialization/conversion. Add the refresh step here as well to keep deserialization behavior consistent with other TF2 modules.
@classmethod
def deserialize(cls, data: dict) -> "DescrptSeAttenV2":
return DescrptSeAttenV2DP.deserialize.__func__(cls, data)
6330a2f
Summary
Validation
Coding agent: Codex
Codex version: codex-cli 0.144.4
Model: gpt-5.6-sol
Reasoning effort: xhigh