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
3 changes: 2 additions & 1 deletion deepks/model/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,7 @@ 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 @@ -246,7 +247,7 @@ def __call__(self, model, sample):
tot_loss = tot_loss + d_loss
loss.append(d_loss)
loss.append(tot_loss)
return loss
return loss, vd_pred

def print_head(self,name,data_keys,align_len=20):
info=f"{name}_energy".rjust(align_len)
Expand Down
24 changes: 17 additions & 7 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
from deepks.model.utils import preprocess, fit_elem_const, make_loss, vd_grad_processor
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,
energy_per_atom=0, vd_divide_by_nlocal=False, vd_grad_process=False, vd_grad_max=1.0, vd_grad_momentum=0.9,
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 @@ -98,12 +98,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,6 +134,9 @@ 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 @@ -142,9 +145,13 @@ def train(model, g_reader, n_epoch=1000, test_reader=None, *,
for sample in g_reader:
model.train()
optimizer.zero_grad()
loss = evaluator(model, sample)
loss, vd_pred = evaluator(model, sample)
if vd_grad_process and vd_pred is not None:
vd_processor.register_hook(vd_pred)
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 @@ -159,7 +166,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 @@ -175,9 +182,12 @@ 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
65 changes: 65 additions & 0 deletions deepks/model/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,3 +167,68 @@ def get_occ(natom):
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)

def remove_hook(self):
self.hook_handle.remove()

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()
Loading