1818def train (model , g_reader , n_epoch = 1000 , test_reader = None , * ,
1919 energy_factor = 1. , force_factor = 0. , stress_factor = 0. , orbital_factor = 0. , v_delta_factor = 0. , phi_factor = 0. ,phi_occ = 0 , band_factor = 0. , band_occ = 0 , bandgap_factor = 0. , bandgap_occ = 0 , density_m_factor = 0. , density_m_occ = 0 , phi_align_factor = 0. , phi_align_occ = 0 , density_factor = 0. ,
2020 energy_loss = None , force_loss = None , stress_loss = None , orbital_loss = None , v_delta_loss = None , phi_loss = None , band_loss = None , bandgap_loss = None , density_m_loss = None , phi_align_loss = None , grad_penalty = 0. ,
21- energy_per_atom = 0 , vd_divide_by_nlocal = False , vd_masked_loss = False , vd_masked_S_threshold = 1e-6 , vd_masked_H_threshold = 1e-6 , use_safe_eigh = False ,
21+ energy_per_atom = 0 , vd_divide_by_nlocal = False , vd_masked_loss = 0 , vd_masked_S_threshold = 1e-6 , vd_masked_H_threshold = 1e-6 , vd_masked_width = 1 , use_safe_eigh = False ,
2222 start_lr = 0.001 , decay_steps = 100 , decay_rate = 0.96 , stop_lr = None , decay_rate_iter = None ,
2323 weight_decay = 0. , fix_embedding = False ,
2424 display_epoch = 100 , display_detail_test = 0 , display_natom_loss = False , ckpt_file = "model.pth" ,
@@ -64,7 +64,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
6464 phi_align_lossfn = phi_align_loss ,
6565 density_factor = density_factor , grad_penalty = grad_penalty ,
6666 energy_per_atom = energy_per_atom , vd_divide_by_nlocal = vd_divide_by_nlocal ,
67- vd_masked_loss = vd_masked_loss , vd_masked_S_threshold = vd_masked_S_threshold , vd_masked_H_threshold = vd_masked_H_threshold ,
67+ vd_masked_loss = vd_masked_loss , vd_masked_S_threshold = vd_masked_S_threshold ,
68+ vd_masked_H_threshold = vd_masked_H_threshold , vd_masked_width = vd_masked_width ,
6869 use_safe_eigh = use_safe_eigh )
6970 if not display_detail_test :
7071 # make test evaluator that only returns l2loss of energy
@@ -88,7 +89,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
8889 density_m_lossfn = density_m_loss , phi_align_lossfn = phi_align_loss ,
8990 density_factor = to_one (density_factor ), grad_penalty = grad_penalty ,
9091 energy_per_atom = energy_per_atom , vd_divide_by_nlocal = vd_divide_by_nlocal ,
91- vd_masked_loss = vd_masked_loss , vd_masked_S_threshold = vd_masked_S_threshold , vd_masked_H_threshold = vd_masked_H_threshold ,
92+ vd_masked_loss = vd_masked_loss , vd_masked_S_threshold = vd_masked_S_threshold ,
93+ vd_masked_H_threshold = vd_masked_H_threshold ,vd_masked_width = vd_masked_width ,
9294 use_safe_eigh = use_safe_eigh )
9395
9496 print ("# epoch trn_err tst_err lr trn_time tst_time" ,end = '' )
0 commit comments