1111from deepmd .dpmodel .loss .loss import (
1212 Loss ,
1313)
14+ from deepmd .dpmodel .loss .reduction import (
15+ masked_atom_mean ,
16+ per_frame_component_mean ,
17+ )
1418from deepmd .utils .data import (
1519 DataRequirementItem ,
1620)
@@ -296,7 +300,7 @@ def call(
296300 if maskf is not None :
297301 # Idiom 2 (extensive): per-frame normalization by real-atom count.
298302 se = xp .square (energy - energy_hat ) # [nf, k]
299- per_frame = xp . mean ( xp . reshape ( se , ( _nf , - 1 )), axis = - 1 ) # [nf]
303+ per_frame = per_frame_component_mean ( se ) # [nf]
300304 if not self .use_huber :
301305 loss += pref_e * xp .mean (per_frame * inv ** norm_exp )
302306 else :
@@ -327,9 +331,7 @@ def call(
327331 l1_ener_loss = xp .mean (xp .abs (energy - energy_hat ))
328332 if maskf is not None :
329333 abs_e = xp .abs (energy - energy_hat ) # [nf, k]
330- per_frame_ae = xp .mean (
331- xp .reshape (abs_e , (_nf , - 1 )), axis = - 1
332- ) # [nf]
334+ per_frame_ae = per_frame_component_mean (abs_e ) # [nf]
333335 l1_ener_masked = xp .mean (per_frame_ae * inv )
334336 loss += pref_e * l1_ener_masked
335337 more_loss ["mae_e" ] = self .display_if_exist (
@@ -346,8 +348,7 @@ def call(
346348 )
347349 if mae :
348350 if maskf is not None :
349- abs_e = xp .abs (energy - energy_hat )
350- per_frame_ae = xp .mean (xp .reshape (abs_e , (_nf , - 1 )), axis = - 1 )
351+ per_frame_ae = per_frame_component_mean (xp .abs (energy - energy_hat ))
351352 mae_e = xp .mean (per_frame_ae * inv )
352353 else :
353354 mae_e = xp .mean (xp .abs (energy - energy_hat )) * atom_norm_ener
@@ -362,10 +363,7 @@ def call(
362363 diff_f_3d = xp .reshape (diff_f , (_nf , _nloc , 3 )) # [nf, nloc, 3]
363364 maskf_col = xp .reshape (maskf , (_nf , _nloc , 1 )) # [nf, nloc, 1]
364365 # Masked MSE computed for rmse_f display regardless of use_huber.
365- sq_f = xp .square (diff_f_3d ) * maskf_col # [nf, nloc, 3]
366- _pfs = xp .sum (xp .reshape (sq_f , (_nf , - 1 )), axis = - 1 ) # [nf]
367- _pfd = xp .sum (maskf , axis = - 1 ) * 3 # [nf]
368- l2_force_masked = xp .mean (_pfs / _pfd )
366+ l2_force_masked = masked_atom_mean (xp .square (diff_f_3d ), maskf , 3 )
369367 if not self .use_huber :
370368 loss += pref_f * l2_force_masked
371369 else :
@@ -435,12 +433,8 @@ def call(
435433 elif self .loss_func == "mae" :
436434 if maskf is not None :
437435 diff_f_3d = xp .reshape (diff_f , (_nf , _nloc , 3 ))
438- maskf_col = xp .reshape (maskf , (_nf , _nloc , 1 ))
439436 if not self .f_use_norm :
440- abs_f = xp .abs (diff_f_3d ) * maskf_col # [nf, nloc, 3]
441- per_frame_sum = xp .sum (xp .reshape (abs_f , (_nf , - 1 )), axis = - 1 )
442- per_frame_dof = xp .sum (maskf , axis = - 1 ) * 3
443- l1_force_masked = xp .mean (per_frame_sum / per_frame_dof )
437+ l1_force_masked = masked_atom_mean (xp .abs (diff_f_3d ), maskf , 3 )
444438 else :
445439 diff_3 = xp .reshape (force_hat - force , (_nf , _nloc , 3 ))
446440 norm_2d = xp .reshape (
@@ -474,11 +468,7 @@ def call(
474468 if mae :
475469 if maskf is not None :
476470 diff_f_3d = xp .reshape (diff_f , (_nf , _nloc , 3 ))
477- maskf_col = xp .reshape (maskf , (_nf , _nloc , 1 ))
478- abs_f = xp .abs (diff_f_3d ) * maskf_col
479- per_frame_sum = xp .sum (xp .reshape (abs_f , (_nf , - 1 )), axis = - 1 )
480- per_frame_dof = xp .sum (maskf , axis = - 1 ) * 3
481- mae_f = xp .mean (per_frame_sum / per_frame_dof )
471+ mae_f = masked_atom_mean (xp .abs (diff_f_3d ), maskf , 3 )
482472 else :
483473 mae_f = xp .mean (xp .abs (diff_f ))
484474 more_loss ["mae_f" ] = self .display_if_exist (mae_f , find_force )
@@ -494,7 +484,7 @@ def call(
494484 v2d = xp .reshape (virial , (_nf , 9 ))
495485 v_hat_2d = xp .reshape (virial_hat , (_nf , 9 ))
496486 se_v = xp .square (v_hat_2d - v2d ) # [nf, 9]
497- per_frame_v = xp . mean (se_v , axis = - 1 ) # [nf]
487+ per_frame_v = per_frame_component_mean (se_v ) # [nf]
498488 if not self .use_huber :
499489 loss += pref_v * xp .mean (per_frame_v * inv ** norm_exp )
500490 else :
@@ -526,8 +516,9 @@ def call(
526516 if maskf is not None :
527517 v2d = xp .reshape (virial , (_nf , 9 ))
528518 v_hat_2d = xp .reshape (virial_hat , (_nf , 9 ))
529- abs_v = xp .abs (v_hat_2d - v2d ) # [nf, 9]
530- per_frame_v = xp .mean (abs_v , axis = - 1 ) # [nf]
519+ per_frame_v = per_frame_component_mean (
520+ xp .abs (v_hat_2d - v2d )
521+ ) # [nf]
531522 l1_virial_masked = xp .mean (per_frame_v * inv )
532523 loss += pref_v * l1_virial_masked
533524 more_loss ["mae_v" ] = self .display_if_exist (
@@ -546,8 +537,7 @@ def call(
546537 if maskf is not None :
547538 v2d = xp .reshape (virial , (_nf , 9 ))
548539 v_hat_2d = xp .reshape (virial_hat , (_nf , 9 ))
549- abs_v = xp .abs (v_hat_2d - v2d )
550- per_frame_v = xp .mean (abs_v , axis = - 1 )
540+ per_frame_v = per_frame_component_mean (xp .abs (v_hat_2d - v2d ))
551541 mae_v = xp .mean (per_frame_v * inv )
552542 else :
553543 mae_v = (
@@ -565,10 +555,10 @@ def call(
565555 # Idiom 1 (per-atom masked mean, ncomp=1).
566556 ae_2d = xp .reshape (atom_ener , (_nf , _nloc ))
567557 ae_hat_2d = xp .reshape (atom_ener_hat , (_nf , _nloc ))
568- sq_ae = xp .square (ae_hat_2d - ae_2d ) * maskf # [nf, nloc]
569- per_frame_sum = xp .sum (sq_ae , axis = - 1 ) # [nf]
570558 per_frame_dof = xp .sum (maskf , axis = - 1 ) # [nf]
571- l2_ae_masked = xp .mean (per_frame_sum / per_frame_dof )
559+ l2_ae_masked = masked_atom_mean (
560+ xp .square (ae_hat_2d - ae_2d )[:, :, None ], maskf , 1
561+ )
572562 if not self .use_huber :
573563 loss += pref_ae * l2_ae_masked
574564 else :
@@ -609,10 +599,9 @@ def call(
609599 if maskf is not None :
610600 ae_2d = xp .reshape (atom_ener , (_nf , _nloc ))
611601 ae_hat_2d = xp .reshape (atom_ener_hat , (_nf , _nloc ))
612- abs_ae = xp .abs (ae_hat_2d - ae_2d ) * maskf # [nf, nloc]
613- per_frame_sum = xp .sum (abs_ae , axis = - 1 ) # [nf]
614- per_frame_dof = xp .sum (maskf , axis = - 1 ) # [nf]
615- l1_ae_masked = xp .mean (per_frame_sum / per_frame_dof )
602+ l1_ae_masked = masked_atom_mean (
603+ xp .abs (ae_hat_2d - ae_2d )[:, :, None ], maskf , 1
604+ )
616605 loss += pref_ae * l1_ae_masked
617606 more_loss ["mae_ae" ] = self .display_if_exist (
618607 l1_ae_masked , find_atom_ener
@@ -637,13 +626,9 @@ def call(
637626 # Idiom 1 with pref weight (ncomp=3).
638627 diff_f_3d = xp .reshape (diff_f , (_nf , _nloc , 3 ))
639628 pf_3d = xp .reshape (atom_pref , (_nf , _nloc , 3 ))
640- maskf_col = xp .reshape (maskf , (_nf , _nloc , 1 ))
641- sq_pf = xp .square (diff_f_3d ) * pf_3d * maskf_col # [nf, nloc, 3]
642- per_frame_sum = xp .sum (
643- xp .reshape (sq_pf , (_nf , - 1 )), axis = - 1
644- ) # [nf]
645- per_frame_dof = xp .sum (maskf , axis = - 1 ) * 3 # [nf]
646- l2_pf_masked = xp .mean (per_frame_sum / per_frame_dof )
629+ l2_pf_masked = masked_atom_mean (
630+ xp .square (diff_f_3d ) * pf_3d , maskf , 3
631+ )
647632 loss += pref_pf * l2_pf_masked
648633 more_loss ["rmse_pf" ] = self .display_if_exist (
649634 xp .sqrt (l2_pf_masked ), find_atom_pref
@@ -660,11 +645,7 @@ def call(
660645 if maskf is not None :
661646 diff_f_3d = xp .reshape (diff_f , (_nf , _nloc , 3 ))
662647 pf_3d = xp .reshape (atom_pref , (_nf , _nloc , 3 ))
663- maskf_col = xp .reshape (maskf , (_nf , _nloc , 1 ))
664- abs_pf = xp .abs (diff_f_3d ) * pf_3d * maskf_col # [nf, nloc, 3]
665- per_frame_sum = xp .sum (xp .reshape (abs_pf , (_nf , - 1 )), axis = - 1 )
666- per_frame_dof = xp .sum (maskf , axis = - 1 ) * 3
667- l1_pf_masked = xp .mean (per_frame_sum / per_frame_dof )
648+ l1_pf_masked = masked_atom_mean (xp .abs (diff_f_3d ) * pf_3d , maskf , 3 )
668649 loss += pref_pf * l1_pf_masked
669650 more_loss ["mae_pf" ] = self .display_if_exist (
670651 l1_pf_masked , find_atom_pref
0 commit comments