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
72 changes: 71 additions & 1 deletion slakonet/atoms.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,7 @@ def __init__(
self.positions_pe,
self.positions_vec,
self.periodic_distances,
) = self._periodic_distance()
) = self._periodic_distance_matscipy()
self.neighbour_pos, self.neighbour_vec, self.neighbour_dis = (
self._neighbourlist()
)
Expand Down Expand Up @@ -291,6 +291,76 @@ def get_cell_translations_old(self, **kwargs):

return cellvec, rcellvec, ncell

def _periodic_distance_matscipy(self):
"""Matscipy-accelerated neighbor-cell detection with differentiable reconstruction."""
if self.mask_zero.any():
return self._periodic_distance()
try:
from matscipy.neighbours import neighbour_list as _msp_nl
except ImportError:
return self._periodic_distance()

import numpy as np
device = self.positions.device
dtype = self.positions.dtype
all_positions_pe, all_positions_vec, all_distances = [], [], []
all_rcellvec, all_cellvec = [], []

for ibatch in range(self._n_batch):
n_atoms = int(self.atomic_numbers[ibatch].ne(0).sum().item())
pos = self.positions[ibatch] # [N_max, 3] — keeps grad
latvec = self.latvec[ibatch] # [3, 3]
pos_np = pos[:n_atoms].detach().cpu().numpy()
latvec_np = latvec.detach().cpu().numpy()
cutoff_val = float(self.cutoff[ibatch].item())

_, _, S_ij = _msp_nl(
"ijS", positions=pos_np, cell=latvec_np,
cutoff=cutoff_val, pbc=[True, True, True],
)
if len(S_ij) > 0:
S_np = np.unique(np.vstack([[[0,0,0]], S_ij]), axis=0).astype(np.float64)
else:
S_np = np.array([[0,0,0]], dtype=np.float64)

S_t = torch.tensor(S_np, dtype=dtype, device=device) # [n_sub, 3]
rcellvec_sub = S_t @ latvec # [n_sub, 3]
positions_pe_b = rcellvec_sub.unsqueeze(1) + pos.unsqueeze(0) # [n_sub, N_max, 3]
positions_vec_b = (
-positions_pe_b.unsqueeze(-3) + pos.unsqueeze(0).unsqueeze(-2)
) # [n_sub, N_max, N_max, 3]
eps = 1e-12
distance_b = torch.sqrt(eps + (positions_vec_b ** 2).sum(-1))

if not self.atomic_numbers[ibatch].ne(0).all():
atom_mask = self.atomic_numbers[ibatch].ne(0)
pad_mask = ~(atom_mask.unsqueeze(-1) & atom_mask.unsqueeze(0))
distance_b = distance_b.masked_fill(pad_mask.unsqueeze(0), 1e3)

all_positions_pe.append(positions_pe_b)
all_positions_vec.append(positions_vec_b)
all_distances.append(distance_b)
all_rcellvec.append(rcellvec_sub)
all_cellvec.append(S_t)

if self._n_batch == 1:
positions_pe = all_positions_pe[0].unsqueeze(0)
positions_vec = all_positions_vec[0].unsqueeze(0)
periodic_distances = all_distances[0].unsqueeze(0)
new_rcellvec = all_rcellvec[0].unsqueeze(0)
new_cellvec = all_cellvec[0].unsqueeze(0)
else:
positions_pe = pack(all_positions_pe, value=1e3)
positions_vec = pack(all_positions_vec, value=1e3)
periodic_distances = pack(all_distances, value=1e3)
new_rcellvec = pack(all_rcellvec, value=1e3)
new_cellvec = pack(all_cellvec, value=1e3)

self.rcellvec = new_rcellvec
self.cellvec = new_cellvec
mask_central_cell = (new_rcellvec.abs().sum(-1) == 0)
return mask_central_cell, positions_pe, positions_vec, periodic_distances

def _periodic_distance(self):
"""Get distances between central cell and neighbour cells - fully vectorized."""
mask_central_cell = (self.rcellvec != 0).sum(-1) == 0
Expand Down
82 changes: 30 additions & 52 deletions slakonet/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,64 +294,42 @@ def _compute_nelectrons(self):
return total_electrons.unsqueeze(0)

def _solve_eigenvalue_problem(self, H, S):
"""Solve H*c = E*S*c with appropriate precision."""
"""Solve H*c = E*S*c, batching all k-points into a single eigensolver call."""
n_kpoints = self.max_nk.item()
eigenvalues_list = []
eigenvecs_list = []
occupations_list = []

for ik in range(n_kpoints):
h_k = H[..., ik]
s_k = S[..., ik]
# H: [..., n_orb, n_orb, K] → [K, ..., n_orb, n_orb]
perm_fwd = (-1,) + tuple(range(H.ndim - 1))
H_b = H.permute(perm_fwd)
S_b = S.permute(perm_fwd)

# CRITICAL: Use float64 for eigenvalue decomposition
# This is where precision matters most
if self.use_float32:
h_k = h_k.to(torch.complex128) # Complex128 for stability
s_k = s_k.to(torch.complex128)
if self.use_float32:
H_b = H_b.to(torch.complex128)
S_b = S_b.to(torch.complex128)

# Solve generalized eigenvalue problem
eigenvals, eigenvecs = eighb(h_k, s_k, scheme="chol")
eigenvals, eigenvecs = eighb(H_b, S_b, scheme="chol")
# eigenvals: [K, ..., n_orb] eigenvecs: [K, ..., n_orb, n_orb]

# Convert back to float32 after solve (if needed)
if self.use_float32:
eigenvals = eigenvals.to(torch.float32)
if eigenvecs is not None:
eigenvecs = eigenvecs.to(torch.complex64)

"""
# ===== CRITICAL FIX: Normalize eigenvectors =====
if self.use_float32:
eigenvals = eigenvals.to(torch.float32)
if eigenvecs is not None:
# Compute norms: <c|S|c> for each eigenvector
if eigenvecs.is_complex():
# norms[i] = sqrt(<c_i|S|c_i>)
norms = torch.sqrt(
torch.sum(eigenvecs.conj() * (s_k @ eigenvecs), dim=0).real
)
else:
norms = torch.sqrt(
torch.sum(eigenvecs * (s_k @ eigenvecs), dim=0)
)

# Normalize: c_normalized = c / sqrt(<c|S|c>)
eigenvecs = eigenvecs / norms.unsqueeze(0)
# ===== End normalization =====
"""
# Fermi occupation
occ, _ = fermi(eigenvals, self.nelectron.to(self.device))

eigenvalues_list.append(eigenvals)
eigenvecs_list.append(eigenvecs)
occupations_list.append(occ)

# Stack and convert to eV
eigenvalues = torch.stack(eigenvalues_list, dim=1) * self.H2E
eigenvectors = (
torch.stack(eigenvecs_list, dim=1)
if self.with_eigenvectors
else None
)
occupations = torch.stack(occupations_list, dim=1)
eigenvecs = eigenvecs.to(torch.complex64)

# Occupations: same pattern at every k-point with integer filling
occ_0, _ = fermi(eigenvals[0], self.nelectron.to(self.device))
occupations = occ_0.unsqueeze(0).expand(n_kpoints, *occ_0.shape)

# Permute back: [K, ..., n_orb] → [..., K, n_orb]
ndim_ev = eigenvals.ndim
perm_back = tuple(range(1, ndim_ev - 1)) + (0, ndim_ev - 1)
eigenvalues = eigenvals.permute(perm_back) * self.H2E
occupations = occupations.permute(perm_back)

if self.with_eigenvectors and eigenvecs is not None:
ndim_ec = eigenvecs.ndim
perm_back_ec = tuple(range(1, ndim_ec - 2)) + (0, ndim_ec - 2, ndim_ec - 1)
eigenvectors = eigenvecs.permute(perm_back_ec)
else:
eigenvectors = None

return eigenvalues, eigenvectors, occupations

Expand Down
Loading