diff --git a/slakonet/atoms.py b/slakonet/atoms.py index 4910265..841e26c 100644 --- a/slakonet/atoms.py +++ b/slakonet/atoms.py @@ -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() ) @@ -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 diff --git a/slakonet/main.py b/slakonet/main.py index 1b30540..288f284 100644 --- a/slakonet/main.py +++ b/slakonet/main.py @@ -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: for each eigenvector - if eigenvecs.is_complex(): - # norms[i] = sqrt() - 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() - 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