Skip to content
Open
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
7 changes: 7 additions & 0 deletions conf/config.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
defaults:
- data: default
- training: default
- optimizer: default
- model: default
- eigsolver: cholesky_eigh
- _self_
3 changes: 3 additions & 0 deletions conf/data/default.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
# @package data
xml_folder_path: ???
initial_model_path: ""
3 changes: 3 additions & 0 deletions conf/data/tests.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
# @package data
xml_folder_path: slakonet/tests
initial_model_path: ""
3 changes: 3 additions & 0 deletions conf/eigsolver/cholesky_eigh.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
# @package eigsolver
solver_name: cholesky_eigh
eps: 1.0e-8
8 changes: 8 additions & 0 deletions conf/eigsolver/lanczos.yaml
Original file line number Diff line number Diff line change
@@ -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
5 changes: 5 additions & 0 deletions conf/eigsolver/vqe.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
# @package eigsolver
solver_name: vqe
n_layers: 2
n_shots: 1000
eps: 1.0e-8
6 changes: 6 additions & 0 deletions conf/model/default.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
# @package model
max_Z: 100
kT: 0.025
alpha: 0.1
beta: 0.1
use_float32: true
7 changes: 7 additions & 0 deletions conf/optimizer/default.yaml
Original file line number Diff line number Diff line change
@@ -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
4 changes: 4 additions & 0 deletions conf/predict_config.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
defaults:
- prediction: default
- model: default
- _self_
8 changes: 8 additions & 0 deletions conf/prediction/default.yaml
Original file line number Diff line number Diff line change
@@ -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
13 changes: 13 additions & 0 deletions conf/training/default.yaml
Original file line number Diff line number Diff line change
@@ -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
13 changes: 13 additions & 0 deletions conf/training/quick.yaml
Original file line number Diff line number Diff line change
@@ -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
3 changes: 2 additions & 1 deletion requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -5,4 +5,5 @@ jarvis-tools>=2021.07.19
torch
ase
spglib
pydantic_settings
hydra-core>=1.3.2
omegaconf>=2.3.0
2 changes: 2 additions & 0 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
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
118 changes: 118 additions & 0 deletions slakonet/conf.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
from dataclasses import dataclass, field
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
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)
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:
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)
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)
19 changes: 19 additions & 0 deletions slakonet/eigsolvers/__init__.py
Original file line number Diff line number Diff line change
@@ -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}")
Loading
Loading