diff --git a/deepks/model/evaluator.py b/deepks/model/evaluator.py index 9cd67d7f..906bcfce 100644 --- a/deepks/model/evaluator.py +++ b/deepks/model/evaluator.py @@ -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: @@ -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 = {} @@ -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 @@ -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) \ @@ -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"] @@ -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) diff --git a/deepks/model/reader.py b/deepks/model/reader.py index bbd5ed32..68634cce 100644 --- a/deepks/model/reader.py +++ b/deepks/model/reader.py @@ -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() @@ -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 diff --git a/deepks/model/train.py b/deepks/model/train.py index 85fc4f03..d25219cb 100644 --- a/deepks/model/train.py +++ b/deepks/model/train.py @@ -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", @@ -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 @@ -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() @@ -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 = [] @@ -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) @@ -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() @@ -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: diff --git a/deepks/model/utils.py b/deepks/model/utils.py index 31550105..6b6fc460 100644 --- a/deepks/model/utils.py +++ b/deepks/model/utils.py @@ -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() \ No newline at end of file + @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) \ No newline at end of file