@@ -18,10 +18,12 @@ def __init__(self,
1818 phi_factor = 0. , phi_occ = 0 ,
1919 band_factor = 0. ,band_occ = 0 ,
2020 density_m_factor = 0. ,density_m_occ = 0 ,
21+ phi_align_factor = 0. , phi_align_occ = 0 ,
2122 density_factor = 0. , grad_penalty = 0. ,
2223 energy_lossfn = None , force_lossfn = None ,
2324 stress_lossfn = None , orbital_lossfn = None ,
2425 v_delta_lossfn = None , phi_lossfn = None ,
26+ phi_align_lossfn = None ,
2527 band_lossfn = None , density_m_lossfn = None ,
2628 energy_per_atom = 0 ,vd_divide_by_nlocal = False ):
2729 # energy term
@@ -83,7 +85,15 @@ def __init__(self,
8385 density_m_lossfn = make_loss (** density_m_lossfn )
8486 self .density_m_factor = density_m_factor
8587 self .density_m_lossfn = density_m_lossfn
86- self .get_density_m_occ = get_occ_func (density_m_occ )
88+ self .get_density_m_occ = get_occ_func (density_m_occ )
89+ # phi alignment term
90+ if phi_align_lossfn is None :
91+ phi_align_lossfn = {}
92+ if isinstance (phi_align_lossfn , dict ):
93+ phi_align_lossfn = make_loss (** phi_align_lossfn )
94+ self .phi_align_factor = phi_align_factor
95+ self .phi_align_lossfn = phi_align_lossfn
96+ self .get_phi_align_occ = get_occ_func (phi_align_occ )
8797 # coulomb term of dm; requires head gradient
8898 self .d_factor = density_factor
8999 # gradient penalty, not very useful
@@ -156,9 +166,10 @@ def __call__(self, model, sample):
156166 # print(o_label.shape, op.shape, o_pred.shape, gev.shape)
157167 tot_loss = tot_loss + self .o_factor * self .o_lossfn (o_pred , o_label )
158168 loss .append (self .o_factor * self .o_lossfn (o_pred , o_label ))
159- # optional v_delta/phi/band_energy/density_matrix calculation
169+ # optional v_delta/phi/band_energy/density_matrix/phi_alignment calculation
160170 if (self .vd_factor > 0 and "lb_vd" in sample ) or (self .phi_factor > 0 and "lb_phi" in sample ) \
161- or (self .band_factor > 0 and "lb_band" in sample ) or (self .density_m_factor > 0 and "lb_phi" in sample ):
171+ or (self .band_factor > 0 and "lb_band" in sample ) or (self .density_m_factor > 0 and "lb_phi" in sample ) \
172+ or (self .phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample ):
162173 # cal v_delta
163174 if "vdp" in sample :
164175 vdp = sample ["vdp" ] # can be complex
@@ -212,6 +223,20 @@ def __call__(self, model, sample):
212223 density_m_loss = self .density_m_factor * self .density_m_lossfn (density_m_pred , density_m_label ) * nlocal
213224 tot_loss = tot_loss + density_m_loss
214225 loss .append (density_m_loss )
226+
227+ # optional phi alignment calculation, don't need eigh on vd_pred
228+ if self .phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample :
229+ phi_label = sample ["lb_phi" ]
230+ band_label = sample ["lb_band" ]
231+ occ = self .get_phi_align_occ (natom )
232+ occ_phi_label = phi_label [..., :occ ].clone ()
233+ occ_band_label = band_label [..., :occ ].clone ()
234+ # phi_align_band should close to diagnoal matrix of occ_band_label
235+ phi_align_band = occ_phi_label .mT @ vd_pred @ occ_phi_label
236+ true_diag_band = torch .diag_embed (occ_band_label )
237+ phi_align_loss = self .phi_align_factor * self .phi_align_lossfn (phi_align_band , true_diag_band )
238+ tot_loss = tot_loss + phi_align_loss
239+ loss .append (phi_align_loss )
215240 # density loss with fix head grad
216241 if self .d_factor > 0 and "gldv" in sample :
217242 gldv = sample ["gldv" ]
@@ -245,7 +270,10 @@ def print_head(self,name,data_keys,align_len=20):
245270 info += f"{ name } _band" .rjust (align_len )
246271 # optional density matrix calculation
247272 if self .density_m_factor > 0 and "lb_phi" in data_keys :
248- info += f"{ name } _dm" .rjust (align_len )
273+ info += f"{ name } _dm" .rjust (align_len )
274+ # optional phi alignment calculation
275+ if self .phi_align_factor > 0 and "lb_phi" in data_keys and "lb_band" in data_keys :
276+ info += f"{ name } _phi_align" .rjust (align_len )
249277 # density loss with fix head grad
250278 if self .d_factor > 0 and "gldv" in data_keys :
251279 info += f"{ name } _density" .rjust (align_len )
0 commit comments