Skip to content

Commit 14c0b4d

Browse files
committed
feat: add vd_masked_loss_width
1 parent 0728867 commit 14c0b4d

3 files changed

Lines changed: 45 additions & 9 deletions

File tree

deepks/model/evaluator.py

Lines changed: 10 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, cal_bandgap, get_occ_func, make_loss
11+
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, cal_vd_masked_loss_hs, cal_vd_masked_loss_width, cal_bandgap, get_occ_func, make_loss
1212

1313
class Evaluator:
1414
def __init__(self,
@@ -27,8 +27,9 @@ def __init__(self,
2727
phi_align_lossfn=None,
2828
band_lossfn=None, bandgap_lossfn=None, density_m_lossfn=None,
2929
energy_per_atom=0,vd_divide_by_nlocal=False,
30-
vd_masked_loss=False,
30+
vd_masked_loss=0,
3131
vd_masked_S_threshold=1e-6, vd_masked_H_threshold=1e-6,
32+
vd_masked_width=1,
3233
use_safe_eigh=False):
3334
# energy term
3435
if energy_lossfn is None:
@@ -119,6 +120,8 @@ def __init__(self,
119120
# threshold for vd_masked_loss
120121
self.vd_masked_S_threshold=vd_masked_S_threshold
121122
self.vd_masked_H_threshold=vd_masked_H_threshold
123+
# width for vd_masked_loss_width
124+
self.vd_masked_width=vd_masked_width
122125

123126
def __call__(self, model, sample):
124127
_dref = next(model.parameters()).device
@@ -205,8 +208,11 @@ def __call__(self, model, sample):
205208
# optional v_delta calculation
206209
if self.vd_factor > 0 and "lb_vd" in sample:
207210
vd_label = sample["lb_vd"]
208-
if self.vd_masked_loss and "overlap" in sample:
209-
vd_loss = self.vd_factor * cal_vd_masked_loss(vd_pred, vd_label, sample["overlap"], self.vd_masked_S_threshold, self.vd_masked_H_threshold)
211+
if self.vd_masked_loss :
212+
if self.vd_masked_loss == 1 and "overlap" in sample:
213+
vd_loss = self.vd_factor * cal_vd_masked_loss_hs(vd_pred, vd_label, sample["overlap"], self.vd_masked_S_threshold, self.vd_masked_H_threshold)
214+
elif self.vd_masked_loss == 2:
215+
vd_loss = self.vd_factor * cal_vd_masked_loss_width(vd_pred, vd_label, self.vd_masked_width)
210216
else:
211217
vd_loss = self.vd_factor * self.vd_lossfn(vd_pred, vd_label)
212218
# original: mean method,divide by nlocal**2. vd_divide_by_nlocal:divide by nlocal

deepks/model/train.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
def train(model, g_reader, n_epoch=1000, test_reader=None, *,
1919
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.,
2020
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.,
21-
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,
21+
energy_per_atom=0, vd_divide_by_nlocal=False, vd_masked_loss=0, vd_masked_S_threshold=1e-6, vd_masked_H_threshold=1e-6, vd_masked_width=1, 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,
2424
display_epoch=100, display_detail_test=0, display_natom_loss=False, ckpt_file="model.pth",
@@ -64,7 +64,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
6464
phi_align_lossfn=phi_align_loss,
6565
density_factor=density_factor, grad_penalty=grad_penalty,
6666
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
67-
vd_masked_loss=vd_masked_loss, vd_masked_S_threshold=vd_masked_S_threshold, vd_masked_H_threshold=vd_masked_H_threshold,
67+
vd_masked_loss=vd_masked_loss, vd_masked_S_threshold=vd_masked_S_threshold,
68+
vd_masked_H_threshold=vd_masked_H_threshold, vd_masked_width=vd_masked_width,
6869
use_safe_eigh=use_safe_eigh)
6970
if not display_detail_test:
7071
# make test evaluator that only returns l2loss of energy
@@ -88,7 +89,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
8889
density_m_lossfn=density_m_loss, phi_align_lossfn=phi_align_loss,
8990
density_factor=to_one(density_factor), grad_penalty=grad_penalty,
9091
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
91-
vd_masked_loss=vd_masked_loss, vd_masked_S_threshold=vd_masked_S_threshold, vd_masked_H_threshold=vd_masked_H_threshold,
92+
vd_masked_loss=vd_masked_loss, vd_masked_S_threshold=vd_masked_S_threshold,
93+
vd_masked_H_threshold=vd_masked_H_threshold,vd_masked_width=vd_masked_width,
9294
use_safe_eigh=use_safe_eigh)
9395

9496
print("# epoch trn_err tst_err lr trn_time tst_time",end='')

deepks/model/utils.py

Lines changed: 30 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -157,7 +157,7 @@ def cal_phi_loss(phi_pred,phi_label,phi_occ):
157157
#print("loss.shape:",loss.shape)
158158
return loss
159159

160-
def cal_vd_masked_loss(H_pred, H_label, S_matrix, S_threshold=1e-6, H_threshold=1e-6):
160+
def cal_vd_masked_loss_hs(H_pred, H_label, S_matrix, S_threshold=1e-6, H_threshold=1e-6):
161161
"""
162162
Computes the Mean Squared Error of Hamiltonian elements filtered by
163163
the Overlap matrix and Hamiltonian matrix magnitude. Supports 4D tensors (nframe, nks, nlocal, nlocal).
@@ -175,7 +175,7 @@ def cal_vd_masked_loss(H_pred, H_label, S_matrix, S_threshold=1e-6, H_threshold=
175175
with torch.no_grad():
176176
# Generate mask where |S| > threshold or |H| > threshold.
177177
# Diagonals (S_ii=1) are naturally included if threshold < 1.
178-
mask = ( (torch.abs(S_matrix) > S_threshold) | (torch.abs(H_label) > H_threshold) ).to(H_pred.dtype)
178+
mask = ( (torch.abs(S_matrix) > S_threshold) | (torch.abs(H_label) > H_threshold) ).to(device=H_pred.device, dtype=H_pred.dtype)
179179

180180
# Compute element-wise squared difference
181181
diff_sq = (H_pred - H_label) ** 2
@@ -187,6 +187,34 @@ 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_vd_masked_loss_width(H_pred, H_label, width=1):
191+
"""
192+
Computes the Mean Squared Error of Hamiltonian elements filtered by width .
193+
Args:
194+
H_pred (torch.Tensor): Predicted Hamiltonian.
195+
H_label (torch.Tensor): Label Hamiltonian.
196+
width (int): Width of the mask
197+
Returns:
198+
torch.Tensor: Scalar loss value.
199+
"""
200+
with torch.no_grad():
201+
nlocal = H_pred.size(-1)
202+
i = torch.arange(nlocal, device=H_pred.device).view(nlocal, 1) # (nlocal, 1)
203+
j = torch.arange(nlocal, device=H_pred.device).view(1, nlocal) # (1, nlocal)
204+
205+
# calculate the shortest distance on the circle
206+
diff = torch.abs(i - j)
207+
dist = torch.minimum(diff, nlocal - diff)
208+
209+
# mask is 1 if distance < width, 0 otherwise
210+
mask = (dist < width).to(dtype=H_pred.dtype)
211+
212+
diff_sq = (H_pred - H_label) ** 2
213+
masked_sum = torch.sum(diff_sq * mask)
214+
active_elements = torch.sum(mask)
215+
216+
return masked_sum / (active_elements + 1e-12)
217+
190218
def cal_bandgap(band, occ):
191219
return band[...,occ] - band[...,occ-1]
192220

0 commit comments

Comments
 (0)