Skip to content

Commit 749e9e5

Browse files
authored
feat(lmdb): support mixed-size batches and lazy label availability (#5962)
## Summary - support `mix:N` LMDB batches containing frames with different atom counts, using padded rectangular batches for dense models and a flat ragged node axis for eligible graph models - compact phantom atoms before graph-model evaluation and make loss reductions, validation weighting, and epoch sizing use real atom counts - resolve label availability lazily after data requirements are registered, so required, optional, defaulted, and partially available fields are handled without an eager full-dataset scan - keep non-mixing and native-spin models on their existing rectangular public paths; the ragged regression coverage uses upstream DPA1 and avoids model-specific dependencies ## Behavioral changes - Masked per-atom loss terms now pool all included labels across the batch instead of averaging per-frame means. This intentionally retires the bit-identical reduction guarantee from #5738/#5783 for existing `mixed_type` NPY datasets whose frames have different real atom counts: those frames are weighted by their real label counts rather than equally. Uniform-atom-count batches are unchanged. Hessian pair terms remain normalized per frame so their quadratic component count does not make large structures dominate a batch. - Legacy LMDB files have no exact per-frame label-availability metadata. To avoid an eager O(N) startup scan, the reader uses a bounded probe and conservatively reduces per-frame `find_*` flags at collation. A missed rare signature may discard valid supervision for the affected batch, but default-filled values are never treated as real labels. Recording exact availability metadata when generating LMDB datasets is tracked in #5954. ## Testing - all pre-commit hooks passed for the changed files - 375 passed, 2 skipped, 1 deselected, and 13 subtests passed in the main targeted LMDB/PT/PT-expt/model suite - 15 passed in the isolated loss-reduction and decoder-pool regression suite - review fixes: 100 passed, 2 skipped, and 2 subtests passed in the common loss suite; 60 passed in the PT padding-loss suite; all 25 LMDB training tests passed; padded/unpadded DPA2 graph and Hessian parity tests passed <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **New Features** - Added LMDB batching for frames with different atom counts, including mixed-size and ragged layouts. - Added ragged-batch inference and training for supported energy and spin models. - Added configurable data-source policies for optional labels and parameters. - Added safer handling of padded atoms across neighbor graphs and model outputs. - **Bug Fixes** - Improved per-atom loss normalization for uneven and padded batches. - Prevented padded atoms from affecting neighbor searches, metrics, or losses. - Improved handling of missing labels and default-valued data. - **Documentation** - Documented mixed-size batching, ragged data, and per-atom normalization. <!-- end of auto-generated comment: release notes by coderabbit.ai --> Closes #5965
1 parent adbd6bc commit 749e9e5

63 files changed

Lines changed: 7386 additions & 1633 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

deepmd/dpmodel/loss/dos.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -159,7 +159,7 @@ def call(
159159
)
160160
diff3d = local_pred - local_label # [nf, natoms, numb_dos]
161161
if "mask" in model_dict:
162-
# idiom 1: per-frame masked mean, then average over frames
162+
# Idiom 1 (per-atom masked mean, ncomp=numb_dos).
163163
maskf = xp.astype(model_dict["mask"], diff3d.dtype) # [nf, natoms]
164164
l2_local_loss_dos = masked_atom_mean(
165165
xp.square(diff3d), maskf, self.numb_dos
@@ -184,7 +184,7 @@ def call(
184184
)
185185
diff3d = local_pred_cdf - local_label_cdf # [nf, natoms, numb_dos]
186186
if "mask" in model_dict:
187-
# idiom 1: per-frame masked mean, then average over frames
187+
# Idiom 1 (per-atom masked mean, ncomp=numb_dos).
188188
maskf = xp.astype(model_dict["mask"], diff3d.dtype) # [nf, natoms]
189189
l2_local_loss_cdf = masked_atom_mean(
190190
xp.square(diff3d), maskf, self.numb_dos

deepmd/dpmodel/loss/ener.py

Lines changed: 233 additions & 147 deletions
Large diffs are not rendered by default.

deepmd/dpmodel/loss/ener_spin.py

Lines changed: 99 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,12 @@
1515
masked_atom_mean,
1616
per_frame_component_mean,
1717
)
18+
from deepmd.dpmodel.utils.neighbor_graph.graph import (
19+
frame_id_from_n_node,
20+
)
21+
from deepmd.dpmodel.utils.neighbor_graph.segment import (
22+
segment_sum,
23+
)
1824
from deepmd.utils.data import (
1925
DataRequirementItem,
2026
)
@@ -135,6 +141,11 @@ def __init__(
135141
self.has_v = self.start_pref_v != 0.0 or self.limit_pref_v != 0.0
136142
self.has_ae = self.start_pref_ae != 0.0 or self.limit_pref_ae != 0.0
137143

144+
@property
145+
def supports_ragged_batches(self) -> bool:
146+
"""Whether the configured terms accept a flat per-node batch axis."""
147+
return True
148+
138149
def call(
139150
self,
140151
learning_rate: float,
@@ -162,33 +173,77 @@ def call(
162173
# - norm_exp=1 (intensive_ener_virial=False, legacy): loss uses 1/N scaling, which varies with system size
163174
norm_exp = 2 if self.intensive_ener_virial else 1
164175

165-
# Per-frame mask: recover real-atom count per frame when mask is provided.
166-
# maskf[nf, nloc] = 1.0 for real atoms, 0.0 for ghost padding atoms.
167-
if "mask" in model_dict:
168-
maskf = xp.astype(model_dict["mask"], energy.dtype) # [nf, nloc]
169-
real_natoms = xp.sum(maskf, axis=-1) # [nf]
170-
inv = xp.reshape(1.0 / real_natoms, (-1,)) # [nf]
171-
_nf = maskf.shape[0]
172-
_nloc = maskf.shape[1]
173-
else:
174-
# inv, _nf, _nloc are only read inside ``if maskf is not None`` guards,
175-
# so leaving them unset here is safe (and avoids dead-store warnings).
176-
maskf = None
176+
# Rectangular batches describe valid nodes with ``mask``; ragged ones
177+
# concatenate the node axis and delimit frames with ``n_node``. The
178+
# same included-node counts normalize extensive frame-level terms in
179+
# either layout, while the mask drives pooled per-node reductions.
180+
maskf = (
181+
xp.astype(model_dict["mask"], energy.dtype)
182+
if "mask" in model_dict
183+
else None
184+
)
185+
is_ragged = "n_node" in model_dict
186+
frame_id = None
187+
included_n_node = None
188+
inv = None
189+
if is_ragged:
190+
n_node = model_dict["n_node"]
191+
nframes = n_node.shape[0]
192+
if maskf is None:
193+
included_n_node = xp.astype(n_node, energy.dtype)
194+
else:
195+
frame_id = frame_id_from_n_node(n_node, n_total=maskf.shape[0])
196+
included_n_node = segment_sum(
197+
xp.reshape(maskf, (-1,)), frame_id, nframes
198+
)
199+
elif maskf is not None:
200+
included_n_node = xp.reshape(xp.sum(maskf, axis=-1), (-1,))
201+
if included_n_node is not None:
202+
has_node = included_n_node > 0
203+
safe_n_node = xp.where(
204+
has_node,
205+
included_n_node,
206+
xp.ones_like(included_n_node),
207+
)
208+
inv = xp.where(
209+
has_node,
210+
1.0 / safe_n_node,
211+
xp.zeros_like(included_n_node),
212+
)
213+
214+
def reshape_atomic(value: Array, ncomp: int) -> Array:
215+
"""Return one atomic field on the active rectangular or flat axis."""
216+
if is_ragged:
217+
return xp.reshape(value, (-1, ncomp))
218+
if maskf is not None:
219+
return xp.reshape(value, (*maskf.shape, ncomp))
220+
return xp.reshape(value, (-1, natoms, ncomp))
177221

178222
if self.has_e:
179223
energy_pred = model_dict["energy"]
180224
energy_label = label_dict["energy"]
181225
find_energy = label_dict.get("find_energy", 0.0)
182226
pref_e = pref_e * find_energy
183227
if self.enable_atom_ener_coeff and "atom_energy" in model_dict:
184-
atom_ener_pred = model_dict["atom_energy"]
185-
atom_ener_coeff = label_dict["atom_ener_coeff"]
186-
atom_ener_coeff = xp.reshape(atom_ener_coeff, atom_ener_pred.shape)
187-
energy_pred = xp.sum(atom_ener_coeff * atom_ener_pred, axis=1)
228+
atom_ener_pred = reshape_atomic(model_dict["atom_energy"], 1)
229+
atom_ener_coeff = reshape_atomic(label_dict["atom_ener_coeff"], 1)
230+
weighted_atom_ener = atom_ener_coeff * atom_ener_pred
231+
if maskf is not None:
232+
weighted_atom_ener = weighted_atom_ener * xp.reshape(
233+
maskf,
234+
(*maskf.shape, 1),
235+
)
236+
if is_ragged:
237+
if frame_id is None:
238+
frame_id = frame_id_from_n_node(
239+
n_node, n_total=atom_ener_pred.shape[0]
240+
)
241+
energy_pred = segment_sum(weighted_atom_ener, frame_id, nframes)
242+
else:
243+
energy_pred = xp.sum(weighted_atom_ener, axis=1)
188244
if self.loss_func == "mse":
189245
se_e = xp.square(energy_pred - energy_label) # [nf, k]
190-
if maskf is not None:
191-
# Idiom 2 (extensive): per-frame normalization by real-atom count.
246+
if inv is not None:
192247
per_frame_e = per_frame_component_mean(se_e) # [nf]
193248
loss += pref_e * xp.mean(per_frame_e * inv**norm_exp)
194249
more_loss["rmse_e"] = self.display_if_exist(
@@ -202,8 +257,7 @@ def call(
202257
)
203258
elif self.loss_func == "mae":
204259
l1_ener_loss = xp.mean(xp.abs(energy_pred - energy_label))
205-
if maskf is not None:
206-
# Idiom 2 (extensive) with abs: per-frame normalization by real-atom count.
260+
if inv is not None:
207261
per_frame_ae = per_frame_component_mean(
208262
xp.abs(energy_pred - energy_label)
209263
) # [nf]
@@ -218,7 +272,7 @@ def call(
218272
l1_ener_loss * atom_norm, find_energy
219273
)
220274
if mae:
221-
if maskf is not None:
275+
if inv is not None:
222276
per_frame_ae = per_frame_component_mean(
223277
xp.abs(energy_pred - energy_label)
224278
)
@@ -232,26 +286,18 @@ def call(
232286
if self.has_fr:
233287
find_force = label_dict.get("find_force", 0.0)
234288
pref_fr = pref_fr * find_force
235-
# Reshape to the canonical (nf, natoms, 3) atomic shape: the raw
236-
# data-loader label is flat (nf, natoms * 3), matching the
237-
# ``xp.reshape(label_dict[...], (-1, natoms, ncomp))`` idiom used
238-
# by every other atomic-label loss (see ``dpmodel/loss/dos.py``
239-
# and ``dpmodel/loss/tensor.py``).
240-
force_pred = xp.reshape(model_dict["force"], (-1, natoms, 3))
241-
force_label = xp.reshape(label_dict["force"], (-1, natoms, 3))
289+
force_pred = reshape_atomic(model_dict["force"], 3)
290+
force_label = reshape_atomic(label_dict["force"], 3)
291+
diff_fr = force_label - force_pred
242292
if self.loss_func == "mse":
243-
diff_fr = force_label - force_pred # [nf, nloc, 3]
244293
if maskf is not None:
245-
# Idiom 1 (per-atom masked mean, ncomp=3).
246294
l2_force_real_loss = masked_atom_mean(xp.square(diff_fr), maskf, 3)
247295
loss += pref_fr * l2_force_real_loss
248296
more_loss["rmse_fr"] = self.display_if_exist(
249297
xp.sqrt(l2_force_real_loss), find_force
250298
)
251299
if mae:
252-
mae_fr = masked_atom_mean(
253-
xp.abs(force_label - force_pred), maskf, 3
254-
)
300+
mae_fr = masked_atom_mean(xp.abs(diff_fr), maskf, 3)
255301
more_loss["mae_fr"] = self.display_if_exist(mae_fr, find_force)
256302
else:
257303
l2_force_real_loss = xp.mean(xp.square(diff_fr))
@@ -263,9 +309,8 @@ def call(
263309
mae_fr = xp.mean(xp.abs(force_label - force_pred))
264310
more_loss["mae_fr"] = self.display_if_exist(mae_fr, find_force)
265311
elif self.loss_func == "mae":
266-
abs_diff_fr = xp.abs(force_label - force_pred) # [nf, nloc, 3]
312+
abs_diff_fr = xp.abs(diff_fr)
267313
if maskf is not None:
268-
# Idiom 1 (per-atom masked mean, ncomp=3) with abs.
269314
l1_force_real_masked = masked_atom_mean(abs_diff_fr, maskf, 3)
270315
loss += pref_fr * l1_force_real_masked
271316
more_loss["mae_fr"] = self.display_if_exist(
@@ -281,18 +326,17 @@ def call(
281326
if self.has_fm:
282327
find_force_mag = label_dict.get("find_force_mag", 0.0)
283328
pref_fm = pref_fm * find_force_mag
284-
# Same flat -> (nf, natoms, 3) reshape as the real-force branch above.
285-
force_mag_pred = xp.reshape(model_dict["force_mag"], (-1, natoms, 3))
286-
force_mag_label = xp.reshape(label_dict["force_mag"], (-1, natoms, 3))
287-
mask_mag = model_dict["mask_mag"]
288-
# mask_mag: [nframes, natoms, 1], bool -> use mask multiplication
329+
force_mag_pred = reshape_atomic(model_dict["force_mag"], 3)
330+
force_mag_label = reshape_atomic(label_dict["force_mag"], 3)
331+
mask_mag = reshape_atomic(model_dict["mask_mag"], 1)
332+
if maskf is not None:
333+
mask_mag = xp.logical_and(
334+
mask_mag,
335+
xp.reshape(maskf > 0, (*maskf.shape, 1)),
336+
)
289337
mask_float = xp.astype(mask_mag, force_mag_pred.dtype)
290-
# zero out non-magnetic atoms
291338
diff_fm = (force_mag_label - force_mag_pred) * mask_float
292339
n_valid = xp.sum(mask_float)
293-
# Guard the denominator itself because array backends may evaluate
294-
# both branches of ``where``. This is safe under JAX tracing and
295-
# makes an all-empty magnetic mask contribute exactly zero.
296340
safe_n_valid = xp.where(n_valid > 0, n_valid, xp.ones_like(n_valid))
297341
if self.loss_func == "mse":
298342
l2_force_mag_loss = xp.sum(xp.square(diff_fm)) / (safe_n_valid * 3)
@@ -304,11 +348,7 @@ def call(
304348
mae_fm = xp.sum(xp.abs(diff_fm)) / (safe_n_valid * 3)
305349
more_loss["mae_fm"] = self.display_if_exist(mae_fm, find_force_mag)
306350
elif self.loss_func == "mae":
307-
abs_diff_fm = xp.abs(diff_fm) # [nf, na, 3], zeros for non-magnetic
308-
# Mean over frames, magnetic atoms and xyz (same reduction as
309-
# force_mag MSE, force_real MAE and the displayed mae_fm) so the
310-
# loss is batch-size independent: a 2-frame batch equals the mean
311-
# of the two single-frame losses.
351+
abs_diff_fm = xp.abs(diff_fm)
312352
l1_force_mag_loss = xp.sum(abs_diff_fm) / (safe_n_valid * 3)
313353
loss += pref_fm * l1_force_mag_loss
314354
more_loss["mae_fm"] = self.display_if_exist(
@@ -318,43 +358,34 @@ def call(
318358
if self.has_ae:
319359
find_atom_ener = label_dict.get("find_atom_ener", 0.0)
320360
pref_ae = pref_ae * find_atom_ener
321-
atom_ener = model_dict["atom_energy"]
322-
atom_ener_label = label_dict["atom_ener"]
361+
atom_ener = reshape_atomic(model_dict["atom_energy"], 1)
362+
atom_ener_label = reshape_atomic(label_dict["atom_ener"], 1)
323363
if maskf is not None:
324-
# Idiom 1 (per-atom masked mean, ncomp=1).
325-
ae = xp.reshape(atom_ener, (_nf, _nloc, 1))
326-
ae_label = xp.reshape(atom_ener_label, (_nf, _nloc, 1))
327364
if self.loss_func == "mse":
328365
l2_atom_ener_loss = masked_atom_mean(
329-
xp.square(ae_label - ae), maskf, 1
366+
xp.square(atom_ener_label - atom_ener), maskf, 1
330367
)
331368
loss += pref_ae * l2_atom_ener_loss
332369
more_loss["rmse_ae"] = self.display_if_exist(
333370
xp.sqrt(l2_atom_ener_loss), find_atom_ener
334371
)
335372
elif self.loss_func == "mae":
336373
l1_atom_ener_loss = masked_atom_mean(
337-
xp.abs(ae_label - ae), maskf, 1
374+
xp.abs(atom_ener_label - atom_ener), maskf, 1
338375
)
339376
loss += pref_ae * l1_atom_ener_loss
340377
more_loss["mae_ae"] = self.display_if_exist(
341378
l1_atom_ener_loss, find_atom_ener
342379
)
343380
else:
344-
atom_ener_reshape = xp.reshape(atom_ener, (-1,))
345-
atom_ener_label_reshape = xp.reshape(atom_ener_label, (-1,))
346381
if self.loss_func == "mse":
347-
l2_atom_ener_loss = xp.mean(
348-
xp.square(atom_ener_label_reshape - atom_ener_reshape)
349-
)
382+
l2_atom_ener_loss = xp.mean(xp.square(atom_ener_label - atom_ener))
350383
loss += pref_ae * l2_atom_ener_loss
351384
more_loss["rmse_ae"] = self.display_if_exist(
352385
xp.sqrt(l2_atom_ener_loss), find_atom_ener
353386
)
354387
elif self.loss_func == "mae":
355-
l1_atom_ener_loss = xp.mean(
356-
xp.abs(atom_ener_label_reshape - atom_ener_reshape)
357-
)
388+
l1_atom_ener_loss = xp.mean(xp.abs(atom_ener_label - atom_ener))
358389
loss += pref_ae * l1_atom_ener_loss
359390
more_loss["mae_ae"] = self.display_if_exist(
360391
l1_atom_ener_loss, find_atom_ener
@@ -364,11 +395,10 @@ def call(
364395
find_virial = label_dict.get("find_virial", 0.0)
365396
pref_v = pref_v * find_virial
366397
virial_pred = xp.reshape(model_dict["virial"], (-1, 9))
367-
virial_label = label_dict["virial"]
398+
virial_label = xp.reshape(label_dict["virial"], (-1, 9))
368399
diff_v = virial_label - virial_pred # [nf, 9]
369400
if self.loss_func == "mse":
370-
if maskf is not None:
371-
# Idiom 2 (extensive, k=9): per-frame normalization by real-atom count.
401+
if inv is not None:
372402
per_frame_v = per_frame_component_mean(xp.square(diff_v)) # [nf]
373403
loss += pref_v * xp.mean(per_frame_v * inv**norm_exp)
374404
more_loss["rmse_v"] = self.display_if_exist(
@@ -391,8 +421,7 @@ def call(
391421
more_loss["mae_v"] = self.display_if_exist(mae_v, find_virial)
392422
elif self.loss_func == "mae":
393423
l1_virial_loss = xp.mean(xp.abs(diff_v))
394-
if maskf is not None:
395-
# Idiom 2 (extensive, k=9) with abs: per-frame normalization by real-atom count.
424+
if inv is not None:
396425
per_frame_v = per_frame_component_mean(xp.abs(diff_v)) # [nf]
397426
l1_virial_masked = xp.mean(per_frame_v * inv)
398427
loss += pref_v * l1_virial_masked
@@ -471,6 +500,7 @@ def label_requirement(self) -> list[DataRequirementItem]:
471500
must=False,
472501
high_prec=False,
473502
default=1.0,
503+
source_policy="default",
474504
)
475505
)
476506
return label_requirement

deepmd/dpmodel/loss/loss.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,11 @@ def call(
5151
def label_requirement(self) -> list[DataRequirementItem]:
5252
"""Return data label requirements needed for this loss calculation."""
5353

54+
@property
55+
def supports_ragged_batches(self) -> bool:
56+
"""Whether this objective accepts a flat per-node batch axis."""
57+
return False
58+
5459
@staticmethod
5560
def display_if_exist(loss: Array, find_property: float) -> Array:
5661
"""Display NaN if labeled property is not found.

0 commit comments

Comments
 (0)