Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 35 additions & 7 deletions deepks/model/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -182,9 +193,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
Expand Down Expand Up @@ -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"]
Expand Down Expand Up @@ -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)
Expand Down
24 changes: 17 additions & 7 deletions deepks/model/reader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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):
Expand All @@ -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()
Expand Down Expand Up @@ -184,11 +186,19 @@ 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)
L_inv=torch.linalg.inv(L)
self.t_data["L_inv"]=L_inv\
# 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,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\
Expand Down
8 changes: 6 additions & 2 deletions deepks/model/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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)

Expand Down
Loading