Skip to content

Commit 0728867

Browse files
committed
add bandgap loss
1 parent 612b9d2 commit 0728867

3 files changed

Lines changed: 38 additions & 9 deletions

File tree

deepks/model/evaluator.py

Lines changed: 27 additions & 4 deletions
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.model.reader import generalized_eigh, eigh_wrapper
11-
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, cal_vd_masked_loss, get_occ_func, make_loss
11+
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, cal_vd_masked_loss, cal_bandgap, get_occ_func, make_loss
1212

1313
class Evaluator:
1414
def __init__(self,
@@ -17,14 +17,15 @@ def __init__(self,
1717
v_delta_factor=0.,
1818
phi_factor=0., phi_occ=0,
1919
band_factor=0.,band_occ=0,
20+
bandgap_factor=0.,bandgap_occ=0,
2021
density_m_factor=0.,density_m_occ=0,
2122
phi_align_factor=0., phi_align_occ=0,
2223
density_factor=0., grad_penalty=0.,
2324
energy_lossfn=None, force_lossfn=None,
2425
stress_lossfn=None, orbital_lossfn=None,
2526
v_delta_lossfn=None, phi_lossfn=None,
2627
phi_align_lossfn=None,
27-
band_lossfn=None, density_m_lossfn=None,
28+
band_lossfn=None, bandgap_lossfn=None, density_m_lossfn=None,
2829
energy_per_atom=0,vd_divide_by_nlocal=False,
2930
vd_masked_loss=False,
3031
vd_masked_S_threshold=1e-6, vd_masked_H_threshold=1e-6,
@@ -81,6 +82,14 @@ def __init__(self,
8182
self.band_factor = band_factor
8283
self.band_lossfn = band_lossfn
8384
self.get_band_occ = get_occ_func(band_occ)
85+
# bandgap term
86+
if bandgap_lossfn is None:
87+
bandgap_lossfn = {}
88+
if isinstance(bandgap_lossfn, dict):
89+
bandgap_lossfn = make_loss(**bandgap_lossfn)
90+
self.bandgap_factor = bandgap_factor
91+
self.bandgap_lossfn = bandgap_lossfn
92+
self.get_bandgap_occ = get_occ_func(bandgap_occ)
8493
#density matrix term
8594
if density_m_lossfn is None:
8695
density_m_lossfn = {}
@@ -136,6 +145,7 @@ def __call__(self, model, sample):
136145
or (self.vd_factor > 0 and "lb_vd" in sample)
137146
or (self.phi_factor > 0 and "lb_phi" in sample)
138147
or (self.band_factor > 0 and "lb_band" in sample)
148+
or (self.bandgap_factor > 0 and "lb_band" in sample)
139149
or (self.density_m_factor > 0)
140150
or (self.d_factor > 0 and "gldv" in sample)
141151
or self.g_penalty > 0)
@@ -178,7 +188,8 @@ def __call__(self, model, sample):
178188
loss.append(self.o_factor * self.o_lossfn(o_pred, o_label))
179189
# optional v_delta/phi/band_energy/density_matrix/phi_alignment calculation
180190
if (self.vd_factor > 0 and "lb_vd" in sample) or (self.phi_factor > 0 and "lb_phi" in sample) \
181-
or (self.band_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample) \
191+
or (self.band_factor > 0 and "lb_band" in sample) or (self.bandgap_factor > 0 and "lb_band" in sample) \
192+
or (self.density_m_factor > 0 and "lb_phi" in sample) \
182193
or (self.phi_align_factor > 0 and "lb_phi" in sample and "lb_band" in sample):
183194
# cal v_delta
184195
if "vdp" in sample:
@@ -204,7 +215,7 @@ def __call__(self, model, sample):
204215
tot_loss = tot_loss + vd_loss
205216
loss.append(vd_loss)
206217

207-
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):
218+
if (self.phi_factor > 0 and "lb_phi" in sample) or (self.band_factor > 0 and "lb_band" in sample) or (self.bandgap_factor > 0 and "lb_band" in sample) or (self.density_m_factor > 0 and "lb_phi" in sample):
208219
h_base = sample["h_base"]
209220
if "trans_matrix" in sample:
210221
trans_matrix=sample["trans_matrix"]
@@ -225,6 +236,15 @@ def __call__(self, model, sample):
225236
tot_loss = tot_loss + band_loss
226237
# print("occ_band",band_pred[...,:band_occ],band_label[...,:band_occ])
227238
loss.append(band_loss)
239+
# optional bandgap calculation
240+
if self.bandgap_factor > 0 and "lb_band" in sample:
241+
band_label=sample["lb_band"]
242+
bandgap_occ=self.get_bandgap_occ(natom)
243+
bandgap_label=cal_bandgap(band_label, bandgap_occ)
244+
bandgap_pred=cal_bandgap(band_pred, bandgap_occ)
245+
bandgap_loss = self.bandgap_factor * self.bandgap_lossfn(bandgap_pred, bandgap_label)
246+
tot_loss = tot_loss + bandgap_loss
247+
loss.append(bandgap_loss)
228248
# optional density matrix calculation
229249
if self.density_m_factor > 0 and "lb_phi" in sample:
230250
# calculate density_m_label every time, kind of waste of time
@@ -283,6 +303,9 @@ def print_head(self,name,data_keys,align_len=20):
283303
# optional band energy calculation
284304
if self.band_factor > 0 and "lb_band" in data_keys:
285305
info+=f"{name}_band".rjust(align_len)
306+
# optional bandgap calculation
307+
if self.bandgap_factor > 0 and "lb_band" in data_keys:
308+
info+=f"{name}_bandgap".rjust(align_len)
286309
# optional density matrix calculation
287310
if self.density_m_factor > 0 and "lb_phi" in data_keys:
288311
info+=f"{name}_dm".rjust(align_len)

deepks/model/train.py

Lines changed: 8 additions & 5 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, 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.,
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, bandgap_factor=0., bandgap_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, bandgap_loss=None, density_m_loss=None, phi_align_loss=None, grad_penalty=0.,
2121
energy_per_atom=0, vd_divide_by_nlocal=False, vd_masked_loss=False, vd_masked_S_threshold=1e-6, vd_masked_H_threshold=1e-6, use_safe_eigh=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,
@@ -53,12 +53,14 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
5353
v_delta_factor=v_delta_factor,
5454
phi_factor=phi_factor, phi_occ=phi_occ,
5555
band_factor=band_factor, band_occ=band_occ,
56+
bandgap_factor=bandgap_factor, bandgap_occ=bandgap_occ,
5657
density_m_factor=density_m_factor, density_m_occ=density_m_occ,
5758
phi_align_factor=phi_align_factor, phi_align_occ=phi_align_occ,
5859
energy_lossfn=energy_loss, force_lossfn=force_loss,
5960
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
6061
v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss,
61-
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
62+
band_lossfn=band_loss, bandgap_lossfn=bandgap_loss,
63+
density_m_lossfn=density_m_loss,
6264
phi_align_lossfn=phi_align_loss,
6365
density_factor=density_factor, grad_penalty=grad_penalty,
6466
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
@@ -76,13 +78,14 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
7678
v_delta_factor=to_one(v_delta_factor),
7779
phi_factor=to_one(phi_factor), phi_occ=phi_occ,
7880
band_factor=to_one(band_factor), band_occ=band_occ,
81+
bandgap_factor=to_one(bandgap_factor), bandgap_occ=bandgap_occ,
7982
density_m_factor=to_one(density_m_factor), density_m_occ=density_m_occ,
8083
phi_align_factor=to_one(phi_align_factor), phi_align_occ=phi_align_occ,
8184
energy_lossfn=energy_loss, force_lossfn=force_loss,
8285
stress_lossfn=stress_loss, orbital_lossfn=orbital_loss,
8386
v_delta_lossfn=v_delta_loss,phi_lossfn=phi_loss,
84-
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
85-
phi_align_lossfn=phi_align_loss,
87+
band_lossfn=band_loss, bandgap_lossfn=bandgap_loss,
88+
density_m_lossfn=density_m_loss, phi_align_lossfn=phi_align_loss,
8689
density_factor=to_one(density_factor), grad_penalty=grad_penalty,
8790
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
8891
vd_masked_loss=vd_masked_loss, vd_masked_S_threshold=vd_masked_S_threshold, vd_masked_H_threshold=vd_masked_H_threshold,

deepks/model/utils.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -187,6 +187,9 @@ def cal_vd_masked_loss(H_pred, H_label, S_matrix, S_threshold=1e-6, H_threshold=
187187

188188
return masked_sum / (active_elements + 1e-12)
189189

190+
def cal_bandgap(band, occ):
191+
return band[...,occ] - band[...,occ-1]
192+
190193
def get_occ_func(occ):
191194
# print("type:",type(occ))
192195
if isinstance(occ, int):

0 commit comments

Comments
 (0)