Skip to content

Commit 89d04bd

Browse files
njzjz-botnjzjz-bot
andauthored
fix(model): register magnetic cell derivatives correctly (#5836)
## Summary - register magnetic cell derivatives in the cell-derivative bucket instead of the coordinate bucket - generate the corresponding reduced magnetic virial definition - add the optional DeepEval backend name mapping so spin evaluation request construction remains valid - extend output-definition tests for bucket membership, shape, category, reduction, and atomic metadata ## Why existing tests missed this The existing comprehensive output-definition test only compared the flattened ModelOutputDef key set. The magnetic cell derivative appeared in that flattened set even while stored in the wrong internal bucket, so the test passed. It did not inspect keys_derv_c versus keys_derv_r and did not expect the reduced magnetic virial definition. The new regression asserts both internal bucket membership and the derived reduced definition. ## Validation - ruff format . - ruff check . - complete output-definition test file plus TensorFlow C++ core regression: 17 passed - verified the new reduced output has a DeepEval backend mapping; unsupported backend values keep the existing optional-output NaN/removal behavior An additional pt_expt spin-test collection attempt hit an environment Triton plugin import segmentation fault before test collection; the focused pure-metadata and core tests are unaffected. Closes #5632 Coding agent: Codex Codex version: codex-cli 0.144.4 Model: gpt-5.6-sol Reasoning effort: xhigh <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **Bug Fixes** * Corrected magnetic cell-derivative outputs so they are registered and reduced through the appropriate cell-based pathway. * Added support for magnetic cell-derivative output mapping to backend virial results. * Ensured magnetic derivative outputs have the correct reduced shape and classification. <!-- end of auto-generated comment: release notes by coderabbit.ai --> Co-authored-by: njzjz-bot <njzjz.bot@gmail.com>
1 parent 21df9d6 commit 89d04bd

3 files changed

Lines changed: 55 additions & 1 deletion

File tree

deepmd/dpmodel/output_def.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -512,7 +512,7 @@ def do_derivative(
512512
category=apply_operation(vv, OutputVariableOperation.DERV_C),
513513
)
514514
if vv.magnetic:
515-
def_derv_r[rkcm] = OutputVariableDef(
515+
def_derv_c[rkcm] = OutputVariableDef(
516516
rkcm,
517517
vv.shape + [9], # noqa: RUF005
518518
reducible=True,

deepmd/infer/deep_eval.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,7 @@ class DeepEvalBackend(ABC):
6161
"energy_derv_c": "atom_virial",
6262
"energy_derv_c_mag": "atom_virial_mag",
6363
"energy_derv_c_redu": "virial",
64+
"energy_derv_c_mag_redu": "virial_mag",
6465
"polar": "polar",
6566
"polar_redu": "global_polar",
6667
"polar_derv_r": "force",

source/tests/common/dpmodel/test_output_def.py

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,9 @@
1717
apply_operation,
1818
check_var,
1919
)
20+
from deepmd.infer.deep_eval import (
21+
DeepEvalBackend,
22+
)
2023

2124

2225
class VariableDef:
@@ -32,6 +35,34 @@ def __init__(
3235

3336

3437
class TestDef(unittest.TestCase):
38+
def test_magnetic_cell_derivative_uses_cell_bucket(self) -> None:
39+
"""Magnetic cell derivatives must participate in virial reduction."""
40+
model_def = ModelOutputDef(
41+
FittingOutputDef(
42+
[
43+
OutputVariableDef(
44+
"energy",
45+
[1],
46+
reducible=True,
47+
r_differentiable=True,
48+
c_differentiable=True,
49+
magnetic=True,
50+
)
51+
]
52+
)
53+
)
54+
55+
self.assertIn("energy_derv_c_mag", model_def.keys_derv_c())
56+
self.assertNotIn("energy_derv_c_mag", model_def.keys_derv_r())
57+
self.assertIn("energy_derv_c_mag_redu", model_def.keys_derv_c_redu())
58+
self.assertEqual(
59+
DeepEvalBackend._OUTDEF_DP2BACKEND["energy_derv_c_mag_redu"],
60+
"virial_mag",
61+
)
62+
reduced_def = model_def["energy_derv_c_mag_redu"]
63+
self.assertEqual(reduced_def.shape, [1, 9])
64+
self.assertFalse(reduced_def.atomic)
65+
3566
def test_model_output_def(self) -> None:
3667
defs = [
3768
OutputVariableDef(
@@ -161,6 +192,7 @@ def test_model_output_def(self) -> None:
161192
"energy3_derv_c_redu",
162193
"energy3_derv_r_mag",
163194
"energy3_derv_c_mag",
195+
"energy3_derv_c_mag_redu",
164196
"dos_redu",
165197
"mask",
166198
"mask_mag",
@@ -218,6 +250,7 @@ def test_model_output_def(self) -> None:
218250
self.assertEqual(md["energy3_derv_c_redu"].shape, [1, 9])
219251
self.assertEqual(md["energy3_derv_r_mag"].shape, [1, 3])
220252
self.assertEqual(md["energy3_derv_c_mag"].shape, [1, 9])
253+
self.assertEqual(md["energy3_derv_c_mag_redu"].shape, [1, 9])
221254
self.assertEqual(md["gap"].shape, [13])
222255
self.assertEqual(md["gap_redu"].shape, [13])
223256
# atomic
@@ -240,6 +273,7 @@ def test_model_output_def(self) -> None:
240273
self.assertEqual(md["energy3_derv_r_mag"].atomic, True)
241274
self.assertEqual(md["energy3_derv_c_mag"].atomic, True)
242275
self.assertEqual(md["energy3_derv_c_redu"].atomic, False)
276+
self.assertEqual(md["energy3_derv_c_mag_redu"].atomic, False)
243277
self.assertEqual(md["gap"].atomic, True)
244278
self.assertEqual(md["gap_redu"].atomic, False)
245279
# category
@@ -277,6 +311,10 @@ def test_model_output_def(self) -> None:
277311
self.assertEqual(
278312
md["energy3_derv_c_mag"].category, OutputVariableCategory.DERV_C
279313
)
314+
self.assertEqual(
315+
md["energy3_derv_c_mag_redu"].category,
316+
OutputVariableCategory.DERV_C_REDU,
317+
)
280318
self.assertEqual(md["gap"].category, OutputVariableCategory.OUT)
281319
self.assertEqual(md["gap_redu"].category, OutputVariableCategory.REDU)
282320
# flag
@@ -393,6 +431,15 @@ def test_model_output_def(self) -> None:
393431
md["energy3_derv_c_mag"].category & OVO.DERV_C,
394432
OVO.DERV_C,
395433
)
434+
self.assertEqual(
435+
md["energy3_derv_c_mag_redu"].category & OVO.REDU,
436+
OVO.REDU,
437+
)
438+
self.assertEqual(md["energy3_derv_c_mag_redu"].category & OVO.DERV_R, 0)
439+
self.assertEqual(
440+
md["energy3_derv_c_mag_redu"].category & OVO.DERV_C,
441+
OVO.DERV_C,
442+
)
396443
# apply_operation: energy
397444
self.assertEqual(
398445
apply_operation(md["energy"], OVO.REDU),
@@ -456,6 +503,10 @@ def test_model_output_def(self) -> None:
456503
apply_operation(md["energy3"], OVO.DERV_C),
457504
md["energy3_derv_c_mag"].category,
458505
)
506+
self.assertEqual(
507+
apply_operation(md["energy3_derv_c_mag"], OVO.REDU),
508+
md["energy3_derv_c_mag_redu"].category,
509+
)
459510
# raise ValueError
460511
with self.assertRaises(ValueError):
461512
apply_operation(md["energy_redu"], OVO.REDU)
@@ -481,6 +532,8 @@ def test_model_output_def(self) -> None:
481532
apply_operation(md["energy3_derv_c_redu"], OVO.REDU)
482533
with self.assertRaises(ValueError):
483534
apply_operation(md["energy3_derv_c_mag"], OVO.DERV_C)
535+
with self.assertRaises(ValueError):
536+
apply_operation(md["energy3_derv_c_mag_redu"], OVO.REDU)
484537
# hession
485538
hession_cat = apply_operation(md["energy_derv_r"], OVO.DERV_R)
486539
self.assertEqual(hession_cat & OVO.DERV_R, OVO.DERV_R)

0 commit comments

Comments
 (0)