|
| 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