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
14 changes: 8 additions & 6 deletions deepks/model/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import deepks
except ImportError as e:
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
from deepks.model.reader import generalized_eigh
from deepks.model.reader import generalized_eigh, eigh_wrapper
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, get_occ_func, make_loss

class Evaluator:
Expand All @@ -25,7 +25,8 @@ def __init__(self,
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_per_atom=0,vd_divide_by_nlocal=False,
use_safe_eigh=False):
# energy term
if energy_lossfn is None:
energy_lossfn = {}
Expand Down Expand Up @@ -100,6 +101,8 @@ def __init__(self,
self.g_penalty = grad_penalty
# energy loss divide by 1/natom/natom^2
self.energy_per_atom=energy_per_atom
# use safe_eigh to prevent large grad because of decomposition
self.use_safe_eigh=use_safe_eigh

def __call__(self, model, sample):
_dref = next(model.parameters()).device
Expand Down Expand Up @@ -166,7 +169,6 @@ 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))
vd_pred=None
# 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) \
Expand Down Expand Up @@ -196,9 +198,9 @@ def __call__(self, model, sample):
h_base = sample["h_base"]
if "trans_matrix" in sample:
trans_matrix=sample["trans_matrix"]
band_pred,phi_pred=generalized_eigh(h_base+vd_pred,trans_matrix)
band_pred,phi_pred=generalized_eigh(h_base+vd_pred,trans_matrix, self.use_safe_eigh)
else:
band_pred,phi_pred= torch.linalg.eigh(h_base+vd_pred,UPLO='U')
band_pred,phi_pred= eigh_wrapper(h_base+vd_pred)
# optional phi calculation
if self.phi_factor > 0 and "lb_phi" in sample:
phi_label = sample["lb_phi"]
Expand Down Expand Up @@ -247,7 +249,7 @@ def __call__(self, model, sample):
tot_loss = tot_loss + d_loss
loss.append(d_loss)
loss.append(tot_loss)
return loss, vd_pred
return loss

def print_head(self,name,data_keys,align_len=20):
info=f"{name}_energy".rjust(align_len)
Expand Down
17 changes: 15 additions & 2 deletions deepks/model/reader.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import os,time,sys
import numpy as np
import torch
from deepks.model.utils import safe_eigh

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

def generalized_eigh(h,trans_matrix):
def eigh_wrapper(a, use_safe_eigh=False):
"""
Wrapper for eigendecomposition that supports safe gradients for degenerate cases.
Args:
a: Symmetric/Hermitian matrix.
use_safe_eigh: If True, uses SafeEigh to prevent NaN gradients.
"""
if use_safe_eigh:
return safe_eigh(a)
else:
return torch.linalg.eigh(a, UPLO='U')

def generalized_eigh(h,trans_matrix,use_safe_eigh=False):
symm_h=trans_matrix.mT @ h @ trans_matrix
e,v=torch.linalg.eigh(symm_h)
e,v=eigh_wrapper(symm_h, use_safe_eigh=use_safe_eigh)
phi=trans_matrix @ v
return e,phi

Expand Down
27 changes: 9 additions & 18 deletions deepks/model/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,13 @@
from deepks.model.model import CorrNet
from deepks.model.reader import GroupReader
from deepks.utils import load_dirs, load_elem_table
from deepks.model.utils import preprocess, fit_elem_const, make_loss, vd_grad_processor
from deepks.model.utils import preprocess, fit_elem_const, make_loss
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, 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, vd_grad_process=False, vd_grad_max=1.0, vd_grad_momentum=0.9,
energy_per_atom=0, vd_divide_by_nlocal=False, use_safe_eigh=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,
display_epoch=100, display_detail_test=0, display_natom_loss=False, ckpt_file="model.pth",
Expand Down Expand Up @@ -61,7 +61,8 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
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)
energy_per_atom=energy_per_atom, vd_divide_by_nlocal=vd_divide_by_nlocal,
use_safe_eigh=use_safe_eigh)
if not display_detail_test:
# make test evaluator that only returns l2loss of energy
test_eval = Evaluator(energy_factor=1., energy_lossfn=make_loss(), # default l2 loss
Expand Down Expand Up @@ -98,12 +99,12 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
trn_natom_loss_list=NatomLossList()
tst_natom_loss_list=NatomLossList()
for batch in g_reader.sample_all_batch():
loss,_=evaluator(model,batch)
loss=evaluator(model,batch)
natom=batch["eig"].shape[1]
trn_natom_loss_list.add_loss(natom,loss)
trn_loss=trn_natom_loss_list.avg_loss()
for batch in test_reader.sample_all_batch():
loss,_=test_eval(model,batch)
loss=test_eval(model,batch)
natom=batch["eig"].shape[1]
tst_natom_loss_list.add_loss(natom,loss)
tst_loss=tst_natom_loss_list.avg_loss()
Expand Down Expand Up @@ -134,9 +135,6 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
tst_natom_loss_list.print_avg_atom_loss(align_len)
print('')

if vd_grad_process:
vd_processor = vd_grad_processor(vd_grad_max, vd_grad_momentum)

for epoch in range(1, n_epoch+1):
tic = time()
# loss_list = []
Expand All @@ -145,13 +143,9 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
for sample in g_reader:
model.train()
optimizer.zero_grad()
loss, vd_pred = evaluator(model, sample)
if vd_grad_process and vd_pred is not None:
vd_processor.register_hook(vd_pred)
loss = evaluator(model, sample)
loss[-1].backward()
optimizer.step()
if vd_grad_process and vd_pred is not None:
vd_processor.remove_hook()
# loss_list.append([loss_term.item() for loss_term in loss])
natom=sample["eig"].shape[1]
trn_natom_loss_list.add_loss(natom,loss)
Expand All @@ -166,7 +160,7 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
# tst_loss = np.mean([[loss_term.item() for loss_term in test_eval(model, batch)]
# for batch in test_reader.sample_all_batch()],axis=0)
for batch in test_reader.sample_all_batch():
loss,_=test_eval(model,batch)
loss=test_eval(model,batch)
natom=batch["eig"].shape[1]
tst_natom_loss_list.add_loss(natom,loss)
tst_loss=tst_natom_loss_list.avg_loss()
Expand All @@ -182,12 +176,9 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
trn_natom_loss_list.print_avg_atom_loss(align_len)
tst_natom_loss_list.print_avg_atom_loss(align_len)
print('')
if vd_grad_process:
vd_processor.flush_log(epoch)
if ckpt_file:
model.save(ckpt_file)
if vd_grad_process:
vd_processor.close_log()

if ckpt_file:
model.save(ckpt_file)
if graph_file:
Expand Down
156 changes: 94 additions & 62 deletions deepks/model/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,67 +168,99 @@ def get_occ(natom):
return new_occ[natom]
return get_occ

class vd_grad_processor:
def __init__(self, vd_grad_max, momentum=0.9):
# Threshold for gradient clipping
self.grad_max = vd_grad_max

# EMA (Exponential Moving Average) momentum (0.9 means keeping 90% of history)
self.momentum = momentum
self.running_avg_grad = None

# Buffer logs in memory to avoid IO blocking during training
self.log_buffer = []
self.log_file = open("vd_grad_process_log.txt", "w")

def register_hook(self, vd_pred):
def check_nan_or_large_hook(grad):
# Prevent creating new computation graphs inside the hook
with torch.no_grad():
has_nan = torch.isnan(grad).any()
max_val = torch.max(grad).item() if not has_nan else float('nan')

if has_nan or max_val > self.grad_max:
# Prepare replacement: use historical EMA if available, else small constant
if self.running_avg_grad is not None:
replacement = self.running_avg_grad
ema_max_val = torch.max(self.running_avg_grad).item()
ema_max_str = f"{ema_max_val:.4f}"
else:
# full_like automatically handles device and dtype
replacement = torch.full_like(grad, 1e-4)
ema_max_str = "1e-4 (Default)"

# Create mask and replace outliers
mask = torch.isnan(grad) | (grad > self.grad_max)
grad = torch.where(mask, replacement, grad)

# Buffer the log message
log_msg = (f"Anomaly detected (Max: {max_val:.4f}, NaN: {has_nan}). "
f"Replaced with EMA (Max of EMA: {ema_max_str}).\n")
else:
# Normal gradient: Update EMA
if self.running_avg_grad is None:
self.running_avg_grad = grad.clone().detach()
else:
# In-place update: avg = momentum * avg + (1 - momentum) * grad
self.running_avg_grad.mul_(self.momentum).add_(grad, alpha=1 - self.momentum)

return grad

self.hook_handle = vd_pred.register_hook(check_nan_or_large_hook)
class SafeEigh(torch.autograd.Function):
"""
A custom autograd function for eigendecomposition of real symmetric matrices.
It handles degenerate eigenvalues by masking out the infinite gradients
caused by the term 1/(lambda_i - lambda_j) when lambda_i approx lambda_j.

def remove_hook(self):
self.hook_handle.remove()
Reference:
Derivatives of Partial Eigendecomposition of a Real Symmetric Matrix
for Degenerate Cases (Kasim et al., 2020), Equation (27).
"""

def flush_log(self,epoch):
"""Write buffered logs to disk. Call this periodically (e.g., end of epoch)."""
if self.log_buffer:
self.log_buffer.append(f"--- Epoch {epoch} End ---\n")
self.log_file.writelines(self.log_buffer)
self.log_buffer.clear()
self.log_file.flush()

def close_log(self):
self.flush_log()
self.log_file.close()
@staticmethod
def forward(ctx, a):
"""
Forward pass: Standard eigendecomposition.

Args:
a: Input symmetric matrix. Shape: (..., N, N), supports batching.
Returns:
e: Eigenvalues. Shape: (..., N).
v: Eigenvectors. Shape: (..., N, N).
"""
# Ensure the input is float/complex as required by eigh
# Note: 'U' (Upper) or 'L' (Lower) doesn't matter much for valid symmetric inputs
e, v = torch.linalg.eigh(a)

# Save tensors for the backward pass
ctx.save_for_backward(e, v)
return e, v

@staticmethod
def backward(ctx, grad_e, grad_v):
"""
Backward pass: Computes gradient with respect to input matrix 'a'.

This implementation specifically handles the degeneracy issue where
eigenvalues are identical or very close, which would normally cause
NaNs or Infs in the gradient of eigenvectors.
"""
e, v = ctx.saved_tensors

# 1. Handle cases where gradients might be None
# (e.g., if eigenvalues or eigenvectors are not used in the loss function)
if grad_e is None:
grad_e = torch.zeros_like(e)
if grad_v is None:
grad_v = torch.zeros_like(v)

# 2. Construct the pairwise difference matrix of eigenvalues
# Shape of e: (Batch, N)
# Use unsqueeze to broadcast: (Batch, N, 1) - (Batch, 1, N) -> (Batch, N, N)
# e_diff[..., i, j] = e[..., i] - e[..., j] (column - row)
e_diff = e.unsqueeze(-2) - e.unsqueeze(-1)

# 3. Handle Degeneracy (Masking)
# Define a small threshold to detect degeneracy
epsilon = 1e-8

# Create a mask where |lambda_i - lambda_j| > epsilon
mask = torch.abs(e_diff) > epsilon

# Construct the F matrix: F_ij = 1 / (lambda_j - lambda_i)
# Note: We use the transposed definition implicit in the matrix formula below.
# Here we initialize f_matrix with zeros, effectively ignoring degenerate terms.
f_matrix = torch.zeros_like(e_diff)

# Only compute division for non-degenerate pairs
# This prevents division by zero and corresponds to setting the gradient
# contribution of degenerate subspaces to zero (Gauge Invariance).
f_matrix[mask] = 1.0 / e_diff[mask]

# 4. Compute the gradient w.r.t. the input matrix 'a'
# Formula: grad_a = v @ (diag(grad_e) + F * (v^T @ grad_v)) @ v^T

# Projection of gradients onto the eigenvector basis: v^T @ grad_v
# transpose(-2, -1) handles the last two dimensions for batch processing
vt = v.transpose(-2, -1)
v_t_grad_v = vt @ grad_v

# The middle term: diag(grad_e) + F * (v^T @ grad_v)
# torch.diag_embed creates a diagonal matrix from the eigenvalue gradients
# f_matrix * v_t_grad_v performs element-wise multiplication (Hadamard product)
mid_term = torch.diag_embed(grad_e) + f_matrix * v_t_grad_v

# Transform back to the original basis
grad_a = v @ mid_term @ vt

# 5. Enforce symmetry
# Since the input 'a' is symmetric, its gradient must also be symmetric.
grad_a = 0.5 * (grad_a + grad_a.transpose(-2, -1))

return grad_a

# Wrapper function for easy usage
def safe_eigh(input_tensor):
return SafeEigh.apply(input_tensor)
Loading