88except ImportError as e :
99 sys .path .append (os .path .dirname (os .path .realpath (__file__ )) + "/../../" )
1010from deepks .model .reader import generalized_eigh , eigh_wrapper
11- from deepks .model .utils import get_density_matrix , cal_phi_loss , cal_v_delta , cal_vd_masked_loss , get_occ_func , make_loss
11+ from deepks .model .utils import get_density_matrix , cal_phi_loss , cal_v_delta , cal_vd_masked_loss , cal_bandgap , get_occ_func , make_loss
1212
1313class Evaluator :
1414 def __init__ (self ,
@@ -17,14 +17,15 @@ def __init__(self,
1717 v_delta_factor = 0. ,
1818 phi_factor = 0. , phi_occ = 0 ,
1919 band_factor = 0. ,band_occ = 0 ,
20+ bandgap_factor = 0. ,bandgap_occ = 0 ,
2021 density_m_factor = 0. ,density_m_occ = 0 ,
2122 phi_align_factor = 0. , phi_align_occ = 0 ,
2223 density_factor = 0. , grad_penalty = 0. ,
2324 energy_lossfn = None , force_lossfn = None ,
2425 stress_lossfn = None , orbital_lossfn = None ,
2526 v_delta_lossfn = None , phi_lossfn = None ,
2627 phi_align_lossfn = None ,
27- band_lossfn = None , density_m_lossfn = None ,
28+ band_lossfn = None , bandgap_lossfn = None , density_m_lossfn = None ,
2829 energy_per_atom = 0 ,vd_divide_by_nlocal = False ,
2930 vd_masked_loss = False ,
3031 vd_masked_S_threshold = 1e-6 , vd_masked_H_threshold = 1e-6 ,
@@ -81,6 +82,14 @@ def __init__(self,
8182 self .band_factor = band_factor
8283 self .band_lossfn = band_lossfn
8384 self .get_band_occ = get_occ_func (band_occ )
85+ # bandgap term
86+ if bandgap_lossfn is None :
87+ bandgap_lossfn = {}
88+ if isinstance (bandgap_lossfn , dict ):
89+ bandgap_lossfn = make_loss (** bandgap_lossfn )
90+ self .bandgap_factor = bandgap_factor
91+ self .bandgap_lossfn = bandgap_lossfn
92+ self .get_bandgap_occ = get_occ_func (bandgap_occ )
8493 #density matrix term
8594 if density_m_lossfn is None :
8695 density_m_lossfn = {}
@@ -136,6 +145,7 @@ def __call__(self, model, sample):
136145 or (self .vd_factor > 0 and "lb_vd" in sample )
137146 or (self .phi_factor > 0 and "lb_phi" in sample )
138147 or (self .band_factor > 0 and "lb_band" in sample )
148+ or (self .bandgap_factor > 0 and "lb_band" in sample )
139149 or (self .density_m_factor > 0 )
140150 or (self .d_factor > 0 and "gldv" in sample )
141151 or self .g_penalty > 0 )
@@ -178,7 +188,8 @@ def __call__(self, model, sample):
178188 loss .append (self .o_factor * self .o_lossfn (o_pred , o_label ))
179189 # optional v_delta/phi/band_energy/density_matrix/phi_alignment calculation
180190 if (self .vd_factor > 0 and "lb_vd" in sample ) or (self .phi_factor > 0 and "lb_phi" in sample ) \
181- or (self .band_factor > 0 and "lb_band" in sample ) or (self .density_m_factor > 0 and "lb_phi" in sample ) \
191+ or (self .band_factor > 0 and "lb_band" in sample ) or (self .bandgap_factor > 0 and "lb_band" in sample ) \
192+ or (self .density_m_factor > 0 and "lb_phi" in sample ) \
182193 or (self .phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample ):
183194 # cal v_delta
184195 if "vdp" in sample :
@@ -204,7 +215,7 @@ def __call__(self, model, sample):
204215 tot_loss = tot_loss + vd_loss
205216 loss .append (vd_loss )
206217
207- if (self .phi_factor > 0 and "lb_phi" in sample ) or (self .band_factor > 0 and "lb_band" in sample ) or (self .density_m_factor > 0 and "lb_phi" in sample ):
218+ if (self .phi_factor > 0 and "lb_phi" in sample ) or (self .band_factor > 0 and "lb_band" in sample ) or (self .bandgap_factor > 0 and "lb_band" in sample ) or ( self . density_m_factor > 0 and "lb_phi" in sample ):
208219 h_base = sample ["h_base" ]
209220 if "trans_matrix" in sample :
210221 trans_matrix = sample ["trans_matrix" ]
@@ -225,6 +236,15 @@ def __call__(self, model, sample):
225236 tot_loss = tot_loss + band_loss
226237 # print("occ_band",band_pred[...,:band_occ],band_label[...,:band_occ])
227238 loss .append (band_loss )
239+ # optional bandgap calculation
240+ if self .bandgap_factor > 0 and "lb_band" in sample :
241+ band_label = sample ["lb_band" ]
242+ bandgap_occ = self .get_bandgap_occ (natom )
243+ bandgap_label = cal_bandgap (band_label , bandgap_occ )
244+ bandgap_pred = cal_bandgap (band_pred , bandgap_occ )
245+ bandgap_loss = self .bandgap_factor * self .bandgap_lossfn (bandgap_pred , bandgap_label )
246+ tot_loss = tot_loss + bandgap_loss
247+ loss .append (bandgap_loss )
228248 # optional density matrix calculation
229249 if self .density_m_factor > 0 and "lb_phi" in sample :
230250 # calculate density_m_label every time, kind of waste of time
@@ -283,6 +303,9 @@ def print_head(self,name,data_keys,align_len=20):
283303 # optional band energy calculation
284304 if self .band_factor > 0 and "lb_band" in data_keys :
285305 info += f"{ name } _band" .rjust (align_len )
306+ # optional bandgap calculation
307+ if self .bandgap_factor > 0 and "lb_band" in data_keys :
308+ info += f"{ name } _bandgap" .rjust (align_len )
286309 # optional density matrix calculation
287310 if self .density_m_factor > 0 and "lb_phi" in data_keys :
288311 info += f"{ name } _dm" .rjust (align_len )
0 commit comments