Skip to content

Commit 2907cf4

Browse files
committed
safe_eig to prevent NAN when decomposition
1 parent e259662 commit 2907cf4

4 files changed

Lines changed: 122 additions & 8 deletions

File tree

deepks/model/evaluator.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import deepks
88
except ImportError as e:
99
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
10-
from deepks.model.reader import generalized_eigh
10+
from deepks.model.reader import generalized_eigh, eigh_wrapper
1111
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, get_occ_func, make_loss
1212

1313
class Evaluator:
@@ -25,7 +25,8 @@ def __init__(self,
2525
v_delta_lossfn=None, phi_lossfn=None,
2626
phi_align_lossfn=None,
2727
band_lossfn=None, density_m_lossfn=None,
28-
energy_per_atom=0,vd_divide_by_nlocal=False):
28+
energy_per_atom=0,vd_divide_by_nlocal=False,
29+
use_safe_eigh=False):
2930
# energy term
3031
if energy_lossfn is None:
3132
energy_lossfn = {}
@@ -100,6 +101,8 @@ def __init__(self,
100101
self.g_penalty = grad_penalty
101102
# energy loss divide by 1/natom/natom^2
102103
self.energy_per_atom=energy_per_atom
104+
# use safe_eigh to prevent large grad because of decomposition
105+
self.use_safe_eigh=use_safe_eigh
103106

104107
def __call__(self, model, sample):
105108
_dref = next(model.parameters()).device
@@ -195,9 +198,9 @@ def __call__(self, model, sample):
195198
h_base = sample["h_base"]
196199
if "trans_matrix" in sample:
197200
trans_matrix=sample["trans_matrix"]
198-
band_pred,phi_pred=generalized_eigh(h_base+vd_pred,trans_matrix)
201+
band_pred,phi_pred=generalized_eigh(h_base+vd_pred,trans_matrix, self.use_safe_eigh)
199202
else:
200-
band_pred,phi_pred= torch.linalg.eigh(h_base+vd_pred,UPLO='U')
203+
band_pred,phi_pred= eigh_wrapper(h_base+vd_pred)
201204
# optional phi calculation
202205
if self.phi_factor > 0 and "lb_phi" in sample:
203206
phi_label = sample["lb_phi"]

deepks/model/reader.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import os,time,sys
22
import numpy as np
33
import torch
4+
from deepks.model.utils import safe_eigh
45

56
def concat_batch(tdicts, dim=0):
67
keys = tdicts[0].keys()
@@ -19,9 +20,21 @@ def split_batch(tdict, size, dim=0):
1920
for i in range(nsecs[0])
2021
]
2122

22-
def generalized_eigh(h,trans_matrix):
23+
def eigh_wrapper(a, use_safe_eigh=False):
24+
"""
25+
Wrapper for eigendecomposition that supports safe gradients for degenerate cases.
26+
Args:
27+
a: Symmetric/Hermitian matrix.
28+
use_safe_eigh: If True, uses SafeEigh to prevent NaN gradients.
29+
"""
30+
if use_safe_eigh:
31+
return safe_eigh(a)
32+
else:
33+
return torch.linalg.eigh(a, UPLO='U')
34+
35+
def generalized_eigh(h,trans_matrix,use_safe_eigh=False):
2336
symm_h=trans_matrix.mT @ h @ trans_matrix
24-
e,v=torch.linalg.eigh(symm_h)
37+
e,v=eigh_wrapper(symm_h, use_safe_eigh=use_safe_eigh)
2538
phi=trans_matrix @ v
2639
return e,phi
2740

deepks/model/train.py

Lines changed: 3 additions & 2 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, 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, density_m_loss=None, phi_align_loss=None, grad_penalty=0.,
21-
energy_per_atom=0, vd_divide_by_nlocal=False,
21+
energy_per_atom=0, vd_divide_by_nlocal=False, 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",
@@ -61,7 +61,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
6161
band_lossfn=band_loss, density_m_lossfn=density_m_loss,
6262
phi_align_lossfn=phi_align_loss,
6363
density_factor=density_factor, grad_penalty=grad_penalty,
64-
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal)
64+
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
65+
use_safe_eigh=use_safe_eigh)
6566
if not display_detail_test:
6667
# make test evaluator that only returns l2loss of energy
6768
test_eval = Evaluator(energy_factor=1., energy_lossfn=make_loss(), # default l2 loss

deepks/model/utils.py

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -167,3 +167,100 @@ def get_occ(natom):
167167
def get_occ(natom):
168168
return new_occ[natom]
169169
return get_occ
170+
171+
class SafeEigh(torch.autograd.Function):
172+
"""
173+
A custom autograd function for eigendecomposition of real symmetric matrices.
174+
It handles degenerate eigenvalues by masking out the infinite gradients
175+
caused by the term 1/(lambda_i - lambda_j) when lambda_i approx lambda_j.
176+
177+
Reference:
178+
Derivatives of Partial Eigendecomposition of a Real Symmetric Matrix
179+
for Degenerate Cases (Kasim et al., 2020), Equation (27).
180+
"""
181+
182+
@staticmethod
183+
def forward(ctx, a):
184+
"""
185+
Forward pass: Standard eigendecomposition.
186+
187+
Args:
188+
a: Input symmetric matrix. Shape: (..., N, N), supports batching.
189+
Returns:
190+
e: Eigenvalues. Shape: (..., N).
191+
v: Eigenvectors. Shape: (..., N, N).
192+
"""
193+
# Ensure the input is float/complex as required by eigh
194+
# Note: 'U' (Upper) or 'L' (Lower) doesn't matter much for valid symmetric inputs
195+
e, v = torch.linalg.eigh(a)
196+
197+
# Save tensors for the backward pass
198+
ctx.save_for_backward(e, v)
199+
return e, v
200+
201+
@staticmethod
202+
def backward(ctx, grad_e, grad_v):
203+
"""
204+
Backward pass: Computes gradient with respect to input matrix 'a'.
205+
206+
This implementation specifically handles the degeneracy issue where
207+
eigenvalues are identical or very close, which would normally cause
208+
NaNs or Infs in the gradient of eigenvectors.
209+
"""
210+
e, v = ctx.saved_tensors
211+
212+
# 1. Handle cases where gradients might be None
213+
# (e.g., if eigenvalues or eigenvectors are not used in the loss function)
214+
if grad_e is None:
215+
grad_e = torch.zeros_like(e)
216+
if grad_v is None:
217+
grad_v = torch.zeros_like(v)
218+
219+
# 2. Construct the pairwise difference matrix of eigenvalues
220+
# Shape of e: (Batch, N)
221+
# Use unsqueeze to broadcast: (Batch, N, 1) - (Batch, 1, N) -> (Batch, N, N)
222+
# e_diff[..., i, j] = e[..., i] - e[..., j] (column - row)
223+
e_diff = e.unsqueeze(-2) - e.unsqueeze(-1)
224+
225+
# 3. Handle Degeneracy (Masking)
226+
# Define a small threshold to detect degeneracy
227+
epsilon = 1e-8
228+
229+
# Create a mask where |lambda_i - lambda_j| > epsilon
230+
mask = torch.abs(e_diff) > epsilon
231+
232+
# Construct the F matrix: F_ij = 1 / (lambda_j - lambda_i)
233+
# Note: We use the transposed definition implicit in the matrix formula below.
234+
# Here we initialize f_matrix with zeros, effectively ignoring degenerate terms.
235+
f_matrix = torch.zeros_like(e_diff)
236+
237+
# Only compute division for non-degenerate pairs
238+
# This prevents division by zero and corresponds to setting the gradient
239+
# contribution of degenerate subspaces to zero (Gauge Invariance).
240+
f_matrix[mask] = 1.0 / e_diff[mask]
241+
242+
# 4. Compute the gradient w.r.t. the input matrix 'a'
243+
# Formula: grad_a = v @ (diag(grad_e) + F * (v^T @ grad_v)) @ v^T
244+
245+
# Projection of gradients onto the eigenvector basis: v^T @ grad_v
246+
# transpose(-2, -1) handles the last two dimensions for batch processing
247+
vt = v.transpose(-2, -1)
248+
v_t_grad_v = vt @ grad_v
249+
250+
# The middle term: diag(grad_e) + F * (v^T @ grad_v)
251+
# torch.diag_embed creates a diagonal matrix from the eigenvalue gradients
252+
# f_matrix * v_t_grad_v performs element-wise multiplication (Hadamard product)
253+
mid_term = torch.diag_embed(grad_e) + f_matrix * v_t_grad_v
254+
255+
# Transform back to the original basis
256+
grad_a = v @ mid_term @ vt
257+
258+
# 5. Enforce symmetry
259+
# Since the input 'a' is symmetric, its gradient must also be symmetric.
260+
grad_a = 0.5 * (grad_a + grad_a.transpose(-2, -1))
261+
262+
return grad_a
263+
264+
# Wrapper function for easy usage
265+
def safe_eigh(input_tensor):
266+
return SafeEigh.apply(input_tensor)

0 commit comments

Comments
 (0)