@@ -84,9 +84,10 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
8484 data_keys = g_reader .readers [0 ].sample_all ().keys ()
8585 # L_inv_in=1 if "L_inv" in data_keys else 0
8686 # print("if L_inv in sample:",L_inv_in)
87- evaluator .print_head ("trn_loss" ,data_keys )
87+ align_len = 20
88+ evaluator .print_head ("trn_loss" ,data_keys ,align_len )
8889 if display_detail_test :
89- test_eval .print_head ("tst_loss" ,data_keys )
90+ test_eval .print_head ("tst_loss" ,data_keys , align_len )
9091 # print("")
9192
9293 tic = time ()
@@ -109,24 +110,24 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
109110 tst_time = time () - tic
110111 if display_natom_loss :
111112 for natom in trn_natom_loss_list .natoms ():
112- evaluator .print_head (str (natom )+ "_trn" ,data_keys )
113+ evaluator .print_head (str (natom )+ "_trn" ,data_keys , align_len )
113114 for natom in tst_natom_loss_list .natoms ():
114115 if display_detail_test :
115- test_eval .print_head (str (natom )+ "_tst" ,data_keys )
116+ test_eval .print_head (str (natom )+ "_tst" ,data_keys , align_len )
116117 else :
117- test_eval .print_head (str (natom )+ "_tst" ,[])#just energy
118+ test_eval .print_head (str (natom )+ "_tst" ,[], align_len )#just energy
118119 print ("" )
119120
120121 print (f" { 0 :<8d} { np .sqrt (np .abs (trn_loss [- 1 ])):>.2e} { np .sqrt (np .abs (tst_loss [- 1 ])):>.2e} "
121122 f" { start_lr :>.2e} { 0 :>8.2f} { tst_time :>8.2f} " ,end = '' )
122123 for loss_term in trn_loss [:- 1 ]:
123- print (f"{ loss_term :>18 .4e} " ,end = '' )
124+ print (f"{ loss_term :>{ align_len } .4e} " ,end = '' )
124125 if display_detail_test :
125126 for loss_term in tst_loss [:- 1 ]:
126- print (f"{ loss_term :>18 .4e} " ,end = '' )
127+ print (f"{ loss_term :>{ align_len } .4e} " ,end = '' )
127128 if display_natom_loss :
128- trn_natom_loss_list .print_avg_atom_loss ()
129- tst_natom_loss_list .print_avg_atom_loss ()
129+ trn_natom_loss_list .print_avg_atom_loss (align_len )
130+ tst_natom_loss_list .print_avg_atom_loss (align_len )
130131 print ('' )
131132
132133 for epoch in range (1 , n_epoch + 1 ):
@@ -162,13 +163,13 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
162163 print (f" { epoch :<8d} { np .sqrt (np .abs (trn_loss [- 1 ])):>.2e} { np .sqrt (np .abs (tst_loss [- 1 ])):>.2e} "
163164 f" { scheduler .get_last_lr ()[0 ]:>.2e} { trn_time :>8.2f} { tst_time :8.2f} " ,end = '' )
164165 for loss_term in trn_loss [:- 1 ]:
165- print (f"{ loss_term :>18 .4e} " ,end = '' )
166+ print (f"{ loss_term :>{ align_len } .4e} " ,end = '' )
166167 if display_detail_test and epoch % (display_detail_test * display_epoch ) == 0 :
167168 for loss_term in tst_loss [:- 1 ]:
168- print (f"{ loss_term :>18 .4e} " ,end = '' )
169+ print (f"{ loss_term :>{ align_len } .4e} " ,end = '' )
169170 if display_natom_loss :
170- trn_natom_loss_list .print_avg_atom_loss ()
171- tst_natom_loss_list .print_avg_atom_loss ()
171+ trn_natom_loss_list .print_avg_atom_loss (align_len )
172+ tst_natom_loss_list .print_avg_atom_loss (align_len )
172173 print ('' )
173174 if ckpt_file :
174175 model .save (ckpt_file )
0 commit comments