Skip to content

Commit 4229fc3

Browse files
committed
Adjust the output way to ensure that log.train is well aligned
1 parent e5e0644 commit 4229fc3

2 files changed

Lines changed: 27 additions & 27 deletions

File tree

deepks/model/evaluator.py

Lines changed: 13 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -220,35 +220,34 @@ def __call__(self, model, sample):
220220
loss.append(tot_loss)
221221
return loss
222222

223-
def print_head(self,name,data_keys):
224-
len=20
225-
info=f"{name}_energy".rjust(len)
223+
def print_head(self,name,data_keys,align_len=20):
224+
info=f"{name}_energy".rjust(align_len)
226225
if self.g_penalty > 0 and "eg0" in data_keys:
227-
info+=f"{name}_grad".rjust(len)
226+
info+=f"{name}_grad".rjust(align_len)
228227
# optional force calculation
229228
if self.f_factor > 0 and "lb_f" in data_keys:
230-
info+=f"{name}_force".rjust(len)
229+
info+=f"{name}_force".rjust(align_len)
231230
# optional stress calculation
232231
if self.s_factor > 0 and "lb_s" in data_keys:
233-
info+=f"{name}_stress".rjust(len)
232+
info+=f"{name}_stress".rjust(align_len)
234233
# optional orbital(bandgap) calculation
235234
if self.o_factor > 0 and "lb_o" in data_keys:
236-
info+=f"{name}_bandgap".rjust(len)
235+
info+=f"{name}_bandgap".rjust(align_len)
237236
# optional v_delta calculation
238237
if self.vd_factor > 0 and "lb_vd" in data_keys:
239-
info+=f"{name}_v_delta".rjust(len)
238+
info+=f"{name}_v_delta".rjust(align_len)
240239
# optional phi calculation
241240
if self.phi_factor > 0 and "lb_phi" in data_keys:
242-
info+=f"{name}_phi".rjust(len)
241+
info+=f"{name}_phi".rjust(align_len)
243242
# optional band energy calculation
244243
if self.band_factor > 0 and "lb_band" in data_keys:
245-
info+=f"{name}_band".rjust(len)
244+
info+=f"{name}_band".rjust(align_len)
246245
# optional density matrix calculation
247246
if self.density_m_factor > 0 and "lb_phi" in data_keys:
248-
info+=f"{name}_dm".rjust(len)
247+
info+=f"{name}_dm".rjust(align_len)
249248
# density loss with fix head grad
250249
if self.d_factor > 0 and "gldv" in data_keys:
251-
info+=f"{name}_density".rjust(len)
250+
info+=f"{name}_density".rjust(align_len)
252251
print(info,end='')
253252

254253
class NatomLossList:
@@ -275,11 +274,11 @@ def avg_atom_loss(self):
275274
# avg upon data
276275
return {natom:np.mean(losses,axis=0) for (natom,losses) in self.natom_loss_list.items()}
277276

278-
def print_avg_atom_loss(self):
277+
def print_avg_atom_loss(self,align_len=20):
279278
avg_atom_loss = sorted(self.avg_atom_loss().items(), key=lambda x: x[0])
280279
for (atom,aal) in avg_atom_loss:
281280
for avg_atom_loss_term in aal[:-1]:
282-
print(f"{avg_atom_loss_term:>18.4e}",end='')
281+
print(f"{avg_atom_loss_term:>{align_len}.4e}",end='')
283282

284283
def avg_loss(self):
285284
# avg upon data and natom

deepks/model/train.py

Lines changed: 14 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)