From f0927e916c649356ff6e923f559bbae72e3fad33 Mon Sep 17 00:00:00 2001 From: crhysc Date: Fri, 24 Apr 2026 11:50:35 -0400 Subject: [PATCH 1/3] Optimize neighbor list and k-point eigensolver MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace O(N_cells·N²) dense distance allocation with matscipy cell-list neighbor detection (O(N·z̄)), then rebuild only the relevant cells differentiably from positions so autograd is preserved for forces/stress. Batch all k-points into a single eighb call instead of a Python for-loop, removing serial overhead and letting cuBLAS/LAPACK parallelize across k-points. --- slakonet/atoms.py | 72 ++++++++++++++++++++++++++++++++++++++++- slakonet/main.py | 82 +++++++++++++++++------------------------------ 2 files changed, 101 insertions(+), 53 deletions(-) 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 From 9d05a9dc8bd75fb930e7e08cf26b590660ea6bb7 Mon Sep 17 00:00:00 2001 From: Charles Campbell Date: Fri, 8 May 2026 13:23:20 -0400 Subject: [PATCH 2/3] Migrate configuration to Hydra + OmegaConf Replaces JSON config + Pydantic BaseSettings + argparse with a structured Hydra config hierarchy, making all hardcoded hyperparameters composable YAML and ensuring training/prediction entry points are unit-testable without the Hydra runtime. Key changes: - slakonet/conf.py: dataclass-based structured configs (DataConfig, TrainingConfig, OptimizerConfig, ModelConfig, PredictionConfig) that serve as the schema for both YAML files and direct OmegaConf instantiation in tests - conf/: YAML config tree with groups data/, training/, optimizer/, model/, prediction/; includes training/quick.yaml for fast debug runs - train_slakonet.py: run_training(cfg: DictConfig) callable + @hydra.main - predict_slakonet.py: run_prediction(cfg: DictConfig) callable + @hydra.main; energy_range is now a typed list, not a split string - optim.py: set_random_seed() extracted from module-level; kpoints, scheduler factor/patience, and regularization weights added as params to train_multi_vasp_skf_parameters; get_cache_dir fallback for jarvis compat - tests/conftest.py: session-scoped model fixture, minimal_train_cfg via OmegaConf.structured (no Hydra runtime or network needed) - tests/test_config.py: 6 schema/composition tests running in <0.25s - tests/test_bands.py: default_model() moved to session fixture; tests use fixtures and run_training() directly --- conf/config.yaml | 6 ++ conf/data/default.yaml | 3 + conf/data/tests.yaml | 3 + conf/model/default.yaml | 6 ++ conf/optimizer/default.yaml | 7 ++ conf/predict_config.yaml | 4 ++ conf/prediction/default.yaml | 8 +++ conf/training/default.yaml | 13 ++++ conf/training/quick.yaml | 13 ++++ requirements.txt | 3 +- setup.py | 2 + slakonet/conf.py | 81 +++++++++++++++++++++ slakonet/optim.py | 55 +++++++++----- slakonet/predict_slakonet.py | 75 +++++++------------- slakonet/tests/conftest.py | 35 +++++++++ slakonet/tests/test_bands.py | 105 +++++++-------------------- slakonet/tests/test_config.py | 87 +++++++++++++++++++++++ slakonet/train_slakonet.py | 130 +++++++++++++--------------------- 18 files changed, 407 insertions(+), 229 deletions(-) create mode 100644 conf/config.yaml create mode 100644 conf/data/default.yaml create mode 100644 conf/data/tests.yaml create mode 100644 conf/model/default.yaml create mode 100644 conf/optimizer/default.yaml create mode 100644 conf/predict_config.yaml create mode 100644 conf/prediction/default.yaml create mode 100644 conf/training/default.yaml create mode 100644 conf/training/quick.yaml create mode 100644 slakonet/conf.py create mode 100644 slakonet/tests/conftest.py create mode 100644 slakonet/tests/test_config.py diff --git a/conf/config.yaml b/conf/config.yaml new file mode 100644 index 0000000..bf6cca0 --- /dev/null +++ b/conf/config.yaml @@ -0,0 +1,6 @@ +defaults: + - data: default + - training: default + - optimizer: default + - model: default + - _self_ diff --git a/conf/data/default.yaml b/conf/data/default.yaml new file mode 100644 index 0000000..cf329f5 --- /dev/null +++ b/conf/data/default.yaml @@ -0,0 +1,3 @@ +# @package data +xml_folder_path: ??? +initial_model_path: "" diff --git a/conf/data/tests.yaml b/conf/data/tests.yaml new file mode 100644 index 0000000..67bb73f --- /dev/null +++ b/conf/data/tests.yaml @@ -0,0 +1,3 @@ +# @package data +xml_folder_path: slakonet/tests +initial_model_path: "" diff --git a/conf/model/default.yaml b/conf/model/default.yaml new file mode 100644 index 0000000..98ee46a --- /dev/null +++ b/conf/model/default.yaml @@ -0,0 +1,6 @@ +# @package model +max_Z: 100 +kT: 0.025 +alpha: 0.1 +beta: 0.1 +use_float32: true diff --git a/conf/optimizer/default.yaml b/conf/optimizer/default.yaml new file mode 100644 index 0000000..c65b5ad --- /dev/null +++ b/conf/optimizer/default.yaml @@ -0,0 +1,7 @@ +# @package optimizer +kpoints: [5, 5, 5] +scheduler_factor: 0.8 +scheduler_patience: 10 +hs_regularization: 1.0e-10 +rep_regularization: 1.0e-8 +random_seed: 42 diff --git a/conf/predict_config.yaml b/conf/predict_config.yaml new file mode 100644 index 0000000..5841414 --- /dev/null +++ b/conf/predict_config.yaml @@ -0,0 +1,4 @@ +defaults: + - prediction: default + - model: default + - _self_ diff --git a/conf/prediction/default.yaml b/conf/prediction/default.yaml new file mode 100644 index 0000000..6df1fe1 --- /dev/null +++ b/conf/prediction/default.yaml @@ -0,0 +1,8 @@ +# @package prediction +model_path: null +file_format: poscar +file_path: null +output_filename: slakonet_bands_dos.png +energy_range: [-8.0, 8.0] +jid: JVASP-107 +cutoff: 10.0 diff --git a/conf/training/default.yaml b/conf/training/default.yaml new file mode 100644 index 0000000..dd9e731 --- /dev/null +++ b/conf/training/default.yaml @@ -0,0 +1,13 @@ +# @package training +num_epochs: 100 +learning_rate: 0.00001 +batch_size: null +save_directory: out +plot_frequency: 5 +weight_by_system_size: true +early_stopping_patience: 20 +target_property: bandgap +force_weight: 0.0 +energy_weight: 1.0 +bandgap_weight: 1.0 +cutoff: 10.0 diff --git a/conf/training/quick.yaml b/conf/training/quick.yaml new file mode 100644 index 0000000..f4267fd --- /dev/null +++ b/conf/training/quick.yaml @@ -0,0 +1,13 @@ +# @package training +num_epochs: 3 +learning_rate: 0.001 +batch_size: 2 +save_directory: out +plot_frequency: 3 +weight_by_system_size: true +early_stopping_patience: 20 +target_property: bandgap +force_weight: 0.0 +energy_weight: 1.0 +bandgap_weight: 1.0 +cutoff: 10.0 diff --git a/requirements.txt b/requirements.txt index ad97c34..09eda32 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,4 +5,5 @@ jarvis-tools>=2021.07.19 torch ase spglib -pydantic_settings +hydra-core>=1.3.2 +omegaconf>=2.3.0 diff --git a/setup.py b/setup.py index 098fff4..c492333 100644 --- a/setup.py +++ b/setup.py @@ -17,6 +17,8 @@ long_description_content_type="text/markdown", url="https://github.com/atomgptlab/slakonet", packages=setuptools.find_packages(), + package_data={"slakonet": ["../conf/**/*.yaml"]}, + include_package_data=True, entry_points={ "console_scripts": [ "predict_slakonet=slakonet.predict_slakonet:main", diff --git a/slakonet/conf.py b/slakonet/conf.py new file mode 100644 index 0000000..4591232 --- /dev/null +++ b/slakonet/conf.py @@ -0,0 +1,81 @@ +from dataclasses import dataclass, field +from typing import Optional, List +from omegaconf import MISSING +from hydra.core.config_store import ConfigStore + + +@dataclass +class DataConfig: + xml_folder_path: str = MISSING + initial_model_path: str = "" + + +@dataclass +class TrainingConfig: + num_epochs: int = 100 + learning_rate: float = 1e-5 + batch_size: Optional[int] = None + save_directory: str = "out" + plot_frequency: int = 5 + weight_by_system_size: bool = True + early_stopping_patience: int = 20 + target_property: str = "bandgap" + force_weight: float = 0.0 + energy_weight: float = 1.0 + bandgap_weight: float = 1.0 + cutoff: float = 10.0 + + +@dataclass +class OptimizerConfig: + kpoints: List[int] = field(default_factory=lambda: [5, 5, 5]) + scheduler_factor: float = 0.8 + scheduler_patience: int = 10 + hs_regularization: float = 1e-10 + rep_regularization: float = 1e-8 + random_seed: int = 42 + + +@dataclass +class ModelConfig: + max_Z: int = 100 + kT: float = 0.025 + alpha: float = 0.1 + beta: float = 0.1 + use_float32: bool = True + + +@dataclass +class PredictionConfig: + model_path: Optional[str] = None + file_format: str = "poscar" + file_path: Optional[str] = None + output_filename: str = "slakonet_bands_dos.png" + energy_range: List[float] = field(default_factory=lambda: [-8.0, 8.0]) + jid: str = "JVASP-107" + cutoff: float = 10.0 + + +@dataclass +class SlakoNetTrainConfig: + data: DataConfig = field(default_factory=DataConfig) + training: TrainingConfig = field(default_factory=TrainingConfig) + optimizer: OptimizerConfig = field(default_factory=OptimizerConfig) + model: ModelConfig = field(default_factory=ModelConfig) + + +@dataclass +class SlakoNetPredictConfig: + prediction: PredictionConfig = field(default_factory=PredictionConfig) + model: ModelConfig = field(default_factory=ModelConfig) + + +def register_configs() -> None: + cs = ConfigStore.instance() + cs.store(name="train_config", node=SlakoNetTrainConfig) + cs.store(name="predict_config", node=SlakoNetPredictConfig) + cs.store(group="data", name="default", node=DataConfig) + cs.store(group="training", name="default", node=TrainingConfig) + cs.store(group="optimizer", name="default", node=OptimizerConfig) + cs.store(group="model", name="default", node=ModelConfig) + cs.store(group="prediction", name="default", node=PredictionConfig) diff --git a/slakonet/optim.py b/slakonet/optim.py index 57088a4..013bd71 100644 --- a/slakonet/optim.py +++ b/slakonet/optim.py @@ -37,28 +37,40 @@ import zipfile import requests import io -from jarvis.core.utils import get_cache_dir +try: + from jarvis.core.utils import get_cache_dir +except ImportError: + import pathlib + + def get_cache_dir(name: str) -> str: + d = pathlib.Path.home() / ".cache" / name + d.mkdir(parents=True, exist_ok=True) + return str(d) matplotlib.rcParams["figure.max_open_warning"] = 50 # torch.set_default_dtype(torch.float32) # torch.set_default_dtype(torch.float32) -random_seed = 42 -random.seed(random_seed) -torch.manual_seed(random_seed) -np.random.seed(random_seed) -torch.cuda.manual_seed_all(random_seed) -try: - import torch_xla.core.xla_model as xm +def set_random_seed(seed: int = 42) -> None: + """Set all relevant random seeds for reproducibility.""" + random.seed(seed) + torch.manual_seed(seed) + np.random.seed(seed) + torch.cuda.manual_seed_all(seed) + try: + import torch_xla.core.xla_model as xm + + xm.set_rng_state(seed) + except ImportError: + pass + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + os.environ["PYTHONHASHSEED"] = str(seed) + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" + torch.use_deterministic_algorithms(True) - xm.set_rng_state(random_seed) -except ImportError: - pass -torch.backends.cudnn.deterministic = True -torch.backends.cudnn.benchmark = False -os.environ["PYTHONHASHSEED"] = str(random_seed) -os.environ["CUBLAS_WORKSPACE_CONFIG"] = str(":4096:8") -torch.use_deterministic_algorithms(True) + +set_random_seed(42) torch.autograd.set_detect_anomaly(True) @@ -2522,6 +2534,11 @@ def train_multi_vasp_skf_parameters( bandgap_weight=1.0, device="cuda", cutoff=10.0, + kpoints=None, + scheduler_factor=0.8, + scheduler_patience=10, + hs_regularization=1e-10, + rep_regularization=1e-8, ): """ Enhanced training function for multiple VASP datasets with flexible loss targets @@ -2594,14 +2611,14 @@ def train_multi_vasp_skf_parameters( # Setup training shell_dict = generate_shell_dict_upto_Z65() - kpoints = torch.tensor([5, 5, 5]) + kpoints = torch.tensor(kpoints if kpoints is not None else [5, 5, 5]) # Setup optimizer and scheduler optimizer = optim.AdamW( multi_element_optimizer.parameters(), lr=learning_rate ) scheduler = optim.lr_scheduler.ReduceLROnPlateau( - optimizer, mode="min", factor=0.8, patience=10 + optimizer, mode="min", factor=scheduler_factor, patience=scheduler_patience ) print(f"\nStarting multi-VASP training:") @@ -2758,7 +2775,7 @@ def train_multi_vasp_skf_parameters( total_rep_reg += (opt.r_tail_coef**2).sum() regularization = ( - 1e-10 * (total_h_reg + total_s_reg) + 1e-8 * total_rep_reg + hs_regularization * (total_h_reg + total_s_reg) + rep_regularization * total_rep_reg ) # Final loss diff --git a/slakonet/predict_slakonet.py b/slakonet/predict_slakonet.py index 6be1cca..360cacc 100644 --- a/slakonet/predict_slakonet.py +++ b/slakonet/predict_slakonet.py @@ -12,52 +12,20 @@ from slakonet.main import generate_shell_dict_upto_Z65 import torch from jarvis.core.atoms import Atoms -import argparse from jarvis.io.vasp.inputs import Poscar -import argparse -import sys import time from jarvis.db.jsonutils import dumpjson import pprint +import hydra +from omegaconf import DictConfig +from hydra.utils import to_absolute_path + +from slakonet.conf import register_configs plt.rcParams.update({"font.size": 14}) H2E = 27.211 -parser = argparse.ArgumentParser(description="SlakoNet Pretrained Models") -parser.add_argument( - "--model_path", - default=None, - # default="slakonet/tests/slakonet_v1_sic", - help="Provide model path ", -) -parser.add_argument( - "--file_format", default="poscar", help="poscar/cif/xyz/pdb file format." -) -parser.add_argument( - "--file_path", - default=None, - help="Path to atomic structure file.", -) -parser.add_argument( - "--output_filename", - default="slakonet_bands_dos.png", - help="Path to desired output file name", -) -parser.add_argument( - "--energy_range", - default="-8 8", - help="Energy range for bandstructure and DOS plots", -) -parser.add_argument( - "--jid", - default="JVASP-107", - help="JARVIS-DFT Identifier", -) -parser.add_argument( - "--cutoff", - default="10", - help="Pairwise cutoff", -) +register_configs() device = "cuda" if torch.cuda.is_available() else "cpu" @@ -510,24 +478,24 @@ def plot_band_dos_atoms( return fig, properties, atom_pdos, energy_grid -# Usage -if __name__ == "__main__": - args = parser.parse_args(sys.argv[1:]) - model_path = args.model_path +def run_prediction(cfg: DictConfig): + """Execute prediction from a DictConfig. Callable without the Hydra runtime.""" + model_path = cfg.prediction.model_path model = None atoms = None if model_path is None: model = default_model() - file_path = args.file_path - file_format = args.file_format - output_filename = args.output_filename - energy_range = np.array(args.energy_range.split(" "), dtype="float") - jid = args.jid - cutoff = float(args.cutoff) + file_path = cfg.prediction.file_path + file_format = cfg.prediction.file_format + output_filename = cfg.prediction.output_filename + energy_range = np.array(list(cfg.prediction.energy_range), dtype="float") + jid = cfg.prediction.jid + cutoff = float(cfg.prediction.cutoff) if file_path is not None: + file_path = to_absolute_path(file_path) if file_format == "poscar": atoms = Atoms.from_poscar(file_path) elif file_format == "cif": @@ -541,7 +509,6 @@ def plot_band_dos_atoms( "File format not implemented", file_format ) - # fig, properties, atom_pdos, energy_grid = plot_band_dos_atoms(jid='JVASP-107') t1 = time.time() fig, properties, atom_pdos, energy_grid = plot_band_dos_atoms( atoms=atoms, @@ -554,3 +521,13 @@ def plot_band_dos_atoms( ) t2 = time.time() print("Time(s)", t2 - t1) + return fig, properties, atom_pdos, energy_grid + + +@hydra.main(config_path="../conf", config_name="predict_config", version_base=None) +def main(cfg: DictConfig) -> None: + run_prediction(cfg) + + +if __name__ == "__main__": + main() diff --git a/slakonet/tests/conftest.py b/slakonet/tests/conftest.py new file mode 100644 index 0000000..e61008e --- /dev/null +++ b/slakonet/tests/conftest.py @@ -0,0 +1,35 @@ +import pytest +import os +from omegaconf import OmegaConf + +TEST_DIR = os.path.dirname(os.path.abspath(__file__)) + + +@pytest.fixture +def minimal_train_cfg(): + """Structured train config with test-dir data — no Hydra runtime or network needed.""" + from slakonet.conf import SlakoNetTrainConfig + cfg = OmegaConf.structured(SlakoNetTrainConfig) + cfg.data.xml_folder_path = TEST_DIR + cfg.training.num_epochs = 1 + cfg.training.batch_size = 2 + cfg.training.target_property = "bandgap" + cfg.optimizer.random_seed = 0 + return cfg + + +@pytest.fixture +def minimal_predict_cfg(): + """Structured predict config pointing at a known JID — no Hydra runtime needed.""" + from slakonet.conf import SlakoNetPredictConfig + cfg = OmegaConf.structured(SlakoNetPredictConfig) + cfg.prediction.jid = "JVASP-1002" + cfg.prediction.cutoff = 10.0 + return cfg + + +@pytest.fixture(scope="session") +def default_model_fixture(): + """Session-scoped model download so tests share a single download.""" + from slakonet.optim import default_model + return default_model() diff --git a/slakonet/tests/test_bands.py b/slakonet/tests/test_bands.py index 72ef258..81057ef 100644 --- a/slakonet/tests/test_bands.py +++ b/slakonet/tests/test_bands.py @@ -1,117 +1,64 @@ +import os +import pytest +import torch from slakonet.get_bands import get_gap from slakonet.optim import ( MultiElementSkfParameterOptimizer, get_atoms, kpts_to_klines, - default_model, - multi_vasp_training ) -import os -import glob from slakonet.atoms import Geometry from slakonet.main import SimpleDftb, generate_shell_dict_upto_Z65 -import torch test_dir = os.path.dirname(os.path.abspath(__file__)) -model = default_model() -def test_basic(): - jid = "JVASP-107" +def test_basic(default_model_fixture): + model = default_model_fixture jid = "JVASP-1002" - atoms, opt_gap, mbj_gap = get_atoms( - jid - ) # Atoms.from_dict(get_jid_data(jid=jid,dataset='dft_3d')['atoms']) - # atoms=Atoms.from_poscar("tests/POSCAR") - # atoms=Atoms.from_poscar("tests/POSCAR-SiC.vasp") + atoms, opt_gap, mbj_gap = get_atoms(jid) geometry = Geometry.from_ase_atoms([atoms.ase_converter()]) - # Generate shell dictionary shell_dict = generate_shell_dict_upto_Z65() updated_skfs = model.get_updated_skfs() - - # Create comprehensive HS feeds that include ALL element pairs h_feed = model._create_comprehensive_feed(updated_skfs, shell_dict, "H") s_feed = model._create_comprehensive_feed(updated_skfs, shell_dict, "S") - - # Calculate total electron count for the system nelectron = model._calculate_system_electrons(geometry, updated_skfs) kpoints = torch.tensor([3, 3, 3]) - device = "cpu" - with_eigenvectors = True calc = SimpleDftb( geometry, model, - # shell_dict=shell_dict, kpoints=kpoints, - #h_feed=h_feed, - #s_feed=s_feed, - #nelectron=nelectron, - device=device, - with_eigenvectors=with_eigenvectors, + device="cpu", + with_eigenvectors=True, ) - - #info_ev = calc.calculate_ev_curve( - # method="polynomial", - #) - #eigenvalues = calc() calc.calculate() energy = calc.energy - #forces = calc._compute_forces_finite_diff() - # freqs,ds = calc.calculate_phonon_modes() print("energy", energy) - #print("forces", forces) - #print("eigenvalues", eigenvalues) - #print("ev info", info_ev) - import sys + assert energy is not None - # sys.exit() - """ - properties, success = model.compute_multi_element_properties( - geometry=geometry, - shell_dict=shell_dict, - kpoints=kpoints, - get_fermi=True, - get_bulk_mod=True, - get_forces=True, - device=device, - with_eigenvectors=True, - ) - """ -""" -def test_si(): - # Find the test file directory +def test_si(default_model_fixture): bandgap, opt_gap, mbj_gap, calc, formula = get_gap( - jid="JVASP-1002", model=model, plot=True + jid="JVASP-1002", model=default_model_fixture, plot=True ) - print( - "bandgap, opt_gap, mbj_gap, calc, formula", - bandgap, - opt_gap, - mbj_gap, - calc, - formula, - ) - + print("bandgap, opt_gap, mbj_gap, formula", bandgap, opt_gap, mbj_gap, formula) + assert bandgap is not None -def test_training(): - vasprun_files = [ - os.path.join(test_dir, "vasprun-107.xml"), - os.path.join(test_dir, "vasprun-1002.xml"), - os.path.join(test_dir, "vasprun-1002.xml"), - os.path.join(test_dir, "vasprun-1002.xml"), - ] - # vasprun_files = [] - # for i in glob.glob("vasprun/*.xml"): - # vasprun_files.append(i) - print("vasprun_files", vasprun_files) - multi_vasp_training(vasprun_files, model=model, batch_size=2) -""" +def test_training(minimal_train_cfg, default_model_fixture, tmp_path): + """Train for 1 epoch on the test vasprun files using the Hydra config.""" + from slakonet.train_slakonet import run_training + from omegaconf import OmegaConf -# test_basic() -# test_si() -# test_training() + cfg = OmegaConf.merge( + minimal_train_cfg, + OmegaConf.create({"training": {"save_directory": str(tmp_path / "out")}}), + ) + # Use preloaded model to avoid repeated downloads + import unittest.mock as mock + with mock.patch("slakonet.train_slakonet.default_model", return_value=default_model_fixture): + result = run_training(cfg) + assert result is not None diff --git a/slakonet/tests/test_config.py b/slakonet/tests/test_config.py new file mode 100644 index 0000000..cf98e5f --- /dev/null +++ b/slakonet/tests/test_config.py @@ -0,0 +1,87 @@ +"""Config schema and composition tests — no GPU or network required.""" +import os +import pytest +from omegaconf import OmegaConf + +from slakonet.conf import ( + SlakoNetTrainConfig, + SlakoNetPredictConfig, + TrainingConfig, + OptimizerConfig, +) + + +def test_train_config_defaults(): + cfg = OmegaConf.structured(SlakoNetTrainConfig) + assert cfg.training.cutoff == pytest.approx(10.0) + assert list(cfg.optimizer.kpoints) == [5, 5, 5] + assert cfg.optimizer.hs_regularization == pytest.approx(1e-10) + assert cfg.optimizer.rep_regularization == pytest.approx(1e-8) + assert cfg.optimizer.scheduler_factor == pytest.approx(0.8) + assert cfg.optimizer.scheduler_patience == 10 + assert cfg.optimizer.random_seed == 42 + assert cfg.model.max_Z == 100 + assert cfg.model.kT == pytest.approx(0.025) + + +def test_training_config_defaults(): + cfg = OmegaConf.structured(TrainingConfig) + assert cfg.num_epochs == 100 + assert cfg.learning_rate == pytest.approx(1e-5) + assert cfg.batch_size is None + assert cfg.target_property == "bandgap" + assert cfg.force_weight == pytest.approx(0.0) + assert cfg.energy_weight == pytest.approx(1.0) + assert cfg.bandgap_weight == pytest.approx(1.0) + + +def test_predict_config_energy_range(): + cfg = OmegaConf.structured(SlakoNetPredictConfig) + er = list(cfg.prediction.energy_range) + assert er == pytest.approx([-8.0, 8.0]) + assert cfg.prediction.file_format == "poscar" + assert cfg.prediction.cutoff == pytest.approx(10.0) + assert cfg.prediction.jid == "JVASP-107" + + +def test_config_merge_override(): + cfg = OmegaConf.structured(SlakoNetTrainConfig) + override = OmegaConf.create( + {"training": {"num_epochs": 5, "learning_rate": 0.0001}} + ) + merged = OmegaConf.merge(cfg, override) + assert merged.training.num_epochs == 5 + assert merged.training.learning_rate == pytest.approx(0.0001) + assert merged.optimizer.kpoints == [5, 5, 5] # unaffected + + +def test_config_to_yaml_roundtrip(): + cfg = OmegaConf.structured(SlakoNetTrainConfig) + # Set the required MISSING field before serializing + cfg.data.xml_folder_path = "/tmp/test" + yaml_str = OmegaConf.to_yaml(cfg) + loaded = OmegaConf.create(yaml_str) + assert loaded.training.num_epochs == cfg.training.num_epochs + assert loaded.optimizer.scheduler_factor == pytest.approx( + cfg.optimizer.scheduler_factor + ) + + +def test_hydra_yaml_compose(): + """Integration test: load real YAML files via hydra.compose.""" + from hydra import initialize_config_dir, compose + from slakonet.conf import register_configs + + conf_dir = os.path.abspath( + os.path.join(os.path.dirname(__file__), "..", "..", "conf") + ) + register_configs() + with initialize_config_dir(config_dir=conf_dir, version_base=None): + cfg = compose( + config_name="config", + overrides=["data.xml_folder_path=slakonet/tests"], + ) + assert cfg.optimizer.scheduler_factor == pytest.approx(0.8) + assert cfg.optimizer.rep_regularization == pytest.approx(1e-8) + assert cfg.training.early_stopping_patience == 20 + assert cfg.data.xml_folder_path == "slakonet/tests" diff --git a/slakonet/train_slakonet.py b/slakonet/train_slakonet.py index 558058e..60f58df 100644 --- a/slakonet/train_slakonet.py +++ b/slakonet/train_slakonet.py @@ -1,111 +1,79 @@ -from slakonet.get_bands import get_gap +import os +import hydra +from omegaconf import DictConfig, OmegaConf +from hydra.utils import to_absolute_path + +from slakonet.conf import register_configs from slakonet.optim import ( MultiElementSkfParameterOptimizer, train_multi_vasp_skf_parameters, default_model, + set_random_seed, ) -import os -import glob -import argparse -import sys -from pydantic_settings import BaseSettings -from jarvis.db.jsonutils import loadjson - - -class SlakoNetConfig(BaseSettings): - # Required paths - initial_model_path: str = "" - xml_folder_path: str = "" - - # Training parameters - num_epochs: int = 3 - batch_size: int = None - save_directory: str = "out" - learning_rate: float = 0.001 - plot_frequency: int = 5 - weight_by_system_size: bool = True - early_stopping_patience: int = 20 - - # Loss configuration - target_property: str = ( - "bandgap" # "energy", "forces", "bandgap", or "both" - ) - # target_property: str = "bandgap" # "energy", "forces", "bandgap", or "both" - force_weight: float = 0.0 - energy_weight: float = 1.0 - bandgap_weight: float = 1.0 - cutoff: float = 10.0 - H2E = 27.211 -parser = argparse.ArgumentParser(description="SlakoNet Train Model") -parser.add_argument( - "--config_name", - default="slakonet/examples/config_example.json", - help="Path of the config file", -) - +register_configs() -if __name__ == "__main__": - args = parser.parse_args(sys.argv[1:]) - config_dict = loadjson(args.config_name) - config = SlakoNetConfig(**config_dict) - # Collect XML files - xml_folder_path = config.xml_folder_path - xml_files = [] - for i in os.listdir(xml_folder_path): - if ".xml" in i: - pth = os.path.join(config.xml_folder_path, i) - xml_files.append(pth) +def run_training(cfg: DictConfig): + """Execute training from a DictConfig. Callable without the Hydra runtime.""" + set_random_seed(cfg.optimizer.random_seed) - print(f"Found {len(xml_files)} XML files in {xml_folder_path}") + xml_folder = to_absolute_path(cfg.data.xml_folder_path) + xml_files = [ + os.path.join(xml_folder, f) + for f in os.listdir(xml_folder) + if f.endswith(".xml") + ] + print(f"Found {len(xml_files)} XML files in {xml_folder}") - # Load model - model_path = config.initial_model_path - if model_path == "": + model_path = cfg.data.initial_model_path + if not model_path: print("Loading default model...") model = default_model() else: print(f"Loading model from {model_path}...") model = MultiElementSkfParameterOptimizer.load_model( - model_path, + to_absolute_path(model_path), method="compact", ) - # Print training configuration print("\n" + "=" * 70) print("TRAINING CONFIGURATION") print("=" * 70) - print(f"Model: {model_path if model_path else 'default'}") - print(f"XML files: {len(xml_files)}") - print(f"Epochs: {config.num_epochs}") - print(f"Learning rate: {config.learning_rate}") - print(f"Batch size: {config.batch_size or 'all'}") - print(f"Target property: {config.target_property}") - print( - f"Weights: energy={config.energy_weight}, bandgap={config.bandgap_weight}, force={config.force_weight}" - ) - print(f"Save directory: {config.save_directory}") + print(OmegaConf.to_yaml(cfg)) print("=" * 70 + "\n") - # Train model trained_optimizer, history, data_loader = train_multi_vasp_skf_parameters( multi_element_optimizer=model, vasprun_paths=xml_files, - num_epochs=config.num_epochs, - learning_rate=config.learning_rate, - batch_size=config.batch_size, - plot_frequency=config.plot_frequency, - save_directory=config.save_directory, - weight_by_system_size=config.weight_by_system_size, - early_stopping_patience=config.early_stopping_patience, - target_property=config.target_property, - force_weight=config.force_weight, - energy_weight=config.energy_weight, - bandgap_weight=config.bandgap_weight, - cutoff=config.cutoff, + num_epochs=cfg.training.num_epochs, + learning_rate=cfg.training.learning_rate, + batch_size=cfg.training.batch_size, + plot_frequency=cfg.training.plot_frequency, + save_directory=to_absolute_path(cfg.training.save_directory), + weight_by_system_size=cfg.training.weight_by_system_size, + early_stopping_patience=cfg.training.early_stopping_patience, + target_property=cfg.training.target_property, + force_weight=cfg.training.force_weight, + energy_weight=cfg.training.energy_weight, + bandgap_weight=cfg.training.bandgap_weight, + cutoff=cfg.training.cutoff, + kpoints=list(cfg.optimizer.kpoints), + scheduler_factor=cfg.optimizer.scheduler_factor, + scheduler_patience=cfg.optimizer.scheduler_patience, + hs_regularization=cfg.optimizer.hs_regularization, + rep_regularization=cfg.optimizer.rep_regularization, ) + return trained_optimizer, history, data_loader - print("\n✅ Training completed successfully!") + +@hydra.main(config_path="../conf", config_name="config", version_base=None) +def main(cfg: DictConfig) -> None: + run_training(cfg) + print("\nTraining completed successfully!") + + +if __name__ == "__main__": + main() From 059cbce20fd289eb3177555c63f13c8295e18e31 Mon Sep 17 00:00:00 2001 From: crhysc Date: Sat, 9 May 2026 13:30:58 -0400 Subject: [PATCH 3/3] Add pluggable eigensolver system via Hydra config Introduces a SOLID-compliant plugin architecture for the generalized eigenvalue solver used in SimpleDftb's forward pass. The solver algorithm is now a first-class Hydra config group (eigsolver: cholesky_eigh | lanczos | vqe), selectable at runtime without code changes. - slakonet/eigsolvers/: new module with _EigSolver ABC (mirrors _SkFeed), CholeskyEighSolver (the correctly-named former 'QR' solver), LanczosSolver (finds m lowest eigenvalues via Krylov subspace with normal/modified GS or selective reorthogonalization), VqeSolver stub, and make_eigsolver factory - conf/eigsolver/: three YAML config files for each variant - conf.py: CholeskyEighConfig, LanczosConfig, VqeConfig dataclasses registered as eigsolver Hydra group - SimpleDftb accepts eigsolver=None (defaults to CholeskyEighSolver for backward compatibility); _solve_eigenvalue_problem delegates to self.eigsolver.solve - Lanczos seed v_0 is detached; gradients flow through matrix-vector products with H_tilde and eigh(T) back to H and S --- conf/config.yaml | 1 + conf/eigsolver/cholesky_eigh.yaml | 3 + conf/eigsolver/lanczos.yaml | 8 ++ conf/eigsolver/vqe.yaml | 5 + slakonet/conf.py | 39 +++++++- slakonet/eigsolvers/__init__.py | 19 ++++ slakonet/eigsolvers/base.py | 74 ++++++++++++++ slakonet/eigsolvers/cholesky_eigh.py | 45 +++++++++ slakonet/eigsolvers/lanczos.py | 140 +++++++++++++++++++++++++++ slakonet/eigsolvers/vqe.py | 21 ++++ slakonet/main.py | 10 +- slakonet/optim.py | 8 +- slakonet/train_slakonet.py | 4 + 13 files changed, 374 insertions(+), 3 deletions(-) create mode 100644 conf/eigsolver/cholesky_eigh.yaml create mode 100644 conf/eigsolver/lanczos.yaml create mode 100644 conf/eigsolver/vqe.yaml create mode 100644 slakonet/eigsolvers/__init__.py create mode 100644 slakonet/eigsolvers/base.py create mode 100644 slakonet/eigsolvers/cholesky_eigh.py create mode 100644 slakonet/eigsolvers/lanczos.py create mode 100644 slakonet/eigsolvers/vqe.py diff --git a/conf/config.yaml b/conf/config.yaml index bf6cca0..ae9b8d1 100644 --- a/conf/config.yaml +++ b/conf/config.yaml @@ -3,4 +3,5 @@ defaults: - training: default - optimizer: default - model: default + - eigsolver: cholesky_eigh - _self_ diff --git a/conf/eigsolver/cholesky_eigh.yaml b/conf/eigsolver/cholesky_eigh.yaml new file mode 100644 index 0000000..a31ab9c --- /dev/null +++ b/conf/eigsolver/cholesky_eigh.yaml @@ -0,0 +1,3 @@ +# @package eigsolver +solver_name: cholesky_eigh +eps: 1.0e-8 diff --git a/conf/eigsolver/lanczos.yaml b/conf/eigsolver/lanczos.yaml new file mode 100644 index 0000000..77b95e0 --- /dev/null +++ b/conf/eigsolver/lanczos.yaml @@ -0,0 +1,8 @@ +# @package eigsolver +solver_name: lanczos +m: 20 +tol: 1.0e-6 +max_iter: 300 +reorthogonalization: true +reorthogonalization_type: modified_gram_schmidt +eps: 1.0e-8 diff --git a/conf/eigsolver/vqe.yaml b/conf/eigsolver/vqe.yaml new file mode 100644 index 0000000..a5919b2 --- /dev/null +++ b/conf/eigsolver/vqe.yaml @@ -0,0 +1,5 @@ +# @package eigsolver +solver_name: vqe +n_layers: 2 +n_shots: 1000 +eps: 1.0e-8 diff --git a/slakonet/conf.py b/slakonet/conf.py index 4591232..ee22172 100644 --- a/slakonet/conf.py +++ b/slakonet/conf.py @@ -1,9 +1,41 @@ from dataclasses import dataclass, field -from typing import Optional, List +from typing import Optional, List, Any from omegaconf import MISSING from hydra.core.config_store import ConfigStore +@dataclass +class EigsolverConfig: + """Base config for all eigensolver variants. Not used directly.""" + solver_name: str = MISSING + eps: float = 1e-8 + + +@dataclass +class CholeskyEighConfig(EigsolverConfig): + """Generalized eigenproblem via Cholesky whitening + torch.linalg.eigh.""" + solver_name: str = "cholesky_eigh" + + +@dataclass +class LanczosConfig(EigsolverConfig): + """Lanczos algorithm: finds the m lowest eigenvalues via a Krylov subspace.""" + solver_name: str = "lanczos" + m: int = 20 + tol: float = 1e-6 + max_iter: int = 300 + reorthogonalization: bool = True + reorthogonalization_type: str = "modified_gram_schmidt" + + +@dataclass +class VqeConfig(EigsolverConfig): + """Variational Quantum Eigensolver — stub pending paper implementation.""" + solver_name: str = "vqe" + n_layers: int = 2 + n_shots: int = 1000 + + @dataclass class DataConfig: xml_folder_path: str = MISSING @@ -62,12 +94,14 @@ class SlakoNetTrainConfig: training: TrainingConfig = field(default_factory=TrainingConfig) optimizer: OptimizerConfig = field(default_factory=OptimizerConfig) model: ModelConfig = field(default_factory=ModelConfig) + eigsolver: Any = field(default_factory=CholeskyEighConfig) @dataclass class SlakoNetPredictConfig: prediction: PredictionConfig = field(default_factory=PredictionConfig) model: ModelConfig = field(default_factory=ModelConfig) + eigsolver: Any = field(default_factory=CholeskyEighConfig) def register_configs() -> None: @@ -79,3 +113,6 @@ def register_configs() -> None: cs.store(group="optimizer", name="default", node=OptimizerConfig) cs.store(group="model", name="default", node=ModelConfig) cs.store(group="prediction", name="default", node=PredictionConfig) + cs.store(group="eigsolver", name="cholesky_eigh", node=CholeskyEighConfig) + cs.store(group="eigsolver", name="lanczos", node=LanczosConfig) + cs.store(group="eigsolver", name="vqe", node=VqeConfig) diff --git a/slakonet/eigsolvers/__init__.py b/slakonet/eigsolvers/__init__.py new file mode 100644 index 0000000..00598d0 --- /dev/null +++ b/slakonet/eigsolvers/__init__.py @@ -0,0 +1,19 @@ +from slakonet.eigsolvers.base import _EigSolver +from slakonet.eigsolvers.cholesky_eigh import CholeskyEighSolver +from slakonet.eigsolvers.lanczos import LanczosSolver +from slakonet.eigsolvers.vqe import VqeSolver + +__all__ = ["_EigSolver", "CholeskyEighSolver", "LanczosSolver", "VqeSolver", "make_eigsolver"] + + +def make_eigsolver(cfg) -> _EigSolver: + """Instantiate an eigensolver from a config dataclass.""" + name = cfg.solver_name + if name == "cholesky_eigh": + return CholeskyEighSolver(cfg) + elif name == "lanczos": + return LanczosSolver(cfg) + elif name == "vqe": + return VqeSolver(cfg) + else: + raise ValueError(f"Unknown eigsolver solver_name: {name!r}") diff --git a/slakonet/eigsolvers/base.py b/slakonet/eigsolvers/base.py new file mode 100644 index 0000000..aaef1bd --- /dev/null +++ b/slakonet/eigsolvers/base.py @@ -0,0 +1,74 @@ +from abc import ABC +from inspect import getfullargspec +from warnings import warn + +import torch +from torch import Tensor + + +class _EigSolver(ABC): + """ABC for objects responsible for solving the generalized eigenvalue problem. + + Subclasses solve H|ψ⟩ = E·S|ψ⟩ and return eigenvalues (and optionally + eigenvectors). The interface mirrors `_SkFeed` for consistency. + + Arguments: + device: Device on which tensors reside. + dtype: Floating point dtype used by the solver. + """ + + def __init__(self, device: torch.device, dtype: torch.dtype): + self.__device = device + self.__dtype = dtype + + def __init_subclass__(cls, check_sig: bool = True): + """Warn if subclasses' `solve` method is missing required arguments.""" + + def check(func, required_args): + sig = getfullargspec(func) + name = func.__qualname__ + if check_sig: + missing = ", ".join(required_args - set(sig.args)) + if missing: + warn( + f'Signature Warning: keyword argument(s) "{missing}"' + f' missing from method "{name}"', + stacklevel=4, + ) + if sig.varkw is None: + warn( + f'Signature Warning: method "{name}" must accept ' + f"arbitrary keyword arguments, i.e. **kwargs.", + stacklevel=4, + ) + + if hasattr(cls, "solve"): + check(cls.solve, {"H", "S"}) + + @property + def device(self) -> torch.device: + return self.__device + + @device.setter + def device(self, value): + name = self.__class__.__name__ + raise AttributeError( + f"{name} object's device can only be modified via the '.to' method." + ) + + @property + def dtype(self) -> torch.dtype: + return self.__dtype + + def solve(self, H: Tensor, S: Tensor, **kwargs) -> tuple: + """Solve H|ψ⟩ = E·S|ψ⟩. + + Arguments: + H: Hamiltonian matrix, shape [..., n, n]. + S: Overlap matrix, shape [..., n, n]. + + Returns: + eigenvalues: Shape [..., n] (or [..., m] for partial solvers). + eigenvectors: Shape [..., n, n] (or [..., n, m]), or None. + """ + raise NotImplementedError diff --git a/slakonet/eigsolvers/cholesky_eigh.py b/slakonet/eigsolvers/cholesky_eigh.py new file mode 100644 index 0000000..43725d2 --- /dev/null +++ b/slakonet/eigsolvers/cholesky_eigh.py @@ -0,0 +1,45 @@ +import torch +from torch import Tensor + +from slakonet.eigsolvers.base import _EigSolver + + +class CholeskyEighSolver(_EigSolver): + """Generalized eigensolver via Cholesky whitening + torch.linalg.eigh. + + Reduces H|ψ⟩ = E·S|ψ⟩ to a standard symmetric problem by computing + the Cholesky decomposition of S, transforming H into + H_tilde = L⁻¹ H L⁻ᵀ, and calling torch.linalg.eigh on H_tilde. + Falls back to torch.linalg.eig on Cholesky failure. + + This is O(n³) and returns all n eigenvalues/eigenvectors. + """ + + def __init__(self, cfg, device=None, dtype=None): + device = device or torch.device("cpu") + dtype = dtype or torch.float64 + super().__init__(device, dtype) + self.eps = cfg.eps + + def solve(self, H: Tensor, S: Tensor, **kwargs) -> tuple: + n = H.shape[-1] + device = H.device + dtype = H.dtype + + eye = torch.eye(n, device=device, dtype=dtype) + S_reg = S + self.eps * eye + + try: + L = torch.linalg.cholesky(S_reg) + L_inv = torch.linalg.inv(L) + H_tilde = L_inv @ H @ L_inv.mH + eigenvals, eigenvecs_tilde = torch.linalg.eigh(H_tilde) + eigenvecs = L_inv.mH @ eigenvecs_tilde + except RuntimeError as e: + print(f"Cholesky failed: {e}, falling back to eig") + eigenvals, eigenvecs = torch.linalg.eig( + torch.linalg.solve(S_reg, H) + ) + eigenvals = eigenvals.real + + return eigenvals, eigenvecs diff --git a/slakonet/eigsolvers/lanczos.py b/slakonet/eigsolvers/lanczos.py new file mode 100644 index 0000000..d851879 --- /dev/null +++ b/slakonet/eigsolvers/lanczos.py @@ -0,0 +1,140 @@ +import torch +from torch import Tensor + +from slakonet.eigsolvers.base import _EigSolver + +_REORTH_TYPES = {"normal_gram_schmidt", "modified_gram_schmidt", "selective"} + + +class LanczosSolver(_EigSolver): + """Partial eigensolver using the Lanczos algorithm. + + Finds the m lowest eigenvalues of H|ψ⟩ = E·S|ψ⟩ via a Krylov subspace + of dimension m. The generalized problem is first reduced to a standard + symmetric form using the same Cholesky preamble as CholeskyEighSolver. + + All operations use GPU-compatible PyTorch primitives only. + + Gradient flow: + The seed vector v_0 is detached before the Krylov loop — it is an + arbitrary random direction with no physical meaning. Gradients flow + from the Ritz eigenvalues through eigh(T) → tridiagonal entries + (alpha, beta) → matrix-vector products with H_tilde → H and S. + """ + + def __init__(self, cfg, device=None, dtype=None): + device = device or torch.device("cpu") + dtype = dtype or torch.float64 + super().__init__(device, dtype) + self.m = cfg.m + self.tol = cfg.tol + self.max_iter = cfg.max_iter + self.reorthogonalization = cfg.reorthogonalization + self.reorthogonalization_type = cfg.reorthogonalization_type + self.eps = cfg.eps + + if self.reorthogonalization_type not in _REORTH_TYPES: + raise ValueError( + f"reorthogonalization_type must be one of {_REORTH_TYPES}, " + f"got {self.reorthogonalization_type!r}" + ) + + def solve(self, H: Tensor, S: Tensor, **kwargs) -> tuple: + device = H.device + dtype = H.dtype + n = H.shape[-1] + batch_shape = H.shape[:-2] + m = min(self.m, n) + + eye = torch.eye(n, device=device, dtype=dtype) + S_reg = S + self.eps * eye + + L = torch.linalg.cholesky(S_reg) + L_inv = torch.linalg.inv(L) + H_tilde = L_inv @ H @ L_inv.mH + + alphas, betas, V = self._krylov(H_tilde, m, device, dtype, batch_shape, n) + + T = self._build_tridiagonal(alphas, betas, m, device, dtype, batch_shape) + eigenvalues_T, Z = torch.linalg.eigh(T) + + eigenvecs_full = V @ Z + eigenvecs = L_inv.mH @ eigenvecs_full + + return eigenvalues_T, eigenvecs + + def _krylov(self, A, m, device, dtype, batch_shape, n): + """Build the Lanczos Krylov basis and collect tridiagonal entries.""" + v = torch.randn(*batch_shape, n, device=device, dtype=dtype) + v = v.detach() + v = v / torch.linalg.norm(v, dim=-1, keepdim=True) + + alphas = [] + betas = [] + V = [] + v_prev = None + beta_prev = None + + for j in range(m): + V.append(v) + + w = (A @ v.unsqueeze(-1)).squeeze(-1) + + alpha = (v * w).sum(dim=-1) + alphas.append(alpha) + + w = w - alpha.unsqueeze(-1) * v + if j > 0: + w = w - beta_prev.unsqueeze(-1) * v_prev + + if self.reorthogonalization and j > 0: + w = self._reorthogonalize(w, V) + + beta = torch.linalg.norm(w, dim=-1) + betas.append(beta) + + v_prev = v + beta_prev = beta + v = w / beta.unsqueeze(-1).clamp(min=1e-30) + + return alphas, betas, torch.stack(V, dim=-1) + + def _reorthogonalize(self, w: Tensor, V: list) -> Tensor: + if self.reorthogonalization_type == "normal_gram_schmidt": + V_mat = torch.stack(V, dim=-1) + coeffs = (V_mat.mH @ w.unsqueeze(-1)).squeeze(-1) + w = w - (V_mat @ coeffs.unsqueeze(-1)).squeeze(-1) + + elif self.reorthogonalization_type == "modified_gram_schmidt": + for vk in V: + coeff = (vk * w).sum(dim=-1) + w = w - coeff.unsqueeze(-1) * vk + + elif self.reorthogonalization_type == "selective": + # Reorthogonalize only against vectors with small residuals, i.e. + # those whose contribution to w is above a numerical noise floor. + eps_mach = torch.finfo(w.dtype).eps + thresh = eps_mach ** 0.5 + for vk in V: + coeff = (vk * w).sum(dim=-1) + mask = coeff.abs() > thresh + if mask.any(): + w = w - (coeff * mask.to(coeff.dtype)).unsqueeze(-1) * vk + + return w + + @staticmethod + def _build_tridiagonal(alphas, betas, m, device, dtype, batch_shape): + """Assemble the m×m symmetric tridiagonal matrix T.""" + T = torch.zeros(*batch_shape, m, m, device=device, dtype=dtype) + alpha_stack = torch.stack(alphas, dim=-1) + diag_idx = torch.arange(m, device=device) + T[..., diag_idx, diag_idx] = alpha_stack + + beta_stack = torch.stack(betas[:-1], dim=-1) if m > 1 else None + if beta_stack is not None: + off_idx = torch.arange(m - 1, device=device) + T[..., off_idx, off_idx + 1] = beta_stack + T[..., off_idx + 1, off_idx] = beta_stack + + return T diff --git a/slakonet/eigsolvers/vqe.py b/slakonet/eigsolvers/vqe.py new file mode 100644 index 0000000..afa67e5 --- /dev/null +++ b/slakonet/eigsolvers/vqe.py @@ -0,0 +1,21 @@ +import torch +from torch import Tensor + +from slakonet.eigsolvers.base import _EigSolver + + +class VqeSolver(_EigSolver): + """Variational Quantum Eigensolver — stub pending paper implementation.""" + + def __init__(self, cfg, device=None, dtype=None): + device = device or torch.device("cpu") + dtype = dtype or torch.float64 + super().__init__(device, dtype) + self.n_layers = cfg.n_layers + self.n_shots = cfg.n_shots + self.eps = cfg.eps + + def solve(self, H: Tensor, S: Tensor, **kwargs) -> tuple: + raise NotImplementedError( + "VQE solver is not yet implemented; see paper in development." + ) diff --git a/slakonet/main.py b/slakonet/main.py index 288f284..71d27f3 100644 --- a/slakonet/main.py +++ b/slakonet/main.py @@ -73,6 +73,7 @@ def __init__( beta=0.1, updated_skfs=None, fermi_surface=False, + eigsolver=None, # shell_dict=None, # h_feed=None, # s_feed=None, @@ -135,6 +136,13 @@ def __init__( self.beta = beta self.fermi_surface = fermi_surface + if eigsolver is not None: + self.eigsolver = eigsolver + else: + from slakonet.eigsolvers import make_eigsolver + from slakonet.conf import CholeskyEighConfig + self.eigsolver = make_eigsolver(CholeskyEighConfig()) + def _generate_shell_dict_from_skfs(self): """ Build shell_dict entirely from the loaded SKFs — no hardcoded Z<=65 fallback. @@ -306,7 +314,7 @@ def _solve_eigenvalue_problem(self, H, S): H_b = H_b.to(torch.complex128) S_b = S_b.to(torch.complex128) - eigenvals, eigenvecs = eighb(H_b, S_b, scheme="chol") + eigenvals, eigenvecs = self.eigsolver.solve(H_b, S_b) # eigenvals: [K, ..., n_orb] eigenvecs: [K, ..., n_orb, n_orb] if self.use_float32: diff --git a/slakonet/optim.py b/slakonet/optim.py index 013bd71..c5586da 100644 --- a/slakonet/optim.py +++ b/slakonet/optim.py @@ -1521,6 +1521,7 @@ def compute_multi_element_properties( device=None, with_eigenvectors=False, cutoff=10.0, + eigsolver=None, ): """Compute DFTB properties for multi-element systems using ALL available optimizers""" if device is None: @@ -1553,6 +1554,7 @@ def compute_multi_element_properties( compute_forces=get_forces, with_eigenvectors=with_eigenvectors, cutoff=cutoff, + eigsolver=eigsolver, ) else: calc = SimpleDftb( @@ -1567,6 +1569,7 @@ def compute_multi_element_properties( compute_forces=get_forces, with_eigenvectors=with_eigenvectors, cutoff=cutoff, + eigsolver=eigsolver, ) # Compute properties @@ -2539,6 +2542,7 @@ def train_multi_vasp_skf_parameters( scheduler_patience=10, hs_regularization=1e-10, rep_regularization=1e-8, + eigsolver=None, ): """ Enhanced training function for multiple VASP datasets with flexible loss targets @@ -2679,6 +2683,7 @@ def train_multi_vasp_skf_parameters( get_forces=True, device=device, cutoff=cutoff, + eigsolver=eigsolver, ) ) print(f" Properties computed, success={success}") @@ -3301,7 +3306,7 @@ def default_model_old(dir_path=None, model_name="slakonet_v0"): return model -def get_hamiltonian(jarvis_atoms=None, kpts=[1, 4, 4], model=None): +def get_hamiltonian(jarvis_atoms=None, kpts=[1, 4, 4], model=None, eigsolver=None): geometry = Geometry.from_ase_atoms([jarvis_atoms.ase_converter()]) @@ -3313,6 +3318,7 @@ def get_hamiltonian(jarvis_atoms=None, kpts=[1, 4, 4], model=None): model=model, compute_forces=False, include_dos_data=False, + eigsolver=eigsolver, ) calc.calculate() diff --git a/slakonet/train_slakonet.py b/slakonet/train_slakonet.py index 60f58df..57b6088 100644 --- a/slakonet/train_slakonet.py +++ b/slakonet/train_slakonet.py @@ -20,6 +20,9 @@ def run_training(cfg: DictConfig): """Execute training from a DictConfig. Callable without the Hydra runtime.""" set_random_seed(cfg.optimizer.random_seed) + from slakonet.eigsolvers import make_eigsolver + eigsolver = make_eigsolver(cfg.eigsolver) + xml_folder = to_absolute_path(cfg.data.xml_folder_path) xml_files = [ os.path.join(xml_folder, f) @@ -65,6 +68,7 @@ def run_training(cfg: DictConfig): scheduler_patience=cfg.optimizer.scheduler_patience, hs_regularization=cfg.optimizer.hs_regularization, rep_regularization=cfg.optimizer.rep_regularization, + eigsolver=eigsolver, ) return trained_optimizer, history, data_loader