From 7c77176b7605eba3a537c4fe2ddc35cd2aa64e6a Mon Sep 17 00:00:00 2001 From: Liangxuan <1351620715@qq.com> Date: Fri, 24 Oct 2025 16:02:30 +0800 Subject: [PATCH 1/3] use trans_matrix instead of L_inv, modify generalized_eigh accordingly --- deepks/model/evaluator.py | 6 +++--- deepks/model/reader.py | 12 ++++++------ 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/deepks/model/evaluator.py b/deepks/model/evaluator.py index 9582daef..1d3315d2 100644 --- a/deepks/model/evaluator.py +++ b/deepks/model/evaluator.py @@ -182,9 +182,9 @@ def __call__(self, model, sample): if (self.phi_factor > 0 and "lb_phi" in sample) or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample): h_base = sample["h_base"] - if "L_inv" in sample: - L_inv=sample["L_inv"] - band_pred,phi_pred=generalized_eigh(h_base+vd_pred,L_inv) + if "trans_matrix" in sample: + trans_matrix=sample["trans_matrix"] + band_pred,phi_pred=generalized_eigh(h_base+vd_pred,trans_matrix) else: band_pred,phi_pred= torch.linalg.eigh(h_base+vd_pred,UPLO='U') # optional phi calculation diff --git a/deepks/model/reader.py b/deepks/model/reader.py index 25be4ddc..5b785dc7 100644 --- a/deepks/model/reader.py +++ b/deepks/model/reader.py @@ -19,10 +19,10 @@ def split_batch(tdict, size, dim=0): for i in range(nsecs[0]) ] -def generalized_eigh(h,L_inv): - symm_h=L_inv @ h @ L_inv.mT +def generalized_eigh(h,trans_matrix): + symm_h=trans_matrix.mT @ h @ trans_matrix e,v=torch.linalg.eigh(symm_h) - phi=L_inv.mT @ v + phi=trans_matrix @ v return e,phi class Reader(object): @@ -185,10 +185,10 @@ def prepare(self): #print("use generalized eigh") overlap=torch.tensor(np.load(self.overlap_path)) L=torch.linalg.cholesky(overlap) - L_inv=torch.linalg.inv(L) - self.t_data["L_inv"]=L_inv\ + trans_matrix=torch.linalg.inv(L).mT + self.t_data["trans_matrix"]=trans_matrix\ .reshape(raw_nframes, -1, self.nlocal, self.nlocal)[conv].clone() - band_ref,phi_ref=generalized_eigh(h_ref,L_inv) + band_ref,phi_ref=generalized_eigh(h_ref,trans_matrix) else: band_ref,phi_ref=torch.linalg.eigh(h_ref,UPLO='U') # U for upper triangle self.t_data["lb_band"]=band_ref\ From b864680c692aa9ebf3a6f00fabe01a0133bebfa0 Mon Sep 17 00:00:00 2001 From: Liangxuan <1351620715@qq.com> Date: Mon, 27 Oct 2025 15:41:54 +0800 Subject: [PATCH 2/3] Add a new eigh_method by setting eigh_method as 2. It can enhance the precision in computing smaller eigenvalues --- deepks/model/reader.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/deepks/model/reader.py b/deepks/model/reader.py index 5b785dc7..1d65bba9 100644 --- a/deepks/model/reader.py +++ b/deepks/model/reader.py @@ -35,6 +35,7 @@ def __init__(self, data_path, batch_size, phialpha_name="phialpha",gevdm_name="grad_evdm", h_base_name="h_base",h_ref_name="hamiltonian", read_overlap = False, overlap_name="overlap", + eigh_method = 1, eg_name="eg_base", gveg_name="grad_veg", gldv_name="grad_ldv", conv_name="conv", atom_name="atom", **kwargs): @@ -61,6 +62,7 @@ def __init__(self, data_path, batch_size, self.c_path = self.check_exist(conv_name+".npy") self.a_path = self.check_exist(atom_name+".npy") self.read_overlap = read_overlap + self.eigh_method = eigh_method # load data self.load_meta() self.prepare() @@ -184,8 +186,16 @@ def prepare(self): if self.read_overlap is True and self.overlap_path is not None: #print("use generalized eigh") overlap=torch.tensor(np.load(self.overlap_path)) - L=torch.linalg.cholesky(overlap) - trans_matrix=torch.linalg.inv(L).mT + # When overlap matrix is ill-conditioned, the eigenvalues (i.e. band) can suffer from significant roundoff errors. + if self.eigh_method == 1: + L=torch.linalg.cholesky(overlap) + trans_matrix=torch.linalg.inv(L).mT + # Substitute cholesky with eigen decomposition. + # This modification effectively reorders the entries of symm_h, placing larger values towards the upper left-hand corner, thereby enhancing the precision in computing smaller eigenvalues + elif self.eigh_method == 2: + overlap_eigenvalue,overlap_eigenvector=torch.linalg.eigh(overlap) + sigma_inv_sqrt = torch.diag_embed(1.0 / torch.sqrt(overlap_eigenvalue)) + trans_matrix=overlap_eigenvector @ sigma_inv_sqrt self.t_data["trans_matrix"]=trans_matrix\ .reshape(raw_nframes, -1, self.nlocal, self.nlocal)[conv].clone() band_ref,phi_ref=generalized_eigh(h_ref,trans_matrix) From 168d1b11973a7f2dd01cda700cf21f81be564ab9 Mon Sep 17 00:00:00 2001 From: Liangxuan <1351620715@qq.com> Date: Mon, 27 Oct 2025 20:10:09 +0800 Subject: [PATCH 3/3] Add phi alignment loss. Set phi_align_factor and phi_align_occ to use it. --- deepks/model/evaluator.py | 36 ++++++++++++++++++++++++++++++++---- deepks/model/train.py | 8 ++++++-- 2 files changed, 38 insertions(+), 6 deletions(-) diff --git a/deepks/model/evaluator.py b/deepks/model/evaluator.py index 1d3315d2..ccb34fea 100644 --- a/deepks/model/evaluator.py +++ b/deepks/model/evaluator.py @@ -18,10 +18,12 @@ def __init__(self, phi_factor=0., phi_occ=0, band_factor=0.,band_occ=0, density_m_factor=0.,density_m_occ=0, + phi_align_factor=0., phi_align_occ=0, density_factor=0., grad_penalty=0., energy_lossfn=None, force_lossfn=None, stress_lossfn=None, orbital_lossfn=None, v_delta_lossfn=None, phi_lossfn=None, + phi_align_lossfn=None, band_lossfn=None, density_m_lossfn=None, energy_per_atom=0,vd_divide_by_nlocal=False): # energy term @@ -83,7 +85,15 @@ def __init__(self, density_m_lossfn = make_loss(**density_m_lossfn) self.density_m_factor = density_m_factor self.density_m_lossfn = density_m_lossfn - self.get_density_m_occ = get_occ_func(density_m_occ) + self.get_density_m_occ = get_occ_func(density_m_occ) + # phi alignment term + if phi_align_lossfn is None: + phi_align_lossfn = {} + if isinstance(phi_align_lossfn, dict): + phi_align_lossfn = make_loss(**phi_align_lossfn) + self.phi_align_factor = phi_align_factor + self.phi_align_lossfn = phi_align_lossfn + self.get_phi_align_occ = get_occ_func(phi_align_occ) # coulomb term of dm; requires head gradient self.d_factor = density_factor # gradient penalty, not very useful @@ -156,9 +166,10 @@ def __call__(self, model, sample): # print(o_label.shape, op.shape, o_pred.shape, gev.shape) tot_loss = tot_loss + self.o_factor * self.o_lossfn(o_pred, o_label) loss.append(self.o_factor * self.o_lossfn(o_pred, o_label)) - # optional v_delta/phi/band_energy/density_matrix calculation + # optional v_delta/phi/band_energy/density_matrix/phi_alignment calculation if (self.vd_factor > 0 and "lb_vd" in sample) or (self.phi_factor > 0 and "lb_phi" in sample) \ - or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample): + or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample) \ + or (self.phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample): # cal v_delta if "vdp" in sample: vdp = sample["vdp"] # can be complex @@ -212,6 +223,20 @@ def __call__(self, model, sample): density_m_loss = self.density_m_factor * self.density_m_lossfn(density_m_pred, density_m_label) * nlocal tot_loss = tot_loss + density_m_loss loss.append(density_m_loss) + + # optional phi alignment calculation, don't need eigh on vd_pred + if self.phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample: + phi_label = sample["lb_phi"] + band_label = sample["lb_band"] + occ = self.get_phi_align_occ(natom) + occ_phi_label = phi_label[..., :occ].clone() + occ_band_label = band_label[..., :occ].clone() + # phi_align_band should close to diagnoal matrix of occ_band_label + phi_align_band = occ_phi_label.mT @ vd_pred @ occ_phi_label + true_diag_band = torch.diag_embed(occ_band_label) + phi_align_loss = self.phi_align_factor * self.phi_align_lossfn(phi_align_band, true_diag_band) + tot_loss = tot_loss + phi_align_loss + loss.append(phi_align_loss) # density loss with fix head grad if self.d_factor > 0 and "gldv" in sample: gldv = sample["gldv"] @@ -245,7 +270,10 @@ def print_head(self,name,data_keys,align_len=20): info+=f"{name}_band".rjust(align_len) # optional density matrix calculation if self.density_m_factor > 0 and "lb_phi" in data_keys: - info+=f"{name}_dm".rjust(align_len) + info+=f"{name}_dm".rjust(align_len) + # optional phi alignment calculation + if self.phi_align_factor > 0 and "lb_phi" in data_keys and "lb_band" in data_keys: + info+=f"{name}_phi_align".rjust(align_len) # density loss with fix head grad if self.d_factor > 0 and "gldv" in data_keys: info+=f"{name}_density".rjust(align_len) diff --git a/deepks/model/train.py b/deepks/model/train.py index 24f98b87..c6357f36 100644 --- a/deepks/model/train.py +++ b/deepks/model/train.py @@ -16,8 +16,8 @@ from deepks.model.evaluator import Evaluator, NatomLossList def train(model, g_reader, n_epoch=1000, test_reader=None, *, - 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, density_m_factor=0., density_m_occ=0, density_factor=0., - energy_loss=None, force_loss=None, stress_loss=None, orbital_loss=None, v_delta_loss=None, phi_loss=None, band_loss=None, density_m_loss=None, grad_penalty=0., + 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, density_m_factor=0., density_m_occ=0, phi_align_factor=0., phi_align_occ=0, density_factor=0., + energy_loss=None, force_loss=None, stress_loss=None, orbital_loss=None, v_delta_loss=None, phi_loss=None, band_loss=None, density_m_loss=None, phi_align_loss=None, grad_penalty=0., energy_per_atom=0, vd_divide_by_nlocal=False, start_lr=0.001, decay_steps=100, decay_rate=0.96, stop_lr=None, decay_rate_iter=None, weight_decay=0., fix_embedding=False, @@ -54,10 +54,12 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *, phi_factor=phi_factor, phi_occ=phi_occ, band_factor=band_factor, band_occ=band_occ, density_m_factor=density_m_factor, density_m_occ=density_m_occ, + phi_align_factor=phi_align_factor, phi_align_occ=phi_align_occ, energy_lossfn=energy_loss, force_lossfn=force_loss, stress_lossfn=stress_loss, orbital_lossfn=orbital_loss, v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss, band_lossfn=band_loss, density_m_lossfn=density_m_loss, + phi_align_lossfn=phi_align_loss, density_factor=density_factor, grad_penalty=grad_penalty, energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal) if not display_detail_test: @@ -73,10 +75,12 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *, phi_factor=to_one(phi_factor), phi_occ=phi_occ, band_factor=to_one(band_factor), band_occ=band_occ, density_m_factor=to_one(density_m_factor), density_m_occ=density_m_occ, + phi_align_factor=to_one(phi_align_factor), phi_align_occ=phi_align_occ, energy_lossfn=energy_loss, force_lossfn=force_loss, stress_lossfn=stress_loss, orbital_lossfn=orbital_loss, v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss, band_lossfn=band_loss, density_m_lossfn=density_m_loss, + phi_align_lossfn=phi_align_loss, density_factor=to_one(density_factor), grad_penalty=grad_penalty, energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)