Skip to content

Commit 168d1b1

Browse files
committed
Add phi alignment loss. Set phi_align_factor and phi_align_occ to use it.
1 parent b864680 commit 168d1b1

2 files changed

Lines changed: 38 additions & 6 deletions

File tree

deepks/model/evaluator.py

Lines changed: 32 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -18,10 +18,12 @@ def __init__(self,
1818
phi_factor=0., phi_occ=0,
1919
band_factor=0.,band_occ=0,
2020
density_m_factor=0.,density_m_occ=0,
21+
phi_align_factor=0., phi_align_occ=0,
2122
density_factor=0., grad_penalty=0.,
2223
energy_lossfn=None, force_lossfn=None,
2324
stress_lossfn=None, orbital_lossfn=None,
2425
v_delta_lossfn=None, phi_lossfn=None,
26+
phi_align_lossfn=None,
2527
band_lossfn=None, density_m_lossfn=None,
2628
energy_per_atom=0,vd_divide_by_nlocal=False):
2729
# energy term
@@ -83,7 +85,15 @@ def __init__(self,
8385
density_m_lossfn = make_loss(**density_m_lossfn)
8486
self.density_m_factor = density_m_factor
8587
self.density_m_lossfn = density_m_lossfn
86-
self.get_density_m_occ = get_occ_func(density_m_occ)
88+
self.get_density_m_occ = get_occ_func(density_m_occ)
89+
# phi alignment term
90+
if phi_align_lossfn is None:
91+
phi_align_lossfn = {}
92+
if isinstance(phi_align_lossfn, dict):
93+
phi_align_lossfn = make_loss(**phi_align_lossfn)
94+
self.phi_align_factor = phi_align_factor
95+
self.phi_align_lossfn = phi_align_lossfn
96+
self.get_phi_align_occ = get_occ_func(phi_align_occ)
8797
# coulomb term of dm; requires head gradient
8898
self.d_factor = density_factor
8999
# gradient penalty, not very useful
@@ -156,9 +166,10 @@ def __call__(self, model, sample):
156166
# print(o_label.shape, op.shape, o_pred.shape, gev.shape)
157167
tot_loss = tot_loss + self.o_factor * self.o_lossfn(o_pred, o_label)
158168
loss.append(self.o_factor * self.o_lossfn(o_pred, o_label))
159-
# optional v_delta/phi/band_energy/density_matrix calculation
169+
# optional v_delta/phi/band_energy/density_matrix/phi_alignment calculation
160170
if (self.vd_factor > 0 and "lb_vd" in sample) or (self.phi_factor > 0 and "lb_phi" in sample) \
161-
or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample):
171+
or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample) \
172+
or (self.phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample):
162173
# cal v_delta
163174
if "vdp" in sample:
164175
vdp = sample["vdp"] # can be complex
@@ -212,6 +223,20 @@ def __call__(self, model, sample):
212223
density_m_loss = self.density_m_factor * self.density_m_lossfn(density_m_pred, density_m_label) * nlocal
213224
tot_loss = tot_loss + density_m_loss
214225
loss.append(density_m_loss)
226+
227+
# optional phi alignment calculation, don't need eigh on vd_pred
228+
if self.phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample:
229+
phi_label = sample["lb_phi"]
230+
band_label = sample["lb_band"]
231+
occ = self.get_phi_align_occ(natom)
232+
occ_phi_label = phi_label[..., :occ].clone()
233+
occ_band_label = band_label[..., :occ].clone()
234+
# phi_align_band should close to diagnoal matrix of occ_band_label
235+
phi_align_band = occ_phi_label.mT @ vd_pred @ occ_phi_label
236+
true_diag_band = torch.diag_embed(occ_band_label)
237+
phi_align_loss = self.phi_align_factor * self.phi_align_lossfn(phi_align_band, true_diag_band)
238+
tot_loss = tot_loss + phi_align_loss
239+
loss.append(phi_align_loss)
215240
# density loss with fix head grad
216241
if self.d_factor > 0 and "gldv" in sample:
217242
gldv = sample["gldv"]
@@ -245,7 +270,10 @@ def print_head(self,name,data_keys,align_len=20):
245270
info+=f"{name}_band".rjust(align_len)
246271
# optional density matrix calculation
247272
if self.density_m_factor > 0 and "lb_phi" in data_keys:
248-
info+=f"{name}_dm".rjust(align_len)
273+
info+=f"{name}_dm".rjust(align_len)
274+
# optional phi alignment calculation
275+
if self.phi_align_factor > 0 and "lb_phi" in data_keys and "lb_band" in data_keys:
276+
info+=f"{name}_phi_align".rjust(align_len)
249277
# density loss with fix head grad
250278
if self.d_factor > 0 and "gldv" in data_keys:
251279
info+=f"{name}_density".rjust(align_len)

deepks/model/train.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,8 @@
1616
from deepks.model.evaluator import Evaluator, NatomLossList
1717

1818
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., 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, phi_loss=None, band_loss=None, density_m_loss=None, grad_penalty=0.,
19+
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.,
20+
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.,
2121
energy_per_atom=0, vd_divide_by_nlocal=False,
2222
start_lr=0.001, decay_steps=100, decay_rate=0.96, stop_lr=None, decay_rate_iter=None,
2323
weight_decay=0., fix_embedding=False,
@@ -54,10 +54,12 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
5454
phi_factor=phi_factor, phi_occ=phi_occ,
5555
band_factor=band_factor, band_occ=band_occ,
5656
density_m_factor=density_m_factor, density_m_occ=density_m_occ,
57+
phi_align_factor=phi_align_factor, phi_align_occ=phi_align_occ,
5758
energy_lossfn=energy_loss, force_lossfn=force_loss,
5859
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
5960
v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss,
6061
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
62+
phi_align_lossfn=phi_align_loss,
6163
density_factor=density_factor, grad_penalty=grad_penalty,
6264
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)
6365
if not display_detail_test:
@@ -73,10 +75,12 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
7375
phi_factor=to_one(phi_factor), phi_occ=phi_occ,
7476
band_factor=to_one(band_factor), band_occ=band_occ,
7577
density_m_factor=to_one(density_m_factor), density_m_occ=density_m_occ,
78+
phi_align_factor=to_one(phi_align_factor), phi_align_occ=phi_align_occ,
7679
energy_lossfn=energy_loss, force_lossfn=force_loss,
7780
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
7881
v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss,
7982
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
83+
phi_align_lossfn=phi_align_loss,
8084
density_factor=to_one(density_factor), grad_penalty=grad_penalty,
8185
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)
8286

0 commit comments

Comments
 (0)