Skip to content

Commit f8e1365

Browse files
committed
refactor(phase3): migrate model and scf implementations to core
1 parent 4ded701 commit f8e1365

37 files changed

Lines changed: 2792 additions & 2720 deletions

deepks/core/ml/eval/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
11
"""Evaluation components for DeepKS core ML layer."""
22

33
from .evaluator import * # noqa: F401,F403
4+
from .test import * # noqa: F401,F403

deepks/core/ml/eval/evaluator.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
except ImportError as e:
99
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
1010
from deepks.io.transforms.linalg import generalized_eigh
11-
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, get_occ_func, make_loss, get_gedm, cal_vdr, loss_hr
11+
from deepks.core.ml.utils import get_density_matrix, cal_phi_loss, cal_v_delta, get_occ_func, make_loss, get_gedm, cal_vdr, loss_hr
1212

1313
class Evaluator:
1414
def __init__(self,

deepks/core/ml/eval/test.py

Lines changed: 93 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,93 @@
1+
import os
2+
import numpy as np
3+
import torch
4+
import torch.nn as nn
5+
try:
6+
import deepks
7+
except ImportError as e:
8+
import sys
9+
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
10+
from deepks.default import DEVICE
11+
from deepks.core.ml.models.corrnet import CorrNet
12+
from deepks.io.readers.group_reader import GroupReader
13+
from deepks.utils import load_yaml, load_dirs, check_list
14+
15+
16+
def test(model, g_reader, dump_prefix="test", group=False):
17+
model.eval()
18+
loss_fn=nn.MSELoss()
19+
label_list = []
20+
pred_list = []
21+
22+
for i in range(g_reader.nsystems):
23+
sample = g_reader.sample_all(i)
24+
nframes = sample["lb_e"].shape[0]
25+
for k, v in sample.items():
26+
if isinstance(v, list):
27+
sample[k] = [vv.to(DEVICE, non_blocking=True) for vv in v]
28+
elif not torch.is_complex(v):
29+
sample[k] = v.to(DEVICE, non_blocking=True)
30+
else:
31+
if k == "phialpha":
32+
sample[k] = v.to("cpu", dtype=torch.complex128, non_blocking=True)
33+
else:
34+
sample[k] = v.to(DEVICE, dtype=torch.complex128, non_blocking=True)
35+
label, data = sample["lb_e"], sample["eig"]
36+
pred = model(data)
37+
error = torch.sqrt(loss_fn(pred, label))
38+
39+
error_np = error.item()
40+
label_np = label.cpu().numpy().reshape(nframes, -1).sum(axis=1)
41+
pred_np = pred.detach().cpu().numpy().reshape(nframes, -1).sum(axis=1)
42+
error_l1 = np.mean(np.abs(label_np - pred_np))
43+
label_list.append(label_np)
44+
pred_list.append(pred_np)
45+
46+
if not group and dump_prefix is not None:
47+
nd = max(len(str(g_reader.nsystems)), 2)
48+
dump_res = np.stack([label_np, pred_np], axis=1)
49+
header = f"{g_reader.path_list[i]}\nmean l1 error: {error_l1}\nmean l2 error: {error_np}\nreal_ene pred_ene"
50+
filename = f"{dump_prefix}.{i:0{nd}}.out"
51+
np.savetxt(filename, dump_res, header=header)
52+
# print(f"system {i} finished")
53+
54+
all_label = np.concatenate(label_list, axis=0)
55+
all_pred = np.concatenate(pred_list, axis=0)
56+
all_err_l1 = np.mean(np.abs(all_label - all_pred))
57+
all_err_l2 = np.sqrt(np.mean((all_label - all_pred) ** 2))
58+
info = f"all systems mean l1 error: {all_err_l1}\nall systems mean l2 error: {all_err_l2}"
59+
print(info)
60+
if dump_prefix is not None and group:
61+
np.savetxt(f"{dump_prefix}.out", np.stack([all_label, all_pred], axis=1),
62+
header=info + "\nreal_ene pred_ene")
63+
return all_err_l1, all_err_l2
64+
65+
66+
def main(data_paths, model_file="model.pth",
67+
output_prefix='test', group=False,
68+
e_name='l_e_delta', d_name=['dm_eig']):
69+
data_paths = load_dirs(data_paths)
70+
if len(d_name) == 1:
71+
d_name = d_name[0]
72+
g_reader = GroupReader(data_paths, e_name=e_name, d_name=d_name,
73+
conv_filter=False, extra_label=True)
74+
model_file = check_list(model_file)
75+
for f in model_file:
76+
print(f)
77+
p = os.path.dirname(f)
78+
model = CorrNet.load(f).double().to(DEVICE)
79+
dump = os.path.join(p, output_prefix)
80+
dir_name = os.path.dirname(dump)
81+
if dir_name:
82+
os.makedirs(dir_name, exist_ok=True)
83+
if model.elem_table is not None:
84+
elist, econst = model.elem_table
85+
g_reader.collect_elems(elist)
86+
g_reader.subtract_elem_const(econst)
87+
test(model, g_reader, dump_prefix=dump, group=group)
88+
g_reader.revert_elem_const()
89+
90+
91+
if __name__ == "__main__":
92+
from deepks.main import test_cli as cli
93+
cli()

deepks/core/ml/train/__init__.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,3 @@
1-
"""Scaffold package for refactor architecture."""
1+
"""Training components for DeepKS core ML layer."""
2+
3+
from .train import * # noqa: F401,F403

deepks/core/ml/train/train.py

Lines changed: 249 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,249 @@
1+
import os
2+
import sys
3+
import numpy as np
4+
import torch
5+
import torch.optim as optim
6+
from time import time
7+
try:
8+
import deepks
9+
except ImportError as e:
10+
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
11+
from deepks.default import DEVICE
12+
from deepks.core.ml.models.corrnet import CorrNet
13+
from deepks.io.readers.group_reader import GroupReader
14+
from deepks.utils import load_dirs, load_elem_table
15+
from deepks.core.ml.utils import preprocess, fit_elem_const, make_loss
16+
from deepks.core.ml.eval.evaluator import Evaluator, NatomLossList
17+
18+
def train(model, g_reader, n_epoch=1000, test_reader=None, *,
19+
energy_factor=1., force_factor=0., stress_factor=0., orbital_factor=0., v_delta_factor=0., v_delta_r_factor=0., phi_factor=0.,phi_occ=0, band_factor=0., band_occ=0, density_m_factor=0., density_m_occ=0, density_factor=0.,
20+
energy_loss=None, force_loss=None, stress_loss=None, orbital_loss=None, v_delta_loss=None, v_delta_r_loss=None, phi_loss=None, band_loss=None, density_m_loss=None, grad_penalty=0.,
21+
energy_per_atom=0, vd_divide_by_nlocal=False,
22+
start_lr=0.001, decay_steps=100, decay_rate=0.96, stop_lr=None, decay_rate_iter=None,
23+
weight_decay=0., fix_embedding=False,
24+
display_epoch=100, display_detail_test=0, display_natom_loss=False, ckpt_file="model.pth",
25+
graph_file=None, device=DEVICE):
26+
27+
model = model.to(device)
28+
model.eval()
29+
print("# working on device:", device)
30+
if test_reader is None:
31+
test_reader = g_reader
32+
# fix parameters if needed
33+
if fix_embedding and model.embedder is not None:
34+
model.embedder.requires_grad_(False)
35+
# set up optimizer and lr scheduler
36+
if decay_rate_iter is not None:
37+
# decay_rate of start_lr for iterations, often start from iter.00
38+
current_dir=os.getcwd()
39+
current_iter=current_dir.split("/")[-2].split(".")[-1]
40+
if current_iter != "init": # no need to change
41+
current_iter=int(current_iter)
42+
start_lr=start_lr*(decay_rate_iter**current_iter)
43+
print(f"# resetting start_lr to {start_lr:.2e} because of decay_rate_iter")
44+
optimizer = optim.Adam(model.parameters(), lr=start_lr, weight_decay=weight_decay)
45+
if stop_lr is not None:
46+
decay_rate = (stop_lr / start_lr) ** (1 / (n_epoch // decay_steps))
47+
print(f"# resetting decay_rate: {decay_rate:.4f} "
48+
+ f"to satisfy stop_lr: {stop_lr:.2e}")
49+
scheduler = optim.lr_scheduler.StepLR(optimizer, decay_steps, decay_rate)
50+
# make evaluators for training
51+
evaluator = Evaluator(energy_factor=energy_factor, force_factor=force_factor,
52+
stress_factor=stress_factor, orbital_factor=orbital_factor,
53+
v_delta_factor=v_delta_factor, v_delta_r_factor=v_delta_r_factor,
54+
phi_factor=phi_factor, phi_occ=phi_occ,
55+
band_factor=band_factor, band_occ=band_occ,
56+
density_m_factor=density_m_factor, density_m_occ=density_m_occ,
57+
energy_lossfn=energy_loss, force_lossfn=force_loss,
58+
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
59+
v_delta_lossfn=v_delta_loss, v_delta_r_lossfn=v_delta_r_loss, phi_lossfn=phi_loss,
60+
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
61+
density_factor=density_factor, grad_penalty=grad_penalty,
62+
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)
63+
if not display_detail_test:
64+
# make test evaluator that only returns l2loss of energy
65+
test_eval = Evaluator(energy_factor=1., energy_lossfn=make_loss(), # default l2 loss
66+
force_factor=0., density_factor=0., grad_penalty=0.,energy_per_atom=energy_per_atom)
67+
else:
68+
# make test evaluator that returns loss of every concerned items, but all with factor==1
69+
to_one = lambda x: 0. if x == 0. else 1.
70+
test_eval = Evaluator(energy_factor=to_one(energy_factor), force_factor=to_one(force_factor),
71+
stress_factor=to_one(stress_factor), orbital_factor=to_one(orbital_factor),
72+
v_delta_factor=to_one(v_delta_factor), v_delta_r_factor=to_one(v_delta_r_factor),
73+
phi_factor=to_one(phi_factor), phi_occ=phi_occ,
74+
band_factor=to_one(band_factor), band_occ=band_occ,
75+
density_m_factor=to_one(density_m_factor), density_m_occ=density_m_occ,
76+
energy_lossfn=energy_loss, force_lossfn=force_loss,
77+
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
78+
v_delta_lossfn=v_delta_loss, v_delta_r_lossfn=v_delta_r_loss, phi_lossfn=phi_loss,
79+
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
80+
density_factor=to_one(density_factor), grad_penalty=grad_penalty,
81+
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)
82+
83+
print("# epoch trn_err tst_err lr trn_time tst_time",end='')
84+
data_keys = g_reader.readers[0].sample_all().keys()
85+
# L_inv_in=1 if "L_inv" in data_keys else 0
86+
# print("if L_inv in sample:",L_inv_in)
87+
align_len=20
88+
evaluator.print_head("trn_loss",data_keys,align_len)
89+
if display_detail_test:
90+
test_eval.print_head("tst_loss",data_keys,align_len)
91+
# print("")
92+
93+
tic = time()
94+
trn_natom_loss_list=NatomLossList()
95+
tst_natom_loss_list=NatomLossList()
96+
for batch in g_reader.sample_all_batch():
97+
loss=evaluator(model,batch)
98+
natom=batch["eig"].shape[1]
99+
trn_natom_loss_list.add_loss(natom,loss)
100+
trn_loss=trn_natom_loss_list.avg_loss()
101+
for batch in test_reader.sample_all_batch():
102+
loss=test_eval(model,batch)
103+
natom=batch["eig"].shape[1]
104+
tst_natom_loss_list.add_loss(natom,loss)
105+
tst_loss=tst_natom_loss_list.avg_loss()
106+
# trn_loss = np.mean([[loss_term.item() for loss_term in evaluator(model, batch)]
107+
# for batch in g_reader.sample_all_batch()],axis=0)
108+
# tst_loss = np.mean([[loss_term.item() for loss_term in test_eval(model, batch)]
109+
# for batch in test_reader.sample_all_batch()],axis=0)
110+
tst_time = time() - tic
111+
if display_natom_loss:
112+
for natom in trn_natom_loss_list.natoms():
113+
evaluator.print_head(str(natom)+"_trn",data_keys,align_len)
114+
for natom in tst_natom_loss_list.natoms():
115+
if display_detail_test:
116+
test_eval.print_head(str(natom)+"_tst",data_keys,align_len)
117+
else:
118+
test_eval.print_head(str(natom)+"_tst",[],align_len)#just energy
119+
print("")
120+
121+
print(f" {0:<8d} {np.sqrt(np.abs(trn_loss[-1])):>.2e} {np.sqrt(np.abs(tst_loss[-1])):>.2e}"
122+
f" {start_lr:>.2e} {0:>8.2f} {tst_time:>8.2f}",end='')
123+
for loss_term in trn_loss[:-1]:
124+
print(f"{loss_term:>{align_len}.4e}",end='')
125+
if display_detail_test:
126+
for loss_term in tst_loss[:-1]:
127+
print(f"{loss_term:>{align_len}.4e}",end='')
128+
if display_natom_loss:
129+
trn_natom_loss_list.print_avg_atom_loss(align_len)
130+
tst_natom_loss_list.print_avg_atom_loss(align_len)
131+
print('')
132+
133+
for epoch in range(1, n_epoch+1):
134+
tic = time()
135+
# loss_list = []
136+
trn_natom_loss_list.clear_loss()
137+
tst_natom_loss_list.clear_loss()
138+
for sample in g_reader:
139+
model.train()
140+
optimizer.zero_grad()
141+
loss = evaluator(model, sample)
142+
loss[-1].backward()
143+
# print("vdr_pred grad:",evaluator.vdr_pred.grad)
144+
# print("e_loss grad:",evaluator.e_loss.grad)
145+
# print("vdr_loss grad:",evaluator.vdr_loss.grad)
146+
# print("tot_loss grad:",evaluator.tot_loss.grad)
147+
optimizer.step()
148+
# loss_list.append([loss_term.item() for loss_term in loss])
149+
natom=sample["eig"].shape[1]
150+
trn_natom_loss_list.add_loss(natom,loss)
151+
scheduler.step()
152+
153+
if epoch % display_epoch == 0:
154+
model.eval()
155+
# trn_loss = np.mean(loss_list,axis=0)
156+
trn_loss=trn_natom_loss_list.avg_loss()
157+
trn_time = time() - tic
158+
tic = time()
159+
# tst_loss = np.mean([[loss_term.item() for loss_term in test_eval(model, batch)]
160+
# for batch in test_reader.sample_all_batch()],axis=0)
161+
for batch in test_reader.sample_all_batch():
162+
loss=test_eval(model,batch)
163+
natom=batch["eig"].shape[1]
164+
tst_natom_loss_list.add_loss(natom,loss)
165+
tst_loss=tst_natom_loss_list.avg_loss()
166+
tst_time = time() - tic
167+
print(f" {epoch:<8d} {np.sqrt(np.abs(trn_loss[-1])):>.2e} {np.sqrt(np.abs(tst_loss[-1])):>.2e}"
168+
f" {scheduler.get_last_lr()[0]:>.2e} {trn_time:>8.2f} {tst_time:8.2f}",end='')
169+
for loss_term in trn_loss[:-1]:
170+
print(f"{loss_term:>{align_len}.4e}",end='')
171+
if display_detail_test and epoch%(display_detail_test*display_epoch) == 0:
172+
for loss_term in tst_loss[:-1]:
173+
print(f"{loss_term:>{align_len}.4e}",end='')
174+
if display_natom_loss:
175+
trn_natom_loss_list.print_avg_atom_loss(align_len)
176+
tst_natom_loss_list.print_avg_atom_loss(align_len)
177+
print('')
178+
if ckpt_file:
179+
model.save(ckpt_file)
180+
181+
if ckpt_file:
182+
model.save(ckpt_file)
183+
if graph_file:
184+
model.compile_save(graph_file)
185+
186+
187+
def main(train_paths, test_paths=None,
188+
restart=None, ckpt_file=None,
189+
model_args=None, data_args=None,
190+
preprocess_args=None, train_args=None,
191+
proj_basis=None, fit_elem=False,
192+
seed=None, device=None):
193+
194+
if seed is None:
195+
seed = np.random.randint(0, 2**32)
196+
print(f'# using seed: {seed}')
197+
np.random.seed(seed)
198+
torch.manual_seed(seed)
199+
200+
if model_args is None: model_args = {}
201+
if data_args is None: data_args = {}
202+
if preprocess_args is None: preprocess_args = {}
203+
if train_args is None: train_args = {}
204+
if proj_basis is not None:
205+
model_args["proj_basis"] = proj_basis
206+
if ckpt_file is not None:
207+
train_args["ckpt_file"] = ckpt_file
208+
if device is not None:
209+
train_args["device"] = device
210+
211+
train_paths = load_dirs(train_paths)
212+
# print(f'# training with {len(train_paths)} system(s)')
213+
g_reader = GroupReader(train_paths, **data_args)
214+
if test_paths is not None:
215+
test_paths = load_dirs(test_paths)
216+
# print(f'# testing with {len(test_paths)} system(s)')
217+
test_reader = GroupReader(test_paths, **data_args)
218+
else:
219+
print('# testing with training set')
220+
test_reader = None
221+
222+
if restart is not None:
223+
model = CorrNet.load(restart)
224+
if model.elem_table is not None:
225+
fit_elem_const(g_reader, test_reader, model.elem_table)
226+
else:
227+
input_dim = g_reader.ndesc
228+
if model_args.get("input_dim", input_dim) != input_dim:
229+
print(f"# `input_dim` in `model_args` does not match data",
230+
f"({input_dim}).", "Use the one in data.", file=sys.stderr)
231+
model_args["input_dim"] = input_dim
232+
if fit_elem:
233+
elem_table = model_args.get("elem_table", None)
234+
if isinstance(elem_table, str):
235+
elem_table = load_elem_table(elem_table)
236+
elem_table = fit_elem_const(g_reader, test_reader, elem_table)
237+
model_args["elem_table"] = elem_table
238+
model = CorrNet(**model_args).double()
239+
240+
preprocess(model, g_reader, **preprocess_args)
241+
# start=time()
242+
train(model, g_reader, test_reader=test_reader, **train_args)
243+
# end=time()
244+
# print("all train time:",end-start)
245+
246+
247+
if __name__ == "__main__":
248+
from deepks.main import train_cli as cli
249+
cli()

0 commit comments

Comments
 (0)