From 766234a06d55784d2f5e32e8310e449321b51ce1 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 31 Jul 2026 02:34:17 +0200 Subject: [PATCH 01/39] implement screening --- environment-cpu.yml | 2 + environment-gpu.yml | 2 + pyproject.toml | 9 + src/skala/pyscf/features.py | 189 ++++++++--- src/skala/pyscf/numint.py | 28 +- tests/test_ao_screening.py | 459 +++++++++++++++++++++++++++ tests/test_ao_screening_benchmark.py | 238 ++++++++++++++ tests/test_gpu4pyscf_ao_screening.py | 170 ++++++++++ 8 files changed, 1052 insertions(+), 45 deletions(-) create mode 100644 tests/test_ao_screening.py create mode 100644 tests/test_ao_screening_benchmark.py create mode 100644 tests/test_gpu4pyscf_ao_screening.py diff --git a/environment-cpu.yml b/environment-cpu.yml index 0d415357..737cdf5c 100644 --- a/environment-cpu.yml +++ b/environment-cpu.yml @@ -17,8 +17,10 @@ dependencies: # Testing and development - pre-commit - pytest + - pytest-benchmark - pytest-cov - pytest-randomly + - pytest-timeout - ruff - mypy - pip: diff --git a/environment-gpu.yml b/environment-gpu.yml index 2abb6ad1..ab454fca 100644 --- a/environment-gpu.yml +++ b/environment-gpu.yml @@ -20,8 +20,10 @@ dependencies: # Testing and development - pre-commit - pytest + - pytest-benchmark - pytest-cov - pytest-randomly + - pytest-timeout - ruff - mypy - pip: diff --git a/pyproject.toml b/pyproject.toml index bfb8fc26..b1a40952 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,8 +33,10 @@ optional-dependencies.dev = [ "pre-commit", "mypy", "pytest", + "pytest-benchmark", "pytest-cov", "pytest-randomly", + "pytest-timeout", ] optional-dependencies.doc = [ "ipywidgets", @@ -94,6 +96,11 @@ detect-same-package = true line-length = 100 [tool.pytest.ini_options] +timeout = 300 +markers = [ + "benchmark: performance measurements collected by pytest-benchmark", + "profiling: single-call performance workloads intended for profilers", +] filterwarnings = [ "error", # PyTorch 2.11 deprecated `torch.jit.load`; Skala's pretrained checkpoints @@ -116,6 +123,8 @@ filterwarnings = [ 'ignore:using cupy as the tensor contraction engine\.:UserWarning', # Deprecation warning in huggingface_hub package 'ignore:hf_xet\.download_files\(\) is deprecated\. Use XetSession\(\)\.new_file_download_group\(\)\.start_download_file\(\) instead\.:DeprecationWarning', + # ASE directly assigns array shapes in `Atoms.new_array`, deprecated by NumPy 2.5. + 'ignore:Setting the shape on a NumPy array has been deprecated in NumPy 2\.5\.:DeprecationWarning', # Upstream deprecations in PySCF / GPU4PySCF are outside this project's control. 'ignore::DeprecationWarning:pyscf\..*', 'ignore::DeprecationWarning:gpu4pyscf\..*', diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 2c5bfc11..b53f9f4c 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -29,6 +29,7 @@ DEFAULT_FEATURES = ["density", "kin", "grad", "grid_coords", "grid_weights"] DEFAULT_FEATURES_SET = set(DEFAULT_FEATURES) +CPU_AO_SCREENING_BLOCK_SIZE = 9 * dft.gen_grid.BLKSIZE # Features that require per-atom grid decomposition. _ATOMIC_GRID_FEATURES = { @@ -38,6 +39,25 @@ } +def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: + active_shells = np.any(screen_index, axis=0) + ao_loc = mol.ao_loc_nr() + return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) + + +def _spatially_group_atom_grids( + mol: gto.Mole, coords: np.ndarray, atomic_grid_sizes: Tensor +) -> np.ndarray: + sort_indices = [] + start = 0 + for size in atomic_grid_sizes.tolist(): + stop = start + size + atom_sort_indices = dft.gen_grid.arg_group_grids(mol, coords[start:stop]) + sort_indices.append(start + atom_sort_indices) + start = stop + return np.concatenate(sort_indices) + + def maybe_expand_and_divide( feature: torch.Tensor, expand: bool, divisor: float ) -> torch.Tensor: @@ -59,6 +79,7 @@ def chunked_features( max_memory_in_mb: int | None = None, safety_fraction: float = 0.8, compile_feature_function: bool = False, + screen_aos: bool = False, ) -> Iterator[dict[str, Tensor]]: """ Chunked feature generation for a given molecule. The density features are generated in chunks to avoid memory issues. @@ -72,6 +93,7 @@ def chunked_features( max_memory_in_mb: The maximum memory to use for each chunk in megabytes (MB). If None, the maximum memory is determined automatically. safety_fraction: The fraction of the available memory to use for each chunk. compile_feature_function: Whether to compile the feature function. + screen_aos: Whether to evaluate each atom chunk through backend AO screening. Yields: A dictionary of features for each chunk. @@ -82,6 +104,8 @@ def chunked_features( raise ValueError( "The current implementation of chunked_features requires 'atomic_grid_sizes' to be in the requested features." ) + if grids.coords is None or grids.weights is None: + raise ValueError("Grids must be built before generating chunked features.") # if dm is a 3D tensor, then we have a spin-polarized system with_spin = True if len(dm.shape) == 3 else False @@ -143,16 +167,60 @@ def chunked_features( if with_mgga_feature: assert ff is not None - feat_tensor = non_chunk( - dm.double(), - mol, - grids.coords[grid_slice], - ff, - compile_feature_function=compile_feature_function, - gpu=dm.device.type == "cuda", - ) + gpu = dm.device.type == "cuda" + if screen_aos: + chunk_grids = copy(grids) + chunk_grids.coords = grids.coords[grid_slice] + chunk_grids.weights = grids.weights[grid_slice] + if gpu: + chunk_grids._non0ao_idx = None + else: + grid_sort_indices = _spatially_group_atom_grids( + mol, + chunk_grids.coords, + feature_chunk["atomic_grid_sizes"], + ) + chunk_grids.coords = chunk_grids.coords[grid_sort_indices] + chunk_grids.weights = chunk_grids.weights[grid_sort_indices] + grid_sort_indices_t = torch.as_tensor( + grid_sort_indices, device=dm.device + ) + for feat_name in ( + "grid_coords", + "grid_weights", + "atomic_grid_weights", + ): + if feat_name in feature_chunk: + feature_chunk[feat_name] = feature_chunk[feat_name][ + grid_sort_indices_t + ] + chunk_grids.non0tab = dft.gen_grid.make_screen_index( + mol, + chunk_grids.coords, + cutoff=chunk_grids.cutoff, + ) + feat_tensor = ChunkEvalForward.apply( + dm.double(), + mol, + chunk_grids, + ff, + None if gpu else CPU_AO_SCREENING_BLOCK_SIZE, + compile_feature_function, + gpu, + ) + mgga_features = ff.to_dict(feat_tensor) + else: + feat_tensor = non_chunk( + dm.double(), + mol, + grids.coords[grid_slice], + ff, + compile_feature_function=compile_feature_function, + gpu=gpu, + ) + mgga_features = ff.to_dict(feat_tensor) - for k, v in ff.to_dict(feat_tensor).items(): + for k, v in mgga_features.items(): feature_chunk[k] = maybe_expand_and_divide(v, not with_spin, 2) yield feature_chunk @@ -601,7 +669,7 @@ def setup_context( gto.Mole, Grid, FeatureFunction, - int, + int | None, int, bool, bool, @@ -627,7 +695,7 @@ def forward( mol: gto.Mole, grids: Grid, feature_function: FeatureFunction, - blksize: int, + blksize: int | None, compile_feature_function: bool, gpu: bool, *vectors_jvp: torch.Tensor, @@ -637,6 +705,7 @@ def forward( block_loop_kwargs = { "deriv": feature_function.deriv, "blksize": blksize if not gpu else None, + "non0tab": None if gpu else getattr(grids, "non0tab", None), } if gpu: check_gpu_imports_were_successful() @@ -669,24 +738,41 @@ def forward( *block_loop_args, **block_loop_kwargs ): start, end = end, end + weights.size - # Mask dm to only include the relevant AOs - if mask is None or not gpu: - mask = torch.arange(mol.nao_nr(), device=dm.device) - else: - mask = torch.from_dlpack(mask) - masked_dm = dm_sorted[..., mask[:, None], mask[None, :]] + ao = from_numpy_or_cupy( + ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu + ) + if gpu and mask is not None: + mask = from_numpy_or_cupy(mask, device=dm.device, dtype=torch.long) + elif not gpu and mask is not None: + num_screen_rows = ( + weights.size + dft.gen_grid.BLKSIZE - 1 + ) // dft.gen_grid.BLKSIZE + mask = mask[:num_screen_rows] + mask = torch.as_tensor( + _active_cpu_aos(mol, mask), device=dm.device, dtype=torch.long + ) + ao = ao[..., mask, :] + if mask is not None and mask.numel() == 0: + continue + masked_dm = ( + dm_sorted + if mask is None + else dm_sorted[..., mask[:, None], mask[None, :]] + ) # Apply chain rule for this particular block partial_func = partial_feature_function_over_aos( feature_function, - from_numpy_or_cupy( - ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu - ), + ao, ) for v_sorted in vectors_jvp_sorted: partial_func = partial_jvp_function_over_tangents( partial_func, - v_sorted[..., mask[:, None], mask[None, :]], + ( + v_sorted + if mask is None + else v_sorted[..., mask[:, None], mask[None, :]] + ), ) # Compute feature (or its jvp) for this block with masked dm @@ -699,7 +785,7 @@ def forward( return features @staticmethod - def jvp(ctx: FunctionCtx, grad_input: torch.Tensor) -> torch.Tensor: + def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor) -> torch.Tensor: # Chain rule for the jvp return ChunkEvalForward.apply( ctx.dm, @@ -710,7 +796,7 @@ def jvp(ctx: FunctionCtx, grad_input: torch.Tensor) -> torch.Tensor: ctx.compile_feature_function, ctx.gpu, *ctx.vectors_jvp, - grad_input, + grad_inputs[0], ) @staticmethod @@ -776,7 +862,7 @@ def setup_context( Grid, FeatureFunction, list[str], - int, + int | None, bool, bool, torch.Tensor, @@ -803,7 +889,7 @@ def forward( grids: Grid, feature_function: FeatureFunction, derivative_types: list[str], - blksize: int, + blksize: int | None, compile_feature_function: bool, gpu: bool, *vectors: torch.Tensor, @@ -812,6 +898,7 @@ def forward( block_loop_kwargs = { "deriv": feature_function.deriv, "blksize": blksize if not gpu else None, + "non0tab": None if gpu else getattr(grids, "non0tab", None), } if gpu: check_gpu_imports_were_successful() @@ -842,19 +929,28 @@ def forward( ): start, end = end, end + weights.size - # Mask to only include the relevant AOs - if mask is None or not gpu: - mask = torch.arange(mol.nao_nr(), device=dm.device) - else: + ao = from_numpy_or_cupy( + ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu + ) + if gpu and mask is not None: mask = from_numpy_or_cupy(mask, device=dm.device, dtype=torch.long) + elif not gpu and mask is not None: + num_screen_rows = ( + weights.size + dft.gen_grid.BLKSIZE - 1 + ) // dft.gen_grid.BLKSIZE + mask = mask[:num_screen_rows] + mask = torch.as_tensor( + _active_cpu_aos(mol, mask), device=dm.device, dtype=torch.long + ) + ao = ao[..., mask, :] + if mask is not None and mask.numel() == 0: + continue # Apply chain rule for this particular block # but be careful with signature change upon first vjp partial_func = partial_feature_function_over_aos( feature_function, - from_numpy_or_cupy( - ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu - ), + ao, ) for derivative_type, vector, v_sorted in zip( derivative_types, vectors, vectors_sorted, strict=True @@ -862,12 +958,20 @@ def forward( if derivative_type == "jvp": partial_func = partial_jvp_function_over_tangents( partial_func, - v_sorted[..., mask[:, None], mask[None, :]], + ( + v_sorted + if mask is None + else v_sorted[..., mask[:, None], mask[None, :]] + ), ) elif derivative_type == "vjp": partial_func = partial_vjp_function_over_tangents( partial_func, - v_sorted[..., mask[:, None], mask[None, :]], + ( + v_sorted + if mask is None + else v_sorted[..., mask[:, None], mask[None, :]] + ), ) elif derivative_type == "first_vjp": partial_func = partial_vjp_function_over_tangents( @@ -877,14 +981,19 @@ def forward( raise ValueError( f"Unknown derivative {derivative_type} (must be one of 'jvp', 'vjp', 'first_vjp')" ) + masked_dm = ( + dm_sorted + if mask is None + else dm_sorted[..., mask[:, None], mask[None, :]] + ) if compile_feature_function: - out[..., mask[:, None], mask[None, :]] += torch.compile(partial_func)( - dm_sorted[..., mask[:, None], mask[None, :]] - ) + block_result = torch.compile(partial_func)(masked_dm) else: - out[..., mask[:, None], mask[None, :]] += partial_func( - dm_sorted[..., mask[:, None], mask[None, :]] - ) + block_result = partial_func(masked_dm) + if mask is None: + out += block_result + else: + out[..., mask[:, None], mask[None, :]] += block_result return out[..., unsort_idx, :][..., unsort_idx] @staticmethod diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 558f4d22..8bb81fe0 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -1,10 +1,11 @@ # SPDX-License-Identifier: MIT from collections.abc import Callable -from typing import Any, Generic, Protocol +from typing import Any, Generic, Protocol, overload import torch from pyscf import dft, gto +from pyscf.dft import numint as pyscf_numint from torch import Tensor from skala.functional.base import ExcFunctionalBase @@ -20,6 +21,13 @@ from skala.pyscf.features import chunked_features, generate_features +def _should_screen_aos(mol: gto.Mole) -> bool: + # Keep the compatibility fallback here, not at call sites. PySCF uses this + # crossover before selecting sparse density/Vxc contractions. + switch_size = pyscf_numint.SWITCH_SIZE + return mol.nao_nr() > switch_size + + class LibXCSpec(Protocol): __version__: str | None __references__: str | None @@ -134,9 +142,15 @@ def from_backend( ) -> Tensor: return from_numpy_or_cupy(x, device=device or self.device, transpose=transpose) + @overload + def to_backend(self, x: Tensor) -> Array: ... + + @overload + def to_backend(self, x: list[Tensor]) -> list[Array]: ... + def to_backend(self, x: Tensor | list[Tensor]) -> Array | list[Array]: if isinstance(x, list): - return [self.to_backend(y) for y in x] # type: ignore + return [self.to_backend(y) for y in x] if self.device.type == "cuda": return to_cupy(x) @@ -160,7 +174,7 @@ def get_rho( max_memory=max_memory, gpu=self.device.type == "cuda", ) - return self.to_backend(mol_features["density"].sum(0)) # type: ignore + return self.to_backend(mol_features["density"].sum(0)) def __call__( self, @@ -192,6 +206,7 @@ def __call__( if self._functional_supports_atom_chunking(): dm = dm.detach().requires_grad_() + screen_aos = _should_screen_aos(mol) tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) V_xc = torch.zeros_like(dm) @@ -203,6 +218,7 @@ def __call__( func_deriv=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, safety_fraction=0.8, # tends to be faster for large chunks + screen_aos=screen_aos, ): E_xc_chunk = self.func.get_exc(mol_features) (V_xc_chunk,) = torch.autograd.grad( @@ -260,7 +276,7 @@ def nr_rks( N, E_xc, V_xc = self( mol, grids, xc_code, self.from_backend(dm), max_memory=max_memory ) - return N.sum().item(), E_xc.item(), self.to_backend(V_xc) # type: ignore + return N.sum().item(), E_xc.item(), self.to_backend(V_xc) def nr_uks( self, @@ -275,7 +291,7 @@ def nr_uks( N, E_xc, V_xc = self( mol, grids, xc_code, self.from_backend(dm), max_memory=max_memory ) - return self.to_backend(N), E_xc.item(), self.to_backend(V_xc) # type: ignore + return self.to_backend(N), E_xc.item(), self.to_backend(V_xc) class libxc: __version__ = None @@ -314,6 +330,7 @@ def gen_response( if self._functional_supports_atom_chunking(): dm0 = dm0.requires_grad_() + screen_aos = _should_screen_aos(ks.mol) def hessian_vector_product_atom_chunked(dm1: Array) -> Array: dm1_tensor = self.from_backend(dm1) @@ -330,6 +347,7 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: safety_fraction=kwargs.get( "safety_fraction", 0.0 ), # Force small chunks (single atoms) because it's empirically fastest. + screen_aos=screen_aos, ): E_xc_chunk = self.func.get_exc(mol_features) (V_xc_chunk,) = torch.autograd.grad( diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py new file mode 100644 index 00000000..b89ae032 --- /dev/null +++ b/tests/test_ao_screening.py @@ -0,0 +1,459 @@ +from collections.abc import Callable, Iterator + +import numpy as np +import pytest +import torch +from pyscf import dft, gto +from pyscf.dft import numint as pyscf_numint + +from skala.functional.base import ExcFunctionalBase +from skala.pyscf import features as features_module +from skala.pyscf import numint as numint_module +from skala.pyscf.features import ( + CPU_AO_SCREENING_BLOCK_SIZE, + ChunkEvalForward, + MGGAFeatureFunction, + _active_cpu_aos, + chunked_features, +) +from skala.pyscf.numint import SkalaNumInt, _should_screen_aos + + +@pytest.fixture +def carbon() -> gto.Mole: + return gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0) + + +@pytest.mark.parametrize( + ("switch_offset", "expected"), [(1, False), (0, False), (-1, True)] +) +def test_should_screen_aos_at_crossover( + carbon: gto.Mole, + monkeypatch: pytest.MonkeyPatch, + switch_offset: int, + expected: bool, +) -> None: + monkeypatch.setattr( + pyscf_numint, + "SWITCH_SIZE", + carbon.nao_nr() + switch_offset, + ) + + assert _should_screen_aos(carbon) is expected + + +def test_active_cpu_aos(carbon: gto.Mole) -> None: + ao_loc = carbon.ao_loc_nr() + screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) + screen_index[0, 0] = 1 + screen_index[1, -1] = 1 + + expected = np.concatenate( + ( + np.arange(ao_loc[0], ao_loc[1]), + np.arange(ao_loc[-2], ao_loc[-1]), + ) + ) + + assert np.array_equal(_active_cpu_aos(carbon, screen_index), expected) + + empty = _active_cpu_aos(carbon, np.zeros_like(screen_index)) + assert empty.dtype == np.int64 + assert empty.size == 0 + + +@pytest.mark.parametrize("screen_aos", [False, True]) +def test_chunked_features_routes_screening( + carbon: gto.Mole, + monkeypatch: pytest.MonkeyPatch, + screen_aos: bool, +) -> None: + ngrids = dft.gen_grid.BLKSIZE + grids = dft.Grids(carbon) + grids.coords = np.zeros((ngrids, 3)) + grids.weights = np.ones(ngrids) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + calls: list[str] = [] + + monkeypatch.setattr( + features_module, + "get_grid_features", + lambda *args, **kwargs: { + "atomic_grid_sizes": torch.tensor([ngrids]), + "grid_weights": torch.ones(ngrids, dtype=torch.float64), + }, + ) + monkeypatch.setattr( + features_module, + "estimate_max_grid_chunk_size", + lambda *args, **kwargs: ngrids, + ) + monkeypatch.setattr( + dft.gen_grid, + "make_screen_index", + lambda *args, **kwargs: np.ones((1, carbon.nbas), dtype=np.uint8), + ) + + def fake_non_chunk(*args: object, **kwargs: object) -> torch.Tensor: + calls.append("non_chunk") + return torch.zeros((1, ngrids), dtype=torch.float64) + + def fake_chunk_eval(*args: object, **kwargs: object) -> torch.Tensor: + calls.append("ChunkEval") + return torch.zeros((1, ngrids), dtype=torch.float64) + + monkeypatch.setattr(features_module, "non_chunk", fake_non_chunk) + monkeypatch.setattr(ChunkEvalForward, "apply", fake_chunk_eval) + + list( + chunked_features( + carbon, + dm, + grids, + {"atomic_grid_sizes", "density", "grid_weights"}, + func_deriv=1, + screen_aos=screen_aos, + ) + ) + + assert calls == ["ChunkEval" if screen_aos else "non_chunk"] + + if screen_aos: + assert CPU_AO_SCREENING_BLOCK_SIZE == 504 + + +def test_cpu_screening_spatially_groups_each_atom( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + atom_grid_size = dft.gen_grid.BLKSIZE + ngrids = 2 * atom_grid_size + coords = np.zeros((ngrids, 3)) + coords[:, 0] = np.arange(ngrids) + weights = np.arange(ngrids, dtype=np.float64) + 100 + atomic_grid_weights = torch.arange(ngrids, dtype=torch.float64) + 200 + grids = dft.Grids(mol) + grids.coords = coords + grids.weights = weights + grouped_slices: list[np.ndarray] = [] + + monkeypatch.setattr( + features_module, + "get_grid_features", + lambda *args, **kwargs: { + "atomic_grid_sizes": torch.tensor([atom_grid_size, atom_grid_size]), + "grid_coords": torch.from_numpy(coords.copy()), + "grid_weights": torch.from_numpy(weights.copy()), + "atomic_grid_weights": atomic_grid_weights, + }, + ) + monkeypatch.setattr( + features_module, + "estimate_max_grid_chunk_size", + lambda *args, **kwargs: ngrids, + ) + + def fake_group_grids(mol_arg: gto.Mole, atom_coords: np.ndarray) -> np.ndarray: + assert mol_arg is mol + grouped_slices.append(atom_coords.copy()) + return np.arange(atom_grid_size - 1, -1, -1) + + monkeypatch.setattr(dft.gen_grid, "arg_group_grids", fake_group_grids) + + sort_indices = np.concatenate( + ( + np.arange(atom_grid_size - 1, -1, -1), + np.arange(ngrids - 1, atom_grid_size - 1, -1), + ) + ) + + def fake_make_screen_index( + mol_arg: gto.Mole, sorted_coords: np.ndarray, cutoff: float + ) -> np.ndarray: + assert mol_arg is mol + assert np.array_equal(sorted_coords, coords[sort_indices]) + return np.ones((2, mol.nbas), dtype=np.uint8) + + monkeypatch.setattr(dft.gen_grid, "make_screen_index", fake_make_screen_index) + + def fake_chunk_eval( + dm: torch.Tensor, + mol_arg: gto.Mole, + sorted_grids: dft.Grids, + feature_function: MGGAFeatureFunction, + block_size: int, + compile_feature_function: bool, + gpu: bool, + ) -> torch.Tensor: + assert mol_arg is mol + assert feature_function.with_density + assert block_size == CPU_AO_SCREENING_BLOCK_SIZE + assert not compile_feature_function + assert not gpu + assert np.array_equal(sorted_grids.coords, coords[sort_indices]) + assert np.array_equal(sorted_grids.weights, weights[sort_indices]) + return torch.from_numpy(sorted_grids.coords[:, 0]).to(dm).unsqueeze(0) + + monkeypatch.setattr(ChunkEvalForward, "apply", fake_chunk_eval) + + (feature_chunk,) = list( + chunked_features( + mol, + torch.eye(mol.nao_nr(), dtype=torch.float64), + grids, + { + "atomic_grid_sizes", + "atomic_grid_weights", + "density", + "grid_coords", + "grid_weights", + }, + func_deriv=1, + screen_aos=True, + ) + ) + + assert len(grouped_slices) == 2 + assert np.array_equal(grouped_slices[0], coords[:atom_grid_size]) + assert np.array_equal(grouped_slices[1], coords[atom_grid_size:]) + assert torch.equal( + feature_chunk["atomic_grid_sizes"], + torch.tensor([atom_grid_size, atom_grid_size]), + ) + assert torch.equal( + feature_chunk["grid_coords"], torch.from_numpy(coords[sort_indices]) + ) + assert torch.equal( + feature_chunk["grid_weights"], torch.from_numpy(weights[sort_indices]) + ) + assert torch.equal( + feature_chunk["atomic_grid_weights"], atomic_grid_weights[sort_indices] + ) + expected_density = torch.from_numpy(coords[sort_indices, 0]) / 2 + assert torch.equal(feature_chunk["density"][0], expected_density) + assert torch.equal(feature_chunk["density"][1], expected_density) + + +class QuadraticDensityFunctional(ExcFunctionalBase): + def __init__(self) -> None: + super().__init__() + self.features = ["atomic_grid_sizes", "density", "grid_weights"] + + def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + return (mol["density"].square() * mol["grid_weights"]).sum() + + +class FakeKS: + def __init__(self, mol: gto.Mole, grids: object | None = None) -> None: + self.mol = mol + self.grids = grids or object() + self.max_memory = 100 + + def make_rdm1(self, mo_coeff: np.ndarray, mo_occ: np.ndarray) -> np.ndarray: + return np.eye(self.mol.nao_nr()) + + def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: + return np.zeros_like(dm) + + +@pytest.mark.parametrize("expected", [False, True]) +def test_first_and_second_order_use_same_screening_decision( + carbon: gto.Mole, + monkeypatch: pytest.MonkeyPatch, + expected: bool, +) -> None: + switch_size = carbon.nao_nr() - 1 if expected else carbon.nao_nr() + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", switch_size) + decisions: list[bool] = [] + + def fake_chunked_features( + mol: gto.Mole, + dm: torch.Tensor, + grids: object, + features: set[str], + func_deriv: int, + *, + screen_aos: bool, + **kwargs: object, + ) -> Iterator[dict[str, torch.Tensor]]: + decisions.append(screen_aos) + density = dm.square().sum().reshape(1).expand(2, 1) / 2 + yield { + "atomic_grid_sizes": torch.tensor([1]), + "density": density, + "grid_weights": torch.ones(1, dtype=dm.dtype), + } + + monkeypatch.setattr(numint_module, "chunked_features", fake_chunked_features) + numint = SkalaNumInt(QuadraticDensityFunctional()) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + + numint(carbon, object(), None, dm) + + ks = FakeKS(carbon) + response = numint.gen_response( + np.eye(carbon.nao_nr()), np.ones(carbon.nao_nr()), ks=ks + ) + result = response(np.eye(carbon.nao_nr())) + + assert result.shape == (carbon.nao_nr(), carbon.nao_nr()) + assert decisions == [expected, expected] + + +def test_cpu_screening_slices_and_scatters_full_derivatives( + carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch +) -> None: + ngrids = dft.gen_grid.BLKSIZE + grids = dft.Grids(carbon) + grids.coords = np.zeros((ngrids, 3)) + grids.weights = np.ones(ngrids) + + ao = np.arange(ngrids * carbon.nao_nr(), dtype=np.float64).reshape( + ngrids, carbon.nao_nr() + ) + screen_index = np.zeros((1, carbon.nbas), dtype=np.uint8) + screen_index[0, (0, -1)] = 1 + grids.non0tab = screen_index + active_aos = _active_cpu_aos(carbon, screen_index) + + class FakeNumInt: + def block_loop( + self, *args: object, **kwargs: object + ) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]]: + assert kwargs["non0tab"] is screen_index + yield ao, screen_index, grids.weights, grids.coords + + monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) + + feature_function = MGGAFeatureFunction( + with_density=True, with_grad=False, with_kin=False + ) + dm = torch.diag( + torch.arange(1, carbon.nao_nr() + 1, dtype=torch.float64) + ).requires_grad_() + features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm, carbon, grids, feature_function, ngrids, False, False + ) + + ao_active = torch.from_numpy(ao[:, active_aos]).T + dm_active = dm[..., active_aos[:, None], active_aos[None, :]] + expected = torch.sum((dm_active @ ao_active) * ao_active, dim=0).unsqueeze(0) + assert torch.allclose(features, expected) + + energy = features.square().sum() + (vxc,) = torch.autograd.grad(energy, dm, create_graph=True) + (hvp,) = torch.autograd.grad(vxc, dm, torch.ones_like(dm)) + + inactive_aos = np.setdiff1d(np.arange(carbon.nao_nr()), active_aos) + assert vxc.shape == dm.shape + assert hvp.shape == dm.shape + assert torch.count_nonzero(vxc[inactive_aos]) == 0 + assert torch.count_nonzero(vxc[:, inactive_aos]) == 0 + assert torch.count_nonzero(hvp[inactive_aos]) == 0 + assert torch.count_nonzero(hvp[:, inactive_aos]) == 0 + + +def test_cpu_no_active_aos_returns_full_zero_derivatives( + carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch +) -> None: + ngrids = dft.gen_grid.BLKSIZE + grids = dft.Grids(carbon) + grids.coords = np.zeros((ngrids, 3)) + grids.weights = np.ones(ngrids) + ao = np.ones((ngrids, carbon.nao_nr())) + screen_index = np.zeros((1, carbon.nbas), dtype=np.uint8) + + class FakeNumInt: + def block_loop( + self, *args: object, **kwargs: object + ) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]]: + yield ao, screen_index, grids.weights, grids.coords + + monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) + feature_function = MGGAFeatureFunction( + with_density=True, with_grad=False, with_kin=False + ) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64).requires_grad_() + + features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm, carbon, grids, feature_function, ngrids, False, False + ) + (vxc,) = torch.autograd.grad(features.square().sum(), dm, create_graph=True) + (hvp,) = torch.autograd.grad(vxc, dm, torch.ones_like(dm)) + + assert features.shape == (1, ngrids) + assert vxc.shape == dm.shape + assert hvp.shape == dm.shape + assert torch.count_nonzero(features) == 0 + assert torch.count_nonzero(vxc) == 0 + assert torch.count_nonzero(hvp) == 0 + + +def _minimal_atom_grid(mol: gto.Mole) -> dft.Grids: + grids = dft.Grids(mol) + grids.level = 0 + grids.alignment = 1 + return grids.build(sort_grids=False) + + +@pytest.mark.parametrize("unrestricted", [False, True]) +def test_cpu_rks_uks_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, + load_functional_cached: Callable[..., ExcFunctionalBase | str], + unrestricted: bool, +) -> None: + if unrestricted: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + mean_field = dft.UKS(mol) + else: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + mean_field = dft.RKS(mol) + + functional = load_functional_cached("skala-1.1") + assert isinstance(functional, ExcFunctionalBase) + numint = SkalaNumInt(functional) + grids = _minimal_atom_grid(mol) + dm = mean_field.get_init_guess() + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense = ( + numint.nr_uks(mol, grids, None, dm) + if unrestricted + else numint.nr_rks(mol, grids, None, dm) + ) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened = ( + numint.nr_uks(mol, grids, None, dm) + if unrestricted + else numint.nr_rks(mol, grids, None, dm) + ) + + assert np.allclose(dense[0], screened[0], rtol=1e-10, atol=1e-11) + assert np.isclose(dense[1], screened[1], rtol=1e-9, atol=1e-10) + assert np.allclose(dense[2], screened[2], rtol=1e-8, atol=1e-10) + + +def test_cpu_response_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + grids = _minimal_atom_grid(mol) + ks = FakeKS(mol, grids) + numint = SkalaNumInt(QuadraticDensityFunctional()) + mo_coeff = np.eye(mol.nao_nr()) + mo_occ = np.ones(mol.nao_nr()) + dm1 = np.arange(mol.nao_nr() ** 2, dtype=np.float64).reshape( + mol.nao_nr(), mol.nao_nr() + ) + dm1 += dm1.T + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) + + assert np.allclose( + dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 + ) diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py new file mode 100644 index 00000000..207dab7b --- /dev/null +++ b/tests/test_ao_screening_benchmark.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +from collections.abc import Callable, Iterator +from typing import NamedTuple, cast + +import numpy as np +import pytest +import torch +from pyscf import dft, gto, lib +from pyscf.dft import numint as pyscf_numint +from pytest_benchmark.fixture import BenchmarkFixture + +from skala.functional.base import ExcFunctionalBase +from skala.pyscf.numint import SkalaNumInt, _should_screen_aos + +THREAD_COUNT = 4 +MAX_MEMORY_MB = 2000 + +NAPHTHALENE = """ +C -1.2280 0.7090 0.0 +C -1.2280 -0.7090 0.0 +C 0.0000 1.4180 0.0 +C 0.0000 -1.4180 0.0 +C 1.2280 0.7090 0.0 +C 1.2280 -0.7090 0.0 +C 2.4560 1.4180 0.0 +C 2.4560 -1.4180 0.0 +C 3.6840 0.7090 0.0 +C 3.6840 -0.7090 0.0 +H -2.1700 1.2530 0.0 +H -2.1700 -1.2530 0.0 +H 0.0000 2.5060 0.0 +H 0.0000 -2.5060 0.0 +H 2.4560 2.5060 0.0 +H 2.4560 -2.5060 0.0 +H 4.6260 1.2530 0.0 +H 4.6260 -1.2530 0.0 +""" + +ANTHRACENE = """ +C -1.2280 0.7090 0.0 +C -1.2280 -0.7090 0.0 +C 0.0000 1.4180 0.0 +C 0.0000 -1.4180 0.0 +C 1.2280 0.7090 0.0 +C 1.2280 -0.7090 0.0 +C 2.4560 1.4180 0.0 +C 2.4560 -1.4180 0.0 +C 3.6840 0.7090 0.0 +C 3.6840 -0.7090 0.0 +C 4.9120 1.4180 0.0 +C 4.9120 -1.4180 0.0 +C 6.1400 0.7090 0.0 +C 6.1400 -0.7090 0.0 +H -2.1700 1.2530 0.0 +H -2.1700 -1.2530 0.0 +H 0.0000 2.5060 0.0 +H 0.0000 -2.5060 0.0 +H 2.4560 2.5060 0.0 +H 2.4560 -2.5060 0.0 +H 4.9120 2.5060 0.0 +H 4.9120 -2.5060 0.0 +H 7.0820 1.2530 0.0 +H 7.0820 -1.2530 0.0 +""" + +TETRACENE = """ +C -1.2280 0.7090 0.0 +C -1.2280 -0.7090 0.0 +C 0.0000 1.4180 0.0 +C 0.0000 -1.4180 0.0 +C 1.2280 0.7090 0.0 +C 1.2280 -0.7090 0.0 +C 2.4560 1.4180 0.0 +C 2.4560 -1.4180 0.0 +C 3.6840 0.7090 0.0 +C 3.6840 -0.7090 0.0 +C 4.9120 1.4180 0.0 +C 4.9120 -1.4180 0.0 +C 6.1400 0.7090 0.0 +C 6.1400 -0.7090 0.0 +C 7.3680 1.4180 0.0 +C 7.3680 -1.4180 0.0 +C 8.5960 0.7090 0.0 +C 8.5960 -0.7090 0.0 +H -2.1700 1.2530 0.0 +H -2.1700 -1.2530 0.0 +H 0.0000 2.5060 0.0 +H 0.0000 -2.5060 0.0 +H 2.4560 2.5060 0.0 +H 2.4560 -2.5060 0.0 +H 4.9120 2.5060 0.0 +H 4.9120 -2.5060 0.0 +H 7.3680 2.5060 0.0 +H 7.3680 -2.5060 0.0 +H 9.5380 1.2530 0.0 +H 9.5380 -1.2530 0.0 +""" + + +class BenchmarkSpec(NamedTuple): + name: str + atoms: str + + +class BenchmarkCase(NamedTuple): + mol: gto.Mole + grids: dft.Grids + dm: np.ndarray + numint: SkalaNumInt[np.ndarray] + + +@pytest.fixture(scope="module") +def fixed_cpu_threads() -> Iterator[None]: + previous_pyscf_threads = lib.num_threads() + previous_torch_threads = torch.get_num_threads() + lib.num_threads(THREAD_COUNT) + torch.set_num_threads(THREAD_COUNT) + try: + yield + finally: + torch.set_num_threads(previous_torch_threads) + lib.num_threads(previous_pyscf_threads) + + +@pytest.fixture( + scope="module", + params=[ + pytest.param(BenchmarkSpec("naphthalene", NAPHTHALENE), id="naphthalene"), + pytest.param(BenchmarkSpec("anthracene", ANTHRACENE), id="anthracene"), + pytest.param(BenchmarkSpec("tetracene", TETRACENE), id="tetracene"), + ], +) +def benchmark_case( + request: pytest.FixtureRequest, + fixed_cpu_threads: None, + load_functional_cached: Callable[..., ExcFunctionalBase | str], +) -> BenchmarkCase: + spec = cast(BenchmarkSpec, request.param) + mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) + grids = dft.Grids(mol) + grids.level = 1 + grids.build(sort_grids=False) + dm = dft.RKS(mol).get_init_guess() + functional = load_functional_cached("skala-1.1") + assert isinstance(functional, ExcFunctionalBase) + return BenchmarkCase(mol, grids, dm, SkalaNumInt(functional)) + + +@pytest.fixture +def screened_case(benchmark_case: BenchmarkCase) -> BenchmarkCase: + assert _should_screen_aos(benchmark_case.mol) + return benchmark_case + + +@pytest.fixture +def dense_case( + benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch +) -> BenchmarkCase: + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", benchmark_case.mol.nao_nr()) + assert not _should_screen_aos(benchmark_case.mol) + return benchmark_case + + +def _run_xc(case: BenchmarkCase) -> tuple[float, float, np.ndarray]: + return case.numint.nr_rks( + case.mol, + case.grids, + None, + case.dm, + max_memory=MAX_MEMORY_MB, + ) + + +def _benchmark_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: + pedantic = cast(Callable[..., object], benchmark.pedantic) + pedantic( + _run_xc, + args=(case,), + rounds=1, + iterations=2, + ) + + +@pytest.mark.profiling +def test_screened_and_dense_values_agree( + benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch +) -> None: + assert _should_screen_aos(benchmark_case.mol) + screened = _run_xc(benchmark_case) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", benchmark_case.mol.nao_nr()) + assert not _should_screen_aos(benchmark_case.mol) + dense = _run_xc(benchmark_case) + + assert np.allclose(dense[0], screened[0], rtol=1e-10, atol=1e-11) + assert np.isclose(dense[1], screened[1], rtol=1e-10, atol=1e-11) + vxc_difference = dense[2] - screened[2] + vxc_max_abs_difference = np.max(np.abs(vxc_difference)) + vxc_relative_l2_difference = np.linalg.norm(vxc_difference) / np.linalg.norm( + dense[2] + ) + assert vxc_max_abs_difference < 5e-8 and vxc_relative_l2_difference < 1e-8, ( + f"N: dense={dense[0]:.16g}, screened={screened[0]:.16g}, " + f"abs_diff={abs(dense[0] - screened[0]):.3e}; " + f"E_xc: dense={dense[1]:.16g}, screened={screened[1]:.16g}, " + f"abs_diff={abs(dense[1] - screened[1]):.3e}; " + f"V_xc: max_abs_diff={vxc_max_abs_difference:.3e}, " + f"relative_l2_diff={vxc_relative_l2_difference:.3e}" + ) + + +@pytest.mark.benchmark(group="def2-qzvpp") +def test_with_natural_ao_screening( + benchmark: BenchmarkFixture, screened_case: BenchmarkCase +) -> None: + _benchmark_xc(benchmark, screened_case) + + +@pytest.mark.benchmark(group="def2-qzvpp") +def test_without_ao_screening_by_patching_threshold( + benchmark: BenchmarkFixture, dense_case: BenchmarkCase +) -> None: + _benchmark_xc(benchmark, dense_case) + + +@pytest.mark.profiling +def test_profile_with_natural_ao_screening( + screened_case: BenchmarkCase, +) -> None: + _run_xc(screened_case) + + +@pytest.mark.profiling +def test_profile_without_ao_screening_by_patching_threshold( + dense_case: BenchmarkCase, +) -> None: + _run_xc(dense_case) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py new file mode 100644 index 00000000..540458e3 --- /dev/null +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -0,0 +1,170 @@ +from collections.abc import Callable, Iterator +from types import SimpleNamespace + +import numpy as np +import pytest +import torch +from pyscf import gto +from pyscf.dft import numint as pyscf_numint +from torch.utils.dlpack import from_dlpack + +if not torch.cuda.is_available(): + pytest.skip( + "Skipping gpu4pyscf AO screening tests, because CUDA is not available.", + allow_module_level=True, + ) + +try: + import cupy +except ModuleNotFoundError: + pytest.skip( + "Skipping gpu4pyscf AO screening tests, because CuPy is not available.", + allow_module_level=True, + ) + +from skala.functional.base import ExcFunctionalBase +from skala.gpu4pyscf import SkalaKS +from skala.pyscf.backend import dft_gpu +from skala.pyscf.features import ChunkEvalForward, MGGAFeatureFunction + + +class QuadraticDensityFunctional(ExcFunctionalBase): + def __init__(self) -> None: + super().__init__() + self.features = ["atomic_grid_sizes", "density", "grid_weights"] + + def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + return (mol["density"].square() * mol["grid_weights"]).sum() + + +def _to_numpy(value: object) -> np.ndarray: + return cupy.asnumpy(value) if isinstance(value, cupy.ndarray) else np.asarray(value) + + +@pytest.mark.parametrize("unrestricted", [False, True]) +def test_gpu_rks_uks_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, + load_functional_cached: Callable[..., ExcFunctionalBase | str], + unrestricted: bool, +) -> None: + if unrestricted: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + else: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + + functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(functional, ExcFunctionalBase) + ks = SkalaKS(mol, xc=functional, with_dftd3=False) + ks.grids.level = 0 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + dm = ks.get_init_guess() + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense = ( + ks._numint.nr_uks(mol, ks.grids, None, dm) + if unrestricted + else ks._numint.nr_rks(mol, ks.grids, None, dm) + ) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened = ( + ks._numint.nr_uks(mol, ks.grids, None, dm) + if unrestricted + else ks._numint.nr_rks(mol, ks.grids, None, dm) + ) + + assert np.allclose(_to_numpy(dense[0]), _to_numpy(screened[0]), rtol=1e-9) + assert np.isclose(dense[1], screened[1], rtol=1e-9) + assert np.allclose( + _to_numpy(dense[2]), _to_numpy(screened[2]), rtol=1e-8, atol=2e-9 + ) + + +def test_gpu_response_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) + ks.grids.level = 0 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + mo_coeff = cupy.eye(mol.nao_nr()) + mo_occ = cupy.ones(mol.nao_nr()) + dm1 = cupy.arange(mol.nao_nr() ** 2, dtype=cupy.float64).reshape( + mol.nao_nr(), mol.nao_nr() + ) + dm1 += dm1.T + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + assert np.allclose( + _to_numpy(dense_response(dm1)), + _to_numpy(screened_response(dm1)), + rtol=1e-9, + atol=1e-10, + ) + + +def test_gpu_sparse_mask_sorts_scatters_and_unsorts( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0) + ngrids = 32 + sort_idx = np.array([2, 0, 4, 1, 3]) + active_sorted_aos = np.array([0, 2, 4]) + ao = cupy.arange(active_sorted_aos.size * ngrids, dtype=cupy.float64).reshape( + active_sorted_aos.size, ngrids + ) + weights = cupy.ones(ngrids) + coords = cupy.zeros((ngrids, 3)) + grids = SimpleNamespace(weights=weights, coords=coords) + + class FakeGpuNumInt: + def build(self, mol: gto.Mole, coords: cupy.ndarray) -> "FakeGpuNumInt": + self.gdftopt = SimpleNamespace(_ao_idx=sort_idx) + return self + + def block_loop( + self, *args: object, **kwargs: object + ) -> Iterator[tuple[object, object, object, object]]: + yield ao, cupy.asarray(active_sorted_aos), weights, coords + + assert dft_gpu is not None + monkeypatch.setattr(dft_gpu.numint, "NumInt", FakeGpuNumInt) + + feature_function = MGGAFeatureFunction( + with_density=True, with_grad=False, with_kin=False + ) + dm = torch.diag( + torch.arange(1, mol.nao_nr() + 1, dtype=torch.float64, device="cuda") + ).requires_grad_() + features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm, mol, grids, feature_function, None, False, True + ) + + sort_idx_t = torch.as_tensor(sort_idx, device="cuda") + active_t = torch.as_tensor(active_sorted_aos, device="cuda") + dm_sorted = dm[..., sort_idx_t, :][..., sort_idx_t] + dm_active = dm_sorted[..., active_t[:, None], active_t[None, :]] + ao_t = from_dlpack(ao) + expected = torch.sum((dm_active @ ao_t) * ao_t, dim=0).unsqueeze(0) + assert torch.allclose(features, expected) + + energy = features.square().sum() + (vxc,) = torch.autograd.grad(energy, dm, create_graph=True) + (hvp,) = torch.autograd.grad(vxc, dm, torch.ones_like(dm)) + + inactive_original_aos = sort_idx[ + np.setdiff1d(np.arange(mol.nao_nr()), active_sorted_aos) + ] + assert vxc.shape == dm.shape + assert hvp.shape == dm.shape + assert torch.count_nonzero(vxc[inactive_original_aos]) == 0 + assert torch.count_nonzero(vxc[:, inactive_original_aos]) == 0 + assert torch.count_nonzero(hvp[inactive_original_aos]) == 0 + assert torch.count_nonzero(hvp[:, inactive_original_aos]) == 0 From 658c2648089bfa50b46dbfee9214ced8aa8588bd Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 31 Jul 2026 11:20:51 +0200 Subject: [PATCH 02/39] add docstrings --- src/skala/pyscf/features.py | 23 +++++++++++++++++++++++ src/skala/pyscf/numint.py | 9 +++++++++ 2 files changed, 32 insertions(+) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index b53f9f4c..dc297f9c 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -40,6 +40,15 @@ def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: + """Expand a PySCF shell-screening mask into active AO indices. + + Args: + mol: Molecule defining the shell-to-AO ranges. + screen_index: Screening rows whose columns correspond to molecular shells. + + Returns: + Sorted indices of AOs belonging to a shell active in any screening row. + """ active_shells = np.any(screen_index, axis=0) ao_loc = mol.ao_loc_nr() return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) @@ -48,6 +57,20 @@ def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: def _spatially_group_atom_grids( mol: gto.Mole, coords: np.ndarray, atomic_grid_sizes: Tensor ) -> np.ndarray: + """Build a spatial grid permutation independently within each atom. + + Makes screening much more effective as points are spatially grouped + within each atom, while preserving the original atom order and atom boundaries. + + Args: + mol: Molecule used by PySCF to define the spatial grouping boxes. + coords: Atom-major grid coordinates to group. + atomic_grid_sizes: Number of consecutive grid points owned by each atom. + + Returns: + A permutation that spatially groups each atom's points while preserving + atom order and atom boundaries. + """ sort_indices = [] start = 0 for size in atomic_grid_sizes.tolist(): diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 8bb81fe0..b6ed9b11 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -22,6 +22,15 @@ def _should_screen_aos(mol: gto.Mole) -> bool: + """Determine whether a molecule is large enough for AO screening. + + Args: + mol: Molecule whose AO count is compared with PySCF's sparse-contraction + crossover. + + Returns: + Whether the molecule has more AOs than PySCF's screening threshold. + """ # Keep the compatibility fallback here, not at call sites. PySCF uses this # crossover before selecting sparse density/Vxc contractions. switch_size = pyscf_numint.SWITCH_SIZE From eaf5d1c3895e4d995801e198a4cd1947f359532a Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 31 Jul 2026 12:02:14 +0200 Subject: [PATCH 03/39] add failing test for consistency --- tests/test_gpu4pyscf_ao_screening.py | 71 +++++++++++++++++++++++++++- 1 file changed, 70 insertions(+), 1 deletion(-) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 540458e3..b264bbab 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -4,7 +4,7 @@ import numpy as np import pytest import torch -from pyscf import gto +from pyscf import dft, gto from pyscf.dft import numint as pyscf_numint from torch.utils.dlpack import from_dlpack @@ -26,6 +26,14 @@ from skala.gpu4pyscf import SkalaKS from skala.pyscf.backend import dft_gpu from skala.pyscf.features import ChunkEvalForward, MGGAFeatureFunction +from skala.pyscf.numint import SkalaNumInt + +CARBON_CHAIN = """ +C 0.0 0.0 0.0 +C 1.4 0.0 0.0 +C 2.8 0.0 0.0 +C 4.2 0.0 0.0 +""" class QuadraticDensityFunctional(ExcFunctionalBase): @@ -110,6 +118,67 @@ def test_gpu_response_dense_screened_equivalence( ) +def test_gpu_screened_skala_matches_cpu_on_carbon_chain( + monkeypatch: pytest.MonkeyPatch, + load_functional_cached: Callable[..., ExcFunctionalBase | str], +) -> None: + mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) + cpu_grids = dft.Grids(mol) + cpu_grids.level = 1 + cpu_grids.alignment = 1 + cpu_grids.build(sort_grids=False) + gpu_grids = dft_gpu.Grids(mol) + gpu_grids.level = 1 + gpu_grids.alignment = 1 + gpu_grids.build(sort_grids=False) + dm = dft.RKS(mol).get_init_guess() + + np.testing.assert_allclose( + cpu_grids.coords, + cupy.asnumpy(gpu_grids.coords), + rtol=0.0, + atol=0.0, + ) + np.testing.assert_allclose( + cpu_grids.weights, + cupy.asnumpy(gpu_grids.weights), + rtol=1e-12, + atol=1e-12, + ) + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + + cpu_functional = load_functional_cached("skala-1.1", device=torch.device("cpu")) + gpu_functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(cpu_functional, ExcFunctionalBase) + assert isinstance(gpu_functional, ExcFunctionalBase) + cpu_result = SkalaNumInt(cpu_functional, device=torch.device("cpu")).nr_rks( + mol, cpu_grids, None, dm + ) + gpu_result = SkalaNumInt(gpu_functional, device=torch.device("cuda:0")).nr_rks( + mol, gpu_grids, None, cupy.asarray(dm) + ) + + gpu_vxc = cupy.asnumpy(gpu_result[2]) + vxc_difference = cpu_result[2] - gpu_vxc + vxc_max_abs_difference = np.max(np.abs(vxc_difference)) + vxc_relative_l2_difference = np.linalg.norm(vxc_difference) / np.linalg.norm( + cpu_result[2] + ) + assert ( + np.isclose(cpu_result[0], gpu_result[0], rtol=1e-10, atol=1e-11) + and np.isclose(cpu_result[1], gpu_result[1], rtol=1e-10, atol=1e-11) + and vxc_max_abs_difference < 2e-9 + and vxc_relative_l2_difference < 1e-8 + ), ( + f"N: cpu={cpu_result[0]:.16g}, gpu={gpu_result[0]:.16g}, " + f"abs_diff={abs(cpu_result[0] - gpu_result[0]):.3e}; " + f"E_xc: cpu={cpu_result[1]:.16g}, gpu={gpu_result[1]:.16g}, " + f"abs_diff={abs(cpu_result[1] - gpu_result[1]):.3e}; " + f"V_xc: max_abs_diff={vxc_max_abs_difference:.3e}, " + f"relative_l2_diff={vxc_relative_l2_difference:.3e}" + ) + + def test_gpu_sparse_mask_sorts_scatters_and_unsorts( monkeypatch: pytest.MonkeyPatch, ) -> None: From a4dad08822ed9073a6cc4964dee7dbe000a27270 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 31 Jul 2026 12:12:03 +0200 Subject: [PATCH 04/39] add docstring and use dense eval --- tests/test_gpu4pyscf_ao_screening.py | 25 +++++++++++++++++++++++-- 1 file changed, 23 insertions(+), 2 deletions(-) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index b264bbab..983b07fa 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -122,6 +122,25 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> None: + """Catch inaccurate GPU AO screening on spatially diffuse grid blocks. + + GPU4PySCF builds one active-shell mask for each fixed-size coordinate block. That + screening is reliable only when the points in a block are spatially local enough + for the sampled AO values to represent the whole block. Skala currently supplies + an unsorted, atom-major grid, so one GPU block can span a large region around an + atom. This is especially problematic for the AO derivatives used by Skala: an AO + value can be small at the sampled points even though its gradient still makes a + significant contribution. The linear carbon chain and large def2-QZVPP basis + expose this failure in a reasonably small integration test. + + The CPU and GPU calculations use identical coordinates, weights, density matrix, + and Skala 1.1 model. CPU AO evaluation is deliberately forced dense to provide an + independent reference, while GPU AO evaluation is deliberately forced through + screening. Comparing the particle count, XC energy, and complete XC potential + matrix verifies the full feature and VJP path; the potential is particularly + sensitive to omitted derivative contributions. The test should pass once GPU + blocks are spatially grouped before their screening masks are constructed. + """ mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) cpu_grids = dft.Grids(mol) cpu_grids.level = 1 @@ -145,15 +164,17 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( rtol=1e-12, atol=1e-12, ) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - cpu_functional = load_functional_cached("skala-1.1", device=torch.device("cpu")) gpu_functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) assert isinstance(cpu_functional, ExcFunctionalBase) assert isinstance(gpu_functional, ExcFunctionalBase) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) cpu_result = SkalaNumInt(cpu_functional, device=torch.device("cpu")).nr_rks( mol, cpu_grids, None, dm ) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) gpu_result = SkalaNumInt(gpu_functional, device=torch.device("cuda:0")).nr_rks( mol, gpu_grids, None, cupy.asarray(dm) ) From 3fe437100b8d013f42dba33ba90761947c0c1766 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 3 Aug 2026 10:20:18 +0200 Subject: [PATCH 05/39] work on shell screening --- pyproject.toml | 2 + src/skala/pyscf/features.py | 324 ++++++++++++++++++++++++++- src/skala/pyscf/memory_estimators.py | 42 +++- src/skala/pyscf/numint.py | 167 ++++++++++++-- tests/test_ao_screening.py | 206 +++++++++++++++++ tests/test_gpu4pyscf_ao_screening.py | 158 +++++++++++-- tests/test_gpu4pyscf_classes.py | 2 + tests/test_gpu4pyscf_gradients.py | 2 + tests/test_memory_estimators.py | 57 +++++ 9 files changed, 921 insertions(+), 39 deletions(-) create mode 100644 tests/test_memory_estimators.py diff --git a/pyproject.toml b/pyproject.toml index b1a40952..e3ba6a9e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -97,8 +97,10 @@ line-length = 100 [tool.pytest.ini_options] timeout = 300 +addopts = "--benchmark-skip -m 'not profiling'" markers = [ "benchmark: performance measurements collected by pytest-benchmark", + "gpu: requires a CUDA-capable GPU and GPU test dependencies", "profiling: single-call performance workloads intended for profilers", ] filterwarnings = [ diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index dc297f9c..717822d6 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -8,6 +8,8 @@ from abc import ABC, abstractmethod from collections.abc import Callable, Iterator from copy import copy +from dataclasses import dataclass +from typing import Literal, TypeAlias import numpy as np import torch @@ -23,7 +25,10 @@ dft_gpu, from_numpy_or_cupy, ) -from skala.pyscf.memory_estimators import estimate_max_grid_chunk_size +from skala.pyscf.memory_estimators import ( + estimate_global_raw_feature_buffer_memory, + estimate_max_grid_chunk_size, +) LOG = logging.getLogger(__name__) @@ -38,6 +43,40 @@ "atomic_grid_size_bound_shape", } +_Float64Coordinates: TypeAlias = np.ndarray[ + tuple[int, Literal[3]], np.dtype[np.float64] +] +_Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] +_SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" + + +@dataclass(frozen=True) +class _SpatialGridCache: + mol: gto.Mole + source_coords: object + source_weights: object + block_size: int + gpu: bool + sorted_grids: Grid + forward: _Int64Permutation + inverse: _Int64Permutation + + def matches( + self, + mol: gto.Mole, + grids: Grid, + block_size: int, + gpu: bool, + ) -> bool: + """Return whether this entry belongs to the current built grid.""" + return ( + self.mol is mol + and self.source_coords is grids.coords + and self.source_weights is grids.weights + and self.block_size == block_size + and self.gpu is gpu + ) + def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: """Expand a PySCF shell-screening mask into active AO indices. @@ -54,6 +93,125 @@ def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) +def _spatial_grid_permutations( + coords: _Float64Coordinates, block_size: int +) -> tuple[_Int64Permutation, _Int64Permutation]: + """Order a molecular grid into exact-size spatial blocks. + + Recursively partitions points along the longest Cartesian extent. Every left + subtree contains a whole number of evaluator blocks, so all output blocks have + ``block_size`` points except for a possible final remainder. + + Args: + coords: Molecular grid coordinates with shape ``(ngrids, 3)``. + block_size: Fixed number of points consumed by each backend block. + + Returns: + The forward permutation from atom-major to spatial order and its inverse. + + Raises: + ValueError: If the coordinates or block size are invalid. + """ + if coords.ndim != 2 or coords.shape[1] != 3: + raise ValueError("coords must have shape (ngrids, 3)") + if block_size <= 0: + raise ValueError("block_size must be positive") + + def partition(indices: _Int64Permutation) -> list[_Int64Permutation]: + if indices.size <= block_size: + return [indices] + + block_count = (indices.size + block_size - 1) // block_size + left_size = (block_count // 2) * block_size + extents = np.ptp(coords[indices], axis=0) + split_axis = int(np.argmax(extents)) + positions = np.lexsort((indices, coords[indices, split_axis])) + ordered_indices = indices[positions] + return partition(ordered_indices[:left_size]) + partition( + ordered_indices[left_size:] + ) + + ngrids = coords.shape[0] + if ngrids == 0: + empty = np.empty(0, dtype=np.int64) + return empty, empty.copy() + + forward = np.concatenate(partition(np.arange(ngrids, dtype=np.int64))) + inverse = np.empty_like(forward) + inverse[forward] = np.arange(ngrids, dtype=np.int64) + return forward, inverse + + +def _prepare_spatially_sorted_grids( + mol: gto.Mole, + grids: Grid, + block_size: int, + gpu: bool, +) -> tuple[Grid, _Int64Permutation, _Int64Permutation]: + """Copy and spatially order a grid for backend AO screening. + + Preparation is cached on the source grid and reused while its coordinate and + weight arrays, molecule, backend, and evaluator block size remain unchanged. + + Args: + mol: Molecule used to rebuild CPU shell-screening data. + grids: Built CPU or GPU integration grid in atom-major order. + block_size: Fixed number of points consumed by each backend block. + gpu: Whether ``grids`` belongs to GPU4PySCF. + + Returns: + The sorted grid copy, atom-major-to-spatial permutation, and inverse. + """ + if grids.coords is None or grids.weights is None: + raise ValueError("Grids must be built before spatial sorting.") + + cache = getattr(grids, _SPATIAL_GRID_CACHE_ATTRIBUTE, None) + if isinstance(cache, _SpatialGridCache) and cache.matches( + mol, grids, block_size, gpu + ): + return cache.sorted_grids, cache.forward, cache.inverse + + if gpu: + check_gpu_imports_were_successful() + import cupy + + host_coords = cupy.asnumpy(grids.coords) + else: + host_coords = grids.coords + + forward, inverse = _spatial_grid_permutations(host_coords, block_size) + sorted_grids = copy(grids) + vars(sorted_grids).pop(_SPATIAL_GRID_CACHE_ATTRIBUTE, None) + if gpu: + backend_forward = cupy.asarray(forward) + sorted_grids.coords = grids.coords[backend_forward] + sorted_grids.weights = grids.weights[backend_forward] + sorted_grids._non0ao_idx = None + else: + sorted_grids.coords = grids.coords[forward] + sorted_grids.weights = grids.weights[forward] + sorted_grids.non0tab = dft.gen_grid.make_screen_index( + mol, + sorted_grids.coords, + cutoff=sorted_grids.cutoff, + ) + setattr( + grids, + _SPATIAL_GRID_CACHE_ATTRIBUTE, + _SpatialGridCache( + mol=mol, + source_coords=grids.coords, + source_weights=grids.weights, + block_size=block_size, + gpu=gpu, + sorted_grids=sorted_grids, + forward=forward, + inverse=inverse, + ), + ) + return sorted_grids, forward, inverse + + def _spatially_group_atom_grids( mol: gto.Mole, coords: np.ndarray, atomic_grid_sizes: Tensor ) -> np.ndarray: @@ -683,6 +841,168 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: return features.reshape((*dm.shape[:-2], self.nfeats, -1)) +@dataclass +class _GlobalScreenedFeatures: + dm: Tensor + mol: gto.Mole + sorted_grids: Grid + sorted_raw_features: Tensor + atom_major_raw_features: Tensor + forward_permutation: Tensor + inverse_permutation: Tensor + feature_function: MGGAFeatureFunction + block_size: int + compile_feature_function: bool + gpu: bool + grid_features: dict[str, Tensor] + feature_names: set[str] + chunks: list[tuple[slice, slice]] + with_spin: bool + + def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: + """Apply the global raw-feature Jacobian and restore atom-major order.""" + if not self.feature_function.only_linear_feats: + raise NotImplementedError( + "Global screened response requires raw features linear in the density " + "matrix." + ) + sorted_tangent = ChunkEvalForward.apply( + self.dm, + self.mol, + self.sorted_grids, + self.feature_function, + self.block_size, + self.compile_feature_function, + self.gpu, + dm_tangent, + ) + return sorted_tangent.index_select(-1, self.inverse_permutation).detach() + + def build_model_chunk( + self, + raw_features: Tensor, + atom_slice: slice, + grid_slice: slice, + ) -> dict[str, Tensor]: + """Build one atom-aligned model dictionary from raw feature values.""" + feature_chunk: dict[str, Tensor] = {} + for feature_name in ("grid_coords", "grid_weights", "atomic_grid_weights"): + if feature_name in self.feature_names: + feature_chunk[feature_name] = self.grid_features[feature_name][ + grid_slice + ] + + for feature_name in ("coarse_0_atomic_coords", "atomic_grid_sizes"): + if feature_name in self.feature_names: + feature_chunk[feature_name] = self.grid_features[feature_name][ + atom_slice + ] + + if "atomic_grid_size_bound_shape" in self.feature_names: + max_size = int(feature_chunk["atomic_grid_sizes"].max().item()) + feature_chunk["atomic_grid_size_bound_shape"] = torch.zeros( + max_size, + 0, + dtype=torch.long, + device=raw_features.device, + ) + + for feature_name, feature in self.feature_function.to_dict( + raw_features + ).items(): + feature_chunk[feature_name] = maybe_expand_and_divide( + feature, not self.with_spin, 2 + ) + return feature_chunk + + +def _global_screened_features( + mol: gto.Mole, + dm: Tensor, + grids: Grid, + features: set[str], + func_deriv: int, + max_memory_in_mb: int | None = None, + safety_fraction: float = 0.8, + compile_feature_function: bool = False, +) -> _GlobalScreenedFeatures: + """Evaluate raw AO features once on a spatially ordered molecular grid.""" + if "atomic_grid_sizes" not in features: + raise ValueError( + "Global screened features require 'atomic_grid_sizes' for model chunks." + ) + if grids.coords is None or grids.weights is None: + raise ValueError("Grids must be built before generating screened features.") + + feature_function = MGGAFeatureFunction( + with_density="density" in features, + with_grad="grad" in features, + with_kin="kin" in features, + with_lapl="lapl" in features, + ) + grid_features = get_grid_features(mol, dm, grids, features) + max_grid_chunk_size = estimate_max_grid_chunk_size( + dm=dm, + deriv=feature_function.deriv, + max_memory_in_mb=max_memory_in_mb, + safety_fraction=safety_fraction, + func_deriv=func_deriv, + reserved_memory_in_bytes=estimate_global_raw_feature_buffer_memory( + dm, + feature_function.nfeats, + grids.weights.size, + func_deriv, + ), + ) + max_atom_grid = int(grid_features["atomic_grid_sizes"].max().item()) + if max_grid_chunk_size < max_atom_grid: + LOG.warning( + f"Adjusted chunk size {max_grid_chunk_size} to match the largest atomic grid " + f"{max_atom_grid}. Hope for no OOM." + ) + max_grid_chunk_size = max_atom_grid + + gpu = dm.device.type == "cuda" + if gpu: + check_gpu_imports_were_successful() + block_size = int(dft_gpu.numint.MIN_BLK_SIZE) + else: + block_size = CPU_AO_SCREENING_BLOCK_SIZE + sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( + mol, grids, block_size, gpu + ) + sorted_raw_features = ChunkEvalForward.apply( + dm.double(), + mol, + sorted_grids, + feature_function, + block_size, + compile_feature_function, + gpu, + ) + forward_permutation = torch.as_tensor(forward, device=dm.device) + inverse_permutation = torch.as_tensor(inverse, device=dm.device) + atom_major_raw_features = sorted_raw_features.index_select(-1, inverse_permutation) + chunks = make_chunks(grid_features["atomic_grid_sizes"], max_grid_chunk_size) + return _GlobalScreenedFeatures( + dm=dm, + mol=mol, + sorted_grids=sorted_grids, + sorted_raw_features=sorted_raw_features, + atom_major_raw_features=atom_major_raw_features, + forward_permutation=forward_permutation, + inverse_permutation=inverse_permutation, + feature_function=feature_function, + block_size=block_size, + compile_feature_function=compile_feature_function, + gpu=gpu, + grid_features=grid_features, + feature_names=features, + chunks=chunks, + with_spin=dm.ndim == 3, + ) + + class ChunkEvalForward(Function): @staticmethod def setup_context( @@ -727,7 +1047,7 @@ def forward( block_loop_args = (mol, grids, mol.nao) block_loop_kwargs = { "deriv": feature_function.deriv, - "blksize": blksize if not gpu else None, + "blksize": blksize, "non0tab": None if gpu else getattr(grids, "non0tab", None), } if gpu: diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index 6df15ea0..75d34a49 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -9,6 +9,7 @@ def estimate_max_grid_chunk_size( max_memory_in_mb: int | None = None, safety_fraction: float = 0.8, func_deriv: int = 1, + reserved_memory_in_bytes: int = 0, ) -> int: """Heuristically pick a grid chunk size for :func:`chunked_features`. @@ -40,6 +41,9 @@ def estimate_max_grid_chunk_size( (``exc_only``), ``1`` first order (``__call__``/``V_xc``), ``2`` second order (``gen_response``/Hessian-vector product). Selects the calibrated coefficients. + reserved_memory_in_bytes: Memory already committed to allocations whose + size does not depend on the model chunk, such as global raw-feature + and cotangent buffers. Returns: Maximum number of grid points per chunk whose predicted peak memory fits @@ -71,7 +75,7 @@ def estimate_max_grid_chunk_size( ) else: free_bytes = int(max_memory_in_mb * 1000**2) - free_bytes = int(free_bytes * safety_fraction) + free_bytes = int(free_bytes * safety_fraction) - reserved_memory_in_bytes bytes_per_point, fixed_overhead = linear_peak_memory_model( nao=dm.shape[-1], @@ -83,6 +87,42 @@ def estimate_max_grid_chunk_size( return chunk_size +def estimate_global_raw_feature_buffer_memory( + dm: torch.Tensor, + nfeatures: int, + ngrids: int, + func_deriv: int, +) -> int: + """Estimate full-grid raw-feature storage for global screened evaluation. + + First order keeps sorted and atom-major feature values plus atom-major and + sorted cotangents. Second order additionally keeps an atom-major feature JVP, + an atom-major model Hessian action, and its sorted copy. + + Args: + dm: Density matrix whose leading dimensions determine the spin batches. + nfeatures: Number of raw AO-derived features per grid point. + ngrids: Total number of molecular grid points. + func_deriv: Functional derivative order, either first or second. + + Returns: + Estimated bytes occupied by global raw-feature buffers. + + Raises: + ValueError: If ``func_deriv`` is not first or second order. + """ + match func_deriv: + case 1: + buffer_count = 4 + case 2: + buffer_count = 5 + case _: + raise ValueError("Global screened features support func_deriv 1 or 2") + + batch_size = dm.numel() // (dm.shape[-2] * dm.shape[-1]) + return buffer_count * batch_size * nfeatures * ngrids * 8 + + def linear_peak_memory_model( nao: int, deriv: int, diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index b6ed9b11..9cdf44f8 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -18,7 +18,11 @@ to_cupy, to_numpy, ) -from skala.pyscf.features import chunked_features, generate_features +from skala.pyscf.features import ( + _global_screened_features, + chunked_features, + generate_features, +) def _should_screen_aos(mol: gto.Mole) -> bool: @@ -219,6 +223,54 @@ def __call__( tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) V_xc = torch.zeros_like(dm) + + if screen_aos: + screened_features = _global_screened_features( + mol, + dm, + grids, + features=set(self.func.features), + func_deriv=1, + max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, + safety_fraction=0.8, + ) + atom_major_cotangent = torch.zeros_like( + screened_features.atom_major_raw_features + ) + for atom_slice, grid_slice in screened_features.chunks: + local_raw_features = ( + screened_features.atom_major_raw_features[..., grid_slice] + .detach() + .requires_grad_() + ) + mol_features = screened_features.build_model_chunk( + local_raw_features, atom_slice, grid_slice + ) + E_xc_chunk = self.func.get_exc(mol_features) + (local_cotangent,) = torch.autograd.grad( + E_xc_chunk, + local_raw_features, + torch.ones_like(E_xc_chunk), + ) + atom_major_cotangent[..., grid_slice] = local_cotangent.detach() + tot_dens += ( + (mol_features["density"] * mol_features["grid_weights"]) + .sum(dim=-1) + .detach() + ) + E_xc += E_xc_chunk.detach() + del E_xc_chunk, local_cotangent, local_raw_features, mol_features + + sorted_cotangent = atom_major_cotangent.index_select( + -1, screened_features.forward_permutation + ) + (V_xc,) = torch.autograd.grad( + screened_features.sorted_raw_features, + dm, + sorted_cotangent, + ) + return tot_dens, E_xc, V_xc + for mol_features in chunked_features( mol, dm, @@ -227,7 +279,7 @@ def __call__( func_deriv=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, safety_fraction=0.8, # tends to be faster for large chunks - screen_aos=screen_aos, + screen_aos=False, ): E_xc_chunk = self.func.get_exc(mol_features) (V_xc_chunk,) = torch.autograd.grad( @@ -341,10 +393,9 @@ def gen_response( dm0 = dm0.requires_grad_() screen_aos = _should_screen_aos(ks.mol) - def hessian_vector_product_atom_chunked(dm1: Array) -> Array: - dm1_tensor = self.from_backend(dm1) - hvp_total = torch.zeros_like(dm0) - for mol_features in chunked_features( + screened_features = None + if screen_aos: + screened_features = _global_screened_features( ks.mol, dm0, ks.grids, @@ -353,27 +404,97 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, - safety_fraction=kwargs.get( - "safety_fraction", 0.0 - ), # Force small chunks (single atoms) because it's empirically fastest. - screen_aos=screen_aos, - ): - E_xc_chunk = self.func.get_exc(mol_features) - (V_xc_chunk,) = torch.autograd.grad( - E_xc_chunk, - dm0, - torch.ones_like(E_xc_chunk), - retain_graph=True, - create_graph=True, + safety_fraction=kwargs.get("safety_fraction", 0.0), + ) + if not screened_features.feature_function.only_linear_feats: + raise NotImplementedError( + "Global screened response requires raw features linear in " + "the density matrix." + ) + + def hessian_vector_product_atom_chunked(dm1: Array) -> Array: + dm1_tensor = self.from_backend(dm1) + if screened_features is not None: + atom_major_tangent = screened_features.atom_major_jvp(dm1_tensor) + atom_major_hessian_action = torch.zeros_like( + screened_features.atom_major_raw_features + ) + for atom_slice, grid_slice in screened_features.chunks: + local_raw_features = ( + screened_features.atom_major_raw_features[..., grid_slice] + .detach() + .requires_grad_() + ) + mol_features = screened_features.build_model_chunk( + local_raw_features, atom_slice, grid_slice + ) + E_xc_chunk = self.func.get_exc(mol_features) + (local_gradient,) = torch.autograd.grad( + E_xc_chunk, + local_raw_features, + torch.ones_like(E_xc_chunk), + create_graph=True, + ) + if local_gradient.requires_grad: + (local_hessian_action,) = torch.autograd.grad( + local_gradient, + local_raw_features, + atom_major_tangent[..., grid_slice], + ) + else: + local_hessian_action = torch.zeros_like(local_raw_features) + atom_major_hessian_action[..., grid_slice] = ( + local_hessian_action.detach() + ) + del ( + E_xc_chunk, + local_gradient, + local_hessian_action, + local_raw_features, + mol_features, + ) + + sorted_hessian_action = atom_major_hessian_action.index_select( + -1, screened_features.forward_permutation ) - (hvp_chunk,) = torch.autograd.grad( - V_xc_chunk, + (hvp_total,) = torch.autograd.grad( + screened_features.sorted_raw_features, dm0, - dm1_tensor, + sorted_hessian_action, retain_graph=True, ) - hvp_total += hvp_chunk - del E_xc_chunk, V_xc_chunk, hvp_chunk, mol_features + else: + hvp_total = torch.zeros_like(dm0) + for mol_features in chunked_features( + ks.mol, + dm0, + ks.grids, + features=set(self.func.features), + func_deriv=2, + max_memory_in_mb=ks.max_memory + if dm0.device.type == "cpu" + else None, + safety_fraction=kwargs.get( + "safety_fraction", 0.0 + ), # Force small chunks (single atoms) because it's empirically fastest. + screen_aos=False, + ): + E_xc_chunk = self.func.get_exc(mol_features) + (V_xc_chunk,) = torch.autograd.grad( + E_xc_chunk, + dm0, + torch.ones_like(E_xc_chunk), + retain_graph=True, + create_graph=True, + ) + (hvp_chunk,) = torch.autograd.grad( + V_xc_chunk, + dm0, + dm1_tensor, + retain_graph=True, + ) + hvp_total += hvp_chunk + del E_xc_chunk, V_xc_chunk, hvp_chunk, mol_features v1 = self.to_backend(hvp_total) vj = ks.get_j(ks.mol, dm1, hermi=1) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index b89ae032..ce64190a 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -11,9 +11,12 @@ from skala.pyscf import numint as numint_module from skala.pyscf.features import ( CPU_AO_SCREENING_BLOCK_SIZE, + ChunkEvalBackward, ChunkEvalForward, MGGAFeatureFunction, _active_cpu_aos, + _prepare_spatially_sorted_grids, + _spatial_grid_permutations, chunked_features, ) from skala.pyscf.numint import SkalaNumInt, _should_screen_aos @@ -62,6 +65,112 @@ def test_active_cpu_aos(carbon: gto.Mole) -> None: assert empty.size == 0 +@pytest.mark.parametrize(("ngrids", "block_size"), [(0, 4), (3, 4), (8, 4), (10, 4)]) +def test_spatial_grid_permutations_restore_original_order( + ngrids: int, block_size: int +) -> None: + coords = np.arange(3 * ngrids, dtype=np.float64).reshape(ngrids, 3) + + forward, inverse = _spatial_grid_permutations(coords, block_size) + + assert np.array_equal(np.sort(forward), np.arange(ngrids)) + assert np.array_equal(coords[forward][inverse], coords) + assert all( + len(forward[start : start + block_size]) == block_size + for start in range(0, ngrids - block_size + 1, block_size) + ) + assert len(forward) % block_size == ngrids % block_size + + +def test_spatial_grid_permutations_group_interleaved_clusters() -> None: + block_size = 3 + labels = np.tile(np.arange(4), block_size) + offsets = np.repeat(np.arange(block_size), 4) + coords = np.column_stack( + (100.0 * labels + offsets, np.zeros(labels.size), np.zeros(labels.size)) + ) + + forward, _ = _spatial_grid_permutations(coords, block_size) + + grouped_labels = labels[forward].reshape(-1, block_size) + assert np.all(grouped_labels == grouped_labels[:, :1]) + + +def test_prepare_spatially_sorted_cpu_grids( + carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch +) -> None: + coords = np.arange(18, dtype=np.float64).reshape(6, 3) + weights = np.arange(6, dtype=np.float64) + 10 + grids = dft.Grids(carbon) + grids.coords = coords + grids.weights = weights + forward = np.array([4, 2, 0, 5, 3, 1], dtype=np.int64) + inverse = np.argsort(forward) + non0tab = np.ones((1, carbon.nbas), dtype=np.uint8) + partition_calls = 0 + + def fake_spatial_grid_permutations( + coords_arg: np.ndarray, block_size: int + ) -> tuple[np.ndarray, np.ndarray]: + nonlocal partition_calls + partition_calls += 1 + assert coords_arg is grids.coords + assert block_size == 2 + return forward, inverse + + monkeypatch.setattr( + features_module, + "_spatial_grid_permutations", + fake_spatial_grid_permutations, + ) + + screen_index_calls = 0 + + def fake_make_screen_index( + mol_arg: gto.Mole, sorted_coords: np.ndarray, cutoff: float + ) -> np.ndarray: + nonlocal screen_index_calls + screen_index_calls += 1 + assert mol_arg is carbon + assert np.array_equal(sorted_coords, coords[forward]) + assert cutoff == grids.cutoff + return non0tab + + monkeypatch.setattr(dft.gen_grid, "make_screen_index", fake_make_screen_index) + + sorted_grids, actual_forward, actual_inverse = _prepare_spatially_sorted_grids( + carbon, grids, block_size=2, gpu=False + ) + + assert sorted_grids is not grids + assert np.array_equal(grids.coords, coords) + assert np.array_equal(grids.weights, weights) + assert np.array_equal(sorted_grids.coords, coords[forward]) + assert np.array_equal(sorted_grids.weights, weights[forward]) + assert sorted_grids.non0tab is non0tab + assert actual_forward is forward + assert actual_inverse is inverse + + cached_grids, cached_forward, cached_inverse = _prepare_spatially_sorted_grids( + carbon, grids, block_size=2, gpu=False + ) + + assert cached_grids is sorted_grids + assert cached_forward is forward + assert cached_inverse is inverse + assert partition_calls == 1 + assert screen_index_calls == 1 + + grids.coords = grids.coords.copy() + rebuilt_grids, _, _ = _prepare_spatially_sorted_grids( + carbon, grids, block_size=2, gpu=False + ) + + assert rebuilt_grids is not sorted_grids + assert partition_calls == 2 + assert screen_index_calls == 2 + + @pytest.mark.parametrize("screen_aos", [False, True]) def test_chunked_features_routes_screening( carbon: gto.Mole, @@ -284,7 +393,51 @@ def fake_chunked_features( "grid_weights": torch.ones(1, dtype=dm.dtype), } + class FakeGlobalScreenedFeatures: + def __init__(self, dm: torch.Tensor) -> None: + raw_features = dm.sum().reshape(1, 1) + self.feature_function = MGGAFeatureFunction( + with_density=True, + with_grad=False, + with_kin=False, + ) + self.sorted_raw_features = raw_features + self.atom_major_raw_features = raw_features + self.forward_permutation = torch.tensor([0]) + self.chunks = [(slice(0, 1), slice(0, 1))] + + def atom_major_jvp(self, dm_tangent: torch.Tensor) -> torch.Tensor: + return dm_tangent.sum().reshape(1, 1) + + def build_model_chunk( + self, + raw_features: torch.Tensor, + atom_slice: slice, + grid_slice: slice, + ) -> dict[str, torch.Tensor]: + assert atom_slice == slice(0, 1) + assert grid_slice == slice(0, 1) + return { + "atomic_grid_sizes": torch.tensor([1]), + "density": raw_features.expand(2, 1) / 2, + "grid_weights": torch.ones(1, dtype=raw_features.dtype), + } + + def fake_global_screened_features( + mol: gto.Mole, + dm: torch.Tensor, + grids: object, + features: set[str], + func_deriv: int, + **kwargs: object, + ) -> FakeGlobalScreenedFeatures: + decisions.append(True) + return FakeGlobalScreenedFeatures(dm) + monkeypatch.setattr(numint_module, "chunked_features", fake_chunked_features) + monkeypatch.setattr( + numint_module, "_global_screened_features", fake_global_screened_features + ) numint = SkalaNumInt(QuadraticDensityFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) @@ -457,3 +610,56 @@ def test_cpu_response_dense_screened_equivalence( assert np.allclose( dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 ) + assert np.allclose( + dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 + ) + + +@pytest.mark.parametrize("func_deriv", [1, 2]) +def test_global_screened_ao_traversals_are_independent_of_model_chunks( + monkeypatch: pytest.MonkeyPatch, + func_deriv: int, +) -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + grids = _minimal_atom_grid(mol) + atom_grid_size = grids.weights.size // mol.natm + monkeypatch.setattr( + features_module, + "estimate_max_grid_chunk_size", + lambda *args, **kwargs: atom_grid_size, + ) + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + + forward_calls = 0 + backward_calls = 0 + original_forward_apply = ChunkEvalForward.apply + original_backward_apply = ChunkEvalBackward.apply + + def counting_forward_apply(*args: object) -> torch.Tensor: + nonlocal forward_calls + forward_calls += 1 + return original_forward_apply(*args) + + def counting_backward_apply(*args: object) -> torch.Tensor: + nonlocal backward_calls + backward_calls += 1 + return original_backward_apply(*args) + + monkeypatch.setattr(ChunkEvalForward, "apply", counting_forward_apply) + monkeypatch.setattr(ChunkEvalBackward, "apply", counting_backward_apply) + functional = QuadraticDensityFunctional() + numint = SkalaNumInt(functional) + + if func_deriv == 1: + dm = dft.RKS(mol).get_init_guess() + numint.nr_rks(mol, grids, None, dm) + assert forward_calls == 1 + else: + ks = FakeKS(mol, grids) + response = numint.gen_response( + np.eye(mol.nao_nr()), np.ones(mol.nao_nr()), ks=ks + ) + response(np.eye(mol.nao_nr())) + assert forward_calls == 2 + + assert backward_calls == 1 diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 983b07fa..808a9a8c 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -8,6 +8,8 @@ from pyscf.dft import numint as pyscf_numint from torch.utils.dlpack import from_dlpack +pytestmark = pytest.mark.gpu + if not torch.cuda.is_available(): pytest.skip( "Skipping gpu4pyscf AO screening tests, because CUDA is not available.", @@ -25,7 +27,11 @@ from skala.functional.base import ExcFunctionalBase from skala.gpu4pyscf import SkalaKS from skala.pyscf.backend import dft_gpu -from skala.pyscf.features import ChunkEvalForward, MGGAFeatureFunction +from skala.pyscf.features import ( + ChunkEvalForward, + MGGAFeatureFunction, + _prepare_spatially_sorted_grids, +) from skala.pyscf.numint import SkalaNumInt CARBON_CHAIN = """ @@ -45,10 +51,75 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: return (mol["density"].square() * mol["grid_weights"]).sum() +class QuadraticMGGAFunctional(ExcFunctionalBase): + def __init__(self) -> None: + super().__init__() + self.features = [ + "atomic_grid_sizes", + "density", + "grad", + "kin", + "grid_weights", + ] + + def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + energy_density = ( + mol["density"].square() + + mol["grad"].square().sum(dim=-2) + + mol["kin"].square() + ) + return (energy_density * mol["grid_weights"]).sum() + + def _to_numpy(value: object) -> np.ndarray: return cupy.asnumpy(value) if isinstance(value, cupy.ndarray) else np.asarray(value) +def test_prepare_spatially_sorted_gpu_grids() -> None: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + coords = cupy.asarray( + [[20.0, 0.0, 0.0], [0.0, 0.0, 0.0], [21.0, 0.0, 0.0], [1.0, 0.0, 0.0]] + ) + weights = cupy.arange(coords.shape[0], dtype=cupy.float64) + grids = dft_gpu.Grids(mol) + grids.coords = coords + grids.weights = weights + original_screening_cache = cupy.arange(1) + grids._non0ao_idx = original_screening_cache + + sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( + mol, grids, block_size=2, gpu=True + ) + + assert sorted_grids is not grids + assert grids.coords is coords + assert grids.weights is weights + assert grids._non0ao_idx is original_screening_cache + assert isinstance(sorted_grids.coords, cupy.ndarray) + assert isinstance(sorted_grids.weights, cupy.ndarray) + assert sorted_grids._non0ao_idx is None + assert np.array_equal( + cupy.asnumpy(sorted_grids.coords), cupy.asnumpy(coords)[forward] + ) + assert np.array_equal( + cupy.asnumpy(sorted_grids.weights), cupy.asnumpy(weights)[forward] + ) + assert np.array_equal( + cupy.asnumpy(sorted_grids.coords)[inverse], cupy.asnumpy(coords) + ) + + sorted_screening_cache = object() + sorted_grids._non0ao_idx = sorted_screening_cache + cached_grids, cached_forward, cached_inverse = _prepare_spatially_sorted_grids( + mol, grids, block_size=2, gpu=True + ) + + assert cached_grids is sorted_grids + assert cached_forward is forward + assert cached_inverse is inverse + assert cached_grids._non0ao_idx is sorted_screening_cache + + @pytest.mark.parametrize("unrestricted", [False, True]) def test_gpu_rks_uks_dense_screened_equivalence( monkeypatch: pytest.MonkeyPatch, @@ -116,30 +187,91 @@ def test_gpu_response_dense_screened_equivalence( rtol=1e-9, atol=1e-10, ) + assert np.allclose( + _to_numpy(dense_response(dm1)), + _to_numpy(screened_response(dm1)), + rtol=1e-9, + atol=1e-10, + ) + + +def test_gpu_uks_response_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) + ks.grids.level = 0 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + mo_coeff = cupy.stack((cupy.eye(mol.nao_nr()), cupy.eye(mol.nao_nr()))) + mo_occ = cupy.ones((2, mol.nao_nr())) + dm1 = cupy.ones((2, mol.nao_nr(), mol.nao_nr())) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + np.testing.assert_allclose( + _to_numpy(screened_response(dm1)), + _to_numpy(dense_response(dm1)), + rtol=1e-9, + atol=1e-10, + ) + + +def test_gpu_multiblock_mgga_response_dense_screened_equivalence( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) + ks = SkalaKS(mol, xc=QuadraticMGGAFunctional(), with_dftd3=False) + ks.grids.level = 1 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + assert ks.grids.weights.size > dft_gpu.numint.MIN_BLK_SIZE + mo_coeff = cupy.eye(mol.nao_nr()) + mo_occ = cupy.ones(mol.nao_nr()) + dm1 = cupy.eye(mol.nao_nr()) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + + np.testing.assert_allclose( + _to_numpy(screened_response(dm1)), + _to_numpy(dense_response(dm1)), + rtol=1e-9, + atol=1e-9, + ) def test_gpu_screened_skala_matches_cpu_on_carbon_chain( monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> None: - """Catch inaccurate GPU AO screening on spatially diffuse grid blocks. + """Prevent inaccurate GPU AO screening on spatially diffuse grid blocks. GPU4PySCF builds one active-shell mask for each fixed-size coordinate block. That screening is reliable only when the points in a block are spatially local enough - for the sampled AO values to represent the whole block. Skala currently supplies - an unsorted, atom-major grid, so one GPU block can span a large region around an - atom. This is especially problematic for the AO derivatives used by Skala: an AO - value can be small at the sampled points even though its gradient still makes a - significant contribution. The linear carbon chain and large def2-QZVPP basis - expose this failure in a reasonably small integration test. + for the sampled AO values to represent the whole block. Passing Skala's unsorted, + atom-major grid directly would allow one GPU block to span a large region around + an atom. This is especially problematic for the AO derivatives used by Skala: an + AO value can be small at the sampled points even though its gradient still makes + a significant contribution. The implementation therefore partitions the whole + molecular grid into exact-size spatial blocks for AO evaluation, then restores + atom-major feature order before evaluating the model. The linear carbon chain and + large def2-QZVPP basis expose regressions in that ordering on a reasonably small + system. The CPU and GPU calculations use identical coordinates, weights, density matrix, and Skala 1.1 model. CPU AO evaluation is deliberately forced dense to provide an independent reference, while GPU AO evaluation is deliberately forced through screening. Comparing the particle count, XC energy, and complete XC potential matrix verifies the full feature and VJP path; the potential is particularly - sensitive to omitted derivative contributions. The test should pass once GPU - blocks are spatially grouped before their screening masks are constructed. + sensitive to omitted derivative contributions. """ mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) cpu_grids = dft.Grids(mol) @@ -187,9 +319,9 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( ) assert ( np.isclose(cpu_result[0], gpu_result[0], rtol=1e-10, atol=1e-11) - and np.isclose(cpu_result[1], gpu_result[1], rtol=1e-10, atol=1e-11) - and vxc_max_abs_difference < 2e-9 - and vxc_relative_l2_difference < 1e-8 + and np.isclose(cpu_result[1], gpu_result[1], rtol=1e-8, atol=1e-9) + and vxc_max_abs_difference < 2e-7 + and vxc_relative_l2_difference < 1e-7 ), ( f"N: cpu={cpu_result[0]:.16g}, gpu={gpu_result[0]:.16g}, " f"abs_diff={abs(cpu_result[0] - gpu_result[0]):.3e}; " diff --git a/tests/test_gpu4pyscf_classes.py b/tests/test_gpu4pyscf_classes.py index 7aa28865..19920d59 100644 --- a/tests/test_gpu4pyscf_classes.py +++ b/tests/test_gpu4pyscf_classes.py @@ -4,6 +4,8 @@ import torch from pyscf import gto +pytestmark = pytest.mark.gpu + if not torch.cuda.is_available(): pytest.skip( "Skipping gpu4pyscf classes tests, because CUDA is not available.", diff --git a/tests/test_gpu4pyscf_gradients.py b/tests/test_gpu4pyscf_gradients.py index 03e3c0cd..f2c5e365 100644 --- a/tests/test_gpu4pyscf_gradients.py +++ b/tests/test_gpu4pyscf_gradients.py @@ -3,6 +3,8 @@ import pytest import torch +pytestmark = pytest.mark.gpu + if not torch.cuda.is_available(): pytest.skip( "Skipping gpu4pyscf gradients tests, because CUDA is not available.", diff --git a/tests/test_memory_estimators.py b/tests/test_memory_estimators.py new file mode 100644 index 00000000..f95850ab --- /dev/null +++ b/tests/test_memory_estimators.py @@ -0,0 +1,57 @@ +import pytest +import torch + +from skala.pyscf.memory_estimators import ( + estimate_global_raw_feature_buffer_memory, + estimate_max_grid_chunk_size, + linear_peak_memory_model, +) + + +@pytest.mark.parametrize( + ("dm_shape", "func_deriv", "buffer_count"), + [((10, 10), 1, 4), ((2, 10, 10), 1, 4), ((10, 10), 2, 5)], +) +def test_global_raw_feature_buffer_memory( + dm_shape: tuple[int, ...], func_deriv: int, buffer_count: int +) -> None: + dm = torch.zeros(dm_shape, dtype=torch.float64) + nfeatures = 5 + ngrids = 123 + batch_size = dm.numel() // (dm.shape[-2] * dm.shape[-1]) + + actual = estimate_global_raw_feature_buffer_memory( + dm, nfeatures, ngrids, func_deriv + ) + + assert actual == buffer_count * batch_size * nfeatures * ngrids * 8 + + +def test_global_raw_feature_buffer_memory_rejects_unsupported_order() -> None: + with pytest.raises(ValueError, match="func_deriv 1 or 2"): + estimate_global_raw_feature_buffer_memory( + torch.eye(2, dtype=torch.float64), 1, 1, func_deriv=0 + ) + + +def test_reserved_memory_reduces_grid_chunk_size() -> None: + dm = torch.eye(10, dtype=torch.float64) + bytes_per_point, _ = linear_peak_memory_model(nao=10, deriv=1, func_deriv=1) + base_chunk_size = estimate_max_grid_chunk_size( + dm, + deriv=1, + max_memory_in_mb=100, + safety_fraction=1.0, + func_deriv=1, + ) + reserved_points = 123 + reserved_chunk_size = estimate_max_grid_chunk_size( + dm, + deriv=1, + max_memory_in_mb=100, + safety_fraction=1.0, + func_deriv=1, + reserved_memory_in_bytes=int(bytes_per_point * reserved_points), + ) + + assert base_chunk_size - reserved_chunk_size == reserved_points From abee701e6932b4870579eb8077eb2e26c30755a9 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 3 Aug 2026 10:57:06 +0200 Subject: [PATCH 06/39] add nuclear gradient test --- src/skala/ase/__init__.py | 2 +- src/skala/functional/__init__.py | 6 +- src/skala/functional/load.py | 4 +- src/skala/functional/utils/irreps.py | 2 +- src/skala/gpu4pyscf/gradients.py | 2 +- src/skala/pyscf/features.py | 4 +- src/skala/pyscf/gradients.py | 2 +- tests/test_gpu4pyscf_ao_screening.py | 10 +-- tests/test_gpu4pyscf_classes.py | 10 +-- tests/test_gpu4pyscf_gradients.py | 107 ++++++++++++++++++++++++--- 10 files changed, 118 insertions(+), 31 deletions(-) diff --git a/src/skala/ase/__init__.py b/src/skala/ase/__init__.py index ea3cd2a1..4bb5c6ae 100644 --- a/src/skala/ase/__init__.py +++ b/src/skala/ase/__init__.py @@ -8,6 +8,6 @@ ) from e -from skala.ase.calculator import Skala # noqa: F401 +from skala.ase.calculator import Skala __all__ = ["Skala"] diff --git a/src/skala/functional/__init__.py b/src/skala/functional/__init__.py index 36c64bc3..4be42f1f 100644 --- a/src/skala/functional/__init__.py +++ b/src/skala/functional/__init__.py @@ -30,9 +30,6 @@ ) __all__ = [ - "ExcFunctionalBase", - "SkalaFunctional", - "TracedFunctional", "LDA", "PBE", "R2SCAN", @@ -40,6 +37,9 @@ "SCAN", "SPW92", "TPSS", + "ExcFunctionalBase", + "SkalaFunctional", + "TracedFunctional", "load_functional", ] diff --git a/src/skala/functional/load.py b/src/skala/functional/load.py index 1edcb9c8..ca3baa7e 100644 --- a/src/skala/functional/load.py +++ b/src/skala/functional/load.py @@ -125,7 +125,7 @@ def load( raise RuntimeError( "metadata in traced functional extra_files does not have the correct format (dict)." ) - if not all([isinstance(key, str) for key in _metadata]): + if not all(isinstance(key, str) for key in _metadata): raise RuntimeError("metadata keys in traced functional must be strings.") metadata = cast(dict[str, Any], _metadata) @@ -134,7 +134,7 @@ def load( raise RuntimeError( "features in traced functional extra_files does not have the correct format (list)." ) - if not all([isinstance(feat, str) for feat in _features]): + if not all(isinstance(feat, str) for feat in _features): raise RuntimeError( "features in traced functional must be a list of strings." ) diff --git a/src/skala/functional/utils/irreps.py b/src/skala/functional/utils/irreps.py index 88eddf8c..22dca47f 100644 --- a/src/skala/functional/utils/irreps.py +++ b/src/skala/functional/utils/irreps.py @@ -111,7 +111,7 @@ def __getitem__(self, i: int) -> int: class MulIr: - __slots__ = ("_mul", "_ir") + __slots__ = ("_ir", "_mul") _mul: int _ir: Irrep diff --git a/src/skala/gpu4pyscf/gradients.py b/src/skala/gpu4pyscf/gradients.py index 54aa23ef..7d1fc0ad 100644 --- a/src/skala/gpu4pyscf/gradients.py +++ b/src/skala/gpu4pyscf/gradients.py @@ -93,7 +93,7 @@ def veff_and_expl_nuc_grad( nuc_feat_names = list(nuc_grad_feats) # ensure specific order nuc_feat_tensors = [mol_feats[feat] for feat in nuc_feat_names] other_feats = { - feat: mol_feats[feat] for feat in mol_feats.keys() if feat not in nuc_grad_feats + feat: mol_feats[feat] for feat in mol_feats if feat not in nuc_grad_feats } def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 717822d6..5eb93315 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -289,7 +289,7 @@ def chunked_features( raise ValueError("Grids must be built before generating chunked features.") # if dm is a 3D tensor, then we have a spin-polarized system - with_spin = True if len(dm.shape) == 3 else False + with_spin = len(dm.shape) == 3 grid_features = get_grid_features(mol, dm, grids, features) with_mgga_feature = ( @@ -495,7 +495,7 @@ def generate_features( features = features or DEFAULT_FEATURES_SET # if dm is a 3D tensor, then we have a spin-polarized system - with_spin = True if len(dm.shape) == 3 else False + with_spin = len(dm.shape) == 3 if gpu and dm.device.type != "cuda": raise ValueError("Density matrix must be on the GPU when gpu=True.") diff --git a/src/skala/pyscf/gradients.py b/src/skala/pyscf/gradients.py index 02c0a97b..a1037d36 100644 --- a/src/skala/pyscf/gradients.py +++ b/src/skala/pyscf/gradients.py @@ -88,7 +88,7 @@ def veff_and_expl_nuc_grad( nuc_feat_names = list(nuc_grad_feats) # ensure specific order nuc_feat_tensors = [mol_feats[feat] for feat in nuc_feat_names] other_feats = { - feat: mol_feats[feat] for feat in mol_feats.keys() if feat not in nuc_grad_feats + feat: mol_feats[feat] for feat in mol_feats if feat not in nuc_grad_feats } def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 808a9a8c..aa36c242 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -24,15 +24,15 @@ allow_module_level=True, ) -from skala.functional.base import ExcFunctionalBase -from skala.gpu4pyscf import SkalaKS -from skala.pyscf.backend import dft_gpu -from skala.pyscf.features import ( +from skala.functional.base import ExcFunctionalBase # noqa: E402 +from skala.gpu4pyscf import SkalaKS # noqa: E402 +from skala.pyscf.backend import dft_gpu # noqa: E402 +from skala.pyscf.features import ( # noqa: E402 ChunkEvalForward, MGGAFeatureFunction, _prepare_spatially_sorted_grids, ) -from skala.pyscf.numint import SkalaNumInt +from skala.pyscf.numint import SkalaNumInt # noqa: E402 CARBON_CHAIN = """ C 0.0 0.0 0.0 diff --git a/tests/test_gpu4pyscf_classes.py b/tests/test_gpu4pyscf_classes.py index 19920d59..a46f7cbd 100644 --- a/tests/test_gpu4pyscf_classes.py +++ b/tests/test_gpu4pyscf_classes.py @@ -12,11 +12,11 @@ allow_module_level=True, ) -from skala.functional.base import ExcFunctionalBase -from skala.gpu4pyscf import SkalaKS -from skala.gpu4pyscf.dft import SkalaRKS, SkalaUKS -from skala.gpu4pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient -from skala.gpu4pyscf.grids import UnsortableGrids +from skala.functional.base import ExcFunctionalBase # noqa: E402 +from skala.gpu4pyscf import SkalaKS # noqa: E402 +from skala.gpu4pyscf.dft import SkalaRKS, SkalaUKS # noqa: E402 +from skala.gpu4pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient # noqa: E402 +from skala.gpu4pyscf.grids import UnsortableGrids # noqa: E402 @pytest.fixture(params=["skala-1.0", "skala-1.1"]) diff --git a/tests/test_gpu4pyscf_gradients.py b/tests/test_gpu4pyscf_gradients.py index f2c5e365..b9eb1f69 100644 --- a/tests/test_gpu4pyscf_gradients.py +++ b/tests/test_gpu4pyscf_gradients.py @@ -18,21 +18,41 @@ allow_module_level=True, ) -from _ridders import num_grad_ridders -from gpu4pyscf import dft, scf -from pyscf import gto -from test_pyscf_gradients import FULL_GRAD_REF - -from skala.functional.base import ExcFunctionalBase -from skala.gpu4pyscf import SkalaKS -from skala.gpu4pyscf.gradients import ( +from _ridders import num_grad_ridders # noqa: E402 +from gpu4pyscf import dft, scf # noqa: E402 +from pyscf import gto # noqa: E402 +from pyscf.dft import numint as pyscf_numint # noqa: E402 +from test_pyscf_gradients import FULL_GRAD_REF # noqa: E402 + +from skala.functional.base import ExcFunctionalBase # noqa: E402 +from skala.gpu4pyscf import SkalaKS # noqa: E402 +from skala.gpu4pyscf.gradients import ( # noqa: E402 SkalaRKSGradient, SkalaUKSGradient, nuc_grad_from_veff, veff_and_expl_nuc_grad, ) -from skala.pyscf.features import generate_features -from skala.utils import torch_allocator +from skala.pyscf import SkalaKS as CpuSkalaKS # noqa: E402 +from skala.pyscf.features import generate_features # noqa: E402 +from skala.pyscf.gradients import SkalaRKSGradient as CpuSkalaRKSGradient # noqa: E402 +from skala.pyscf.numint import _should_screen_aos # noqa: E402 +from skala.utils import torch_allocator # noqa: E402 + +H2_SKALA_1_1_GRAD_REF = torch.tensor( + [ + [ + 2.6170957571276746e-10, + 2.1217813541405875e-10, + -1.345246115431109e-02, + ], + [ + -2.6170957571276746e-10, + -2.1217813541405844e-10, + 1.3452461154311535e-02, + ], + ], + dtype=torch.float64, +) def test_torch_allocator_is_active_after_import() -> None: @@ -441,6 +461,73 @@ def test_full_grad( ) +def test_nuclear_gradient_cpu_gpu_dense_screened_agree( + monkeypatch: pytest.MonkeyPatch, + load_functional_cached: Callable[..., ExcFunctionalBase | str], +) -> None: + """Compare complete nuclear gradients across backend and SCF screening routes. + + AO screening controls the SCF feature/Vxc path that produces the converged + density. The analytic nuclear-gradient contraction then uses the same atom-major + implementation for the dense and screened densities on each backend. + """ + cpu_functional = load_functional_cached("skala-1.1") + gpu_functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(cpu_functional, ExcFunctionalBase) + assert isinstance(gpu_functional, ExcFunctionalBase) + + gradients: dict[str, torch.Tensor] = {} + for backend, functional in ( + ("cpu", cpu_functional), + ("gpu", gpu_functional), + ): + for screened in (False, True): + mol = gto.M( + atom="H 0 0 0; H 0 0 0.74", + basis="sto-3g", + verbose=0, + ) + monkeypatch.setattr( + pyscf_numint, + "SWITCH_SIZE", + mol.nao_nr() - int(screened), + ) + assert _should_screen_aos(mol) is screened + if backend == "cpu": + mean_field = CpuSkalaKS(mol, xc=functional, with_dftd3=False) + gradient_type = CpuSkalaRKSGradient + else: + mean_field = SkalaKS(mol, xc=functional, with_dftd3=False) + gradient_type = SkalaRKSGradient + mean_field.grids.level = 0 + mean_field.grids.build(mol, sort_grids=False) + mean_field.conv_tol = 1e-10 + mean_field.kernel() + assert mean_field.converged + gradient = gradient_type(mean_field).kernel() + route = f"{backend}-{'screened' if screened else 'dense'}" + gradients[route] = torch.from_numpy(gradient) + + for route, gradient in gradients.items(): + torch.testing.assert_close( + gradient, + H2_SKALA_1_1_GRAD_REF, + rtol=1e-7, + atol=1e-8, + msg=f"{route} does not match the stored nuclear-gradient reference", + ) + + cpu_dense = gradients["cpu-dense"] + for route, gradient in gradients.items(): + torch.testing.assert_close( + gradient, + cpu_dense, + rtol=1e-7, + atol=1e-8, + msg=f"{route} does not match cpu-dense", + ) + + def test_cuda_kernel_memory_stability() -> None: """Checks that repeated calls do not increase Torch's allocated CUDA memory.""" From 9428b67e2e62d29bfb6b4818591ebdd1ec7ed80c Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 3 Aug 2026 23:54:45 +0200 Subject: [PATCH 07/39] add memory benchmark --- pyproject.toml | 1 + src/skala/pyscf/numint.py | 8 ++ tests/test_ao_screening_benchmark.py | 167 +++++++++++++++++++++++++-- 3 files changed, 165 insertions(+), 11 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index e3ba6a9e..bda1dd73 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -32,6 +32,7 @@ dependencies = [ optional-dependencies.dev = [ "pre-commit", "mypy", + "memray", "pytest", "pytest-benchmark", "pytest-cov", diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 9cdf44f8..ee367c69 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -234,10 +234,12 @@ def __call__( max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, safety_fraction=0.8, ) + # Store only full-grid feature cotangents; model activations remain chunk-local. atom_major_cotangent = torch.zeros_like( screened_features.atom_major_raw_features ) for atom_slice, grid_slice in screened_features.chunks: + # Break the reorder graph so model backprop retains only this chunk. local_raw_features = ( screened_features.atom_major_raw_features[..., grid_slice] .detach() @@ -261,9 +263,11 @@ def __call__( E_xc += E_xc_chunk.detach() del E_xc_chunk, local_cotangent, local_raw_features, mol_features + # Reorder detached cotangents explicitly instead of backpropagating through it. sorted_cotangent = atom_major_cotangent.index_select( -1, screened_features.forward_permutation ) + # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. (V_xc,) = torch.autograd.grad( screened_features.sorted_raw_features, dm, @@ -416,10 +420,12 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: dm1_tensor = self.from_backend(dm1) if screened_features is not None: atom_major_tangent = screened_features.atom_major_jvp(dm1_tensor) + # Store the full-grid model Hessian action, not per-chunk model graphs. atom_major_hessian_action = torch.zeros_like( screened_features.atom_major_raw_features ) for atom_slice, grid_slice in screened_features.chunks: + # Isolate the second-order model graph to the current atomic chunk. local_raw_features = ( screened_features.atom_major_raw_features[..., grid_slice] .detach() @@ -454,9 +460,11 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: mol_features, ) + # Restore box order after all chunk-local Hessian actions are detached. sorted_hessian_action = atom_major_hessian_action.index_select( -1, screened_features.forward_permutation ) + # The custom VJP traverses AO blocks sequentially and retains no AO graph. (hvp_total,) = torch.autograd.grad( screened_features.sorted_raw_features, dm0, diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 207dab7b..fe9b0c9b 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -1,6 +1,11 @@ from __future__ import annotations +import multiprocessing as mp +import tempfile +import traceback from collections.abc import Callable, Iterator +from multiprocessing.connection import Connection +from pathlib import Path from typing import NamedTuple, cast import numpy as np @@ -10,11 +15,13 @@ from pyscf.dft import numint as pyscf_numint from pytest_benchmark.fixture import BenchmarkFixture +from skala.functional import load_functional from skala.functional.base import ExcFunctionalBase from skala.pyscf.numint import SkalaNumInt, _should_screen_aos THREAD_COUNT = 4 MAX_MEMORY_MB = 2000 +MEMORY_WORKER_TIMEOUT_SECONDS = 240 NAPHTHALENE = """ C -1.2280 0.7090 0.0 @@ -110,6 +117,24 @@ class BenchmarkCase(NamedTuple): numint: SkalaNumInt[np.ndarray] +BENCHMARK_SPECS = [ + pytest.param(BenchmarkSpec("naphthalene", NAPHTHALENE), id="naphthalene"), + pytest.param(BenchmarkSpec("anthracene", ANTHRACENE), id="anthracene"), + pytest.param(BenchmarkSpec("tetracene", TETRACENE), id="tetracene"), +] + + +def _make_benchmark_case( + spec: BenchmarkSpec, functional: ExcFunctionalBase +) -> BenchmarkCase: + mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) + grids = dft.Grids(mol) + grids.level = 1 + grids.build(sort_grids=False) + dm = dft.RKS(mol).get_init_guess() + return BenchmarkCase(mol, grids, dm, SkalaNumInt(functional)) + + @pytest.fixture(scope="module") def fixed_cpu_threads() -> Iterator[None]: previous_pyscf_threads = lib.num_threads() @@ -125,11 +150,7 @@ def fixed_cpu_threads() -> Iterator[None]: @pytest.fixture( scope="module", - params=[ - pytest.param(BenchmarkSpec("naphthalene", NAPHTHALENE), id="naphthalene"), - pytest.param(BenchmarkSpec("anthracene", ANTHRACENE), id="anthracene"), - pytest.param(BenchmarkSpec("tetracene", TETRACENE), id="tetracene"), - ], + params=BENCHMARK_SPECS, ) def benchmark_case( request: pytest.FixtureRequest, @@ -137,14 +158,9 @@ def benchmark_case( load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> BenchmarkCase: spec = cast(BenchmarkSpec, request.param) - mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) - grids = dft.Grids(mol) - grids.level = 1 - grids.build(sort_grids=False) - dm = dft.RKS(mol).get_init_guess() functional = load_functional_cached("skala-1.1") assert isinstance(functional, ExcFunctionalBase) - return BenchmarkCase(mol, grids, dm, SkalaNumInt(functional)) + return _make_benchmark_case(spec, functional) @pytest.fixture @@ -182,6 +198,100 @@ def _benchmark_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: ) +def _run_gpu_xc(spec: BenchmarkSpec, screened: bool) -> int: + from skala.gpu4pyscf import SkalaKS + + mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) + functional = load_functional("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(functional, ExcFunctionalBase) + ks = SkalaKS(mol, xc=functional, with_dftd3=False) + ks.grids.level = 1 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + dm = ks.get_init_guess() + if not screened: + pyscf_numint.SWITCH_SIZE = mol.nao_nr() + assert _should_screen_aos(mol) is screened + + torch.cuda.synchronize() + torch.cuda.empty_cache() + baseline_bytes = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + ks._numint.nr_rks(mol, ks.grids, None, dm, max_memory=MAX_MEMORY_MB) + torch.cuda.synchronize() + return torch.cuda.max_memory_allocated() - baseline_bytes + + +def _memory_worker( + spec: BenchmarkSpec, + screened: bool, + backend: str, + control: Connection, +) -> None: + """Measure one route in an isolated process and return its allocation peak.""" + try: + lib.num_threads(THREAD_COUNT) + torch.set_num_threads(THREAD_COUNT) + if backend == "cpu": + import memray + + functional = load_functional("skala-1.1") + assert isinstance(functional, ExcFunctionalBase) + case = _make_benchmark_case(spec, functional) + if not screened: + pyscf_numint.SWITCH_SIZE = case.mol.nao_nr() + assert _should_screen_aos(case.mol) is screened + + with tempfile.TemporaryDirectory() as tmpdir: + profile_path = Path(tmpdir) / "allocations.bin" + with memray.Tracker(profile_path): + _run_xc(case) + peak_bytes = memray.FileReader(profile_path).metadata.peak_memory + elif backend == "cuda": + peak_bytes = _run_gpu_xc(spec, screened) + else: + raise ValueError(f"Unknown memory benchmark backend: {backend}") + control.send(("done", peak_bytes)) + except Exception: # noqa: BLE001 - forward worker failures to the parent + control.send(("error", traceback.format_exc())) + finally: + control.close() + + +def _measure_peak_memory(spec: BenchmarkSpec, screened: bool, backend: str) -> int: + """Return peak allocations for one isolated CPU or CUDA evaluation.""" + context = mp.get_context("spawn") + control, worker_control = context.Pipe() + worker = context.Process( + target=_memory_worker, + args=(spec, screened, backend, worker_control), + ) + worker.start() + worker_control.close() + + try: + if not control.poll(MEMORY_WORKER_TIMEOUT_SECONDS): + raise TimeoutError("Memory benchmark worker timed out") + status, detail = cast(tuple[str, int | str], control.recv()) + if status == "error": + raise RuntimeError(f"Memory benchmark worker failed:\n{detail}") + if status != "done": + raise RuntimeError(f"Unexpected memory benchmark status: {status}") + + worker.join() + if worker.exitcode != 0: + raise RuntimeError( + f"Memory benchmark worker exited with code {worker.exitcode}" + ) + assert isinstance(detail, int) + return detail + finally: + if worker.is_alive(): + worker.terminate() + worker.join() + control.close() + + @pytest.mark.profiling def test_screened_and_dense_values_agree( benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch @@ -224,6 +334,41 @@ def test_without_ao_screening_by_patching_threshold( _benchmark_xc(benchmark, dense_case) +@pytest.mark.profiling +@pytest.mark.parametrize("spec", BENCHMARK_SPECS) +@pytest.mark.parametrize( + "backend", + ["cpu", pytest.param("cuda", marks=pytest.mark.gpu)], +) +def test_screened_and_dense_peak_memory( + spec: BenchmarkSpec, + backend: str, + record_property: Callable[[str, object], None], +) -> None: + if backend == "cuda": + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + pytest.importorskip("cupy") + pytest.importorskip("gpu4pyscf") + + screened_peak_bytes = _measure_peak_memory(spec, screened=True, backend=backend) + dense_peak_bytes = _measure_peak_memory(spec, screened=False, backend=backend) + mib = 1024**2 + screened_peak_mib = screened_peak_bytes / mib + dense_peak_mib = dense_peak_bytes / mib + peak_ratio = screened_peak_bytes / dense_peak_bytes + + record_property("backend", backend) + record_property("screened_peak_allocations_mib", screened_peak_mib) + record_property("dense_peak_allocations_mib", dense_peak_mib) + record_property("screened_to_dense_peak_ratio", peak_ratio) + print( + f"\n{spec.name} {backend} peak allocations: " + f"screened={screened_peak_mib:.1f} MiB, dense={dense_peak_mib:.1f} MiB, " + f"ratio={peak_ratio:.3f}" + ) + + @pytest.mark.profiling def test_profile_with_natural_ao_screening( screened_case: BenchmarkCase, From 8ad6f335fc1ed6a04b5b85a0d59e381d03dfff4b Mon Sep 17 00:00:00 2001 From: jenswehner Date: Tue, 4 Aug 2026 02:07:31 +0200 Subject: [PATCH 08/39] remove chunked case --- src/skala/pyscf/features.py | 446 ++++++++------------------- src/skala/pyscf/memory_estimators.py | 2 +- src/skala/pyscf/numint.py | 332 ++++++++------------ tests/test_ao_screening.py | 202 +----------- tests/test_ao_screening_benchmark.py | 186 ++++++++--- 5 files changed, 428 insertions(+), 740 deletions(-) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 5eb93315..4448d8db 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -212,33 +212,6 @@ def _prepare_spatially_sorted_grids( return sorted_grids, forward, inverse -def _spatially_group_atom_grids( - mol: gto.Mole, coords: np.ndarray, atomic_grid_sizes: Tensor -) -> np.ndarray: - """Build a spatial grid permutation independently within each atom. - - Makes screening much more effective as points are spatially grouped - within each atom, while preserving the original atom order and atom boundaries. - - Args: - mol: Molecule used by PySCF to define the spatial grouping boxes. - coords: Atom-major grid coordinates to group. - atomic_grid_sizes: Number of consecutive grid points owned by each atom. - - Returns: - A permutation that spatially groups each atom's points while preserving - atom order and atom boundaries. - """ - sort_indices = [] - start = 0 - for size in atomic_grid_sizes.tolist(): - stop = start + size - atom_sort_indices = dft.gen_grid.arg_group_grids(mol, coords[start:stop]) - sort_indices.append(start + atom_sort_indices) - start = stop - return np.concatenate(sort_indices) - - def maybe_expand_and_divide( feature: torch.Tensor, expand: bool, divisor: float ) -> torch.Tensor: @@ -251,162 +224,6 @@ def maybe_expand_and_divide( return feature -def chunked_features( - mol: gto.Mole, - dm: Tensor, - grids: Grid, - features: set[str], - func_deriv: int, - max_memory_in_mb: int | None = None, - safety_fraction: float = 0.8, - compile_feature_function: bool = False, - screen_aos: bool = False, -) -> Iterator[dict[str, Tensor]]: - """ - Chunked feature generation for a given molecule. The density features are generated in chunks to avoid memory issues. - - Input: - mol: The molecule for which to generate features. - dm: The density matrix. - grids: The grid points. - features: The set of features to generate. - func_deriv: The order of the functional derivative. - max_memory_in_mb: The maximum memory to use for each chunk in megabytes (MB). If None, the maximum memory is determined automatically. - safety_fraction: The fraction of the available memory to use for each chunk. - compile_feature_function: Whether to compile the feature function. - screen_aos: Whether to evaluate each atom chunk through backend AO screening. - - Yields: - A dictionary of features for each chunk. - """ - - features = features or DEFAULT_FEATURES_SET - if "atomic_grid_sizes" not in features: - raise ValueError( - "The current implementation of chunked_features requires 'atomic_grid_sizes' to be in the requested features." - ) - if grids.coords is None or grids.weights is None: - raise ValueError("Grids must be built before generating chunked features.") - - # if dm is a 3D tensor, then we have a spin-polarized system - with_spin = len(dm.shape) == 3 - - grid_features = get_grid_features(mol, dm, grids, features) - with_mgga_feature = ( - "density" in features - or "grad" in features - or "kin" in features - or "lapl" in features - ) - - # Build the feature function once; it is reused for every chunk. - ff = None - if with_mgga_feature: - ff = MGGAFeatureFunction( - with_density="density" in features, - with_grad="grad" in features, - with_kin="kin" in features, - with_lapl="lapl" in features, - ) - - # Determine the chunk size automatically when not explicitly provided. - if ff is not None: - max_grid_chunk_size = estimate_max_grid_chunk_size( - dm=dm, - deriv=ff.deriv, - max_memory_in_mb=max_memory_in_mb, - safety_fraction=safety_fraction, - func_deriv=func_deriv, - ) - if max_grid_chunk_size < ( - max_atom_grid := int(grid_features["atomic_grid_sizes"].max().item()) - ): - LOG.warning( - f"Adjusted chunk size {max_grid_chunk_size} to match the largest atomic grid {max_atom_grid}. Hope for no OOM." - ) - max_grid_chunk_size = max_atom_grid - else: # no feature function is available, use the full grid. - max_grid_chunk_size = grid_features["grid_weights"].shape[0] - - for atom_slice, grid_slice in make_chunks( - grid_features["atomic_grid_sizes"], max_grid_chunk_size - ): - feature_chunk = {} - for feat_name in ["grid_coords", "grid_weights", "atomic_grid_weights"]: - if feat_name in features: - feature_chunk[feat_name] = grid_features[feat_name][grid_slice] - - for feat_name in ["coarse_0_atomic_coords", "atomic_grid_sizes"]: - if feat_name in features: - feature_chunk[feat_name] = grid_features[feat_name][atom_slice] - - if "atomic_grid_size_bound_shape" in features: - max_size = int(feature_chunk["atomic_grid_sizes"].max().item()) - feature_chunk["atomic_grid_size_bound_shape"] = torch.zeros( - max_size, 0, dtype=torch.long, device=dm.device - ) - - if with_mgga_feature: - assert ff is not None - gpu = dm.device.type == "cuda" - if screen_aos: - chunk_grids = copy(grids) - chunk_grids.coords = grids.coords[grid_slice] - chunk_grids.weights = grids.weights[grid_slice] - if gpu: - chunk_grids._non0ao_idx = None - else: - grid_sort_indices = _spatially_group_atom_grids( - mol, - chunk_grids.coords, - feature_chunk["atomic_grid_sizes"], - ) - chunk_grids.coords = chunk_grids.coords[grid_sort_indices] - chunk_grids.weights = chunk_grids.weights[grid_sort_indices] - grid_sort_indices_t = torch.as_tensor( - grid_sort_indices, device=dm.device - ) - for feat_name in ( - "grid_coords", - "grid_weights", - "atomic_grid_weights", - ): - if feat_name in feature_chunk: - feature_chunk[feat_name] = feature_chunk[feat_name][ - grid_sort_indices_t - ] - chunk_grids.non0tab = dft.gen_grid.make_screen_index( - mol, - chunk_grids.coords, - cutoff=chunk_grids.cutoff, - ) - feat_tensor = ChunkEvalForward.apply( - dm.double(), - mol, - chunk_grids, - ff, - None if gpu else CPU_AO_SCREENING_BLOCK_SIZE, - compile_feature_function, - gpu, - ) - mgga_features = ff.to_dict(feat_tensor) - else: - feat_tensor = non_chunk( - dm.double(), - mol, - grids.coords[grid_slice], - ff, - compile_feature_function=compile_feature_function, - gpu=gpu, - ) - mgga_features = ff.to_dict(feat_tensor) - - for k, v in mgga_features.items(): - feature_chunk[k] = maybe_expand_and_divide(v, not with_spin, 2) - - yield feature_chunk - - def make_chunks( atomic_grid_sizes: Tensor, max_grid_chunk_size: int ) -> list[tuple[slice, slice]]: @@ -1003,6 +820,107 @@ def _global_screened_features( ) +@dataclass(frozen=True) +class _AOBlock: + ao: Tensor + active_aos: Tensor | None + grid_slice: slice + + def select_aos(self, matrix: Tensor) -> Tensor: + if self.active_aos is None: + return matrix + return matrix[..., self.active_aos[:, None], self.active_aos[None, :]] + + def add_to(self, matrix: Tensor, block_result: Tensor) -> None: + if self.active_aos is None: + matrix += block_result + else: + matrix[..., self.active_aos[:, None], self.active_aos[None, :]] += ( + block_result + ) + + +class _AOBlockLoop: + def __init__( + self, + dm: Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: FeatureFunction, + blksize: int | None, + gpu: bool, + ) -> None: + self.dm = dm + self.mol = mol + self.grids = grids + self.feature_function = feature_function + self.blksize = blksize + self.gpu = gpu + self.sort_idx: Tensor | None + self.unsort_idx: Tensor | None + + if gpu: + check_gpu_imports_were_successful() + self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) + self.numint.grid_blksize = blksize + self.sort_idx = torch.as_tensor( + self.numint.gdftopt._ao_idx, device=dm.device + ) + self.unsort_idx = torch.argsort(self.sort_idx) + else: + self.numint = dft.numint.NumInt() + self.sort_idx = None + self.unsort_idx = None + + def order_aos(self, matrix: Tensor) -> Tensor: + if self.sort_idx is None: + return matrix + return matrix[..., self.sort_idx, :][..., self.sort_idx] + + def restore_ao_order(self, matrix: Tensor) -> Tensor: + if self.unsort_idx is None: + return matrix + return matrix[..., self.unsort_idx, :][..., self.unsort_idx] + + def __iter__(self) -> Iterator[_AOBlock]: + end = 0 + for ao_block, mask, weights, _ in self.numint.block_loop( + mol=self.mol, + grids=self.grids, + nao=self.mol.nao, + deriv=self.feature_function.deriv, + blksize=self.blksize, + non0tab=(None if self.gpu else getattr(self.grids, "non0tab", None)), + ): + start, end = end, end + weights.size + ao = from_numpy_or_cupy( + ao_block, + device=self.dm.device, + dtype=self.dm.dtype, + transpose=not self.gpu, + ) + active_aos: Tensor | None + if mask is None: + active_aos = None + elif self.gpu: + active_aos = from_numpy_or_cupy( + mask, device=self.dm.device, dtype=torch.long + ) + else: + num_screen_rows = ( + weights.size + dft.gen_grid.BLKSIZE - 1 + ) // dft.gen_grid.BLKSIZE + active_aos = torch.as_tensor( + _active_cpu_aos(self.mol, mask[:num_screen_rows]), + device=self.dm.device, + dtype=torch.long, + ) + ao = ao[..., active_aos, :] + if active_aos is not None and active_aos.numel() == 0: + continue + yield _AOBlock(ao, active_aos, slice(start, end)) + + class ChunkEvalForward(Function): @staticmethod def setup_context( @@ -1044,20 +962,9 @@ def forward( *vectors_jvp: torch.Tensor, ) -> torch.Tensor: ngrids = grids.weights.size - block_loop_args = (mol, grids, mol.nao) - block_loop_kwargs = { - "deriv": feature_function.deriv, - "blksize": blksize, - "non0tab": None if gpu else getattr(grids, "non0tab", None), - } - if gpu: - check_gpu_imports_were_successful() - ni = dft_gpu.numint.NumInt().build(mol, grids.coords) - ni.grid_blksize = blksize - sort_idx = ni.gdftopt._ao_idx - else: - ni = dft.numint.NumInt() - sort_idx = np.arange(mol.nao_nr()) + block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + dm_ordered = block_loop.order_aos(dm) + vectors_jvp_ordered = [block_loop.order_aos(vector) for vector in vectors_jvp] features = torch.zeros( *dm.shape[:-2], @@ -1069,62 +976,26 @@ def forward( if len(vectors_jvp) > 1 and feature_function.only_linear_feats: return features - # Pre-sort DM and JVP vectors once (sort_idx is constant across blocks) - sort_idx_t = torch.as_tensor(sort_idx, device=dm.device) - dm_sorted = dm[..., sort_idx_t, :][..., sort_idx_t] - vectors_jvp_sorted = [ - v[..., sort_idx_t, :][..., sort_idx_t] for v in vectors_jvp - ] - - end = 0 - for ao_block, mask, weights, _ in ni.block_loop( - *block_loop_args, **block_loop_kwargs - ): - start, end = end, end + weights.size - ao = from_numpy_or_cupy( - ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu - ) - if gpu and mask is not None: - mask = from_numpy_or_cupy(mask, device=dm.device, dtype=torch.long) - elif not gpu and mask is not None: - num_screen_rows = ( - weights.size + dft.gen_grid.BLKSIZE - 1 - ) // dft.gen_grid.BLKSIZE - mask = mask[:num_screen_rows] - mask = torch.as_tensor( - _active_cpu_aos(mol, mask), device=dm.device, dtype=torch.long - ) - ao = ao[..., mask, :] - if mask is not None and mask.numel() == 0: - continue - masked_dm = ( - dm_sorted - if mask is None - else dm_sorted[..., mask[:, None], mask[None, :]] - ) - + for block in block_loop: # Apply chain rule for this particular block partial_func = partial_feature_function_over_aos( feature_function, - ao, + block.ao, ) - for v_sorted in vectors_jvp_sorted: + for vector in vectors_jvp_ordered: partial_func = partial_jvp_function_over_tangents( partial_func, - ( - v_sorted - if mask is None - else v_sorted[..., mask[:, None], mask[None, :]] - ), + block.select_aos(vector), ) # Compute feature (or its jvp) for this block with masked dm + active_dm = block.select_aos(dm_ordered) if compile_feature_function: - temp_feature = torch.compile(partial_func)(masked_dm) + temp_feature = torch.compile(partial_func)(active_dm) else: - temp_feature = partial_func(masked_dm) + temp_feature = partial_func(active_dm) - features[..., start:end] = temp_feature + features[..., block.grid_slice] = temp_feature return features @staticmethod @@ -1237,107 +1108,54 @@ def forward( gpu: bool, *vectors: torch.Tensor, ) -> torch.Tensor: - block_loop_args = (mol, grids, mol.nao) - block_loop_kwargs = { - "deriv": feature_function.deriv, - "blksize": blksize if not gpu else None, - "non0tab": None if gpu else getattr(grids, "non0tab", None), - } - if gpu: - check_gpu_imports_were_successful() - ni = dft_gpu.numint.NumInt().build(mol, grids.coords) - ni.grid_blksize = blksize - sort_idx = ni.gdftopt._ao_idx - else: - ni = dft.numint.NumInt() - sort_idx = np.arange(mol.nao_nr()) + block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + dm_ordered = block_loop.order_aos(dm) + vectors_ordered = [ + block_loop.order_aos(vector) + if derivative_type in ("jvp", "vjp") + else vector + for derivative_type, vector in zip(derivative_types, vectors, strict=True) + ] - end: int = 0 out = torch.zeros_like(dm) if len(vectors) > 1 and feature_function.only_linear_feats: return out - # Pre-sort DM and derivative vectors once (sort_idx is constant across blocks) - sort_idx_t = torch.as_tensor(sort_idx, device=dm.device) - unsort_idx = torch.argsort(sort_idx_t) - dm_sorted = dm[..., sort_idx_t, :][..., sort_idx_t] - vectors_sorted = [ - v[..., sort_idx_t, :][..., sort_idx_t] if dt in ("jvp", "vjp") else v - for dt, v in zip(derivative_types, vectors, strict=True) - ] - - for ao_block, mask, weights, _ in ni.block_loop( - *block_loop_args, - **block_loop_kwargs, - ): - start, end = end, end + weights.size - - ao = from_numpy_or_cupy( - ao_block, device=dm.device, dtype=dm.dtype, transpose=not gpu - ) - if gpu and mask is not None: - mask = from_numpy_or_cupy(mask, device=dm.device, dtype=torch.long) - elif not gpu and mask is not None: - num_screen_rows = ( - weights.size + dft.gen_grid.BLKSIZE - 1 - ) // dft.gen_grid.BLKSIZE - mask = mask[:num_screen_rows] - mask = torch.as_tensor( - _active_cpu_aos(mol, mask), device=dm.device, dtype=torch.long - ) - ao = ao[..., mask, :] - if mask is not None and mask.numel() == 0: - continue - + for block in block_loop: # Apply chain rule for this particular block # but be careful with signature change upon first vjp partial_func = partial_feature_function_over_aos( feature_function, - ao, + block.ao, ) - for derivative_type, vector, v_sorted in zip( - derivative_types, vectors, vectors_sorted, strict=True + for derivative_type, vector, vector_ordered in zip( + derivative_types, vectors, vectors_ordered, strict=True ): if derivative_type == "jvp": partial_func = partial_jvp_function_over_tangents( partial_func, - ( - v_sorted - if mask is None - else v_sorted[..., mask[:, None], mask[None, :]] - ), + block.select_aos(vector_ordered), ) elif derivative_type == "vjp": partial_func = partial_vjp_function_over_tangents( partial_func, - ( - v_sorted - if mask is None - else v_sorted[..., mask[:, None], mask[None, :]] - ), + block.select_aos(vector_ordered), ) elif derivative_type == "first_vjp": partial_func = partial_vjp_function_over_tangents( - partial_func, vector[..., start:end] + partial_func, vector[..., block.grid_slice] ) else: raise ValueError( f"Unknown derivative {derivative_type} (must be one of 'jvp', 'vjp', 'first_vjp')" ) - masked_dm = ( - dm_sorted - if mask is None - else dm_sorted[..., mask[:, None], mask[None, :]] - ) + active_dm = block.select_aos(dm_ordered) if compile_feature_function: - block_result = torch.compile(partial_func)(masked_dm) - else: - block_result = partial_func(masked_dm) - if mask is None: - out += block_result + block_result = torch.compile(partial_func)(active_dm) else: - out[..., mask[:, None], mask[None, :]] += block_result - return out[..., unsort_idx, :][..., unsort_idx] + block_result = partial_func(active_dm) + block.add_to(out, block_result) + return block_loop.restore_ao_order(out) @staticmethod def jvp(ctx: FunctionCtx, *grad_input: torch.Tensor) -> torch.Tensor: diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index 75d34a49..56607e31 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -11,7 +11,7 @@ def estimate_max_grid_chunk_size( func_deriv: int = 1, reserved_memory_in_bytes: int = 0, ) -> int: - """Heuristically pick a grid chunk size for :func:`chunked_features`. + """Heuristically pick a model grid chunk size for screened feature evaluation. The dominant per-chunk allocation is the atomic-orbital matrix evaluated by ``non_chunk`` (shape ``(ncomp, nao, n)`` in float64, with no AO screening), diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index ee367c69..f097d1a3 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -20,7 +20,6 @@ ) from skala.pyscf.features import ( _global_screened_features, - chunked_features, generate_features, ) @@ -217,117 +216,86 @@ def __call__( f"Density matrix device {dm.device} does not match functional device {self.device}" ) - if self._functional_supports_atom_chunking(): + if self._functional_supports_atom_chunking() and _should_screen_aos(mol): dm = dm.detach().requires_grad_() - screen_aos = _should_screen_aos(mol) tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) - V_xc = torch.zeros_like(dm) - - if screen_aos: - screened_features = _global_screened_features( - mol, - dm, - grids, - features=set(self.func.features), - func_deriv=1, - max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, - safety_fraction=0.8, - ) - # Store only full-grid feature cotangents; model activations remain chunk-local. - atom_major_cotangent = torch.zeros_like( - screened_features.atom_major_raw_features - ) - for atom_slice, grid_slice in screened_features.chunks: - # Break the reorder graph so model backprop retains only this chunk. - local_raw_features = ( - screened_features.atom_major_raw_features[..., grid_slice] - .detach() - .requires_grad_() - ) - mol_features = screened_features.build_model_chunk( - local_raw_features, atom_slice, grid_slice - ) - E_xc_chunk = self.func.get_exc(mol_features) - (local_cotangent,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - ) - atom_major_cotangent[..., grid_slice] = local_cotangent.detach() - tot_dens += ( - (mol_features["density"] * mol_features["grid_weights"]) - .sum(dim=-1) - .detach() - ) - E_xc += E_xc_chunk.detach() - del E_xc_chunk, local_cotangent, local_raw_features, mol_features - - # Reorder detached cotangents explicitly instead of backpropagating through it. - sorted_cotangent = atom_major_cotangent.index_select( - -1, screened_features.forward_permutation - ) - # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. - (V_xc,) = torch.autograd.grad( - screened_features.sorted_raw_features, - dm, - sorted_cotangent, - ) - return tot_dens, E_xc, V_xc - - for mol_features in chunked_features( + screened_features = _global_screened_features( mol, dm, grids, features=set(self.func.features), func_deriv=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, - safety_fraction=0.8, # tends to be faster for large chunks - screen_aos=False, - ): + safety_fraction=0.8, + ) + # Store only full-grid feature cotangents; model activations remain chunk-local. + atom_major_cotangent = torch.zeros_like( + screened_features.atom_major_raw_features + ) + for atom_slice, grid_slice in screened_features.chunks: + # Break the reorder graph so model backprop retains only this chunk. + local_raw_features = ( + screened_features.atom_major_raw_features[..., grid_slice] + .detach() + .requires_grad_() + ) + mol_features = screened_features.build_model_chunk( + local_raw_features, atom_slice, grid_slice + ) E_xc_chunk = self.func.get_exc(mol_features) - (V_xc_chunk,) = torch.autograd.grad( + (local_cotangent,) = torch.autograd.grad( E_xc_chunk, - dm, + local_raw_features, torch.ones_like(E_xc_chunk), ) + atom_major_cotangent[..., grid_slice] = local_cotangent.detach() tot_dens += ( (mol_features["density"] * mol_features["grid_weights"]) .sum(dim=-1) .detach() ) E_xc += E_xc_chunk.detach() - V_xc += V_xc_chunk.detach() - del E_xc_chunk, V_xc_chunk, mol_features + del E_xc_chunk, local_cotangent, local_raw_features, mol_features - return tot_dens, E_xc, V_xc - else: - dm = dm.requires_grad_() - mol_features = generate_features( - mol, - dm, - grids, - set(self.func.features), - chunk_size=self.chunk_size, - max_memory=max_memory, - gpu=self.device.type == "cuda", + # Reorder detached cotangents explicitly instead of backpropagating through it. + sorted_cotangent = atom_major_cotangent.index_select( + -1, screened_features.forward_permutation ) - E_xc = self.func.get_exc(mol_features) + # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. (V_xc,) = torch.autograd.grad( - E_xc, + screened_features.sorted_raw_features, dm, - torch.ones_like(E_xc), - retain_graph=second_order, - create_graph=second_order, - ) - - rho = mol_features["density"] - grid_weights = mol_features.get( - "grid_weights", self.from_backend(grids.weights) + sorted_cotangent, ) - tot_dens = (rho * grid_weights).sum(dim=-1) return tot_dens, E_xc, V_xc + dm = dm.requires_grad_() + mol_features = generate_features( + mol, + dm, + grids, + set(self.func.features), + chunk_size=self.chunk_size, + max_memory=max_memory, + gpu=self.device.type == "cuda", + ) + E_xc = self.func.get_exc(mol_features) + (V_xc,) = torch.autograd.grad( + E_xc, + dm, + torch.ones_like(E_xc), + retain_graph=second_order, + create_graph=second_order, + ) + + rho = mol_features["density"] + grid_weights = mol_features.get( + "grid_weights", self.from_backend(grids.weights) + ) + tot_dens = (rho * grid_weights).sum(dim=-1) + return tot_dens, E_xc, V_xc + def nr_rks( self, mol: gto.Mole, @@ -393,116 +361,77 @@ def gen_response( dm0 = self.from_backend(ks.make_rdm1(mo_coeff, mo_occ)) - if self._functional_supports_atom_chunking(): + if self._functional_supports_atom_chunking() and _should_screen_aos(ks.mol): dm0 = dm0.requires_grad_() - screen_aos = _should_screen_aos(ks.mol) - - screened_features = None - if screen_aos: - screened_features = _global_screened_features( - ks.mol, - dm0, - ks.grids, - features=set(self.func.features), - func_deriv=2, - max_memory_in_mb=ks.max_memory - if dm0.device.type == "cpu" - else None, - safety_fraction=kwargs.get("safety_fraction", 0.0), + screened_features = _global_screened_features( + ks.mol, + dm0, + ks.grids, + features=set(self.func.features), + func_deriv=2, + max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, + safety_fraction=kwargs.get("safety_fraction", 0.0), + ) + if not screened_features.feature_function.only_linear_feats: + raise NotImplementedError( + "Global screened response requires raw features linear in " + "the density matrix." ) - if not screened_features.feature_function.only_linear_feats: - raise NotImplementedError( - "Global screened response requires raw features linear in " - "the density matrix." - ) def hessian_vector_product_atom_chunked(dm1: Array) -> Array: dm1_tensor = self.from_backend(dm1) - if screened_features is not None: - atom_major_tangent = screened_features.atom_major_jvp(dm1_tensor) - # Store the full-grid model Hessian action, not per-chunk model graphs. - atom_major_hessian_action = torch.zeros_like( - screened_features.atom_major_raw_features + atom_major_tangent = screened_features.atom_major_jvp(dm1_tensor) + # Store the full-grid model Hessian action, not per-chunk model graphs. + atom_major_hessian_action = torch.zeros_like( + screened_features.atom_major_raw_features + ) + for atom_slice, grid_slice in screened_features.chunks: + # Isolate the second-order model graph to the current atomic chunk. + local_raw_features = ( + screened_features.atom_major_raw_features[..., grid_slice] + .detach() + .requires_grad_() ) - for atom_slice, grid_slice in screened_features.chunks: - # Isolate the second-order model graph to the current atomic chunk. - local_raw_features = ( - screened_features.atom_major_raw_features[..., grid_slice] - .detach() - .requires_grad_() - ) - mol_features = screened_features.build_model_chunk( - local_raw_features, atom_slice, grid_slice - ) - E_xc_chunk = self.func.get_exc(mol_features) - (local_gradient,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - create_graph=True, - ) - if local_gradient.requires_grad: - (local_hessian_action,) = torch.autograd.grad( - local_gradient, - local_raw_features, - atom_major_tangent[..., grid_slice], - ) - else: - local_hessian_action = torch.zeros_like(local_raw_features) - atom_major_hessian_action[..., grid_slice] = ( - local_hessian_action.detach() - ) - del ( - E_xc_chunk, + mol_features = screened_features.build_model_chunk( + local_raw_features, atom_slice, grid_slice + ) + E_xc_chunk = self.func.get_exc(mol_features) + (local_gradient,) = torch.autograd.grad( + E_xc_chunk, + local_raw_features, + torch.ones_like(E_xc_chunk), + create_graph=True, + ) + if local_gradient.requires_grad: + (local_hessian_action,) = torch.autograd.grad( local_gradient, - local_hessian_action, local_raw_features, - mol_features, + atom_major_tangent[..., grid_slice], ) - - # Restore box order after all chunk-local Hessian actions are detached. - sorted_hessian_action = atom_major_hessian_action.index_select( - -1, screened_features.forward_permutation + else: + local_hessian_action = torch.zeros_like(local_raw_features) + atom_major_hessian_action[..., grid_slice] = ( + local_hessian_action.detach() ) - # The custom VJP traverses AO blocks sequentially and retains no AO graph. - (hvp_total,) = torch.autograd.grad( - screened_features.sorted_raw_features, - dm0, - sorted_hessian_action, - retain_graph=True, + del ( + E_xc_chunk, + local_gradient, + local_hessian_action, + local_raw_features, + mol_features, ) - else: - hvp_total = torch.zeros_like(dm0) - for mol_features in chunked_features( - ks.mol, - dm0, - ks.grids, - features=set(self.func.features), - func_deriv=2, - max_memory_in_mb=ks.max_memory - if dm0.device.type == "cpu" - else None, - safety_fraction=kwargs.get( - "safety_fraction", 0.0 - ), # Force small chunks (single atoms) because it's empirically fastest. - screen_aos=False, - ): - E_xc_chunk = self.func.get_exc(mol_features) - (V_xc_chunk,) = torch.autograd.grad( - E_xc_chunk, - dm0, - torch.ones_like(E_xc_chunk), - retain_graph=True, - create_graph=True, - ) - (hvp_chunk,) = torch.autograd.grad( - V_xc_chunk, - dm0, - dm1_tensor, - retain_graph=True, - ) - hvp_total += hvp_chunk - del E_xc_chunk, V_xc_chunk, hvp_chunk, mol_features + + # Restore box order after all chunk-local Hessian actions are detached. + sorted_hessian_action = atom_major_hessian_action.index_select( + -1, screened_features.forward_permutation + ) + # The custom VJP traverses AO blocks sequentially and retains no AO graph. + (hvp_total,) = torch.autograd.grad( + screened_features.sorted_raw_features, + dm0, + sorted_hessian_action, + retain_graph=True, + ) v1 = self.to_backend(hvp_total) vj = ks.get_j(ks.mol, dm1, hermi=1) @@ -514,27 +443,26 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: return hessian_vector_product_atom_chunked - else: - # caching V_xc saves a forward pass in each iteration - dm0 = dm0.requires_grad_() - V_xc = self(ks.mol, ks.grids, None, dm0, second_order=True)[2] + # caching V_xc saves a forward pass in each iteration + dm0 = dm0.requires_grad_() + V_xc = self(ks.mol, ks.grids, None, dm0, second_order=True)[2] - def hessian_vector_product(dm1: Array) -> Array: - v1 = self.to_backend( - torch.autograd.grad( - V_xc, dm0, self.from_backend(dm1), retain_graph=True - )[0] - ) - vj = ks.get_j(ks.mol, dm1, hermi=1) + def hessian_vector_product(dm1: Array) -> Array: + v1 = self.to_backend( + torch.autograd.grad( + V_xc, dm0, self.from_backend(dm1), retain_graph=True + )[0] + ) + vj = ks.get_j(ks.mol, dm1, hermi=1) - if ks.mol.spin == 0: - v1 += vj - else: - v1 += vj[0] + vj[1] + if ks.mol.spin == 0: + v1 += vj + else: + v1 += vj[0] + vj[1] - return v1 + return v1 - return hessian_vector_product + return hessian_vector_product def _functional_supports_atom_chunking(self) -> bool: return "atomic_grid_sizes" in self.func.features diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index ce64190a..8ea6d118 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -10,14 +10,12 @@ from skala.pyscf import features as features_module from skala.pyscf import numint as numint_module from skala.pyscf.features import ( - CPU_AO_SCREENING_BLOCK_SIZE, ChunkEvalBackward, ChunkEvalForward, MGGAFeatureFunction, _active_cpu_aos, _prepare_spatially_sorted_grids, _spatial_grid_permutations, - chunked_features, ) from skala.pyscf.numint import SkalaNumInt, _should_screen_aos @@ -171,178 +169,6 @@ def fake_make_screen_index( assert screen_index_calls == 2 -@pytest.mark.parametrize("screen_aos", [False, True]) -def test_chunked_features_routes_screening( - carbon: gto.Mole, - monkeypatch: pytest.MonkeyPatch, - screen_aos: bool, -) -> None: - ngrids = dft.gen_grid.BLKSIZE - grids = dft.Grids(carbon) - grids.coords = np.zeros((ngrids, 3)) - grids.weights = np.ones(ngrids) - dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) - calls: list[str] = [] - - monkeypatch.setattr( - features_module, - "get_grid_features", - lambda *args, **kwargs: { - "atomic_grid_sizes": torch.tensor([ngrids]), - "grid_weights": torch.ones(ngrids, dtype=torch.float64), - }, - ) - monkeypatch.setattr( - features_module, - "estimate_max_grid_chunk_size", - lambda *args, **kwargs: ngrids, - ) - monkeypatch.setattr( - dft.gen_grid, - "make_screen_index", - lambda *args, **kwargs: np.ones((1, carbon.nbas), dtype=np.uint8), - ) - - def fake_non_chunk(*args: object, **kwargs: object) -> torch.Tensor: - calls.append("non_chunk") - return torch.zeros((1, ngrids), dtype=torch.float64) - - def fake_chunk_eval(*args: object, **kwargs: object) -> torch.Tensor: - calls.append("ChunkEval") - return torch.zeros((1, ngrids), dtype=torch.float64) - - monkeypatch.setattr(features_module, "non_chunk", fake_non_chunk) - monkeypatch.setattr(ChunkEvalForward, "apply", fake_chunk_eval) - - list( - chunked_features( - carbon, - dm, - grids, - {"atomic_grid_sizes", "density", "grid_weights"}, - func_deriv=1, - screen_aos=screen_aos, - ) - ) - - assert calls == ["ChunkEval" if screen_aos else "non_chunk"] - - if screen_aos: - assert CPU_AO_SCREENING_BLOCK_SIZE == 504 - - -def test_cpu_screening_spatially_groups_each_atom( - monkeypatch: pytest.MonkeyPatch, -) -> None: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) - atom_grid_size = dft.gen_grid.BLKSIZE - ngrids = 2 * atom_grid_size - coords = np.zeros((ngrids, 3)) - coords[:, 0] = np.arange(ngrids) - weights = np.arange(ngrids, dtype=np.float64) + 100 - atomic_grid_weights = torch.arange(ngrids, dtype=torch.float64) + 200 - grids = dft.Grids(mol) - grids.coords = coords - grids.weights = weights - grouped_slices: list[np.ndarray] = [] - - monkeypatch.setattr( - features_module, - "get_grid_features", - lambda *args, **kwargs: { - "atomic_grid_sizes": torch.tensor([atom_grid_size, atom_grid_size]), - "grid_coords": torch.from_numpy(coords.copy()), - "grid_weights": torch.from_numpy(weights.copy()), - "atomic_grid_weights": atomic_grid_weights, - }, - ) - monkeypatch.setattr( - features_module, - "estimate_max_grid_chunk_size", - lambda *args, **kwargs: ngrids, - ) - - def fake_group_grids(mol_arg: gto.Mole, atom_coords: np.ndarray) -> np.ndarray: - assert mol_arg is mol - grouped_slices.append(atom_coords.copy()) - return np.arange(atom_grid_size - 1, -1, -1) - - monkeypatch.setattr(dft.gen_grid, "arg_group_grids", fake_group_grids) - - sort_indices = np.concatenate( - ( - np.arange(atom_grid_size - 1, -1, -1), - np.arange(ngrids - 1, atom_grid_size - 1, -1), - ) - ) - - def fake_make_screen_index( - mol_arg: gto.Mole, sorted_coords: np.ndarray, cutoff: float - ) -> np.ndarray: - assert mol_arg is mol - assert np.array_equal(sorted_coords, coords[sort_indices]) - return np.ones((2, mol.nbas), dtype=np.uint8) - - monkeypatch.setattr(dft.gen_grid, "make_screen_index", fake_make_screen_index) - - def fake_chunk_eval( - dm: torch.Tensor, - mol_arg: gto.Mole, - sorted_grids: dft.Grids, - feature_function: MGGAFeatureFunction, - block_size: int, - compile_feature_function: bool, - gpu: bool, - ) -> torch.Tensor: - assert mol_arg is mol - assert feature_function.with_density - assert block_size == CPU_AO_SCREENING_BLOCK_SIZE - assert not compile_feature_function - assert not gpu - assert np.array_equal(sorted_grids.coords, coords[sort_indices]) - assert np.array_equal(sorted_grids.weights, weights[sort_indices]) - return torch.from_numpy(sorted_grids.coords[:, 0]).to(dm).unsqueeze(0) - - monkeypatch.setattr(ChunkEvalForward, "apply", fake_chunk_eval) - - (feature_chunk,) = list( - chunked_features( - mol, - torch.eye(mol.nao_nr(), dtype=torch.float64), - grids, - { - "atomic_grid_sizes", - "atomic_grid_weights", - "density", - "grid_coords", - "grid_weights", - }, - func_deriv=1, - screen_aos=True, - ) - ) - - assert len(grouped_slices) == 2 - assert np.array_equal(grouped_slices[0], coords[:atom_grid_size]) - assert np.array_equal(grouped_slices[1], coords[atom_grid_size:]) - assert torch.equal( - feature_chunk["atomic_grid_sizes"], - torch.tensor([atom_grid_size, atom_grid_size]), - ) - assert torch.equal( - feature_chunk["grid_coords"], torch.from_numpy(coords[sort_indices]) - ) - assert torch.equal( - feature_chunk["grid_weights"], torch.from_numpy(weights[sort_indices]) - ) - assert torch.equal( - feature_chunk["atomic_grid_weights"], atomic_grid_weights[sort_indices] - ) - expected_density = torch.from_numpy(coords[sort_indices, 0]) / 2 - assert torch.equal(feature_chunk["density"][0], expected_density) - assert torch.equal(feature_chunk["density"][1], expected_density) - - class QuadraticDensityFunctional(ExcFunctionalBase): def __init__(self) -> None: super().__init__() @@ -373,21 +199,18 @@ def test_first_and_second_order_use_same_screening_decision( ) -> None: switch_size = carbon.nao_nr() - 1 if expected else carbon.nao_nr() monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", switch_size) - decisions: list[bool] = [] + routes: list[str] = [] - def fake_chunked_features( + def fake_generate_features( mol: gto.Mole, dm: torch.Tensor, grids: object, - features: set[str], - func_deriv: int, - *, - screen_aos: bool, + features: set[str] | None = None, **kwargs: object, - ) -> Iterator[dict[str, torch.Tensor]]: - decisions.append(screen_aos) + ) -> dict[str, torch.Tensor]: + routes.append("dense") density = dm.square().sum().reshape(1).expand(2, 1) / 2 - yield { + return { "atomic_grid_sizes": torch.tensor([1]), "density": density, "grid_weights": torch.ones(1, dtype=dm.dtype), @@ -431,26 +254,29 @@ def fake_global_screened_features( func_deriv: int, **kwargs: object, ) -> FakeGlobalScreenedFeatures: - decisions.append(True) + routes.append("screened") return FakeGlobalScreenedFeatures(dm) - monkeypatch.setattr(numint_module, "chunked_features", fake_chunked_features) + monkeypatch.setattr(numint_module, "generate_features", fake_generate_features) monkeypatch.setattr( numint_module, "_global_screened_features", fake_global_screened_features ) numint = SkalaNumInt(QuadraticDensityFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + grids = dft.Grids(carbon) + grids.weights = np.ones(1) - numint(carbon, object(), None, dm) + numint(carbon, grids, None, dm) - ks = FakeKS(carbon) + ks = FakeKS(carbon, grids) response = numint.gen_response( np.eye(carbon.nao_nr()), np.ones(carbon.nao_nr()), ks=ks ) result = response(np.eye(carbon.nao_nr())) assert result.shape == (carbon.nao_nr(), carbon.nao_nr()) - assert decisions == [expected, expected] + expected_route = "screened" if expected else "dense" + assert routes == [expected_route, expected_route] def test_cpu_screening_slices_and_scatters_full_derivatives( diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index fe9b0c9b..564eaa69 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -6,7 +6,7 @@ from collections.abc import Callable, Iterator from multiprocessing.connection import Connection from pathlib import Path -from typing import NamedTuple, cast +from typing import Any, NamedTuple, cast import numpy as np import pytest @@ -14,6 +14,7 @@ from pyscf import dft, gto, lib from pyscf.dft import numint as pyscf_numint from pytest_benchmark.fixture import BenchmarkFixture +from torch.utils.dlpack import from_dlpack from skala.functional import load_functional from skala.functional.base import ExcFunctionalBase @@ -117,6 +118,16 @@ class BenchmarkCase(NamedTuple): numint: SkalaNumInt[np.ndarray] +DeviceResult = tuple[float, float, object] + + +class DeviceBenchmarkCase(NamedTuple): + backend: str + mol: gto.Mole + run: Callable[[], DeviceResult] + synchronize: Callable[[], None] + + BENCHMARK_SPECS = [ pytest.param(BenchmarkSpec("naphthalene", NAPHTHALENE), id="naphthalene"), pytest.param(BenchmarkSpec("anthracene", ANTHRACENE), id="anthracene"), @@ -135,6 +146,35 @@ def _make_benchmark_case( return BenchmarkCase(mol, grids, dm, SkalaNumInt(functional)) +def _make_gpu_benchmark_case( + spec: BenchmarkSpec, functional: ExcFunctionalBase +) -> DeviceBenchmarkCase: + import cupy + + from skala.gpu4pyscf import SkalaKS + + mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) + ks = SkalaKS(mol, xc=functional, with_dftd3=False) + ks.grids.level = 1 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + dm = cupy.asarray(dft.RKS(mol).get_init_guess()) + assert _should_screen_aos(mol) + + def run() -> DeviceResult: + result = ks._numint.nr_rks( + mol, + ks.grids, + None, + dm, + max_memory=MAX_MEMORY_MB, + ) + torch.cuda.synchronize() + return cast(DeviceResult, result) + + return DeviceBenchmarkCase("cuda", mol, run, torch.cuda.synchronize) + + @pytest.fixture(scope="module") def fixed_cpu_threads() -> Iterator[None]: previous_pyscf_threads = lib.num_threads() @@ -148,19 +188,48 @@ def fixed_cpu_threads() -> Iterator[None]: lib.num_threads(previous_pyscf_threads) -@pytest.fixture( - scope="module", - params=BENCHMARK_SPECS, -) +@pytest.fixture(scope="module", params=BENCHMARK_SPECS) +def benchmark_spec(request: pytest.FixtureRequest) -> BenchmarkSpec: + return cast(BenchmarkSpec, request.param) + + +@pytest.fixture(scope="module") def benchmark_case( + benchmark_spec: BenchmarkSpec, request: pytest.FixtureRequest, fixed_cpu_threads: None, load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> BenchmarkCase: - spec = cast(BenchmarkSpec, request.param) functional = load_functional_cached("skala-1.1") assert isinstance(functional, ExcFunctionalBase) - return _make_benchmark_case(spec, functional) + return _make_benchmark_case(benchmark_spec, functional) + + +@pytest.fixture( + scope="module", + params=["cpu", pytest.param("cuda", marks=pytest.mark.gpu)], +) +def device_benchmark_case( + request: pytest.FixtureRequest, + benchmark_spec: BenchmarkSpec, + fixed_cpu_threads: None, + load_functional_cached: Callable[..., ExcFunctionalBase | str], +) -> DeviceBenchmarkCase: + backend = cast(str, request.param) + if backend == "cpu": + functional = load_functional_cached("skala-1.1") + assert isinstance(functional, ExcFunctionalBase) + case = _make_benchmark_case(benchmark_spec, functional) + assert _should_screen_aos(case.mol) + return DeviceBenchmarkCase("cpu", case.mol, lambda: _run_xc(case), lambda: None) + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + pytest.importorskip("cupy") + pytest.importorskip("gpu4pyscf") + functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(functional, ExcFunctionalBase) + return _make_gpu_benchmark_case(benchmark_spec, functional) @pytest.fixture @@ -198,26 +267,30 @@ def _benchmark_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: ) -def _run_gpu_xc(spec: BenchmarkSpec, screened: bool) -> int: - from skala.gpu4pyscf import SkalaKS +def _benchmark_device_xc( + benchmark: BenchmarkFixture, case: DeviceBenchmarkCase +) -> None: + case.synchronize() + pedantic = cast(Callable[..., object], benchmark.pedantic) + pedantic( + case.run, + rounds=1, + iterations=2, + ) - mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) + +def _run_gpu_xc(spec: BenchmarkSpec, screened: bool) -> int: functional = load_functional("skala-1.1", device=torch.device("cuda:0")) assert isinstance(functional, ExcFunctionalBase) - ks = SkalaKS(mol, xc=functional, with_dftd3=False) - ks.grids.level = 1 - ks.grids.alignment = 1 - ks.grids.build(sort_grids=False) - dm = ks.get_init_guess() + case = _make_gpu_benchmark_case(spec, functional) if not screened: - pyscf_numint.SWITCH_SIZE = mol.nao_nr() - assert _should_screen_aos(mol) is screened + pyscf_numint.SWITCH_SIZE = 10**9 torch.cuda.synchronize() torch.cuda.empty_cache() baseline_bytes = torch.cuda.memory_allocated() torch.cuda.reset_peak_memory_stats() - ks._numint.nr_rks(mol, ks.grids, None, dm, max_memory=MAX_MEMORY_MB) + case.run() torch.cuda.synchronize() return torch.cuda.max_memory_allocated() - baseline_bytes @@ -294,23 +367,50 @@ def _measure_peak_memory(spec: BenchmarkSpec, screened: bool, backend: str) -> i @pytest.mark.profiling def test_screened_and_dense_values_agree( - benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch + device_benchmark_case: DeviceBenchmarkCase, + benchmark_spec: BenchmarkSpec, + load_functional_cached: Callable[..., ExcFunctionalBase | str], + monkeypatch: pytest.MonkeyPatch, ) -> None: - assert _should_screen_aos(benchmark_case.mol) - screened = _run_xc(benchmark_case) - - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", benchmark_case.mol.nao_nr()) - assert not _should_screen_aos(benchmark_case.mol) - dense = _run_xc(benchmark_case) - - assert np.allclose(dense[0], screened[0], rtol=1e-10, atol=1e-11) - assert np.isclose(dense[1], screened[1], rtol=1e-10, atol=1e-11) - vxc_difference = dense[2] - screened[2] + case = device_benchmark_case + assert _should_screen_aos(case.mol) + screened = case.run() + + monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", case.mol.nao_nr()) + assert not _should_screen_aos(case.mol) + if case.backend == "cpu": + dense = case.run() + else: + cpu_functional = load_functional_cached("skala-1.1", device=torch.device("cpu")) + assert isinstance(cpu_functional, ExcFunctionalBase) + dense = _run_xc(_make_benchmark_case(benchmark_spec, cpu_functional)) + + scalar_rtol = 2e-10 if case.backend == "cpu" else 1e-8 + density_close = np.allclose(dense[0], screened[0], rtol=scalar_rtol, atol=1e-11) + energy_close = np.isclose(dense[1], screened[1], rtol=scalar_rtol, atol=1e-10) + dense_vxc = ( + dense[2] + if isinstance(dense[2], np.ndarray) + else from_dlpack(cast(Any, dense[2])).cpu().numpy() + ) + screened_vxc = ( + screened[2] + if isinstance(screened[2], np.ndarray) + else from_dlpack(cast(Any, screened[2])).cpu().numpy() + ) + vxc_difference = dense_vxc - screened_vxc vxc_max_abs_difference = np.max(np.abs(vxc_difference)) vxc_relative_l2_difference = np.linalg.norm(vxc_difference) / np.linalg.norm( - dense[2] + dense_vxc ) - assert vxc_max_abs_difference < 5e-8 and vxc_relative_l2_difference < 1e-8, ( + vxc_max_atol = 5e-8 if case.backend == "cpu" else 2e-7 + vxc_relative_rtol = 1e-8 if case.backend == "cpu" else 1e-7 + assert ( + density_close + and energy_close + and vxc_max_abs_difference < vxc_max_atol + and vxc_relative_l2_difference < vxc_relative_rtol + ), ( f"N: dense={dense[0]:.16g}, screened={screened[0]:.16g}, " f"abs_diff={abs(dense[0] - screened[0]):.3e}; " f"E_xc: dense={dense[1]:.16g}, screened={screened[1]:.16g}, " @@ -334,6 +434,14 @@ def test_without_ao_screening_by_patching_threshold( _benchmark_xc(benchmark, dense_case) +@pytest.mark.benchmark(group="device-def2-qzvpp-screened") +def test_screened_runtime_by_device( + benchmark: BenchmarkFixture, + device_benchmark_case: DeviceBenchmarkCase, +) -> None: + _benchmark_device_xc(benchmark, device_benchmark_case) + + @pytest.mark.profiling @pytest.mark.parametrize("spec", BENCHMARK_SPECS) @pytest.mark.parametrize( @@ -371,13 +479,21 @@ def test_screened_and_dense_peak_memory( @pytest.mark.profiling def test_profile_with_natural_ao_screening( - screened_case: BenchmarkCase, + device_benchmark_case: DeviceBenchmarkCase, ) -> None: - _run_xc(screened_case) + assert _should_screen_aos(device_benchmark_case.mol) + device_benchmark_case.run() @pytest.mark.profiling def test_profile_without_ao_screening_by_patching_threshold( - dense_case: BenchmarkCase, + device_benchmark_case: DeviceBenchmarkCase, + monkeypatch: pytest.MonkeyPatch, ) -> None: - _run_xc(dense_case) + monkeypatch.setattr( + pyscf_numint, + "SWITCH_SIZE", + device_benchmark_case.mol.nao_nr(), + ) + assert not _should_screen_aos(device_benchmark_case.mol) + device_benchmark_case.run() From 0343e37080e1df274afb3e517c0e925a718ccd1f Mon Sep 17 00:00:00 2001 From: Jens Date: Tue, 4 Aug 2026 09:57:59 +0200 Subject: [PATCH 09/39] Apply suggestions from code review Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- src/skala/pyscf/numint.py | 2 +- tests/test_ao_screening.py | 3 --- tests/test_gpu4pyscf_ao_screening.py | 6 ------ 3 files changed, 1 insertion(+), 10 deletions(-) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index f097d1a3..c30644f5 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -421,7 +421,7 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: mol_features, ) - # Restore box order after all chunk-local Hessian actions are detached. + # Restore block order after all chunk-local Hessian actions are detached. sorted_hessian_action = atom_major_hessian_action.index_select( -1, screened_features.forward_permutation ) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 8ea6d118..55136411 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -436,9 +436,6 @@ def test_cpu_response_dense_screened_equivalence( assert np.allclose( dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 ) - assert np.allclose( - dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 - ) @pytest.mark.parametrize("func_deriv", [1, 2]) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index aa36c242..40576c2c 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -187,12 +187,6 @@ def test_gpu_response_dense_screened_equivalence( rtol=1e-9, atol=1e-10, ) - assert np.allclose( - _to_numpy(dense_response(dm1)), - _to_numpy(screened_response(dm1)), - rtol=1e-9, - atol=1e-10, - ) def test_gpu_uks_response_dense_screened_equivalence( From 9885e849fbaf3a08fafa8e24d276ca9827ee9af4 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Tue, 4 Aug 2026 16:25:16 +0200 Subject: [PATCH 10/39] refactor tests --- src/skala/functional/base.py | 4 +- src/skala/functional/model.py | 6 +- src/skala/pyscf/features.py | 12 ++- tests/test_ao_screening.py | 5 +- tests/test_ao_screening_benchmark.py | 133 +++++++++++---------------- tests/test_gpu4pyscf_ao_screening.py | 78 ++++++++++++++-- tests/test_pyscf_gradients.py | 2 +- 7 files changed, 143 insertions(+), 97 deletions(-) diff --git a/src/skala/functional/base.py b/src/skala/functional/base.py index cd6137bf..cbb40705 100644 --- a/src/skala/functional/base.py +++ b/src/skala/functional/base.py @@ -8,7 +8,7 @@ """ from collections.abc import Callable -from typing import Any +from typing import Any, ClassVar import torch from torch import nn @@ -25,7 +25,7 @@ class ExcFunctionalBase(nn.Module): energy density from molecular features. """ - features: list[str] + features: ClassVar[list[str]] """List of features that this functional requires.""" def get_d3_settings(self) -> str | None: diff --git a/src/skala/functional/model.py b/src/skala/functional/model.py index 531709bc..ca817efd 100644 --- a/src/skala/functional/model.py +++ b/src/skala/functional/model.py @@ -9,7 +9,7 @@ """ import math -from typing import Any, cast +from typing import Any, ClassVar, cast import torch from e3nn import o3 @@ -57,7 +57,7 @@ def _prepare_features_raw( class SemiLocalFeatures(nn.Module): """Compute semi-local (ab, ba) feature pairs with a pre-buffered permutation index.""" - _PERM = [1, 0, 3, 2, 5, 4, 6] + _PERM: ClassVar[list[int]] = [1, 0, 3, 2, 5, 4, 6] _feature_perm: torch.Tensor def __init__(self) -> None: @@ -912,7 +912,7 @@ def _o3_linear_codegen( for i_in, i_out in instr ] - outs: list[Any] = list() + outs: list[Any] = [] for (i_in, i_out), w in zip(instr, weights, strict=True): x1_i = x1[:, slices[0][i_in][0] : slices[0][i_in][1]] # type: ignore outs.append( diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 4448d8db..d115460c 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -9,7 +9,7 @@ from collections.abc import Callable, Iterator from copy import copy from dataclasses import dataclass -from typing import Literal, TypeAlias +from typing import TypeAlias import numpy as np import torch @@ -43,9 +43,7 @@ "atomic_grid_size_bound_shape", } -_Float64Coordinates: TypeAlias = np.ndarray[ - tuple[int, Literal[3]], np.dtype[np.float64] -] +_Float64Coordinates: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.float64]] _Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] _SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" @@ -883,6 +881,11 @@ def restore_ao_order(self, matrix: Tensor) -> Tensor: return matrix[..., self.unsort_idx, :][..., self.unsort_idx] def __iter__(self) -> Iterator[_AOBlock]: + block_loop_options: dict[str, bool] = {} + if self.gpu: + # GPU4PySCF otherwise omits zero-AO blocks, shifting all later grid slices. + block_loop_options["strict_grid_order"] = True + end = 0 for ao_block, mask, weights, _ in self.numint.block_loop( mol=self.mol, @@ -891,6 +894,7 @@ def __iter__(self) -> Iterator[_AOBlock]: deriv=self.feature_function.deriv, blksize=self.blksize, non0tab=(None if self.gpu else getattr(self.grids, "non0tab", None)), + **block_loop_options, ): start, end = end, end + weights.size ao = from_numpy_or_cupy( diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 55136411..d058a008 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -300,6 +300,7 @@ def block_loop( self, *args: object, **kwargs: object ) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]]: assert kwargs["non0tab"] is screen_index + assert "strict_grid_order" not in kwargs yield ao, screen_index, grids.weights, grids.coords monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) @@ -461,12 +462,12 @@ def test_global_screened_ao_traversals_are_independent_of_model_chunks( def counting_forward_apply(*args: object) -> torch.Tensor: nonlocal forward_calls forward_calls += 1 - return original_forward_apply(*args) + return original_forward_apply(*args) # type: ignore[no-untyped-call] def counting_backward_apply(*args: object) -> torch.Tensor: nonlocal backward_calls backward_calls += 1 - return original_backward_apply(*args) + return original_backward_apply(*args) # type: ignore[no-untyped-call] monkeypatch.setattr(ChunkEvalForward, "apply", counting_forward_apply) monkeypatch.setattr(ChunkEvalBackward, "apply", counting_backward_apply) diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 564eaa69..03246ef0 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -111,17 +111,10 @@ class BenchmarkSpec(NamedTuple): atoms: str -class BenchmarkCase(NamedTuple): - mol: gto.Mole - grids: dft.Grids - dm: np.ndarray - numint: SkalaNumInt[np.ndarray] - - DeviceResult = tuple[float, float, object] -class DeviceBenchmarkCase(NamedTuple): +class BenchmarkCase(NamedTuple): backend: str mol: gto.Mole run: Callable[[], DeviceResult] @@ -136,43 +129,46 @@ class DeviceBenchmarkCase(NamedTuple): def _make_benchmark_case( - spec: BenchmarkSpec, functional: ExcFunctionalBase + spec: BenchmarkSpec, functional: ExcFunctionalBase, backend: str ) -> BenchmarkCase: mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) - grids = dft.Grids(mol) - grids.level = 1 - grids.build(sort_grids=False) - dm = dft.RKS(mol).get_init_guess() - return BenchmarkCase(mol, grids, dm, SkalaNumInt(functional)) + initial_dm = dft.RKS(mol).get_init_guess() - -def _make_gpu_benchmark_case( - spec: BenchmarkSpec, functional: ExcFunctionalBase -) -> DeviceBenchmarkCase: - import cupy - - from skala.gpu4pyscf import SkalaKS - - mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0) - ks = SkalaKS(mol, xc=functional, with_dftd3=False) - ks.grids.level = 1 - ks.grids.alignment = 1 - ks.grids.build(sort_grids=False) - dm = cupy.asarray(dft.RKS(mol).get_init_guess()) - assert _should_screen_aos(mol) + if backend == "cpu": + grids = dft.Grids(mol) + grids.level = 1 + grids.build(sort_grids=False) + dm: Any = initial_dm + numint: Any = SkalaNumInt(functional) + synchronize: Callable[[], None] = lambda: None # noqa: E731 + elif backend == "cuda": + import cupy + + from skala.gpu4pyscf import SkalaKS + + ks = SkalaKS(mol, xc=functional, with_dftd3=False) + ks.grids.level = 1 + ks.grids.alignment = 1 + ks.grids.build(sort_grids=False) + grids = ks.grids + dm = cupy.asarray(initial_dm) + numint = ks._numint + synchronize = torch.cuda.synchronize + else: + raise ValueError(f"Unknown benchmark backend: {backend}") def run() -> DeviceResult: - result = ks._numint.nr_rks( + result = numint.nr_rks( mol, - ks.grids, + grids, None, dm, max_memory=MAX_MEMORY_MB, ) - torch.cuda.synchronize() + synchronize() return cast(DeviceResult, result) - return DeviceBenchmarkCase("cuda", mol, run, torch.cuda.synchronize) + return BenchmarkCase(backend, mol, run, synchronize) @pytest.fixture(scope="module") @@ -202,7 +198,7 @@ def benchmark_case( ) -> BenchmarkCase: functional = load_functional_cached("skala-1.1") assert isinstance(functional, ExcFunctionalBase) - return _make_benchmark_case(benchmark_spec, functional) + return _make_benchmark_case(benchmark_spec, functional, "cpu") @pytest.fixture( @@ -214,22 +210,22 @@ def device_benchmark_case( benchmark_spec: BenchmarkSpec, fixed_cpu_threads: None, load_functional_cached: Callable[..., ExcFunctionalBase | str], -) -> DeviceBenchmarkCase: +) -> BenchmarkCase: backend = cast(str, request.param) if backend == "cpu": functional = load_functional_cached("skala-1.1") assert isinstance(functional, ExcFunctionalBase) - case = _make_benchmark_case(benchmark_spec, functional) - assert _should_screen_aos(case.mol) - return DeviceBenchmarkCase("cpu", case.mol, lambda: _run_xc(case), lambda: None) - - if not torch.cuda.is_available(): - pytest.skip("CUDA is not available") - pytest.importorskip("cupy") - pytest.importorskip("gpu4pyscf") - functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) - assert isinstance(functional, ExcFunctionalBase) - return _make_gpu_benchmark_case(benchmark_spec, functional) + else: + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + pytest.importorskip("cupy") + pytest.importorskip("gpu4pyscf") + functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) + assert isinstance(functional, ExcFunctionalBase) + + case = _make_benchmark_case(benchmark_spec, functional, backend) + assert _should_screen_aos(case.mol) + return case @pytest.fixture @@ -247,29 +243,7 @@ def dense_case( return benchmark_case -def _run_xc(case: BenchmarkCase) -> tuple[float, float, np.ndarray]: - return case.numint.nr_rks( - case.mol, - case.grids, - None, - case.dm, - max_memory=MAX_MEMORY_MB, - ) - - -def _benchmark_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: - pedantic = cast(Callable[..., object], benchmark.pedantic) - pedantic( - _run_xc, - args=(case,), - rounds=1, - iterations=2, - ) - - -def _benchmark_device_xc( - benchmark: BenchmarkFixture, case: DeviceBenchmarkCase -) -> None: +def _benchmark_device_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: case.synchronize() pedantic = cast(Callable[..., object], benchmark.pedantic) pedantic( @@ -282,7 +256,8 @@ def _benchmark_device_xc( def _run_gpu_xc(spec: BenchmarkSpec, screened: bool) -> int: functional = load_functional("skala-1.1", device=torch.device("cuda:0")) assert isinstance(functional, ExcFunctionalBase) - case = _make_gpu_benchmark_case(spec, functional) + case = _make_benchmark_case(spec, functional, "cuda") + assert _should_screen_aos(case.mol) if not screened: pyscf_numint.SWITCH_SIZE = 10**9 @@ -310,7 +285,7 @@ def _memory_worker( functional = load_functional("skala-1.1") assert isinstance(functional, ExcFunctionalBase) - case = _make_benchmark_case(spec, functional) + case = _make_benchmark_case(spec, functional, "cpu") if not screened: pyscf_numint.SWITCH_SIZE = case.mol.nao_nr() assert _should_screen_aos(case.mol) is screened @@ -318,7 +293,7 @@ def _memory_worker( with tempfile.TemporaryDirectory() as tmpdir: profile_path = Path(tmpdir) / "allocations.bin" with memray.Tracker(profile_path): - _run_xc(case) + case.run() peak_bytes = memray.FileReader(profile_path).metadata.peak_memory elif backend == "cuda": peak_bytes = _run_gpu_xc(spec, screened) @@ -367,7 +342,7 @@ def _measure_peak_memory(spec: BenchmarkSpec, screened: bool, backend: str) -> i @pytest.mark.profiling def test_screened_and_dense_values_agree( - device_benchmark_case: DeviceBenchmarkCase, + device_benchmark_case: BenchmarkCase, benchmark_spec: BenchmarkSpec, load_functional_cached: Callable[..., ExcFunctionalBase | str], monkeypatch: pytest.MonkeyPatch, @@ -383,7 +358,7 @@ def test_screened_and_dense_values_agree( else: cpu_functional = load_functional_cached("skala-1.1", device=torch.device("cpu")) assert isinstance(cpu_functional, ExcFunctionalBase) - dense = _run_xc(_make_benchmark_case(benchmark_spec, cpu_functional)) + dense = _make_benchmark_case(benchmark_spec, cpu_functional, "cpu").run() scalar_rtol = 2e-10 if case.backend == "cpu" else 1e-8 density_close = np.allclose(dense[0], screened[0], rtol=scalar_rtol, atol=1e-11) @@ -424,20 +399,20 @@ def test_screened_and_dense_values_agree( def test_with_natural_ao_screening( benchmark: BenchmarkFixture, screened_case: BenchmarkCase ) -> None: - _benchmark_xc(benchmark, screened_case) + _benchmark_device_xc(benchmark, screened_case) @pytest.mark.benchmark(group="def2-qzvpp") def test_without_ao_screening_by_patching_threshold( benchmark: BenchmarkFixture, dense_case: BenchmarkCase ) -> None: - _benchmark_xc(benchmark, dense_case) + _benchmark_device_xc(benchmark, dense_case) @pytest.mark.benchmark(group="device-def2-qzvpp-screened") def test_screened_runtime_by_device( benchmark: BenchmarkFixture, - device_benchmark_case: DeviceBenchmarkCase, + device_benchmark_case: BenchmarkCase, ) -> None: _benchmark_device_xc(benchmark, device_benchmark_case) @@ -479,7 +454,7 @@ def test_screened_and_dense_peak_memory( @pytest.mark.profiling def test_profile_with_natural_ao_screening( - device_benchmark_case: DeviceBenchmarkCase, + device_benchmark_case: BenchmarkCase, ) -> None: assert _should_screen_aos(device_benchmark_case.mol) device_benchmark_case.run() @@ -487,7 +462,7 @@ def test_profile_with_natural_ao_screening( @pytest.mark.profiling def test_profile_without_ao_screening_by_patching_threshold( - device_benchmark_case: DeviceBenchmarkCase, + device_benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr( diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 40576c2c..7513de9c 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -31,6 +31,7 @@ ChunkEvalForward, MGGAFeatureFunction, _prepare_spatially_sorted_grids, + non_chunk, ) from skala.pyscf.numint import SkalaNumInt # noqa: E402 @@ -326,15 +327,66 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( ) +def test_gpu_empty_ao_block_matches_dense_reference() -> None: + """Preserve grid alignment when GPU4PySCF finds no AOs in a block. + + GPU4PySCF normally omits fixed-size grid blocks whose screening mask contains + no active atomic orbitals. Skala assigns each yielded result to a cumulative + grid slice, so omitting an empty block would shift every later result into the + wrong positions. ``strict_grid_order=True`` makes the backend yield the empty + block and allows Skala to advance that slice before processing active blocks. + + The first block is placed far from the molecule to make its active-AO set + empty, while the second block samples the molecular region. Comparing MGGA + features and their density-matrix VJP with dense AO evaluation verifies both + forward placement and backward slicing through the real GPU backend. + """ + mol = gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0) + assert dft_gpu is not None + block_size = int(dft_gpu.numint.MIN_BLK_SIZE) + far_coords = cupy.full((block_size, 3), 100.0, dtype=cupy.float64) + near_coords = cupy.linspace(-0.5, 0.5, block_size * 3, dtype=cupy.float64).reshape( + block_size, 3 + ) + coords = cupy.concatenate((far_coords, near_coords)) + grids = dft_gpu.Grids(mol) + grids.coords = coords + grids.weights = cupy.ones(coords.shape[0], dtype=cupy.float64) + + screening_numint = dft_gpu.numint.NumInt().build(mol, coords) + active_ao_counts = [ + len(block[1]) for block in grids.get_non0ao_idx(screening_numint.gdftopt) + ] + assert active_ao_counts[0] == 0 + assert active_ao_counts[1] > 0 + grids._non0ao_idx = None + + feature_function = MGGAFeatureFunction( + with_density=True, with_grad=True, with_kin=True + ) + dm = torch.eye(mol.nao_nr(), dtype=torch.float64, device="cuda", requires_grad=True) + screened = ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm, mol, grids, feature_function, block_size, False, True + ) + dense = non_chunk(dm, mol, coords, feature_function, gpu=True) + + torch.testing.assert_close(screened, dense, rtol=1e-12, atol=1e-12) + (screened_vjp,) = torch.autograd.grad(screened.square().sum(), dm) + (dense_vjp,) = torch.autograd.grad(dense.square().sum(), dm) + torch.testing.assert_close(screened_vjp, dense_vjp, rtol=1e-12, atol=5e-8) + + def test_gpu_sparse_mask_sorts_scatters_and_unsorts( monkeypatch: pytest.MonkeyPatch, ) -> None: mol = gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0) - ngrids = 32 + assert dft_gpu is not None + block_size = int(dft_gpu.numint.MIN_BLK_SIZE) + ngrids = 2 * block_size sort_idx = np.array([2, 0, 4, 1, 3]) active_sorted_aos = np.array([0, 2, 4]) - ao = cupy.arange(active_sorted_aos.size * ngrids, dtype=cupy.float64).reshape( - active_sorted_aos.size, ngrids + ao = cupy.arange(active_sorted_aos.size * block_size, dtype=cupy.float64).reshape( + active_sorted_aos.size, block_size ) weights = cupy.ones(ngrids) coords = cupy.zeros((ngrids, 3)) @@ -348,9 +400,20 @@ def build(self, mol: gto.Mole, coords: cupy.ndarray) -> "FakeGpuNumInt": def block_loop( self, *args: object, **kwargs: object ) -> Iterator[tuple[object, object, object, object]]: - yield ao, cupy.asarray(active_sorted_aos), weights, coords + assert kwargs["strict_grid_order"] is True + yield ( + cupy.empty((0, block_size), dtype=cupy.float64), + cupy.empty(0, dtype=cupy.int64), + weights[:block_size], + coords[:block_size], + ) + yield ( + ao, + cupy.asarray(active_sorted_aos), + weights[block_size:], + coords[block_size:], + ) - assert dft_gpu is not None monkeypatch.setattr(dft_gpu.numint, "NumInt", FakeGpuNumInt) feature_function = MGGAFeatureFunction( @@ -368,7 +431,10 @@ def block_loop( dm_sorted = dm[..., sort_idx_t, :][..., sort_idx_t] dm_active = dm_sorted[..., active_t[:, None], active_t[None, :]] ao_t = from_dlpack(ao) - expected = torch.sum((dm_active @ ao_t) * ao_t, dim=0).unsqueeze(0) + expected = torch.zeros_like(features) + expected[..., block_size:] = torch.sum((dm_active @ ao_t) * ao_t, dim=0).unsqueeze( + 0 + ) assert torch.allclose(features, expected) energy = features.square().sum() diff --git a/tests/test_pyscf_gradients.py b/tests/test_pyscf_gradients.py index 4f086b53..48d031d6 100644 --- a/tests/test_pyscf_gradients.py +++ b/tests/test_pyscf_gradients.py @@ -556,7 +556,7 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: exc_test = TestFunc() # Explicitly pass all features including integer ones — should auto-discard them - veff, nuc_grad = veff_and_expl_nuc_grad( + _vexc, nuc_grad = veff_and_expl_nuc_grad( exc_test, mol, grid, rdm1, nuc_grad_feats=set(exc_test.features) ) assert nuc_grad.shape == (mol.natm, 3) From a02c507d9a07a0302c64a8f5b665a94a28e09a01 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 09:52:44 +0200 Subject: [PATCH 11/39] fix consistent screened response memory policy --- src/skala/pyscf/memory_estimators.py | 8 ++++++-- src/skala/pyscf/numint.py | 2 +- tests/test_ao_screening.py | 23 ++++++++++++++++++++++- tests/test_memory_estimators.py | 15 +++++++++++++++ 4 files changed, 44 insertions(+), 4 deletions(-) diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index 56607e31..1efb029c 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -52,10 +52,14 @@ def estimate_max_grid_chunk_size( clamp it to at least the largest atomic grid size. Raises: - ValueError: If ``max_memory_in_mb`` is ``None`` and ``dm`` lives on a device - type other than ``cuda`` or ``cpu`` (supply ``max_memory_in_mb`` instead). + ValueError: If ``safety_fraction`` is outside ``(0, 1]``, or if + ``max_memory_in_mb`` is ``None`` and ``dm`` lives on a device type + other than ``cuda`` or ``cpu`` (supply ``max_memory_in_mb`` instead). RuntimeError: If CPU host memory cannot be determined automatically. """ + if not 0 < safety_fraction <= 1: + raise ValueError("safety_fraction must be greater than 0 and at most 1") + if max_memory_in_mb is None: match dm.device.type: case "cuda": diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index c30644f5..c326a449 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -370,7 +370,7 @@ def gen_response( features=set(self.func.features), func_deriv=2, max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, - safety_fraction=kwargs.get("safety_fraction", 0.0), + safety_fraction=kwargs.get("safety_fraction", 0.8), ) if not screened_features.feature_function.only_linear_feats: raise NotImplementedError( diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index d058a008..6515508f 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -192,14 +192,17 @@ def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: @pytest.mark.parametrize("expected", [False, True]) +@pytest.mark.parametrize("response_safety_fraction", [None, 0.6]) def test_first_and_second_order_use_same_screening_decision( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch, expected: bool, + response_safety_fraction: float | None, ) -> None: switch_size = carbon.nao_nr() - 1 if expected else carbon.nao_nr() monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", switch_size) routes: list[str] = [] + safety_fractions: list[float] = [] def fake_generate_features( mol: gto.Mole, @@ -255,6 +258,9 @@ def fake_global_screened_features( **kwargs: object, ) -> FakeGlobalScreenedFeatures: routes.append("screened") + safety_fraction = kwargs["safety_fraction"] + assert isinstance(safety_fraction, float) + safety_fractions.append(safety_fraction) return FakeGlobalScreenedFeatures(dm) monkeypatch.setattr(numint_module, "generate_features", fake_generate_features) @@ -269,14 +275,29 @@ def fake_global_screened_features( numint(carbon, grids, None, dm) ks = FakeKS(carbon, grids) + response_kwargs = ( + {} + if response_safety_fraction is None + else {"safety_fraction": response_safety_fraction} + ) response = numint.gen_response( - np.eye(carbon.nao_nr()), np.ones(carbon.nao_nr()), ks=ks + np.eye(carbon.nao_nr()), + np.ones(carbon.nao_nr()), + ks=ks, + **response_kwargs, ) result = response(np.eye(carbon.nao_nr())) assert result.shape == (carbon.nao_nr(), carbon.nao_nr()) expected_route = "screened" if expected else "dense" assert routes == [expected_route, expected_route] + if expected: + assert safety_fractions == [ + 0.8, + 0.8 if response_safety_fraction is None else response_safety_fraction, + ] + else: + assert safety_fractions == [] def test_cpu_screening_slices_and_scatters_full_derivatives( diff --git a/tests/test_memory_estimators.py b/tests/test_memory_estimators.py index f95850ab..79821633 100644 --- a/tests/test_memory_estimators.py +++ b/tests/test_memory_estimators.py @@ -55,3 +55,18 @@ def test_reserved_memory_reduces_grid_chunk_size() -> None: ) assert base_chunk_size - reserved_chunk_size == reserved_points + + +@pytest.mark.parametrize("safety_fraction", [-0.1, 0.0, 1.1]) +def test_grid_chunk_size_rejects_invalid_safety_fraction( + safety_fraction: float, +) -> None: + with pytest.raises( + ValueError, match="safety_fraction must be greater than 0 and at most 1" + ): + estimate_max_grid_chunk_size( + torch.eye(2, dtype=torch.float64), + deriv=1, + max_memory_in_mb=100, + safety_fraction=safety_fraction, + ) From 0fc36bea98d067659c9c82f95319e4e965c868af Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 10:17:48 +0200 Subject: [PATCH 12/39] refactor remove unsupported kinetic tensor features --- src/skala/pyscf/features.py | 87 +++------------------------ src/skala/pyscf/numint.py | 5 -- tests/test_ao_screening.py | 117 ++++++++++++++++++++++++++++++++++++ 3 files changed, 124 insertions(+), 85 deletions(-) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index d115460c..da085a40 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -463,7 +463,6 @@ def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: class FeatureFunction(nn.Module, ABC): deriv: int nfeats: int - only_linear_feats: bool @abstractmethod def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ... @@ -477,8 +476,6 @@ class MGGAFeatureFunction(FeatureFunction): with_grad: bool with_kin: bool with_lapl: bool - with_ked_var: bool - with_ked_det: bool def __init__( self, @@ -486,8 +483,6 @@ def __init__( with_grad: bool = True, with_kin: bool = True, with_lapl: bool = False, - with_ked_var: bool = False, - with_ked_det: bool = False, ): super().__init__() @@ -495,29 +490,18 @@ def __init__( self.with_grad = with_grad self.with_kin = with_kin self.with_lapl = with_lapl - self.with_ked_var = with_ked_var - self.with_ked_det = with_ked_det self.deriv = 0 - if with_grad or with_kin or with_ked_var or with_ked_det: + if with_grad or with_kin: self.deriv = 1 if with_lapl: self.deriv = 2 - self.nfeats = ( - with_density - + with_grad * 3 - + with_kin - + with_lapl - + with_ked_var - + with_ked_det - ) + self.nfeats = with_density + with_grad * 3 + with_kin + with_lapl if self.nfeats == 0: raise ValueError("At least one feature must be selected.") - self.only_linear_feats = not (with_ked_var or with_ked_det) - def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: """Convert the features to a dictionary with the keys being the feature names.""" feature_index = 0 @@ -534,17 +518,9 @@ def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: if self.with_lapl: feature_dict["lapl"] = features[..., feature_index, :] feature_index += 1 - if self.with_ked_var: - feature_dict["ked_var"] = features[..., feature_index, :] - feature_index += 1 - if self.with_ked_det: - feature_dict["ked_det"] = features[..., feature_index, :] - feature_index += 1 return feature_dict def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: - with_Q: bool = self.with_ked_var or self.with_ked_det - # Flatten all but the last two dimensions # then restore the original shape at the end dm_view = dm.view(-1, dm.shape[-2], dm.shape[-1]) @@ -580,7 +556,7 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ) feat_idx += 1 - if (self.with_kin or self.with_lapl) and not with_Q: + if self.with_kin or self.with_lapl: for i in range(3): ci = dm_view @ ao[i + 1] features[..., feat_idx, :] += 0.5 * torch.sum( @@ -604,52 +580,6 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: features[..., feat_idx, :] += 2 * torch.sum( c0 * ao[i][None, :, :], dim=-2 ) - - if with_Q: - Q = torch.zeros( - (dm_view.shape[0], ao.shape[-1], 3, 3), device=dm.device, dtype=dm.dtype - ) - - for i in range(3): - ci = dm_view @ ao[i + 1] - for j in range(i, 3): - Q = torch.sum(ci * ao[j + 1][None, :, :], dim=-2) - - if self.with_kin: - features[..., feat_idx, :] = 0.5 * torch.einsum("...ii->...", Q) - feat_idx += 1 - - if self.with_lapl: - features[..., feat_idx, :] = 2 * torch.einsum("...ii->...", Q) - # 0 is without derivative - # 1 2 3 are x y z derivatives - # 4 5 6 are xx xy xz derivatives - # 7 8 9 are yy yz zz derivatives - for i in (4, 7, 9): - features[..., feat_idx, :] += 2 * torch.sum( - c0 * ao[i][None, :, :], dim=-2 - ) - feat_idx += 1 - - if self.with_ked_var: - if not self.with_kin: - trace = torch.einsum("...ii->...", Q) - else: - trace = 2 * features[:, feat_idx - 1, :] - features[..., feat_idx, :] = 0.5 * torch.sum( - ( - trace[:, None, None] - * torch.eye(3, device=dm.device, dtype=dm.dtype)[None, :, :] - - Q - ) - ** 2, - dim=(-2, -1), - ) - feat_idx += 1 - - if self.with_ked_det: - features[..., feat_idx, :] = torch.det(Q) - feat_idx += 1 if len(dm.shape) == 2: return features.reshape((self.nfeats, -1)) else: @@ -676,11 +606,6 @@ class _GlobalScreenedFeatures: def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: """Apply the global raw-feature Jacobian and restore atom-major order.""" - if not self.feature_function.only_linear_feats: - raise NotImplementedError( - "Global screened response requires raw features linear in the density " - "matrix." - ) sorted_tangent = ChunkEvalForward.apply( self.dm, self.mol, @@ -977,7 +902,8 @@ def forward( device=dm.device, dtype=dm.dtype, ) - if len(vectors_jvp) > 1 and feature_function.only_linear_feats: + # Raw AO features are linear in dm, so derivatives above first order vanish. + if len(vectors_jvp) > 1: return features for block in block_loop: @@ -1122,7 +1048,8 @@ def forward( ] out = torch.zeros_like(dm) - if len(vectors) > 1 and feature_function.only_linear_feats: + # Raw AO features are linear in dm, so derivatives above first order vanish. + if len(vectors) > 1: return out for block in block_loop: diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index c326a449..384801ee 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -372,11 +372,6 @@ def gen_response( max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, safety_fraction=kwargs.get("safety_fraction", 0.8), ) - if not screened_features.feature_function.only_linear_feats: - raise NotImplementedError( - "Global screened response requires raw features linear in " - "the density matrix." - ) def hessian_vector_product_atom_chunked(dm1: Array) -> Array: dm1_tensor = self.from_backend(dm1) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 6515508f..5a1c0353 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -25,6 +25,123 @@ def carbon() -> gto.Mole: return gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0) +@pytest.mark.parametrize( + ("options", "expected_deriv", "expected_nfeats", "expected_names"), + [ + ( + { + "with_density": True, + "with_grad": False, + "with_kin": False, + "with_lapl": False, + }, + 0, + 1, + {"density"}, + ), + ( + { + "with_density": False, + "with_grad": True, + "with_kin": False, + "with_lapl": False, + }, + 1, + 3, + {"grad"}, + ), + ( + { + "with_density": False, + "with_grad": False, + "with_kin": True, + "with_lapl": False, + }, + 1, + 1, + {"kin"}, + ), + ( + { + "with_density": False, + "with_grad": False, + "with_kin": False, + "with_lapl": True, + }, + 2, + 1, + {"lapl"}, + ), + ( + { + "with_density": True, + "with_grad": True, + "with_kin": True, + "with_lapl": True, + }, + 2, + 6, + {"density", "grad", "kin", "lapl"}, + ), + ], +) +def test_mgga_supported_features_are_linear_in_density_matrix( + options: dict[str, bool], + expected_deriv: int, + expected_nfeats: int, + expected_names: set[str], +) -> None: + """Check each supported feature layout and its linear dependence on ``dm``. + + Linearity requires the first JVP to equal direct feature evaluation on the + tangent and the second JVP to vanish. + """ + feature_function = MGGAFeatureFunction(**options) + ncomp = (expected_deriv + 1) * (expected_deriv + 2) * (expected_deriv + 3) // 6 + ao = torch.arange(1, ncomp * 2 * 3 + 1, dtype=torch.float64).reshape(ncomp, 2, 3) + if expected_deriv == 0: + ao = ao[0] + dm = torch.tensor([[2.0, 0.5], [0.5, 1.0]], dtype=torch.float64) + tangent = torch.tensor([[0.2, -0.1], [-0.1, 0.3]], dtype=torch.float64) + + features = feature_function(dm, ao) + _, feature_jvp = torch.func.jvp( + lambda value: feature_function(value, ao), + (dm,), + (tangent,), + ) + + def first_jvp(value: torch.Tensor) -> torch.Tensor: + return torch.func.jvp( + lambda inner: feature_function(inner, ao), + (value,), + (tangent,), + )[1] + + _, second_jvp = torch.func.jvp( + first_jvp, + (dm,), + (torch.ones_like(dm),), + ) + + assert feature_function.deriv == expected_deriv + assert feature_function.nfeats == expected_nfeats + assert features.shape == (expected_nfeats, 3) + assert set(feature_function.to_dict(features)) == expected_names + torch.testing.assert_close(feature_jvp, feature_function(tangent, ao)) + torch.testing.assert_close(second_jvp, torch.zeros_like(second_jvp)) + + +def test_mgga_requires_at_least_one_feature() -> None: + with pytest.raises(ValueError, match="At least one feature must be selected"): + MGGAFeatureFunction( + with_density=False, + with_grad=False, + with_kin=False, + with_lapl=False, + ) + + @pytest.mark.parametrize( ("switch_offset", "expected"), [(1, False), (0, False), (-1, True)] ) From a425a10df791ead3edfc12354a738a6ba6a8daf4 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 10:46:29 +0200 Subject: [PATCH 13/39] refactor extract chunk evaluation helpers --- src/skala/pyscf/features.py | 254 +++++++++++++----------------------- tests/test_ao_screening.py | 89 +++++++++++++ 2 files changed, 179 insertions(+), 164 deletions(-) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index da085a40..50ca2678 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -420,7 +420,7 @@ def partial_feature_function_over_aos( """Returns a function that computes the feature function with the given ao, but not the dm already passed to the function. - Purpose is to allow for chaining of derivatives. + Purpose is to allow evaluating a block-local VJP. """ def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: @@ -429,22 +429,6 @@ def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: return partial_feature_function -def partial_jvp_function_over_tangents( - func: Callable[[torch.Tensor], torch.Tensor], - tangents: torch.Tensor, -) -> Callable[[torch.Tensor], torch.Tensor]: - """Returns a function that computes the jvp of the given function with tangents, - but not primals already passed to the function. - - Purpose is to allow for chaining of derivatives over primals.""" - - def reduced_jvp(primals: torch.Tensor) -> torch.Tensor: - _, tangent = torch.func.jvp(func, (primals,), (tangents,)) - return tangent - - return reduced_jvp - - def partial_vjp_function_over_tangents( func: Callable[[torch.Tensor], torch.Tensor], tangents: torch.Tensor, @@ -452,7 +436,7 @@ def partial_vjp_function_over_tangents( """Returns a function that computes the vjp of the given function with tangents, but not primals already passed to the function. - Purpose is to allow for chaining of derivatives over primals.""" + Purpose is to evaluate the feature-space adjoint for one AO block.""" def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: return torch.func.vjp(func, primals)[1](tangents)[0] @@ -763,6 +747,25 @@ def add_to(self, matrix: Tensor, block_result: Tensor) -> None: ) +def _evaluate_feature_block( + feature_function: FeatureFunction, + block: _AOBlock, + active_dm: Tensor, + compile_feature_function: bool, + feature_cotangent: Tensor | None = None, +) -> Tensor: + """Evaluate one active-AO feature block or its feature-space VJP.""" + partial_func = partial_feature_function_over_aos(feature_function, block.ao) + if feature_cotangent is not None: + partial_func = partial_vjp_function_over_tangents( + partial_func, feature_cotangent[..., block.grid_slice] + ) + + if compile_feature_function: + return torch.compile(partial_func)(active_dm) + return partial_func(active_dm) + + class _AOBlockLoop: def __init__( self, @@ -892,8 +895,6 @@ def forward( ) -> torch.Tensor: ngrids = grids.weights.size block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) - dm_ordered = block_loop.order_aos(dm) - vectors_jvp_ordered = [block_loop.order_aos(vector) for vector in vectors_jvp] features = torch.zeros( *dm.shape[:-2], @@ -906,31 +907,41 @@ def forward( if len(vectors_jvp) > 1: return features + # Since the raw feature map is linear, its JVP is direct evaluation on + # the tangent density matrix. + evaluation_dm = vectors_jvp[0] if vectors_jvp else dm + evaluation_dm_ordered = block_loop.order_aos(evaluation_dm) for block in block_loop: - # Apply chain rule for this particular block - partial_func = partial_feature_function_over_aos( + active_dm = block.select_aos(evaluation_dm_ordered) + temp_feature = _evaluate_feature_block( feature_function, - block.ao, + block, + active_dm, + compile_feature_function, ) - for vector in vectors_jvp_ordered: - partial_func = partial_jvp_function_over_tangents( - partial_func, - block.select_aos(vector), - ) - - # Compute feature (or its jvp) for this block with masked dm - active_dm = block.select_aos(dm_ordered) - if compile_feature_function: - temp_feature = torch.compile(partial_func)(active_dm) - else: - temp_feature = partial_func(active_dm) features[..., block.grid_slice] = temp_feature return features @staticmethod - def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor) -> torch.Tensor: - # Chain rule for the jvp + def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor | None) -> torch.Tensor: + if len(ctx.vectors_jvp) > 1: + return torch.zeros( + *ctx.dm.shape[:-2], + ctx.feature_function.nfeats, + ctx.grids.weights.size, + device=ctx.dm.device, + dtype=ctx.dm.dtype, + ) + vector_tangent = grad_inputs[7] if ctx.vectors_jvp else grad_inputs[0] + if vector_tangent is None: + return torch.zeros( + *ctx.dm.shape[:-2], + ctx.feature_function.nfeats, + ctx.grids.weights.size, + device=ctx.dm.device, + dtype=ctx.dm.dtype, + ) return ChunkEvalForward.apply( ctx.dm, ctx.mol, @@ -939,32 +950,28 @@ def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor) -> torch.Tensor: ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors_jvp, - grad_inputs[0], + vector_tangent, ) @staticmethod def backward( ctx: FunctionCtx, *grad_outputs: torch.Tensor ) -> tuple[torch.Tensor | None, ...]: - # After one vjp (backward) the signature of the function changes from dm.shape -> (*dm.shape[:-2], nfeats, ngrid) to dm.shape -> dm.shape - # therefore we move to a different function that does essentially the same thing, but with the new signature - - # Derivative to dm - grads = [ - ChunkEvalBackward.apply( + feature_cotangent = grad_outputs[0] + if ctx.vectors_jvp: + dm_grad = ctx.dm * 0 + else: + dm_grad = ChunkEvalBackward.apply( ctx.dm, ctx.mol, ctx.grids, ctx.feature_function, - ["jvp"] * len(ctx.vectors_jvp) + ["first_vjp"], ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors_jvp, - *grad_outputs, + feature_cotangent, ) - ] + grads = [dm_grad] # We need to provide None for the gradients of the non-differentiable inputs # these are mol (1), grids (2), feature_function (3), blksize (4), @@ -973,25 +980,22 @@ def backward( grads += [None] * num_non_differentiable_inputs - # Gradients of earlier tangents - for i in range(len(ctx.vectors_jvp)): - derivative_types = ["jvp"] * len(ctx.vectors_jvp) - derivative_types[i] = "first_vjp" - grads.append( - ChunkEvalBackward.apply( + # A first JVP is linear in its tangent; higher JVPs are identically zero. + for vector in ctx.vectors_jvp: + if len(ctx.vectors_jvp) == 1: + vector_grad = ChunkEvalBackward.apply( ctx.dm, ctx.mol, ctx.grids, ctx.feature_function, - derivative_types, ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors_jvp[:i], - *grad_outputs, - *ctx.vectors_jvp[i + 1 :], + feature_cotangent, ) - ) + else: + vector_grad = vector * 0 + grads.append(vector_grad) return tuple(grads) @@ -1005,24 +1009,22 @@ def setup_context( gto.Mole, Grid, FeatureFunction, - list[str], int | None, bool, bool, torch.Tensor, ], - output: tuple[torch.Tensor, ...], + output: torch.Tensor, ) -> None: ( ctx.dm, ctx.mol, ctx.grids, ctx.feature_function, - ctx.derivative_types, ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors, + ctx.feature_cotangent, ) = inputs ctx.save_for_backward(ctx.dm) @@ -1032,144 +1034,68 @@ def forward( mol: gto.Mole, grids: Grid, feature_function: FeatureFunction, - derivative_types: list[str], blksize: int | None, compile_feature_function: bool, gpu: bool, - *vectors: torch.Tensor, + feature_cotangent: torch.Tensor, ) -> torch.Tensor: block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) dm_ordered = block_loop.order_aos(dm) - vectors_ordered = [ - block_loop.order_aos(vector) - if derivative_type in ("jvp", "vjp") - else vector - for derivative_type, vector in zip(derivative_types, vectors, strict=True) - ] out = torch.zeros_like(dm) - # Raw AO features are linear in dm, so derivatives above first order vanish. - if len(vectors) > 1: - return out - for block in block_loop: - # Apply chain rule for this particular block - # but be careful with signature change upon first vjp - partial_func = partial_feature_function_over_aos( + active_dm = block.select_aos(dm_ordered) + block_result = _evaluate_feature_block( feature_function, - block.ao, + block, + active_dm, + compile_feature_function, + feature_cotangent, ) - for derivative_type, vector, vector_ordered in zip( - derivative_types, vectors, vectors_ordered, strict=True - ): - if derivative_type == "jvp": - partial_func = partial_jvp_function_over_tangents( - partial_func, - block.select_aos(vector_ordered), - ) - elif derivative_type == "vjp": - partial_func = partial_vjp_function_over_tangents( - partial_func, - block.select_aos(vector_ordered), - ) - elif derivative_type == "first_vjp": - partial_func = partial_vjp_function_over_tangents( - partial_func, vector[..., block.grid_slice] - ) - else: - raise ValueError( - f"Unknown derivative {derivative_type} (must be one of 'jvp', 'vjp', 'first_vjp')" - ) - active_dm = block.select_aos(dm_ordered) - if compile_feature_function: - block_result = torch.compile(partial_func)(active_dm) - else: - block_result = partial_func(active_dm) block.add_to(out, block_result) return block_loop.restore_ao_order(out) @staticmethod - def jvp(ctx: FunctionCtx, *grad_input: torch.Tensor) -> torch.Tensor: - # Chain rule for the jvp + def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor | None) -> torch.Tensor: + feature_cotangent_tangent = grad_inputs[7] + if feature_cotangent_tangent is None: + return torch.zeros_like(ctx.dm) return ChunkEvalBackward.apply( ctx.dm, ctx.mol, ctx.grids, ctx.feature_function, - ctx.derivative_types + ["jvp"], ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors, - grad_input, + feature_cotangent_tangent, ) @staticmethod def backward( ctx: FunctionCtx, *grad_outputs: torch.Tensor ) -> tuple[torch.Tensor | None, ...]: - # Chain rule for the vjp + # The raw feature Jacobian is constant in dm. The only nonzero gradient + # propagates through the feature-space cotangent. + grads = [ctx.dm * 0] + # We need to provide None for the gradients of the non-differentiable inputs + # these are mol (1), grids (2), feature_function (3), blksize (4), + # compile_feature_function (5), gpu (6) + num_non_differentiable_inputs = 6 - # Gradient corresponding to dm - grads = [ - ChunkEvalBackward.apply( + grads += [None] * num_non_differentiable_inputs + grads.append( + ChunkEvalForward.apply( ctx.dm, ctx.mol, ctx.grids, ctx.feature_function, - ctx.derivative_types + ["vjp"], ctx.blksize, ctx.compile_feature_function, ctx.gpu, - *ctx.vectors, - *grad_outputs, + grad_outputs[0], ) - ] - # We need to provide None for the gradients of the non-differentiable inputs - # these are mol (1), grids (2), feature_function (3), derivative_types (4), blksize (5), - # compile_feature_function (6), gpu (7) - num_non_differentiable_inputs = 7 - - grads += [None] * num_non_differentiable_inputs - # Gradients of gradients - for i, derivative_type in enumerate(ctx.derivative_types): - derivative_types = copy(ctx.derivative_types) - if derivative_type == "jvp" or derivative_type == "vjp": - derivative_types[i] = "vjp" - grads.append( - ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - derivative_types, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - *ctx.vectors[:i], - *grad_outputs, - *ctx.vectors[i + 1 :], - ) - ) - elif derivative_type == "first_vjp": - grads.append( - ChunkEvalForward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - *ctx.vectors[:i], - *grad_outputs, - *ctx.vectors[i + 1 :], - ) - ) - else: - raise ValueError( - f"Unknown derivative {derivative_type} (must be one of 'jvp', 'vjp', 'first_vjp')" - ) + ) return tuple(grads) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 5a1c0353..3e621583 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -14,6 +14,8 @@ ChunkEvalForward, MGGAFeatureFunction, _active_cpu_aos, + _AOBlock, + _evaluate_feature_block, _prepare_spatially_sorted_grids, _spatial_grid_permutations, ) @@ -417,6 +419,93 @@ def fake_global_screened_features( assert safety_fractions == [] +def test_feature_block_helper_localizes_derivative_vectors() -> None: + """Use AO slices for linear JVPs and grid slices for feature VJPs.""" + feature_function = MGGAFeatureFunction( + with_density=True, + with_grad=False, + with_kin=False, + ) + block = _AOBlock( + ao=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), + active_aos=torch.tensor([0, 2]), + grid_slice=slice(1, 3), + ) + dm_ordered = torch.tensor( + [[2.0, 0.1, 0.2], [0.1, 1.0, 0.3], [0.2, 0.3, 3.0]], + dtype=torch.float64, + ) + tangent_ordered = torch.tensor( + [[0.5, 1.0, -0.2], [1.0, 0.4, 0.3], [-0.2, 0.3, 0.7]], + dtype=torch.float64, + ) + active_dm = block.select_aos(dm_ordered) + + feature_jvp = _evaluate_feature_block( + feature_function, + block, + block.select_aos(tangent_ordered), + compile_feature_function=False, + ) + expected_jvp = feature_function(block.select_aos(tangent_ordered), block.ao) + torch.testing.assert_close(feature_jvp, expected_jvp) + + full_grid_cotangent = torch.tensor([[10.0, 0.25, -0.5, 20.0]], dtype=torch.float64) + feature_vjp = _evaluate_feature_block( + feature_function, + block, + active_dm, + compile_feature_function=False, + feature_cotangent=full_grid_cotangent, + ) + local_cotangent = full_grid_cotangent[0, block.grid_slice] + expected_vjp = torch.einsum("g,ig,jg->ij", local_cotangent, block.ao, block.ao) + torch.testing.assert_close(feature_vjp, expected_vjp) + + +def test_chunk_eval_transforms_follow_linear_operator(carbon: gto.Mole) -> None: + """Check first and second JVPs and the feature-cotangent adjoint JVP.""" + grids = _minimal_atom_grid(carbon) + feature_function = MGGAFeatureFunction( + with_density=True, + with_grad=False, + with_kin=False, + ) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + tangent = torch.arange(1, dm.numel() + 1, dtype=dm.dtype).reshape(dm.shape) + + def evaluate(value: torch.Tensor) -> torch.Tensor: + return ChunkEvalForward.apply( # type: ignore[no-untyped-call] + value, carbon, grids, feature_function, None, False, False + ) + + features, feature_tangent = torch.func.jvp(evaluate, (dm,), (tangent,)) + torch.testing.assert_close(feature_tangent, evaluate(tangent)) + + def first_jvp(value: torch.Tensor) -> torch.Tensor: + return torch.func.jvp(evaluate, (value,), (tangent,))[1] + + _, second_jvp = torch.func.jvp(first_jvp, (dm,), (torch.ones_like(dm),)) + torch.testing.assert_close(second_jvp, torch.zeros_like(features)) + + feature_cotangent = torch.arange( + 1, features.numel() + 1, dtype=features.dtype + ).reshape(features.shape) + cotangent_tangent = torch.flip(feature_cotangent, dims=(-1,)) + + def apply_adjoint(value: torch.Tensor) -> torch.Tensor: + return ChunkEvalBackward.apply( # type: ignore[no-untyped-call] + dm, carbon, grids, feature_function, None, False, False, value + ) + + _, adjoint_tangent = torch.func.jvp( + apply_adjoint, + (feature_cotangent,), + (cotangent_tangent,), + ) + torch.testing.assert_close(adjoint_tangent, apply_adjoint(cotangent_tangent)) + + def test_cpu_screening_slices_and_scatters_full_derivatives( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: From a6676cf1d16f4a3c02fd46d87d4ead36369c19e8 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 12:16:52 +0200 Subject: [PATCH 14/39] refactor centralize feature evaluation policy --- .../cpp/cpp_integration/prepare_inputs.py | 12 +- src/skala/functional/base.py | 4 +- src/skala/pyscf/dft.py | 4 +- src/skala/pyscf/evaluation.py | 86 ++++++++++++ src/skala/pyscf/features.py | 125 ++++++------------ src/skala/pyscf/numint.py | 29 ++-- tests/test_ao_screening.py | 105 +++------------ tests/test_evaluation.py | 84 ++++++++++++ tests/test_gpu4pyscf_ao_screening.py | 9 +- 9 files changed, 257 insertions(+), 201 deletions(-) create mode 100644 src/skala/pyscf/evaluation.py create mode 100644 tests/test_evaluation.py diff --git a/examples/cpp/cpp_integration/prepare_inputs.py b/examples/cpp/cpp_integration/prepare_inputs.py index c5d45628..d3c103fe 100755 --- a/examples/cpp/cpp_integration/prepare_inputs.py +++ b/examples/cpp/cpp_integration/prepare_inputs.py @@ -7,12 +7,9 @@ from pyscf import dft, gto from pyscf.dft import gen_grid +from skala.functional.model import SkalaFunctional from skala.functional.traditional import LDA -from skala.pyscf.features import ( - _ATOMIC_GRID_FEATURES, - DEFAULT_FEATURES_SET, - generate_features, -) +from skala.pyscf.features import generate_features def main() -> None: @@ -45,12 +42,9 @@ def main() -> None: grid.level = 3 grid.build(sort_grids=False) features = generate_features( - molecule, dm, grid, features=DEFAULT_FEATURES_SET | _ATOMIC_GRID_FEATURES + molecule, dm, grid, features=set(SkalaFunctional.features) ) - # Add a feature called `coarse_0_atomic_coords` containing the atomic coordinates. - features["coarse_0_atomic_coords"] = torch.from_numpy(molecule.atom_coords()) - # Save all features as individual .pt files. for key, value in features.items(): torch.save(value, str(args.output_dir / f"{key}.pt")) diff --git a/src/skala/functional/base.py b/src/skala/functional/base.py index cbb40705..cd6137bf 100644 --- a/src/skala/functional/base.py +++ b/src/skala/functional/base.py @@ -8,7 +8,7 @@ """ from collections.abc import Callable -from typing import Any, ClassVar +from typing import Any import torch from torch import nn @@ -25,7 +25,7 @@ class ExcFunctionalBase(nn.Module): energy density from molecular features. """ - features: ClassVar[list[str]] + features: list[str] """List of features that this functional requires.""" def get_d3_settings(self) -> str | None: diff --git a/src/skala/pyscf/dft.py b/src/skala/pyscf/dft.py index 62786922..499bf6f8 100644 --- a/src/skala/pyscf/dft.py +++ b/src/skala/pyscf/dft.py @@ -62,7 +62,7 @@ from pyscf.df import df_jk from skala.functional.base import ExcFunctionalBase -from skala.pyscf.features import _ATOMIC_GRID_FEATURES +from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient from skala.pyscf.grids import UnsortableGrids from skala.pyscf.numint import SkalaNumInt @@ -73,7 +73,7 @@ def _needs_unsorted_grids(func: ExcFunctionalBase) -> bool: """Return True when the functional needs per-atom grid ordering.""" - return bool(set(func.features) & _ATOMIC_GRID_FEATURES) + return FeatureSpec(func.features).requires_atomic_layout def _build_grids_unsorted( diff --git a/src/skala/pyscf/evaluation.py b/src/skala/pyscf/evaluation.py new file mode 100644 index 00000000..9c4eb495 --- /dev/null +++ b/src/skala/pyscf/evaluation.py @@ -0,0 +1,86 @@ +# SPDX-License-Identifier: MIT + +"""Feature requirements and numerical-evaluation policy.""" + +from collections.abc import Iterable +from dataclasses import dataclass + +_MGGA_FEATURES = frozenset({"density", "grad", "kin", "lapl"}) +_ATOMIC_LAYOUT_FEATURES = frozenset( + { + "atomic_grid_weights", + "atomic_grid_sizes", + "atomic_grid_size_bound_shape", + } +) + + +@dataclass(frozen=True, init=False) +class FeatureSpec: + """Normalized feature names and their evaluation requirements.""" + + names: frozenset[str] + + def __init__(self, names: Iterable[str]) -> None: + object.__setattr__(self, "names", frozenset(names)) + + def requests(self, feature: str) -> bool: + """Return whether a feature is requested.""" + return feature in self.names + + @property + def with_density(self) -> bool: + """Return whether density is requested.""" + return self.requests("density") + + @property + def with_grad(self) -> bool: + """Return whether the density gradient is requested.""" + return self.requests("grad") + + @property + def with_kin(self) -> bool: + """Return whether kinetic-energy density is requested.""" + return self.requests("kin") + + @property + def with_lapl(self) -> bool: + """Return whether the density Laplacian is requested.""" + return self.requests("lapl") + + @property + def requires_mgga(self) -> bool: + """Return whether AO-based meta-GGA features are requested.""" + return bool(self.names & _MGGA_FEATURES) + + @property + def mgga_feature_count(self) -> int: + """Return the scalar width of the requested meta-GGA features.""" + return self.with_density + 3 * self.with_grad + self.with_kin + self.with_lapl + + @property + def ao_derivative_order(self) -> int: + """Return the highest AO derivative order needed by the features.""" + if "lapl" in self.names: + return 2 + if self.names & {"grad", "kin"}: + return 1 + return 0 + + @property + def requires_atomic_layout(self) -> bool: + """Return whether grid points must retain per-atom ordering.""" + return bool(self.names & _ATOMIC_LAYOUT_FEATURES) + + @property + def supports_screened_evaluation(self) -> bool: + """Return whether atom-aligned screened evaluation is supported.""" + return "atomic_grid_sizes" in self.names + + +@dataclass(frozen=True) +class EvaluationPolicy: + """Settings shared by dense and screened AO feature evaluation.""" + + ao_block_size: int | None = None + safety_fraction: float = 0.8 diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 50ca2678..0a59eb62 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -25,6 +25,7 @@ dft_gpu, from_numpy_or_cupy, ) +from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec from skala.pyscf.memory_estimators import ( estimate_global_raw_feature_buffer_memory, estimate_max_grid_chunk_size, @@ -36,13 +37,6 @@ DEFAULT_FEATURES_SET = set(DEFAULT_FEATURES) CPU_AO_SCREENING_BLOCK_SIZE = 9 * dft.gen_grid.BLKSIZE -# Features that require per-atom grid decomposition. -_ATOMIC_GRID_FEATURES = { - "atomic_grid_weights", - "atomic_grid_sizes", - "atomic_grid_size_bound_shape", -} - _Float64Coordinates: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.float64]] _Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] _SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" @@ -307,7 +301,8 @@ def generate_features( A dictionary containing the requested features. The keys are the feature names, and the values are the corresponding tensors. """ - features = features or DEFAULT_FEATURES_SET + feature_spec = FeatureSpec(DEFAULT_FEATURES_SET if features is None else features) + evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) # if dm is a 3D tensor, then we have a spin-polarized system with_spin = len(dm.shape) == 3 @@ -315,28 +310,17 @@ def generate_features( if gpu and dm.device.type != "cuda": raise ValueError("Density matrix must be on the GPU when gpu=True.") - mol_features = get_grid_features(mol, dm, grids, features) + mol_features = get_grid_features(mol, dm, grids, feature_spec) - with_mgga_feature = ( - "density" in features - or "grad" in features - or "kin" in features - or "lapl" in features - ) - if with_mgga_feature: + if feature_spec.requires_mgga: mgga_features = auto_chunk( dm, mol, grids, - MGGAFeatureFunction( - with_density="density" in features, - with_grad="grad" in features, - with_kin="kin" in features, - with_lapl="lapl" in features, - ), - block_size=chunk_size, + MGGAFeatureFunction(feature_spec), + block_size=evaluation_policy.ao_block_size, max_memory=max_memory, - fix_block_size=chunk_size is None, + fix_block_size=evaluation_policy.ao_block_size is None, gpu=gpu, ) @@ -352,26 +336,26 @@ def get_grid_features( mol: gto.Mole, dm: Tensor, grids: Grid, - requested_features: set[str], + feature_spec: FeatureSpec, ) -> dict[str, Tensor]: grid_features = {} - if "grid_coords" in requested_features: + if feature_spec.requests("grid_coords"): grid_features["grid_coords"] = from_numpy_or_cupy( grids.coords, device=dm.device, dtype=dm.dtype ) - if "grid_weights" in requested_features: + if feature_spec.requests("grid_weights"): grid_features["grid_weights"] = from_numpy_or_cupy( grids.weights, device=dm.device, dtype=dm.dtype ) - if "coarse_0_atomic_coords" in requested_features: + if feature_spec.requests("coarse_0_atomic_coords"): grid_features["coarse_0_atomic_coords"] = from_numpy_or_cupy( mol.atom_coords(), device=dm.device, dtype=dm.dtype ) - if requested_features & _ATOMIC_GRID_FEATURES: + if feature_spec.requires_atomic_layout: atom_grids_tab = grids.gen_atomic_grids( mol, grids.atom_grid, grids.radi_method, grids.level, grids.prune ) @@ -387,18 +371,18 @@ def get_grid_features( f"Set grids.alignment = 1 before building grids to disable padding." ) - if "atomic_grid_sizes" in requested_features: + if feature_spec.requests("atomic_grid_sizes"): grid_features["atomic_grid_sizes"] = torch.tensor( sizes, dtype=torch.long, device=dm.device ) - if "atomic_grid_size_bound_shape" in requested_features: + if feature_spec.requests("atomic_grid_size_bound_shape"): max_size = max(sizes) grid_features["atomic_grid_size_bound_shape"] = torch.zeros( max_size, 0, dtype=torch.long, device=dm.device ) - if "atomic_grid_weights" in requested_features: + if feature_spec.requests("atomic_grid_weights"): raw_weights = np.concatenate( [atom_grids_tab[mol.atom_symbol(ia)][1] for ia in range(mol.natm)] ) @@ -456,50 +440,29 @@ def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: ... class MGGAFeatureFunction(FeatureFunction): - with_density: bool - with_grad: bool - with_kin: bool - with_lapl: bool - - def __init__( - self, - with_density: bool = True, - with_grad: bool = True, - with_kin: bool = True, - with_lapl: bool = False, - ): + def __init__(self, feature_spec: FeatureSpec): super().__init__() - self.with_density = with_density - self.with_grad = with_grad - self.with_kin = with_kin - self.with_lapl = with_lapl - - self.deriv = 0 - if with_grad or with_kin: - self.deriv = 1 - if with_lapl: - self.deriv = 2 - - self.nfeats = with_density + with_grad * 3 + with_kin + with_lapl - - if self.nfeats == 0: + if not feature_spec.requires_mgga: raise ValueError("At least one feature must be selected.") + self.feature_spec = feature_spec + self.deriv = feature_spec.ao_derivative_order + self.nfeats = feature_spec.mgga_feature_count def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: """Convert the features to a dictionary with the keys being the feature names.""" feature_index = 0 feature_dict: dict[str, torch.Tensor] = {} - if self.with_density: + if self.feature_spec.with_density: feature_dict["density"] = features[..., feature_index, :] feature_index += 1 - if self.with_grad: + if self.feature_spec.with_grad: feature_dict["grad"] = features[..., feature_index : feature_index + 3, :] feature_index += 3 - if self.with_kin: + if self.feature_spec.with_kin: feature_dict["kin"] = features[..., feature_index, :] feature_index += 1 - if self.with_lapl: + if self.feature_spec.with_lapl: feature_dict["lapl"] = features[..., feature_index, :] feature_index += 1 return feature_dict @@ -529,33 +492,33 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: c0 = dm_view @ ao[0] feat_idx = 0 - if self.with_density: + if self.feature_spec.with_density: features[..., feat_idx, :] = torch.sum(c0 * ao[0][None, :, :], dim=-2) feat_idx += 1 - if self.with_grad: + if self.feature_spec.with_grad: for i in range(3): features[..., feat_idx, :] = 2 * torch.sum( c0 * ao[i + 1][None, :, :], dim=-2 ) feat_idx += 1 - if self.with_kin or self.with_lapl: + if self.feature_spec.with_kin or self.feature_spec.with_lapl: for i in range(3): ci = dm_view @ ao[i + 1] features[..., feat_idx, :] += 0.5 * torch.sum( ci * ao[i + 1][None, :, :], dim=-2 ) - if self.with_kin: + if self.feature_spec.with_kin: feat_idx += 1 - if self.with_lapl: + if self.feature_spec.with_lapl: features[..., feat_idx, :] = 4 * features[..., feat_idx - 1, :] else: # Multiply times four for the laplacian features[..., feat_idx, :] *= 4.0 - if self.with_lapl: + if self.feature_spec.with_lapl: # 0 is without derivative # 1 2 3 are x y z derivatives # 4 5 6 are xx xy xz derivatives @@ -584,7 +547,7 @@ class _GlobalScreenedFeatures: compile_feature_function: bool gpu: bool grid_features: dict[str, Tensor] - feature_names: set[str] + feature_spec: FeatureSpec chunks: list[tuple[slice, slice]] with_spin: bool @@ -611,18 +574,18 @@ def build_model_chunk( """Build one atom-aligned model dictionary from raw feature values.""" feature_chunk: dict[str, Tensor] = {} for feature_name in ("grid_coords", "grid_weights", "atomic_grid_weights"): - if feature_name in self.feature_names: + if self.feature_spec.requests(feature_name): feature_chunk[feature_name] = self.grid_features[feature_name][ grid_slice ] for feature_name in ("coarse_0_atomic_coords", "atomic_grid_sizes"): - if feature_name in self.feature_names: + if self.feature_spec.requests(feature_name): feature_chunk[feature_name] = self.grid_features[feature_name][ atom_slice ] - if "atomic_grid_size_bound_shape" in self.feature_names: + if self.feature_spec.requests("atomic_grid_size_bound_shape"): max_size = int(feature_chunk["atomic_grid_sizes"].max().item()) feature_chunk["atomic_grid_size_bound_shape"] = torch.zeros( max_size, @@ -644,27 +607,25 @@ def _global_screened_features( mol: gto.Mole, dm: Tensor, grids: Grid, - features: set[str], + features: FeatureSpec | set[str], func_deriv: int, max_memory_in_mb: int | None = None, safety_fraction: float = 0.8, compile_feature_function: bool = False, ) -> _GlobalScreenedFeatures: """Evaluate raw AO features once on a spatially ordered molecular grid.""" - if "atomic_grid_sizes" not in features: + feature_spec = ( + features if isinstance(features, FeatureSpec) else FeatureSpec(features) + ) + if not feature_spec.supports_screened_evaluation: raise ValueError( "Global screened features require 'atomic_grid_sizes' for model chunks." ) if grids.coords is None or grids.weights is None: raise ValueError("Grids must be built before generating screened features.") - feature_function = MGGAFeatureFunction( - with_density="density" in features, - with_grad="grad" in features, - with_kin="kin" in features, - with_lapl="lapl" in features, - ) - grid_features = get_grid_features(mol, dm, grids, features) + feature_function = MGGAFeatureFunction(feature_spec) + grid_features = get_grid_features(mol, dm, grids, feature_spec) max_grid_chunk_size = estimate_max_grid_chunk_size( dm=dm, deriv=feature_function.deriv, @@ -721,7 +682,7 @@ def _global_screened_features( compile_feature_function=compile_feature_function, gpu=gpu, grid_features=grid_features, - feature_names=features, + feature_spec=feature_spec, chunks=chunks, with_spin=dm.ndim == 3, ) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 384801ee..5cab4754 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -18,6 +18,7 @@ to_cupy, to_numpy, ) +from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec from skala.pyscf.features import ( _global_screened_features, generate_features, @@ -144,7 +145,8 @@ def __init__( check_gpu_imports_were_successful() self.func = functional.to(device=self.device) - self.chunk_size = chunk_size + self.feature_spec = FeatureSpec(self.func.features) + self.evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) def from_backend( self, @@ -182,7 +184,7 @@ def get_rho( self.from_backend(dm), grids, features={"density"}, - chunk_size=self.chunk_size, + chunk_size=self.evaluation_policy.ao_block_size, max_memory=max_memory, gpu=self.device.type == "cuda", ) @@ -216,7 +218,7 @@ def __call__( f"Density matrix device {dm.device} does not match functional device {self.device}" ) - if self._functional_supports_atom_chunking() and _should_screen_aos(mol): + if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): dm = dm.detach().requires_grad_() tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) @@ -224,10 +226,10 @@ def __call__( mol, dm, grids, - features=set(self.func.features), + features=self.feature_spec, func_deriv=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, - safety_fraction=0.8, + safety_fraction=self.evaluation_policy.safety_fraction, ) # Store only full-grid feature cotangents; model activations remain chunk-local. atom_major_cotangent = torch.zeros_like( @@ -275,8 +277,8 @@ def __call__( mol, dm, grids, - set(self.func.features), - chunk_size=self.chunk_size, + set(self.feature_spec.names), + chunk_size=self.evaluation_policy.ao_block_size, max_memory=max_memory, gpu=self.device.type == "cuda", ) @@ -361,16 +363,20 @@ def gen_response( dm0 = self.from_backend(ks.make_rdm1(mo_coeff, mo_occ)) - if self._functional_supports_atom_chunking() and _should_screen_aos(ks.mol): + if self.feature_spec.supports_screened_evaluation and _should_screen_aos( + ks.mol + ): dm0 = dm0.requires_grad_() screened_features = _global_screened_features( ks.mol, dm0, ks.grids, - features=set(self.func.features), + features=self.feature_spec, func_deriv=2, max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, - safety_fraction=kwargs.get("safety_fraction", 0.8), + safety_fraction=kwargs.get( + "safety_fraction", self.evaluation_policy.safety_fraction + ), ) def hessian_vector_product_atom_chunked(dm1: Array) -> Array: @@ -458,6 +464,3 @@ def hessian_vector_product(dm1: Array) -> Array: return v1 return hessian_vector_product - - def _functional_supports_atom_chunking(self) -> bool: - return "atomic_grid_sizes" in self.func.features diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 3e621583..657a31d6 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -9,6 +9,7 @@ from skala.functional.base import ExcFunctionalBase from skala.pyscf import features as features_module from skala.pyscf import numint as numint_module +from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.features import ( ChunkEvalBackward, ChunkEvalForward, @@ -28,77 +29,27 @@ def carbon() -> gto.Mole: @pytest.mark.parametrize( - ("options", "expected_deriv", "expected_nfeats", "expected_names"), + ("feature_names", "expected_deriv", "expected_nfeats"), [ - ( - { - "with_density": True, - "with_grad": False, - "with_kin": False, - "with_lapl": False, - }, - 0, - 1, - {"density"}, - ), - ( - { - "with_density": False, - "with_grad": True, - "with_kin": False, - "with_lapl": False, - }, - 1, - 3, - {"grad"}, - ), - ( - { - "with_density": False, - "with_grad": False, - "with_kin": True, - "with_lapl": False, - }, - 1, - 1, - {"kin"}, - ), - ( - { - "with_density": False, - "with_grad": False, - "with_kin": False, - "with_lapl": True, - }, - 2, - 1, - {"lapl"}, - ), - ( - { - "with_density": True, - "with_grad": True, - "with_kin": True, - "with_lapl": True, - }, - 2, - 6, - {"density", "grad", "kin", "lapl"}, - ), + ({"density"}, 0, 1), + ({"grad"}, 1, 3), + ({"kin"}, 1, 1), + ({"lapl"}, 2, 1), + ({"density", "grad", "kin", "lapl"}, 2, 6), ], ) def test_mgga_supported_features_are_linear_in_density_matrix( - options: dict[str, bool], + feature_names: set[str], expected_deriv: int, expected_nfeats: int, - expected_names: set[str], ) -> None: """Check each supported feature layout and its linear dependence on ``dm``. Linearity requires the first JVP to equal direct feature evaluation on the tangent and the second JVP to vanish. """ - feature_function = MGGAFeatureFunction(**options) + feature_spec = FeatureSpec(feature_names) + feature_function = MGGAFeatureFunction(feature_spec) ncomp = (expected_deriv + 1) * (expected_deriv + 2) * (expected_deriv + 3) // 6 ao = torch.arange(1, ncomp * 2 * 3 + 1, dtype=torch.float64).reshape(ncomp, 2, 3) if expected_deriv == 0: @@ -128,20 +79,16 @@ def first_jvp(value: torch.Tensor) -> torch.Tensor: assert feature_function.deriv == expected_deriv assert feature_function.nfeats == expected_nfeats + assert feature_function.feature_spec is feature_spec assert features.shape == (expected_nfeats, 3) - assert set(feature_function.to_dict(features)) == expected_names + assert set(feature_function.to_dict(features)) == feature_names torch.testing.assert_close(feature_jvp, feature_function(tangent, ao)) torch.testing.assert_close(second_jvp, torch.zeros_like(second_jvp)) def test_mgga_requires_at_least_one_feature() -> None: with pytest.raises(ValueError, match="At least one feature must be selected"): - MGGAFeatureFunction( - with_density=False, - with_grad=False, - with_kin=False, - with_lapl=False, - ) + MGGAFeatureFunction(FeatureSpec([])) @pytest.mark.parametrize( @@ -341,11 +288,7 @@ def fake_generate_features( class FakeGlobalScreenedFeatures: def __init__(self, dm: torch.Tensor) -> None: raw_features = dm.sum().reshape(1, 1) - self.feature_function = MGGAFeatureFunction( - with_density=True, - with_grad=False, - with_kin=False, - ) + self.feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) self.sorted_raw_features = raw_features self.atom_major_raw_features = raw_features self.forward_permutation = torch.tensor([0]) @@ -421,11 +364,7 @@ def fake_global_screened_features( def test_feature_block_helper_localizes_derivative_vectors() -> None: """Use AO slices for linear JVPs and grid slices for feature VJPs.""" - feature_function = MGGAFeatureFunction( - with_density=True, - with_grad=False, - with_kin=False, - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) block = _AOBlock( ao=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), active_aos=torch.tensor([0, 2]), @@ -466,11 +405,7 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: def test_chunk_eval_transforms_follow_linear_operator(carbon: gto.Mole) -> None: """Check first and second JVPs and the feature-cotangent adjoint JVP.""" grids = _minimal_atom_grid(carbon) - feature_function = MGGAFeatureFunction( - with_density=True, - with_grad=False, - with_kin=False, - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) tangent = torch.arange(1, dm.numel() + 1, dtype=dm.dtype).reshape(dm.shape) @@ -532,9 +467,7 @@ def block_loop( monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) - feature_function = MGGAFeatureFunction( - with_density=True, with_grad=False, with_kin=False - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) dm = torch.diag( torch.arange(1, carbon.nao_nr() + 1, dtype=torch.float64) ).requires_grad_() @@ -577,9 +510,7 @@ def block_loop( yield ao, screen_index, grids.weights, grids.coords monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) - feature_function = MGGAFeatureFunction( - with_density=True, with_grad=False, with_kin=False - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py new file mode 100644 index 00000000..159f2488 --- /dev/null +++ b/tests/test_evaluation.py @@ -0,0 +1,84 @@ +from dataclasses import FrozenInstanceError + +import pytest +import torch + +from skala.functional.base import ExcFunctionalBase +from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec +from skala.pyscf.features import generate_features +from skala.pyscf.numint import SkalaNumInt + + +@pytest.mark.parametrize( + ("features", "expected_order"), + [ + ([], 0), + (["density"], 0), + (["grad"], 1), + (["kin"], 1), + (["lapl"], 2), + (["density", "grad", "kin", "lapl"], 2), + ], +) +def test_feature_spec_derives_mgga_requirements( + features: list[str], expected_order: int +) -> None: + spec = FeatureSpec(features) + + assert spec.requires_mgga is bool(features) + assert spec.ao_derivative_order == expected_order + assert spec.with_density is ("density" in features) + assert spec.with_grad is ("grad" in features) + assert spec.with_kin is ("kin" in features) + assert spec.with_lapl is ("lapl" in features) + + +@pytest.mark.parametrize( + ("feature", "supports_screened_evaluation"), + [ + ("atomic_grid_weights", False), + ("atomic_grid_sizes", True), + ("atomic_grid_size_bound_shape", False), + ], +) +def test_feature_spec_derives_atomic_layout_requirements( + feature: str, supports_screened_evaluation: bool +) -> None: + spec = FeatureSpec([feature, feature]) + + assert spec.names == frozenset({feature}) + assert spec.requires_atomic_layout + assert spec.supports_screened_evaluation is supports_screened_evaluation + + +def test_evaluation_policy_defaults_and_is_immutable() -> None: + policy = EvaluationPolicy() + + assert policy.ao_block_size is None + assert policy.safety_fraction == 0.8 + with pytest.raises(FrozenInstanceError): + policy.safety_fraction = 0.5 # type: ignore[misc] + + +def test_explicit_empty_feature_set_stays_empty() -> None: + features = generate_features( + mol=object(), + dm=torch.eye(1), + grids=object(), + features=set(), + ) + + assert features == {} + + +class DensityFunctional(ExcFunctionalBase): + def __init__(self) -> None: + super().__init__() + self.features = ["density"] + + +def test_numint_translates_chunk_size_into_evaluation_policy() -> None: + numint = SkalaNumInt(DensityFunctional(), chunk_size=96) + + assert numint.feature_spec == FeatureSpec(["density"]) + assert numint.evaluation_policy == EvaluationPolicy(ao_block_size=96) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 7513de9c..977b5ebd 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -27,6 +27,7 @@ from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.pyscf.backend import dft_gpu # noqa: E402 +from skala.pyscf.evaluation import FeatureSpec # noqa: E402 from skala.pyscf.features import ( # noqa: E402 ChunkEvalForward, MGGAFeatureFunction, @@ -361,9 +362,7 @@ def test_gpu_empty_ao_block_matches_dense_reference() -> None: assert active_ao_counts[1] > 0 grids._non0ao_idx = None - feature_function = MGGAFeatureFunction( - with_density=True, with_grad=True, with_kin=True - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density", "grad", "kin"])) dm = torch.eye(mol.nao_nr(), dtype=torch.float64, device="cuda", requires_grad=True) screened = ChunkEvalForward.apply( # type: ignore[no-untyped-call] dm, mol, grids, feature_function, block_size, False, True @@ -416,9 +415,7 @@ def block_loop( monkeypatch.setattr(dft_gpu.numint, "NumInt", FakeGpuNumInt) - feature_function = MGGAFeatureFunction( - with_density=True, with_grad=False, with_kin=False - ) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) dm = torch.diag( torch.arange(1, mol.nao_nr() + 1, dtype=torch.float64, device="cuda") ).requires_grad_() From 5783c7a76725cbc5844a5b34dc938b3388d47782 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 13:33:15 +0200 Subject: [PATCH 15/39] refactor isolate feature mathematics --- src/skala/pyscf/feature_math.py | 119 +++++++++++++++++++++ src/skala/pyscf/features.py | 151 +++------------------------ tests/test_ao_screening.py | 2 +- tests/test_gpu4pyscf_ao_screening.py | 2 +- 4 files changed, 136 insertions(+), 138 deletions(-) create mode 100644 src/skala/pyscf/feature_math.py diff --git a/src/skala/pyscf/feature_math.py b/src/skala/pyscf/feature_math.py new file mode 100644 index 00000000..5c1cf617 --- /dev/null +++ b/src/skala/pyscf/feature_math.py @@ -0,0 +1,119 @@ +# SPDX-License-Identifier: MIT + +"""Raw density-feature mathematics and model formatting.""" + +from abc import ABC, abstractmethod + +import torch +from torch import nn + +from skala.pyscf.evaluation import FeatureSpec + + +def maybe_expand_and_divide( + feature: torch.Tensor, expand: bool, divisor: float +) -> torch.Tensor: + """Expand a feature across spin channels and divide it when requested.""" + if expand: + return torch.stack([feature / divisor, feature / divisor], dim=0) + return feature + + +class FeatureFunction(nn.Module, ABC): + """Base class for raw features evaluated from density and AO tensors.""" + + deriv: int + nfeats: int + + @abstractmethod + def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ... + + @abstractmethod + def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: ... + + +class MGGAFeatureFunction(FeatureFunction): + """Evaluate the requested linear meta-GGA density features.""" + + def __init__(self, feature_spec: FeatureSpec): + super().__init__() + + if not feature_spec.requires_mgga: + raise ValueError("At least one feature must be selected.") + self.feature_spec = feature_spec + self.deriv = feature_spec.ao_derivative_order + self.nfeats = feature_spec.mgga_feature_count + + def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: + """Convert a packed feature tensor to its named feature tensors.""" + feature_index = 0 + feature_dict: dict[str, torch.Tensor] = {} + if self.feature_spec.with_density: + feature_dict["density"] = features[..., feature_index, :] + feature_index += 1 + if self.feature_spec.with_grad: + feature_dict["grad"] = features[..., feature_index : feature_index + 3, :] + feature_index += 3 + if self.feature_spec.with_kin: + feature_dict["kin"] = features[..., feature_index, :] + feature_index += 1 + if self.feature_spec.with_lapl: + feature_dict["lapl"] = features[..., feature_index, :] + return feature_dict + + def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: + dm_view = dm.view(-1, dm.shape[-2], dm.shape[-1]) + dm_view = 0.5 * (dm_view + dm_view.transpose(-1, -2)) + + features = torch.zeros( + (dm_view.shape[0], self.nfeats, ao.shape[-1]), + device=dm.device, + dtype=dm.dtype, + ) + + if self.deriv == 0: + c0 = dm_view @ ao + features[..., 0, :] = torch.sum(c0 * ao[None, :, :], dim=-2) + if len(dm.shape) == 2: + return features.reshape((self.nfeats, -1)) + return features.reshape((*dm.shape[:-2], self.nfeats, -1)) + + c0 = dm_view @ ao[0] + + feature_index = 0 + if self.feature_spec.with_density: + features[..., feature_index, :] = torch.sum(c0 * ao[0][None, :, :], dim=-2) + feature_index += 1 + + if self.feature_spec.with_grad: + for component in range(3): + features[..., feature_index, :] = 2 * torch.sum( + c0 * ao[component + 1][None, :, :], dim=-2 + ) + feature_index += 1 + + if self.feature_spec.with_kin or self.feature_spec.with_lapl: + for component in range(3): + ci = dm_view @ ao[component + 1] + features[..., feature_index, :] += 0.5 * torch.sum( + ci * ao[component + 1][None, :, :], dim=-2 + ) + + if self.feature_spec.with_kin: + feature_index += 1 + if self.feature_spec.with_lapl: + features[..., feature_index, :] = ( + 4 * features[..., feature_index - 1, :] + ) + else: + features[..., feature_index, :] *= 4.0 + + if self.feature_spec.with_lapl: + for component in (4, 7, 9): + features[..., feature_index, :] += 2 * torch.sum( + c0 * ao[component][None, :, :], dim=-2 + ) + + if len(dm.shape) == 2: + return features.reshape((self.nfeats, -1)) + return features.reshape((*dm.shape[:-2], self.nfeats, -1)) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 0a59eb62..d08ba162 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -5,7 +5,6 @@ """ import logging -from abc import ABC, abstractmethod from collections.abc import Callable, Iterator from copy import copy from dataclasses import dataclass @@ -14,10 +13,11 @@ import numpy as np import torch from pyscf import dft, gto -from torch import Tensor, nn +from torch import Tensor from torch.autograd import Function from torch.autograd.function import FunctionCtx +from skala.pyscf import feature_math from skala.pyscf.backend import ( Array, Grid, @@ -204,18 +204,6 @@ def _prepare_spatially_sorted_grids( return sorted_grids, forward, inverse -def maybe_expand_and_divide( - feature: torch.Tensor, expand: bool, divisor: float -) -> torch.Tensor: - """ - Expand feature along spin channels and divide its value by divisor if expand is True. - """ - if expand: - return torch.stack([feature / divisor, feature / divisor], dim=0) - else: - return feature - - def make_chunks( atomic_grid_sizes: Tensor, max_grid_chunk_size: int ) -> list[tuple[slice, slice]]: @@ -317,7 +305,7 @@ def generate_features( dm, mol, grids, - MGGAFeatureFunction(feature_spec), + feature_math.MGGAFeatureFunction(feature_spec), block_size=evaluation_policy.ao_block_size, max_memory=max_memory, fix_block_size=evaluation_policy.ao_block_size is None, @@ -325,7 +313,7 @@ def generate_features( ) for feature in mgga_features: - mol_features[feature] = maybe_expand_and_divide( + mol_features[feature] = feature_math.maybe_expand_and_divide( mgga_features[feature], not with_spin, 2 ) @@ -393,10 +381,6 @@ def get_grid_features( return grid_features -def is_density_feature(feature: str) -> bool: - return feature in {"density", "grad", "kin"} - - def partial_feature_function_over_aos( feature_function: Callable[[torch.Tensor, torch.Tensor], torch.Tensor], ao: torch.Tensor, @@ -428,111 +412,6 @@ def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: return reduced_vjp -class FeatureFunction(nn.Module, ABC): - deriv: int - nfeats: int - - @abstractmethod - def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ... - - @abstractmethod - def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: ... - - -class MGGAFeatureFunction(FeatureFunction): - def __init__(self, feature_spec: FeatureSpec): - super().__init__() - - if not feature_spec.requires_mgga: - raise ValueError("At least one feature must be selected.") - self.feature_spec = feature_spec - self.deriv = feature_spec.ao_derivative_order - self.nfeats = feature_spec.mgga_feature_count - - def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: - """Convert the features to a dictionary with the keys being the feature names.""" - feature_index = 0 - feature_dict: dict[str, torch.Tensor] = {} - if self.feature_spec.with_density: - feature_dict["density"] = features[..., feature_index, :] - feature_index += 1 - if self.feature_spec.with_grad: - feature_dict["grad"] = features[..., feature_index : feature_index + 3, :] - feature_index += 3 - if self.feature_spec.with_kin: - feature_dict["kin"] = features[..., feature_index, :] - feature_index += 1 - if self.feature_spec.with_lapl: - feature_dict["lapl"] = features[..., feature_index, :] - feature_index += 1 - return feature_dict - - def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: - # Flatten all but the last two dimensions - # then restore the original shape at the end - dm_view = dm.view(-1, dm.shape[-2], dm.shape[-1]) - # Explicit symmetrization for autodiff - dm_view = 0.5 * (dm_view + dm_view.transpose(-1, -2)) - - features = torch.zeros( - (dm_view.shape[0], self.nfeats, ao.shape[-1]), - device=dm.device, - dtype=dm.dtype, - ) - - # Handle the density only case, where ao has one dim less - if self.deriv == 0: - c0 = dm_view @ ao - features[..., 0, :] = torch.sum(c0 * ao[None, :, :], dim=-2) - if len(dm.shape) == 2: - return features.reshape((self.nfeats, -1)) - else: - return features.reshape((*dm.shape[:-2], self.nfeats, -1)) - - c0 = dm_view @ ao[0] - - feat_idx = 0 - if self.feature_spec.with_density: - features[..., feat_idx, :] = torch.sum(c0 * ao[0][None, :, :], dim=-2) - feat_idx += 1 - - if self.feature_spec.with_grad: - for i in range(3): - features[..., feat_idx, :] = 2 * torch.sum( - c0 * ao[i + 1][None, :, :], dim=-2 - ) - feat_idx += 1 - - if self.feature_spec.with_kin or self.feature_spec.with_lapl: - for i in range(3): - ci = dm_view @ ao[i + 1] - features[..., feat_idx, :] += 0.5 * torch.sum( - ci * ao[i + 1][None, :, :], dim=-2 - ) - - if self.feature_spec.with_kin: - feat_idx += 1 - if self.feature_spec.with_lapl: - features[..., feat_idx, :] = 4 * features[..., feat_idx - 1, :] - else: - # Multiply times four for the laplacian - features[..., feat_idx, :] *= 4.0 - - if self.feature_spec.with_lapl: - # 0 is without derivative - # 1 2 3 are x y z derivatives - # 4 5 6 are xx xy xz derivatives - # 7 8 9 are yy yz zz derivatives - for i in (4, 7, 9): - features[..., feat_idx, :] += 2 * torch.sum( - c0 * ao[i][None, :, :], dim=-2 - ) - if len(dm.shape) == 2: - return features.reshape((self.nfeats, -1)) - else: - return features.reshape((*dm.shape[:-2], self.nfeats, -1)) - - @dataclass class _GlobalScreenedFeatures: dm: Tensor @@ -542,7 +421,7 @@ class _GlobalScreenedFeatures: atom_major_raw_features: Tensor forward_permutation: Tensor inverse_permutation: Tensor - feature_function: MGGAFeatureFunction + feature_function: feature_math.MGGAFeatureFunction block_size: int compile_feature_function: bool gpu: bool @@ -597,7 +476,7 @@ def build_model_chunk( for feature_name, feature in self.feature_function.to_dict( raw_features ).items(): - feature_chunk[feature_name] = maybe_expand_and_divide( + feature_chunk[feature_name] = feature_math.maybe_expand_and_divide( feature, not self.with_spin, 2 ) return feature_chunk @@ -624,7 +503,7 @@ def _global_screened_features( if grids.coords is None or grids.weights is None: raise ValueError("Grids must be built before generating screened features.") - feature_function = MGGAFeatureFunction(feature_spec) + feature_function = feature_math.MGGAFeatureFunction(feature_spec) grid_features = get_grid_features(mol, dm, grids, feature_spec) max_grid_chunk_size = estimate_max_grid_chunk_size( dm=dm, @@ -709,7 +588,7 @@ def add_to(self, matrix: Tensor, block_result: Tensor) -> None: def _evaluate_feature_block( - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, block: _AOBlock, active_dm: Tensor, compile_feature_function: bool, @@ -733,7 +612,7 @@ def __init__( dm: Tensor, mol: gto.Mole, grids: Grid, - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, blksize: int | None, gpu: bool, ) -> None: @@ -822,7 +701,7 @@ def setup_context( torch.Tensor, gto.Mole, Grid, - FeatureFunction, + feature_math.FeatureFunction, int | None, int, bool, @@ -848,7 +727,7 @@ def forward( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, blksize: int | None, compile_feature_function: bool, gpu: bool, @@ -969,7 +848,7 @@ def setup_context( torch.Tensor, gto.Mole, Grid, - FeatureFunction, + feature_math.FeatureFunction, int | None, bool, bool, @@ -994,7 +873,7 @@ def forward( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, blksize: int | None, compile_feature_function: bool, gpu: bool, @@ -1064,7 +943,7 @@ def non_chunk( dm: torch.Tensor, mol: gto.Mole, coords: Array, - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, compile_feature_function: bool = False, gpu: bool = False, ) -> torch.Tensor: @@ -1089,7 +968,7 @@ def auto_chunk( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: FeatureFunction, + feature_function: feature_math.FeatureFunction, block_size: int | None = None, max_memory: int = 2000, fix_block_size: bool = True, diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 657a31d6..539cffa4 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -10,10 +10,10 @@ from skala.pyscf import features as features_module from skala.pyscf import numint as numint_module from skala.pyscf.evaluation import FeatureSpec +from skala.pyscf.feature_math import MGGAFeatureFunction from skala.pyscf.features import ( ChunkEvalBackward, ChunkEvalForward, - MGGAFeatureFunction, _active_cpu_aos, _AOBlock, _evaluate_feature_block, diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 977b5ebd..c115a7ec 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -28,9 +28,9 @@ from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.pyscf.backend import dft_gpu # noqa: E402 from skala.pyscf.evaluation import FeatureSpec # noqa: E402 +from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 from skala.pyscf.features import ( # noqa: E402 ChunkEvalForward, - MGGAFeatureFunction, _prepare_spatially_sorted_grids, non_chunk, ) From d29938947f7fd11599825e50b44e07bffc849b71 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 13:55:52 +0200 Subject: [PATCH 16/39] refactor isolate ao autograd evaluation --- pyproject.toml | 7 +- src/skala/pyscf/ao_evaluation.py | 530 +++++++++++++++++++++++++ src/skala/pyscf/features.py | 557 +-------------------------- tests/test_ao_screening.py | 43 ++- tests/test_gpu4pyscf_ao_screening.py | 10 +- 5 files changed, 582 insertions(+), 565 deletions(-) create mode 100644 src/skala/pyscf/ao_evaluation.py diff --git a/pyproject.toml b/pyproject.toml index bda1dd73..6f9d4951 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -62,11 +62,10 @@ python_version = "3.12" exclude = ["third_party/"] disable_error_code = ["no-any-return"] -# torch.autograd.Function uses dynamic attributes on FunctionCtx (ctx.save_for_backward pattern) -# and Function.apply() is untyped. These are fundamental to PyTorch's autograd API. +# torch.autograd.Function.apply() is untyped in PyTorch. [[tool.mypy.overrides]] -module = "skala.pyscf.features" -disable_error_code = ["no-any-return", "attr-defined", "no-untyped-call"] +module = ["skala.pyscf.ao_evaluation", "skala.pyscf.features"] +disable_error_code = ["no-untyped-call"] [tool.ruff] target-version = "py311" diff --git a/src/skala/pyscf/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py new file mode 100644 index 00000000..18be440e --- /dev/null +++ b/src/skala/pyscf/ao_evaluation.py @@ -0,0 +1,530 @@ +# SPDX-License-Identifier: MIT + +"""Blockwise atomic-orbital feature evaluation and custom autograd.""" + +from collections.abc import Callable, Iterator +from dataclasses import dataclass +from typing import Protocol, cast + +import numpy as np +import torch +from pyscf import dft, gto +from torch import Tensor +from torch.autograd import Function +from torch.autograd.function import FunctionCtx +from typing_extensions import Unpack + +from skala.pyscf import feature_math +from skala.pyscf.backend import ( + Array, + Grid, + check_gpu_imports_were_successful, + dft_gpu, + from_numpy_or_cupy, +) + + +class _ChunkEvalForwardContext(Protocol): + dm: Tensor + mol: gto.Mole + grids: Grid + feature_function: feature_math.FeatureFunction + blksize: int | None + compile_feature_function: bool + gpu: bool + vectors_jvp: tuple[Tensor, ...] + + +class _ChunkEvalBackwardContext(Protocol): + dm: Tensor + mol: gto.Mole + grids: Grid + feature_function: feature_math.FeatureFunction + blksize: int | None + compile_feature_function: bool + gpu: bool + + +def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: + """Expand a PySCF shell-screening mask into active AO indices.""" + active_shells = np.any(screen_index, axis=0) + ao_loc = mol.ao_loc_nr() + return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) + + +def partial_feature_function_over_aos( + feature_function: Callable[[torch.Tensor, torch.Tensor], torch.Tensor], + ao: torch.Tensor, +) -> Callable[[torch.Tensor], torch.Tensor]: + """Bind an AO block to a feature function for block-local evaluation.""" + + def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: + return feature_function(dm, ao) + + return partial_feature_function + + +def partial_vjp_function_over_tangents( + func: Callable[[torch.Tensor], torch.Tensor], + tangents: torch.Tensor, +) -> Callable[[torch.Tensor], torch.Tensor]: + """Bind feature cotangents to a function for block-local VJP evaluation.""" + + def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: + return torch.func.vjp(func, primals)[1](tangents)[0] + + return reduced_vjp + + +@dataclass(frozen=True) +class _AOBlock: + ao: Tensor + active_aos: Tensor | None + grid_slice: slice + + def select_aos(self, matrix: Tensor) -> Tensor: + if self.active_aos is None: + return matrix + return matrix[..., self.active_aos[:, None], self.active_aos[None, :]] + + def add_to(self, matrix: Tensor, block_result: Tensor) -> None: + if self.active_aos is None: + matrix += block_result + else: + matrix[..., self.active_aos[:, None], self.active_aos[None, :]] += ( + block_result + ) + + +def _evaluate_feature_block( + feature_function: feature_math.FeatureFunction, + block: _AOBlock, + active_dm: Tensor, + compile_feature_function: bool, + feature_cotangent: Tensor | None = None, +) -> Tensor: + """Evaluate one active-AO feature block or its feature-space VJP.""" + partial_func = partial_feature_function_over_aos(feature_function, block.ao) + if feature_cotangent is not None: + partial_func = partial_vjp_function_over_tangents( + partial_func, feature_cotangent[..., block.grid_slice] + ) + + if compile_feature_function: + return torch.compile(partial_func)(active_dm) + return partial_func(active_dm) + + +class _AOBlockLoop: + def __init__( + self, + dm: Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + blksize: int | None, + gpu: bool, + ) -> None: + self.dm = dm + self.mol = mol + self.grids = grids + self.feature_function = feature_function + self.blksize = blksize + self.gpu = gpu + self.sort_idx: Tensor | None + self.unsort_idx: Tensor | None + + if gpu: + check_gpu_imports_were_successful() + self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) + self.numint.grid_blksize = blksize + self.sort_idx = torch.as_tensor( + self.numint.gdftopt._ao_idx, device=dm.device + ) + self.unsort_idx = torch.argsort(self.sort_idx) + else: + self.numint = dft.numint.NumInt() + self.sort_idx = None + self.unsort_idx = None + + def order_aos(self, matrix: Tensor) -> Tensor: + if self.sort_idx is None: + return matrix + return matrix[..., self.sort_idx, :][..., self.sort_idx] + + def restore_ao_order(self, matrix: Tensor) -> Tensor: + if self.unsort_idx is None: + return matrix + return matrix[..., self.unsort_idx, :][..., self.unsort_idx] + + def __iter__(self) -> Iterator[_AOBlock]: + block_loop_options: dict[str, bool] = {} + if self.gpu: + # GPU4PySCF otherwise omits zero-AO blocks, shifting all later grid slices. + block_loop_options["strict_grid_order"] = True + + end = 0 + for ao_block, mask, weights, _ in self.numint.block_loop( + mol=self.mol, + grids=self.grids, + nao=self.mol.nao, + deriv=self.feature_function.deriv, + blksize=self.blksize, + non0tab=(None if self.gpu else getattr(self.grids, "non0tab", None)), + **block_loop_options, + ): + start, end = end, end + weights.size + ao = from_numpy_or_cupy( + ao_block, + device=self.dm.device, + dtype=self.dm.dtype, + transpose=not self.gpu, + ) + active_aos: Tensor | None + if mask is None: + active_aos = None + elif self.gpu: + active_aos = from_numpy_or_cupy( + mask, device=self.dm.device, dtype=torch.long + ) + else: + num_screen_rows = ( + weights.size + dft.gen_grid.BLKSIZE - 1 + ) // dft.gen_grid.BLKSIZE + active_aos = torch.as_tensor( + _active_cpu_aos(self.mol, mask[:num_screen_rows]), + device=self.dm.device, + dtype=torch.long, + ) + ao = ao[..., active_aos, :] + if active_aos is not None and active_aos.numel() == 0: + continue + yield _AOBlock(ao, active_aos, slice(start, end)) + + +class ChunkEvalForward(Function): + @staticmethod + def setup_context( + ctx: FunctionCtx, + inputs: tuple[ + Tensor, + gto.Mole, + Grid, + feature_math.FeatureFunction, + int | None, + bool, + bool, + # The starred spelling requires Python 3.11. + Unpack[tuple[Tensor, ...]], # noqa: UP044 + ], + output: torch.Tensor, + ) -> None: + if len(inputs) < 7: + raise ValueError("ChunkEvalForward requires seven fixed inputs.") + context = cast(_ChunkEvalForwardContext, ctx) + ( + context.dm, + context.mol, + context.grids, + context.feature_function, + context.blksize, + context.compile_feature_function, + context.gpu, + *vectors_jvp, + ) = inputs + context.vectors_jvp = tuple(vectors_jvp) + ctx.save_for_backward(context.dm) + + @staticmethod + def forward( + dm: torch.Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + blksize: int | None, + compile_feature_function: bool, + gpu: bool, + *vectors_jvp: torch.Tensor, + ) -> torch.Tensor: + ngrids = grids.weights.size + block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + + features = torch.zeros( + *dm.shape[:-2], + feature_function.nfeats, + ngrids, + device=dm.device, + dtype=dm.dtype, + ) + # Raw AO features are linear in dm, so derivatives above first order vanish. + if len(vectors_jvp) > 1: + return features + + evaluation_dm = vectors_jvp[0] if vectors_jvp else dm + evaluation_dm_ordered = block_loop.order_aos(evaluation_dm) + for block in block_loop: + active_dm = block.select_aos(evaluation_dm_ordered) + temp_feature = _evaluate_feature_block( + feature_function, + block, + active_dm, + compile_feature_function, + ) + features[..., block.grid_slice] = temp_feature + return features + + @staticmethod + def jvp( + ctx: _ChunkEvalForwardContext, *grad_inputs: torch.Tensor | None + ) -> torch.Tensor: + if len(ctx.vectors_jvp) > 1: + return torch.zeros( + *ctx.dm.shape[:-2], + ctx.feature_function.nfeats, + ctx.grids.weights.size, + device=ctx.dm.device, + dtype=ctx.dm.dtype, + ) + vector_tangent = grad_inputs[7] if ctx.vectors_jvp else grad_inputs[0] + if vector_tangent is None: + return torch.zeros( + *ctx.dm.shape[:-2], + ctx.feature_function.nfeats, + ctx.grids.weights.size, + device=ctx.dm.device, + dtype=ctx.dm.dtype, + ) + return ChunkEvalForward.apply( + ctx.dm, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ctx.gpu, + vector_tangent, + ) + + @staticmethod + def backward( + ctx: _ChunkEvalForwardContext, *grad_outputs: torch.Tensor + ) -> tuple[torch.Tensor | None, ...]: + feature_cotangent = grad_outputs[0] + if ctx.vectors_jvp: + dm_grad = ctx.dm * 0 + else: + dm_grad = ChunkEvalBackward.apply( + ctx.dm, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ctx.gpu, + feature_cotangent, + ) + grads: list[Tensor | None] = [dm_grad] + grads += [None] * 6 + + for vector in ctx.vectors_jvp: + if len(ctx.vectors_jvp) == 1: + vector_grad = ChunkEvalBackward.apply( + ctx.dm, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ctx.gpu, + feature_cotangent, + ) + else: + vector_grad = vector * 0 + grads.append(vector_grad) + + return tuple(grads) + + +class ChunkEvalBackward(Function): + @staticmethod + def setup_context( + ctx: FunctionCtx, + inputs: tuple[ + torch.Tensor, + gto.Mole, + Grid, + feature_math.FeatureFunction, + int | None, + bool, + bool, + torch.Tensor, + ], + output: torch.Tensor, + ) -> None: + context = cast(_ChunkEvalBackwardContext, ctx) + ( + context.dm, + context.mol, + context.grids, + context.feature_function, + context.blksize, + context.compile_feature_function, + context.gpu, + _feature_cotangent, + ) = inputs + ctx.save_for_backward(context.dm) + + @staticmethod + def forward( + dm: torch.Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + blksize: int | None, + compile_feature_function: bool, + gpu: bool, + feature_cotangent: torch.Tensor, + ) -> torch.Tensor: + block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + dm_ordered = block_loop.order_aos(dm) + + out = torch.zeros_like(dm) + for block in block_loop: + active_dm = block.select_aos(dm_ordered) + block_result = _evaluate_feature_block( + feature_function, + block, + active_dm, + compile_feature_function, + feature_cotangent, + ) + block.add_to(out, block_result) + return block_loop.restore_ao_order(out) + + @staticmethod + def jvp( + ctx: _ChunkEvalBackwardContext, *grad_inputs: torch.Tensor | None + ) -> torch.Tensor: + feature_cotangent_tangent = grad_inputs[7] + if feature_cotangent_tangent is None: + return torch.zeros_like(ctx.dm) + return ChunkEvalBackward.apply( + ctx.dm, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ctx.gpu, + feature_cotangent_tangent, + ) + + @staticmethod + def backward( + ctx: _ChunkEvalBackwardContext, *grad_outputs: torch.Tensor + ) -> tuple[torch.Tensor | None, ...]: + grads: list[Tensor | None] = [ctx.dm * 0] + grads += [None] * 6 + grads.append( + ChunkEvalForward.apply( + ctx.dm, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ctx.gpu, + grad_outputs[0], + ) + ) + return tuple(grads) + + +def non_chunk( + dm: torch.Tensor, + mol: gto.Mole, + coords: Array, + feature_function: feature_math.FeatureFunction, + compile_feature_function: bool = False, + gpu: bool = False, +) -> torch.Tensor: + """Evaluate raw features over the full grid without block chunking.""" + if gpu: + check_gpu_imports_were_successful() + ni = dft_gpu.numint.NumInt().build(mol, coords) + else: + ni = dft.numint.NumInt() + ao = from_numpy_or_cupy( + ni.eval_ao(mol, coords, deriv=feature_function.deriv, non0tab=None), + device=dm.device, + dtype=dm.dtype, + transpose=True, + ) + if compile_feature_function: + return torch.compile(feature_function.forward)(dm, ao) + return feature_function.forward(dm, ao) + + +def _resolve_ao_block_size( + mol: gto.Mole, + feature_function: feature_math.FeatureFunction, + block_size: int | None, + max_memory: int, + gpu: bool, +) -> int | None: + """Resolve an aligned CPU block size or delegate GPU sizing to its backend.""" + if gpu: + if block_size is not None: + raise ValueError("Setting custom block size is not supported on GPU.") + return None + + if block_size is None: + nao = mol.nao_nr() + comp = ( + (feature_function.deriv + 1) + * (feature_function.deriv + 2) + * (feature_function.deriv + 3) + // 6 + ) + backend_block_size = dft.gen_grid.BLKSIZE + block_size = int(max_memory * 1e6 / ((comp + 1) * nao * 8 * backend_block_size)) + block_size = max(4, min(block_size, 1200)) * backend_block_size + + return block_size - block_size % dft.gen_grid.BLKSIZE + + +def auto_chunk( + dm: torch.Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + block_size: int | None = None, + max_memory: int = 2000, + gpu: bool = False, +) -> dict[str, torch.Tensor]: + """Evaluate raw features with a memory-derived or explicit AO block size.""" + if gpu: + check_gpu_imports_were_successful() + if dm.device.type != "cuda": + raise ValueError("Density matrix must be on the GPU when gpu=True.") + + blksize = _resolve_ao_block_size(mol, feature_function, block_size, max_memory, gpu) + + if blksize is not None and blksize >= grids.weights.shape[0]: + features = non_chunk( + dm.double(), + mol, + grids.coords, + feature_function, + ) + else: + features = ChunkEvalForward.apply( + dm.double(), + mol, + grids, + feature_function, + blksize, + False, + gpu, + ) + return feature_function.to_dict(features) diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index d08ba162..3b5e74fe 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -5,7 +5,6 @@ """ import logging -from collections.abc import Callable, Iterator from copy import copy from dataclasses import dataclass from typing import TypeAlias @@ -14,12 +13,9 @@ import torch from pyscf import dft, gto from torch import Tensor -from torch.autograd import Function -from torch.autograd.function import FunctionCtx -from skala.pyscf import feature_math +from skala.pyscf import ao_evaluation, feature_math from skala.pyscf.backend import ( - Array, Grid, check_gpu_imports_were_successful, dft_gpu, @@ -70,21 +66,6 @@ def matches( ) -def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: - """Expand a PySCF shell-screening mask into active AO indices. - - Args: - mol: Molecule defining the shell-to-AO ranges. - screen_index: Screening rows whose columns correspond to molecular shells. - - Returns: - Sorted indices of AOs belonging to a shell active in any screening row. - """ - active_shells = np.any(screen_index, axis=0) - ao_loc = mol.ao_loc_nr() - return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) - - def _spatial_grid_permutations( coords: _Float64Coordinates, block_size: int ) -> tuple[_Int64Permutation, _Int64Permutation]: @@ -301,14 +282,13 @@ def generate_features( mol_features = get_grid_features(mol, dm, grids, feature_spec) if feature_spec.requires_mgga: - mgga_features = auto_chunk( + mgga_features = ao_evaluation.auto_chunk( dm, mol, grids, feature_math.MGGAFeatureFunction(feature_spec), block_size=evaluation_policy.ao_block_size, max_memory=max_memory, - fix_block_size=evaluation_policy.ao_block_size is None, gpu=gpu, ) @@ -381,37 +361,6 @@ def get_grid_features( return grid_features -def partial_feature_function_over_aos( - feature_function: Callable[[torch.Tensor, torch.Tensor], torch.Tensor], - ao: torch.Tensor, -) -> Callable[[torch.Tensor], torch.Tensor]: - """Returns a function that computes the feature function with the given ao, - but not the dm already passed to the function. - - Purpose is to allow evaluating a block-local VJP. - """ - - def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: - return feature_function(dm, ao) - - return partial_feature_function - - -def partial_vjp_function_over_tangents( - func: Callable[[torch.Tensor], torch.Tensor], - tangents: torch.Tensor, -) -> Callable[[torch.Tensor], torch.Tensor]: - """Returns a function that computes the vjp of the given function with tangents, - but not primals already passed to the function. - - Purpose is to evaluate the feature-space adjoint for one AO block.""" - - def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: - return torch.func.vjp(func, primals)[1](tangents)[0] - - return reduced_vjp - - @dataclass class _GlobalScreenedFeatures: dm: Tensor @@ -432,7 +381,7 @@ class _GlobalScreenedFeatures: def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: """Apply the global raw-feature Jacobian and restore atom-major order.""" - sorted_tangent = ChunkEvalForward.apply( + sorted_tangent = ao_evaluation.ChunkEvalForward.apply( self.dm, self.mol, self.sorted_grids, @@ -535,7 +484,7 @@ def _global_screened_features( sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( mol, grids, block_size, gpu ) - sorted_raw_features = ChunkEvalForward.apply( + sorted_raw_features = ao_evaluation.ChunkEvalForward.apply( dm.double(), mol, sorted_grids, @@ -565,501 +514,3 @@ def _global_screened_features( chunks=chunks, with_spin=dm.ndim == 3, ) - - -@dataclass(frozen=True) -class _AOBlock: - ao: Tensor - active_aos: Tensor | None - grid_slice: slice - - def select_aos(self, matrix: Tensor) -> Tensor: - if self.active_aos is None: - return matrix - return matrix[..., self.active_aos[:, None], self.active_aos[None, :]] - - def add_to(self, matrix: Tensor, block_result: Tensor) -> None: - if self.active_aos is None: - matrix += block_result - else: - matrix[..., self.active_aos[:, None], self.active_aos[None, :]] += ( - block_result - ) - - -def _evaluate_feature_block( - feature_function: feature_math.FeatureFunction, - block: _AOBlock, - active_dm: Tensor, - compile_feature_function: bool, - feature_cotangent: Tensor | None = None, -) -> Tensor: - """Evaluate one active-AO feature block or its feature-space VJP.""" - partial_func = partial_feature_function_over_aos(feature_function, block.ao) - if feature_cotangent is not None: - partial_func = partial_vjp_function_over_tangents( - partial_func, feature_cotangent[..., block.grid_slice] - ) - - if compile_feature_function: - return torch.compile(partial_func)(active_dm) - return partial_func(active_dm) - - -class _AOBlockLoop: - def __init__( - self, - dm: Tensor, - mol: gto.Mole, - grids: Grid, - feature_function: feature_math.FeatureFunction, - blksize: int | None, - gpu: bool, - ) -> None: - self.dm = dm - self.mol = mol - self.grids = grids - self.feature_function = feature_function - self.blksize = blksize - self.gpu = gpu - self.sort_idx: Tensor | None - self.unsort_idx: Tensor | None - - if gpu: - check_gpu_imports_were_successful() - self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) - self.numint.grid_blksize = blksize - self.sort_idx = torch.as_tensor( - self.numint.gdftopt._ao_idx, device=dm.device - ) - self.unsort_idx = torch.argsort(self.sort_idx) - else: - self.numint = dft.numint.NumInt() - self.sort_idx = None - self.unsort_idx = None - - def order_aos(self, matrix: Tensor) -> Tensor: - if self.sort_idx is None: - return matrix - return matrix[..., self.sort_idx, :][..., self.sort_idx] - - def restore_ao_order(self, matrix: Tensor) -> Tensor: - if self.unsort_idx is None: - return matrix - return matrix[..., self.unsort_idx, :][..., self.unsort_idx] - - def __iter__(self) -> Iterator[_AOBlock]: - block_loop_options: dict[str, bool] = {} - if self.gpu: - # GPU4PySCF otherwise omits zero-AO blocks, shifting all later grid slices. - block_loop_options["strict_grid_order"] = True - - end = 0 - for ao_block, mask, weights, _ in self.numint.block_loop( - mol=self.mol, - grids=self.grids, - nao=self.mol.nao, - deriv=self.feature_function.deriv, - blksize=self.blksize, - non0tab=(None if self.gpu else getattr(self.grids, "non0tab", None)), - **block_loop_options, - ): - start, end = end, end + weights.size - ao = from_numpy_or_cupy( - ao_block, - device=self.dm.device, - dtype=self.dm.dtype, - transpose=not self.gpu, - ) - active_aos: Tensor | None - if mask is None: - active_aos = None - elif self.gpu: - active_aos = from_numpy_or_cupy( - mask, device=self.dm.device, dtype=torch.long - ) - else: - num_screen_rows = ( - weights.size + dft.gen_grid.BLKSIZE - 1 - ) // dft.gen_grid.BLKSIZE - active_aos = torch.as_tensor( - _active_cpu_aos(self.mol, mask[:num_screen_rows]), - device=self.dm.device, - dtype=torch.long, - ) - ao = ao[..., active_aos, :] - if active_aos is not None and active_aos.numel() == 0: - continue - yield _AOBlock(ao, active_aos, slice(start, end)) - - -class ChunkEvalForward(Function): - @staticmethod - def setup_context( - ctx: FunctionCtx, - inputs: tuple[ - torch.Tensor, - gto.Mole, - Grid, - feature_math.FeatureFunction, - int | None, - int, - bool, - bool, - torch.Tensor, - ], - output: torch.Tensor, - ) -> None: - ( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - *ctx.vectors_jvp, - ) = inputs - ctx.save_for_backward(ctx.dm) - - @staticmethod - def forward( - dm: torch.Tensor, - mol: gto.Mole, - grids: Grid, - feature_function: feature_math.FeatureFunction, - blksize: int | None, - compile_feature_function: bool, - gpu: bool, - *vectors_jvp: torch.Tensor, - ) -> torch.Tensor: - ngrids = grids.weights.size - block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) - - features = torch.zeros( - *dm.shape[:-2], - feature_function.nfeats, - ngrids, - device=dm.device, - dtype=dm.dtype, - ) - # Raw AO features are linear in dm, so derivatives above first order vanish. - if len(vectors_jvp) > 1: - return features - - # Since the raw feature map is linear, its JVP is direct evaluation on - # the tangent density matrix. - evaluation_dm = vectors_jvp[0] if vectors_jvp else dm - evaluation_dm_ordered = block_loop.order_aos(evaluation_dm) - for block in block_loop: - active_dm = block.select_aos(evaluation_dm_ordered) - temp_feature = _evaluate_feature_block( - feature_function, - block, - active_dm, - compile_feature_function, - ) - - features[..., block.grid_slice] = temp_feature - return features - - @staticmethod - def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor | None) -> torch.Tensor: - if len(ctx.vectors_jvp) > 1: - return torch.zeros( - *ctx.dm.shape[:-2], - ctx.feature_function.nfeats, - ctx.grids.weights.size, - device=ctx.dm.device, - dtype=ctx.dm.dtype, - ) - vector_tangent = grad_inputs[7] if ctx.vectors_jvp else grad_inputs[0] - if vector_tangent is None: - return torch.zeros( - *ctx.dm.shape[:-2], - ctx.feature_function.nfeats, - ctx.grids.weights.size, - device=ctx.dm.device, - dtype=ctx.dm.dtype, - ) - return ChunkEvalForward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - vector_tangent, - ) - - @staticmethod - def backward( - ctx: FunctionCtx, *grad_outputs: torch.Tensor - ) -> tuple[torch.Tensor | None, ...]: - feature_cotangent = grad_outputs[0] - if ctx.vectors_jvp: - dm_grad = ctx.dm * 0 - else: - dm_grad = ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent, - ) - grads = [dm_grad] - - # We need to provide None for the gradients of the non-differentiable inputs - # these are mol (1), grids (2), feature_function (3), blksize (4), - # compile_feature_function (5), gpu (6) - num_non_differentiable_inputs = 6 - - grads += [None] * num_non_differentiable_inputs - - # A first JVP is linear in its tangent; higher JVPs are identically zero. - for vector in ctx.vectors_jvp: - if len(ctx.vectors_jvp) == 1: - vector_grad = ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent, - ) - else: - vector_grad = vector * 0 - grads.append(vector_grad) - - return tuple(grads) - - -class ChunkEvalBackward(Function): - @staticmethod - def setup_context( - ctx: FunctionCtx, - inputs: tuple[ - torch.Tensor, - gto.Mole, - Grid, - feature_math.FeatureFunction, - int | None, - bool, - bool, - torch.Tensor, - ], - output: torch.Tensor, - ) -> None: - ( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - ctx.feature_cotangent, - ) = inputs - ctx.save_for_backward(ctx.dm) - - @staticmethod - def forward( - dm: torch.Tensor, - mol: gto.Mole, - grids: Grid, - feature_function: feature_math.FeatureFunction, - blksize: int | None, - compile_feature_function: bool, - gpu: bool, - feature_cotangent: torch.Tensor, - ) -> torch.Tensor: - block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) - dm_ordered = block_loop.order_aos(dm) - - out = torch.zeros_like(dm) - for block in block_loop: - active_dm = block.select_aos(dm_ordered) - block_result = _evaluate_feature_block( - feature_function, - block, - active_dm, - compile_feature_function, - feature_cotangent, - ) - block.add_to(out, block_result) - return block_loop.restore_ao_order(out) - - @staticmethod - def jvp(ctx: FunctionCtx, *grad_inputs: torch.Tensor | None) -> torch.Tensor: - feature_cotangent_tangent = grad_inputs[7] - if feature_cotangent_tangent is None: - return torch.zeros_like(ctx.dm) - return ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent_tangent, - ) - - @staticmethod - def backward( - ctx: FunctionCtx, *grad_outputs: torch.Tensor - ) -> tuple[torch.Tensor | None, ...]: - # The raw feature Jacobian is constant in dm. The only nonzero gradient - # propagates through the feature-space cotangent. - grads = [ctx.dm * 0] - # We need to provide None for the gradients of the non-differentiable inputs - # these are mol (1), grids (2), feature_function (3), blksize (4), - # compile_feature_function (5), gpu (6) - num_non_differentiable_inputs = 6 - - grads += [None] * num_non_differentiable_inputs - grads.append( - ChunkEvalForward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - grad_outputs[0], - ) - ) - return tuple(grads) - - -def non_chunk( - dm: torch.Tensor, - mol: gto.Mole, - coords: Array, - feature_function: feature_math.FeatureFunction, - compile_feature_function: bool = False, - gpu: bool = False, -) -> torch.Tensor: - if gpu: - check_gpu_imports_were_successful() - ni = dft_gpu.numint.NumInt().build(mol, coords) - else: - ni = dft.numint.NumInt() - ao = from_numpy_or_cupy( - ni.eval_ao(mol, coords, deriv=feature_function.deriv, non0tab=None), - device=dm.device, - dtype=dm.dtype, - transpose=True, - ) - if compile_feature_function: - return torch.compile(feature_function.forward)(dm, ao) - else: - return feature_function.forward(dm, ao) - - -def auto_chunk( - dm: torch.Tensor, - mol: gto.Mole, - grids: Grid, - feature_function: feature_math.FeatureFunction, - block_size: int | None = None, - max_memory: int = 2000, - fix_block_size: bool = True, - compile_feature_function: bool = False, - gpu: bool = False, -) -> dict[str, torch.Tensor]: - """ - Automatically splits feature evaluation into smaller chunks if needed. - - This function determines the appropriate chunk size for evaluating a feature - function on molecular grids, based on available memory and number of basis - functions. If the computed chunk size is larger than the size of the grid, or - if a fixed block size was provided, it uses a non-chunked approach. - - Parameters - ---------- - dm: torch.Tensor - Density matrix or set of density matrices used for - evaluating the feature function. - mol: gto.Mole - PySCF molecule object representing the system of interest. - grids: Grid - Grids object defining the points in space on which - the feature function is evaluated. - feature_function: FeatureFunction - The object representing the feature function to evaluate. The number of derivatives (deriv) determines - how many components to compute. - gpu: bool, optional - Whether to use GPU for computation. Defaults to False. - block_size: int | None, optional - Manually specified block size for chunking. (CPU only) - Defaults to None. - max_memory: int, optional - Maximum memory in MB to use for chunking (CPU only) - fix_block_size: bool, optional - Whether to fix the block size or compute it - automatically based on system resources. Defaults to True. (CPU only) - compile_feature_function: bool, optional - If True, compiles the feature function for efficiency. Defaults to False. - - Returns - ------- - dict[str, torch.Tensor]: - The evaluated feature function on the specified grids, either - computed in smaller chunks or in a single pass, depending on the block size. - """ - - if gpu: - check_gpu_imports_were_successful() - if dm.device.type != "cuda": - raise ValueError("Density matrix must be on the GPU when gpu=True.") - - blksize: int | None - - if gpu and block_size is not None: - raise ValueError("Setting custom block size is not supported on GPU.") - - if block_size is None and fix_block_size and not gpu: - nao = mol.nao_nr() - comp = ( - (feature_function.deriv + 1) - * (feature_function.deriv + 2) - * (feature_function.deriv + 3) - // 6 - ) - BLKSIZE = dft.gen_grid.BLKSIZE - blksize = int(max_memory * 1e6 / ((comp + 1) * nao * 8 * BLKSIZE)) - blksize = max(4, min(blksize, 1200)) * BLKSIZE - else: - blksize = block_size - - if blksize is not None and not gpu: - blksize = blksize - blksize % dft.gen_grid.BLKSIZE - - if blksize is not None and blksize >= grids.weights.shape[0]: - features = non_chunk( - dm.double(), - mol, - grids.coords, - feature_function, - compile_feature_function=compile_feature_function, - gpu=gpu, - ) - else: - features = ChunkEvalForward.apply( - dm.double(), - mol, - grids, - feature_function, - blksize, - compile_feature_function, - gpu, - ) - return feature_function.to_dict(features) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 539cffa4..c63935aa 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -9,14 +9,17 @@ from skala.functional.base import ExcFunctionalBase from skala.pyscf import features as features_module from skala.pyscf import numint as numint_module -from skala.pyscf.evaluation import FeatureSpec -from skala.pyscf.feature_math import MGGAFeatureFunction -from skala.pyscf.features import ( +from skala.pyscf.ao_evaluation import ( ChunkEvalBackward, ChunkEvalForward, _active_cpu_aos, _AOBlock, _evaluate_feature_block, + _resolve_ao_block_size, +) +from skala.pyscf.evaluation import FeatureSpec +from skala.pyscf.feature_math import MGGAFeatureFunction +from skala.pyscf.features import ( _prepare_spatially_sorted_grids, _spatial_grid_permutations, ) @@ -129,6 +132,40 @@ def test_active_cpu_aos(carbon: gto.Mole) -> None: assert empty.size == 0 +def test_resolve_ao_block_size_modes(carbon: gto.Mole) -> None: + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + backend_block_size = dft.gen_grid.BLKSIZE + + # CPU sizes are aligned locally; GPU sizing is delegated unless explicitly invalid. + automatic = _resolve_ao_block_size( + carbon, feature_function, block_size=None, max_memory=0, gpu=False + ) + explicit = _resolve_ao_block_size( + carbon, + feature_function, + block_size=backend_block_size + 1, + max_memory=0, + gpu=False, + ) + + assert automatic == 4 * backend_block_size + assert explicit == backend_block_size + assert ( + _resolve_ao_block_size( + carbon, feature_function, block_size=None, max_memory=0, gpu=True + ) + is None + ) + with pytest.raises(ValueError, match="custom block size"): + _resolve_ao_block_size( + carbon, + feature_function, + block_size=backend_block_size, + max_memory=0, + gpu=True, + ) + + @pytest.mark.parametrize(("ngrids", "block_size"), [(0, 4), (3, 4), (8, 4), (10, 4)]) def test_spatial_grid_permutations_restore_original_order( ngrids: int, block_size: int diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index c115a7ec..260c705f 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -26,14 +26,14 @@ from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 -from skala.pyscf.backend import dft_gpu # noqa: E402 -from skala.pyscf.evaluation import FeatureSpec # noqa: E402 -from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 -from skala.pyscf.features import ( # noqa: E402 +from skala.pyscf.ao_evaluation import ( # noqa: E402 ChunkEvalForward, - _prepare_spatially_sorted_grids, non_chunk, ) +from skala.pyscf.backend import dft_gpu # noqa: E402 +from skala.pyscf.evaluation import FeatureSpec # noqa: E402 +from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 +from skala.pyscf.features import _prepare_spatially_sorted_grids # noqa: E402 from skala.pyscf.numint import SkalaNumInt # noqa: E402 CARBON_CHAIN = """ From 683e3507ec9f03eff64e3d4cbd87f1b02440ac08 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 14:18:15 +0200 Subject: [PATCH 17/39] refactor isolate screened feature evaluation --- pyproject.toml | 2 +- src/skala/pyscf/features.py | 377 +-------------------------- src/skala/pyscf/memory_estimators.py | 4 +- src/skala/pyscf/model_chunking.py | 196 ++++++++++++++ src/skala/pyscf/numint.py | 57 ++-- src/skala/pyscf/screening.py | 293 +++++++++++++++++++++ tests/test_ao_screening.py | 92 ++++--- tests/test_gpu4pyscf_ao_screening.py | 2 +- tests/test_memory_estimators.py | 10 +- 9 files changed, 584 insertions(+), 449 deletions(-) create mode 100644 src/skala/pyscf/model_chunking.py create mode 100644 src/skala/pyscf/screening.py diff --git a/pyproject.toml b/pyproject.toml index 6f9d4951..ffa8c876 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -64,7 +64,7 @@ disable_error_code = ["no-any-return"] # torch.autograd.Function.apply() is untyped in PyTorch. [[tool.mypy.overrides]] -module = ["skala.pyscf.ao_evaluation", "skala.pyscf.features"] +module = ["skala.pyscf.ao_evaluation", "skala.pyscf.screening"] disable_error_code = ["no-untyped-call"] [tool.ruff] diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index 3b5e74fe..e27fe2c4 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -4,235 +4,17 @@ Methods for generating and manipulating density features. """ -import logging -from copy import copy -from dataclasses import dataclass -from typing import TypeAlias - import numpy as np import torch -from pyscf import dft, gto +from pyscf import gto from torch import Tensor from skala.pyscf import ao_evaluation, feature_math -from skala.pyscf.backend import ( - Grid, - check_gpu_imports_were_successful, - dft_gpu, - from_numpy_or_cupy, -) +from skala.pyscf.backend import Grid, from_numpy_or_cupy from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec -from skala.pyscf.memory_estimators import ( - estimate_global_raw_feature_buffer_memory, - estimate_max_grid_chunk_size, -) - -LOG = logging.getLogger(__name__) DEFAULT_FEATURES = ["density", "kin", "grad", "grid_coords", "grid_weights"] DEFAULT_FEATURES_SET = set(DEFAULT_FEATURES) -CPU_AO_SCREENING_BLOCK_SIZE = 9 * dft.gen_grid.BLKSIZE - -_Float64Coordinates: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.float64]] -_Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] -_SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" - - -@dataclass(frozen=True) -class _SpatialGridCache: - mol: gto.Mole - source_coords: object - source_weights: object - block_size: int - gpu: bool - sorted_grids: Grid - forward: _Int64Permutation - inverse: _Int64Permutation - - def matches( - self, - mol: gto.Mole, - grids: Grid, - block_size: int, - gpu: bool, - ) -> bool: - """Return whether this entry belongs to the current built grid.""" - return ( - self.mol is mol - and self.source_coords is grids.coords - and self.source_weights is grids.weights - and self.block_size == block_size - and self.gpu is gpu - ) - - -def _spatial_grid_permutations( - coords: _Float64Coordinates, block_size: int -) -> tuple[_Int64Permutation, _Int64Permutation]: - """Order a molecular grid into exact-size spatial blocks. - - Recursively partitions points along the longest Cartesian extent. Every left - subtree contains a whole number of evaluator blocks, so all output blocks have - ``block_size`` points except for a possible final remainder. - - Args: - coords: Molecular grid coordinates with shape ``(ngrids, 3)``. - block_size: Fixed number of points consumed by each backend block. - - Returns: - The forward permutation from atom-major to spatial order and its inverse. - - Raises: - ValueError: If the coordinates or block size are invalid. - """ - if coords.ndim != 2 or coords.shape[1] != 3: - raise ValueError("coords must have shape (ngrids, 3)") - if block_size <= 0: - raise ValueError("block_size must be positive") - - def partition(indices: _Int64Permutation) -> list[_Int64Permutation]: - if indices.size <= block_size: - return [indices] - - block_count = (indices.size + block_size - 1) // block_size - left_size = (block_count // 2) * block_size - extents = np.ptp(coords[indices], axis=0) - split_axis = int(np.argmax(extents)) - positions = np.lexsort((indices, coords[indices, split_axis])) - ordered_indices = indices[positions] - return partition(ordered_indices[:left_size]) + partition( - ordered_indices[left_size:] - ) - - ngrids = coords.shape[0] - if ngrids == 0: - empty = np.empty(0, dtype=np.int64) - return empty, empty.copy() - - forward = np.concatenate(partition(np.arange(ngrids, dtype=np.int64))) - inverse = np.empty_like(forward) - inverse[forward] = np.arange(ngrids, dtype=np.int64) - return forward, inverse - - -def _prepare_spatially_sorted_grids( - mol: gto.Mole, - grids: Grid, - block_size: int, - gpu: bool, -) -> tuple[Grid, _Int64Permutation, _Int64Permutation]: - """Copy and spatially order a grid for backend AO screening. - - Preparation is cached on the source grid and reused while its coordinate and - weight arrays, molecule, backend, and evaluator block size remain unchanged. - - Args: - mol: Molecule used to rebuild CPU shell-screening data. - grids: Built CPU or GPU integration grid in atom-major order. - block_size: Fixed number of points consumed by each backend block. - gpu: Whether ``grids`` belongs to GPU4PySCF. - - Returns: - The sorted grid copy, atom-major-to-spatial permutation, and inverse. - """ - if grids.coords is None or grids.weights is None: - raise ValueError("Grids must be built before spatial sorting.") - - cache = getattr(grids, _SPATIAL_GRID_CACHE_ATTRIBUTE, None) - if isinstance(cache, _SpatialGridCache) and cache.matches( - mol, grids, block_size, gpu - ): - return cache.sorted_grids, cache.forward, cache.inverse - - if gpu: - check_gpu_imports_were_successful() - import cupy - - host_coords = cupy.asnumpy(grids.coords) - else: - host_coords = grids.coords - - forward, inverse = _spatial_grid_permutations(host_coords, block_size) - sorted_grids = copy(grids) - vars(sorted_grids).pop(_SPATIAL_GRID_CACHE_ATTRIBUTE, None) - if gpu: - backend_forward = cupy.asarray(forward) - sorted_grids.coords = grids.coords[backend_forward] - sorted_grids.weights = grids.weights[backend_forward] - sorted_grids._non0ao_idx = None - else: - sorted_grids.coords = grids.coords[forward] - sorted_grids.weights = grids.weights[forward] - sorted_grids.non0tab = dft.gen_grid.make_screen_index( - mol, - sorted_grids.coords, - cutoff=sorted_grids.cutoff, - ) - setattr( - grids, - _SPATIAL_GRID_CACHE_ATTRIBUTE, - _SpatialGridCache( - mol=mol, - source_coords=grids.coords, - source_weights=grids.weights, - block_size=block_size, - gpu=gpu, - sorted_grids=sorted_grids, - forward=forward, - inverse=inverse, - ), - ) - return sorted_grids, forward, inverse - - -def make_chunks( - atomic_grid_sizes: Tensor, max_grid_chunk_size: int -) -> list[tuple[slice, slice]]: - """ - Generate chunks of atomic and grid indices based on the maximum grid chunk size. - Input: - atomic_grid_sizes: A tensor of atomic grid sizes. - max_grid_chunk_size: The maximum size of each grid chunk. - Returns: - A list of tuples, where each tuple contains a slice for the atomic indices and a slice for the grid indices. - """ - - if max_grid_chunk_size < atomic_grid_sizes.max().item(): - raise ValueError( - "max_grid_chunk_size must be at least the maximum atomic grid size" - ) - - atom_and_grid_slices = [] - atom_start = 0 - grid_start = 0 - chunk_size = 0 - - for i, atom_grid_size in enumerate(atomic_grid_sizes): - chunk_size += atom_grid_size.item() - if chunk_size > max_grid_chunk_size: - atom_and_grid_slices.append( - ( - slice(atom_start, i), - slice(grid_start, grid_start + chunk_size - atom_grid_size.item()), - ) - ) - atom_start = i - grid_start += chunk_size - atom_grid_size.item() - chunk_size = atom_grid_size.item() - - if chunk_size > 0: - atom_and_grid_slices.append( - ( - slice(atom_start, len(atomic_grid_sizes)), - slice(grid_start, grid_start + chunk_size), - ) - ) - - LOG.debug( - f"Generated {len(atom_and_grid_slices)} chunks of grid sizes: {[g.stop - g.start for _, g in atom_and_grid_slices]}" - ) - - return atom_and_grid_slices def generate_features( @@ -359,158 +141,3 @@ def get_grid_features( ) return grid_features - - -@dataclass -class _GlobalScreenedFeatures: - dm: Tensor - mol: gto.Mole - sorted_grids: Grid - sorted_raw_features: Tensor - atom_major_raw_features: Tensor - forward_permutation: Tensor - inverse_permutation: Tensor - feature_function: feature_math.MGGAFeatureFunction - block_size: int - compile_feature_function: bool - gpu: bool - grid_features: dict[str, Tensor] - feature_spec: FeatureSpec - chunks: list[tuple[slice, slice]] - with_spin: bool - - def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: - """Apply the global raw-feature Jacobian and restore atom-major order.""" - sorted_tangent = ao_evaluation.ChunkEvalForward.apply( - self.dm, - self.mol, - self.sorted_grids, - self.feature_function, - self.block_size, - self.compile_feature_function, - self.gpu, - dm_tangent, - ) - return sorted_tangent.index_select(-1, self.inverse_permutation).detach() - - def build_model_chunk( - self, - raw_features: Tensor, - atom_slice: slice, - grid_slice: slice, - ) -> dict[str, Tensor]: - """Build one atom-aligned model dictionary from raw feature values.""" - feature_chunk: dict[str, Tensor] = {} - for feature_name in ("grid_coords", "grid_weights", "atomic_grid_weights"): - if self.feature_spec.requests(feature_name): - feature_chunk[feature_name] = self.grid_features[feature_name][ - grid_slice - ] - - for feature_name in ("coarse_0_atomic_coords", "atomic_grid_sizes"): - if self.feature_spec.requests(feature_name): - feature_chunk[feature_name] = self.grid_features[feature_name][ - atom_slice - ] - - if self.feature_spec.requests("atomic_grid_size_bound_shape"): - max_size = int(feature_chunk["atomic_grid_sizes"].max().item()) - feature_chunk["atomic_grid_size_bound_shape"] = torch.zeros( - max_size, - 0, - dtype=torch.long, - device=raw_features.device, - ) - - for feature_name, feature in self.feature_function.to_dict( - raw_features - ).items(): - feature_chunk[feature_name] = feature_math.maybe_expand_and_divide( - feature, not self.with_spin, 2 - ) - return feature_chunk - - -def _global_screened_features( - mol: gto.Mole, - dm: Tensor, - grids: Grid, - features: FeatureSpec | set[str], - func_deriv: int, - max_memory_in_mb: int | None = None, - safety_fraction: float = 0.8, - compile_feature_function: bool = False, -) -> _GlobalScreenedFeatures: - """Evaluate raw AO features once on a spatially ordered molecular grid.""" - feature_spec = ( - features if isinstance(features, FeatureSpec) else FeatureSpec(features) - ) - if not feature_spec.supports_screened_evaluation: - raise ValueError( - "Global screened features require 'atomic_grid_sizes' for model chunks." - ) - if grids.coords is None or grids.weights is None: - raise ValueError("Grids must be built before generating screened features.") - - feature_function = feature_math.MGGAFeatureFunction(feature_spec) - grid_features = get_grid_features(mol, dm, grids, feature_spec) - max_grid_chunk_size = estimate_max_grid_chunk_size( - dm=dm, - deriv=feature_function.deriv, - max_memory_in_mb=max_memory_in_mb, - safety_fraction=safety_fraction, - func_deriv=func_deriv, - reserved_memory_in_bytes=estimate_global_raw_feature_buffer_memory( - dm, - feature_function.nfeats, - grids.weights.size, - func_deriv, - ), - ) - max_atom_grid = int(grid_features["atomic_grid_sizes"].max().item()) - if max_grid_chunk_size < max_atom_grid: - LOG.warning( - f"Adjusted chunk size {max_grid_chunk_size} to match the largest atomic grid " - f"{max_atom_grid}. Hope for no OOM." - ) - max_grid_chunk_size = max_atom_grid - - gpu = dm.device.type == "cuda" - if gpu: - check_gpu_imports_were_successful() - block_size = int(dft_gpu.numint.MIN_BLK_SIZE) - else: - block_size = CPU_AO_SCREENING_BLOCK_SIZE - sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( - mol, grids, block_size, gpu - ) - sorted_raw_features = ao_evaluation.ChunkEvalForward.apply( - dm.double(), - mol, - sorted_grids, - feature_function, - block_size, - compile_feature_function, - gpu, - ) - forward_permutation = torch.as_tensor(forward, device=dm.device) - inverse_permutation = torch.as_tensor(inverse, device=dm.device) - atom_major_raw_features = sorted_raw_features.index_select(-1, inverse_permutation) - chunks = make_chunks(grid_features["atomic_grid_sizes"], max_grid_chunk_size) - return _GlobalScreenedFeatures( - dm=dm, - mol=mol, - sorted_grids=sorted_grids, - sorted_raw_features=sorted_raw_features, - atom_major_raw_features=atom_major_raw_features, - forward_permutation=forward_permutation, - inverse_permutation=inverse_permutation, - feature_function=feature_function, - block_size=block_size, - compile_feature_function=compile_feature_function, - gpu=gpu, - grid_features=grid_features, - feature_spec=feature_spec, - chunks=chunks, - with_spin=dm.ndim == 3, - ) diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index 1efb029c..d6f203a4 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -3,7 +3,7 @@ import torch -def estimate_max_grid_chunk_size( +def estimate_max_model_grid_points( dm: torch.Tensor, deriv: int, max_memory_in_mb: int | None = None, @@ -11,7 +11,7 @@ def estimate_max_grid_chunk_size( func_deriv: int = 1, reserved_memory_in_bytes: int = 0, ) -> int: - """Heuristically pick a model grid chunk size for screened feature evaluation. + """Heuristically limit grid points per atom-aligned model evaluation. The dominant per-chunk allocation is the atomic-orbital matrix evaluated by ``non_chunk`` (shape ``(ncomp, nao, n)`` in float64, with no AO screening), diff --git a/src/skala/pyscf/model_chunking.py b/src/skala/pyscf/model_chunking.py new file mode 100644 index 00000000..ed362cb7 --- /dev/null +++ b/src/skala/pyscf/model_chunking.py @@ -0,0 +1,196 @@ +# SPDX-License-Identifier: MIT + +"""Build atom-aligned model feature chunks from globally evaluated raw features. + +This module controls how many complete atomic grids are fed through the functional +model at once. Atomic grids are never split because model features may depend on +atom-local shapes and coordinates. +""" + +import logging +from collections.abc import Iterator +from dataclasses import dataclass + +import torch +from pyscf import gto +from torch import Tensor + +from skala.pyscf import feature_math +from skala.pyscf.backend import Grid +from skala.pyscf.features import get_grid_features +from skala.pyscf.memory_estimators import ( + estimate_global_raw_feature_buffer_memory, + estimate_max_model_grid_points, +) + +LOG = logging.getLogger(__name__) + + +@dataclass(frozen=True) +class AtomGridChunk: + """Matching atom and grid slices for one model evaluation chunk.""" + + atom_slice: slice + grid_slice: slice + + +def _make_atom_grid_chunks( + atomic_grid_sizes: Tensor, max_model_grid_points: int +) -> list[AtomGridChunk]: + """Build atom-aligned slices up to the requested model grid chunk size.""" + if max_model_grid_points < atomic_grid_sizes.max().item(): + raise ValueError( + "max_model_grid_points must be at least the maximum atomic grid size" + ) + + chunks: list[AtomGridChunk] = [] + atom_start = 0 + grid_start = 0 + chunk_size = 0 + + for atom_index, atom_grid_size in enumerate(atomic_grid_sizes): + chunk_size += atom_grid_size.item() + if chunk_size > max_model_grid_points: + chunks.append( + AtomGridChunk( + atom_slice=slice(atom_start, atom_index), + grid_slice=slice( + grid_start, grid_start + chunk_size - atom_grid_size.item() + ), + ) + ) + atom_start = atom_index + grid_start += chunk_size - atom_grid_size.item() + chunk_size = atom_grid_size.item() + + if chunk_size > 0: + chunks.append( + AtomGridChunk( + atom_slice=slice(atom_start, len(atomic_grid_sizes)), + grid_slice=slice(grid_start, grid_start + chunk_size), + ) + ) + + LOG.debug( + "Generated %d model chunks of grid sizes: %s", + len(chunks), + [chunk.grid_slice.stop - chunk.grid_slice.start for chunk in chunks], + ) + return chunks + + +@dataclass(frozen=True) +class ModelFeatureChunk: + """Chunk-local raw features and the corresponding model input dictionary.""" + + grid_slice: slice + raw_features: Tensor + model_features: dict[str, Tensor] + + +@dataclass(frozen=True) +class ModelFeatureChunker: + """Reusable atom-aligned partition of raw and model features.""" + + atom_major_raw_features: Tensor + grid_features: dict[str, Tensor] + feature_function: feature_math.MGGAFeatureFunction + chunk_layouts: list[AtomGridChunk] + with_spin: bool + + def __iter__(self) -> Iterator[ModelFeatureChunk]: + """Yield detached raw features paired with atom-aligned model inputs.""" + feature_spec = self.feature_function.feature_spec + for layout in self.chunk_layouts: + raw_features = ( + self.atom_major_raw_features[..., layout.grid_slice] + .detach() + .requires_grad_() + ) + model_features: dict[str, Tensor] = {} + for feature_name in ( + "grid_coords", + "grid_weights", + "atomic_grid_weights", + ): + if feature_spec.requests(feature_name): + model_features[feature_name] = self.grid_features[feature_name][ + layout.grid_slice + ] + + for feature_name in ("coarse_0_atomic_coords", "atomic_grid_sizes"): + if feature_spec.requests(feature_name): + model_features[feature_name] = self.grid_features[feature_name][ + layout.atom_slice + ] + + if feature_spec.requests("atomic_grid_size_bound_shape"): + max_size = int(model_features["atomic_grid_sizes"].max().item()) + model_features["atomic_grid_size_bound_shape"] = torch.zeros( + max_size, + 0, + dtype=torch.long, + device=raw_features.device, + ) + + for feature_name, feature in self.feature_function.to_dict( + raw_features + ).items(): + model_features[feature_name] = feature_math.maybe_expand_and_divide( + feature, not self.with_spin, 2 + ) + yield ModelFeatureChunk( + grid_slice=layout.grid_slice, + raw_features=raw_features, + model_features=model_features, + ) + + +def prepare_model_feature_chunks( + mol: gto.Mole, + dm: Tensor, + grids: Grid, + atom_major_raw_features: Tensor, + feature_function: feature_math.MGGAFeatureFunction, + func_deriv: int, + max_memory_in_mb: int | None = None, + safety_fraction: float = 0.8, +) -> ModelFeatureChunker: + """Prepare memory-sized, atom-aligned chunks for functional model evaluation.""" + feature_spec = feature_function.feature_spec + if not feature_spec.supports_screened_evaluation: + raise ValueError("Atom-aligned model chunking requires 'atomic_grid_sizes'.") + + grid_features = get_grid_features(mol, dm, grids, feature_spec) + max_model_grid_points = estimate_max_model_grid_points( + dm=dm, + deriv=feature_function.deriv, + max_memory_in_mb=max_memory_in_mb, + safety_fraction=safety_fraction, + func_deriv=func_deriv, + reserved_memory_in_bytes=estimate_global_raw_feature_buffer_memory( + dm, + feature_function.nfeats, + atom_major_raw_features.shape[-1], + func_deriv, + ), + ) + max_atom_grid = int(grid_features["atomic_grid_sizes"].max().item()) + if max_model_grid_points < max_atom_grid: + LOG.warning( + "Adjusted model chunk size %d to match the largest atomic grid %d. " + "Hope for no OOM.", + max_model_grid_points, + max_atom_grid, + ) + max_model_grid_points = max_atom_grid + + return ModelFeatureChunker( + atom_major_raw_features=atom_major_raw_features, + grid_features=grid_features, + feature_function=feature_function, + chunk_layouts=_make_atom_grid_chunks( + grid_features["atomic_grid_sizes"], max_model_grid_points + ), + with_spin=dm.ndim == 3, + ) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 5cab4754..a699888e 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -19,10 +19,9 @@ to_numpy, ) from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec -from skala.pyscf.features import ( - _global_screened_features, - generate_features, -) +from skala.pyscf.features import generate_features +from skala.pyscf.model_chunking import prepare_model_feature_chunks +from skala.pyscf.screening import prepare_screened_feature_buffer def _should_screen_aos(mol: gto.Mole) -> bool: @@ -222,11 +221,18 @@ def __call__( dm = dm.detach().requires_grad_() tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) - screened_features = _global_screened_features( + screened_features = prepare_screened_feature_buffer( mol, dm, grids, features=self.feature_spec, + ) + model_chunks = prepare_model_feature_chunks( + mol, + dm, + grids, + atom_major_raw_features=screened_features.atom_major_raw_features, + feature_function=screened_features.feature_function, func_deriv=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, safety_fraction=self.evaluation_policy.safety_fraction, @@ -235,23 +241,16 @@ def __call__( atom_major_cotangent = torch.zeros_like( screened_features.atom_major_raw_features ) - for atom_slice, grid_slice in screened_features.chunks: - # Break the reorder graph so model backprop retains only this chunk. - local_raw_features = ( - screened_features.atom_major_raw_features[..., grid_slice] - .detach() - .requires_grad_() - ) - mol_features = screened_features.build_model_chunk( - local_raw_features, atom_slice, grid_slice - ) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features E_xc_chunk = self.func.get_exc(mol_features) (local_cotangent,) = torch.autograd.grad( E_xc_chunk, local_raw_features, torch.ones_like(E_xc_chunk), ) - atom_major_cotangent[..., grid_slice] = local_cotangent.detach() + atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() tot_dens += ( (mol_features["density"] * mol_features["grid_weights"]) .sum(dim=-1) @@ -367,11 +366,18 @@ def gen_response( ks.mol ): dm0 = dm0.requires_grad_() - screened_features = _global_screened_features( + screened_features = prepare_screened_feature_buffer( ks.mol, dm0, ks.grids, features=self.feature_spec, + ) + model_chunks = prepare_model_feature_chunks( + ks.mol, + dm0, + ks.grids, + atom_major_raw_features=screened_features.atom_major_raw_features, + feature_function=screened_features.feature_function, func_deriv=2, max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, safety_fraction=kwargs.get( @@ -386,16 +392,9 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: atom_major_hessian_action = torch.zeros_like( screened_features.atom_major_raw_features ) - for atom_slice, grid_slice in screened_features.chunks: - # Isolate the second-order model graph to the current atomic chunk. - local_raw_features = ( - screened_features.atom_major_raw_features[..., grid_slice] - .detach() - .requires_grad_() - ) - mol_features = screened_features.build_model_chunk( - local_raw_features, atom_slice, grid_slice - ) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features E_xc_chunk = self.func.get_exc(mol_features) (local_gradient,) = torch.autograd.grad( E_xc_chunk, @@ -407,11 +406,11 @@ def hessian_vector_product_atom_chunked(dm1: Array) -> Array: (local_hessian_action,) = torch.autograd.grad( local_gradient, local_raw_features, - atom_major_tangent[..., grid_slice], + atom_major_tangent[..., chunk.grid_slice], ) else: local_hessian_action = torch.zeros_like(local_raw_features) - atom_major_hessian_action[..., grid_slice] = ( + atom_major_hessian_action[..., chunk.grid_slice] = ( local_hessian_action.detach() ) del ( diff --git a/src/skala/pyscf/screening.py b/src/skala/pyscf/screening.py new file mode 100644 index 00000000..7f53f8b9 --- /dev/null +++ b/src/skala/pyscf/screening.py @@ -0,0 +1,293 @@ +# SPDX-License-Identifier: MIT + +"""Extend PySCF and GPU4PySCF grids for screened Skala evaluation. + +Skala's AO evaluator benefits from spatially local grid blocks, while PySCF and +GPU4PySCF provide integration grids in atom-major order with backend-specific AO +screening metadata. This module is the extension layer interposed between those +backend-owned grid objects and Skala's feature evaluation. It deliberately avoids +subclassing either grid implementation so the same screened path can serve both. + +Grid preparation attaches a :class:`SpatialGridLayout` cache to the source grid. +The source coordinate and weight arrays retain their order, but the extra cache +attribute modifies the grid object. A shallow grid copy receives spatially reordered +coordinates and weights. For PySCF, its ``non0tab`` shell-screening table is rebuilt; +for GPU4PySCF, ``_non0ao_idx`` is cleared so the backend can rebuild it for the new +order. The cached forward and inverse permutations bridge spatial AO evaluation and +the atom-major layout expected by model features. + +:func:`prepare_screened_feature_buffer` orchestrates this grid extension and one +global raw-feature evaluation. The returned :class:`ScreenedFeatureBuffer` retains +the spatial ordering needed for Jacobian and adjoint AO passes. Atom-aligned model +batching is a separate process owned by :mod:`skala.pyscf.model_chunking`. +""" + +from copy import copy +from dataclasses import dataclass +from typing import Generic, TypeAlias, cast + +import numpy as np +import torch +from pyscf import dft, gto +from torch import Tensor + +from skala.pyscf import ao_evaluation, feature_math +from skala.pyscf.backend import ( + Array, + Grid, + check_gpu_imports_were_successful, +) +from skala.pyscf.evaluation import FeatureSpec + +CPU_AO_SCREENING_BLOCK_SIZE = 9 * dft.gen_grid.BLKSIZE + +_Float64Coordinates: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.float64]] +_Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] +_SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" + + +@dataclass(frozen=True) +class SpatialGridLayout(Generic[Array]): + """Cached spatial ordering derived from an atom-major integration grid.""" + + mol: gto.Mole + source_coords: Array + source_weights: Array + block_size: int + gpu: bool + sorted_grids: Grid + forward: _Int64Permutation + inverse: _Int64Permutation + + def matches( + self, + mol: gto.Mole, + grids: Grid, + block_size: int, + gpu: bool, + ) -> bool: + """Return whether this layout belongs to the current built grid.""" + return ( + self.mol is mol + and self.source_coords is grids.coords + and self.source_weights is grids.weights + and self.block_size == block_size + and self.gpu is gpu + ) + + +def _decompose_grid_into_spatial_blocks( + coords: _Float64Coordinates, block_size: int +) -> tuple[_Int64Permutation, _Int64Permutation]: + """Decompose a molecular grid into spatial blocks and return its permutations. + + Recursively partitions points along the longest Cartesian extent. Every left + subtree contains a whole number of evaluator blocks, so all output blocks have + ``block_size`` points except for a possible final remainder. + + Args: + coords: Molecular grid coordinates with shape ``(ngrids, 3)``. + block_size: Fixed number of points consumed by each backend block. + + Returns: + The forward permutation from atom-major to spatial order and its inverse. + + Raises: + ValueError: If the coordinates or block size are invalid. + """ + if coords.ndim != 2 or coords.shape[1] != 3: + raise ValueError("coords must have shape (ngrids, 3)") + if block_size <= 0: + raise ValueError("block_size must be positive") + + def partition(indices: _Int64Permutation) -> list[_Int64Permutation]: + if indices.size <= block_size: + return [indices] + + block_count = (indices.size + block_size - 1) // block_size + left_size = (block_count // 2) * block_size + extents = np.ptp(coords[indices], axis=0) + split_axis = int(np.argmax(extents)) + positions = np.lexsort((indices, coords[indices, split_axis])) + ordered_indices = indices[positions] + return partition(ordered_indices[:left_size]) + partition( + ordered_indices[left_size:] + ) + + ngrids = coords.shape[0] + if ngrids == 0: + empty = np.empty(0, dtype=np.int64) + return empty, empty.copy() + + forward = np.concatenate(partition(np.arange(ngrids, dtype=np.int64))) + inverse = np.empty_like(forward) + inverse[forward] = np.arange(ngrids, dtype=np.int64) + return forward, inverse + + +def _prepare_spatially_sorted_grids( + mol: gto.Mole, + grids: Grid, + block_size: int, + gpu: bool, +) -> tuple[Grid, _Int64Permutation, _Int64Permutation]: + """Copy and spatially order a grid for backend AO screening. + + Preparation is cached on the source grid and reused while its coordinate and + weight arrays, molecule, backend, and evaluator block size remain unchanged. + + Args: + mol: Molecule used to rebuild CPU shell-screening data. + grids: Built CPU or GPU integration grid in atom-major order. + block_size: Fixed number of points consumed by each backend block. + gpu: Whether ``grids`` belongs to GPU4PySCF. + + Returns: + The sorted grid copy, atom-major-to-spatial permutation, and inverse. + """ + if grids.coords is None or grids.weights is None: + raise ValueError("Grids must be built before spatial sorting.") + + cache = getattr(grids, _SPATIAL_GRID_CACHE_ATTRIBUTE, None) + if isinstance(cache, SpatialGridLayout) and cache.matches( + mol, grids, block_size, gpu + ): + return cache.sorted_grids, cache.forward, cache.inverse + + if gpu: + check_gpu_imports_were_successful() + import cupy + + host_coords = cast(_Float64Coordinates, cupy.asnumpy(grids.coords)) + else: + host_coords = cast(_Float64Coordinates, grids.coords) + + forward, inverse = _decompose_grid_into_spatial_blocks(host_coords, block_size) + sorted_grids = copy(grids) + vars(sorted_grids).pop(_SPATIAL_GRID_CACHE_ATTRIBUTE, None) + if gpu: + import cupy + + backend_forward = cupy.asarray(forward) + sorted_grids.coords = grids.coords[backend_forward] + sorted_grids.weights = grids.weights[backend_forward] + sorted_grids._non0ao_idx = None + else: + sorted_grids.coords = grids.coords[forward] + sorted_grids.weights = grids.weights[forward] + sorted_grids.non0tab = dft.gen_grid.make_screen_index( + mol, + sorted_grids.coords, + cutoff=sorted_grids.cutoff, + ) + setattr( + grids, + _SPATIAL_GRID_CACHE_ATTRIBUTE, + SpatialGridLayout( + mol=mol, + source_coords=grids.coords, + source_weights=grids.weights, + block_size=block_size, + gpu=gpu, + sorted_grids=sorted_grids, + forward=forward, + inverse=inverse, + ), + ) + return sorted_grids, forward, inverse + + +@dataclass +class ScreenedFeatureBuffer: + """Globally screened raw features and transformations between grid orders.""" + + dm: Tensor + mol: gto.Mole + sorted_grids: Grid + sorted_raw_features: Tensor + atom_major_raw_features: Tensor + forward_permutation: Tensor + inverse_permutation: Tensor + feature_function: feature_math.MGGAFeatureFunction + block_size: int + compile_feature_function: bool + gpu: bool + + def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: + """Apply the global raw-feature Jacobian and restore atom-major order.""" + sorted_tangent = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( + self.dm, + self.mol, + self.sorted_grids, + self.feature_function, + self.block_size, + self.compile_feature_function, + self.gpu, + dm_tangent, + ), + ) + return sorted_tangent.index_select(-1, self.inverse_permutation).detach() + + +def prepare_screened_feature_buffer( + mol: gto.Mole, + dm: Tensor, + grids: Grid, + features: FeatureSpec | set[str], + compile_feature_function: bool = False, +) -> ScreenedFeatureBuffer: + """Prepare the backend grid extension and its screened feature buffer. + + The source grid receives a cached :class:`SpatialGridLayout`; its coordinate and + weight arrays remain atom-major. Feature evaluation runs once on a spatially + reordered grid copy, and the resulting buffer restores atom-major order for model + chunk construction. + """ + feature_spec = ( + features if isinstance(features, FeatureSpec) else FeatureSpec(features) + ) + if grids.coords is None or grids.weights is None: + raise ValueError("Grids must be built before generating screened features.") + + feature_function = feature_math.MGGAFeatureFunction(feature_spec) + gpu = dm.device.type == "cuda" + if gpu: + check_gpu_imports_were_successful() + from gpu4pyscf.dft import numint as dft_gpu_numint + + block_size = int(dft_gpu_numint.MIN_BLK_SIZE) + else: + block_size = CPU_AO_SCREENING_BLOCK_SIZE + sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( + mol, grids, block_size, gpu + ) + sorted_raw_features = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( + dm.double(), + mol, + sorted_grids, + feature_function, + block_size, + compile_feature_function, + gpu, + ), + ) + forward_permutation = torch.as_tensor(forward, device=dm.device) + inverse_permutation = torch.as_tensor(inverse, device=dm.device) + atom_major_raw_features = sorted_raw_features.index_select(-1, inverse_permutation) + return ScreenedFeatureBuffer( + dm=dm, + mol=mol, + sorted_grids=sorted_grids, + sorted_raw_features=sorted_raw_features, + atom_major_raw_features=atom_major_raw_features, + forward_permutation=forward_permutation, + inverse_permutation=inverse_permutation, + feature_function=feature_function, + block_size=block_size, + compile_feature_function=compile_feature_function, + gpu=gpu, + ) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index c63935aa..ae653d8b 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -7,8 +7,9 @@ from pyscf.dft import numint as pyscf_numint from skala.functional.base import ExcFunctionalBase -from skala.pyscf import features as features_module +from skala.pyscf import model_chunking as model_chunking_module from skala.pyscf import numint as numint_module +from skala.pyscf import screening as screening_module from skala.pyscf.ao_evaluation import ( ChunkEvalBackward, ChunkEvalForward, @@ -19,11 +20,12 @@ ) from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.feature_math import MGGAFeatureFunction -from skala.pyscf.features import ( +from skala.pyscf.model_chunking import ModelFeatureChunk +from skala.pyscf.numint import SkalaNumInt, _should_screen_aos +from skala.pyscf.screening import ( + _decompose_grid_into_spatial_blocks, _prepare_spatially_sorted_grids, - _spatial_grid_permutations, ) -from skala.pyscf.numint import SkalaNumInt, _should_screen_aos @pytest.fixture @@ -167,12 +169,12 @@ def test_resolve_ao_block_size_modes(carbon: gto.Mole) -> None: @pytest.mark.parametrize(("ngrids", "block_size"), [(0, 4), (3, 4), (8, 4), (10, 4)]) -def test_spatial_grid_permutations_restore_original_order( +def test_decompose_grid_into_spatial_blocks_restores_original_order( ngrids: int, block_size: int ) -> None: coords = np.arange(3 * ngrids, dtype=np.float64).reshape(ngrids, 3) - forward, inverse = _spatial_grid_permutations(coords, block_size) + forward, inverse = _decompose_grid_into_spatial_blocks(coords, block_size) assert np.array_equal(np.sort(forward), np.arange(ngrids)) assert np.array_equal(coords[forward][inverse], coords) @@ -183,7 +185,7 @@ def test_spatial_grid_permutations_restore_original_order( assert len(forward) % block_size == ngrids % block_size -def test_spatial_grid_permutations_group_interleaved_clusters() -> None: +def test_decompose_grid_into_spatial_blocks_groups_interleaved_clusters() -> None: block_size = 3 labels = np.tile(np.arange(4), block_size) offsets = np.repeat(np.arange(block_size), 4) @@ -191,7 +193,7 @@ def test_spatial_grid_permutations_group_interleaved_clusters() -> None: (100.0 * labels + offsets, np.zeros(labels.size), np.zeros(labels.size)) ) - forward, _ = _spatial_grid_permutations(coords, block_size) + forward, _ = _decompose_grid_into_spatial_blocks(coords, block_size) grouped_labels = labels[forward].reshape(-1, block_size) assert np.all(grouped_labels == grouped_labels[:, :1]) @@ -210,7 +212,7 @@ def test_prepare_spatially_sorted_cpu_grids( non0tab = np.ones((1, carbon.nbas), dtype=np.uint8) partition_calls = 0 - def fake_spatial_grid_permutations( + def fake_decompose_grid_into_spatial_blocks( coords_arg: np.ndarray, block_size: int ) -> tuple[np.ndarray, np.ndarray]: nonlocal partition_calls @@ -220,9 +222,9 @@ def fake_spatial_grid_permutations( return forward, inverse monkeypatch.setattr( - features_module, - "_spatial_grid_permutations", - fake_spatial_grid_permutations, + screening_module, + "_decompose_grid_into_spatial_blocks", + fake_decompose_grid_into_spatial_blocks, ) screen_index_calls = 0 @@ -322,49 +324,67 @@ def fake_generate_features( "grid_weights": torch.ones(1, dtype=dm.dtype), } - class FakeGlobalScreenedFeatures: + class FakeScreenedFeatureBuffer: def __init__(self, dm: torch.Tensor) -> None: raw_features = dm.sum().reshape(1, 1) self.feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) self.sorted_raw_features = raw_features self.atom_major_raw_features = raw_features self.forward_permutation = torch.tensor([0]) - self.chunks = [(slice(0, 1), slice(0, 1))] def atom_major_jvp(self, dm_tangent: torch.Tensor) -> torch.Tensor: return dm_tangent.sum().reshape(1, 1) - def build_model_chunk( - self, - raw_features: torch.Tensor, - atom_slice: slice, - grid_slice: slice, - ) -> dict[str, torch.Tensor]: - assert atom_slice == slice(0, 1) - assert grid_slice == slice(0, 1) - return { - "atomic_grid_sizes": torch.tensor([1]), - "density": raw_features.expand(2, 1) / 2, - "grid_weights": torch.ones(1, dtype=raw_features.dtype), - } - - def fake_global_screened_features( + class FakeModelFeatureChunks: + def __init__(self, raw_features: torch.Tensor) -> None: + self.raw_features = raw_features + + def __iter__(self) -> Iterator[ModelFeatureChunk]: + raw_features = self.raw_features.detach().requires_grad_() + yield ModelFeatureChunk( + grid_slice=slice(0, 1), + raw_features=raw_features, + model_features={ + "atomic_grid_sizes": torch.tensor([1]), + "density": raw_features.expand(2, 1) / 2, + "grid_weights": torch.ones(1, dtype=raw_features.dtype), + }, + ) + + def fake_prepare_screened_feature_buffer( mol: gto.Mole, dm: torch.Tensor, grids: object, features: set[str], - func_deriv: int, **kwargs: object, - ) -> FakeGlobalScreenedFeatures: + ) -> FakeScreenedFeatureBuffer: routes.append("screened") + return FakeScreenedFeatureBuffer(dm) + + def fake_prepare_model_feature_chunks( + mol: gto.Mole, + dm: torch.Tensor, + grids: object, + atom_major_raw_features: torch.Tensor, + feature_function: MGGAFeatureFunction, + func_deriv: int, + **kwargs: object, + ) -> FakeModelFeatureChunks: safety_fraction = kwargs["safety_fraction"] assert isinstance(safety_fraction, float) safety_fractions.append(safety_fraction) - return FakeGlobalScreenedFeatures(dm) + return FakeModelFeatureChunks(atom_major_raw_features) monkeypatch.setattr(numint_module, "generate_features", fake_generate_features) monkeypatch.setattr( - numint_module, "_global_screened_features", fake_global_screened_features + numint_module, + "prepare_screened_feature_buffer", + fake_prepare_screened_feature_buffer, + ) + monkeypatch.setattr( + numint_module, + "prepare_model_feature_chunks", + fake_prepare_model_feature_chunks, ) numint = SkalaNumInt(QuadraticDensityFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) @@ -635,7 +655,7 @@ def test_cpu_response_dense_screened_equivalence( @pytest.mark.parametrize("func_deriv", [1, 2]) -def test_global_screened_ao_traversals_are_independent_of_model_chunks( +def test_screened_ao_traversals_are_independent_of_model_chunking( monkeypatch: pytest.MonkeyPatch, func_deriv: int, ) -> None: @@ -643,8 +663,8 @@ def test_global_screened_ao_traversals_are_independent_of_model_chunks( grids = _minimal_atom_grid(mol) atom_grid_size = grids.weights.size // mol.natm monkeypatch.setattr( - features_module, - "estimate_max_grid_chunk_size", + model_chunking_module, + "estimate_max_model_grid_points", lambda *args, **kwargs: atom_grid_size, ) monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 260c705f..3f0ac566 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -33,8 +33,8 @@ from skala.pyscf.backend import dft_gpu # noqa: E402 from skala.pyscf.evaluation import FeatureSpec # noqa: E402 from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 -from skala.pyscf.features import _prepare_spatially_sorted_grids # noqa: E402 from skala.pyscf.numint import SkalaNumInt # noqa: E402 +from skala.pyscf.screening import _prepare_spatially_sorted_grids # noqa: E402 CARBON_CHAIN = """ C 0.0 0.0 0.0 diff --git a/tests/test_memory_estimators.py b/tests/test_memory_estimators.py index 79821633..05f5914f 100644 --- a/tests/test_memory_estimators.py +++ b/tests/test_memory_estimators.py @@ -3,7 +3,7 @@ from skala.pyscf.memory_estimators import ( estimate_global_raw_feature_buffer_memory, - estimate_max_grid_chunk_size, + estimate_max_model_grid_points, linear_peak_memory_model, ) @@ -37,7 +37,7 @@ def test_global_raw_feature_buffer_memory_rejects_unsupported_order() -> None: def test_reserved_memory_reduces_grid_chunk_size() -> None: dm = torch.eye(10, dtype=torch.float64) bytes_per_point, _ = linear_peak_memory_model(nao=10, deriv=1, func_deriv=1) - base_chunk_size = estimate_max_grid_chunk_size( + base_chunk_size = estimate_max_model_grid_points( dm, deriv=1, max_memory_in_mb=100, @@ -45,7 +45,7 @@ def test_reserved_memory_reduces_grid_chunk_size() -> None: func_deriv=1, ) reserved_points = 123 - reserved_chunk_size = estimate_max_grid_chunk_size( + reserved_chunk_size = estimate_max_model_grid_points( dm, deriv=1, max_memory_in_mb=100, @@ -58,13 +58,13 @@ def test_reserved_memory_reduces_grid_chunk_size() -> None: @pytest.mark.parametrize("safety_fraction", [-0.1, 0.0, 1.1]) -def test_grid_chunk_size_rejects_invalid_safety_fraction( +def test_model_grid_point_limit_rejects_invalid_safety_fraction( safety_fraction: float, ) -> None: with pytest.raises( ValueError, match="safety_fraction must be greater than 0 and at most 1" ): - estimate_max_grid_chunk_size( + estimate_max_model_grid_points( torch.eye(2, dtype=torch.float64), deriv=1, max_memory_in_mb=100, From 92ef053affacdb8d3bd593e6410e3328cfbe9603 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 15:04:15 +0200 Subject: [PATCH 18/39] refactor make screened grid cache explicit --- src/skala/pyscf/numint.py | 369 +++++++++++++++++---------- src/skala/pyscf/screening.py | 205 ++++----------- tests/test_ao_screening.py | 172 ++++++++++--- tests/test_gpu4pyscf_ao_screening.py | 23 +- 4 files changed, 418 insertions(+), 351 deletions(-) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index a699888e..8cb12cde 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -1,14 +1,15 @@ # SPDX-License-Identifier: MIT from collections.abc import Callable -from typing import Any, Generic, Protocol, overload +from typing import Any, Generic, Protocol, cast, overload import torch -from pyscf import dft, gto +from pyscf import gto from pyscf.dft import numint as pyscf_numint from torch import Tensor from skala.functional.base import ExcFunctionalBase +from skala.pyscf import ao_evaluation, feature_math from skala.pyscf.backend import ( KS, Array, @@ -21,7 +22,12 @@ from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec from skala.pyscf.features import generate_features from skala.pyscf.model_chunking import prepare_model_feature_chunks -from skala.pyscf.screening import prepare_screened_feature_buffer +from skala.pyscf.screening import ( + CPU_AO_SCREENING_BLOCK_SIZE, + SpatialGridLayout, + prepare_spatial_grid_layout, + screened_feature_jvp, +) def _should_screen_aos(mol: gto.Mole) -> bool: @@ -135,10 +141,7 @@ def __init__( chunk_size: int | None = None, device: torch.device | None = None, ): - if device is None: - self.device = torch.get_default_device() - else: - self.device = device + self.device = device or torch.get_default_device() if self.device.type == "cuda": check_gpu_imports_were_successful() @@ -147,6 +150,40 @@ def __init__( self.feature_spec = FeatureSpec(self.func.features) self.evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) + def reset(self) -> "SkalaNumInt[Array]": + """Return this integrator; spatial layouts are owned by grid objects.""" + return self + + def _get_spatial_grid_layout( + self, + mol: gto.Mole, + grids: Grid, + ) -> SpatialGridLayout: + grid_state = vars(grids) + spatial_grid_layout = cast( + SpatialGridLayout | None, + grid_state.get("_skala_spatial_grid_layout"), + ) + if spatial_grid_layout is not None: + return spatial_grid_layout + + if self.device.type == "cuda": + check_gpu_imports_were_successful() + from gpu4pyscf.dft import numint as dft_gpu_numint + + block_size = int(dft_gpu_numint.MIN_BLK_SIZE) + else: + block_size = CPU_AO_SCREENING_BLOCK_SIZE + + spatial_grid_layout = prepare_spatial_grid_layout( + mol, + grids, + block_size, + self.device, + ) + grid_state["_skala_spatial_grid_layout"] = spatial_grid_layout + return spatial_grid_layout + def from_backend( self, x: Array, @@ -192,25 +229,17 @@ def get_rho( def __call__( self, mol: gto.Mole, - grids: dft.Grids, + grids: Grid, xc_code: str | None, dm: Tensor, second_order: bool = False, max_memory: int = 2000, ) -> tuple[Tensor, Tensor, Tensor]: - """ - Evaluate the XC functional for the given molecule and density matrix. - Input: - mol: The molecule. - grids: The grid. - xc_code: The XC code (not used in the reimplementation). - dm: The density matrix. - second_order: Whether to compute second-order derivatives. - max_memory: The maximum memory to use for each chunk in megabytes (MB). If None, the maximum memory is determined automatically. - - Returns: - A tuple of the total integrated density, the XC energy, and the XC potential. - """ + """Evaluate the XC functional for a molecule and density matrix.""" + if second_order: + raise NotImplementedError( + "Direct second-order evaluation is not supported; use gen_response()." + ) if self.device != dm.device: raise ValueError( @@ -218,58 +247,87 @@ def __call__( ) if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): - dm = dm.detach().requires_grad_() - tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) - E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) - screened_features = prepare_screened_feature_buffer( - mol, - dm, - grids, - features=self.feature_spec, - ) - model_chunks = prepare_model_feature_chunks( + return self._call_screened(mol, grids, dm, max_memory) + return self._call_dense(mol, grids, dm, max_memory) + + def _call_screened( + self, + mol: gto.Mole, + grids: Grid, + dm: Tensor, + max_memory: int, + ) -> tuple[Tensor, Tensor, Tensor]: + dm = dm.detach().requires_grad_() + tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) + E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) + feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) + spatial_grid_layout = self._get_spatial_grid_layout(mol, grids) + sorted_raw_features = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm.double(), mol, - dm, - grids, - atom_major_raw_features=screened_features.atom_major_raw_features, - feature_function=screened_features.feature_function, - func_deriv=1, - max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, - safety_fraction=self.evaluation_policy.safety_fraction, + spatial_grid_layout.sorted_grids, + feature_function, + spatial_grid_layout.block_size, + False, + dm.device.type == "cuda", + ), + ) + atom_major_raw_features = sorted_raw_features.index_select( + -1, spatial_grid_layout.inverse_permutation + ) + model_chunks = prepare_model_feature_chunks( + mol, + dm, + grids, + atom_major_raw_features=atom_major_raw_features, + feature_function=feature_function, + func_deriv=1, + max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, + safety_fraction=self.evaluation_policy.safety_fraction, + ) + # Store only full-grid feature cotangents; model activations remain chunk-local. + atom_major_cotangent = torch.zeros_like(atom_major_raw_features) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features + E_xc_chunk = self.func.get_exc(mol_features) + (local_cotangent,) = torch.autograd.grad( + E_xc_chunk, + local_raw_features, + torch.ones_like(E_xc_chunk), ) - # Store only full-grid feature cotangents; model activations remain chunk-local. - atom_major_cotangent = torch.zeros_like( - screened_features.atom_major_raw_features + atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() + tot_dens += ( + (mol_features["density"] * mol_features["grid_weights"]) + .sum(dim=-1) + .detach() ) - for chunk in model_chunks: - local_raw_features = chunk.raw_features - mol_features = chunk.model_features - E_xc_chunk = self.func.get_exc(mol_features) - (local_cotangent,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - ) - atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() - tot_dens += ( - (mol_features["density"] * mol_features["grid_weights"]) - .sum(dim=-1) - .detach() - ) - E_xc += E_xc_chunk.detach() - del E_xc_chunk, local_cotangent, local_raw_features, mol_features + E_xc += E_xc_chunk.detach() + del E_xc_chunk, local_cotangent, local_raw_features, mol_features - # Reorder detached cotangents explicitly instead of backpropagating through it. - sorted_cotangent = atom_major_cotangent.index_select( - -1, screened_features.forward_permutation - ) - # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. - (V_xc,) = torch.autograd.grad( - screened_features.sorted_raw_features, - dm, - sorted_cotangent, - ) - return tot_dens, E_xc, V_xc + # Reorder detached cotangents explicitly instead of backpropagating through it. + sorted_cotangent = atom_major_cotangent.index_select( + -1, spatial_grid_layout.forward_permutation + ) + # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. + (V_xc,) = torch.autograd.grad( + sorted_raw_features, + dm, + sorted_cotangent, + ) + return tot_dens, E_xc, V_xc + + def _call_dense( + self, + mol: gto.Mole, + grids: Grid, + dm: Tensor, + max_memory: int, + *, + create_graph: bool = False, + ) -> tuple[Tensor, Tensor, Tensor]: dm = dm.requires_grad_() mol_features = generate_features( @@ -286,8 +344,8 @@ def __call__( E_xc, dm, torch.ones_like(E_xc), - retain_graph=second_order, - create_graph=second_order, + retain_graph=create_graph, + create_graph=create_graph, ) rho = mol_features["density"] @@ -365,87 +423,126 @@ def gen_response( if self.feature_spec.supports_screened_evaluation and _should_screen_aos( ks.mol ): - dm0 = dm0.requires_grad_() - screened_features = prepare_screened_feature_buffer( - ks.mol, + return self._gen_response_screened( + ks, dm0, - ks.grids, - features=self.feature_spec, - ) - model_chunks = prepare_model_feature_chunks( - ks.mol, - dm0, - ks.grids, - atom_major_raw_features=screened_features.atom_major_raw_features, - feature_function=screened_features.feature_function, - func_deriv=2, - max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, safety_fraction=kwargs.get( "safety_fraction", self.evaluation_policy.safety_fraction ), ) + return self._gen_response_dense(ks, dm0) + + def _gen_response_screened( + self, + ks: KS, + dm0: Tensor, + *, + safety_fraction: float, + ) -> Callable[[Array], Array]: + dm0 = dm0.requires_grad_() + feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) + spatial_grid_layout = self._get_spatial_grid_layout(ks.mol, ks.grids) + sorted_raw_features = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm0.double(), + ks.mol, + spatial_grid_layout.sorted_grids, + feature_function, + spatial_grid_layout.block_size, + False, + dm0.device.type == "cuda", + ), + ) + atom_major_raw_features = sorted_raw_features.index_select( + -1, spatial_grid_layout.inverse_permutation + ) + model_chunks = prepare_model_feature_chunks( + ks.mol, + dm0, + ks.grids, + atom_major_raw_features=atom_major_raw_features, + feature_function=feature_function, + func_deriv=2, + max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, + safety_fraction=safety_fraction, + ) - def hessian_vector_product_atom_chunked(dm1: Array) -> Array: - dm1_tensor = self.from_backend(dm1) - atom_major_tangent = screened_features.atom_major_jvp(dm1_tensor) - # Store the full-grid model Hessian action, not per-chunk model graphs. - atom_major_hessian_action = torch.zeros_like( - screened_features.atom_major_raw_features + def hessian_vector_product_atom_chunked(dm1: Array) -> Array: + dm1_tensor = self.from_backend(dm1) + atom_major_tangent = screened_feature_jvp( + dm0, + dm1_tensor, + ks.mol, + spatial_grid_layout, + feature_function, + ) + # Store the full-grid model Hessian action, not per-chunk model graphs. + atom_major_hessian_action = torch.zeros_like(atom_major_raw_features) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features + E_xc_chunk = self.func.get_exc(mol_features) + (local_gradient,) = torch.autograd.grad( + E_xc_chunk, + local_raw_features, + torch.ones_like(E_xc_chunk), + create_graph=True, ) - for chunk in model_chunks: - local_raw_features = chunk.raw_features - mol_features = chunk.model_features - E_xc_chunk = self.func.get_exc(mol_features) - (local_gradient,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - create_graph=True, - ) - if local_gradient.requires_grad: - (local_hessian_action,) = torch.autograd.grad( - local_gradient, - local_raw_features, - atom_major_tangent[..., chunk.grid_slice], - ) - else: - local_hessian_action = torch.zeros_like(local_raw_features) - atom_major_hessian_action[..., chunk.grid_slice] = ( - local_hessian_action.detach() - ) - del ( - E_xc_chunk, + if local_gradient.requires_grad: + (local_hessian_action,) = torch.autograd.grad( local_gradient, - local_hessian_action, local_raw_features, - mol_features, + atom_major_tangent[..., chunk.grid_slice], ) - - # Restore block order after all chunk-local Hessian actions are detached. - sorted_hessian_action = atom_major_hessian_action.index_select( - -1, screened_features.forward_permutation + else: + local_hessian_action = torch.zeros_like(local_raw_features) + atom_major_hessian_action[..., chunk.grid_slice] = ( + local_hessian_action.detach() ) - # The custom VJP traverses AO blocks sequentially and retains no AO graph. - (hvp_total,) = torch.autograd.grad( - screened_features.sorted_raw_features, - dm0, - sorted_hessian_action, - retain_graph=True, + del ( + E_xc_chunk, + local_gradient, + local_hessian_action, + local_raw_features, + mol_features, ) - v1 = self.to_backend(hvp_total) - vj = ks.get_j(ks.mol, dm1, hermi=1) - if ks.mol.spin == 0: - v1 += vj - else: - v1 += vj[0] + vj[1] - return v1 + # Restore block order after all chunk-local Hessian actions are detached. + sorted_hessian_action = atom_major_hessian_action.index_select( + -1, spatial_grid_layout.forward_permutation + ) + # The custom VJP traverses AO blocks sequentially and retains no AO graph. + (hvp_total,) = torch.autograd.grad( + sorted_raw_features, + dm0, + sorted_hessian_action, + retain_graph=True, + ) + + v1 = self.to_backend(hvp_total) + vj = ks.get_j(ks.mol, dm1, hermi=1) + if ks.mol.spin == 0: + v1 += vj + else: + v1 += vj[0] + vj[1] + return v1 - return hessian_vector_product_atom_chunked + return hessian_vector_product_atom_chunked - # caching V_xc saves a forward pass in each iteration + def _gen_response_dense( + self, + ks: KS, + dm0: Tensor, + ) -> Callable[[Array], Array]: dm0 = dm0.requires_grad_() - V_xc = self(ks.mol, ks.grids, None, dm0, second_order=True)[2] + _, _, V_xc = self._call_dense( + ks.mol, + ks.grids, + dm0, + ks.max_memory, + create_graph=True, + ) def hessian_vector_product(dm1: Array) -> Array: v1 = self.to_backend( diff --git a/src/skala/pyscf/screening.py b/src/skala/pyscf/screening.py index 7f53f8b9..1f75976a 100644 --- a/src/skala/pyscf/screening.py +++ b/src/skala/pyscf/screening.py @@ -8,23 +8,23 @@ backend-owned grid objects and Skala's feature evaluation. It deliberately avoids subclassing either grid implementation so the same screened path can serve both. -Grid preparation attaches a :class:`SpatialGridLayout` cache to the source grid. -The source coordinate and weight arrays retain their order, but the extra cache -attribute modifies the grid object. A shallow grid copy receives spatially reordered -coordinates and weights. For PySCF, its ``non0tab`` shell-screening table is rebuilt; -for GPU4PySCF, ``_non0ao_idx`` is cleared so the backend can rebuild it for the new -order. The cached forward and inverse permutations bridge spatial AO evaluation and -the atom-major layout expected by model features. - -:func:`prepare_screened_feature_buffer` orchestrates this grid extension and one -global raw-feature evaluation. The returned :class:`ScreenedFeatureBuffer` retains -the spatial ordering needed for Jacobian and adjoint AO passes. Atom-aligned model -batching is a separate process owned by :mod:`skala.pyscf.model_chunking`. +Grid preparation produces a :class:`SpatialGridLayout`, which is cached on the source +grid for later evaluations. The source grid's integration data remains unchanged. A +shallow grid copy receives spatially reordered coordinates and weights. For PySCF, +its ``non0tab`` shell-screening table is rebuilt; for GPU4PySCF, ``_non0ao_idx`` is +cleared so the backend can rebuild it for the new order. The cached forward and +inverse permutations bridge spatial AO evaluation and the atom-major layout expected +by model features. + +:func:`prepare_spatial_grid_layout` owns this reusable grid extension. The integrator +attaches it to the source grid and owns density-dependent feature evaluation. +Atom-aligned model batching is a separate process owned by +:mod:`skala.pyscf.model_chunking`. """ from copy import copy from dataclasses import dataclass -from typing import Generic, TypeAlias, cast +from typing import TypeAlias, cast import numpy as np import torch @@ -32,48 +32,22 @@ from torch import Tensor from skala.pyscf import ao_evaluation, feature_math -from skala.pyscf.backend import ( - Array, - Grid, - check_gpu_imports_were_successful, -) -from skala.pyscf.evaluation import FeatureSpec +from skala.pyscf.backend import Grid, check_gpu_imports_were_successful CPU_AO_SCREENING_BLOCK_SIZE = 9 * dft.gen_grid.BLKSIZE _Float64Coordinates: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.float64]] _Int64Permutation: TypeAlias = np.ndarray[tuple[int], np.dtype[np.int64]] -_SPATIAL_GRID_CACHE_ATTRIBUTE = "_skala_spatial_grid_cache" @dataclass(frozen=True) -class SpatialGridLayout(Generic[Array]): - """Cached spatial ordering derived from an atom-major integration grid.""" +class SpatialGridLayout: + """Evaluation-ready spatial ordering derived from an atom-major grid.""" - mol: gto.Mole - source_coords: Array - source_weights: Array block_size: int - gpu: bool sorted_grids: Grid - forward: _Int64Permutation - inverse: _Int64Permutation - - def matches( - self, - mol: gto.Mole, - grids: Grid, - block_size: int, - gpu: bool, - ) -> bool: - """Return whether this layout belongs to the current built grid.""" - return ( - self.mol is mol - and self.source_coords is grids.coords - and self.source_weights is grids.weights - and self.block_size == block_size - and self.gpu is gpu - ) + forward_permutation: Tensor + inverse_permutation: Tensor def _decompose_grid_into_spatial_blocks( @@ -125,35 +99,27 @@ def partition(indices: _Int64Permutation) -> list[_Int64Permutation]: return forward, inverse -def _prepare_spatially_sorted_grids( +def prepare_spatial_grid_layout( mol: gto.Mole, grids: Grid, block_size: int, - gpu: bool, -) -> tuple[Grid, _Int64Permutation, _Int64Permutation]: - """Copy and spatially order a grid for backend AO screening. - - Preparation is cached on the source grid and reused while its coordinate and - weight arrays, molecule, backend, and evaluator block size remain unchanged. + device: torch.device, +) -> SpatialGridLayout: + """Build a spatially ordered grid layout for backend AO screening. Args: mol: Molecule used to rebuild CPU shell-screening data. grids: Built CPU or GPU integration grid in atom-major order. block_size: Fixed number of points consumed by each backend block. - gpu: Whether ``grids`` belongs to GPU4PySCF. + device: Torch device used for permutation tensors. Returns: - The sorted grid copy, atom-major-to-spatial permutation, and inverse. + An evaluation-ready layout containing the sorted grid and both permutations. """ if grids.coords is None or grids.weights is None: raise ValueError("Grids must be built before spatial sorting.") - cache = getattr(grids, _SPATIAL_GRID_CACHE_ATTRIBUTE, None) - if isinstance(cache, SpatialGridLayout) and cache.matches( - mol, grids, block_size, gpu - ): - return cache.sorted_grids, cache.forward, cache.inverse - + gpu = device.type == "cuda" if gpu: check_gpu_imports_were_successful() import cupy @@ -164,7 +130,6 @@ def _prepare_spatially_sorted_grids( forward, inverse = _decompose_grid_into_spatial_blocks(host_coords, block_size) sorted_grids = copy(grids) - vars(sorted_grids).pop(_SPATIAL_GRID_CACHE_ATTRIBUTE, None) if gpu: import cupy @@ -180,114 +145,36 @@ def _prepare_spatially_sorted_grids( sorted_grids.coords, cutoff=sorted_grids.cutoff, ) - setattr( - grids, - _SPATIAL_GRID_CACHE_ATTRIBUTE, - SpatialGridLayout( - mol=mol, - source_coords=grids.coords, - source_weights=grids.weights, - block_size=block_size, - gpu=gpu, - sorted_grids=sorted_grids, - forward=forward, - inverse=inverse, - ), + return SpatialGridLayout( + block_size=block_size, + sorted_grids=sorted_grids, + forward_permutation=torch.as_tensor(forward, device=device), + inverse_permutation=torch.as_tensor(inverse, device=device), ) - return sorted_grids, forward, inverse - - -@dataclass -class ScreenedFeatureBuffer: - """Globally screened raw features and transformations between grid orders.""" - - dm: Tensor - mol: gto.Mole - sorted_grids: Grid - sorted_raw_features: Tensor - atom_major_raw_features: Tensor - forward_permutation: Tensor - inverse_permutation: Tensor - feature_function: feature_math.MGGAFeatureFunction - block_size: int - compile_feature_function: bool - gpu: bool - - def atom_major_jvp(self, dm_tangent: Tensor) -> Tensor: - """Apply the global raw-feature Jacobian and restore atom-major order.""" - sorted_tangent = cast( - Tensor, - ao_evaluation.ChunkEvalForward.apply( - self.dm, - self.mol, - self.sorted_grids, - self.feature_function, - self.block_size, - self.compile_feature_function, - self.gpu, - dm_tangent, - ), - ) - return sorted_tangent.index_select(-1, self.inverse_permutation).detach() -def prepare_screened_feature_buffer( - mol: gto.Mole, +def screened_feature_jvp( dm: Tensor, - grids: Grid, - features: FeatureSpec | set[str], + dm_tangent: Tensor, + mol: gto.Mole, + spatial_grid_layout: SpatialGridLayout, + feature_function: feature_math.MGGAFeatureFunction, compile_feature_function: bool = False, -) -> ScreenedFeatureBuffer: - """Prepare the backend grid extension and its screened feature buffer. - - The source grid receives a cached :class:`SpatialGridLayout`; its coordinate and - weight arrays remain atom-major. Feature evaluation runs once on a spatially - reordered grid copy, and the resulting buffer restores atom-major order for model - chunk construction. - """ - feature_spec = ( - features if isinstance(features, FeatureSpec) else FeatureSpec(features) - ) - if grids.coords is None or grids.weights is None: - raise ValueError("Grids must be built before generating screened features.") - - feature_function = feature_math.MGGAFeatureFunction(feature_spec) - gpu = dm.device.type == "cuda" - if gpu: - check_gpu_imports_were_successful() - from gpu4pyscf.dft import numint as dft_gpu_numint - - block_size = int(dft_gpu_numint.MIN_BLK_SIZE) - else: - block_size = CPU_AO_SCREENING_BLOCK_SIZE - sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( - mol, grids, block_size, gpu - ) - sorted_raw_features = cast( +) -> Tensor: + """Apply the raw-feature Jacobian and restore atom-major grid order.""" + sorted_tangent = cast( Tensor, ao_evaluation.ChunkEvalForward.apply( - dm.double(), + dm, mol, - sorted_grids, + spatial_grid_layout.sorted_grids, feature_function, - block_size, + spatial_grid_layout.block_size, compile_feature_function, - gpu, + dm.device.type == "cuda", + dm_tangent, ), ) - forward_permutation = torch.as_tensor(forward, device=dm.device) - inverse_permutation = torch.as_tensor(inverse, device=dm.device) - atom_major_raw_features = sorted_raw_features.index_select(-1, inverse_permutation) - return ScreenedFeatureBuffer( - dm=dm, - mol=mol, - sorted_grids=sorted_grids, - sorted_raw_features=sorted_raw_features, - atom_major_raw_features=atom_major_raw_features, - forward_permutation=forward_permutation, - inverse_permutation=inverse_permutation, - feature_function=feature_function, - block_size=block_size, - compile_feature_function=compile_feature_function, - gpu=gpu, - ) + return sorted_tangent.index_select( + -1, spatial_grid_layout.inverse_permutation + ).detach() diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index ae653d8b..513f37e6 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -23,8 +23,9 @@ from skala.pyscf.model_chunking import ModelFeatureChunk from skala.pyscf.numint import SkalaNumInt, _should_screen_aos from skala.pyscf.screening import ( + SpatialGridLayout, _decompose_grid_into_spatial_blocks, - _prepare_spatially_sorted_grids, + prepare_spatial_grid_layout, ) @@ -211,14 +212,15 @@ def test_prepare_spatially_sorted_cpu_grids( inverse = np.argsort(forward) non0tab = np.ones((1, carbon.nbas), dtype=np.uint8) partition_calls = 0 + decomposition_block_sizes: list[int] = [] def fake_decompose_grid_into_spatial_blocks( coords_arg: np.ndarray, block_size: int ) -> tuple[np.ndarray, np.ndarray]: nonlocal partition_calls partition_calls += 1 + decomposition_block_sizes.append(block_size) assert coords_arg is grids.coords - assert block_size == 2 return forward, inverse monkeypatch.setattr( @@ -228,22 +230,23 @@ def fake_decompose_grid_into_spatial_blocks( ) screen_index_calls = 0 + screened_molecules: list[gto.Mole] = [] def fake_make_screen_index( mol_arg: gto.Mole, sorted_coords: np.ndarray, cutoff: float ) -> np.ndarray: nonlocal screen_index_calls screen_index_calls += 1 - assert mol_arg is carbon + screened_molecules.append(mol_arg) assert np.array_equal(sorted_coords, coords[forward]) assert cutoff == grids.cutoff return non0tab monkeypatch.setattr(dft.gen_grid, "make_screen_index", fake_make_screen_index) - sorted_grids, actual_forward, actual_inverse = _prepare_spatially_sorted_grids( - carbon, grids, block_size=2, gpu=False - ) + device = torch.device("cpu") + layout = prepare_spatial_grid_layout(carbon, grids, block_size=2, device=device) + sorted_grids = layout.sorted_grids assert sorted_grids is not grids assert np.array_equal(grids.coords, coords) @@ -251,27 +254,17 @@ def fake_make_screen_index( assert np.array_equal(sorted_grids.coords, coords[forward]) assert np.array_equal(sorted_grids.weights, weights[forward]) assert sorted_grids.non0tab is non0tab - assert actual_forward is forward - assert actual_inverse is inverse - - cached_grids, cached_forward, cached_inverse = _prepare_spatially_sorted_grids( - carbon, grids, block_size=2, gpu=False + torch.testing.assert_close( + layout.forward_permutation, torch.as_tensor(forward, device=device) + ) + torch.testing.assert_close( + layout.inverse_permutation, torch.as_tensor(inverse, device=device) ) - - assert cached_grids is sorted_grids - assert cached_forward is forward - assert cached_inverse is inverse assert partition_calls == 1 assert screen_index_calls == 1 - - grids.coords = grids.coords.copy() - rebuilt_grids, _, _ = _prepare_spatially_sorted_grids( - carbon, grids, block_size=2, gpu=False - ) - - assert rebuilt_grids is not sorted_grids - assert partition_calls == 2 - assert screen_index_calls == 2 + assert decomposition_block_sizes == [2] + assert screened_molecules == [carbon] + assert not hasattr(grids, "_skala_spatial_grid_layout") class QuadraticDensityFunctional(ExcFunctionalBase): @@ -283,6 +276,54 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: return (mol["density"].square() * mol["grid_weights"]).sum() +def test_grid_reuses_spatial_grid_layout_across_numints( + carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch +) -> None: + grids = dft.Grids(carbon) + grids.coords = np.arange(18, dtype=np.float64).reshape(6, 3) + grids.weights = np.arange(6, dtype=np.float64) + other_grids = dft.Grids(carbon) + other_grids.coords = grids.coords.copy() + other_grids.weights = grids.weights.copy() + layouts: list[SpatialGridLayout] = [] + + def fake_prepare_spatial_grid_layout( + mol: gto.Mole, + grids: object, + block_size: int, + device: torch.device, + ) -> SpatialGridLayout: + layout = SpatialGridLayout( + block_size=block_size, + sorted_grids=grids, + forward_permutation=torch.arange(6, device=device), + inverse_permutation=torch.arange(6, device=device), + ) + layouts.append(layout) + return layout + + monkeypatch.setattr( + numint_module, + "prepare_spatial_grid_layout", + fake_prepare_spatial_grid_layout, + ) + numint = SkalaNumInt(QuadraticDensityFunctional()) + other_numint = SkalaNumInt(QuadraticDensityFunctional()) + + layout = numint._get_spatial_grid_layout(carbon, grids) + assert other_numint._get_spatial_grid_layout(carbon, grids) is layout + assert vars(grids)["_skala_spatial_grid_layout"] is layout + assert len(layouts) == 1 + + numint.reset() + assert numint._get_spatial_grid_layout(carbon, grids) is layout + + other_layout = numint._get_spatial_grid_layout(carbon, other_grids) + assert other_layout is not layout + assert vars(other_grids)["_skala_spatial_grid_layout"] is other_layout + assert len(layouts) == 2 + + class FakeKS: def __init__(self, mol: gto.Mole, grids: object | None = None) -> None: self.mol = mol @@ -296,6 +337,19 @@ def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: return np.zeros_like(dm) +def test_call_rejects_second_order_evaluation(carbon: gto.Mole) -> None: + numint = SkalaNumInt(QuadraticDensityFunctional()) + + with pytest.raises(NotImplementedError, match="second-order evaluation"): + numint( + carbon, + dft.Grids(carbon), + None, + torch.eye(carbon.nao_nr(), dtype=torch.float64), + second_order=True, + ) + + @pytest.mark.parametrize("expected", [False, True]) @pytest.mark.parametrize("response_safety_fraction", [None, 0.6]) def test_first_and_second_order_use_same_screening_decision( @@ -324,16 +378,13 @@ def fake_generate_features( "grid_weights": torch.ones(1, dtype=dm.dtype), } - class FakeScreenedFeatureBuffer: - def __init__(self, dm: torch.Tensor) -> None: - raw_features = dm.sum().reshape(1, 1) - self.feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) - self.sorted_raw_features = raw_features - self.atom_major_raw_features = raw_features - self.forward_permutation = torch.tensor([0]) + class FakeSpatialGridLayout: + block_size = 1 + forward_permutation = torch.tensor([0]) + inverse_permutation = torch.tensor([0]) - def atom_major_jvp(self, dm_tangent: torch.Tensor) -> torch.Tensor: - return dm_tangent.sum().reshape(1, 1) + def __init__(self, sorted_grids: object) -> None: + self.sorted_grids = sorted_grids class FakeModelFeatureChunks: def __init__(self, raw_features: torch.Tensor) -> None: @@ -351,15 +402,29 @@ def __iter__(self) -> Iterator[ModelFeatureChunk]: }, ) - def fake_prepare_screened_feature_buffer( + def fake_prepare_spatial_grid_layout( mol: gto.Mole, - dm: torch.Tensor, grids: object, - features: set[str], - **kwargs: object, - ) -> FakeScreenedFeatureBuffer: + block_size: int, + device: torch.device, + ) -> FakeSpatialGridLayout: + return FakeSpatialGridLayout(grids) + + def fake_chunk_eval_forward( + dm: torch.Tensor, + *args: object, + ) -> torch.Tensor: routes.append("screened") - return FakeScreenedFeatureBuffer(dm) + return dm.sum().reshape(1, 1) + + def fake_screened_feature_jvp( + dm: torch.Tensor, + dm_tangent: torch.Tensor, + mol: gto.Mole, + spatial_grid_layout: object, + feature_function: MGGAFeatureFunction, + ) -> torch.Tensor: + return dm_tangent.sum().reshape(1, 1) def fake_prepare_model_feature_chunks( mol: gto.Mole, @@ -378,14 +443,24 @@ def fake_prepare_model_feature_chunks( monkeypatch.setattr(numint_module, "generate_features", fake_generate_features) monkeypatch.setattr( numint_module, - "prepare_screened_feature_buffer", - fake_prepare_screened_feature_buffer, + "prepare_spatial_grid_layout", + fake_prepare_spatial_grid_layout, + ) + monkeypatch.setattr( + ChunkEvalForward, + "apply", + staticmethod(fake_chunk_eval_forward), ) monkeypatch.setattr( numint_module, "prepare_model_feature_chunks", fake_prepare_model_feature_chunks, ) + monkeypatch.setattr( + numint_module, + "screened_feature_jvp", + fake_screened_feature_jvp, + ) numint = SkalaNumInt(QuadraticDensityFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) grids = dft.Grids(carbon) @@ -591,6 +666,21 @@ def _minimal_atom_grid(mol: gto.Mole) -> dft.Grids: return grids.build(sort_grids=False) +def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> None: + numint = SkalaNumInt(QuadraticDensityFunctional()) + grids = _minimal_atom_grid(carbon) + spatial_grid_layout = prepare_spatial_grid_layout( + carbon, + grids, + block_size=dft.gen_grid.BLKSIZE, + device=torch.device("cpu"), + ) + vars(grids)["_skala_spatial_grid_layout"] = spatial_grid_layout + + assert numint.reset() is numint + assert vars(grids)["_skala_spatial_grid_layout"] is spatial_grid_layout + + @pytest.mark.parametrize("unrestricted", [False, True]) def test_cpu_rks_uks_dense_screened_equivalence( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 3f0ac566..b0e06ac9 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -34,7 +34,7 @@ from skala.pyscf.evaluation import FeatureSpec # noqa: E402 from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 from skala.pyscf.numint import SkalaNumInt # noqa: E402 -from skala.pyscf.screening import _prepare_spatially_sorted_grids # noqa: E402 +from skala.pyscf.screening import prepare_spatial_grid_layout # noqa: E402 CARBON_CHAIN = """ C 0.0 0.0 0.0 @@ -89,9 +89,11 @@ def test_prepare_spatially_sorted_gpu_grids() -> None: original_screening_cache = cupy.arange(1) grids._non0ao_idx = original_screening_cache - sorted_grids, forward, inverse = _prepare_spatially_sorted_grids( - mol, grids, block_size=2, gpu=True - ) + device = torch.device("cuda") + layout = prepare_spatial_grid_layout(mol, grids, block_size=2, device=device) + sorted_grids = layout.sorted_grids + forward = layout.forward_permutation.cpu().numpy() + inverse = layout.inverse_permutation.cpu().numpy() assert sorted_grids is not grids assert grids.coords is coords @@ -109,17 +111,8 @@ def test_prepare_spatially_sorted_gpu_grids() -> None: assert np.array_equal( cupy.asnumpy(sorted_grids.coords)[inverse], cupy.asnumpy(coords) ) - - sorted_screening_cache = object() - sorted_grids._non0ao_idx = sorted_screening_cache - cached_grids, cached_forward, cached_inverse = _prepare_spatially_sorted_grids( - mol, grids, block_size=2, gpu=True - ) - - assert cached_grids is sorted_grids - assert cached_forward is forward - assert cached_inverse is inverse - assert cached_grids._non0ao_idx is sorted_screening_cache + assert layout.forward_permutation.device.type == "cuda" + assert layout.inverse_permutation.device.type == "cuda" @pytest.mark.parametrize("unrestricted", [False, True]) From 60385d500ccc42a1cad2399722013ea4cbc3d289 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 16:28:21 +0200 Subject: [PATCH 19/39] refactor separate xc integration from pyscf adapter --- src/skala/pyscf/numint.py | 355 ++++---------------------- src/skala/pyscf/xc_integrator.py | 363 +++++++++++++++++++++++++++ tests/test_ao_screening.py | 25 +- tests/test_ao_screening_benchmark.py | 3 +- tests/test_gpu4pyscf_gradients.py | 2 +- tests/test_xc_integrator.py | 107 ++++++++ 6 files changed, 530 insertions(+), 325 deletions(-) create mode 100644 src/skala/pyscf/xc_integrator.py create mode 100644 tests/test_xc_integrator.py diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 8cb12cde..38e80143 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -1,49 +1,23 @@ # SPDX-License-Identifier: MIT from collections.abc import Callable -from typing import Any, Generic, Protocol, cast, overload +from typing import Any, Generic, Protocol, overload import torch from pyscf import gto -from pyscf.dft import numint as pyscf_numint from torch import Tensor from skala.functional.base import ExcFunctionalBase -from skala.pyscf import ao_evaluation, feature_math from skala.pyscf.backend import ( KS, Array, Grid, - check_gpu_imports_were_successful, from_numpy_or_cupy, to_cupy, to_numpy, ) from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec -from skala.pyscf.features import generate_features -from skala.pyscf.model_chunking import prepare_model_feature_chunks -from skala.pyscf.screening import ( - CPU_AO_SCREENING_BLOCK_SIZE, - SpatialGridLayout, - prepare_spatial_grid_layout, - screened_feature_jvp, -) - - -def _should_screen_aos(mol: gto.Mole) -> bool: - """Determine whether a molecule is large enough for AO screening. - - Args: - mol: Molecule whose AO count is compared with PySCF's sparse-contraction - crossover. - - Returns: - Whether the molecule has more AOs than PySCF's screening threshold. - """ - # Keep the compatibility fallback here, not at call sites. PySCF uses this - # crossover before selecting sparse density/Vxc contractions. - switch_size = pyscf_numint.SWITCH_SIZE - return mol.nao_nr() > switch_size +from skala.pyscf.xc_integrator import XCIntegrator class LibXCSpec(Protocol): @@ -133,57 +107,38 @@ class SkalaNumInt(PySCFNumInt[Array]): -1.1425799... """ - device: torch.device - def __init__( self, functional: ExcFunctionalBase, chunk_size: int | None = None, device: torch.device | None = None, ): - self.device = device or torch.get_default_device() + self.integrator = XCIntegrator(functional, chunk_size=chunk_size, device=device) - if self.device.type == "cuda": - check_gpu_imports_were_successful() + @property + def device(self) -> torch.device: + """Torch device used by the XC integrator.""" + return self.integrator.device + + @property + def func(self) -> ExcFunctionalBase: + """Functional retained for gradient-adapter compatibility.""" + return self.integrator.functional - self.func = functional.to(device=self.device) - self.feature_spec = FeatureSpec(self.func.features) - self.evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) + @property + def feature_spec(self) -> FeatureSpec: + """Feature requirements owned by the XC integrator.""" + return self.integrator.feature_spec + + @property + def evaluation_policy(self) -> EvaluationPolicy: + """Numerical policy owned by the XC integrator.""" + return self.integrator.evaluation_policy def reset(self) -> "SkalaNumInt[Array]": """Return this integrator; spatial layouts are owned by grid objects.""" return self - def _get_spatial_grid_layout( - self, - mol: gto.Mole, - grids: Grid, - ) -> SpatialGridLayout: - grid_state = vars(grids) - spatial_grid_layout = cast( - SpatialGridLayout | None, - grid_state.get("_skala_spatial_grid_layout"), - ) - if spatial_grid_layout is not None: - return spatial_grid_layout - - if self.device.type == "cuda": - check_gpu_imports_were_successful() - from gpu4pyscf.dft import numint as dft_gpu_numint - - block_size = int(dft_gpu_numint.MIN_BLK_SIZE) - else: - block_size = CPU_AO_SCREENING_BLOCK_SIZE - - spatial_grid_layout = prepare_spatial_grid_layout( - mol, - grids, - block_size, - self.device, - ) - grid_state["_skala_spatial_grid_layout"] = spatial_grid_layout - return spatial_grid_layout - def from_backend( self, x: Array, @@ -194,10 +149,8 @@ def from_backend( @overload def to_backend(self, x: Tensor) -> Array: ... - @overload def to_backend(self, x: list[Tensor]) -> list[Array]: ... - def to_backend(self, x: Tensor | list[Tensor]) -> Array | list[Array]: if isinstance(x, list): return [self.to_backend(y) for y in x] @@ -215,16 +168,13 @@ def get_rho( max_memory: int = 2000, verbose: int = 0, ) -> Array: - mol_features = generate_features( + density = self.integrator.density( mol, self.from_backend(dm), grids, - features={"density"}, - chunk_size=self.evaluation_policy.ao_block_size, max_memory=max_memory, - gpu=self.device.type == "cuda", ) - return self.to_backend(mol_features["density"].sum(0)) + return self.to_backend(density) def __call__( self, @@ -235,125 +185,26 @@ def __call__( second_order: bool = False, max_memory: int = 2000, ) -> tuple[Tensor, Tensor, Tensor]: - """Evaluate the XC functional for a molecule and density matrix.""" + """ + Evaluate the XC functional for the given molecule and density matrix. + Input: + mol: The molecule. + grids: The grid. + xc_code: The XC code (not used in the reimplementation). + dm: The density matrix. + second_order: Unsupported; use ``gen_response`` for response evaluation. + max_memory: The maximum memory to use for each chunk in megabytes (MB). + + Returns: + A tuple of the total integrated density, the XC energy, and the XC potential. + """ if second_order: raise NotImplementedError( "Direct second-order evaluation is not supported; use gen_response()." ) - if self.device != dm.device: - raise ValueError( - f"Density matrix device {dm.device} does not match functional device {self.device}" - ) - - if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): - return self._call_screened(mol, grids, dm, max_memory) - return self._call_dense(mol, grids, dm, max_memory) - - def _call_screened( - self, - mol: gto.Mole, - grids: Grid, - dm: Tensor, - max_memory: int, - ) -> tuple[Tensor, Tensor, Tensor]: - dm = dm.detach().requires_grad_() - tot_dens = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) - E_xc = torch.tensor(0.0, device=self.device, dtype=dm.dtype) - feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) - spatial_grid_layout = self._get_spatial_grid_layout(mol, grids) - sorted_raw_features = cast( - Tensor, - ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm.double(), - mol, - spatial_grid_layout.sorted_grids, - feature_function, - spatial_grid_layout.block_size, - False, - dm.device.type == "cuda", - ), - ) - atom_major_raw_features = sorted_raw_features.index_select( - -1, spatial_grid_layout.inverse_permutation - ) - model_chunks = prepare_model_feature_chunks( - mol, - dm, - grids, - atom_major_raw_features=atom_major_raw_features, - feature_function=feature_function, - func_deriv=1, - max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, - safety_fraction=self.evaluation_policy.safety_fraction, - ) - # Store only full-grid feature cotangents; model activations remain chunk-local. - atom_major_cotangent = torch.zeros_like(atom_major_raw_features) - for chunk in model_chunks: - local_raw_features = chunk.raw_features - mol_features = chunk.model_features - E_xc_chunk = self.func.get_exc(mol_features) - (local_cotangent,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - ) - atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() - tot_dens += ( - (mol_features["density"] * mol_features["grid_weights"]) - .sum(dim=-1) - .detach() - ) - E_xc += E_xc_chunk.detach() - del E_xc_chunk, local_cotangent, local_raw_features, mol_features - - # Reorder detached cotangents explicitly instead of backpropagating through it. - sorted_cotangent = atom_major_cotangent.index_select( - -1, spatial_grid_layout.forward_permutation - ) - # The custom VJP reevaluates AO blocks sequentially without a full-grid AO graph. - (V_xc,) = torch.autograd.grad( - sorted_raw_features, - dm, - sorted_cotangent, - ) - return tot_dens, E_xc, V_xc - - def _call_dense( - self, - mol: gto.Mole, - grids: Grid, - dm: Tensor, - max_memory: int, - *, - create_graph: bool = False, - ) -> tuple[Tensor, Tensor, Tensor]: - - dm = dm.requires_grad_() - mol_features = generate_features( - mol, - dm, - grids, - set(self.feature_spec.names), - chunk_size=self.evaluation_policy.ao_block_size, - max_memory=max_memory, - gpu=self.device.type == "cuda", - ) - E_xc = self.func.get_exc(mol_features) - (V_xc,) = torch.autograd.grad( - E_xc, - dm, - torch.ones_like(E_xc), - retain_graph=create_graph, - create_graph=create_graph, - ) - - rho = mol_features["density"] - grid_weights = mol_features.get( - "grid_weights", self.from_backend(grids.weights) - ) - tot_dens = (rho * grid_weights).sum(dim=-1) - return tot_dens, E_xc, V_xc + result = self.integrator(mol, grids, dm, max_memory=max_memory) + return result.electron_count, result.energy, result.potential def nr_rks( self, @@ -406,6 +257,7 @@ def gen_response( ks: KS, **kwargs: Any, ) -> Callable[[Array], Array]: + """Generates the response function for the functional.""" assert mo_coeff is not None assert mo_occ is not None if kwargs is not None: @@ -419,137 +271,16 @@ def gen_response( assert kwargs["with_j"] dm0 = self.from_backend(ks.make_rdm1(mo_coeff, mo_occ)) - - if self.feature_spec.supports_screened_evaluation and _should_screen_aos( - ks.mol - ): - return self._gen_response_screened( - ks, - dm0, - safety_fraction=kwargs.get( - "safety_fraction", self.evaluation_policy.safety_fraction - ), - ) - return self._gen_response_dense(ks, dm0) - - def _gen_response_screened( - self, - ks: KS, - dm0: Tensor, - *, - safety_fraction: float, - ) -> Callable[[Array], Array]: - dm0 = dm0.requires_grad_() - feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) - spatial_grid_layout = self._get_spatial_grid_layout(ks.mol, ks.grids) - sorted_raw_features = cast( - Tensor, - ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm0.double(), - ks.mol, - spatial_grid_layout.sorted_grids, - feature_function, - spatial_grid_layout.block_size, - False, - dm0.device.type == "cuda", - ), - ) - atom_major_raw_features = sorted_raw_features.index_select( - -1, spatial_grid_layout.inverse_permutation - ) - model_chunks = prepare_model_feature_chunks( - ks.mol, - dm0, - ks.grids, - atom_major_raw_features=atom_major_raw_features, - feature_function=feature_function, - func_deriv=2, - max_memory_in_mb=ks.max_memory if dm0.device.type == "cpu" else None, - safety_fraction=safety_fraction, - ) - - def hessian_vector_product_atom_chunked(dm1: Array) -> Array: - dm1_tensor = self.from_backend(dm1) - atom_major_tangent = screened_feature_jvp( - dm0, - dm1_tensor, - ks.mol, - spatial_grid_layout, - feature_function, - ) - # Store the full-grid model Hessian action, not per-chunk model graphs. - atom_major_hessian_action = torch.zeros_like(atom_major_raw_features) - for chunk in model_chunks: - local_raw_features = chunk.raw_features - mol_features = chunk.model_features - E_xc_chunk = self.func.get_exc(mol_features) - (local_gradient,) = torch.autograd.grad( - E_xc_chunk, - local_raw_features, - torch.ones_like(E_xc_chunk), - create_graph=True, - ) - if local_gradient.requires_grad: - (local_hessian_action,) = torch.autograd.grad( - local_gradient, - local_raw_features, - atom_major_tangent[..., chunk.grid_slice], - ) - else: - local_hessian_action = torch.zeros_like(local_raw_features) - atom_major_hessian_action[..., chunk.grid_slice] = ( - local_hessian_action.detach() - ) - del ( - E_xc_chunk, - local_gradient, - local_hessian_action, - local_raw_features, - mol_features, - ) - - # Restore block order after all chunk-local Hessian actions are detached. - sorted_hessian_action = atom_major_hessian_action.index_select( - -1, spatial_grid_layout.forward_permutation - ) - # The custom VJP traverses AO blocks sequentially and retains no AO graph. - (hvp_total,) = torch.autograd.grad( - sorted_raw_features, - dm0, - sorted_hessian_action, - retain_graph=True, - ) - - v1 = self.to_backend(hvp_total) - vj = ks.get_j(ks.mol, dm1, hermi=1) - if ks.mol.spin == 0: - v1 += vj - else: - v1 += vj[0] + vj[1] - return v1 - - return hessian_vector_product_atom_chunked - - def _gen_response_dense( - self, - ks: KS, - dm0: Tensor, - ) -> Callable[[Array], Array]: - dm0 = dm0.requires_grad_() - _, _, V_xc = self._call_dense( + xc_response = self.integrator.gen_response( ks.mol, ks.grids, dm0, - ks.max_memory, - create_graph=True, + max_memory=ks.max_memory, + safety_fraction=kwargs.get("safety_fraction"), ) def hessian_vector_product(dm1: Array) -> Array: - v1 = self.to_backend( - torch.autograd.grad( - V_xc, dm0, self.from_backend(dm1), retain_graph=True - )[0] - ) + v1 = self.to_backend(xc_response(self.from_backend(dm1))) vj = ks.get_j(ks.mol, dm1, hermi=1) if ks.mol.spin == 0: diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py new file mode 100644 index 00000000..b2bdeea5 --- /dev/null +++ b/src/skala/pyscf/xc_integrator.py @@ -0,0 +1,363 @@ +# SPDX-License-Identifier: MIT + +"""Tensor-level exchange-correlation integration.""" + +from collections.abc import Callable +from dataclasses import dataclass +from typing import cast + +import torch +from pyscf import gto +from pyscf.dft import numint as pyscf_numint +from torch import Tensor + +from skala.functional.base import ExcFunctionalBase +from skala.pyscf import ao_evaluation, feature_math +from skala.pyscf.backend import Grid, check_gpu_imports_were_successful +from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec +from skala.pyscf.features import generate_features +from skala.pyscf.model_chunking import prepare_model_feature_chunks +from skala.pyscf.screening import ( + CPU_AO_SCREENING_BLOCK_SIZE, + SpatialGridLayout, + prepare_spatial_grid_layout, + screened_feature_jvp, +) + + +def _should_screen_aos(mol: gto.Mole) -> bool: + """Return whether PySCF's sparse-contraction crossover is exceeded.""" + return mol.nao_nr() > pyscf_numint.SWITCH_SIZE + + +@dataclass(frozen=True) +class XCResult: + """Tensor-valued result of exchange-correlation integration.""" + + electron_count: Tensor + energy: Tensor + potential: Tensor + + +class XCIntegrator: + """Evaluate XC energies, potentials, and potential responses in Torch.""" + + def __init__( + self, + functional: ExcFunctionalBase, + chunk_size: int | None = None, + device: torch.device | None = None, + ) -> None: + self.device = device or torch.get_default_device() + if self.device.type == "cuda": + check_gpu_imports_were_successful() + + self.functional = functional.to(device=self.device) + self.feature_spec = FeatureSpec(self.functional.features) + self.evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) + + def density( + self, + mol: gto.Mole, + dm: Tensor, + grids: Grid, + max_memory: int = 2000, + ) -> Tensor: + """Evaluate the total density on each grid point.""" + mol_features = generate_features( + mol, + dm, + grids, + features={"density"}, + chunk_size=self.evaluation_policy.ao_block_size, + max_memory=max_memory, + gpu=self.device.type == "cuda", + ) + return mol_features["density"].sum(0) + + def __call__( + self, + mol: gto.Mole, + grids: Grid, + dm: Tensor, + max_memory: int = 2000, + ) -> XCResult: + """Evaluate electron count, XC energy, and XC potential.""" + self._validate_device(dm) + if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): + return self._integrate_screened(mol, grids, dm, max_memory) + return self._integrate_dense(mol, grids, dm, max_memory) + + def gen_response( + self, + mol: gto.Mole, + grids: Grid, + dm0: Tensor, + max_memory: int = 2000, + safety_fraction: float | None = None, + ) -> Callable[[Tensor], Tensor]: + """Build an XC-only Hessian-vector product callable.""" + self._validate_device(dm0) + if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): + return self._gen_response_screened( + mol, + grids, + dm0, + max_memory=max_memory, + safety_fraction=( + self.evaluation_policy.safety_fraction + if safety_fraction is None + else safety_fraction + ), + ) + return self._gen_response_dense(mol, grids, dm0, max_memory=max_memory) + + def _validate_device(self, dm: Tensor) -> None: + if self.device != dm.device: + raise ValueError( + f"Density matrix device {dm.device} does not match functional device {self.device}" + ) + + def _get_spatial_grid_layout( + self, + mol: gto.Mole, + grids: Grid, + ) -> SpatialGridLayout: + grid_state = vars(grids) + spatial_grid_layout = cast( + SpatialGridLayout | None, + grid_state.get("_skala_spatial_grid_layout"), + ) + if spatial_grid_layout is not None: + return spatial_grid_layout + + if self.device.type == "cuda": + check_gpu_imports_were_successful() + from gpu4pyscf.dft import numint as dft_gpu_numint + + block_size = int(dft_gpu_numint.MIN_BLK_SIZE) + else: + block_size = CPU_AO_SCREENING_BLOCK_SIZE + + spatial_grid_layout = prepare_spatial_grid_layout( + mol, + grids, + block_size, + self.device, + ) + grid_state["_skala_spatial_grid_layout"] = spatial_grid_layout + return spatial_grid_layout + + def _integrate_screened( + self, + mol: gto.Mole, + grids: Grid, + dm: Tensor, + max_memory: int, + ) -> XCResult: + dm = dm.detach().requires_grad_() + electron_count = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) + energy = torch.tensor(0.0, device=self.device, dtype=dm.dtype) + feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) + spatial_grid_layout = self._get_spatial_grid_layout(mol, grids) + sorted_raw_features = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm.double(), + mol, + spatial_grid_layout.sorted_grids, + feature_function, + spatial_grid_layout.block_size, + False, + dm.device.type == "cuda", + ), + ) + atom_major_raw_features = sorted_raw_features.index_select( + -1, spatial_grid_layout.inverse_permutation + ) + model_chunks = prepare_model_feature_chunks( + mol, + dm, + grids, + atom_major_raw_features=atom_major_raw_features, + feature_function=feature_function, + func_deriv=1, + max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, + safety_fraction=self.evaluation_policy.safety_fraction, + ) + atom_major_cotangent = torch.zeros_like(atom_major_raw_features) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features + energy_chunk = self.functional.get_exc(mol_features) + (local_cotangent,) = torch.autograd.grad( + energy_chunk, + local_raw_features, + torch.ones_like(energy_chunk), + ) + atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() + electron_count += ( + (mol_features["density"] * mol_features["grid_weights"]) + .sum(dim=-1) + .detach() + ) + energy += energy_chunk.detach() + del energy_chunk, local_cotangent, local_raw_features, mol_features + + sorted_cotangent = atom_major_cotangent.index_select( + -1, spatial_grid_layout.forward_permutation + ) + (potential,) = torch.autograd.grad( + sorted_raw_features, + dm, + sorted_cotangent, + ) + return XCResult(electron_count, energy, potential) + + def _integrate_dense( + self, + mol: gto.Mole, + grids: Grid, + dm: Tensor, + max_memory: int, + *, + create_graph: bool = False, + ) -> XCResult: + dm = dm.requires_grad_() + mol_features = generate_features( + mol, + dm, + grids, + set(self.feature_spec.names) | {"density", "grid_weights"}, + chunk_size=self.evaluation_policy.ao_block_size, + max_memory=max_memory, + gpu=self.device.type == "cuda", + ) + energy = self.functional.get_exc(mol_features) + (potential,) = torch.autograd.grad( + energy, + dm, + torch.ones_like(energy), + retain_graph=create_graph, + create_graph=create_graph, + ) + electron_count = (mol_features["density"] * mol_features["grid_weights"]).sum( + dim=-1 + ) + return XCResult(electron_count, energy, potential) + + def _gen_response_screened( + self, + mol: gto.Mole, + grids: Grid, + dm0: Tensor, + *, + max_memory: int, + safety_fraction: float, + ) -> Callable[[Tensor], Tensor]: + dm0 = dm0.requires_grad_() + feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) + spatial_grid_layout = self._get_spatial_grid_layout(mol, grids) + sorted_raw_features = cast( + Tensor, + ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] + dm0.double(), + mol, + spatial_grid_layout.sorted_grids, + feature_function, + spatial_grid_layout.block_size, + False, + dm0.device.type == "cuda", + ), + ) + atom_major_raw_features = sorted_raw_features.index_select( + -1, spatial_grid_layout.inverse_permutation + ) + model_chunks = prepare_model_feature_chunks( + mol, + dm0, + grids, + atom_major_raw_features=atom_major_raw_features, + feature_function=feature_function, + func_deriv=2, + max_memory_in_mb=max_memory if dm0.device.type == "cpu" else None, + safety_fraction=safety_fraction, + ) + + def hessian_vector_product(dm1: Tensor) -> Tensor: + atom_major_tangent = screened_feature_jvp( + dm0, + dm1, + mol, + spatial_grid_layout, + feature_function, + ) + atom_major_hessian_action = torch.zeros_like(atom_major_raw_features) + for chunk in model_chunks: + local_raw_features = chunk.raw_features + mol_features = chunk.model_features + energy_chunk = self.functional.get_exc(mol_features) + (local_gradient,) = torch.autograd.grad( + energy_chunk, + local_raw_features, + torch.ones_like(energy_chunk), + create_graph=True, + ) + if local_gradient.requires_grad: + (local_hessian_action,) = torch.autograd.grad( + local_gradient, + local_raw_features, + atom_major_tangent[..., chunk.grid_slice], + ) + else: + local_hessian_action = torch.zeros_like(local_raw_features) + atom_major_hessian_action[..., chunk.grid_slice] = ( + local_hessian_action.detach() + ) + del ( + energy_chunk, + local_gradient, + local_hessian_action, + local_raw_features, + mol_features, + ) + + sorted_hessian_action = atom_major_hessian_action.index_select( + -1, spatial_grid_layout.forward_permutation + ) + (hvp_total,) = torch.autograd.grad( + sorted_raw_features, + dm0, + sorted_hessian_action, + retain_graph=True, + ) + return hvp_total + + return hessian_vector_product + + def _gen_response_dense( + self, + mol: gto.Mole, + grids: Grid, + dm0: Tensor, + *, + max_memory: int, + ) -> Callable[[Tensor], Tensor]: + dm0 = dm0.requires_grad_() + potential = self._integrate_dense( + mol, + grids, + dm0, + max_memory, + create_graph=True, + ).potential + + def hessian_vector_product(dm1: Tensor) -> Tensor: + return torch.autograd.grad( + potential, + dm0, + dm1, + retain_graph=True, + )[0] + + return hessian_vector_product diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 513f37e6..f71faa11 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -8,8 +8,8 @@ from skala.functional.base import ExcFunctionalBase from skala.pyscf import model_chunking as model_chunking_module -from skala.pyscf import numint as numint_module from skala.pyscf import screening as screening_module +from skala.pyscf import xc_integrator as xc_integrator_module from skala.pyscf.ao_evaluation import ( ChunkEvalBackward, ChunkEvalForward, @@ -21,12 +21,13 @@ from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.feature_math import MGGAFeatureFunction from skala.pyscf.model_chunking import ModelFeatureChunk -from skala.pyscf.numint import SkalaNumInt, _should_screen_aos +from skala.pyscf.numint import SkalaNumInt from skala.pyscf.screening import ( SpatialGridLayout, _decompose_grid_into_spatial_blocks, prepare_spatial_grid_layout, ) +from skala.pyscf.xc_integrator import _should_screen_aos @pytest.fixture @@ -303,22 +304,22 @@ def fake_prepare_spatial_grid_layout( return layout monkeypatch.setattr( - numint_module, + xc_integrator_module, "prepare_spatial_grid_layout", fake_prepare_spatial_grid_layout, ) numint = SkalaNumInt(QuadraticDensityFunctional()) other_numint = SkalaNumInt(QuadraticDensityFunctional()) - layout = numint._get_spatial_grid_layout(carbon, grids) - assert other_numint._get_spatial_grid_layout(carbon, grids) is layout + layout = numint.integrator._get_spatial_grid_layout(carbon, grids) + assert other_numint.integrator._get_spatial_grid_layout(carbon, grids) is layout assert vars(grids)["_skala_spatial_grid_layout"] is layout assert len(layouts) == 1 numint.reset() - assert numint._get_spatial_grid_layout(carbon, grids) is layout + assert numint.integrator._get_spatial_grid_layout(carbon, grids) is layout - other_layout = numint._get_spatial_grid_layout(carbon, other_grids) + other_layout = numint.integrator._get_spatial_grid_layout(carbon, other_grids) assert other_layout is not layout assert vars(other_grids)["_skala_spatial_grid_layout"] is other_layout assert len(layouts) == 2 @@ -440,9 +441,11 @@ def fake_prepare_model_feature_chunks( safety_fractions.append(safety_fraction) return FakeModelFeatureChunks(atom_major_raw_features) - monkeypatch.setattr(numint_module, "generate_features", fake_generate_features) monkeypatch.setattr( - numint_module, + xc_integrator_module, "generate_features", fake_generate_features + ) + monkeypatch.setattr( + xc_integrator_module, "prepare_spatial_grid_layout", fake_prepare_spatial_grid_layout, ) @@ -452,12 +455,12 @@ def fake_prepare_model_feature_chunks( staticmethod(fake_chunk_eval_forward), ) monkeypatch.setattr( - numint_module, + xc_integrator_module, "prepare_model_feature_chunks", fake_prepare_model_feature_chunks, ) monkeypatch.setattr( - numint_module, + xc_integrator_module, "screened_feature_jvp", fake_screened_feature_jvp, ) diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 03246ef0..1996e671 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -18,7 +18,8 @@ from skala.functional import load_functional from skala.functional.base import ExcFunctionalBase -from skala.pyscf.numint import SkalaNumInt, _should_screen_aos +from skala.pyscf.numint import SkalaNumInt +from skala.pyscf.xc_integrator import _should_screen_aos THREAD_COUNT = 4 MAX_MEMORY_MB = 2000 diff --git a/tests/test_gpu4pyscf_gradients.py b/tests/test_gpu4pyscf_gradients.py index b9eb1f69..cd2154a3 100644 --- a/tests/test_gpu4pyscf_gradients.py +++ b/tests/test_gpu4pyscf_gradients.py @@ -35,7 +35,7 @@ from skala.pyscf import SkalaKS as CpuSkalaKS # noqa: E402 from skala.pyscf.features import generate_features # noqa: E402 from skala.pyscf.gradients import SkalaRKSGradient as CpuSkalaRKSGradient # noqa: E402 -from skala.pyscf.numint import _should_screen_aos # noqa: E402 +from skala.pyscf.xc_integrator import _should_screen_aos # noqa: E402 from skala.utils import torch_allocator # noqa: E402 H2_SKALA_1_1_GRAD_REF = torch.tensor( diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py new file mode 100644 index 00000000..7991154a --- /dev/null +++ b/tests/test_xc_integrator.py @@ -0,0 +1,107 @@ +from typing import Any, cast + +import numpy as np +import pytest +import torch +from pyscf import dft, gto + +from skala.functional.base import ExcFunctionalBase +from skala.pyscf import xc_integrator as xc_integrator_module +from skala.pyscf.numint import SkalaNumInt +from skala.pyscf.xc_integrator import XCIntegrator, XCResult + + +class QuadraticDensityFunctional(ExcFunctionalBase): + def __init__(self) -> None: + super().__init__() + self.features = ["density"] + + def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + return (mol["density"].square() * mol["grid_weights"]).sum() + + +def test_xc_integrator_returns_tensors_and_xc_only_response( + monkeypatch: pytest.MonkeyPatch, +) -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + grids = dft.Grids(mol) + + def fake_generate_features( + mol: gto.Mole, + dm: torch.Tensor, + grids: object, + features: set[str], + **kwargs: object, + ) -> dict[str, torch.Tensor]: + assert features == {"density", "grid_weights"} + return { + "density": dm.sum().reshape(1), + "grid_weights": torch.tensor([2.0], dtype=dm.dtype), + } + + monkeypatch.setattr( + xc_integrator_module, + "generate_features", + fake_generate_features, + ) + integrator = XCIntegrator(QuadraticDensityFunctional()) + dm = torch.tensor([[1.0, 2.0], [2.0, 3.0]], dtype=torch.float64) + + result = integrator(mol, grids, dm) + response = integrator.gen_response(mol, grids, dm.detach().clone()) + + assert isinstance(result, XCResult) + torch.testing.assert_close(result.electron_count, dm.new_tensor(16.0)) + torch.testing.assert_close(result.energy, dm.new_tensor(128.0)) + torch.testing.assert_close(result.potential, torch.full_like(dm, 32.0)) + torch.testing.assert_close(response(torch.ones_like(dm)), torch.full_like(dm, 16.0)) + + +class FakeKS: + def __init__(self, mol: gto.Mole) -> None: + self.mol = mol + self.grids = dft.Grids(mol) + self.max_memory = 123 + + def make_rdm1(self, mo_coeff: np.ndarray, mo_occ: np.ndarray) -> np.ndarray: + return np.eye(self.mol.nao_nr()) + + def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: + assert hermi == 1 + return np.full_like(dm, 3.0) + + +class FakeXCIntegrator: + device = torch.device("cpu") + + def __init__(self) -> None: + self.calls: list[tuple[int, float | None]] = [] + + def gen_response( + self, + mol: gto.Mole, + grids: object, + dm0: torch.Tensor, + max_memory: int, + safety_fraction: float | None, + ) -> Any: + self.calls.append((max_memory, safety_fraction)) + return lambda dm1: 2 * dm1 + + +def test_numint_response_adds_coulomb_to_xc_response() -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + ks = FakeKS(mol) + numint = SkalaNumInt(QuadraticDensityFunctional()) + fake_integrator = FakeXCIntegrator() + numint.integrator = cast(XCIntegrator, fake_integrator) + + response = numint.gen_response( + np.eye(mol.nao_nr()), + np.ones(mol.nao_nr()), + ks=cast(Any, ks), + safety_fraction=0.6, + ) + + np.testing.assert_allclose(response(np.ones((mol.nao_nr(), mol.nao_nr()))), 5.0) + assert fake_integrator.calls == [(123, 0.6)] From 13077363357edff45cda742dd2ffc8508b604ad9 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Wed, 5 Aug 2026 20:23:52 +0200 Subject: [PATCH 20/39] make functions private --- src/skala/pyscf/numint.py | 37 ++++++++++++++----------------------- 1 file changed, 14 insertions(+), 23 deletions(-) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 38e80143..fb86b7e8 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -135,25 +135,16 @@ def evaluation_policy(self) -> EvaluationPolicy: """Numerical policy owned by the XC integrator.""" return self.integrator.evaluation_policy - def reset(self) -> "SkalaNumInt[Array]": - """Return this integrator; spatial layouts are owned by grid objects.""" - return self - - def from_backend( - self, - x: Array, - device: torch.device | None = None, - transpose: bool = False, - ) -> Tensor: - return from_numpy_or_cupy(x, device=device or self.device, transpose=transpose) + def _from_backend(self, x: Array) -> Tensor: + return from_numpy_or_cupy(x, device=self.device) @overload - def to_backend(self, x: Tensor) -> Array: ... + def _to_backend(self, x: Tensor) -> Array: ... @overload - def to_backend(self, x: list[Tensor]) -> list[Array]: ... - def to_backend(self, x: Tensor | list[Tensor]) -> Array | list[Array]: + def _to_backend(self, x: list[Tensor]) -> list[Array]: ... + def _to_backend(self, x: Tensor | list[Tensor]) -> Array | list[Array]: if isinstance(x, list): - return [self.to_backend(y) for y in x] + return [self._to_backend(y) for y in x] if self.device.type == "cuda": return to_cupy(x) @@ -170,11 +161,11 @@ def get_rho( ) -> Array: density = self.integrator.density( mol, - self.from_backend(dm), + self._from_backend(dm), grids, max_memory=max_memory, ) - return self.to_backend(density) + return self._to_backend(density) def __call__( self, @@ -217,9 +208,9 @@ def nr_rks( """Restricted Kohn-Sham method, applicable if both spin-densities as equal.""" assert len(dm.shape) == 2 N, E_xc, V_xc = self( - mol, grids, xc_code, self.from_backend(dm), max_memory=max_memory + mol, grids, xc_code, self._from_backend(dm), max_memory=max_memory ) - return N.sum().item(), E_xc.item(), self.to_backend(V_xc) + return N.sum().item(), E_xc.item(), self._to_backend(V_xc) def nr_uks( self, @@ -232,9 +223,9 @@ def nr_uks( """Unrestricted Kohn-Sham method, spin densities can be different.""" assert len(dm.shape) == 3 and dm.shape[0] == 2 N, E_xc, V_xc = self( - mol, grids, xc_code, self.from_backend(dm), max_memory=max_memory + mol, grids, xc_code, self._from_backend(dm), max_memory=max_memory ) - return self.to_backend(N), E_xc.item(), self.to_backend(V_xc) + return self._to_backend(N), E_xc.item(), self._to_backend(V_xc) class libxc: __version__ = None @@ -270,7 +261,7 @@ def gen_response( if "with_j" in kwargs: assert kwargs["with_j"] - dm0 = self.from_backend(ks.make_rdm1(mo_coeff, mo_occ)) + dm0 = self._from_backend(ks.make_rdm1(mo_coeff, mo_occ)) xc_response = self.integrator.gen_response( ks.mol, ks.grids, @@ -280,7 +271,7 @@ def gen_response( ) def hessian_vector_product(dm1: Array) -> Array: - v1 = self.to_backend(xc_response(self.from_backend(dm1))) + v1 = self._to_backend(xc_response(self._from_backend(dm1))) vj = ks.get_j(ks.mol, dm1, hermi=1) if ks.mol.spin == 0: From 428c3932e6a955452bc58b8a6a410782e71b5762 Mon Sep 17 00:00:00 2001 From: Jens Date: Thu, 6 Aug 2026 17:16:33 +0200 Subject: [PATCH 21/39] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- src/skala/pyscf/xc_integrator.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index b2bdeea5..0aae8f4c 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -156,14 +156,15 @@ def _integrate_screened( max_memory: int, ) -> XCResult: dm = dm.detach().requires_grad_() - electron_count = torch.tensor((0.0, 0.0), device=self.device, dtype=dm.dtype) - energy = torch.tensor(0.0, device=self.device, dtype=dm.dtype) + dm_eval = dm.double() + electron_count = torch.zeros(2, device=self.device, dtype=dm_eval.dtype) + energy = torch.tensor(0.0, device=self.device, dtype=dm_eval.dtype) feature_function = feature_math.MGGAFeatureFunction(self.feature_spec) spatial_grid_layout = self._get_spatial_grid_layout(mol, grids) sorted_raw_features = cast( Tensor, ao_evaluation.ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm.double(), + dm_eval, mol, spatial_grid_layout.sorted_grids, feature_function, From 5c0d8db2725e1f5a13b2f65dff4bc175433369f2 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Thu, 6 Aug 2026 19:17:23 +0200 Subject: [PATCH 22/39] fix screening in the cpu case --- src/skala/pyscf/ao_evaluation.py | 287 ++++++++++++++++++++++--------- src/skala/pyscf/xc_integrator.py | 3 +- tests/test_ao_screening.py | 136 +++++++++++---- 3 files changed, 311 insertions(+), 115 deletions(-) diff --git a/src/skala/pyscf/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py index 18be440e..c7ac53db 100644 --- a/src/skala/pyscf/ao_evaluation.py +++ b/src/skala/pyscf/ao_evaluation.py @@ -45,21 +45,26 @@ class _ChunkEvalBackwardContext(Protocol): gpu: bool -def _active_cpu_aos(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: - """Expand a PySCF shell-screening mask into active AO indices.""" +def _active_cpu_ao_indices(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: + """Expand active shells in a PySCF screen-index slice to AO indices. + + A shell is active for the grid block if it is nonzero in any of the + ``BLKSIZE``-point rows covered by that block. ``ao_loc_nr`` maps each shell + to its contiguous range in PySCF's AO ordering. + """ active_shells = np.any(screen_index, axis=0) ao_loc = mol.ao_loc_nr() return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) -def partial_feature_function_over_aos( +def partial_feature_function_over_ao_values( feature_function: Callable[[torch.Tensor, torch.Tensor], torch.Tensor], - ao: torch.Tensor, + ao_values: torch.Tensor, ) -> Callable[[torch.Tensor], torch.Tensor]: - """Bind an AO block to a feature function for block-local evaluation.""" + """Bind evaluated AO values to a feature function for one grid block.""" def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: - return feature_function(dm, ao) + return feature_function(dm, ao_values) return partial_feature_function @@ -78,44 +83,82 @@ def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: @dataclass(frozen=True) class _AOBlock: - ao: Tensor - active_aos: Tensor | None + """Evaluated AO data and index metadata for one contiguous grid block. + + ``ao_values`` contains only the active AO rows when screening is enabled. + ``active_ao_indices`` identifies those rows in the backend's current AO + ordering; ``None`` means that ``ao_values`` contains every AO. The CPU + backend uses PySCF's native AO order, while the GPU backend uses + GPU4PySCF's sorted AO order until the completed matrix is restored. + """ + + ao_values: Tensor + active_ao_indices: Tensor | None grid_slice: slice - def select_aos(self, matrix: Tensor) -> Tensor: - if self.active_aos is None: + def select_active_ao_submatrix(self, matrix: Tensor) -> Tensor: + """Gather the square matrix corresponding to this block's AO values.""" + if self.active_ao_indices is None: return matrix - return matrix[..., self.active_aos[:, None], self.active_aos[None, :]] + return matrix[ + ..., self.active_ao_indices[:, None], self.active_ao_indices[None, :] + ] - def add_to(self, matrix: Tensor, block_result: Tensor) -> None: - if self.active_aos is None: + def add_active_ao_submatrix(self, matrix: Tensor, block_result: Tensor) -> None: + """Add a block result into its active rows and columns in ``matrix``.""" + if self.active_ao_indices is None: matrix += block_result else: - matrix[..., self.active_aos[:, None], self.active_aos[None, :]] += ( - block_result - ) + matrix[ + ..., self.active_ao_indices[:, None], self.active_ao_indices[None, :] + ] += block_result def _evaluate_feature_block( feature_function: feature_math.FeatureFunction, block: _AOBlock, - active_dm: Tensor, + active_dm_submatrix: Tensor, compile_feature_function: bool, feature_cotangent: Tensor | None = None, ) -> Tensor: """Evaluate one active-AO feature block or its feature-space VJP.""" - partial_func = partial_feature_function_over_aos(feature_function, block.ao) + partial_func = partial_feature_function_over_ao_values( + feature_function, block.ao_values + ) if feature_cotangent is not None: partial_func = partial_vjp_function_over_tangents( partial_func, feature_cotangent[..., block.grid_slice] ) if compile_feature_function: - return torch.compile(partial_func)(active_dm) - return partial_func(active_dm) - + return torch.compile(partial_func)(active_dm_submatrix) + return partial_func(active_dm_submatrix) + + +class _CPUAOBlockLoop: + """Yield CPU AO values screened with the exact PySCF screen-index table. + + PySCF evaluates AOs with ``grids.non0tab``, whose rows each describe one + ``dft.gen_grid.BLKSIZE``-point range and whose columns describe shells. The + loop converts the rows covered by each yielded grid block into AO indices, + slices the evaluated AO tensor, and records those indices for density-matrix + gathering and result scattering. If every shell is active for a particular + block, the loop keeps the full AO tensor and records ``None`` instead of an + identity index. Whether a block is dense can therefore vary across the + rows of one ``non0tab`` table. + + The second item yielded by ``NumInt.block_loop`` is intentionally ignored. + Despite being called ``mask`` by PySCF, it is not the authoritative + screening table for that block. After AO evaluation, PySCF may replace it + with ``None`` to request dense downstream contractions. That policy depends + on the total grid's ``ALIGNMENT_UNIT`` divisibility and PySCF's sparsity + heuristic, not on whether shells were screened during AO evaluation. Using + that yielded value would therefore make Skala's active AO set depend on + contraction policy and grid alignment. Reading the exact rows from + ``grids.non0tab`` preserves the screening information actually used for AO + evaluation. + """ -class _AOBlockLoop: def __init__( self, dm: Tensor, @@ -123,83 +166,157 @@ def __init__( grids: Grid, feature_function: feature_math.FeatureFunction, blksize: int | None, - gpu: bool, ) -> None: self.dm = dm self.mol = mol + assert isinstance(grids, dft.Grids) self.grids = grids self.feature_function = feature_function self.blksize = blksize - self.gpu = gpu - self.sort_idx: Tensor | None - self.unsort_idx: Tensor | None - - if gpu: - check_gpu_imports_were_successful() - self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) - self.numint.grid_blksize = blksize - self.sort_idx = torch.as_tensor( - self.numint.gdftopt._ao_idx, device=dm.device - ) - self.unsort_idx = torch.argsort(self.sort_idx) - else: - self.numint = dft.numint.NumInt() - self.sort_idx = None - self.unsort_idx = None + self.numint = dft.numint.NumInt() def order_aos(self, matrix: Tensor) -> Tensor: - if self.sort_idx is None: - return matrix - return matrix[..., self.sort_idx, :][..., self.sort_idx] + return matrix def restore_ao_order(self, matrix: Tensor) -> Tensor: - if self.unsort_idx is None: - return matrix - return matrix[..., self.unsort_idx, :][..., self.unsort_idx] + return matrix + + def _active_ao_indices( + self, + non0tab: np.ndarray, + grid_start: int, + grid_end: int, + ) -> Tensor | None: + """Create active AO indices for the exact rows covering a grid block. + + ``NumInt.block_loop`` requires CPU block sizes to be integer multiples + of ``dft.gen_grid.BLKSIZE``. Consequently every non-final block starts + and ends on screen-index row boundaries; the ceiling for ``grid_end`` + also includes the final partial row. All shells active in any covered + row are included because one AO tensor is shared by the whole grid + block. + + Returns ``None`` when the covered rows activate every AO. ``_AOBlock`` + uses that value as its dense sentinel, avoiding identity indexing of AO + values and density matrices. An empty tensor means that no AO is active + and the caller can omit the block entirely. + + This method requires the authoritative screen-index table and must not + consume the mask yielded by ``NumInt.block_loop``. PySCF may set that + yielded mask to ``None`` after AO evaluation when sparse contraction is + unsuitable, even though ``non0tab`` still contains the exact + shell-screening data. The caller handles a missing ``non0tab`` as the + genuinely dense case. + """ + row_start = grid_start // dft.gen_grid.BLKSIZE + row_end = (grid_end + dft.gen_grid.BLKSIZE - 1) // dft.gen_grid.BLKSIZE + block_non0tab = non0tab[row_start:row_end] + if np.all(np.any(block_non0tab, axis=0)): + return None + return torch.as_tensor( + _active_cpu_ao_indices(self.mol, block_non0tab), + device=self.dm.device, + dtype=torch.long, + ) def __iter__(self) -> Iterator[_AOBlock]: - block_loop_options: dict[str, bool] = {} - if self.gpu: - # GPU4PySCF otherwise omits zero-AO blocks, shifting all later grid slices. - block_loop_options["strict_grid_order"] = True + non0tab = self.grids.non0tab end = 0 - for ao_block, mask, weights, _ in self.numint.block_loop( + for backend_ao_values, _, block_weights, _ in self.numint.block_loop( mol=self.mol, grids=self.grids, nao=self.mol.nao, deriv=self.feature_function.deriv, blksize=self.blksize, - non0tab=(None if self.gpu else getattr(self.grids, "non0tab", None)), - **block_loop_options, + non0tab=non0tab, ): - start, end = end, end + weights.size - ao = from_numpy_or_cupy( - ao_block, - device=self.dm.device, - dtype=self.dm.dtype, - transpose=not self.gpu, + start, end = end, end + block_weights.size + ao_values = ( + torch.from_numpy(backend_ao_values) + .to(device=self.dm.device, dtype=self.dm.dtype) + .transpose(-1, -2) ) - active_aos: Tensor | None - if mask is None: - active_aos = None - elif self.gpu: - active_aos = from_numpy_or_cupy( - mask, device=self.dm.device, dtype=torch.long - ) - else: - num_screen_rows = ( - weights.size + dft.gen_grid.BLKSIZE - 1 - ) // dft.gen_grid.BLKSIZE - active_aos = torch.as_tensor( - _active_cpu_aos(self.mol, mask[:num_screen_rows]), - device=self.dm.device, - dtype=torch.long, - ) - ao = ao[..., active_aos, :] - if active_aos is not None and active_aos.numel() == 0: + active_ao_indices = ( + None + if non0tab is None + else self._active_ao_indices(non0tab, start, end) + ) + if active_ao_indices is None: + yield _AOBlock(ao_values, None, slice(start, end)) + continue + + if active_ao_indices.numel() == 0: continue - yield _AOBlock(ao, active_aos, slice(start, end)) + ao_values = ao_values[..., active_ao_indices, :] + yield _AOBlock(ao_values, active_ao_indices, slice(start, end)) + + +class _GPUAOBlockLoop: + """Yield GPU4PySCF AO values and compact indices in sorted AO order.""" + + def __init__( + self, + dm: Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + blksize: int | None, + ) -> None: + check_gpu_imports_were_successful() + self.dm = dm + self.mol = mol + self.grids = grids + self.feature_function = feature_function + self.blksize = blksize + self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) + self.numint.grid_blksize = blksize + self.sort_idx = torch.as_tensor(self.numint.gdftopt._ao_idx, device=dm.device) + self.unsort_idx = torch.argsort(self.sort_idx) + + def order_aos(self, matrix: Tensor) -> Tensor: + return matrix[..., self.sort_idx[:, None], self.sort_idx[None, :]] + + def restore_ao_order(self, matrix: Tensor) -> Tensor: + return matrix[..., self.unsort_idx[:, None], self.unsort_idx[None, :]] + + def __iter__(self) -> Iterator[_AOBlock]: + end = 0 + for ( + backend_ao_values, + active_ao_indices, + block_weights, + _, + ) in self.numint.block_loop( + mol=self.mol, + grids=self.grids, + nao=self.mol.nao, + deriv=self.feature_function.deriv, + blksize=self.blksize, + non0tab=None, + # GPU4PySCF otherwise omits zero-AO blocks, shifting later grid slices. + strict_grid_order=True, + ): + start, end = end, end + block_weights.size + if active_ao_indices.size == 0: + continue + yield _AOBlock( + torch.from_dlpack(backend_ao_values), + torch.from_dlpack(active_ao_indices), + slice(start, end), + ) + + +def _make_ao_block_loop( + dm: Tensor, + mol: gto.Mole, + grids: Grid, + feature_function: feature_math.FeatureFunction, + blksize: int | None, + gpu: bool, +) -> _CPUAOBlockLoop | _GPUAOBlockLoop: + loop_type = _GPUAOBlockLoop if gpu else _CPUAOBlockLoop + return loop_type(dm, mol, grids, feature_function, blksize) class ChunkEvalForward(Function): @@ -247,7 +364,7 @@ def forward( *vectors_jvp: torch.Tensor, ) -> torch.Tensor: ngrids = grids.weights.size - block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + block_loop = _make_ao_block_loop(dm, mol, grids, feature_function, blksize, gpu) features = torch.zeros( *dm.shape[:-2], @@ -263,11 +380,13 @@ def forward( evaluation_dm = vectors_jvp[0] if vectors_jvp else dm evaluation_dm_ordered = block_loop.order_aos(evaluation_dm) for block in block_loop: - active_dm = block.select_aos(evaluation_dm_ordered) + active_dm_submatrix = block.select_active_ao_submatrix( + evaluation_dm_ordered + ) temp_feature = _evaluate_feature_block( feature_function, block, - active_dm, + active_dm_submatrix, compile_feature_function, ) features[..., block.grid_slice] = temp_feature @@ -385,20 +504,20 @@ def forward( gpu: bool, feature_cotangent: torch.Tensor, ) -> torch.Tensor: - block_loop = _AOBlockLoop(dm, mol, grids, feature_function, blksize, gpu) + block_loop = _make_ao_block_loop(dm, mol, grids, feature_function, blksize, gpu) dm_ordered = block_loop.order_aos(dm) out = torch.zeros_like(dm) for block in block_loop: - active_dm = block.select_aos(dm_ordered) + active_dm_submatrix = block.select_active_ao_submatrix(dm_ordered) block_result = _evaluate_feature_block( feature_function, block, - active_dm, + active_dm_submatrix, compile_feature_function, feature_cotangent, ) - block.add_to(out, block_result) + block.add_active_ao_submatrix(out, block_result) return block_loop.restore_ao_order(out) @staticmethod diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index 0aae8f4c..096f4bef 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -27,7 +27,8 @@ def _should_screen_aos(mol: gto.Mole) -> bool: """Return whether PySCF's sparse-contraction crossover is exceeded.""" - return mol.nao_nr() > pyscf_numint.SWITCH_SIZE + # we use a smaller threshold because for MetaGGAs the AO evaluation is more expensive + return 2 * mol.nao_nr() > pyscf_numint.SWITCH_SIZE @dataclass(frozen=True) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index f71faa11..22806b78 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -13,8 +13,9 @@ from skala.pyscf.ao_evaluation import ( ChunkEvalBackward, ChunkEvalForward, - _active_cpu_aos, + _active_cpu_ao_indices, _AOBlock, + _CPUAOBlockLoop, _evaluate_feature_block, _resolve_ao_block_size, ) @@ -110,13 +111,13 @@ def test_should_screen_aos_at_crossover( monkeypatch.setattr( pyscf_numint, "SWITCH_SIZE", - carbon.nao_nr() + switch_offset, + 2 * carbon.nao_nr() + switch_offset, ) assert _should_screen_aos(carbon) is expected -def test_active_cpu_aos(carbon: gto.Mole) -> None: +def test_active_cpu_ao_indices(carbon: gto.Mole) -> None: ao_loc = carbon.ao_loc_nr() screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) screen_index[0, 0] = 1 @@ -129,9 +130,9 @@ def test_active_cpu_aos(carbon: gto.Mole) -> None: ) ) - assert np.array_equal(_active_cpu_aos(carbon, screen_index), expected) + assert np.array_equal(_active_cpu_ao_indices(carbon, screen_index), expected) - empty = _active_cpu_aos(carbon, np.zeros_like(screen_index)) + empty = _active_cpu_ao_indices(carbon, np.zeros_like(screen_index)) assert empty.dtype == np.int64 assert empty.size == 0 @@ -359,7 +360,7 @@ def test_first_and_second_order_use_same_screening_decision( expected: bool, response_safety_fraction: float | None, ) -> None: - switch_size = carbon.nao_nr() - 1 if expected else carbon.nao_nr() + switch_size = 2 * carbon.nao_nr() - 1 if expected else 2 * carbon.nao_nr() monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", switch_size) routes: list[str] = [] safety_fractions: list[float] = [] @@ -501,8 +502,8 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: """Use AO slices for linear JVPs and grid slices for feature VJPs.""" feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) block = _AOBlock( - ao=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), - active_aos=torch.tensor([0, 2]), + ao_values=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), + active_ao_indices=torch.tensor([0, 2]), grid_slice=slice(1, 3), ) dm_ordered = torch.tensor( @@ -513,27 +514,31 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: [[0.5, 1.0, -0.2], [1.0, 0.4, 0.3], [-0.2, 0.3, 0.7]], dtype=torch.float64, ) - active_dm = block.select_aos(dm_ordered) + active_dm_submatrix = block.select_active_ao_submatrix(dm_ordered) feature_jvp = _evaluate_feature_block( feature_function, block, - block.select_aos(tangent_ordered), + block.select_active_ao_submatrix(tangent_ordered), compile_feature_function=False, ) - expected_jvp = feature_function(block.select_aos(tangent_ordered), block.ao) + expected_jvp = feature_function( + block.select_active_ao_submatrix(tangent_ordered), block.ao_values + ) torch.testing.assert_close(feature_jvp, expected_jvp) full_grid_cotangent = torch.tensor([[10.0, 0.25, -0.5, 20.0]], dtype=torch.float64) feature_vjp = _evaluate_feature_block( feature_function, block, - active_dm, + active_dm_submatrix, compile_feature_function=False, feature_cotangent=full_grid_cotangent, ) local_cotangent = full_grid_cotangent[0, block.grid_slice] - expected_vjp = torch.einsum("g,ig,jg->ij", local_cotangent, block.ao, block.ao) + expected_vjp = torch.einsum( + "g,ig,jg->ij", local_cotangent, block.ao_values, block.ao_values + ) torch.testing.assert_close(feature_vjp, expected_vjp) @@ -579,7 +584,8 @@ def apply_adjoint(value: torch.Tensor) -> torch.Tensor: def test_cpu_screening_slices_and_scatters_full_derivatives( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: - ngrids = dft.gen_grid.BLKSIZE + block_size = dft.gen_grid.BLKSIZE + ngrids = 2 * block_size grids = dft.Grids(carbon) grids.coords = np.zeros((ngrids, 3)) grids.weights = np.ones(ngrids) @@ -587,18 +593,26 @@ def test_cpu_screening_slices_and_scatters_full_derivatives( ao = np.arange(ngrids * carbon.nao_nr(), dtype=np.float64).reshape( ngrids, carbon.nao_nr() ) - screen_index = np.zeros((1, carbon.nbas), dtype=np.uint8) - screen_index[0, (0, -1)] = 1 + screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) + screen_index[0, 0] = 1 + screen_index[1, -1] = 1 grids.non0tab = screen_index - active_aos = _active_cpu_aos(carbon, screen_index) + active_ao_indices = _active_cpu_ao_indices(carbon, screen_index) class FakeNumInt: def block_loop( self, *args: object, **kwargs: object - ) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]]: + ) -> Iterator[tuple[np.ndarray, None, np.ndarray, np.ndarray]]: assert kwargs["non0tab"] is screen_index assert "strict_grid_order" not in kwargs - yield ao, screen_index, grids.weights, grids.coords + for start in range(0, ngrids, block_size): + grid_slice = slice(start, start + block_size) + yield ( + ao[grid_slice], + None, + grids.weights[grid_slice], + grids.coords[grid_slice], + ) monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) @@ -607,25 +621,85 @@ def block_loop( torch.arange(1, carbon.nao_nr() + 1, dtype=torch.float64) ).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm, carbon, grids, feature_function, ngrids, False, False + dm, carbon, grids, feature_function, block_size, False, False ) - ao_active = torch.from_numpy(ao[:, active_aos]).T - dm_active = dm[..., active_aos[:, None], active_aos[None, :]] - expected = torch.sum((dm_active @ ao_active) * ao_active, dim=0).unsqueeze(0) + expected_blocks = [] + for block_index, start in enumerate(range(0, ngrids, block_size)): + block_active_ao_indices = _active_cpu_ao_indices( + carbon, screen_index[block_index : block_index + 1] + ) + grid_slice = slice(start, start + block_size) + active_ao_values = torch.from_numpy( + ao[grid_slice][:, block_active_ao_indices] + ).T + active_dm_submatrix = dm[ + ..., + block_active_ao_indices[:, None], + block_active_ao_indices[None, :], + ] + expected_blocks.append( + torch.sum( + (active_dm_submatrix @ active_ao_values) * active_ao_values, dim=0 + ) + ) + expected = torch.cat(expected_blocks).unsqueeze(0) assert torch.allclose(features, expected) energy = features.square().sum() (vxc,) = torch.autograd.grad(energy, dm, create_graph=True) (hvp,) = torch.autograd.grad(vxc, dm, torch.ones_like(dm)) - inactive_aos = np.setdiff1d(np.arange(carbon.nao_nr()), active_aos) + inactive_ao_indices = np.setdiff1d(np.arange(carbon.nao_nr()), active_ao_indices) assert vxc.shape == dm.shape assert hvp.shape == dm.shape - assert torch.count_nonzero(vxc[inactive_aos]) == 0 - assert torch.count_nonzero(vxc[:, inactive_aos]) == 0 - assert torch.count_nonzero(hvp[inactive_aos]) == 0 - assert torch.count_nonzero(hvp[:, inactive_aos]) == 0 + assert torch.count_nonzero(vxc[inactive_ao_indices]) == 0 + assert torch.count_nonzero(vxc[:, inactive_ao_indices]) == 0 + assert torch.count_nonzero(hvp[inactive_ao_indices]) == 0 + assert torch.count_nonzero(hvp[:, inactive_ao_indices]) == 0 + + +def test_cpu_all_active_block_uses_dense_sentinel( + carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch +) -> None: + block_size = dft.gen_grid.BLKSIZE + ngrids = 2 * block_size + grids = dft.Grids(carbon) + grids.coords = np.zeros((ngrids, 3)) + grids.weights = np.ones(ngrids) + screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) + screen_index[0] = 1 + screen_index[1, 0] = 1 + grids.non0tab = screen_index + ao_values = np.ones((ngrids, carbon.nao_nr())) + + class FakeNumInt: + def block_loop( + self, *args: object, **kwargs: object + ) -> Iterator[tuple[np.ndarray, None, np.ndarray, np.ndarray]]: + for start in range(0, ngrids, block_size): + grid_slice = slice(start, start + block_size) + yield ( + ao_values[grid_slice], + None, + grids.weights[grid_slice], + grids.coords[grid_slice], + ) + + monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) + feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + + blocks = list(_CPUAOBlockLoop(dm, carbon, grids, feature_function, block_size)) + + assert len(blocks) == 2 + assert blocks[0].active_ao_indices is None + assert blocks[0].ao_values.shape == (carbon.nao_nr(), block_size) + expected_sparse_indices = torch.as_tensor( + _active_cpu_ao_indices(carbon, screen_index[1:]), dtype=torch.long + ) + torch.testing.assert_close(blocks[1].active_ao_indices, expected_sparse_indices) + assert blocks[1].ao_values.shape == (expected_sparse_indices.numel(), block_size) def test_cpu_no_active_aos_returns_full_zero_derivatives( @@ -637,12 +711,14 @@ def test_cpu_no_active_aos_returns_full_zero_derivatives( grids.weights = np.ones(ngrids) ao = np.ones((ngrids, carbon.nao_nr())) screen_index = np.zeros((1, carbon.nbas), dtype=np.uint8) + grids.non0tab = screen_index class FakeNumInt: def block_loop( self, *args: object, **kwargs: object - ) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]]: - yield ao, screen_index, grids.weights, grids.coords + ) -> Iterator[tuple[np.ndarray, None, np.ndarray, np.ndarray]]: + assert kwargs["non0tab"] is screen_index + yield ao, None, grids.weights, grids.coords monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) From eec82b4ee9e2f9b19d467ae51a36fc914fa5364d Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 10:50:28 +0200 Subject: [PATCH 23/39] use pca grid decomposition --- src/skala/pyscf/screening.py | 32 ++++++++++++++++++++++++++------ tests/test_ao_screening.py | 26 ++++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 6 deletions(-) diff --git a/src/skala/pyscf/screening.py b/src/skala/pyscf/screening.py index 1f75976a..1359231e 100644 --- a/src/skala/pyscf/screening.py +++ b/src/skala/pyscf/screening.py @@ -55,9 +55,10 @@ def _decompose_grid_into_spatial_blocks( ) -> tuple[_Int64Permutation, _Int64Permutation]: """Decompose a molecular grid into spatial blocks and return its permutations. - Recursively partitions points along the longest Cartesian extent. Every left - subtree contains a whole number of evaluator blocks, so all output blocks have - ``block_size`` points except for a possible final remainder. + Recursively partitions points along their principal spatial direction. Every + left subtree contains a whole number of evaluator blocks, so all output blocks + have ``block_size`` points except for a possible final remainder. Degenerate + principal directions fall back to the longest Cartesian extent. Args: coords: Molecular grid coordinates with shape ``(ngrids, 3)``. @@ -74,15 +75,34 @@ def _decompose_grid_into_spatial_blocks( if block_size <= 0: raise ValueError("block_size must be positive") + def split_projections(indices: _Int64Permutation) -> np.ndarray: + point_coords = coords[indices] + centered_coords = point_coords - point_coords.mean(axis=0) + scatter = centered_coords.T @ centered_coords + eigenvalues, eigenvectors = np.linalg.eigh(scatter) + eigenvalue_scale = max(abs(eigenvalues[-1]), abs(eigenvalues[-2])) + if np.isclose( + eigenvalues[-1], + eigenvalues[-2], + rtol=1e-12, + atol=np.finfo(np.float64).eps * eigenvalue_scale, + ): + split_axis = int(np.argmax(np.ptp(point_coords, axis=0))) + return point_coords[:, split_axis] + + principal_direction = eigenvectors[:, -1] + largest_component = int(np.argmax(np.abs(principal_direction))) + if principal_direction[largest_component] < 0: + principal_direction = -principal_direction + return centered_coords @ principal_direction + def partition(indices: _Int64Permutation) -> list[_Int64Permutation]: if indices.size <= block_size: return [indices] block_count = (indices.size + block_size - 1) // block_size left_size = (block_count // 2) * block_size - extents = np.ptp(coords[indices], axis=0) - split_axis = int(np.argmax(extents)) - positions = np.lexsort((indices, coords[indices, split_axis])) + positions = np.lexsort((indices, split_projections(indices))) ordered_indices = indices[positions] return partition(ordered_indices[:left_size]) + partition( ordered_indices[left_size:] diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 22806b78..94dbe126 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -202,6 +202,32 @@ def test_decompose_grid_into_spatial_blocks_groups_interleaved_clusters() -> Non assert np.all(grouped_labels == grouped_labels[:, :1]) +def test_decompose_grid_into_spatial_blocks_uses_principal_direction() -> None: + longitudinal = np.arange(-3.5, 4.0) + transverse = 0.45 * (np.square(longitudinal) - np.mean(np.square(longitudinal))) + coords = np.column_stack( + ( + longitudinal + transverse, + longitudinal - transverse, + np.zeros(longitudinal.size), + ) + ) + + forward, _ = _decompose_grid_into_spatial_blocks(coords, block_size=2) + + assert set(forward[:4]) == set(range(4)) + assert set(forward[4:]) == set(range(4, 8)) + + +def test_decompose_grid_into_spatial_blocks_handles_identical_points() -> None: + coords = np.ones((10, 3), dtype=np.float64) + + forward, inverse = _decompose_grid_into_spatial_blocks(coords, block_size=4) + + assert np.array_equal(forward, np.arange(coords.shape[0])) + assert np.array_equal(inverse, forward) + + def test_prepare_spatially_sorted_cpu_grids( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: From cda801018de654db08ee181efd687a34fa3a6a69 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 13:37:18 +0200 Subject: [PATCH 24/39] use enum and rework switch size patching --- src/skala/functional/base.py | 10 +- src/skala/functional/density.py | 14 +- src/skala/functional/load.py | 7 +- src/skala/functional/model.py | 92 +++++++------- src/skala/functional/traditional.py | 109 ++++++++++------ src/skala/gpu4pyscf/gradients.py | 61 ++++----- src/skala/pyscf/ao_evaluation.py | 101 +++++---------- src/skala/pyscf/evaluation.py | 57 ++++++--- src/skala/pyscf/feature_math.py | 29 +++-- src/skala/pyscf/features.py | 47 ++++--- src/skala/pyscf/gradients.py | 61 ++++----- src/skala/pyscf/memory_estimators.py | 4 +- src/skala/pyscf/model_chunking.py | 59 +++++---- src/skala/pyscf/numint.py | 11 -- src/skala/pyscf/xc_integrator.py | 25 ++-- tests/test_ao_screening.py | 184 ++++++++++++++------------- tests/test_ao_screening_benchmark.py | 76 +++++------ tests/test_evaluation.py | 71 ++++------- tests/test_gpu4pyscf_ao_screening.py | 122 +++++++++--------- tests/test_gpu4pyscf_gradients.py | 103 ++++++++------- tests/test_memory_estimators.py | 8 +- tests/test_model.py | 61 +++++---- tests/test_pyscf_gradients.py | 87 +++++++------ tests/test_xc_integrator.py | 17 +-- 24 files changed, 712 insertions(+), 704 deletions(-) diff --git a/src/skala/functional/base.py b/src/skala/functional/base.py index cd6137bf..06ec0154 100644 --- a/src/skala/functional/base.py +++ b/src/skala/functional/base.py @@ -13,6 +13,8 @@ import torch from torch import nn +from skala.features import Feature, FeatureMap + VxcType = tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] @@ -25,7 +27,7 @@ class ExcFunctionalBase(nn.Module): energy density from molecular features. """ - features: list[str] + features: list[Feature] """List of features that this functional requires.""" def get_d3_settings(self) -> str | None: @@ -35,7 +37,7 @@ def get_d3_settings(self) -> str | None: """ return None - def get_exc_density(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc_density(self, mol: FeatureMap) -> torch.Tensor: """ Returns the exchange-correlation density for the given molecule. It should return a tensor of shape (G,) where G is the number of grid points @@ -46,7 +48,7 @@ def get_exc_density(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: "get_exc_density not implemented for this functional." ) - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """ Compute the exchange-correlation energy. @@ -62,7 +64,7 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: The total exchange-correlation energy. """ exc_density = self.get_exc_density(mol).double() - grid_weights = mol["grid_weights"].double() + grid_weights = mol[Feature.GRID_WEIGHTS].double() return (exc_density * grid_weights).sum() diff --git a/src/skala/functional/density.py b/src/skala/functional/density.py index 69a16ca9..292b97dc 100644 --- a/src/skala/functional/density.py +++ b/src/skala/functional/density.py @@ -14,13 +14,13 @@ import torch from torch import Tensor +from skala.features import Feature, FeatureMap + EPS = 1e-10 -IMMUTABLES = frozenset(["grid_coords", "grid_weights"]) +IMMUTABLES: frozenset[str] = frozenset([Feature.GRID_COORDS, Feature.GRID_WEIGHTS]) -def _map( - mol_features: dict[str, Tensor], f: Callable[[Tensor], Tensor] -) -> dict[str, Tensor]: +def _map(mol_features: FeatureMap, f: Callable[[Tensor], Tensor]) -> FeatureMap: """ Apply a function to mutable molecular features. @@ -43,8 +43,8 @@ def _map( def separate( - mol_features: dict[str, Tensor], -) -> tuple[dict[str, Tensor], dict[str, Tensor]]: + mol_features: FeatureMap, +) -> tuple[FeatureMap, FeatureMap]: """ Separate molecular features into spin-up and spin-down components. @@ -74,7 +74,7 @@ def separate( return mol_a, mol_b -def scale_by(mol_features: dict[str, Tensor], factor: float) -> dict[str, Tensor]: +def scale_by(mol_features: FeatureMap, factor: float) -> FeatureMap: """ Scale molecular features by a constant factor. diff --git a/src/skala/functional/load.py b/src/skala/functional/load.py index ca3baa7e..ec081e6e 100644 --- a/src/skala/functional/load.py +++ b/src/skala/functional/load.py @@ -13,6 +13,7 @@ import torch +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase PROTOCOL_VERSION = 2 @@ -46,7 +47,7 @@ def __init__( super().__init__() self._traced_model = traced_model self.metadata = dict(metadata) - self.features = list(features) + self.features = [Feature(feature) for feature in features] self.expected_d3_settings = expected_d3_settings def get_d3_settings(self) -> str | None: @@ -56,10 +57,10 @@ def get_d3_settings(self) -> str | None: """ return self.expected_d3_settings - def get_exc_density(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc_density(self, mol: FeatureMap) -> torch.Tensor: return self._traced_model.get_exc_density(mol) - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: return self._traced_model.get_exc(mol) @property diff --git a/src/skala/functional/model.py b/src/skala/functional/model.py index ca817efd..d3c45092 100644 --- a/src/skala/functional/model.py +++ b/src/skala/functional/model.py @@ -16,6 +16,7 @@ from opt_einsum_fx import jitable, optimize_einsums_full from torch import fx, nn +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase, enhancement_density_inner_product from skala.functional.layers import ScaledSigmoid from skala.functional.utils.irreps import Irreps @@ -25,9 +26,7 @@ ANGSTROM_TO_BOHR = 1.88973 -def _prepare_features_raw( - mol: dict[str, torch.Tensor], eps: float = 1e-5 -) -> torch.Tensor: +def _prepare_features_raw(mol: FeatureMap, eps: float = 1e-5) -> torch.Tensor: """Compute log-space semi-local features from packed density data. Args: @@ -39,10 +38,10 @@ def _prepare_features_raw( """ x = torch.cat( [ - mol["density"].permute(1, 2, 0), - (mol["grad"] ** 2).sum(1).permute(1, 2, 0), - mol["kin"].permute(1, 2, 0), - (mol["grad"].sum(0) ** 2).sum(0).unsqueeze(-1), + mol[Feature.DENSITY].permute(1, 2, 0), + (mol[Feature.GRAD] ** 2).sum(1).permute(1, 2, 0), + mol[Feature.KIN].permute(1, 2, 0), + (mol[Feature.GRAD].sum(0) ** 2).sum(0).unsqueeze(-1), ], dim=-1, ) @@ -68,9 +67,7 @@ def __init__(self) -> None: persistent=False, ) - def forward( - self, mol: dict[str, torch.Tensor] - ) -> tuple[torch.Tensor, torch.Tensor]: + def forward(self, mol: FeatureMap) -> tuple[torch.Tensor, torch.Tensor]: features = _prepare_features_raw(mol) features_ab = features features_ba = features.index_select(-1, self._feature_perm) @@ -127,15 +124,15 @@ class SkalaFunctional(ExcFunctionalBase): """ features = [ - "density", - "kin", - "grad", - "grid_coords", - "grid_weights", - "atomic_grid_weights", - "atomic_grid_sizes", - "coarse_0_atomic_coords", - "atomic_grid_size_bound_shape", + Feature.DENSITY, + Feature.KIN, + Feature.GRAD, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.ATOMIC_GRID_SIZES, + Feature.COARSE_0_ATOMIC_COORDS, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, ] def __init__( @@ -250,9 +247,7 @@ def _init_weights(self) -> None: def dtype(self) -> torch.dtype: return cast(nn.Linear, self.input_model[0]).weight.dtype - def pack_features( - self, mol_feats: dict[str, torch.Tensor] - ) -> dict[str, torch.Tensor]: + def pack_features(self, mol_feats: FeatureMap) -> FeatureMap: """Pack flat features into dense (grid_per_atom, atoms, …) layout. Args: @@ -261,49 +256,52 @@ def pack_features( Returns: Packed features dictionary. """ - atomic_grid_sizes = mol_feats["atomic_grid_sizes"] - size_bound = mol_feats["atomic_grid_size_bound_shape"].shape[0] + atomic_grid_sizes = mol_feats[Feature.ATOMIC_GRID_SIZES] + size_bound = mol_feats[Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE].shape[0] - packed_mol_feats: dict[str, torch.Tensor] = {} + packed_mol_feats: FeatureMap = {} for key in self.features: - if key == "atomic_grid_weights": + if key == Feature.ATOMIC_GRID_WEIGHTS: packed_mol_feats[key] = pad_ragged( mol_feats[key], atomic_grid_sizes, size_bound ).T # (max_grid_size, num_atoms) - elif key == "grid_weights": + elif key == Feature.GRID_WEIGHTS: continue - elif key == "grid_coords": + elif key == Feature.GRID_COORDS: packed_mol_feats[key] = pad_ragged( mol_feats[key], atomic_grid_sizes, size_bound ).permute(1, 0, 2) # (max_grid_size, num_atoms, 3) - elif key == "coarse_0_atomic_coords": + elif key == Feature.COARSE_0_ATOMIC_COORDS: packed_mol_feats[key] = mol_feats[key] - elif key == "density": + elif key == Feature.DENSITY: packed_mol_feats[key] = pad_ragged( mol_feats[key].T, atomic_grid_sizes, size_bound ).permute(2, 1, 0) # (2, max_grid_size, num_atoms) - elif key == "grad": + elif key == Feature.GRAD: packed_mol_feats[key] = pad_ragged( mol_feats[key].permute(2, 0, 1), atomic_grid_sizes, size_bound ).permute(2, 3, 1, 0) # (2, 3, max_grid_size, num_atoms) - elif key == "kin": + elif key == Feature.KIN: packed_mol_feats[key] = pad_ragged( mol_feats[key].T, atomic_grid_sizes, size_bound ).permute(2, 1, 0) # (2, max_grid_size, num_atoms) - elif key in ("atomic_grid_sizes", "atomic_grid_size_bound_shape"): + elif key in ( + Feature.ATOMIC_GRID_SIZES, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, + ): continue else: raise ValueError(f"Unexpected key: {key}") return packed_mol_feats - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: exc_density = self._get_exc_density_padded(mol).double() grid_weights = ( pad_ragged( - mol["grid_weights"], - mol["atomic_grid_sizes"], - mol["atomic_grid_size_bound_shape"].shape[0], + mol[Feature.GRID_WEIGHTS], + mol[Feature.ATOMIC_GRID_SIZES], + mol[Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE].shape[0], ) .T.double() .reshape(-1) @@ -311,20 +309,20 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: return (exc_density * grid_weights).sum() - def get_exc_density(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc_density(self, mol: FeatureMap) -> torch.Tensor: padded = self._get_exc_density_padded(mol) - sizes = mol["atomic_grid_sizes"] - size_bound = mol["atomic_grid_size_bound_shape"].shape[0] + sizes = mol[Feature.ATOMIC_GRID_SIZES] + size_bound = mol[Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE].shape[0] num_atoms = sizes.shape[0] - total_grid_points = mol["grid_weights"].shape[0] + total_grid_points = mol[Feature.GRID_WEIGHTS].shape[0] padded_2d = padded.reshape(size_bound, num_atoms).T return unpad_ragged(padded_2d, sizes, total_grid_points) - def _get_exc_density_padded(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def _get_exc_density_padded(self, mol: FeatureMap) -> torch.Tensor: mol = self.pack_features(mol) - grid_coords = mol["grid_coords"] - atomic_grid_weights = mol["atomic_grid_weights"] - coarse_coords = mol["coarse_0_atomic_coords"] + grid_coords = mol[Feature.GRID_COORDS] + atomic_grid_weights = mol[Feature.ATOMIC_GRID_WEIGHTS] + coarse_coords = mol[Feature.COARSE_0_ATOMIC_COORDS] features_ab, features_ba = self.semi_local_features(mol) # Learned symmetrized features @@ -352,7 +350,7 @@ def _get_exc_density_padded(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: directions ) # (num_fine, num_coarse, (lmax+1)^2) - exp_m1_rho_total = torch.exp(-mol["density"].sum(0).unsqueeze(-1)).to( + exp_m1_rho_total = torch.exp(-mol[Feature.DENSITY].sum(0).unsqueeze(-1)).to( self.dtype ) @@ -368,7 +366,7 @@ def _get_exc_density_padded(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: enhancement_factor = self.output_model(features) return enhancement_density_inner_product( enhancement_factor=enhancement_factor.view(-1, 1), - density=mol["density"].reshape(2, -1), + density=mol[Feature.DENSITY].reshape(2, -1), ) def reset_parameters(self) -> None: diff --git a/src/skala/functional/traditional.py b/src/skala/functional/traditional.py index 04fa05a6..0097752f 100644 --- a/src/skala/functional/traditional.py +++ b/src/skala/functional/traditional.py @@ -12,6 +12,7 @@ import torch from torch import Tensor, nn +from skala.features import Feature, FeatureMap from skala.functional import density from skala.functional.base import ExcFunctionalBase @@ -27,7 +28,7 @@ class SpinScaledXCFunctional(ExcFunctionalBase): def get_d3_settings(self) -> str: return self.__class__.__name__.lower() - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: + def exchange(self, mol_features: FeatureMap) -> Tensor: """ Compute the exchange energy density. @@ -43,7 +44,7 @@ def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: """ raise NotImplementedError() - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: + def correlation_density(self, mol_features: FeatureMap) -> Tensor: """ Compute the correlation energy density. @@ -59,7 +60,7 @@ def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: """ raise NotImplementedError() - def correlation(self, mol_features: dict[str, Tensor]) -> Tensor: + def correlation(self, mol_features: FeatureMap) -> Tensor: """ Compute the correlation energy. @@ -73,10 +74,10 @@ def correlation(self, mol_features: dict[str, Tensor]) -> Tensor: Tensor Correlation energy. """ - rho_total = mol_features["density"].sum(0) + rho_total = mol_features[Feature.DENSITY].sum(0) return rho_total * self.correlation_density(mol_features) - def get_exc_density(self, mol: dict[str, Tensor]) -> Tensor: + def get_exc_density(self, mol: FeatureMap) -> Tensor: exch = self.exchange(density.scale_by(mol, 2)).sum(0) / 2 corr = self.correlation(mol) return exch + corr @@ -90,15 +91,21 @@ class LDA(SpinScaledXCFunctional): Exchange: E_x[ρ] = -3/4 * (3/π)^(1/3) * ρ^(4/3) """ - features = ["density", "grid_weights"] + features = [ + Feature.DENSITY, + Feature.GRID_WEIGHTS, + ] - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: + def exchange(self, mol_features: FeatureMap) -> Tensor: return ( - -3 / 4 * (3 / math.pi) ** (1 / 3) * mol_features["density"].abs() ** (4 / 3) + -3 + / 4 + * (3 / math.pi) ** (1 / 3) + * mol_features[Feature.DENSITY].abs() ** (4 / 3) ) - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: - return mol_features["density"].new_zeros((1,)) + def correlation_density(self, mol_features: FeatureMap) -> Tensor: + return mol_features[Feature.DENSITY].new_zeros((1,)) class SPW92(SpinScaledXCFunctional): @@ -109,14 +116,20 @@ class SPW92(SpinScaledXCFunctional): correlation energy of the uniform electron gas. """ - features = ["density", "grid_weights"] + features = [ + Feature.DENSITY, + Feature.GRID_WEIGHTS, + ] - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: + def exchange(self, mol_features: FeatureMap) -> Tensor: return ( - -3 / 4 * (3 / math.pi) ** (1 / 3) * mol_features["density"].abs() ** (4 / 3) + -3 + / 4 + * (3 / math.pi) ** (1 / 3) + * mol_features[Feature.DENSITY].abs() ** (4 / 3) ) - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: + def correlation_density(self, mol_features: FeatureMap) -> Tensor: def Gamma( rs: Tensor, A: float, a1: float, b1: float, b2: float, b3: float, b4: float ) -> Tensor: @@ -124,7 +137,7 @@ def Gamma( poly = (b1 + (b2 + (b3 + b4 * rs_sq) * rs_sq) * rs_sq) * rs_sq return -2 * A * (1 + a1 * rs) * torch.log(1 + 0.5 / (A * poly)) - rho = mol_features["density"] + rho = mol_features[Feature.DENSITY] zeta, rho_total = density.zeta(rho), rho.sum(0) ff0 = 1.709921 ff = ((1 + zeta) ** (4 / 3) + (1 - zeta) ** (4 / 3) - 2) / (2 ** (4 / 3) - 2) @@ -147,7 +160,11 @@ class PBE(SpinScaledXCFunctional): and correlation gradient corrections to the local density approximation. """ - features = ["density", "grad", "grid_weights"] + features = [ + Feature.DENSITY, + Feature.GRAD, + Feature.GRID_WEIGHTS, + ] def __init__(self) -> None: super().__init__() @@ -156,9 +173,9 @@ def __init__(self) -> None: self.kappa = nn.Parameter(torch.tensor(0.804), requires_grad=False) self.mu = self.beta * (math.pi**2 / 3) - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: - rho = mol_features["density"] - grad = mol_features["grad"] + def exchange(self, mol_features: FeatureMap) -> Tensor: + rho = mol_features[Feature.DENSITY] + grad = mol_features[Feature.GRAD] FX = ( 1 + self.kappa @@ -167,10 +184,10 @@ def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: ) return self.lda.exchange(mol_features) * FX - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: + def correlation_density(self, mol_features: FeatureMap) -> Tensor: eps_c_unif = self.lda.correlation_density(mol_features) - rho = mol_features["density"] - grad = mol_features["grad"] + rho = mol_features[Feature.DENSITY] + grad = mol_features[Feature.GRAD] rho_total, grad_total = rho.sum(0), grad.sum(0) zeta = density.zeta(rho) ks = torch.sqrt(4 * density.kF(rho_total) / math.pi) @@ -200,7 +217,12 @@ class TPSS(SpinScaledXCFunctional): exact constraints of density functional theory. """ - features = ["density", "kin", "grad", "grid_weights"] + features = [ + Feature.DENSITY, + Feature.KIN, + Feature.GRAD, + Feature.GRID_WEIGHTS, + ] def __init__(self) -> None: super().__init__() @@ -211,10 +233,10 @@ def __init__(self) -> None: self.b = nn.Parameter(torch.tensor(0.40), requires_grad=False) self.d = nn.Parameter(torch.tensor(2.8), requires_grad=False) - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: - rho = mol_features["density"] - grad = mol_features["grad"] - kin = mol_features["kin"] + def exchange(self, mol_features: FeatureMap) -> Tensor: + rho = mol_features[Feature.DENSITY] + grad = mol_features[Feature.GRAD] + kin = mol_features[Feature.KIN] # p is the reduced gradient squared, z is the zeta value p, z = density.reduced_gradient(rho, grad) ** 2, density.z(rho, grad, kin) alpha = (5 * p / 3) * (1 / torch.clamp(z, density.EPS) - 1) @@ -233,10 +255,10 @@ def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: FX = 1 + kappa - kappa / (1 + x / kappa) return self.lda.exchange(mol_features) * FX - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: - rho = mol_features["density"] - grad = mol_features["grad"] - kin = mol_features["kin"] + def correlation_density(self, mol_features: FeatureMap) -> Tensor: + rho = mol_features[Feature.DENSITY] + grad = mol_features[Feature.GRAD] + kin = mol_features[Feature.KIN] rho_total, grad_total, kin_total = rho.sum(0), grad.sum(0), kin.sum(0) zeta, grad_zeta = density.zeta(rho), density.grad_zeta(rho, grad).norm(dim=-2) @@ -253,7 +275,7 @@ def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: z = density.z(rho_total, grad_total, kin_total) mols = density.separate(mol_features) eps_c_revpkzb = eps_c_pbe * (1 + Czetaxi * z**2) - (1 + Czetaxi) * z**2 * sum( - (mols[spin]["density"][spin] / rho_total) + (mols[spin][Feature.DENSITY][spin] / rho_total) * torch.max(eps_c_pbe, self.pbe.correlation_density(mols[spin])) for spin in range(2) ) @@ -261,7 +283,12 @@ def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: class _SCANLikeFunctional(SpinScaledXCFunctional): - features = ["density", "kin", "grad", "grid_weights"] + features = [ + Feature.DENSITY, + Feature.KIN, + Feature.GRAD, + Feature.GRID_WEIGHTS, + ] def __init__( self, alpha_mode: int, interpolation_mode: int, gradient_correction_mode: int @@ -677,16 +704,16 @@ def _scan_correlation_per_particle( energy = ec1 + ief * (ec0 - ec1) return torch.where(total_density > 0, energy, torch.zeros_like(energy)) - def exchange(self, mol_features: dict[str, Tensor]) -> Tensor: - rho = torch.clamp(mol_features["density"], min=0.0) - grad_norm = density.grad_norm(mol_features["grad"]) - kin = torch.clamp(mol_features["kin"], min=0.0) + def exchange(self, mol_features: FeatureMap) -> Tensor: + rho = torch.clamp(mol_features[Feature.DENSITY], min=0.0) + grad_norm = density.grad_norm(mol_features[Feature.GRAD]) + kin = torch.clamp(mol_features[Feature.KIN], min=0.0) return self._scan_exchange_density(rho, grad_norm, kin) - def correlation_density(self, mol_features: dict[str, Tensor]) -> Tensor: - rho = torch.clamp(mol_features["density"], min=0.0) - grad = mol_features["grad"] - kin = torch.clamp(mol_features["kin"], min=0.0) + def correlation_density(self, mol_features: FeatureMap) -> Tensor: + rho = torch.clamp(mol_features[Feature.DENSITY], min=0.0) + grad = mol_features[Feature.GRAD] + kin = torch.clamp(mol_features[Feature.KIN], min=0.0) return self._scan_correlation_per_particle(rho, grad, kin) diff --git a/src/skala/gpu4pyscf/gradients.py b/src/skala/gpu4pyscf/gradients.py index 7d1fc0ad..cc251124 100644 --- a/src/skala/gpu4pyscf/gradients.py +++ b/src/skala/gpu4pyscf/gradients.py @@ -18,6 +18,7 @@ from torch.utils.dlpack import from_dlpack import skala.pyscf.features as feature +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase LOG = logging.getLogger(__name__) @@ -28,7 +29,7 @@ def veff_and_expl_nuc_grad( mol: gto.Mole, grid: dft.Grids, rdm1: torch.Tensor, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ returns: @@ -37,21 +38,21 @@ def veff_and_expl_nuc_grad( """ SUPPORTED_FEATS = { - "density", - "grad", - "kin", - "grid_coords", - "grid_weights", - "atomic_grid_weights", - "coarse_0_atomic_coords", + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.COARSE_0_ATOMIC_COORDS, } if nuc_grad_feats is None: # generate feature list from functional features nuc_grad_feats = set(functional.features) # Integer-valued features have no nuclear gradient — always discard them - nuc_grad_feats.discard("atomic_grid_sizes") - nuc_grad_feats.discard("atomic_grid_size_bound_shape") + nuc_grad_feats.discard(Feature.ATOMIC_GRID_SIZES) + nuc_grad_feats.discard(Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE) # check for unsupported features unsupported_feats = {feat for feat in nuc_grad_feats if feat not in SUPPORTED_FEATS} @@ -63,9 +64,9 @@ def veff_and_expl_nuc_grad( LOG.debug("nuc_grad_feats = %s", nuc_grad_feats) # determine the maximum ao derivative needed - if "grad" in nuc_grad_feats or "kin" in nuc_grad_feats: + if Feature.GRAD in nuc_grad_feats or Feature.KIN in nuc_grad_feats: ao_deriv = 2 - elif "density" in nuc_grad_feats: + elif Feature.DENSITY in nuc_grad_feats: ao_deriv = 1 else: ao_deriv = 0 @@ -87,7 +88,7 @@ def veff_and_expl_nuc_grad( # Discard atomic_grid_weights from VJP features: d(atomic_grid_weights)/dR = 0 # because they are raw quadrature weights that depend only on the radial/angular # grid rule, not on nuclear positions. They still pass through as other_feats. - nuc_grad_feats.discard("atomic_grid_weights") + nuc_grad_feats.discard(Feature.ATOMIC_GRID_WEIGHTS) # Get required derivatives nuc_feat_names = list(nuc_grad_feats) # ensure specific order @@ -118,7 +119,7 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: ) else: dExc_tuple = () - dExc: dict[str, torch.Tensor] = {} + dExc: FeatureMap = {} for i in range(len(dExc_tuple)): dExc[nuc_feat_names[i]] = dExc_tuple[i].detach() @@ -146,16 +147,16 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: # Calculate the contribution to veff for this atomic grid veff_atm = torch.zeros((2, 3, nao, nao), dtype=rdm1.dtype, device=rdm1.device) - if "density" in nuc_grad_feats: + if Feature.DENSITY in nuc_grad_feats: veff_atm += torch.einsum( "si, xip, iq -> sxpq", - dExc["density"][:, atm_start:atm_end], + dExc[Feature.DENSITY][:, atm_start:atm_end], ao[1:4], ao[0], ) - if "grad" in nuc_grad_feats: - Exc_dgrad_atm = dExc["grad"][:, :, atm_start:atm_end] + if Feature.GRAD in nuc_grad_feats: + Exc_dgrad_atm = dExc[Feature.GRAD][:, :, atm_start:atm_end] veff_atm += torch.einsum( "syi, xip, yiq -> sxpq", Exc_dgrad_atm, ao[1:4], ao[1:4] @@ -191,8 +192,8 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[9], ao[0] ) - if "kin" in nuc_grad_feats: - Exc_dkin_atm = dExc["kin"][:, atm_start:atm_end] + if Feature.KIN in nuc_grad_feats: + Exc_dkin_atm = dExc[Feature.KIN][:, atm_start:atm_end] # XX, XY, XZ = 4, 5, 6 veff_atm[:, 0] += ( torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[4], ao[1]) / 2 @@ -224,12 +225,12 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[9], ao[3]) / 2 ) - if "grid_coords" in nuc_grad_feats: + if Feature.GRID_COORDS in nuc_grad_feats: # also add the explicit grid coordinate dependence - nuc_grad[atm_id] += dExc["grid_coords"][atm_start:atm_end].sum(dim=0) + nuc_grad[atm_id] += dExc[Feature.GRID_COORDS][atm_start:atm_end].sum(dim=0) - if "grid_weights" in nuc_grad_feats: - Exc_dgw = dExc["grid_weights"][atm_start:atm_end] + if Feature.GRID_WEIGHTS in nuc_grad_feats: + Exc_dgw = dExc[Feature.GRID_WEIGHTS][atm_start:atm_end] nuc_grad += from_dlpack(weight1) @ Exc_dgw # add the grid coordinate dependence via the density-like quantities to the nuclear gradient # we get those from the veff block. This tends to largely cancel with the grid_weights derivative, @@ -242,8 +243,8 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: veff += veff_atm atm_start = atm_end - if "coarse_0_atomic_coords" in nuc_grad_feats: - nuc_grad += dExc["coarse_0_atomic_coords"] + if Feature.COARSE_0_ATOMIC_COORDS in nuc_grad_feats: + nuc_grad += dExc[Feature.COARSE_0_ATOMIC_COORDS] # finalize if len(rdm1.shape) == 2: @@ -270,7 +271,7 @@ def nuc_grad_from_veff( class SkalaRKSGradient(RHFGradient): # type: ignore[misc] functional: ExcFunctionalBase """Skala functional""" - nuc_grad_feats: set[str] | None + nuc_grad_feats: set[Feature] | None """Which partial derivatives to take into account. None defaults to all.""" veff_nuc_grad_: torch.Tensor | None """Contribution of the coordinate dependence of density, grad, kin, etc.""" @@ -281,7 +282,7 @@ def __init__( self, ks: SCF, verbose: bool = False, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ): super().__init__(ks) self.functional = ks._numint.func @@ -369,7 +370,7 @@ def reset(self, mol: gto.Mole | None = None) -> "SkalaRKSGradient": class SkalaUKSGradient(UHFGradient): # type: ignore[misc] functional: ExcFunctionalBase """Skala functional""" - nuc_grad_feats: set[str] | None + nuc_grad_feats: set[Feature] | None """Which partial derivatives to take into account. None defaults to all.""" veff_nuc_grad_: torch.Tensor | None """Contribution of the coordinate dependence of density, grad, kin, etc.""" @@ -380,7 +381,7 @@ def __init__( self, ks: SCF, verbose: bool = False, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ): super().__init__(ks) self.functional = ks._numint.func diff --git a/src/skala/pyscf/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py index c7ac53db..2cdaf1c8 100644 --- a/src/skala/pyscf/ao_evaluation.py +++ b/src/skala/pyscf/ao_evaluation.py @@ -3,8 +3,7 @@ """Blockwise atomic-orbital feature evaluation and custom autograd.""" from collections.abc import Callable, Iterator -from dataclasses import dataclass -from typing import Protocol, cast +from typing import NamedTuple, Protocol, TypeAlias, cast import numpy as np import torch @@ -12,8 +11,10 @@ from torch import Tensor from torch.autograd import Function from torch.autograd.function import FunctionCtx +from torch.utils.dlpack import from_dlpack from typing_extensions import Unpack +from skala.features import FeatureMap from skala.pyscf import feature_math from skala.pyscf.backend import ( Array, @@ -23,6 +24,9 @@ from_numpy_or_cupy, ) +_ScreenIndex: TypeAlias = np.ndarray[tuple[int, int], np.dtype[np.uint8]] +_AOIndices: TypeAlias = np.ndarray[tuple[int], np.dtype[np.intp]] + class _ChunkEvalForwardContext(Protocol): dm: Tensor @@ -45,7 +49,7 @@ class _ChunkEvalBackwardContext(Protocol): gpu: bool -def _active_cpu_ao_indices(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarray: +def _active_cpu_ao_indices(mol: gto.Mole, screen_index: _ScreenIndex) -> _AOIndices: """Expand active shells in a PySCF screen-index slice to AO indices. A shell is active for the grid block if it is nonzero in any of the @@ -57,32 +61,7 @@ def _active_cpu_ao_indices(mol: gto.Mole, screen_index: np.ndarray) -> np.ndarra return np.flatnonzero(np.repeat(active_shells, np.diff(ao_loc))) -def partial_feature_function_over_ao_values( - feature_function: Callable[[torch.Tensor, torch.Tensor], torch.Tensor], - ao_values: torch.Tensor, -) -> Callable[[torch.Tensor], torch.Tensor]: - """Bind evaluated AO values to a feature function for one grid block.""" - - def partial_feature_function(dm: torch.Tensor) -> torch.Tensor: - return feature_function(dm, ao_values) - - return partial_feature_function - - -def partial_vjp_function_over_tangents( - func: Callable[[torch.Tensor], torch.Tensor], - tangents: torch.Tensor, -) -> Callable[[torch.Tensor], torch.Tensor]: - """Bind feature cotangents to a function for block-local VJP evaluation.""" - - def reduced_vjp(primals: torch.Tensor) -> torch.Tensor: - return torch.func.vjp(func, primals)[1](tangents)[0] - - return reduced_vjp - - -@dataclass(frozen=True) -class _AOBlock: +class _AOBlock(NamedTuple): """Evaluated AO data and index metadata for one contiguous grid block. ``ao_values`` contains only the active AO rows when screening is enabled. @@ -122,17 +101,22 @@ def _evaluate_feature_block( feature_cotangent: Tensor | None = None, ) -> Tensor: """Evaluate one active-AO feature block or its feature-space VJP.""" - partial_func = partial_feature_function_over_ao_values( - feature_function, block.ao_values - ) + + def evaluate_features(dm: Tensor) -> Tensor: + return feature_function(dm, block.ao_values) + + evaluation_function: Callable[[Tensor], Tensor] = evaluate_features if feature_cotangent is not None: - partial_func = partial_vjp_function_over_tangents( - partial_func, feature_cotangent[..., block.grid_slice] - ) + local_cotangent = feature_cotangent[..., block.grid_slice] + + def evaluate_vjp(dm: Tensor) -> Tensor: + return torch.func.vjp(evaluate_features, dm)[1](local_cotangent)[0] + + evaluation_function = evaluate_vjp if compile_feature_function: - return torch.compile(partial_func)(active_dm_submatrix) - return partial_func(active_dm_submatrix) + return torch.compile(evaluation_function)(active_dm_submatrix) + return evaluation_function(active_dm_submatrix) class _CPUAOBlockLoop: @@ -183,7 +167,7 @@ def restore_ao_order(self, matrix: Tensor) -> Tensor: def _active_ao_indices( self, - non0tab: np.ndarray, + non0tab: _ScreenIndex, grid_start: int, grid_end: int, ) -> Tensor | None: @@ -213,11 +197,7 @@ def _active_ao_indices( block_non0tab = non0tab[row_start:row_end] if np.all(np.any(block_non0tab, axis=0)): return None - return torch.as_tensor( - _active_cpu_ao_indices(self.mol, block_non0tab), - device=self.dm.device, - dtype=torch.long, - ) + return torch.from_numpy(_active_cpu_ao_indices(self.mol, block_non0tab)) def __iter__(self) -> Iterator[_AOBlock]: non0tab = self.grids.non0tab @@ -232,11 +212,7 @@ def __iter__(self) -> Iterator[_AOBlock]: non0tab=non0tab, ): start, end = end, end + block_weights.size - ao_values = ( - torch.from_numpy(backend_ao_values) - .to(device=self.dm.device, dtype=self.dm.dtype) - .transpose(-1, -2) - ) + ao_values = torch.from_numpy(backend_ao_values).transpose(-1, -2) active_ao_indices = ( None if non0tab is None @@ -301,8 +277,8 @@ def __iter__(self) -> Iterator[_AOBlock]: if active_ao_indices.size == 0: continue yield _AOBlock( - torch.from_dlpack(backend_ao_values), - torch.from_dlpack(active_ao_indices), + from_dlpack(backend_ao_values), + from_dlpack(active_ao_indices), slice(start, end), ) @@ -559,7 +535,7 @@ def backward( return tuple(grads) -def non_chunk( +def evaluate_full_grid( dm: torch.Tensor, mol: gto.Mole, coords: Array, @@ -593,9 +569,9 @@ def _resolve_ao_block_size( ) -> int | None: """Resolve an aligned CPU block size or delegate GPU sizing to its backend.""" if gpu: - if block_size is not None: - raise ValueError("Setting custom block size is not supported on GPU.") - return None + if block_size is None: + return None + raise ValueError("Setting custom block size is not supported on GPU.") if block_size is None: nao = mol.nao_nr() @@ -620,7 +596,7 @@ def auto_chunk( block_size: int | None = None, max_memory: int = 2000, gpu: bool = False, -) -> dict[str, torch.Tensor]: +) -> FeatureMap: """Evaluate raw features with a memory-derived or explicit AO block size.""" if gpu: check_gpu_imports_were_successful() @@ -630,20 +606,9 @@ def auto_chunk( blksize = _resolve_ao_block_size(mol, feature_function, block_size, max_memory, gpu) if blksize is not None and blksize >= grids.weights.shape[0]: - features = non_chunk( - dm.double(), - mol, - grids.coords, - feature_function, - ) + features = evaluate_full_grid(dm.double(), mol, grids.coords, feature_function) else: features = ChunkEvalForward.apply( - dm.double(), - mol, - grids, - feature_function, - blksize, - False, - gpu, + dm.double(), mol, grids, feature_function, blksize, False, gpu ) return feature_function.to_dict(features) diff --git a/src/skala/pyscf/evaluation.py b/src/skala/pyscf/evaluation.py index 9c4eb495..042fd915 100644 --- a/src/skala/pyscf/evaluation.py +++ b/src/skala/pyscf/evaluation.py @@ -5,53 +5,72 @@ from collections.abc import Iterable from dataclasses import dataclass -_MGGA_FEATURES = frozenset({"density", "grad", "kin", "lapl"}) +from skala.features import Feature + +_AO_FEATURES = frozenset( + { + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.LAPL, + } +) _ATOMIC_LAYOUT_FEATURES = frozenset( { - "atomic_grid_weights", - "atomic_grid_sizes", - "atomic_grid_size_bound_shape", + Feature.ATOMIC_GRID_WEIGHTS, + Feature.ATOMIC_GRID_SIZES, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, } ) -@dataclass(frozen=True, init=False) class FeatureSpec: """Normalized feature names and their evaluation requirements.""" - names: frozenset[str] + def __init__(self, names: Iterable[Feature]) -> None: + self._names = frozenset(names) + + @property + def names(self) -> frozenset[Feature]: + """Return the normalized feature names.""" + return self._names + + def __eq__(self, other: object) -> bool: + if not isinstance(other, FeatureSpec): + return NotImplemented + return self.names == other.names - def __init__(self, names: Iterable[str]) -> None: - object.__setattr__(self, "names", frozenset(names)) + def __hash__(self) -> int: + return hash(self.names) - def requests(self, feature: str) -> bool: + def requests(self, feature: Feature) -> bool: """Return whether a feature is requested.""" return feature in self.names @property def with_density(self) -> bool: """Return whether density is requested.""" - return self.requests("density") + return self.requests(Feature.DENSITY) @property def with_grad(self) -> bool: """Return whether the density gradient is requested.""" - return self.requests("grad") + return self.requests(Feature.GRAD) @property def with_kin(self) -> bool: """Return whether kinetic-energy density is requested.""" - return self.requests("kin") + return self.requests(Feature.KIN) @property def with_lapl(self) -> bool: """Return whether the density Laplacian is requested.""" - return self.requests("lapl") + return self.requests(Feature.LAPL) @property - def requires_mgga(self) -> bool: - """Return whether AO-based meta-GGA features are requested.""" - return bool(self.names & _MGGA_FEATURES) + def requires_ao_evaluation(self) -> bool: + """Return whether AO-derived features are requested.""" + return bool(self.names & _AO_FEATURES) @property def mgga_feature_count(self) -> int: @@ -61,9 +80,9 @@ def mgga_feature_count(self) -> int: @property def ao_derivative_order(self) -> int: """Return the highest AO derivative order needed by the features.""" - if "lapl" in self.names: + if Feature.LAPL in self.names: return 2 - if self.names & {"grad", "kin"}: + if self.names & {Feature.GRAD, Feature.KIN}: return 1 return 0 @@ -75,7 +94,7 @@ def requires_atomic_layout(self) -> bool: @property def supports_screened_evaluation(self) -> bool: """Return whether atom-aligned screened evaluation is supported.""" - return "atomic_grid_sizes" in self.names + return Feature.ATOMIC_GRID_SIZES in self.names @dataclass(frozen=True) diff --git a/src/skala/pyscf/feature_math.py b/src/skala/pyscf/feature_math.py index 5c1cf617..f16d278c 100644 --- a/src/skala/pyscf/feature_math.py +++ b/src/skala/pyscf/feature_math.py @@ -7,6 +7,7 @@ import torch from torch import nn +from skala.features import Feature, FeatureMap from skala.pyscf.evaluation import FeatureSpec @@ -29,7 +30,7 @@ class FeatureFunction(nn.Module, ABC): def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ... @abstractmethod - def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: ... + def to_dict(self, features: torch.Tensor) -> FeatureMap: ... class MGGAFeatureFunction(FeatureFunction): @@ -38,27 +39,29 @@ class MGGAFeatureFunction(FeatureFunction): def __init__(self, feature_spec: FeatureSpec): super().__init__() - if not feature_spec.requires_mgga: - raise ValueError("At least one feature must be selected.") + if not feature_spec.requires_ao_evaluation: + raise ValueError("At least one AO-derived feature must be selected.") self.feature_spec = feature_spec self.deriv = feature_spec.ao_derivative_order self.nfeats = feature_spec.mgga_feature_count - def to_dict(self, features: torch.Tensor) -> dict[str, torch.Tensor]: + def to_dict(self, features: torch.Tensor) -> FeatureMap: """Convert a packed feature tensor to its named feature tensors.""" feature_index = 0 - feature_dict: dict[str, torch.Tensor] = {} + feature_dict: FeatureMap = {} if self.feature_spec.with_density: - feature_dict["density"] = features[..., feature_index, :] + feature_dict[Feature.DENSITY] = features[..., feature_index, :] feature_index += 1 if self.feature_spec.with_grad: - feature_dict["grad"] = features[..., feature_index : feature_index + 3, :] + feature_dict[Feature.GRAD] = features[ + ..., feature_index : feature_index + 3, : + ] feature_index += 3 if self.feature_spec.with_kin: - feature_dict["kin"] = features[..., feature_index, :] + feature_dict[Feature.KIN] = features[..., feature_index, :] feature_index += 1 if self.feature_spec.with_lapl: - feature_dict["lapl"] = features[..., feature_index, :] + feature_dict[Feature.LAPL] = features[..., feature_index, :] return feature_dict def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: @@ -82,13 +85,13 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: feature_index = 0 if self.feature_spec.with_density: - features[..., feature_index, :] = torch.sum(c0 * ao[0][None, :, :], dim=-2) + features[..., feature_index, :] = torch.sum(c0 * ao[0, None, :, :], dim=-2) feature_index += 1 if self.feature_spec.with_grad: for component in range(3): features[..., feature_index, :] = 2 * torch.sum( - c0 * ao[component + 1][None, :, :], dim=-2 + c0 * ao[component + 1, None, :, :], dim=-2 ) feature_index += 1 @@ -96,7 +99,7 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: for component in range(3): ci = dm_view @ ao[component + 1] features[..., feature_index, :] += 0.5 * torch.sum( - ci * ao[component + 1][None, :, :], dim=-2 + ci * ao[component + 1, None, :, :], dim=-2 ) if self.feature_spec.with_kin: @@ -111,7 +114,7 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: if self.feature_spec.with_lapl: for component in (4, 7, 9): features[..., feature_index, :] += 2 * torch.sum( - c0 * ao[component][None, :, :], dim=-2 + c0 * ao[component, None, :, :], dim=-2 ) if len(dm.shape) == 2: diff --git a/src/skala/pyscf/features.py b/src/skala/pyscf/features.py index e27fe2c4..751c7dfd 100644 --- a/src/skala/pyscf/features.py +++ b/src/skala/pyscf/features.py @@ -9,11 +9,18 @@ from pyscf import gto from torch import Tensor +from skala.features import Feature, FeatureMap from skala.pyscf import ao_evaluation, feature_math from skala.pyscf.backend import Grid, from_numpy_or_cupy from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec -DEFAULT_FEATURES = ["density", "kin", "grad", "grid_coords", "grid_weights"] +DEFAULT_FEATURES = [ + Feature.DENSITY, + Feature.KIN, + Feature.GRAD, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, +] DEFAULT_FEATURES_SET = set(DEFAULT_FEATURES) @@ -21,11 +28,11 @@ def generate_features( mol: gto.Mole, dm: Tensor, grids: Grid, - features: set[str] | None = None, + features: set[Feature] | None = None, chunk_size: int | None = None, max_memory: int = 2000, gpu: bool = False, -) -> dict[str, Tensor]: +) -> FeatureMap: """Generate density features for a given molecule. The density features are stored in a dictionary with the keys matching the requested features. @@ -56,14 +63,14 @@ def generate_features( evaluation_policy = EvaluationPolicy(ao_block_size=chunk_size) # if dm is a 3D tensor, then we have a spin-polarized system - with_spin = len(dm.shape) == 3 + is_spin_polarized = len(dm.shape) == 3 if gpu and dm.device.type != "cuda": raise ValueError("Density matrix must be on the GPU when gpu=True.") mol_features = get_grid_features(mol, dm, grids, feature_spec) - if feature_spec.requires_mgga: + if feature_spec.requires_ao_evaluation: mgga_features = ao_evaluation.auto_chunk( dm, mol, @@ -76,7 +83,7 @@ def generate_features( for feature in mgga_features: mol_features[feature] = feature_math.maybe_expand_and_divide( - mgga_features[feature], not with_spin, 2 + mgga_features[feature], not is_spin_polarized, 2 ) return mol_features @@ -87,21 +94,21 @@ def get_grid_features( dm: Tensor, grids: Grid, feature_spec: FeatureSpec, -) -> dict[str, Tensor]: - grid_features = {} +) -> FeatureMap: + grid_features: FeatureMap = {} - if feature_spec.requests("grid_coords"): - grid_features["grid_coords"] = from_numpy_or_cupy( + if feature_spec.requests(Feature.GRID_COORDS): + grid_features[Feature.GRID_COORDS] = from_numpy_or_cupy( grids.coords, device=dm.device, dtype=dm.dtype ) - if feature_spec.requests("grid_weights"): - grid_features["grid_weights"] = from_numpy_or_cupy( + if feature_spec.requests(Feature.GRID_WEIGHTS): + grid_features[Feature.GRID_WEIGHTS] = from_numpy_or_cupy( grids.weights, device=dm.device, dtype=dm.dtype ) - if feature_spec.requests("coarse_0_atomic_coords"): - grid_features["coarse_0_atomic_coords"] = from_numpy_or_cupy( + if feature_spec.requests(Feature.COARSE_0_ATOMIC_COORDS): + grid_features[Feature.COARSE_0_ATOMIC_COORDS] = from_numpy_or_cupy( mol.atom_coords(), device=dm.device, dtype=dm.dtype ) @@ -121,22 +128,22 @@ def get_grid_features( f"Set grids.alignment = 1 before building grids to disable padding." ) - if feature_spec.requests("atomic_grid_sizes"): - grid_features["atomic_grid_sizes"] = torch.tensor( + if feature_spec.requests(Feature.ATOMIC_GRID_SIZES): + grid_features[Feature.ATOMIC_GRID_SIZES] = torch.tensor( sizes, dtype=torch.long, device=dm.device ) - if feature_spec.requests("atomic_grid_size_bound_shape"): + if feature_spec.requests(Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE): max_size = max(sizes) - grid_features["atomic_grid_size_bound_shape"] = torch.zeros( + grid_features[Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE] = torch.zeros( max_size, 0, dtype=torch.long, device=dm.device ) - if feature_spec.requests("atomic_grid_weights"): + if feature_spec.requests(Feature.ATOMIC_GRID_WEIGHTS): raw_weights = np.concatenate( [atom_grids_tab[mol.atom_symbol(ia)][1] for ia in range(mol.natm)] ) - grid_features["atomic_grid_weights"] = from_numpy_or_cupy( + grid_features[Feature.ATOMIC_GRID_WEIGHTS] = from_numpy_or_cupy( raw_weights, device=dm.device, dtype=dm.dtype ) diff --git a/src/skala/pyscf/gradients.py b/src/skala/pyscf/gradients.py index a1037d36..abd883e2 100644 --- a/src/skala/pyscf/gradients.py +++ b/src/skala/pyscf/gradients.py @@ -15,6 +15,7 @@ from pyscf.scf.hf import SCF import skala.pyscf.features as feature +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase LOG = logging.getLogger(__name__) @@ -25,7 +26,7 @@ def veff_and_expl_nuc_grad( mol: gto.Mole, grid: dft.Grids, rdm1: torch.Tensor, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ returns: @@ -34,21 +35,21 @@ def veff_and_expl_nuc_grad( """ SUPPORTED_FEATS = { - "density", - "grad", - "kin", - "grid_coords", - "grid_weights", - "atomic_grid_weights", - "coarse_0_atomic_coords", + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.COARSE_0_ATOMIC_COORDS, } if nuc_grad_feats is None: # generate feature list from functional features nuc_grad_feats = set(functional.features) # Integer-valued features have no nuclear gradient — always discard them - nuc_grad_feats.discard("atomic_grid_sizes") - nuc_grad_feats.discard("atomic_grid_size_bound_shape") + nuc_grad_feats.discard(Feature.ATOMIC_GRID_SIZES) + nuc_grad_feats.discard(Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE) # check for unsupported features unsupported_feats = {feat for feat in nuc_grad_feats if feat not in SUPPORTED_FEATS} @@ -60,9 +61,9 @@ def veff_and_expl_nuc_grad( LOG.debug("nuc_grad_feats = %s", nuc_grad_feats) # determine the maximum ao derivative needed - if "grad" in nuc_grad_feats or "kin" in nuc_grad_feats: + if Feature.GRAD in nuc_grad_feats or Feature.KIN in nuc_grad_feats: ao_deriv = 2 - elif "density" in nuc_grad_feats: + elif Feature.DENSITY in nuc_grad_feats: ao_deriv = 1 else: ao_deriv = 0 @@ -82,7 +83,7 @@ def veff_and_expl_nuc_grad( # Discard atomic_grid_weights from VJP features: d(atomic_grid_weights)/dR = 0 # because they are raw quadrature weights that depend only on the radial/angular # grid rule, not on nuclear positions. They still pass through as other_feats. - nuc_grad_feats.discard("atomic_grid_weights") + nuc_grad_feats.discard(Feature.ATOMIC_GRID_WEIGHTS) # Get required derivatives nuc_feat_names = list(nuc_grad_feats) # ensure specific order @@ -99,7 +100,7 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: _, dExc_func = torch.func.vjp(exc_feat_func, *nuc_feat_tensors) dExc_tuple = dExc_func(torch.tensor(1.0, dtype=rdm1.dtype)) - dExc: dict[str, torch.Tensor] = {} + dExc: FeatureMap = {} for i in range(len(dExc_tuple)): dExc[nuc_feat_names[i]] = dExc_tuple[i].detach() @@ -124,16 +125,16 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: # Calculate the contribution to veff for this atomic grid veff_atm = torch.zeros((2, 3, nao, nao), dtype=rdm1.dtype) - if "density" in nuc_grad_feats: + if Feature.DENSITY in nuc_grad_feats: veff_atm += torch.einsum( "si, xip, iq -> sxpq", - dExc["density"][:, atm_start:atm_end], + dExc[Feature.DENSITY][:, atm_start:atm_end], ao[1:4], ao[0], ) - if "grad" in nuc_grad_feats: - Exc_dgrad_atm = dExc["grad"][:, :, atm_start:atm_end] + if Feature.GRAD in nuc_grad_feats: + Exc_dgrad_atm = dExc[Feature.GRAD][:, :, atm_start:atm_end] veff_atm += torch.einsum( "syi, xip, yiq -> sxpq", Exc_dgrad_atm, ao[1:4], ao[1:4] @@ -169,8 +170,8 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[9], ao[0] ) - if "kin" in nuc_grad_feats: - Exc_dkin_atm = dExc["kin"][:, atm_start:atm_end] + if Feature.KIN in nuc_grad_feats: + Exc_dkin_atm = dExc[Feature.KIN][:, atm_start:atm_end] # XX, XY, XZ = 4, 5, 6 veff_atm[:, 0] += ( torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[4], ao[1]) / 2 @@ -202,12 +203,12 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[9], ao[3]) / 2 ) - if "grid_coords" in nuc_grad_feats: + if Feature.GRID_COORDS in nuc_grad_feats: # also add the explicit grid coordinate dependence - nuc_grad[atm_id] += dExc["grid_coords"][atm_start:atm_end].sum(dim=0) + nuc_grad[atm_id] += dExc[Feature.GRID_COORDS][atm_start:atm_end].sum(dim=0) - if "grid_weights" in nuc_grad_feats: - Exc_dgw = dExc["grid_weights"][atm_start:atm_end] + if Feature.GRID_WEIGHTS in nuc_grad_feats: + Exc_dgw = dExc[Feature.GRID_WEIGHTS][atm_start:atm_end] nuc_grad += torch.from_numpy(weight1) @ Exc_dgw # add the grid coordinate dependence via the density-like quantities to the nuclear gradient # we get those from the veff block. This tends to largely cancel with the grid_weights derivative, @@ -220,8 +221,8 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: veff += veff_atm atm_start = atm_end - if "coarse_0_atomic_coords" in nuc_grad_feats: - nuc_grad += dExc["coarse_0_atomic_coords"] + if Feature.COARSE_0_ATOMIC_COORDS in nuc_grad_feats: + nuc_grad += dExc[Feature.COARSE_0_ATOMIC_COORDS] # finalize if len(rdm1.shape) == 2: @@ -235,7 +236,7 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: class SkalaRKSGradient(RHFGradient): # type: ignore[misc] functional: ExcFunctionalBase """LivDFT functional""" - nuc_grad_feats: set[str] | None + nuc_grad_feats: set[Feature] | None """Which partial derivatives to take into account. None defaults to all.""" veff_nuc_grad_: torch.Tensor """Contribution of the coordinate dependence of density, grad, kin, etc.""" @@ -246,7 +247,7 @@ def __init__( self, ks: SCF, verbose: bool = False, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ): super().__init__(ks) self.functional = ks._numint.func @@ -312,7 +313,7 @@ def extra_force(self, atom_id: int, envs: dict[str, Any]) -> int: class SkalaUKSGradient(UHFGradient): # type: ignore[misc] functional: ExcFunctionalBase """LivDFT functional""" - nuc_grad_feats: set[str] | None + nuc_grad_feats: set[Feature] | None """Which partial derivatives to take into account. None defaults to all.""" veff_nuc_grad_: torch.Tensor """Contribution of the coordinate dependence of density, grad, kin, etc.""" @@ -323,7 +324,7 @@ def __init__( self, ks: SCF, verbose: bool = False, - nuc_grad_feats: set[str] | None = None, + nuc_grad_feats: set[Feature] | None = None, ): super().__init__(ks) self.functional = ks._numint.func diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index d6f203a4..e13dfb6a 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -3,7 +3,7 @@ import torch -def estimate_max_model_grid_points( +def estimate_max_gridpoint_chunk_size( dm: torch.Tensor, deriv: int, max_memory_in_mb: int | None = None, @@ -14,7 +14,7 @@ def estimate_max_model_grid_points( """Heuristically limit grid points per atom-aligned model evaluation. The dominant per-chunk allocation is the atomic-orbital matrix evaluated by - ``non_chunk`` (shape ``(ncomp, nao, n)`` in float64, with no AO screening), + ``evaluate_full_grid`` (shape ``(ncomp, nao, n)`` in float64, with no AO screening), together with the ``c0``/``ci`` products formed inside the feature function and retained by autograd for the backward pass. Peak memory is therefore modelled as affine in the number of grid points ``n`` (see diff --git a/src/skala/pyscf/model_chunking.py b/src/skala/pyscf/model_chunking.py index ed362cb7..6f08c88b 100644 --- a/src/skala/pyscf/model_chunking.py +++ b/src/skala/pyscf/model_chunking.py @@ -8,26 +8,27 @@ """ import logging -from collections.abc import Iterator +from collections.abc import Iterator, Mapping, Sequence from dataclasses import dataclass +from typing import NamedTuple import torch from pyscf import gto from torch import Tensor +from skala.features import Feature, FeatureMap from skala.pyscf import feature_math from skala.pyscf.backend import Grid from skala.pyscf.features import get_grid_features from skala.pyscf.memory_estimators import ( estimate_global_raw_feature_buffer_memory, - estimate_max_model_grid_points, + estimate_max_gridpoint_chunk_size, ) LOG = logging.getLogger(__name__) -@dataclass(frozen=True) -class AtomGridChunk: +class AtomGridChunk(NamedTuple): """Matching atom and grid slices for one model evaluation chunk.""" atom_slice: slice @@ -79,13 +80,12 @@ def _make_atom_grid_chunks( return chunks -@dataclass(frozen=True) -class ModelFeatureChunk: +class ModelFeatureChunk(NamedTuple): """Chunk-local raw features and the corresponding model input dictionary.""" grid_slice: slice raw_features: Tensor - model_features: dict[str, Tensor] + model_features: FeatureMap @dataclass(frozen=True) @@ -93,10 +93,10 @@ class ModelFeatureChunker: """Reusable atom-aligned partition of raw and model features.""" atom_major_raw_features: Tensor - grid_features: dict[str, Tensor] + grid_features: Mapping[Feature, Tensor] feature_function: feature_math.MGGAFeatureFunction - chunk_layouts: list[AtomGridChunk] - with_spin: bool + chunk_layouts: Sequence[AtomGridChunk] + is_spin_polarized: bool def __iter__(self) -> Iterator[ModelFeatureChunk]: """Yield detached raw features paired with atom-aligned model inputs.""" @@ -107,26 +107,29 @@ def __iter__(self) -> Iterator[ModelFeatureChunk]: .detach() .requires_grad_() ) - model_features: dict[str, Tensor] = {} + model_features: FeatureMap = {} for feature_name in ( - "grid_coords", - "grid_weights", - "atomic_grid_weights", + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, ): if feature_spec.requests(feature_name): model_features[feature_name] = self.grid_features[feature_name][ layout.grid_slice ] - for feature_name in ("coarse_0_atomic_coords", "atomic_grid_sizes"): + for feature_name in ( + Feature.COARSE_0_ATOMIC_COORDS, + Feature.ATOMIC_GRID_SIZES, + ): if feature_spec.requests(feature_name): model_features[feature_name] = self.grid_features[feature_name][ layout.atom_slice ] - if feature_spec.requests("atomic_grid_size_bound_shape"): - max_size = int(model_features["atomic_grid_sizes"].max().item()) - model_features["atomic_grid_size_bound_shape"] = torch.zeros( + if feature_spec.requests(Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE): + max_size = int(model_features[Feature.ATOMIC_GRID_SIZES].max().item()) + model_features[Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE] = torch.zeros( max_size, 0, dtype=torch.long, @@ -137,7 +140,7 @@ def __iter__(self) -> Iterator[ModelFeatureChunk]: raw_features ).items(): model_features[feature_name] = feature_math.maybe_expand_and_divide( - feature, not self.with_spin, 2 + feature, not self.is_spin_polarized, 2 ) yield ModelFeatureChunk( grid_slice=layout.grid_slice, @@ -152,30 +155,32 @@ def prepare_model_feature_chunks( grids: Grid, atom_major_raw_features: Tensor, feature_function: feature_math.MGGAFeatureFunction, - func_deriv: int, + deriv_order: int, max_memory_in_mb: int | None = None, safety_fraction: float = 0.8, ) -> ModelFeatureChunker: """Prepare memory-sized, atom-aligned chunks for functional model evaluation.""" feature_spec = feature_function.feature_spec if not feature_spec.supports_screened_evaluation: - raise ValueError("Atom-aligned model chunking requires 'atomic_grid_sizes'.") + raise ValueError( + f"Atom-aligned model chunking requires {Feature.ATOMIC_GRID_SIZES.value!r}." + ) grid_features = get_grid_features(mol, dm, grids, feature_spec) - max_model_grid_points = estimate_max_model_grid_points( + max_model_grid_points = estimate_max_gridpoint_chunk_size( dm=dm, deriv=feature_function.deriv, max_memory_in_mb=max_memory_in_mb, safety_fraction=safety_fraction, - func_deriv=func_deriv, + func_deriv=deriv_order, reserved_memory_in_bytes=estimate_global_raw_feature_buffer_memory( dm, feature_function.nfeats, atom_major_raw_features.shape[-1], - func_deriv, + deriv_order, ), ) - max_atom_grid = int(grid_features["atomic_grid_sizes"].max().item()) + max_atom_grid = int(grid_features[Feature.ATOMIC_GRID_SIZES].max().item()) if max_model_grid_points < max_atom_grid: LOG.warning( "Adjusted model chunk size %d to match the largest atomic grid %d. " @@ -190,7 +195,7 @@ def prepare_model_feature_chunks( grid_features=grid_features, feature_function=feature_function, chunk_layouts=_make_atom_grid_chunks( - grid_features["atomic_grid_sizes"], max_model_grid_points + grid_features[Feature.ATOMIC_GRID_SIZES], max_model_grid_points ), - with_spin=dm.ndim == 3, + is_spin_polarized=dm.ndim == 3, ) diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index fb86b7e8..12ef9bcb 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -16,7 +16,6 @@ to_cupy, to_numpy, ) -from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec from skala.pyscf.xc_integrator import XCIntegrator @@ -125,16 +124,6 @@ def func(self) -> ExcFunctionalBase: """Functional retained for gradient-adapter compatibility.""" return self.integrator.functional - @property - def feature_spec(self) -> FeatureSpec: - """Feature requirements owned by the XC integrator.""" - return self.integrator.feature_spec - - @property - def evaluation_policy(self) -> EvaluationPolicy: - """Numerical policy owned by the XC integrator.""" - return self.integrator.evaluation_policy - def _from_backend(self, x: Array) -> Tensor: return from_numpy_or_cupy(x, device=self.device) diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index 096f4bef..ad2ea8fa 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -3,14 +3,14 @@ """Tensor-level exchange-correlation integration.""" from collections.abc import Callable -from dataclasses import dataclass -from typing import cast +from typing import NamedTuple, cast import torch from pyscf import gto from pyscf.dft import numint as pyscf_numint from torch import Tensor +from skala.features import Feature from skala.functional.base import ExcFunctionalBase from skala.pyscf import ao_evaluation, feature_math from skala.pyscf.backend import Grid, check_gpu_imports_were_successful @@ -31,8 +31,7 @@ def _should_screen_aos(mol: gto.Mole) -> bool: return 2 * mol.nao_nr() > pyscf_numint.SWITCH_SIZE -@dataclass(frozen=True) -class XCResult: +class XCResult(NamedTuple): """Tensor-valued result of exchange-correlation integration.""" electron_count: Tensor @@ -69,12 +68,12 @@ def density( mol, dm, grids, - features={"density"}, + features={Feature.DENSITY}, chunk_size=self.evaluation_policy.ao_block_size, max_memory=max_memory, gpu=self.device.type == "cuda", ) - return mol_features["density"].sum(0) + return mol_features[Feature.DENSITY].sum(0) def __call__( self, @@ -183,7 +182,7 @@ def _integrate_screened( grids, atom_major_raw_features=atom_major_raw_features, feature_function=feature_function, - func_deriv=1, + deriv_order=1, max_memory_in_mb=max_memory if dm.device.type == "cpu" else None, safety_fraction=self.evaluation_policy.safety_fraction, ) @@ -199,7 +198,7 @@ def _integrate_screened( ) atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() electron_count += ( - (mol_features["density"] * mol_features["grid_weights"]) + (mol_features[Feature.DENSITY] * mol_features[Feature.GRID_WEIGHTS]) .sum(dim=-1) .detach() ) @@ -230,7 +229,7 @@ def _integrate_dense( mol, dm, grids, - set(self.feature_spec.names) | {"density", "grid_weights"}, + set(self.feature_spec.names) | {Feature.DENSITY, Feature.GRID_WEIGHTS}, chunk_size=self.evaluation_policy.ao_block_size, max_memory=max_memory, gpu=self.device.type == "cuda", @@ -243,9 +242,9 @@ def _integrate_dense( retain_graph=create_graph, create_graph=create_graph, ) - electron_count = (mol_features["density"] * mol_features["grid_weights"]).sum( - dim=-1 - ) + electron_count = ( + mol_features[Feature.DENSITY] * mol_features[Feature.GRID_WEIGHTS] + ).sum(dim=-1) return XCResult(electron_count, energy, potential) def _gen_response_screened( @@ -281,7 +280,7 @@ def _gen_response_screened( grids, atom_major_raw_features=atom_major_raw_features, feature_function=feature_function, - func_deriv=2, + deriv_order=2, max_memory_in_mb=max_memory if dm0.device.type == "cpu" else None, safety_fraction=safety_fraction, ) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 94dbe126..aec997a4 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -4,8 +4,9 @@ import pytest import torch from pyscf import dft, gto -from pyscf.dft import numint as pyscf_numint +from utils import patch_ao_screening +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase from skala.pyscf import model_chunking as model_chunking_module from skala.pyscf import screening as screening_module @@ -28,7 +29,6 @@ _decompose_grid_into_spatial_blocks, prepare_spatial_grid_layout, ) -from skala.pyscf.xc_integrator import _should_screen_aos @pytest.fixture @@ -39,15 +39,24 @@ def carbon() -> gto.Mole: @pytest.mark.parametrize( ("feature_names", "expected_deriv", "expected_nfeats"), [ - ({"density"}, 0, 1), - ({"grad"}, 1, 3), - ({"kin"}, 1, 1), - ({"lapl"}, 2, 1), - ({"density", "grad", "kin", "lapl"}, 2, 6), + ({Feature.DENSITY}, 0, 1), + ({Feature.GRAD}, 1, 3), + ({Feature.KIN}, 1, 1), + ({Feature.LAPL}, 2, 1), + ( + { + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.LAPL, + }, + 2, + 6, + ), ], ) def test_mgga_supported_features_are_linear_in_density_matrix( - feature_names: set[str], + feature_names: set[Feature], expected_deriv: int, expected_nfeats: int, ) -> None: @@ -94,27 +103,27 @@ def first_jvp(value: torch.Tensor) -> torch.Tensor: torch.testing.assert_close(second_jvp, torch.zeros_like(second_jvp)) -def test_mgga_requires_at_least_one_feature() -> None: - with pytest.raises(ValueError, match="At least one feature must be selected"): - MGGAFeatureFunction(FeatureSpec([])) +@pytest.mark.parametrize("feature_names", [[], [Feature.GRID_WEIGHTS]]) +def test_mgga_requires_at_least_one_ao_derived_feature( + feature_names: list[Feature], +) -> None: + with pytest.raises( + ValueError, match="At least one AO-derived feature must be selected" + ): + MGGAFeatureFunction(FeatureSpec(feature_names)) -@pytest.mark.parametrize( - ("switch_offset", "expected"), [(1, False), (0, False), (-1, True)] -) -def test_should_screen_aos_at_crossover( - carbon: gto.Mole, - monkeypatch: pytest.MonkeyPatch, - switch_offset: int, - expected: bool, -) -> None: - monkeypatch.setattr( - pyscf_numint, - "SWITCH_SIZE", - 2 * carbon.nao_nr() + switch_offset, - ) +def test_patch_ao_screening_restores_previous_decision(carbon: gto.Mole) -> None: + original_decision = xc_integrator_module._should_screen_aos + + with patch_ao_screening(False): + dense_decision = xc_integrator_module._should_screen_aos + assert not dense_decision(carbon) + with patch_ao_screening(True): + assert xc_integrator_module._should_screen_aos(carbon) + assert xc_integrator_module._should_screen_aos is dense_decision - assert _should_screen_aos(carbon) is expected + assert xc_integrator_module._should_screen_aos is original_decision def test_active_cpu_ao_indices(carbon: gto.Mole) -> None: @@ -138,7 +147,7 @@ def test_active_cpu_ao_indices(carbon: gto.Mole) -> None: def test_resolve_ao_block_size_modes(carbon: gto.Mole) -> None: - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) backend_block_size = dft.gen_grid.BLKSIZE # CPU sizes are aligned locally; GPU sizing is delegated unless explicitly invalid. @@ -298,10 +307,14 @@ def fake_make_screen_index( class QuadraticDensityFunctional(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["atomic_grid_sizes", "density", "grid_weights"] + self.features = [ + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRID_WEIGHTS, + ] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: - return (mol["density"].square() * mol["grid_weights"]).sum() + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() def test_grid_reuses_spatial_grid_layout_across_numints( @@ -386,8 +399,6 @@ def test_first_and_second_order_use_same_screening_decision( expected: bool, response_safety_fraction: float | None, ) -> None: - switch_size = 2 * carbon.nao_nr() - 1 if expected else 2 * carbon.nao_nr() - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", switch_size) routes: list[str] = [] safety_fractions: list[float] = [] @@ -395,15 +406,15 @@ def fake_generate_features( mol: gto.Mole, dm: torch.Tensor, grids: object, - features: set[str] | None = None, + features: set[Feature] | None = None, **kwargs: object, - ) -> dict[str, torch.Tensor]: + ) -> FeatureMap: routes.append("dense") density = dm.square().sum().reshape(1).expand(2, 1) / 2 return { - "atomic_grid_sizes": torch.tensor([1]), - "density": density, - "grid_weights": torch.ones(1, dtype=dm.dtype), + Feature.ATOMIC_GRID_SIZES: torch.tensor([1]), + Feature.DENSITY: density, + Feature.GRID_WEIGHTS: torch.ones(1, dtype=dm.dtype), } class FakeSpatialGridLayout: @@ -424,9 +435,9 @@ def __iter__(self) -> Iterator[ModelFeatureChunk]: grid_slice=slice(0, 1), raw_features=raw_features, model_features={ - "atomic_grid_sizes": torch.tensor([1]), - "density": raw_features.expand(2, 1) / 2, - "grid_weights": torch.ones(1, dtype=raw_features.dtype), + Feature.ATOMIC_GRID_SIZES: torch.tensor([1]), + Feature.DENSITY: raw_features.expand(2, 1) / 2, + Feature.GRID_WEIGHTS: torch.ones(1, dtype=raw_features.dtype), }, ) @@ -460,7 +471,7 @@ def fake_prepare_model_feature_chunks( grids: object, atom_major_raw_features: torch.Tensor, feature_function: MGGAFeatureFunction, - func_deriv: int, + deriv_order: int, **kwargs: object, ) -> FakeModelFeatureChunks: safety_fraction = kwargs["safety_fraction"] @@ -496,21 +507,21 @@ def fake_prepare_model_feature_chunks( grids = dft.Grids(carbon) grids.weights = np.ones(1) - numint(carbon, grids, None, dm) - ks = FakeKS(carbon, grids) response_kwargs = ( {} if response_safety_fraction is None else {"safety_fraction": response_safety_fraction} ) - response = numint.gen_response( - np.eye(carbon.nao_nr()), - np.ones(carbon.nao_nr()), - ks=ks, - **response_kwargs, - ) - result = response(np.eye(carbon.nao_nr())) + with patch_ao_screening(expected): + numint(carbon, grids, None, dm) + response = numint.gen_response( + np.eye(carbon.nao_nr()), + np.ones(carbon.nao_nr()), + ks=ks, + **response_kwargs, + ) + result = response(np.eye(carbon.nao_nr())) assert result.shape == (carbon.nao_nr(), carbon.nao_nr()) expected_route = "screened" if expected else "dense" @@ -526,7 +537,7 @@ def fake_prepare_model_feature_chunks( def test_feature_block_helper_localizes_derivative_vectors() -> None: """Use AO slices for linear JVPs and grid slices for feature VJPs.""" - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) block = _AOBlock( ao_values=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), active_ao_indices=torch.tensor([0, 2]), @@ -571,7 +582,7 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: def test_chunk_eval_transforms_follow_linear_operator(carbon: gto.Mole) -> None: """Check first and second JVPs and the feature-cotangent adjoint JVP.""" grids = _minimal_atom_grid(carbon) - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) tangent = torch.arange(1, dm.numel() + 1, dtype=dm.dtype).reshape(dm.shape) @@ -642,7 +653,7 @@ def block_loop( monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) dm = torch.diag( torch.arange(1, carbon.nao_nr() + 1, dtype=torch.float64) ).requires_grad_() @@ -713,7 +724,7 @@ def block_loop( ) monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) blocks = list(_CPUAOBlockLoop(dm, carbon, grids, feature_function, block_size)) @@ -747,7 +758,7 @@ def block_loop( yield ao, None, grids.weights, grids.coords monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] @@ -788,7 +799,6 @@ def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> No @pytest.mark.parametrize("unrestricted", [False, True]) def test_cpu_rks_uks_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], unrestricted: bool, ) -> None: @@ -805,28 +815,26 @@ def test_cpu_rks_uks_dense_screened_equivalence( grids = _minimal_atom_grid(mol) dm = mean_field.get_init_guess() - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense = ( - numint.nr_uks(mol, grids, None, dm) - if unrestricted - else numint.nr_rks(mol, grids, None, dm) - ) + with patch_ao_screening(False): + dense = ( + numint.nr_uks(mol, grids, None, dm) + if unrestricted + else numint.nr_rks(mol, grids, None, dm) + ) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened = ( - numint.nr_uks(mol, grids, None, dm) - if unrestricted - else numint.nr_rks(mol, grids, None, dm) - ) + with patch_ao_screening(True): + screened = ( + numint.nr_uks(mol, grids, None, dm) + if unrestricted + else numint.nr_rks(mol, grids, None, dm) + ) assert np.allclose(dense[0], screened[0], rtol=1e-10, atol=1e-11) assert np.isclose(dense[1], screened[1], rtol=1e-9, atol=1e-10) assert np.allclose(dense[2], screened[2], rtol=1e-8, atol=1e-10) -def test_cpu_response_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_cpu_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) grids = _minimal_atom_grid(mol) ks = FakeKS(mol, grids) @@ -838,11 +846,11 @@ def test_cpu_response_dense_screened_equivalence( ) dm1 += dm1.T - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(False): + dense_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(True): + screened_response = numint.gen_response(mo_coeff, mo_occ, ks=ks) assert np.allclose( dense_response(dm1), screened_response(dm1), rtol=1e-10, atol=1e-11 @@ -859,10 +867,9 @@ def test_screened_ao_traversals_are_independent_of_model_chunking( atom_grid_size = grids.weights.size // mol.natm monkeypatch.setattr( model_chunking_module, - "estimate_max_model_grid_points", + "estimate_max_gridpoint_chunk_size", lambda *args, **kwargs: atom_grid_size, ) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) forward_calls = 0 backward_calls = 0 @@ -884,16 +891,17 @@ def counting_backward_apply(*args: object) -> torch.Tensor: functional = QuadraticDensityFunctional() numint = SkalaNumInt(functional) - if func_deriv == 1: - dm = dft.RKS(mol).get_init_guess() - numint.nr_rks(mol, grids, None, dm) - assert forward_calls == 1 - else: - ks = FakeKS(mol, grids) - response = numint.gen_response( - np.eye(mol.nao_nr()), np.ones(mol.nao_nr()), ks=ks - ) - response(np.eye(mol.nao_nr())) - assert forward_calls == 2 + with patch_ao_screening(True): + if func_deriv == 1: + dm = dft.RKS(mol).get_init_guess() + numint.nr_rks(mol, grids, None, dm) + assert forward_calls == 1 + else: + ks = FakeKS(mol, grids) + response = numint.gen_response( + np.eye(mol.nao_nr()), np.ones(mol.nao_nr()), ks=ks + ) + response(np.eye(mol.nao_nr())) + assert forward_calls == 2 assert backward_calls == 1 diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 1996e671..39426665 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -12,14 +12,13 @@ import pytest import torch from pyscf import dft, gto, lib -from pyscf.dft import numint as pyscf_numint from pytest_benchmark.fixture import BenchmarkFixture from torch.utils.dlpack import from_dlpack +from utils import patch_ao_screening from skala.functional import load_functional from skala.functional.base import ExcFunctionalBase from skala.pyscf.numint import SkalaNumInt -from skala.pyscf.xc_integrator import _should_screen_aos THREAD_COUNT = 4 MAX_MEMORY_MB = 2000 @@ -211,7 +210,7 @@ def device_benchmark_case( benchmark_spec: BenchmarkSpec, fixed_cpu_threads: None, load_functional_cached: Callable[..., ExcFunctionalBase | str], -) -> BenchmarkCase: +) -> Iterator[BenchmarkCase]: backend = cast(str, request.param) if backend == "cpu": functional = load_functional_cached("skala-1.1") @@ -225,23 +224,20 @@ def device_benchmark_case( assert isinstance(functional, ExcFunctionalBase) case = _make_benchmark_case(benchmark_spec, functional, backend) - assert _should_screen_aos(case.mol) - return case + with patch_ao_screening(True): + yield case @pytest.fixture -def screened_case(benchmark_case: BenchmarkCase) -> BenchmarkCase: - assert _should_screen_aos(benchmark_case.mol) - return benchmark_case +def screened_case(benchmark_case: BenchmarkCase) -> Iterator[BenchmarkCase]: + with patch_ao_screening(True): + yield benchmark_case @pytest.fixture -def dense_case( - benchmark_case: BenchmarkCase, monkeypatch: pytest.MonkeyPatch -) -> BenchmarkCase: - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", benchmark_case.mol.nao_nr()) - assert not _should_screen_aos(benchmark_case.mol) - return benchmark_case +def dense_case(benchmark_case: BenchmarkCase) -> Iterator[BenchmarkCase]: + with patch_ao_screening(False): + yield benchmark_case def _benchmark_device_xc(benchmark: BenchmarkFixture, case: BenchmarkCase) -> None: @@ -258,15 +254,13 @@ def _run_gpu_xc(spec: BenchmarkSpec, screened: bool) -> int: functional = load_functional("skala-1.1", device=torch.device("cuda:0")) assert isinstance(functional, ExcFunctionalBase) case = _make_benchmark_case(spec, functional, "cuda") - assert _should_screen_aos(case.mol) - if not screened: - pyscf_numint.SWITCH_SIZE = 10**9 torch.cuda.synchronize() torch.cuda.empty_cache() baseline_bytes = torch.cuda.memory_allocated() torch.cuda.reset_peak_memory_stats() - case.run() + with patch_ao_screening(screened): + case.run() torch.cuda.synchronize() return torch.cuda.max_memory_allocated() - baseline_bytes @@ -287,13 +281,13 @@ def _memory_worker( functional = load_functional("skala-1.1") assert isinstance(functional, ExcFunctionalBase) case = _make_benchmark_case(spec, functional, "cpu") - if not screened: - pyscf_numint.SWITCH_SIZE = case.mol.nao_nr() - assert _should_screen_aos(case.mol) is screened with tempfile.TemporaryDirectory() as tmpdir: profile_path = Path(tmpdir) / "allocations.bin" - with memray.Tracker(profile_path): + with ( + patch_ao_screening(screened), + memray.Tracker(profile_path), + ): case.run() peak_bytes = memray.FileReader(profile_path).metadata.peak_memory elif backend == "cuda": @@ -346,20 +340,19 @@ def test_screened_and_dense_values_agree( device_benchmark_case: BenchmarkCase, benchmark_spec: BenchmarkSpec, load_functional_cached: Callable[..., ExcFunctionalBase | str], - monkeypatch: pytest.MonkeyPatch, ) -> None: case = device_benchmark_case - assert _should_screen_aos(case.mol) screened = case.run() - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", case.mol.nao_nr()) - assert not _should_screen_aos(case.mol) - if case.backend == "cpu": - dense = case.run() - else: - cpu_functional = load_functional_cached("skala-1.1", device=torch.device("cpu")) - assert isinstance(cpu_functional, ExcFunctionalBase) - dense = _make_benchmark_case(benchmark_spec, cpu_functional, "cpu").run() + with patch_ao_screening(False): + if case.backend == "cpu": + dense = case.run() + else: + cpu_functional = load_functional_cached( + "skala-1.1", device=torch.device("cpu") + ) + assert isinstance(cpu_functional, ExcFunctionalBase) + dense = _make_benchmark_case(benchmark_spec, cpu_functional, "cpu").run() scalar_rtol = 2e-10 if case.backend == "cpu" else 1e-8 density_close = np.allclose(dense[0], screened[0], rtol=scalar_rtol, atol=1e-11) @@ -397,14 +390,14 @@ def test_screened_and_dense_values_agree( @pytest.mark.benchmark(group="def2-qzvpp") -def test_with_natural_ao_screening( +def test_with_ao_screening( benchmark: BenchmarkFixture, screened_case: BenchmarkCase ) -> None: _benchmark_device_xc(benchmark, screened_case) @pytest.mark.benchmark(group="def2-qzvpp") -def test_without_ao_screening_by_patching_threshold( +def test_without_ao_screening_by_patching_decision( benchmark: BenchmarkFixture, dense_case: BenchmarkCase ) -> None: _benchmark_device_xc(benchmark, dense_case) @@ -454,22 +447,15 @@ def test_screened_and_dense_peak_memory( @pytest.mark.profiling -def test_profile_with_natural_ao_screening( +def test_profile_with_ao_screening( device_benchmark_case: BenchmarkCase, ) -> None: - assert _should_screen_aos(device_benchmark_case.mol) device_benchmark_case.run() @pytest.mark.profiling -def test_profile_without_ao_screening_by_patching_threshold( +def test_profile_without_ao_screening_by_patching_decision( device_benchmark_case: BenchmarkCase, - monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr( - pyscf_numint, - "SWITCH_SIZE", - device_benchmark_case.mol.nao_nr(), - ) - assert not _should_screen_aos(device_benchmark_case.mol) - device_benchmark_case.run() + with patch_ao_screening(False): + device_benchmark_case.run() diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py index 159f2488..3d2bfe2f 100644 --- a/tests/test_evaluation.py +++ b/tests/test_evaluation.py @@ -1,48 +1,57 @@ from dataclasses import FrozenInstanceError import pytest -import torch -from skala.functional.base import ExcFunctionalBase +from skala.features import Feature from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec -from skala.pyscf.features import generate_features -from skala.pyscf.numint import SkalaNumInt + + +def test_feature_name_parses_model_metadata_string() -> None: + assert Feature("density") is Feature.DENSITY @pytest.mark.parametrize( ("features", "expected_order"), [ ([], 0), - (["density"], 0), - (["grad"], 1), - (["kin"], 1), - (["lapl"], 2), - (["density", "grad", "kin", "lapl"], 2), + ([Feature.DENSITY], 0), + ([Feature.GRAD], 1), + ([Feature.KIN], 1), + ([Feature.LAPL], 2), + ( + [ + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.LAPL, + ], + 2, + ), ], ) def test_feature_spec_derives_mgga_requirements( - features: list[str], expected_order: int + features: list[Feature], expected_order: int ) -> None: spec = FeatureSpec(features) - assert spec.requires_mgga is bool(features) + assert spec.requires_ao_evaluation is bool(features) assert spec.ao_derivative_order == expected_order - assert spec.with_density is ("density" in features) - assert spec.with_grad is ("grad" in features) - assert spec.with_kin is ("kin" in features) - assert spec.with_lapl is ("lapl" in features) + assert spec.with_density is (Feature.DENSITY in features) + assert spec.with_grad is (Feature.GRAD in features) + assert spec.with_kin is (Feature.KIN in features) + assert spec.with_lapl is (Feature.LAPL in features) @pytest.mark.parametrize( ("feature", "supports_screened_evaluation"), [ - ("atomic_grid_weights", False), - ("atomic_grid_sizes", True), - ("atomic_grid_size_bound_shape", False), + (Feature.ATOMIC_GRID_WEIGHTS, False), + (Feature.ATOMIC_GRID_SIZES, True), + (Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, False), ], ) def test_feature_spec_derives_atomic_layout_requirements( - feature: str, supports_screened_evaluation: bool + feature: Feature, supports_screened_evaluation: bool ) -> None: spec = FeatureSpec([feature, feature]) @@ -58,27 +67,3 @@ def test_evaluation_policy_defaults_and_is_immutable() -> None: assert policy.safety_fraction == 0.8 with pytest.raises(FrozenInstanceError): policy.safety_fraction = 0.5 # type: ignore[misc] - - -def test_explicit_empty_feature_set_stays_empty() -> None: - features = generate_features( - mol=object(), - dm=torch.eye(1), - grids=object(), - features=set(), - ) - - assert features == {} - - -class DensityFunctional(ExcFunctionalBase): - def __init__(self) -> None: - super().__init__() - self.features = ["density"] - - -def test_numint_translates_chunk_size_into_evaluation_policy() -> None: - numint = SkalaNumInt(DensityFunctional(), chunk_size=96) - - assert numint.feature_spec == FeatureSpec(["density"]) - assert numint.evaluation_policy == EvaluationPolicy(ao_block_size=96) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index b0e06ac9..c3c46b8c 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -5,7 +5,6 @@ import pytest import torch from pyscf import dft, gto -from pyscf.dft import numint as pyscf_numint from torch.utils.dlpack import from_dlpack pytestmark = pytest.mark.gpu @@ -24,11 +23,14 @@ allow_module_level=True, ) +from utils import patch_ao_screening # noqa: E402 + +from skala.features import Feature, FeatureMap # noqa: E402 from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.pyscf.ao_evaluation import ( # noqa: E402 ChunkEvalForward, - non_chunk, + evaluate_full_grid, ) from skala.pyscf.backend import dft_gpu # noqa: E402 from skala.pyscf.evaluation import FeatureSpec # noqa: E402 @@ -47,30 +49,34 @@ class QuadraticDensityFunctional(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["atomic_grid_sizes", "density", "grid_weights"] + self.features = [ + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRID_WEIGHTS, + ] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: - return (mol["density"].square() * mol["grid_weights"]).sum() + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() class QuadraticMGGAFunctional(ExcFunctionalBase): def __init__(self) -> None: super().__init__() self.features = [ - "atomic_grid_sizes", - "density", - "grad", - "kin", - "grid_weights", + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.GRID_WEIGHTS, ] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: energy_density = ( - mol["density"].square() - + mol["grad"].square().sum(dim=-2) - + mol["kin"].square() + mol[Feature.DENSITY].square() + + mol[Feature.GRAD].square().sum(dim=-2) + + mol[Feature.KIN].square() ) - return (energy_density * mol["grid_weights"]).sum() + return (energy_density * mol[Feature.GRID_WEIGHTS]).sum() def _to_numpy(value: object) -> np.ndarray: @@ -117,7 +123,6 @@ def test_prepare_spatially_sorted_gpu_grids() -> None: @pytest.mark.parametrize("unrestricted", [False, True]) def test_gpu_rks_uks_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], unrestricted: bool, ) -> None: @@ -134,19 +139,19 @@ def test_gpu_rks_uks_dense_screened_equivalence( ks.grids.build(sort_grids=False) dm = ks.get_init_guess() - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense = ( - ks._numint.nr_uks(mol, ks.grids, None, dm) - if unrestricted - else ks._numint.nr_rks(mol, ks.grids, None, dm) - ) + with patch_ao_screening(False): + dense = ( + ks._numint.nr_uks(mol, ks.grids, None, dm) + if unrestricted + else ks._numint.nr_rks(mol, ks.grids, None, dm) + ) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened = ( - ks._numint.nr_uks(mol, ks.grids, None, dm) - if unrestricted - else ks._numint.nr_rks(mol, ks.grids, None, dm) - ) + with patch_ao_screening(True): + screened = ( + ks._numint.nr_uks(mol, ks.grids, None, dm) + if unrestricted + else ks._numint.nr_rks(mol, ks.grids, None, dm) + ) assert np.allclose(_to_numpy(dense[0]), _to_numpy(screened[0]), rtol=1e-9) assert np.isclose(dense[1], screened[1], rtol=1e-9) @@ -155,9 +160,7 @@ def test_gpu_rks_uks_dense_screened_equivalence( ) -def test_gpu_response_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_gpu_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) ks.grids.level = 0 @@ -170,11 +173,11 @@ def test_gpu_response_dense_screened_equivalence( ) dm1 += dm1.T - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(False): + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(True): + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) assert np.allclose( _to_numpy(dense_response(dm1)), @@ -184,9 +187,7 @@ def test_gpu_response_dense_screened_equivalence( ) -def test_gpu_uks_response_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_gpu_uks_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) ks.grids.level = 0 @@ -196,11 +197,11 @@ def test_gpu_uks_response_dense_screened_equivalence( mo_occ = cupy.ones((2, mol.nao_nr())) dm1 = cupy.ones((2, mol.nao_nr(), mol.nao_nr())) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(False): + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(True): + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) np.testing.assert_allclose( _to_numpy(screened_response(dm1)), @@ -210,9 +211,7 @@ def test_gpu_uks_response_dense_screened_equivalence( ) -def test_gpu_multiblock_mgga_response_dense_screened_equivalence( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_gpu_multiblock_mgga_response_dense_screened_equivalence() -> None: mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) ks = SkalaKS(mol, xc=QuadraticMGGAFunctional(), with_dftd3=False) ks.grids.level = 1 @@ -223,11 +222,11 @@ def test_gpu_multiblock_mgga_response_dense_screened_equivalence( mo_occ = cupy.ones(mol.nao_nr()) dm1 = cupy.eye(mol.nao_nr()) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(False): + dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) + with patch_ao_screening(True): + screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) np.testing.assert_allclose( _to_numpy(screened_response(dm1)), @@ -238,7 +237,6 @@ def test_gpu_multiblock_mgga_response_dense_screened_equivalence( def test_gpu_screened_skala_matches_cpu_on_carbon_chain( - monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> None: """Prevent inaccurate GPU AO screening on spatially diffuse grid blocks. @@ -290,15 +288,15 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( assert isinstance(cpu_functional, ExcFunctionalBase) assert isinstance(gpu_functional, ExcFunctionalBase) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr()) - cpu_result = SkalaNumInt(cpu_functional, device=torch.device("cpu")).nr_rks( - mol, cpu_grids, None, dm - ) + with patch_ao_screening(False): + cpu_result = SkalaNumInt(cpu_functional, device=torch.device("cpu")).nr_rks( + mol, cpu_grids, None, dm + ) - monkeypatch.setattr(pyscf_numint, "SWITCH_SIZE", mol.nao_nr() - 1) - gpu_result = SkalaNumInt(gpu_functional, device=torch.device("cuda:0")).nr_rks( - mol, gpu_grids, None, cupy.asarray(dm) - ) + with patch_ao_screening(True): + gpu_result = SkalaNumInt(gpu_functional, device=torch.device("cuda:0")).nr_rks( + mol, gpu_grids, None, cupy.asarray(dm) + ) gpu_vxc = cupy.asnumpy(gpu_result[2]) vxc_difference = cpu_result[2] - gpu_vxc @@ -355,12 +353,14 @@ def test_gpu_empty_ao_block_matches_dense_reference() -> None: assert active_ao_counts[1] > 0 grids._non0ao_idx = None - feature_function = MGGAFeatureFunction(FeatureSpec(["density", "grad", "kin"])) + feature_function = MGGAFeatureFunction( + FeatureSpec([Feature.DENSITY, Feature.GRAD, Feature.KIN]) + ) dm = torch.eye(mol.nao_nr(), dtype=torch.float64, device="cuda", requires_grad=True) screened = ChunkEvalForward.apply( # type: ignore[no-untyped-call] dm, mol, grids, feature_function, block_size, False, True ) - dense = non_chunk(dm, mol, coords, feature_function, gpu=True) + dense = evaluate_full_grid(dm, mol, coords, feature_function, gpu=True) torch.testing.assert_close(screened, dense, rtol=1e-12, atol=1e-12) (screened_vjp,) = torch.autograd.grad(screened.square().sum(), dm) @@ -408,7 +408,7 @@ def block_loop( monkeypatch.setattr(dft_gpu.numint, "NumInt", FakeGpuNumInt) - feature_function = MGGAFeatureFunction(FeatureSpec(["density"])) + feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) dm = torch.diag( torch.arange(1, mol.nao_nr() + 1, dtype=torch.float64, device="cuda") ).requires_grad_() diff --git a/tests/test_gpu4pyscf_gradients.py b/tests/test_gpu4pyscf_gradients.py index cd2154a3..edf64cb9 100644 --- a/tests/test_gpu4pyscf_gradients.py +++ b/tests/test_gpu4pyscf_gradients.py @@ -2,6 +2,7 @@ import pytest import torch +from torch.utils.dlpack import from_dlpack pytestmark = pytest.mark.gpu @@ -21,9 +22,10 @@ from _ridders import num_grad_ridders # noqa: E402 from gpu4pyscf import dft, scf # noqa: E402 from pyscf import gto # noqa: E402 -from pyscf.dft import numint as pyscf_numint # noqa: E402 from test_pyscf_gradients import FULL_GRAD_REF # noqa: E402 +from utils import patch_ao_screening # noqa: E402 +from skala.features import Feature, FeatureMap # noqa: E402 from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.gpu4pyscf.gradients import ( # noqa: E402 @@ -35,7 +37,6 @@ from skala.pyscf import SkalaKS as CpuSkalaKS # noqa: E402 from skala.pyscf.features import generate_features # noqa: E402 from skala.pyscf.gradients import SkalaRKSGradient as CpuSkalaRKSGradient # noqa: E402 -from skala.pyscf.xc_integrator import _should_screen_aos # noqa: E402 from skala.utils import torch_allocator # noqa: E402 H2_SKALA_1_1_GRAD_REF = torch.tensor( @@ -124,7 +125,7 @@ def get_grid_and_rdm1(mol: gto.Mole) -> tuple[dft.Grids, torch.Tensor]: grids=minimal_grid(mol), ) mf.kernel() - rdm1 = torch.from_dlpack(mf.make_rdm1()) # type: ignore[attr-defined] + rdm1 = from_dlpack(mf.make_rdm1()) return mf.grids, rdm1 # maybe_expand_and_divide(rdm1, len(rdm1.shape) == 2, 2) @@ -132,11 +133,11 @@ def test_grid_coords_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grid_coords"] + self.features = [Feature.GRID_COORDS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return mol["grid_coords"].sum() + return mol[Feature.GRID_COORDS].sum() mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) @@ -159,11 +160,11 @@ def test_coarse_0_atomic_coords_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["coarse_0_atomic_coords"] + self.features = [Feature.COARSE_0_ATOMIC_COORDS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return torch.einsum("nx->", mol["coarse_0_atomic_coords"]) + return torch.einsum("nx->", mol[Feature.COARSE_0_ATOMIC_COORDS]) mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) @@ -182,11 +183,11 @@ def test_grid_weights_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grid_weights"] + self.features = [Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return mol["grid_weights"].sum() + return mol[Feature.GRID_WEIGHTS].sum() def finite_difference_nuc_grad( weight_sum: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -200,7 +201,7 @@ def finite_difference_nuc_grad( def weight_sum_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: """Exc wrapper for the finite difference""" mol.set_geom_(nuc_coords.cpu().numpy(), "bohr", symmetry=None) - mol_feats["grid_weights"] = torch.from_dlpack(minimal_grid(mol).weights) # type: ignore[attr-defined] + mol_feats[Feature.GRID_WEIGHTS] = from_dlpack(minimal_grid(mol).weights) return weight_sum.get_exc(mol_feats) @@ -215,7 +216,7 @@ def weight_sum_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: num_grad, num_err = finite_difference_nuc_grad(exc_test, mol, rdm1) # estimate the minimum expected absolute error eps = ( - exc_test.get_exc({"grid_weights": torch.from_dlpack(grid.weights)}) # type: ignore[attr-defined] + exc_test.get_exc({Feature.GRID_WEIGHTS: from_dlpack(grid.weights)}) * torch.finfo(num_grad.dtype).eps ) @@ -234,11 +235,11 @@ def test_density_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["density", "grid_weights"] + self.features = [Feature.DENSITY, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return (mol["density"] @ mol["grid_weights"]).sum() + return (mol[Feature.DENSITY] @ mol[Feature.GRID_WEIGHTS]).sum() def finite_difference_nuc_grad( dens_sum: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -268,7 +269,7 @@ def dens_sum_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: # calculate analytic result veff = veff_and_expl_nuc_grad( - exc_test, mol, grid, rdm1, nuc_grad_feats={"density"} + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.DENSITY} )[0] ana_grad = 2 * nuc_grad_from_veff(mol, veff, rdm1) @@ -289,15 +290,15 @@ def test_grad_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grad", "grid_weights"] + self.features = [Feature.GRAD, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: return ( - (mol["grad"] ** 2 @ mol["grid_weights"]) + (mol[Feature.GRAD] ** 2 @ mol[Feature.GRID_WEIGHTS]) @ torch.tensor( [1.0, 2.0, 3.0], dtype=torch.float64, - device=mol["grad"].device, + device=mol[Feature.GRAD].device, ) ).sum() @@ -329,7 +330,9 @@ def grad_func_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: num_grad, num_err = finite_difference_nuc_grad(exc_test, mol, rdm1) # calculate analytic result - veff = veff_and_expl_nuc_grad(exc_test, mol, grid, rdm1, nuc_grad_feats={"grad"})[0] + veff = veff_and_expl_nuc_grad( + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.GRAD} + )[0] ana_grad = 2 * nuc_grad_from_veff(mol, veff, rdm1) check_mat = (ana_grad - num_grad).abs() <= torch.max( @@ -349,11 +352,11 @@ def test_kin_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["kin", "grid_weights"] + self.features = [Feature.KIN, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total kinetic energy number""" - return (mol["kin"] @ mol["grid_weights"]).sum() + return (mol[Feature.KIN] @ mol[Feature.GRID_WEIGHTS]).sum() def finite_difference_nuc_grad( kin_func: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -383,7 +386,9 @@ def kin_func_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: num_grad, num_err = finite_difference_nuc_grad(exc_test, mol, rdm1) # calculate analytic result - veff = veff_and_expl_nuc_grad(exc_test, mol, grid, rdm1, nuc_grad_feats={"kin"})[0] + veff = veff_and_expl_nuc_grad( + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.KIN} + )[0] ana_grad = 2 * nuc_grad_from_veff(mol, veff, rdm1) check_mat = (ana_grad - num_grad).abs() <= torch.max( @@ -462,7 +467,6 @@ def test_full_grad( def test_nuclear_gradient_cpu_gpu_dense_screened_agree( - monkeypatch: pytest.MonkeyPatch, load_functional_cached: Callable[..., ExcFunctionalBase | str], ) -> None: """Compare complete nuclear gradients across backend and SCF screening routes. @@ -487,24 +491,19 @@ def test_nuclear_gradient_cpu_gpu_dense_screened_agree( basis="sto-3g", verbose=0, ) - monkeypatch.setattr( - pyscf_numint, - "SWITCH_SIZE", - mol.nao_nr() - int(screened), - ) - assert _should_screen_aos(mol) is screened - if backend == "cpu": - mean_field = CpuSkalaKS(mol, xc=functional, with_dftd3=False) - gradient_type = CpuSkalaRKSGradient - else: - mean_field = SkalaKS(mol, xc=functional, with_dftd3=False) - gradient_type = SkalaRKSGradient - mean_field.grids.level = 0 - mean_field.grids.build(mol, sort_grids=False) - mean_field.conv_tol = 1e-10 - mean_field.kernel() - assert mean_field.converged - gradient = gradient_type(mean_field).kernel() + with patch_ao_screening(screened): + if backend == "cpu": + mean_field = CpuSkalaKS(mol, xc=functional, with_dftd3=False) + gradient_type = CpuSkalaRKSGradient + else: + mean_field = SkalaKS(mol, xc=functional, with_dftd3=False) + gradient_type = SkalaRKSGradient + mean_field.grids.level = 0 + mean_field.grids.build(mol, sort_grids=False) + mean_field.conv_tol = 1e-10 + mean_field.kernel() + assert mean_field.converged + gradient = gradient_type(mean_field).kernel() route = f"{backend}-{'screened' if screened else 'dense'}" gradients[route] = torch.from_numpy(gradient) @@ -537,15 +536,15 @@ def test_cuda_kernel_memory_stability() -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grad", "grid_weights"] + self.features = [Feature.GRAD, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: return ( - (mol["grad"] ** 2 @ mol["grid_weights"]) + (mol[Feature.GRAD] ** 2 @ mol[Feature.GRID_WEIGHTS]) @ torch.tensor( [1.0, 2.0, 3.0], dtype=torch.float64, - device=mol["grad"].device, + device=mol[Feature.GRAD].device, ) ).sum() @@ -554,7 +553,7 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: # Warmup to avoid counting one-time allocations from CUDA runtime/libraries. for _ in range(2): veff = veff_and_expl_nuc_grad( - exc_test, mol, grid, rdm1, nuc_grad_feats={"grad"} + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.GRAD} )[0] _ = 2 * nuc_grad_from_veff(mol, veff, rdm1) @@ -565,7 +564,7 @@ def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: torch.cuda.reset_peak_memory_stats() for _ in range(5): veff = veff_and_expl_nuc_grad( - exc_test, mol, grid, rdm1, nuc_grad_feats={"grad"} + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.GRAD} )[0] _ = 2 * nuc_grad_from_veff(mol, veff, rdm1) torch.cuda.synchronize() diff --git a/tests/test_memory_estimators.py b/tests/test_memory_estimators.py index 05f5914f..ec5b6200 100644 --- a/tests/test_memory_estimators.py +++ b/tests/test_memory_estimators.py @@ -3,7 +3,7 @@ from skala.pyscf.memory_estimators import ( estimate_global_raw_feature_buffer_memory, - estimate_max_model_grid_points, + estimate_max_gridpoint_chunk_size, linear_peak_memory_model, ) @@ -37,7 +37,7 @@ def test_global_raw_feature_buffer_memory_rejects_unsupported_order() -> None: def test_reserved_memory_reduces_grid_chunk_size() -> None: dm = torch.eye(10, dtype=torch.float64) bytes_per_point, _ = linear_peak_memory_model(nao=10, deriv=1, func_deriv=1) - base_chunk_size = estimate_max_model_grid_points( + base_chunk_size = estimate_max_gridpoint_chunk_size( dm, deriv=1, max_memory_in_mb=100, @@ -45,7 +45,7 @@ def test_reserved_memory_reduces_grid_chunk_size() -> None: func_deriv=1, ) reserved_points = 123 - reserved_chunk_size = estimate_max_model_grid_points( + reserved_chunk_size = estimate_max_gridpoint_chunk_size( dm, deriv=1, max_memory_in_mb=100, @@ -64,7 +64,7 @@ def test_model_grid_point_limit_rejects_invalid_safety_fraction( with pytest.raises( ValueError, match="safety_fraction must be greater than 0 and at most 1" ): - estimate_max_model_grid_points( + estimate_max_gridpoint_chunk_size( torch.eye(2, dtype=torch.float64), deriv=1, max_memory_in_mb=100, diff --git a/tests/test_model.py b/tests/test_model.py index 66352583..a785f430 100644 --- a/tests/test_model.py +++ b/tests/test_model.py @@ -15,6 +15,7 @@ import pytest import torch +from skala.features import Feature, FeatureMap from skala.functional import ExcFunctionalBase from skala.functional.model import ( ANGSTROM_TO_BOHR, @@ -53,22 +54,24 @@ def make_mol( grid_per_atom: int, device: str = "cpu", dtype: torch.dtype = torch.float64, -) -> dict[str, torch.Tensor]: +) -> FeatureMap: total_grid = num_atoms * grid_per_atom return { - "density": torch.randn(2, total_grid, dtype=dtype, device=device), - "grad": torch.randn(2, 3, total_grid, dtype=dtype, device=device), - "kin": torch.randn(2, total_grid, dtype=dtype, device=device), - "grid_coords": torch.randn(total_grid, 3, dtype=dtype, device=device), - "grid_weights": torch.randn(total_grid, dtype=dtype, device=device).abs(), - "atomic_grid_weights": torch.randn( + Feature.DENSITY: torch.randn(2, total_grid, dtype=dtype, device=device), + Feature.GRAD: torch.randn(2, 3, total_grid, dtype=dtype, device=device), + Feature.KIN: torch.randn(2, total_grid, dtype=dtype, device=device), + Feature.GRID_COORDS: torch.randn(total_grid, 3, dtype=dtype, device=device), + Feature.GRID_WEIGHTS: torch.randn(total_grid, dtype=dtype, device=device).abs(), + Feature.ATOMIC_GRID_WEIGHTS: torch.randn( total_grid, dtype=dtype, device=device ).abs(), - "atomic_grid_sizes": torch.tensor( + Feature.ATOMIC_GRID_SIZES: torch.tensor( [grid_per_atom] * num_atoms, dtype=torch.int64, device=device ), - "coarse_0_atomic_coords": torch.randn(num_atoms, 3, dtype=dtype, device=device), - "atomic_grid_size_bound_shape": torch.zeros( + Feature.COARSE_0_ATOMIC_COORDS: torch.randn( + num_atoms, 3, dtype=dtype, device=device + ), + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE: torch.zeros( grid_per_atom, 0, dtype=torch.int64, device=device ), } @@ -78,24 +81,26 @@ def make_mol_variable_grid( atomic_grid_sizes: list[int], device: str = "cpu", dtype: torch.dtype = torch.float64, -) -> dict[str, torch.Tensor]: +) -> FeatureMap: """Create a mol dict with variable grid sizes per atom.""" sizes = torch.tensor(atomic_grid_sizes, dtype=torch.int64, device=device) num_atoms = len(atomic_grid_sizes) total_grid = sum(atomic_grid_sizes) size_bound = max(atomic_grid_sizes) return { - "density": torch.randn(2, total_grid, dtype=dtype, device=device), - "grad": torch.randn(2, 3, total_grid, dtype=dtype, device=device), - "kin": torch.randn(2, total_grid, dtype=dtype, device=device), - "grid_coords": torch.randn(total_grid, 3, dtype=dtype, device=device), - "grid_weights": torch.randn(total_grid, dtype=dtype, device=device).abs(), - "atomic_grid_weights": torch.randn( + Feature.DENSITY: torch.randn(2, total_grid, dtype=dtype, device=device), + Feature.GRAD: torch.randn(2, 3, total_grid, dtype=dtype, device=device), + Feature.KIN: torch.randn(2, total_grid, dtype=dtype, device=device), + Feature.GRID_COORDS: torch.randn(total_grid, 3, dtype=dtype, device=device), + Feature.GRID_WEIGHTS: torch.randn(total_grid, dtype=dtype, device=device).abs(), + Feature.ATOMIC_GRID_WEIGHTS: torch.randn( total_grid, dtype=dtype, device=device ).abs(), - "atomic_grid_sizes": sizes, - "coarse_0_atomic_coords": torch.randn(num_atoms, 3, dtype=dtype, device=device), - "atomic_grid_size_bound_shape": torch.zeros( + Feature.ATOMIC_GRID_SIZES: sizes, + Feature.COARSE_0_ATOMIC_COORDS: torch.randn( + num_atoms, 3, dtype=dtype, device=device + ), + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE: torch.zeros( size_bound, 0, dtype=torch.int64, device=device ), } @@ -183,21 +188,21 @@ def test_pack_features_snapshot() -> None: mol = make_mol(4, 10) packed = model.pack_features(mol) - assert packed["density"].shape == (2, 10, 4) - assert packed["kin"].shape == (2, 10, 4) - assert packed["grad"].shape == (2, 3, 10, 4) - assert packed["grid_coords"].shape == (10, 4, 3) - assert packed["atomic_grid_weights"].shape == (10, 4) - assert packed["coarse_0_atomic_coords"].shape == (4, 3) + assert packed[Feature.DENSITY].shape == (2, 10, 4) + assert packed[Feature.KIN].shape == (2, 10, 4) + assert packed[Feature.GRAD].shape == (2, 3, 10, 4) + assert packed[Feature.GRID_COORDS].shape == (10, 4, 3) + assert packed[Feature.ATOMIC_GRID_WEIGHTS].shape == (10, 4) + assert packed[Feature.COARSE_0_ATOMIC_COORDS].shape == (4, 3) torch.testing.assert_close( - packed["density"].sum(), + packed[Feature.DENSITY].sum(), torch.tensor(1.020635438470402e01, dtype=torch.float64), rtol=1e-5, atol=1e-5, ) torch.testing.assert_close( - packed["atomic_grid_weights"].sum(), + packed[Feature.ATOMIC_GRID_WEIGHTS].sum(), torch.tensor(4.032819661608873e01, dtype=torch.float64), rtol=1e-5, atol=1e-5, diff --git a/tests/test_pyscf_gradients.py b/tests/test_pyscf_gradients.py index 48d031d6..d603addf 100644 --- a/tests/test_pyscf_gradients.py +++ b/tests/test_pyscf_gradients.py @@ -5,6 +5,7 @@ from _ridders import num_grad_ridders from pyscf import dft, gto, scf +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase from skala.pyscf import SkalaKS from skala.pyscf.features import generate_features @@ -58,11 +59,11 @@ def test_grid_coords_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grid_coords"] + self.features = [Feature.GRID_COORDS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return mol["grid_coords"].sum() + return mol[Feature.GRID_COORDS].sum() mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) @@ -85,11 +86,11 @@ def test_coarse_0_atomic_coords_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["coarse_0_atomic_coords"] + self.features = [Feature.COARSE_0_ATOMIC_COORDS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return torch.einsum("nx->", mol["coarse_0_atomic_coords"]) + return torch.einsum("nx->", mol[Feature.COARSE_0_ATOMIC_COORDS]) mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) @@ -108,11 +109,11 @@ def test_grid_weights_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grid_weights"] + self.features = [Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return mol["grid_weights"].sum() + return mol[Feature.GRID_WEIGHTS].sum() def finite_difference_nuc_grad( weight_sum: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -126,7 +127,9 @@ def finite_difference_nuc_grad( def weight_sum_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: """Exc wrapper for the finite difference""" mol.set_geom_(nuc_coords.numpy(), "bohr", symmetry=None) - mol_feats["grid_weights"] = torch.from_numpy(minimal_grid(mol).weights) + mol_feats[Feature.GRID_WEIGHTS] = torch.from_numpy( + minimal_grid(mol).weights + ) return weight_sum.get_exc(mol_feats) @@ -171,11 +174,11 @@ def test_density_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["density", "grid_weights"] + self.features = [Feature.DENSITY, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total electron number""" - return (mol["density"] @ mol["grid_weights"]).sum() + return (mol[Feature.DENSITY] @ mol[Feature.GRID_WEIGHTS]).sum() def finite_difference_nuc_grad( dens_sum: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -203,7 +206,7 @@ def dens_sum_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: # calculate analytic result veff = veff_and_expl_nuc_grad( - exc_test, mol, grid, rdm1, nuc_grad_feats={"density"} + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.DENSITY} )[0] ana_grad = nuc_grad_from_veff(mol, veff, rdm1) @@ -224,11 +227,11 @@ def test_grad_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["grad", "grid_weights"] + self.features = [Feature.GRAD, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: return ( - (mol["grad"] ** 2 @ mol["grid_weights"]) + (mol[Feature.GRAD] ** 2 @ mol[Feature.GRID_WEIGHTS]) @ torch.tensor([1.0, 2.0, 3.0], dtype=torch.float64) ).sum() @@ -258,7 +261,9 @@ def grad_func_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: num_grad, num_err = finite_difference_nuc_grad(exc_test, mol, rdm1) # calculate analytic result - veff = veff_and_expl_nuc_grad(exc_test, mol, grid, rdm1, nuc_grad_feats={"grad"})[0] + veff = veff_and_expl_nuc_grad( + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.GRAD} + )[0] ana_grad = nuc_grad_from_veff(mol, veff, rdm1) # This gradient has large-magnitude components whose coarse finite-difference @@ -286,11 +291,11 @@ def test_kin_veff(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["kin", "grid_weights"] + self.features = [Feature.KIN, Feature.GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: """This actually calculates the total kinetic energy number""" - return (mol["kin"] @ mol["grid_weights"]).sum() + return (mol[Feature.KIN] @ mol[Feature.GRID_WEIGHTS]).sum() def finite_difference_nuc_grad( kin_func: ExcFunctionalBase, mol: gto.Mole, rdm1: torch.Tensor @@ -318,7 +323,9 @@ def kin_func_as_nuc_coords_func(nuc_coords: torch.Tensor) -> torch.Tensor: num_grad, num_err = finite_difference_nuc_grad(exc_test, mol, rdm1) # calculate analytic result - veff = veff_and_expl_nuc_grad(exc_test, mol, grid, rdm1, nuc_grad_feats={"kin"})[0] + veff = veff_and_expl_nuc_grad( + exc_test, mol, grid, rdm1, nuc_grad_feats={Feature.KIN} + )[0] ana_grad = nuc_grad_from_veff(mol, veff, rdm1) # Like test_grad_veff, the kinetic-energy gradient has large-magnitude @@ -475,10 +482,10 @@ def test_atomic_grid_weights_gradient(mol_name: str) -> None: class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["atomic_grid_weights"] + self.features = [Feature.ATOMIC_GRID_WEIGHTS] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: - return mol["atomic_grid_weights"].sum() + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + return mol[Feature.ATOMIC_GRID_WEIGHTS].sum() mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) @@ -504,17 +511,17 @@ class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() self.features = [ - "density", - "grid_weights", - "atomic_grid_weights", - "atomic_grid_sizes", - "atomic_grid_size_bound_shape", + Feature.DENSITY, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.ATOMIC_GRID_SIZES, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, ] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: + def get_exc(self, mol: FeatureMap) -> torch.Tensor: # Use density and grid_weights (differentiable) plus atomic_grid_weights (other_feat) - n_electrons = (mol["density"] @ mol["grid_weights"]).sum() - agw_sum = mol["atomic_grid_weights"].sum() + n_electrons = (mol[Feature.DENSITY] @ mol[Feature.GRID_WEIGHTS]).sum() + agw_sum = mol[Feature.ATOMIC_GRID_WEIGHTS].sum() return n_electrons + agw_sum mol = get_mol(mol_name) @@ -540,15 +547,15 @@ class TestFunc(ExcFunctionalBase): def __init__(self) -> None: super().__init__() self.features = [ - "density", - "grid_weights", - "atomic_grid_weights", - "atomic_grid_sizes", - "atomic_grid_size_bound_shape", + Feature.DENSITY, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.ATOMIC_GRID_SIZES, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, ] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: - return (mol["density"] @ mol["grid_weights"]).sum() + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + return (mol[Feature.DENSITY] @ mol[Feature.GRID_WEIGHTS]).sum() mol = get_mol(mol_name) grid, rdm1 = get_grid_and_rdm1(mol) diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py index 7991154a..978696f5 100644 --- a/tests/test_xc_integrator.py +++ b/tests/test_xc_integrator.py @@ -5,6 +5,7 @@ import torch from pyscf import dft, gto +from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase from skala.pyscf import xc_integrator as xc_integrator_module from skala.pyscf.numint import SkalaNumInt @@ -14,10 +15,10 @@ class QuadraticDensityFunctional(ExcFunctionalBase): def __init__(self) -> None: super().__init__() - self.features = ["density"] + self.features = [Feature.DENSITY] - def get_exc(self, mol: dict[str, torch.Tensor]) -> torch.Tensor: - return (mol["density"].square() * mol["grid_weights"]).sum() + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() def test_xc_integrator_returns_tensors_and_xc_only_response( @@ -30,13 +31,13 @@ def fake_generate_features( mol: gto.Mole, dm: torch.Tensor, grids: object, - features: set[str], + features: set[Feature], **kwargs: object, - ) -> dict[str, torch.Tensor]: - assert features == {"density", "grid_weights"} + ) -> FeatureMap: + assert features == {Feature.DENSITY, Feature.GRID_WEIGHTS} return { - "density": dm.sum().reshape(1), - "grid_weights": torch.tensor([2.0], dtype=dm.dtype), + Feature.DENSITY: dm.sum().reshape(1), + Feature.GRID_WEIGHTS: torch.tensor([2.0], dtype=dm.dtype), } monkeypatch.setattr( From 55fe9039cf8ca4c9b5ff58c545718153158e59b8 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 13:44:10 +0200 Subject: [PATCH 25/39] fix enum properly --- src/skala/features.py | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) create mode 100644 src/skala/features.py diff --git a/src/skala/features.py b/src/skala/features.py new file mode 100644 index 00000000..04baf73d --- /dev/null +++ b/src/skala/features.py @@ -0,0 +1,29 @@ +# SPDX-License-Identifier: MIT + +"""Names of built-in molecular features.""" + +from enum import Enum +from typing import TYPE_CHECKING, TypeAlias + +if TYPE_CHECKING: + from torch import Tensor + + +class Feature(str, Enum): # noqa: UP042 - Python 3.10-compatible StrEnum + """String-compatible names of features understood by Skala.""" + + DENSITY = "density" + GRAD = "grad" + KIN = "kin" + LAPL = "lapl" + GRID_COORDS = "grid_coords" + GRID_WEIGHTS = "grid_weights" + ATOMIC_GRID_WEIGHTS = "atomic_grid_weights" + ATOMIC_GRID_SIZES = "atomic_grid_sizes" + ATOMIC_GRID_SIZE_BOUND_SHAPE = "atomic_grid_size_bound_shape" + COARSE_0_ATOMIC_COORDS = "coarse_0_atomic_coords" + + __str__ = str.__str__ + + +FeatureMap: TypeAlias = dict[Feature, "Tensor"] From f5fa1f2424f259b3860c6243f534ea766a27e385 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 13:56:42 +0200 Subject: [PATCH 26/39] make tests nicer --- tests/test_ao_screening.py | 29 +++-------- tests/test_gpu4pyscf_ao_screening.py | 62 +++++++++-------------- tests/test_xc_integrator.py | 74 ++++------------------------ 3 files changed, 41 insertions(+), 124 deletions(-) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index aec997a4..5a96d37c 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -4,7 +4,7 @@ import pytest import torch from pyscf import dft, gto -from utils import patch_ao_screening +from utils import QuadraticFunctional, patch_ao_screening from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase @@ -304,19 +304,6 @@ def fake_make_screen_index( assert not hasattr(grids, "_skala_spatial_grid_layout") -class QuadraticDensityFunctional(ExcFunctionalBase): - def __init__(self) -> None: - super().__init__() - self.features = [ - Feature.ATOMIC_GRID_SIZES, - Feature.DENSITY, - Feature.GRID_WEIGHTS, - ] - - def get_exc(self, mol: FeatureMap) -> torch.Tensor: - return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() - - def test_grid_reuses_spatial_grid_layout_across_numints( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: @@ -348,8 +335,8 @@ def fake_prepare_spatial_grid_layout( "prepare_spatial_grid_layout", fake_prepare_spatial_grid_layout, ) - numint = SkalaNumInt(QuadraticDensityFunctional()) - other_numint = SkalaNumInt(QuadraticDensityFunctional()) + numint = SkalaNumInt(QuadraticFunctional()) + other_numint = SkalaNumInt(QuadraticFunctional()) layout = numint.integrator._get_spatial_grid_layout(carbon, grids) assert other_numint.integrator._get_spatial_grid_layout(carbon, grids) is layout @@ -379,7 +366,7 @@ def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: def test_call_rejects_second_order_evaluation(carbon: gto.Mole) -> None: - numint = SkalaNumInt(QuadraticDensityFunctional()) + numint = SkalaNumInt(QuadraticFunctional()) with pytest.raises(NotImplementedError, match="second-order evaluation"): numint( @@ -502,7 +489,7 @@ def fake_prepare_model_feature_chunks( "screened_feature_jvp", fake_screened_feature_jvp, ) - numint = SkalaNumInt(QuadraticDensityFunctional()) + numint = SkalaNumInt(QuadraticFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) grids = dft.Grids(carbon) grids.weights = np.ones(1) @@ -783,7 +770,7 @@ def _minimal_atom_grid(mol: gto.Mole) -> dft.Grids: def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> None: - numint = SkalaNumInt(QuadraticDensityFunctional()) + numint = SkalaNumInt(QuadraticFunctional()) grids = _minimal_atom_grid(carbon) spatial_grid_layout = prepare_spatial_grid_layout( carbon, @@ -838,7 +825,7 @@ def test_cpu_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) grids = _minimal_atom_grid(mol) ks = FakeKS(mol, grids) - numint = SkalaNumInt(QuadraticDensityFunctional()) + numint = SkalaNumInt(QuadraticFunctional()) mo_coeff = np.eye(mol.nao_nr()) mo_occ = np.ones(mol.nao_nr()) dm1 = np.arange(mol.nao_nr() ** 2, dtype=np.float64).reshape( @@ -888,7 +875,7 @@ def counting_backward_apply(*args: object) -> torch.Tensor: monkeypatch.setattr(ChunkEvalForward, "apply", counting_forward_apply) monkeypatch.setattr(ChunkEvalBackward, "apply", counting_backward_apply) - functional = QuadraticDensityFunctional() + functional = QuadraticFunctional() numint = SkalaNumInt(functional) with patch_ao_screening(True): diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index c3c46b8c..1f1d71a9 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -23,9 +23,9 @@ allow_module_level=True, ) -from utils import patch_ao_screening # noqa: E402 +from utils import QuadraticFunctional, patch_ao_screening # noqa: E402 -from skala.features import Feature, FeatureMap # noqa: E402 +from skala.features import Feature # noqa: E402 from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.pyscf.ao_evaluation import ( # noqa: E402 @@ -46,39 +46,6 @@ """ -class QuadraticDensityFunctional(ExcFunctionalBase): - def __init__(self) -> None: - super().__init__() - self.features = [ - Feature.ATOMIC_GRID_SIZES, - Feature.DENSITY, - Feature.GRID_WEIGHTS, - ] - - def get_exc(self, mol: FeatureMap) -> torch.Tensor: - return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() - - -class QuadraticMGGAFunctional(ExcFunctionalBase): - def __init__(self) -> None: - super().__init__() - self.features = [ - Feature.ATOMIC_GRID_SIZES, - Feature.DENSITY, - Feature.GRAD, - Feature.KIN, - Feature.GRID_WEIGHTS, - ] - - def get_exc(self, mol: FeatureMap) -> torch.Tensor: - energy_density = ( - mol[Feature.DENSITY].square() - + mol[Feature.GRAD].square().sum(dim=-2) - + mol[Feature.KIN].square() - ) - return (energy_density * mol[Feature.GRID_WEIGHTS]).sum() - - def _to_numpy(value: object) -> np.ndarray: return cupy.asnumpy(value) if isinstance(value, cupy.ndarray) else np.asarray(value) @@ -162,7 +129,7 @@ def test_gpu_rks_uks_dense_screened_equivalence( def test_gpu_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) - ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) + ks = SkalaKS(mol, xc=QuadraticFunctional(), with_dftd3=False) ks.grids.level = 0 ks.grids.alignment = 1 ks.grids.build(sort_grids=False) @@ -189,7 +156,7 @@ def test_gpu_response_dense_screened_equivalence() -> None: def test_gpu_uks_response_dense_screened_equivalence() -> None: mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) - ks = SkalaKS(mol, xc=QuadraticDensityFunctional(), with_dftd3=False) + ks = SkalaKS(mol, xc=QuadraticFunctional(), with_dftd3=False) ks.grids.level = 0 ks.grids.alignment = 1 ks.grids.build(sort_grids=False) @@ -212,8 +179,27 @@ def test_gpu_uks_response_dense_screened_equivalence() -> None: def test_gpu_multiblock_mgga_response_dense_screened_equivalence() -> None: + """Exercise screened MGGA Hessian-vector products across multiple GPU blocks. + + Density-only and single-block cases cannot expose errors in block-local JVP + assembly, spatial permutation, or reduction of vector gradient features. The + large basis and grid force multiple GPU AO blocks, while the quadratic density, + gradient, and kinetic terms give a nonzero response for every MGGA feature path. + """ mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) - ks = SkalaKS(mol, xc=QuadraticMGGAFunctional(), with_dftd3=False) + ks = SkalaKS( + mol, + xc=QuadraticFunctional( + [ + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.GRID_WEIGHTS, + ] + ), + with_dftd3=False, + ) ks.grids.level = 1 ks.grids.alignment = 1 ks.grids.build(sort_grids=False) diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py index 978696f5..558038c1 100644 --- a/tests/test_xc_integrator.py +++ b/tests/test_xc_integrator.py @@ -1,29 +1,23 @@ -from typing import Any, cast - -import numpy as np import pytest import torch from pyscf import dft, gto +from utils import QuadraticFunctional from skala.features import Feature, FeatureMap -from skala.functional.base import ExcFunctionalBase from skala.pyscf import xc_integrator as xc_integrator_module -from skala.pyscf.numint import SkalaNumInt from skala.pyscf.xc_integrator import XCIntegrator, XCResult -class QuadraticDensityFunctional(ExcFunctionalBase): - def __init__(self) -> None: - super().__init__() - self.features = [Feature.DENSITY] - - def get_exc(self, mol: FeatureMap) -> torch.Tensor: - return (mol[Feature.DENSITY].square() * mol[Feature.GRID_WEIGHTS]).sum() - - def test_xc_integrator_returns_tensors_and_xc_only_response( monkeypatch: pytest.MonkeyPatch, ) -> None: + """Pin the tensor-level integrator contract independently of PySCF NumInt. + + The synthetic features give closed-form electron count, energy, potential, and + Hessian action values. Checking them here verifies that ``XCIntegrator`` returns + tensors and that ``gen_response`` contains only the XC Hessian action, without + the Coulomb response that the higher-level NumInt wrapper adds. + """ mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) grids = dft.Grids(mol) @@ -45,7 +39,7 @@ def fake_generate_features( "generate_features", fake_generate_features, ) - integrator = XCIntegrator(QuadraticDensityFunctional()) + integrator = XCIntegrator(QuadraticFunctional([Feature.DENSITY])) dm = torch.tensor([[1.0, 2.0], [2.0, 3.0]], dtype=torch.float64) result = integrator(mol, grids, dm) @@ -56,53 +50,3 @@ def fake_generate_features( torch.testing.assert_close(result.energy, dm.new_tensor(128.0)) torch.testing.assert_close(result.potential, torch.full_like(dm, 32.0)) torch.testing.assert_close(response(torch.ones_like(dm)), torch.full_like(dm, 16.0)) - - -class FakeKS: - def __init__(self, mol: gto.Mole) -> None: - self.mol = mol - self.grids = dft.Grids(mol) - self.max_memory = 123 - - def make_rdm1(self, mo_coeff: np.ndarray, mo_occ: np.ndarray) -> np.ndarray: - return np.eye(self.mol.nao_nr()) - - def get_j(self, mol: gto.Mole, dm: np.ndarray, hermi: int) -> np.ndarray: - assert hermi == 1 - return np.full_like(dm, 3.0) - - -class FakeXCIntegrator: - device = torch.device("cpu") - - def __init__(self) -> None: - self.calls: list[tuple[int, float | None]] = [] - - def gen_response( - self, - mol: gto.Mole, - grids: object, - dm0: torch.Tensor, - max_memory: int, - safety_fraction: float | None, - ) -> Any: - self.calls.append((max_memory, safety_fraction)) - return lambda dm1: 2 * dm1 - - -def test_numint_response_adds_coulomb_to_xc_response() -> None: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) - ks = FakeKS(mol) - numint = SkalaNumInt(QuadraticDensityFunctional()) - fake_integrator = FakeXCIntegrator() - numint.integrator = cast(XCIntegrator, fake_integrator) - - response = numint.gen_response( - np.eye(mol.nao_nr()), - np.ones(mol.nao_nr()), - ks=cast(Any, ks), - safety_fraction=0.6, - ) - - np.testing.assert_allclose(response(np.ones((mol.nao_nr(), mol.nao_nr()))), 5.0) - assert fake_integrator.calls == [(123, 0.6)] From 28cf7ef8f8d5fa1ccb0a0b74241eada7c9ac3d5d Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 13:57:01 +0200 Subject: [PATCH 27/39] make tests nicer --- tests/utils.py | 98 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 98 insertions(+) create mode 100644 tests/utils.py diff --git a/tests/utils.py b/tests/utils.py new file mode 100644 index 00000000..3508b9f4 --- /dev/null +++ b/tests/utils.py @@ -0,0 +1,98 @@ +"""Shared functional and route-control helpers for tests.""" + +from collections.abc import Iterable, Iterator +from contextlib import contextmanager +from types import ModuleType +from unittest.mock import patch + +import torch + +from skala.features import Feature, FeatureMap +from skala.functional.base import ExcFunctionalBase +from skala.pyscf import xc_integrator as xc_integrator_module + + +class QuadraticFunctional(ExcFunctionalBase): + """Functional whose energy is a weighted sum of squared AO-derived features.""" + + def __init__( + self, + features: Iterable[Feature] = ( + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRID_WEIGHTS, + ), + ) -> None: + """Initialize the functional with its required model features. + + AO-derived entries contribute quadratic energy terms. Other entries declare + metadata needed by the evaluation route but do not contribute to the energy. + + Args: + features: Features required from the model evaluation. + + Raises: + ValueError: If no AO-derived feature is selected. + """ + super().__init__() + self.features = list(features) + self._quadratic_features = tuple( + feature + for feature in self.features + if feature + in { + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.LAPL, + } + ) + if not self._quadratic_features: + raise ValueError("At least one AO-derived feature must be selected") + + def get_exc(self, mol: FeatureMap) -> torch.Tensor: + """Return the grid-integrated quadratic feature energy. + + The vector components of the density gradient are summed before combining + them with scalar density, kinetic, or Laplacian terms. + + Args: + mol: Model features keyed by their feature identifiers. + + Returns: + Scalar exchange-correlation energy. + """ + grid_weights = mol[Feature.GRID_WEIGHTS] + quadratic_terms = [ + ( + mol[feature].square().sum(dim=-2) + if feature is Feature.GRAD + else mol[feature].square() + ) + for feature in self._quadratic_features + ] + energy_density = torch.stack(quadratic_terms).sum(dim=0) + return (energy_density * grid_weights).sum() + + +@contextmanager +def patch_ao_screening( + enabled: bool, + module: ModuleType = xc_integrator_module, +) -> Iterator[None]: + """Temporarily force the AO-screening route decision. + + Args: + enabled: Whether calls should select screened AO evaluation. + module: Module whose ``_should_screen_aos`` decision function is patched. + + Yields: + Control while the forced decision is active. The previous function is + restored when the context exits. + """ + with patch.object( + module, + "_should_screen_aos", + return_value=enabled, + ): + yield From 8a3b11426826d22ffca9fd0b8034c4bb1a3e3507 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 14:18:00 +0200 Subject: [PATCH 28/39] fix doc tests --- src/skala/functional/__init__.py | 9 +-- src/skala/gpu4pyscf/__init__.py | 8 +-- tests/test_ao_screening.py | 81 +++++++++++++++++------- tests/test_ao_screening_benchmark.py | 7 ++ tests/test_gpu4pyscf_ao_screening.py | 95 +++++++++++----------------- 5 files changed, 112 insertions(+), 88 deletions(-) diff --git a/src/skala/functional/__init__.py b/src/skala/functional/__init__.py index 4be42f1f..b31efae7 100644 --- a/src/skala/functional/__init__.py +++ b/src/skala/functional/__init__.py @@ -80,12 +80,13 @@ def load_functional( name string for PySCF-native functionals. Example: + >>> from skala.features import Feature >>> func = load_functional("skala-1.1") - >>> func.features - ['density', 'kin', 'grad', 'grid_coords', 'grid_weights', ... + >>> func.features[:3] == [Feature.DENSITY, Feature.KIN, Feature.GRAD] + True >>> func = load_functional("lda") - >>> func.features - ['density', 'grid_weights'] + >>> func.features == [Feature.DENSITY, Feature.GRID_WEIGHTS] + True >>> load_functional("b3lyp") 'b3lyp' """ diff --git a/src/skala/gpu4pyscf/__init__.py b/src/skala/gpu4pyscf/__init__.py index bab92804..262ae5fa 100644 --- a/src/skala/gpu4pyscf/__init__.py +++ b/src/skala/gpu4pyscf/__init__.py @@ -86,7 +86,7 @@ def SkalaKS( >>> ks = ks.set(verbose=0) >>> energy = ks.kernel() >>> print(energy) # DOCTEST: Ellipsis - -1.142773... + -1.143024... >>> ks = ks.nuc_grad_method() >>> gradient = ks.kernel() >>> print(abs(gradient).mean()) # DOCTEST: Ellipsis @@ -165,12 +165,12 @@ def SkalaRKS( >>> import torch >>> >>> mol = gto.M(atom="H 0 0 0; H 0 0 1", basis="def2-svp") - >>> ks = SkalaRKS(mol, xc=load_functional("skala-1.1", device=torch.device("cuda:0")), with_density_fit=True)(verbose=0) + >>> ks = SkalaRKS(mol, xc=load_functional("skala-1.1", device=torch.device("cuda:0")), with_density_fit=True, auxbasis="def2-svp-jkfit")(verbose=0) >>> ks # DOCTEST: Ellipsis >>> energy = ks.kernel() >>> print(energy) # DOCTEST: Ellipsis - -1.142773... + -1.143024... """ if isinstance(xc, str): xc = load_functional(xc, device=torch.device("cuda:0")) @@ -247,7 +247,7 @@ def SkalaUKS( >>> energy = ks.kernel() >>> print(energy) # DOCTEST: Ellipsis - -0.499031... + -0.499123... """ if isinstance(xc, str): xc = load_functional(xc, device=torch.device("cuda:0")) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 5a96d37c..288d6f56 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -1,4 +1,5 @@ from collections.abc import Callable, Iterator +from typing import Any import numpy as np import pytest @@ -240,6 +241,13 @@ def test_decompose_grid_into_spatial_blocks_handles_identical_points() -> None: def test_prepare_spatially_sorted_cpu_grids( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: + """Keep block-local AO masks aligned with a reversible spatial grid ordering. + + Screened AO evaluation reorders atom-major grid points into spatial blocks before + PySCF builds its screening mask. Coordinates, weights, and mask must share that + ordering, while the saved inverse permutation restores model features and leaves + the caller's original grid unchanged. + """ coords = np.arange(18, dtype=np.float64).reshape(6, 3) weights = np.arange(6, dtype=np.float64) + 10 grids = dft.Grids(carbon) @@ -307,6 +315,7 @@ def fake_make_screen_index( def test_grid_reuses_spatial_grid_layout_across_numints( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: + """Cache one spatial layout on each grid independently of the NumInt instance.""" grids = dft.Grids(carbon) grids.coords = np.arange(18, dtype=np.float64).reshape(6, 3) grids.weights = np.arange(6, dtype=np.float64) @@ -523,7 +532,13 @@ def fake_prepare_model_feature_chunks( def test_feature_block_helper_localizes_derivative_vectors() -> None: - """Use AO slices for linear JVPs and grid slices for feature VJPs.""" + """Apply derivative vectors in the local coordinate space of one AO block. + + A screened block contains only selected AO rows and a slice of the global grid. + The forward JVP must therefore use the active-AO density submatrix, while the + adjoint calculation must select only this block's grid cotangent. Comparing both + operations with direct local formulas catches mixing up AO and grid localization. + """ feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) block = _AOBlock( ao_values=torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float64), @@ -686,6 +701,7 @@ def block_loop( def test_cpu_all_active_block_uses_dense_sentinel( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: + """Represent an all-active block by ``None`` and a later sparse block by indices.""" block_size = dft.gen_grid.BLKSIZE ngrids = 2 * block_size grids = dft.Grids(carbon) @@ -729,6 +745,7 @@ def block_loop( def test_cpu_no_active_aos_returns_full_zero_derivatives( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: + """Use an empty screen mask and verify full-size zero features, VXC, and HVP.""" ngrids = dft.gen_grid.BLKSIZE grids = dft.Grids(carbon) grids.coords = np.zeros((ngrids, 3)) @@ -784,17 +801,45 @@ def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> No assert vars(grids)["_skala_spatial_grid_layout"] is spatial_grid_layout -@pytest.mark.parametrize("unrestricted", [False, True]) +@pytest.mark.parametrize( + ("atom", "spin", "mean_field_factory", "integration_method"), + [ + pytest.param( + "H 0 0 0; H 0 0 0.74", + 0, + dft.RKS, + SkalaNumInt.nr_rks, + id="rks", + ), + pytest.param( + "H 0 0 0", + 1, + dft.UKS, + SkalaNumInt.nr_uks, + id="uks", + ), + ], +) +@pytest.mark.parametrize( + ("result_index", "rtol", "atol"), + [ + pytest.param(0, 1e-10, 1e-11, id="electron-count"), + pytest.param(1, 1e-9, 1e-10, id="energy"), + pytest.param(2, 1e-8, 1e-10, id="potential"), + ], +) def test_cpu_rks_uks_dense_screened_equivalence( load_functional_cached: Callable[..., ExcFunctionalBase | str], - unrestricted: bool, + atom: str, + spin: int, + mean_field_factory: Callable[[gto.Mole], Any], + integration_method: Callable[..., tuple[float | np.ndarray, float, np.ndarray]], + result_index: int, + rtol: float, + atol: float, ) -> None: - if unrestricted: - mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) - mean_field = dft.UKS(mol) - else: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) - mean_field = dft.RKS(mol) + mol = gto.M(atom=atom, basis="sto-3g", spin=spin, verbose=0) + mean_field = mean_field_factory(mol) functional = load_functional_cached("skala-1.1") assert isinstance(functional, ExcFunctionalBase) @@ -803,22 +848,14 @@ def test_cpu_rks_uks_dense_screened_equivalence( dm = mean_field.get_init_guess() with patch_ao_screening(False): - dense = ( - numint.nr_uks(mol, grids, None, dm) - if unrestricted - else numint.nr_rks(mol, grids, None, dm) - ) + dense = integration_method(numint, mol, grids, None, dm) with patch_ao_screening(True): - screened = ( - numint.nr_uks(mol, grids, None, dm) - if unrestricted - else numint.nr_rks(mol, grids, None, dm) - ) + screened = integration_method(numint, mol, grids, None, dm) - assert np.allclose(dense[0], screened[0], rtol=1e-10, atol=1e-11) - assert np.isclose(dense[1], screened[1], rtol=1e-9, atol=1e-10) - assert np.allclose(dense[2], screened[2], rtol=1e-8, atol=1e-10) + assert np.allclose( + dense[result_index], screened[result_index], rtol=rtol, atol=atol + ) def test_cpu_response_dense_screened_equivalence() -> None: diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 39426665..da0e6048 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -1,3 +1,10 @@ +"""Benchmark dense and screened AO integration across CPU and GPU backends. + +The module compares numerical agreement, runtime, and peak allocations on a small +acene ladder. Profiling workloads run each route in an isolated process so allocator +state and backend initialization do not contaminate the measurements. +""" + from __future__ import annotations import multiprocessing as mp diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 1f1d71a9..e3619d3f 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -46,10 +46,6 @@ """ -def _to_numpy(value: object) -> np.ndarray: - return cupy.asnumpy(value) if isinstance(value, cupy.ndarray) else np.asarray(value) - - def test_prepare_spatially_sorted_gpu_grids() -> None: mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) coords = cupy.asarray( @@ -88,15 +84,20 @@ def test_prepare_spatially_sorted_gpu_grids() -> None: assert layout.inverse_permutation.device.type == "cuda" -@pytest.mark.parametrize("unrestricted", [False, True]) +@pytest.mark.parametrize( + ("atom", "spin", "integration_method_name"), + [ + pytest.param("H 0 0 0; H 0 0 0.74", 0, "nr_rks", id="rks"), + pytest.param("H 0 0 0", 1, "nr_uks", id="uks"), + ], +) def test_gpu_rks_uks_dense_screened_equivalence( load_functional_cached: Callable[..., ExcFunctionalBase | str], - unrestricted: bool, + atom: str, + spin: int, + integration_method_name: str, ) -> None: - if unrestricted: - mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) - else: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + mol = gto.M(atom=atom, basis="sto-3g", spin=spin, verbose=0) functional = load_functional_cached("skala-1.1", device=torch.device("cuda:0")) assert isinstance(functional, ExcFunctionalBase) @@ -105,40 +106,42 @@ def test_gpu_rks_uks_dense_screened_equivalence( ks.grids.alignment = 1 ks.grids.build(sort_grids=False) dm = ks.get_init_guess() + integrate = getattr(ks._numint, integration_method_name) with patch_ao_screening(False): - dense = ( - ks._numint.nr_uks(mol, ks.grids, None, dm) - if unrestricted - else ks._numint.nr_rks(mol, ks.grids, None, dm) - ) + dense = integrate(mol, ks.grids, None, dm) with patch_ao_screening(True): - screened = ( - ks._numint.nr_uks(mol, ks.grids, None, dm) - if unrestricted - else ks._numint.nr_rks(mol, ks.grids, None, dm) - ) + screened = integrate(mol, ks.grids, None, dm) - assert np.allclose(_to_numpy(dense[0]), _to_numpy(screened[0]), rtol=1e-9) + cupy.testing.assert_allclose(dense[0], screened[0], rtol=1e-9) assert np.isclose(dense[1], screened[1], rtol=1e-9) - assert np.allclose( - _to_numpy(dense[2]), _to_numpy(screened[2]), rtol=1e-8, atol=2e-9 - ) + cupy.testing.assert_allclose(dense[2], screened[2], rtol=1e-8, atol=2e-9) -def test_gpu_response_dense_screened_equivalence() -> None: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) +@pytest.mark.parametrize( + ("atom", "spin", "spin_shape"), + [ + pytest.param("H 0 0 0; H 0 0 0.74", 0, (), id="rks"), + pytest.param("H 0 0 0", 1, (2,), id="uks"), + ], +) +def test_gpu_response_dense_screened_equivalence( + atom: str, spin: int, spin_shape: tuple[int, ...] +) -> None: + mol = gto.M(atom=atom, basis="sto-3g", spin=spin, verbose=0) ks = SkalaKS(mol, xc=QuadraticFunctional(), with_dftd3=False) ks.grids.level = 0 ks.grids.alignment = 1 ks.grids.build(sort_grids=False) - mo_coeff = cupy.eye(mol.nao_nr()) - mo_occ = cupy.ones(mol.nao_nr()) + matrix_shape = spin_shape + (mol.nao_nr(), mol.nao_nr()) + mo_coeff = cupy.broadcast_to(cupy.eye(mol.nao_nr()), matrix_shape).copy() + mo_occ = cupy.ones(spin_shape + (mol.nao_nr(),)) dm1 = cupy.arange(mol.nao_nr() ** 2, dtype=cupy.float64).reshape( mol.nao_nr(), mol.nao_nr() ) dm1 += dm1.T + dm1 = cupy.broadcast_to(dm1, matrix_shape).copy() with patch_ao_screening(False): dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) @@ -146,33 +149,9 @@ def test_gpu_response_dense_screened_equivalence() -> None: with patch_ao_screening(True): screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - assert np.allclose( - _to_numpy(dense_response(dm1)), - _to_numpy(screened_response(dm1)), - rtol=1e-9, - atol=1e-10, - ) - - -def test_gpu_uks_response_dense_screened_equivalence() -> None: - mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) - ks = SkalaKS(mol, xc=QuadraticFunctional(), with_dftd3=False) - ks.grids.level = 0 - ks.grids.alignment = 1 - ks.grids.build(sort_grids=False) - mo_coeff = cupy.stack((cupy.eye(mol.nao_nr()), cupy.eye(mol.nao_nr()))) - mo_occ = cupy.ones((2, mol.nao_nr())) - dm1 = cupy.ones((2, mol.nao_nr(), mol.nao_nr())) - - with patch_ao_screening(False): - dense_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - - with patch_ao_screening(True): - screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - - np.testing.assert_allclose( - _to_numpy(screened_response(dm1)), - _to_numpy(dense_response(dm1)), + cupy.testing.assert_allclose( + screened_response(dm1), + dense_response(dm1), rtol=1e-9, atol=1e-10, ) @@ -214,9 +193,9 @@ def test_gpu_multiblock_mgga_response_dense_screened_equivalence() -> None: with patch_ao_screening(True): screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks) - np.testing.assert_allclose( - _to_numpy(screened_response(dm1)), - _to_numpy(dense_response(dm1)), + cupy.testing.assert_allclose( + screened_response(dm1), + dense_response(dm1), rtol=1e-9, atol=1e-9, ) From 740085d8a2d87a3f37b1ea86b1bb9a34f581d41e Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 17:09:19 +0200 Subject: [PATCH 29/39] clean notebooks --- .github/workflows/test.yml | 35 ++++++++++++++++++--- .pre-commit-config.yaml | 15 +++++++++ docs/ase.ipynb | 3 +- docs/pyscf/scf_settings.ipynb | 54 ++++++--------------------------- docs/pyscf/singlepoint.ipynb | 3 +- environment-cpu.yml | 3 +- environment-gpu.yml | 2 +- pyproject.toml | 2 +- src/skala/functional/density.py | 2 +- 9 files changed, 62 insertions(+), 57 deletions(-) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index a20d0b8d..d7d41f89 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -30,14 +30,9 @@ jobs: - name: Run pre-commit hooks run: | - pre-commit install pre-commit run --all-files shell: micromamba-shell {0} - - name: Run mypy - run: mypy . - shell: micromamba-shell {0} - test: runs-on: ubuntu-latest needs: @@ -81,3 +76,33 @@ jobs: --pyargs skala tests/ shell: micromamba-shell {0} + + profiling: + name: "Profiling (Python=3.12 & PySCF=2.9)" + runs-on: ubuntu-latest + needs: + - lint + env: + OMP_NUM_THREADS: 4 + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 + + - name: Setup micromamba + uses: mamba-org/setup-micromamba@4b9113af4fba0e9e1124b252dd6497a419e7396d # v1 + with: + environment-file: environment-cpu.yml + environment-name: skala + cache-environment: true + cache-downloads: true + create-args: >- + python=3.12 + pyscf=2.9 + + - name: Install package in development mode + run: | + pip install -e . --no-deps + shell: micromamba-shell {0} + + - name: Run profiling tests + run: pytest -v -m profiling tests/test_ao_screening_benchmark.py + shell: micromamba-shell {0} diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 29bb12e9..1ea85258 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -10,3 +10,18 @@ repos: # Run the formatter. - id: ruff-format args: [--config, pyproject.toml] + + - repo: https://github.com/srstevenson/nb-clean + rev: f745b986570ef12cfbe0cfe20a7e0271c328914f # frozen: 4.0.1 + hooks: + - id: nb-clean + + - repo: local + hooks: + - id: mypy + name: mypy + entry: mypy + language: system + args: [--config-file, pyproject.toml, --num-workers, "4", .] + pass_filenames: false + always_run: true diff --git a/docs/ase.ipynb b/docs/ase.ipynb index c5d41e5f..26555d8a 100644 --- a/docs/ase.ipynb +++ b/docs/ase.ipynb @@ -339,8 +339,7 @@ "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.13.5" + "pygments_lexer": "ipython3" } }, "nbformat": 4, diff --git a/docs/pyscf/scf_settings.ipynb b/docs/pyscf/scf_settings.ipynb index 34c2e7e1..04aa3ffc 100644 --- a/docs/pyscf/scf_settings.ipynb +++ b/docs/pyscf/scf_settings.ipynb @@ -31,7 +31,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": null, "id": "ed3a3d47", "metadata": {}, "outputs": [], @@ -52,25 +52,10 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "id": "a07500d7", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "converged SCF energy = -1.07091605172654\n", - "**** SCF Summaries ****\n", - "Total Energy = -1.070916051726540\n", - "Nuclear Repulsion Energy = 0.377654773327513\n", - "One-electron Energy = -1.897310624972360\n", - "Two-electron Coulomb Energy = 0.997543909702505\n", - "DFT Exchange-Correlation Energy = -0.548804109784197\n", - "Empirical Dispersion Energy = -0.000328948758201\n" - ] - } - ], + "outputs": [], "source": [ "ks = SkalaKS(mol, xc=\"skala-1.1\")\n", "ks.kernel()\n", @@ -88,19 +73,10 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": null, "id": "bce2aea4", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "1e-09\n", - "None\n" - ] - } - ], + "outputs": [], "source": [ "print(ks.conv_tol)\n", "print(ks.conv_tol_grad)" @@ -116,7 +92,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "id": "471f29ee", "metadata": {}, "outputs": [], @@ -142,19 +118,10 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "id": "482ff510", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "5e-06\n", - "0.001\n" - ] - } - ], + "outputs": [], "source": [ "print(ks.conv_tol)\n", "print(ks.conv_tol_grad)" @@ -184,10 +151,9 @@ "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.13.7" + "pygments_lexer": "ipython3" } }, "nbformat": 4, "nbformat_minor": 5 -} \ No newline at end of file +} diff --git a/docs/pyscf/singlepoint.ipynb b/docs/pyscf/singlepoint.ipynb index 8313c371..8d9aaf3c 100644 --- a/docs/pyscf/singlepoint.ipynb +++ b/docs/pyscf/singlepoint.ipynb @@ -129,8 +129,7 @@ "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.14.2" + "pygments_lexer": "ipython3" } }, "nbformat": 4, diff --git a/environment-cpu.yml b/environment-cpu.yml index 737cdf5c..aa4f4353 100644 --- a/environment-cpu.yml +++ b/environment-cpu.yml @@ -15,6 +15,7 @@ dependencies: - pytorch * cpu_* - qcelemental # Testing and development + - memray - pre-commit - pytest - pytest-benchmark @@ -22,6 +23,6 @@ dependencies: - pytest-randomly - pytest-timeout - ruff - - mypy + - mypy >=2 - pip: - huggingface_hub diff --git a/environment-gpu.yml b/environment-gpu.yml index ab454fca..bda42cf4 100644 --- a/environment-gpu.yml +++ b/environment-gpu.yml @@ -25,7 +25,7 @@ dependencies: - pytest-randomly - pytest-timeout - ruff - - mypy + - mypy >=2 - pip: - huggingface_hub - gpu4pyscf-cuda12x >=1.6,<1.8,!=1.7.1,!=1.7.2 diff --git a/pyproject.toml b/pyproject.toml index ffa8c876..0d389609 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,7 +31,7 @@ dependencies = [ optional-dependencies.dev = [ "pre-commit", - "mypy", + "mypy>=2", "memray", "pytest", "pytest-benchmark", diff --git a/src/skala/functional/density.py b/src/skala/functional/density.py index 292b97dc..f25ad1cd 100644 --- a/src/skala/functional/density.py +++ b/src/skala/functional/density.py @@ -17,7 +17,7 @@ from skala.features import Feature, FeatureMap EPS = 1e-10 -IMMUTABLES: frozenset[str] = frozenset([Feature.GRID_COORDS, Feature.GRID_WEIGHTS]) +IMMUTABLES: frozenset[Feature] = frozenset([Feature.GRID_COORDS, Feature.GRID_WEIGHTS]) def _map(mol_features: FeatureMap, f: Callable[[Tensor], Tensor]) -> FeatureMap: From c0a4a9d4f875767eee676089fe002e75a5b12af2 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Fri, 7 Aug 2026 17:22:46 +0200 Subject: [PATCH 30/39] commit notebooks for analysis --- benchmarks/.gitignore | 1 + .../pyscf_ao_screening_performance.ipynb | 783 +++++++++ ...scf_ao_screening_rotation_comparison.ipynb | 1395 +++++++++++++++++ .../run_pyscf_ao_screening_benchmark.py | 1130 +++++++++++++ ...n_pyscf_ao_screening_rotation_benchmark.py | 409 +++++ benchmarks/vxc_accuracy_grid_grouping.ipynb | 703 +++++++++ 6 files changed, 4421 insertions(+) create mode 100644 benchmarks/.gitignore create mode 100644 benchmarks/pyscf_ao_screening_performance.ipynb create mode 100644 benchmarks/pyscf_ao_screening_rotation_comparison.ipynb create mode 100644 benchmarks/run_pyscf_ao_screening_benchmark.py create mode 100644 benchmarks/run_pyscf_ao_screening_rotation_benchmark.py create mode 100644 benchmarks/vxc_accuracy_grid_grouping.ipynb diff --git a/benchmarks/.gitignore b/benchmarks/.gitignore new file mode 100644 index 00000000..fbca2253 --- /dev/null +++ b/benchmarks/.gitignore @@ -0,0 +1 @@ +results/ diff --git a/benchmarks/pyscf_ao_screening_performance.ipynb b/benchmarks/pyscf_ao_screening_performance.ipynb new file mode 100644 index 00000000..9a6a39a8 --- /dev/null +++ b/benchmarks/pyscf_ao_screening_performance.ipynb @@ -0,0 +1,783 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "236c0ac1", + "metadata": {}, + "source": [ + "# Skala PySCF / GPU4PySCF Benchmark Results\n", + "\n", + "This notebook loads and compares JSON produced by `benchmarks/run_pyscf_ao_screening_benchmark.py`. It does not construct molecules, load Skala, or execute benchmark workloads.\n", + "\n", + "Generate result files from a shell before opening the analysis cells:\n", + "\n", + "```bash\n", + "python benchmarks/run_pyscf_ao_screening_benchmark.py --label mr\n", + "python benchmarks/run_pyscf_ao_screening_benchmark.py \\\n", + " --label main \\\n", + " --source-root /path/to/main-worktree\n", + "```\n", + "\n", + "Add `--smoke` to run only C4H10, or `--preflight-only` to validate the selected checkout and environment without collecting measurements." + ] + }, + { + "cell_type": "markdown", + "id": "e758ba91", + "metadata": {}, + "source": [ + "## Select Result Files\n", + "\n", + "By default, every matching result in `benchmarks/results` is loaded. Replace `SELECTED_RESULT_FILES` with an explicit list when comparing only particular labels or commits." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d6909409", + "metadata": {}, + "outputs": [], + "source": [ + "from __future__ import annotations\n", + "\n", + "import json\n", + "from pathlib import Path\n", + "from typing import Any\n", + "\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "\n", + "MODES = (\"cpu\", \"cpu_dense\", \"gpu\")\n", + "MEASUREMENTS = (\"runtime\", \"memory\")\n", + "TERMINAL_STATUSES = {\"ok\", \"timeout\", \"oom\", \"error\", \"skipped_after_resource_failure\"}\n", + "SCIENTIFIC_CONFIG_KEYS = (\n", + " \"functional\",\n", + " \"basis\",\n", + " \"grid_level\",\n", + " \"grid_alignment\",\n", + " \"max_memory_mb\",\n", + " \"cpu_threads\",\n", + " \"full_carbon_counts\",\n", + " \"expected_ao_counts\",\n", + ")\n", + "\n", + "\n", + "def find_repository_root(start: Path) -> Path:\n", + " for candidate in (start.resolve(), *start.resolve().parents):\n", + " if (candidate / \"pyproject.toml\").is_file() and (\n", + " candidate / \"benchmarks\"\n", + " ).is_dir():\n", + " return candidate\n", + " raise FileNotFoundError(f\"Could not find the Skala repository above {start}\")\n", + "\n", + "\n", + "REPOSITORY_ROOT = find_repository_root(Path.cwd())\n", + "RESULTS_DIR = REPOSITORY_ROOT / \"benchmarks\" / \"results\"\n", + "SELECTED_RESULT_FILES = sorted(RESULTS_DIR.glob(\"skala-pyscf-ao-screening-*.json\"))\n", + "\n", + "print(f\"Selected {len(SELECTED_RESULT_FILES)} result file(s) from {RESULTS_DIR}\")\n", + "for result_file in SELECTED_RESULT_FILES:\n", + " print(f\" {result_file.name}\")" + ] + }, + { + "cell_type": "markdown", + "id": "8284d45b", + "metadata": {}, + "source": [ + "## Load and Validate\n", + "\n", + "The checks below surface incompatible schemas, scientific settings, hardware, routing implementations, dirty checkouts, unexpected statuses, AO counts, and production-versus-CPU-dense fingerprint differences." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "f4de7a3a", + "metadata": {}, + "outputs": [], + "source": [ + "def validate_result_document(document: dict[str, Any]) -> list[str]:\n", + " errors: list[str] = []\n", + " if document.get(\"schema_version\") != 1:\n", + " errors.append(f\"Unsupported schema version: {document.get('schema_version')}\")\n", + " for formula, molecule in document.get(\"molecules\", {}).items():\n", + " observed = molecule.get(\"observed\")\n", + " if observed is not None:\n", + " if observed.get(\"actual_aos\") != molecule.get(\"expected_aos\"):\n", + " errors.append(\n", + " f\"{formula}: expected {molecule.get('expected_aos')} AOs, \"\n", + " f\"observed {observed.get('actual_aos')}\"\n", + " )\n", + " carbon_count = int(molecule[\"carbon_count\"])\n", + " expected_electrons = 8 * carbon_count + 2\n", + " if observed.get(\"electron_count\") != expected_electrons:\n", + " errors.append(\n", + " f\"{formula}: expected {expected_electrons} electrons, \"\n", + " f\"observed {observed.get('electron_count')}\"\n", + " )\n", + " for mode, mode_record in molecule.get(\"modes\", {}).items():\n", + " for measurement in MEASUREMENTS:\n", + " result = mode_record.get(measurement)\n", + " if result is not None and result.get(\"status\") not in TERMINAL_STATUSES:\n", + " errors.append(\n", + " f\"{formula} {mode} {measurement}: unknown status {result.get('status')}\"\n", + " )\n", + " return errors\n", + "\n", + "\n", + "def preferred_fingerprint(mode_record: dict[str, Any]) -> dict[str, float] | None:\n", + " for measurement in MEASUREMENTS:\n", + " result = mode_record.get(measurement, {})\n", + " if result.get(\"status\") == \"ok\" and \"fingerprint\" in result:\n", + " return result[\"fingerprint\"]\n", + " return None\n", + "\n", + "\n", + "def fingerprint_warnings(document: dict[str, Any]) -> list[str]:\n", + " messages: list[str] = []\n", + " for formula, molecule in document[\"molecules\"].items():\n", + " reference = preferred_fingerprint(molecule[\"modes\"][\"cpu_dense\"])\n", + " if reference is None:\n", + " continue\n", + " for production_mode, rtol, atol in (\n", + " (\"cpu\", 1e-8, 5e-8),\n", + " (\"gpu\", 1e-7, 2e-7),\n", + " ):\n", + " production = preferred_fingerprint(molecule[\"modes\"][production_mode])\n", + " if production is None:\n", + " continue\n", + " for key in production:\n", + " if not np.isclose(\n", + " production[key], reference[key], rtol=rtol, atol=atol\n", + " ):\n", + " messages.append(\n", + " f\"{formula} {production_mode}/cpu_dense: {key} differs \"\n", + " f\"({production[key]:.12g} vs {reference[key]:.12g})\"\n", + " )\n", + " return messages\n", + "\n", + "\n", + "def comparison_warnings(documents: list[dict[str, Any]]) -> list[str]:\n", + " messages: list[str] = []\n", + " if not documents:\n", + " return [\"No result documents were selected\"]\n", + " reference = documents[0]\n", + " reference_config = reference[\"configuration\"]\n", + " reference_environment = reference[\"environment\"]\n", + " for document in documents:\n", + " label = document[\"run_label\"]\n", + " if document[\"source\"].get(\"dirty\"):\n", + " messages.append(f\"{label}: source checkout is dirty\")\n", + " messages.extend(\n", + " f\"{label}: {error}\" for error in validate_result_document(document)\n", + " )\n", + " messages.extend(\n", + " f\"{label}: {warning}\" for warning in fingerprint_warnings(document)\n", + " )\n", + " for document in documents[1:]:\n", + " label = document[\"run_label\"]\n", + " for key in SCIENTIFIC_CONFIG_KEYS:\n", + " if document[\"configuration\"].get(key) != reference_config.get(key):\n", + " messages.append(f\"{label}: configuration differs for {key}\")\n", + " for key_path in ((\"hostname\",), (\"platform\",), (\"cuda\", \"device_name\")):\n", + " left: Any = reference_environment\n", + " right: Any = document[\"environment\"]\n", + " for key in key_path:\n", + " left = left.get(key) if isinstance(left, dict) else None\n", + " right = right.get(key) if isinstance(right, dict) else None\n", + " if left != right:\n", + " messages.append(\n", + " f\"{label}: environment differs for {'.'.join(key_path)}\"\n", + " )\n", + "\n", + " route_implementations: dict[str, set[str]] = {mode: set() for mode in MODES}\n", + " for document in documents:\n", + " for molecule in document[\"molecules\"].values():\n", + " for mode in MODES:\n", + " implementation = (\n", + " molecule[\"modes\"][mode].get(\"route\", {}).get(\"implementation\")\n", + " )\n", + " if implementation:\n", + " route_implementations[mode].add(implementation)\n", + " for mode, implementations in route_implementations.items():\n", + " if len(implementations) > 1:\n", + " messages.append(\n", + " f\"{mode}: routing implementations differ: {sorted(implementations)}\"\n", + " )\n", + " return messages\n", + "\n", + "\n", + "def load_result_documents(paths: list[Path]) -> list[dict[str, Any]]:\n", + " documents = [json.loads(path.read_text(encoding=\"utf-8\")) for path in paths]\n", + " messages = comparison_warnings(documents)\n", + " if messages:\n", + " print(\"Comparison warnings:\")\n", + " for message in messages:\n", + " print(f\" WARNING: {message}\")\n", + " return documents\n", + "\n", + "\n", + "def print_status_table(documents: list[dict[str, Any]]) -> None:\n", + " header = f\"{'label':10s} {'formula':9s} {'AOs':>5s} {'mode':10s} {'runtime':12s} {'memory':12s}\"\n", + " print(header)\n", + " print(\"-\" * len(header))\n", + " for document in documents:\n", + " for molecule in document[\"molecules\"].values():\n", + " observed = molecule.get(\"observed\") or {}\n", + " aos = observed.get(\"actual_aos\", molecule[\"expected_aos\"])\n", + " for mode in MODES:\n", + " mode_record = molecule[\"modes\"][mode]\n", + " runtime_status = mode_record.get(\"runtime\", {}).get(\"status\", \"pending\")\n", + " memory_status = mode_record.get(\"memory\", {}).get(\"status\", \"pending\")\n", + " print(\n", + " f\"{document['run_label'][:10]:10s} {molecule['formula']:9s} {aos:5d} \"\n", + " f\"{mode:10s} {runtime_status:12s} {memory_status:12s}\"\n", + " )\n", + "\n", + "\n", + "SELECTED_DOCUMENTS = (\n", + " load_result_documents(SELECTED_RESULT_FILES) if SELECTED_RESULT_FILES else []\n", + ")\n", + "if SELECTED_DOCUMENTS:\n", + " print_status_table(SELECTED_DOCUMENTS)\n", + "else:\n", + " print(\"No benchmark result files were found. Run the benchmark script first.\")" + ] + }, + { + "cell_type": "markdown", + "id": "57fda27a", + "metadata": {}, + "source": [ + "## Visualize\n", + "\n", + "Each metric is rendered in its own notebook output with a single y-axis. Runtime curves show the median of the recorded samples with error bars spanning the observed minimum and maximum; one-sample legacy results therefore have zero-width bounds. Successful observations are plotted against actual spherical AO counts. Failed or timed-out points stay absent from curves and remain visible in the status table." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "caf55d9d", + "metadata": {}, + "outputs": [], + "source": [ + "from matplotlib.axes import Axes\n", + "\n", + "MODE_COLORS = {\n", + " \"cpu\": \"#006D77\",\n", + " \"cpu_dense\": \"#83C5BE\",\n", + " \"gpu\": \"#C44536\",\n", + "}\n", + "REVISION_LINESTYLES = (\"-\", \"--\", \":\", \"-.\")\n", + "RESULT_MARKERS = (\"o\", \"s\", \"^\", \"D\", \"v\", \"P\", \"X\")\n", + "\n", + "\n", + "def measurement_samples(mode_record: dict[str, Any], measurement: str) -> list[float]:\n", + " result = mode_record.get(measurement, {})\n", + " if result.get(\"status\") != \"ok\":\n", + " return []\n", + " if measurement == \"runtime\":\n", + " return [float(value) for value in result.get(\"runtime_samples_seconds\", [])]\n", + " peak_bytes = result.get(\"incremental_peak_bytes\")\n", + " return [float(peak_bytes) / 1024**3] if peak_bytes is not None else []\n", + "\n", + "\n", + "def sample_summary(samples: list[float]) -> tuple[float, float, float] | None:\n", + " if not samples:\n", + " return None\n", + " values = np.asarray(samples, dtype=float)\n", + " center = float(np.median(values))\n", + " return center, center - float(values.min()), float(values.max()) - center\n", + "\n", + "\n", + "def measurement_value(mode_record: dict[str, Any], measurement: str) -> float | None:\n", + " summary = sample_summary(measurement_samples(mode_record, measurement))\n", + " return summary[0] if summary is not None else None\n", + "\n", + "\n", + "def measurement_series(\n", + " document: dict[str, Any], mode: str, measurement: str\n", + ") -> tuple[list[int], list[float], list[float], list[float]]:\n", + " points: list[tuple[int, float, float, float]] = []\n", + " for molecule in document[\"molecules\"].values():\n", + " observed = molecule.get(\"observed\") or {}\n", + " aos = int(observed.get(\"actual_aos\", molecule[\"expected_aos\"]))\n", + " summary = sample_summary(\n", + " measurement_samples(molecule[\"modes\"][mode], measurement)\n", + " )\n", + " if summary is not None and summary[0] > 0.0:\n", + " points.append((aos, *summary))\n", + " points.sort()\n", + " return (\n", + " [point[0] for point in points],\n", + " [point[1] for point in points],\n", + " [point[2] for point in points],\n", + " [point[3] for point in points],\n", + " )\n", + "\n", + "\n", + "def cpu_reference_ratio_series(\n", + " document: dict[str, Any], measurement: str\n", + ") -> tuple[list[int], list[float], list[float], list[float]]:\n", + " points: list[tuple[int, float, float, float]] = []\n", + " for molecule in document[\"molecules\"].values():\n", + " observed = molecule.get(\"observed\") or {}\n", + " aos = int(observed.get(\"actual_aos\", molecule[\"expected_aos\"]))\n", + " production_samples = measurement_samples(molecule[\"modes\"][\"cpu\"], measurement)\n", + " dense_samples = measurement_samples(molecule[\"modes\"][\"cpu_dense\"], measurement)\n", + " production_summary = sample_summary(production_samples)\n", + " dense_summary = sample_summary(dense_samples)\n", + " if (\n", + " production_summary is None\n", + " or dense_summary is None\n", + " or min(production_samples) <= 0.0\n", + " ):\n", + " continue\n", + " center = dense_summary[0] / production_summary[0]\n", + " lower_bound = min(dense_samples) / max(production_samples)\n", + " upper_bound = max(dense_samples) / min(production_samples)\n", + " points.append((aos, center, center - lower_bound, upper_bound - center))\n", + " points.sort()\n", + " return (\n", + " [point[0] for point in points],\n", + " [point[1] for point in points],\n", + " [point[2] for point in points],\n", + " [point[3] for point in points],\n", + " )\n", + "\n", + "\n", + "def endpoint_label(label: str, x_values: list[int]) -> str:\n", + " return (\n", + " f\"{label} (last: {x_values[-1]} AOs)\"\n", + " if x_values\n", + " else f\"{label} (no successful points)\"\n", + " )\n", + "\n", + "\n", + "def style_benchmark_axis(\n", + " axis: Axes, *, title: str, ylabel: str, logarithmic: bool = False\n", + ") -> None:\n", + " axis.set(title=title, xlabel=\"Spherical AO count\", ylabel=ylabel)\n", + " if logarithmic:\n", + " axis.set_yscale(\"log\")\n", + " axis.grid(True, which=\"both\", color=\"#D9D9D9\", linewidth=0.6)\n", + " axis.legend(fontsize=8)\n", + "\n", + "\n", + "def plot_measurement(documents: list[dict[str, Any]], measurement: str) -> None:\n", + " if not documents:\n", + " print(f\"No result files selected; the {measurement} plot was not created.\")\n", + " return\n", + " _, axis = plt.subplots(figsize=(11, 6), constrained_layout=True)\n", + " for document_index, document in enumerate(documents):\n", + " label = document[\"run_label\"]\n", + " line_style = REVISION_LINESTYLES[document_index % len(REVISION_LINESTYLES)]\n", + " marker = RESULT_MARKERS[document_index % len(RESULT_MARKERS)]\n", + " for mode in MODES:\n", + " x_values, y_values, lower_errors, upper_errors = measurement_series(\n", + " document, mode, measurement\n", + " )\n", + " curve_label = f\"{label} {mode}\"\n", + " plot_arguments = {\n", + " \"color\": MODE_COLORS[mode],\n", + " \"linestyle\": line_style,\n", + " \"marker\": marker,\n", + " \"markersize\": 4,\n", + " \"label\": endpoint_label(curve_label, x_values),\n", + " }\n", + " if measurement == \"runtime\":\n", + " axis.errorbar(\n", + " x_values,\n", + " y_values,\n", + " yerr=np.asarray([lower_errors, upper_errors]),\n", + " capsize=3,\n", + " **plot_arguments,\n", + " )\n", + " else:\n", + " axis.plot(x_values, y_values, **plot_arguments)\n", + " if measurement == \"runtime\":\n", + " title = \"One XC/Vxc evaluation (median and min-max)\"\n", + " ylabel = \"Runtime (s)\"\n", + " elif measurement == \"memory\":\n", + " title = \"Incremental allocation peak\"\n", + " ylabel = \"Memory (GiB)\"\n", + " else:\n", + " raise ValueError(f\"Unknown measurement: {measurement}\")\n", + " style_benchmark_axis(axis, title=title, ylabel=ylabel, logarithmic=True)\n", + " plt.show()\n", + "\n", + "\n", + "def plot_cpu_reference_ratio(documents: list[dict[str, Any]], measurement: str) -> None:\n", + " if not documents:\n", + " print(\n", + " f\"No result files selected; the {measurement} ratio plot was not created.\"\n", + " )\n", + " return\n", + " _, axis = plt.subplots(figsize=(11, 6), constrained_layout=True)\n", + " for document_index, document in enumerate(documents):\n", + " x_values, y_values, lower_errors, upper_errors = cpu_reference_ratio_series(\n", + " document, measurement\n", + " )\n", + " line_style = REVISION_LINESTYLES[document_index % len(REVISION_LINESTYLES)]\n", + " marker = RESULT_MARKERS[document_index % len(RESULT_MARKERS)]\n", + " plot_arguments = {\n", + " \"color\": MODE_COLORS[\"cpu\"],\n", + " \"linestyle\": line_style,\n", + " \"marker\": marker,\n", + " \"markersize\": 4,\n", + " \"label\": f\"{document['run_label']} cpu\",\n", + " }\n", + " if measurement == \"runtime\":\n", + " axis.errorbar(\n", + " x_values,\n", + " y_values,\n", + " yerr=np.asarray([lower_errors, upper_errors]),\n", + " capsize=3,\n", + " **plot_arguments,\n", + " )\n", + " else:\n", + " axis.plot(x_values, y_values, **plot_arguments)\n", + " if measurement == \"runtime\":\n", + " title = \"CPU production speedup (median and min-max)\"\n", + " ylabel = \"CPU dense runtime / production runtime\"\n", + " elif measurement == \"memory\":\n", + " title = \"CPU production memory reduction\"\n", + " ylabel = \"CPU dense peak / production peak\"\n", + " else:\n", + " raise ValueError(f\"Unknown measurement: {measurement}\")\n", + " style_benchmark_axis(axis, title=title, ylabel=ylabel)\n", + " axis.axhline(1.0, color=\"#777777\", linewidth=0.8, linestyle=\":\")\n", + " plt.show()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "165b8f50", + "metadata": {}, + "outputs": [], + "source": [ + "plot_measurement(SELECTED_DOCUMENTS, \"runtime\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "b923074a", + "metadata": {}, + "outputs": [], + "source": [ + "plot_measurement(SELECTED_DOCUMENTS, \"memory\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ce3a086b", + "metadata": {}, + "outputs": [], + "source": [ + "plot_cpu_reference_ratio(SELECTED_DOCUMENTS, \"runtime\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "fded2dd7", + "metadata": {}, + "outputs": [], + "source": [ + "plot_cpu_reference_ratio(SELECTED_DOCUMENTS, \"memory\")" + ] + }, + { + "cell_type": "markdown", + "id": "58750d35", + "metadata": {}, + "source": [ + "## Numerical Differences\n", + "\n", + "Production CPU and GPU fingerprints are compared molecule-by-molecule with the CPU-dense reference. The summary reports maximum absolute and relative errors and counts values outside the existing acceptance tolerances. Each per-fingerprint plot shows the signed difference `production - cpu_dense`; the dotted zero line is the CPU-dense reference.\n", + "\n", + "The final six-panel figure compares CPU-dense references across result files. The first selected result is the baseline, and each curve shows `comparison cpu_dense - baseline cpu_dense`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "f4de2998", + "metadata": {}, + "outputs": [], + "source": [ + "FINGERPRINT_LABELS = {\n", + " \"electron_integral\": \"Electron integral\",\n", + " \"xc_energy\": \"XC energy\",\n", + " \"vxc_sum\": \"Vxc sum\",\n", + " \"vxc_trace\": \"Vxc trace\",\n", + " \"vxc_frobenius_norm\": \"Vxc Frobenius norm\",\n", + " \"vxc_max_abs\": \"Vxc max abs\",\n", + "}\n", + "ERROR_TOLERANCES = {\n", + " \"cpu\": (1e-8, 5e-8),\n", + " \"gpu\": (1e-7, 2e-7),\n", + "}\n", + "\n", + "\n", + "def fingerprint_error_records(\n", + " document: dict[str, Any], production_mode: str, fingerprint_key: str\n", + ") -> list[dict[str, Any]]:\n", + " rtol, atol = ERROR_TOLERANCES[production_mode]\n", + " records: list[dict[str, Any]] = []\n", + " for formula, molecule in document[\"molecules\"].items():\n", + " reference = preferred_fingerprint(molecule[\"modes\"][\"cpu_dense\"])\n", + " production = preferred_fingerprint(molecule[\"modes\"][production_mode])\n", + " if (\n", + " reference is None\n", + " or production is None\n", + " or fingerprint_key not in reference\n", + " or fingerprint_key not in production\n", + " ):\n", + " continue\n", + " reference_value = float(reference[fingerprint_key])\n", + " production_value = float(production[fingerprint_key])\n", + " difference = production_value - reference_value\n", + " absolute_error = abs(difference)\n", + " relative_error = absolute_error / max(\n", + " abs(reference_value), np.finfo(float).tiny\n", + " )\n", + " tolerance_scale = atol + rtol * abs(reference_value)\n", + " observed = molecule.get(\"observed\") or {}\n", + " records.append(\n", + " {\n", + " \"formula\": formula,\n", + " \"aos\": int(observed.get(\"actual_aos\", molecule[\"expected_aos\"])),\n", + " \"difference\": difference,\n", + " \"absolute_error\": absolute_error,\n", + " \"relative_error\": relative_error,\n", + " \"tolerance_ratio\": absolute_error / tolerance_scale,\n", + " }\n", + " )\n", + " records.sort(key=lambda record: record[\"aos\"])\n", + " return records\n", + "\n", + "\n", + "def dense_reference_difference_records(\n", + " reference_document: dict[str, Any],\n", + " comparison_document: dict[str, Any],\n", + " fingerprint_key: str,\n", + ") -> list[dict[str, Any]]:\n", + " records: list[dict[str, Any]] = []\n", + " for formula, reference_molecule in reference_document[\"molecules\"].items():\n", + " comparison_molecule = comparison_document[\"molecules\"].get(formula)\n", + " if comparison_molecule is None:\n", + " continue\n", + " reference = preferred_fingerprint(reference_molecule[\"modes\"][\"cpu_dense\"])\n", + " comparison = preferred_fingerprint(comparison_molecule[\"modes\"][\"cpu_dense\"])\n", + " if (\n", + " reference is None\n", + " or comparison is None\n", + " or fingerprint_key not in reference\n", + " or fingerprint_key not in comparison\n", + " ):\n", + " continue\n", + " observed = comparison_molecule.get(\"observed\") or {}\n", + " records.append(\n", + " {\n", + " \"formula\": formula,\n", + " \"aos\": int(\n", + " observed.get(\"actual_aos\", comparison_molecule[\"expected_aos\"])\n", + " ),\n", + " \"difference\": float(comparison[fingerprint_key])\n", + " - float(reference[fingerprint_key]),\n", + " }\n", + " )\n", + " records.sort(key=lambda record: record[\"aos\"])\n", + " return records\n", + "\n", + "\n", + "def print_fingerprint_error_summary(documents: list[dict[str, Any]]) -> None:\n", + " header = (\n", + " f\"{'label':10s} {'mode':4s} {'fingerprint':22s} \"\n", + " f\"{'max abs':>11s} {'max rel':>11s} {'outside':>8s} {'at':>8s}\"\n", + " )\n", + " print(header)\n", + " print(\"-\" * len(header))\n", + " for document in documents:\n", + " for production_mode in ERROR_TOLERANCES:\n", + " for fingerprint_key, fingerprint_label in FINGERPRINT_LABELS.items():\n", + " records = fingerprint_error_records(\n", + " document, production_mode, fingerprint_key\n", + " )\n", + " if not records:\n", + " continue\n", + " max_absolute_error = max(record[\"absolute_error\"] for record in records)\n", + " worst_relative = max(\n", + " records, key=lambda record: record[\"relative_error\"]\n", + " )\n", + " outside_tolerance = sum(\n", + " record[\"tolerance_ratio\"] > 1.0 for record in records\n", + " )\n", + " print(\n", + " f\"{document['run_label'][:10]:10s} {production_mode:4s} \"\n", + " f\"{fingerprint_label:22s} {max_absolute_error:11.3e} \"\n", + " f\"{worst_relative['relative_error']:11.3e} \"\n", + " f\"{outside_tolerance:3d}/{len(records):<4d} \"\n", + " f\"{worst_relative['formula']:>8s}\"\n", + " )\n", + "\n", + "\n", + "def plot_fingerprint_differences(\n", + " documents: list[dict[str, Any]], fingerprint_key: str\n", + ") -> None:\n", + " if fingerprint_key not in FINGERPRINT_LABELS:\n", + " raise ValueError(f\"Unknown fingerprint: {fingerprint_key}\")\n", + " if not documents:\n", + " print(f\"No result files selected; the {fingerprint_key} plot was not created.\")\n", + " return\n", + " _, axis = plt.subplots(figsize=(11, 6), constrained_layout=True)\n", + " for document_index, document in enumerate(documents):\n", + " line_style = REVISION_LINESTYLES[document_index % len(REVISION_LINESTYLES)]\n", + " marker = RESULT_MARKERS[document_index % len(RESULT_MARKERS)]\n", + " for production_mode in ERROR_TOLERANCES:\n", + " records = fingerprint_error_records(\n", + " document, production_mode, fingerprint_key\n", + " )\n", + " x_values = [record[\"aos\"] for record in records]\n", + " curve_label = f\"{document['run_label']} {production_mode}\"\n", + " axis.plot(\n", + " x_values,\n", + " [record[\"difference\"] for record in records],\n", + " color=MODE_COLORS[production_mode],\n", + " linestyle=line_style,\n", + " marker=marker,\n", + " markersize=4,\n", + " label=endpoint_label(curve_label, x_values),\n", + " )\n", + " fingerprint_label = FINGERPRINT_LABELS[fingerprint_key]\n", + " axis.axhline(\n", + " 0.0,\n", + " color=MODE_COLORS[\"cpu_dense\"],\n", + " linewidth=1.0,\n", + " linestyle=\":\",\n", + " label=\"CPU dense reference\",\n", + " )\n", + " style_benchmark_axis(\n", + " axis,\n", + " title=f\"{fingerprint_label} difference from CPU dense\",\n", + " ylabel=f\"{fingerprint_label} - CPU dense reference\",\n", + " )\n", + " axis.ticklabel_format(axis=\"y\", style=\"sci\", scilimits=(0, 0))\n", + " plt.show()\n", + "\n", + "\n", + "def plot_dense_reference_differences(documents: list[dict[str, Any]]) -> None:\n", + " if len(documents) < 2:\n", + " print(\"At least two result files are required to compare CPU-dense references.\")\n", + " return\n", + " reference_document = documents[0]\n", + " figure, axes = plt.subplots(2, 3, figsize=(16, 9), constrained_layout=True)\n", + " for axis, (fingerprint_key, fingerprint_label) in zip(\n", + " axes.flat, FINGERPRINT_LABELS.items(), strict=True\n", + " ):\n", + " plotted_differences: list[float] = []\n", + " for document_index, document in enumerate(documents[1:], start=1):\n", + " records = dense_reference_difference_records(\n", + " reference_document, document, fingerprint_key\n", + " )\n", + " x_values = [record[\"aos\"] for record in records]\n", + " differences = [record[\"difference\"] for record in records]\n", + " plotted_differences.extend(differences)\n", + " curve_label = f\"{document['run_label']} - {reference_document['run_label']}\"\n", + " axis.plot(\n", + " x_values,\n", + " differences,\n", + " color=MODE_COLORS[\"cpu_dense\"],\n", + " linestyle=REVISION_LINESTYLES[\n", + " document_index % len(REVISION_LINESTYLES)\n", + " ],\n", + " marker=RESULT_MARKERS[document_index % len(RESULT_MARKERS)],\n", + " markersize=4,\n", + " label=endpoint_label(curve_label, x_values),\n", + " )\n", + " axis.axhline(0.0, color=\"#777777\", linewidth=0.8, linestyle=\":\")\n", + " axis.set(\n", + " title=fingerprint_label,\n", + " xlabel=\"Spherical AO count\",\n", + " ylabel=\"CPU-dense difference\",\n", + " )\n", + " if plotted_differences and all(value == 0.0 for value in plotted_differences):\n", + " axis.set_ylim(-0.5, 0.5)\n", + " axis.set_yticks([0.0])\n", + " axis.text(\n", + " 0.5,\n", + " 0.54,\n", + " \"All matched differences are exactly zero\",\n", + " color=\"#555555\",\n", + " fontsize=8,\n", + " ha=\"center\",\n", + " transform=axis.transAxes,\n", + " )\n", + " else:\n", + " axis.ticklabel_format(axis=\"y\", style=\"sci\", scilimits=(0, 0))\n", + " axis.grid(True, which=\"both\", color=\"#D9D9D9\", linewidth=0.6)\n", + " axis.legend(fontsize=7)\n", + " figure.suptitle(\n", + " f\"CPU-dense reference differences from {reference_document['run_label']}\"\n", + " )\n", + " plt.show()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "db731370", + "metadata": {}, + "outputs": [], + "source": [ + "print_fingerprint_error_summary(SELECTED_DOCUMENTS)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "148309a1", + "metadata": {}, + "outputs": [], + "source": [ + "for fingerprint_key in FINGERPRINT_LABELS:\n", + " plot_fingerprint_differences(SELECTED_DOCUMENTS, fingerprint_key)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "16648b0a", + "metadata": {}, + "outputs": [], + "source": [ + "plot_dense_reference_differences(SELECTED_DOCUMENTS)" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/benchmarks/pyscf_ao_screening_rotation_comparison.ipynb b/benchmarks/pyscf_ao_screening_rotation_comparison.ipynb new file mode 100644 index 00000000..bf938326 --- /dev/null +++ b/benchmarks/pyscf_ao_screening_rotation_comparison.ipynb @@ -0,0 +1,1395 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "ac79c661", + "metadata": {}, + "outputs": [], + "source": [] + }, + { + "cell_type": "markdown", + "id": "c081221d", + "metadata": {}, + "source": [ + "# Skala AO Screening Rotation Comparison\n", + "\n", + "This notebook loads JSON written by `benchmarks/run_pyscf_ao_screening_rotation_benchmark.py`. It does not construct molecules, load Skala, or execute benchmark workloads.\n", + "\n", + "Generate a full result file before running the analysis:\n", + "\n", + "```bash\n", + "/home/jenswehner/micromamba/envs/skala_gpu_python/bin/python \\\n", + " benchmarks/run_pyscf_ao_screening_rotation_benchmark.py \\\n", + " --label screening\n", + "```\n", + "\n", + "The default run records runtime and incremental peak memory for 72 orientations in each of `gpu`, `cpu_dense`, and `cpu_screened`. Add `--smoke` for one orientation per mode or `--preflight-only` to validate geometry and dependencies without measurements.\n", + "\n", + "## 1. Import Analysis Libraries and Configure Paths\n", + "\n", + "Select one or more rotation result files. Runtime is reported in seconds and incremental peak memory in GiB; CPU dense is the numerical and ratio reference." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3bb528e2", + "metadata": {}, + "outputs": [], + "source": [ + "from __future__ import annotations\n", + "\n", + "import json\n", + "import statistics\n", + "from collections import Counter\n", + "from datetime import UTC, datetime\n", + "from pathlib import Path\n", + "from typing import Any\n", + "\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "import pandas as pd\n", + "import seaborn as sns\n", + "\n", + "MODES = (\"gpu\", \"cpu_dense\", \"cpu_screened\")\n", + "MEASUREMENTS = (\"runtime\", \"memory\")\n", + "REFERENCE_MODE = \"cpu_dense\"\n", + "EXPECTED_AOS = 879\n", + "EXPECTED_ROUTES = {\n", + " \"gpu\": \"global_ao_screening\",\n", + " \"cpu_dense\": \"dense\",\n", + " \"cpu_screened\": \"global_ao_screening\",\n", + "}\n", + "TERMINAL_STATUSES = {\n", + " \"ok\",\n", + " \"timeout\",\n", + " \"oom\",\n", + " \"error\",\n", + " \"skipped_after_resource_failure\",\n", + "}\n", + "FINGERPRINT_LABELS = {\n", + " \"electron_integral\": \"Electron integral\",\n", + " \"xc_energy\": \"XC energy\",\n", + " \"vxc_sum\": \"Vxc sum\",\n", + " \"vxc_trace\": \"Vxc trace\",\n", + " \"vxc_frobenius_norm\": \"Vxc Frobenius norm\",\n", + " \"vxc_max_abs\": \"Vxc max abs\",\n", + "}\n", + "ERROR_TOLERANCES = {\n", + " \"cpu_screened\": (5e-8, 1e-8),\n", + " \"gpu\": (2e-7, 1e-7),\n", + "}\n", + "MODE_COLORS = {\n", + " \"gpu\": \"#C44536\",\n", + " \"cpu_dense\": \"#83C5BE\",\n", + " \"cpu_screened\": \"#006D77\",\n", + "}\n", + "\n", + "\n", + "def find_repository_root(start: Path) -> Path:\n", + " for candidate in (start.resolve(), *start.resolve().parents):\n", + " if (candidate / \"pyproject.toml\").is_file() and (\n", + " candidate / \"benchmarks\"\n", + " ).is_dir():\n", + " return candidate\n", + " raise FileNotFoundError(f\"Could not find the Skala repository above {start}\")\n", + "\n", + "\n", + "REPOSITORY_ROOT = find_repository_root(Path.cwd())\n", + "RESULTS_DIR = REPOSITORY_ROOT / \"benchmarks\" / \"results\"\n", + "ARTIFACT_DIR = RESULTS_DIR / \"rotation_comparison\"\n", + "FIGURE_DIR = ARTIFACT_DIR / \"figures\"\n", + "TABLE_DIR = ARTIFACT_DIR / \"tables\"\n", + "COMPARISON_JSON = ARTIFACT_DIR / \"comparison.json\"\n", + "\n", + "# Replace this list with explicit paths to compare a subset of result files.\n", + "SELECTED_RESULT_FILES = sorted(\n", + " RESULTS_DIR.glob(\"skala-pyscf-ao-screening-rotations-*.json\")\n", + ")\n", + "\n", + "sns.set_theme(style=\"whitegrid\", context=\"notebook\")\n", + "print(f\"Selected {len(SELECTED_RESULT_FILES)} rotation result file(s)\")\n", + "for result_file in SELECTED_RESULT_FILES:\n", + " print(f\" {result_file.name}\")" + ] + }, + { + "cell_type": "markdown", + "id": "556eb2a0", + "metadata": {}, + "source": [ + "## 2. Load and Validate Benchmark JSON Files\n", + "\n", + "Each file is checked for the rotation schema, provenance, configuration, environment, timestamps, and runner hashes. Malformed and partially completed files remain visible in the validation table.\n", + "\n", + "## 3. Normalize Molecule and Execution-Mode Results\n", + "\n", + "The molecule is fixed at C7H16, so normalization produces one row per result file, orientation, and execution mode. Rows include geometry, AO and grid sizes, routes, statuses, measurements, allocator baselines, and both runtime and memory fingerprints.\n", + "\n", + "## 4. Validate Benchmark Completeness and Status\n", + "\n", + "A full run requires 72 orientations, three modes, and successful runtime and memory records. Smoke and partial runs are accepted but explicitly reported.\n", + "\n", + "## 5. Verify AO Counts, Routes, and Screening Thresholds\n", + "\n", + "Observed AO counts must remain 879. Dense CPU must report `dense`; screened CPU and GPU must report `global_ao_screening`, which is expected because 879 exceeds PySCF's switch size of 800." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "50a0ade9", + "metadata": {}, + "outputs": [], + "source": [ + "NORMALIZED_COLUMNS = [\n", + " \"run_label\",\n", + " \"commit\",\n", + " \"branch\",\n", + " \"dirty\",\n", + " \"orientation_key\",\n", + " \"orientation_index\",\n", + " \"azimuth_degrees\",\n", + " \"polar_degrees\",\n", + " \"formula\",\n", + " \"carbon_count\",\n", + " \"expected_aos\",\n", + " \"actual_aos\",\n", + " \"electron_count\",\n", + " \"grid_points\",\n", + " \"mode\",\n", + " \"selected_route\",\n", + " \"route_request\",\n", + " \"switch_size\",\n", + " \"implementation_sha256\",\n", + " \"runtime_status\",\n", + " \"memory_status\",\n", + " \"runtime_samples_seconds\",\n", + " \"runtime_seconds\",\n", + " \"incremental_peak_bytes\",\n", + " \"incremental_peak_gib\",\n", + " \"allocator_baseline_bytes\",\n", + " \"allocator_baseline_gib\",\n", + "] + [\n", + " f\"{measurement}_{fingerprint}\"\n", + " for measurement in MEASUREMENTS\n", + " for fingerprint in FINGERPRINT_LABELS\n", + "]\n", + "\n", + "\n", + "def validate_document(path: Path, document: dict[str, Any]) -> list[str]:\n", + " failures: list[str] = []\n", + " required_fields = {\n", + " \"configuration\",\n", + " \"created_at\",\n", + " \"environment\",\n", + " \"orientations\",\n", + " \"runner_hashes\",\n", + " \"schema_version\",\n", + " \"source\",\n", + " \"updated_at\",\n", + " }\n", + " missing_fields = sorted(required_fields - document.keys())\n", + " if missing_fields:\n", + " failures.append(f\"missing top-level fields: {missing_fields}\")\n", + " if document.get(\"schema_version\") != 1:\n", + " failures.append(f\"unsupported schema version {document.get('schema_version')}\")\n", + " if document.get(\"benchmark\") != \"pyscf_ao_screening_rotations\":\n", + " failures.append(f\"unexpected benchmark marker {document.get('benchmark')!r}\")\n", + "\n", + " configuration = document.get(\"configuration\", {})\n", + " configured_modes = tuple(configuration.get(\"modes\", ()))\n", + " configured_measurements = tuple(configuration.get(\"measurements\", ()))\n", + " if configured_modes and configured_modes != MODES:\n", + " failures.append(f\"configured modes are {configured_modes}, expected {MODES}\")\n", + " if configured_measurements and configured_measurements != MEASUREMENTS:\n", + " failures.append(\n", + " f\"configured measurements are {configured_measurements}, expected {MEASUREMENTS}\"\n", + " )\n", + "\n", + " orientations = document.get(\"orientations\", {})\n", + " expected_count = int(configuration.get(\"orientation_count\", len(orientations)))\n", + " if len(orientations) != expected_count:\n", + " failures.append(\n", + " f\"contains {len(orientations)} orientations, configuration requests {expected_count}\"\n", + " )\n", + " if not configuration.get(\"smoke_run\", False) and len(orientations) != 72:\n", + " failures.append(\n", + " f\"full run contains {len(orientations)} orientations, expected 72\"\n", + " )\n", + " coordinate_hashes = [\n", + " orientation.get(\"coordinate_sha256\") for orientation in orientations.values()\n", + " ]\n", + " if len(set(coordinate_hashes)) != len(coordinate_hashes):\n", + " failures.append(\"orientation coordinate hashes are not unique\")\n", + "\n", + " for orientation_key, orientation in orientations.items():\n", + " observed = orientation.get(\"observed\") or {}\n", + " actual_aos = observed.get(\"actual_aos\")\n", + " if actual_aos is not None and int(actual_aos) != EXPECTED_AOS:\n", + " failures.append(\n", + " f\"{orientation_key}: observed {actual_aos} AOs, expected {EXPECTED_AOS}\"\n", + " )\n", + " modes = orientation.get(\"modes\", {})\n", + " missing_modes = sorted(set(MODES) - modes.keys())\n", + " if missing_modes:\n", + " failures.append(f\"{orientation_key}: missing modes {missing_modes}\")\n", + " for mode in MODES:\n", + " mode_record = modes.get(mode, {})\n", + " route = mode_record.get(\"route\", {})\n", + " selected_route = route.get(\"selected_route\")\n", + " if selected_route is not None and selected_route != EXPECTED_ROUTES[mode]:\n", + " failures.append(\n", + " f\"{orientation_key} {mode}: selected {selected_route}, \"\n", + " f\"expected {EXPECTED_ROUTES[mode]}\"\n", + " )\n", + " switch_size = route.get(\"pyscf_switch_size\")\n", + " if switch_size is not None and EXPECTED_AOS <= int(switch_size):\n", + " failures.append(\n", + " f\"{orientation_key} {mode}: {EXPECTED_AOS} AOs do not exceed \"\n", + " f\"reported switch size {switch_size}\"\n", + " )\n", + " for measurement in MEASUREMENTS:\n", + " result = mode_record.get(measurement)\n", + " if result is None:\n", + " failures.append(f\"{orientation_key} {mode}: missing {measurement}\")\n", + " continue\n", + " status = result.get(\"status\")\n", + " if status not in TERMINAL_STATUSES:\n", + " failures.append(\n", + " f\"{orientation_key} {mode} {measurement}: unknown status {status!r}\"\n", + " )\n", + " elif status != \"ok\":\n", + " failures.append(\n", + " f\"{orientation_key} {mode} {measurement}: status {status}\"\n", + " )\n", + " if (\n", + " measurement == \"runtime\"\n", + " and status == \"ok\"\n", + " and not result.get(\"runtime_samples_seconds\")\n", + " ):\n", + " failures.append(\n", + " f\"{orientation_key} {mode}: successful runtime has no samples\"\n", + " )\n", + " return failures\n", + "\n", + "\n", + "def normalize_document(path: Path, document: dict[str, Any]) -> list[dict[str, Any]]:\n", + " base_molecule = document.get(\"geometry\", {}).get(\"base_molecule\", {})\n", + " source = document.get(\"source\", {})\n", + " run_label = str(document.get(\"run_label\") or path.stem)\n", + " rows: list[dict[str, Any]] = []\n", + " for orientation_key, orientation in document.get(\"orientations\", {}).items():\n", + " observed = orientation.get(\"observed\") or {}\n", + " for mode in MODES:\n", + " mode_record = orientation.get(\"modes\", {}).get(mode, {})\n", + " runtime = mode_record.get(\"runtime\", {})\n", + " memory = mode_record.get(\"memory\", {})\n", + " route = mode_record.get(\"route\", {})\n", + " runtime_samples = [\n", + " float(value) for value in runtime.get(\"runtime_samples_seconds\", [])\n", + " ]\n", + " row: dict[str, Any] = {\n", + " \"run_label\": run_label,\n", + " \"commit\": source.get(\"commit\"),\n", + " \"branch\": source.get(\"branch\"),\n", + " \"dirty\": source.get(\"dirty\"),\n", + " \"orientation_key\": orientation_key,\n", + " \"orientation_index\": orientation.get(\"index\"),\n", + " \"azimuth_degrees\": orientation.get(\"azimuth_degrees\"),\n", + " \"polar_degrees\": orientation.get(\"polar_degrees\"),\n", + " \"formula\": base_molecule.get(\"formula\", observed.get(\"formula\")),\n", + " \"carbon_count\": base_molecule.get(\n", + " \"carbon_count\", observed.get(\"carbon_count\")\n", + " ),\n", + " \"expected_aos\": base_molecule.get(\"expected_aos\", EXPECTED_AOS),\n", + " \"actual_aos\": observed.get(\"actual_aos\"),\n", + " \"electron_count\": observed.get(\"electron_count\"),\n", + " \"grid_points\": observed.get(\"grid_points\"),\n", + " \"mode\": mode,\n", + " \"selected_route\": route.get(\"selected_route\"),\n", + " \"route_request\": route.get(\"request\"),\n", + " \"switch_size\": route.get(\"pyscf_switch_size\"),\n", + " \"implementation_sha256\": route.get(\"implementation_sha256\"),\n", + " \"runtime_status\": runtime.get(\"status\", \"missing\"),\n", + " \"memory_status\": memory.get(\"status\", \"missing\"),\n", + " \"runtime_samples_seconds\": runtime_samples,\n", + " \"runtime_seconds\": (\n", + " statistics.median(runtime_samples) if runtime_samples else np.nan\n", + " ),\n", + " \"incremental_peak_bytes\": memory.get(\"incremental_peak_bytes\"),\n", + " \"incremental_peak_gib\": (\n", + " float(memory[\"incremental_peak_bytes\"]) / 1024**3\n", + " if memory.get(\"incremental_peak_bytes\") is not None\n", + " else np.nan\n", + " ),\n", + " \"allocator_baseline_bytes\": memory.get(\"allocator_baseline_bytes\"),\n", + " \"allocator_baseline_gib\": (\n", + " float(memory[\"allocator_baseline_bytes\"]) / 1024**3\n", + " if memory.get(\"allocator_baseline_bytes\") is not None\n", + " else np.nan\n", + " ),\n", + " }\n", + " for measurement, result in ((\"runtime\", runtime), (\"memory\", memory)):\n", + " fingerprint = result.get(\"fingerprint\", {})\n", + " for fingerprint_key in FINGERPRINT_LABELS:\n", + " row[f\"{measurement}_{fingerprint_key}\"] = fingerprint.get(\n", + " fingerprint_key, np.nan\n", + " )\n", + " rows.append(row)\n", + " return rows\n", + "\n", + "\n", + "DOCUMENTS: list[tuple[Path, dict[str, Any]]] = []\n", + "validation_rows: list[dict[str, str]] = []\n", + "normalized_rows: list[dict[str, Any]] = []\n", + "for result_file in SELECTED_RESULT_FILES:\n", + " try:\n", + " document = json.loads(result_file.read_text(encoding=\"utf-8\"))\n", + " except (OSError, json.JSONDecodeError) as error:\n", + " validation_rows.append(\n", + " {\"run_label\": result_file.stem, \"failure\": f\"could not load: {error}\"}\n", + " )\n", + " continue\n", + " DOCUMENTS.append((result_file, document))\n", + " normalized_rows.extend(normalize_document(result_file, document))\n", + " failures = validate_document(result_file, document)\n", + " run_label = str(document.get(\"run_label\") or result_file.stem)\n", + " validation_rows.extend(\n", + " {\"run_label\": run_label, \"failure\": failure} for failure in failures\n", + " )\n", + "\n", + "run_labels = [\n", + " str(document.get(\"run_label\") or path.stem) for path, document in DOCUMENTS\n", + "]\n", + "duplicate_run_labels = sorted(\n", + " label for label, count in Counter(run_labels).items() if count > 1\n", + ")\n", + "if duplicate_run_labels:\n", + " raise ValueError(f\"Run labels must be unique: {duplicate_run_labels}\")\n", + "\n", + "normalized_df = pd.DataFrame(normalized_rows, columns=NORMALIZED_COLUMNS)\n", + "validation_df = pd.DataFrame(validation_rows, columns=[\"run_label\", \"failure\"])\n", + "if DOCUMENTS:\n", + " print(\n", + " f\"Loaded {len(DOCUMENTS)} document(s) and {len(normalized_df)} normalized rows\"\n", + " )\n", + "else:\n", + " print(\"No rotation benchmark JSON files found. Run the benchmark first.\")\n", + "display(\n", + " validation_df\n", + " if not validation_df.empty\n", + " else pd.DataFrame({\"validation\": [\"passed\"]})\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "c20ff5e3", + "metadata": {}, + "outputs": [], + "source": [ + "if normalized_df.empty:\n", + " status_summary_df = pd.DataFrame()\n", + " route_summary_df = pd.DataFrame()\n", + "else:\n", + " status_rows: list[dict[str, Any]] = []\n", + " for status_column in (\"runtime_status\", \"memory_status\"):\n", + " measurement = status_column.removesuffix(\"_status\")\n", + " grouped = normalized_df.groupby(\n", + " [\"run_label\", \"mode\", status_column], dropna=False\n", + " ).size()\n", + " for (run_label, mode, status), count in grouped.items():\n", + " status_rows.append(\n", + " {\n", + " \"run_label\": run_label,\n", + " \"mode\": mode,\n", + " \"measurement\": measurement,\n", + " \"status\": status,\n", + " \"count\": int(count),\n", + " }\n", + " )\n", + " status_summary_df = pd.DataFrame(status_rows)\n", + " route_summary_df = (\n", + " normalized_df.groupby([\"run_label\", \"mode\", \"selected_route\"], dropna=False)\n", + " .size()\n", + " .rename(\"orientation_count\")\n", + " .reset_index()\n", + " )\n", + "\n", + "print(\"Measurement status counts\")\n", + "display(status_summary_df)\n", + "print(\"Selected route counts\")\n", + "display(route_summary_df)" + ] + }, + { + "cell_type": "markdown", + "id": "b0a59c29", + "metadata": {}, + "source": [ + "## 6. Compare Numerical Fingerprints\n", + "\n", + "For every orientation, the runtime and memory workers are compared within each mode. The runtime fingerprints for GPU and screened CPU are also compared with CPU dense at the same orientation. The table reports signed, absolute, and relative differences for all six recorded quantities and applies configurable mode-specific tolerances." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3d5affd6", + "metadata": {}, + "outputs": [], + "source": [ + "NUMERICAL_COLUMNS = [\n", + " \"run_label\",\n", + " \"orientation_key\",\n", + " \"orientation_index\",\n", + " \"azimuth_degrees\",\n", + " \"polar_degrees\",\n", + " \"comparison\",\n", + " \"mode\",\n", + " \"fingerprint\",\n", + " \"reference_value\",\n", + " \"comparison_value\",\n", + " \"difference\",\n", + " \"absolute_error\",\n", + " \"relative_error\",\n", + " \"tolerance\",\n", + " \"within_tolerance\",\n", + "]\n", + "\n", + "\n", + "def numerical_difference(\n", + " row: pd.Series,\n", + " *,\n", + " comparison: str,\n", + " mode: str,\n", + " fingerprint: str,\n", + " reference_value: float,\n", + " comparison_value: float,\n", + " rtol: float,\n", + " atol: float,\n", + ") -> dict[str, Any] | None:\n", + " if not np.isfinite(reference_value) or not np.isfinite(comparison_value):\n", + " return None\n", + " difference = float(comparison_value - reference_value)\n", + " absolute_error = abs(difference)\n", + " relative_error = absolute_error / max(\n", + " abs(float(reference_value)), np.finfo(float).tiny\n", + " )\n", + " tolerance = atol + rtol * abs(float(reference_value))\n", + " return {\n", + " \"run_label\": row[\"run_label\"],\n", + " \"orientation_key\": row[\"orientation_key\"],\n", + " \"orientation_index\": row[\"orientation_index\"],\n", + " \"azimuth_degrees\": row[\"azimuth_degrees\"],\n", + " \"polar_degrees\": row[\"polar_degrees\"],\n", + " \"comparison\": comparison,\n", + " \"mode\": mode,\n", + " \"fingerprint\": fingerprint,\n", + " \"reference_value\": float(reference_value),\n", + " \"comparison_value\": float(comparison_value),\n", + " \"difference\": difference,\n", + " \"absolute_error\": absolute_error,\n", + " \"relative_error\": relative_error,\n", + " \"tolerance\": tolerance,\n", + " \"within_tolerance\": absolute_error <= tolerance,\n", + " }\n", + "\n", + "\n", + "numerical_rows: list[dict[str, Any]] = []\n", + "for _, row in normalized_df.iterrows():\n", + " mode = str(row[\"mode\"])\n", + " rtol, atol = ERROR_TOLERANCES.get(mode, (5e-8, 1e-8))\n", + " for fingerprint in FINGERPRINT_LABELS:\n", + " record = numerical_difference(\n", + " row,\n", + " comparison=\"runtime_vs_memory\",\n", + " mode=mode,\n", + " fingerprint=fingerprint,\n", + " reference_value=float(row[f\"runtime_{fingerprint}\"]),\n", + " comparison_value=float(row[f\"memory_{fingerprint}\"]),\n", + " rtol=rtol,\n", + " atol=atol,\n", + " )\n", + " if record is not None:\n", + " numerical_rows.append(record)\n", + "\n", + "if not normalized_df.empty:\n", + " indexed = normalized_df.set_index(\n", + " [\"run_label\", \"orientation_key\", \"mode\"], drop=False\n", + " )\n", + " for (run_label, orientation_key), _ in normalized_df.groupby(\n", + " [\"run_label\", \"orientation_key\"]\n", + " ):\n", + " dense_key = (run_label, orientation_key, REFERENCE_MODE)\n", + " if dense_key not in indexed.index:\n", + " continue\n", + " dense_row = indexed.loc[dense_key]\n", + " for mode in (\"cpu_screened\", \"gpu\"):\n", + " production_key = (run_label, orientation_key, mode)\n", + " if production_key not in indexed.index:\n", + " continue\n", + " production_row = indexed.loc[production_key]\n", + " rtol, atol = ERROR_TOLERANCES[mode]\n", + " for fingerprint in FINGERPRINT_LABELS:\n", + " record = numerical_difference(\n", + " production_row,\n", + " comparison=\"mode_vs_cpu_dense\",\n", + " mode=mode,\n", + " fingerprint=fingerprint,\n", + " reference_value=float(dense_row[f\"runtime_{fingerprint}\"]),\n", + " comparison_value=float(production_row[f\"runtime_{fingerprint}\"]),\n", + " rtol=rtol,\n", + " atol=atol,\n", + " )\n", + " if record is not None:\n", + " numerical_rows.append(record)\n", + "\n", + "numerical_df = pd.DataFrame(numerical_rows, columns=NUMERICAL_COLUMNS)\n", + "if numerical_df.empty:\n", + " numerical_summary_df = pd.DataFrame()\n", + "else:\n", + " numerical_summary_df = (\n", + " numerical_df.groupby([\"run_label\", \"comparison\", \"mode\", \"fingerprint\"])\n", + " .agg(\n", + " compared=(\"absolute_error\", \"size\"),\n", + " max_absolute_error=(\"absolute_error\", \"max\"),\n", + " max_relative_error=(\"relative_error\", \"max\"),\n", + " outside_tolerance=(\"within_tolerance\", lambda values: int((~values).sum())),\n", + " )\n", + " .reset_index()\n", + " )\n", + "display(numerical_summary_df)" + ] + }, + { + "cell_type": "markdown", + "id": "1c54906d", + "metadata": {}, + "source": [ + "## 7. Calculate Runtime Metrics and Speedups\n", + "\n", + "Runtime summaries use the median sample as the representative value. The comparison table includes dense-to-screened, dense-to-GPU, screened-CPU-to-GPU, and cross-run speedups at matched orientations.\n", + "\n", + "## 8. Calculate Memory Metrics and Reductions\n", + "\n", + "Incremental peaks and GPU allocator baselines are converted to GiB. Reduction factors use CPU dense as the numerator so values above one indicate improvement.\n", + "\n", + "## 9. Analyze Scaling with Molecular Size\n", + "\n", + "This benchmark intentionally fixes molecular size at C7H16 and 879 AOs, so molecular-size fitting is not meaningful. Instead, a harmonic least-squares model summarizes orientation sensitivity and records fit coefficients, $R^2$, and residual RMSE for runtime and memory." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "56dee578", + "metadata": {}, + "outputs": [], + "source": [ + "def summarize_metric(frame: pd.DataFrame, metric: str, value_name: str) -> pd.DataFrame:\n", + " rows: list[dict[str, Any]] = []\n", + " for (run_label, mode), group in frame.groupby([\"run_label\", \"mode\"]):\n", + " values = group[metric].dropna().astype(float).to_numpy()\n", + " if values.size == 0:\n", + " continue\n", + " mean_value = float(values.mean())\n", + " rows.append(\n", + " {\n", + " \"run_label\": run_label,\n", + " \"mode\": mode,\n", + " \"observation_count\": int(values.size),\n", + " \"minimum\": float(values.min()),\n", + " \"median\": float(np.median(values)),\n", + " \"mean\": mean_value,\n", + " \"standard_deviation\": (\n", + " float(values.std(ddof=1)) if values.size > 1 else 0.0\n", + " ),\n", + " \"coefficient_of_variation\": (\n", + " float(values.std(ddof=1) / mean_value)\n", + " if values.size > 1 and mean_value != 0.0\n", + " else 0.0\n", + " ),\n", + " \"representative\": float(np.median(values)),\n", + " \"unit\": value_name,\n", + " }\n", + " )\n", + " return pd.DataFrame(rows)\n", + "\n", + "\n", + "runtime_statistics_df = summarize_metric(normalized_df, \"runtime_seconds\", \"seconds\")\n", + "if not runtime_statistics_df.empty:\n", + " runtime_sample_counts = (\n", + " normalized_df.assign(\n", + " runtime_sample_count=normalized_df[\"runtime_samples_seconds\"].map(len)\n", + " )\n", + " .groupby([\"run_label\", \"mode\"])[\"runtime_sample_count\"]\n", + " .sum()\n", + " .reset_index()\n", + " )\n", + " runtime_statistics_df = runtime_statistics_df.merge(\n", + " runtime_sample_counts,\n", + " on=[\"run_label\", \"mode\"],\n", + " how=\"left\",\n", + " )\n", + "memory_statistics_df = summarize_metric(normalized_df, \"incremental_peak_gib\", \"GiB\")\n", + "\n", + "\n", + "def finite_ratio(numerator: Any, denominator: Any) -> float:\n", + " numerator_value = float(numerator)\n", + " denominator_value = float(denominator)\n", + " if (\n", + " np.isfinite(numerator_value)\n", + " and np.isfinite(denominator_value)\n", + " and denominator_value > 0.0\n", + " ):\n", + " return numerator_value / denominator_value\n", + " return np.nan\n", + "\n", + "\n", + "comparison_rows: list[dict[str, Any]] = []\n", + "for (run_label, orientation_key), group in normalized_df.groupby(\n", + " [\"run_label\", \"orientation_key\"]\n", + "):\n", + " by_mode = group.set_index(\"mode\")\n", + " if any(mode not in by_mode.index for mode in MODES):\n", + " continue\n", + " dense = by_mode.loc[\"cpu_dense\"]\n", + " screened = by_mode.loc[\"cpu_screened\"]\n", + " gpu = by_mode.loc[\"gpu\"]\n", + " comparison_rows.append(\n", + " {\n", + " \"run_label\": run_label,\n", + " \"orientation_key\": orientation_key,\n", + " \"orientation_index\": dense[\"orientation_index\"],\n", + " \"azimuth_degrees\": dense[\"azimuth_degrees\"],\n", + " \"polar_degrees\": dense[\"polar_degrees\"],\n", + " \"dense_to_screened_runtime_speedup\": finite_ratio(\n", + " dense[\"runtime_seconds\"], screened[\"runtime_seconds\"]\n", + " ),\n", + " \"dense_to_gpu_runtime_speedup\": finite_ratio(\n", + " dense[\"runtime_seconds\"], gpu[\"runtime_seconds\"]\n", + " ),\n", + " \"screened_cpu_to_gpu_runtime_speedup\": finite_ratio(\n", + " screened[\"runtime_seconds\"], gpu[\"runtime_seconds\"]\n", + " ),\n", + " \"dense_to_screened_memory_reduction\": finite_ratio(\n", + " dense[\"incremental_peak_gib\"], screened[\"incremental_peak_gib\"]\n", + " ),\n", + " \"dense_to_gpu_memory_reduction\": finite_ratio(\n", + " dense[\"incremental_peak_gib\"], gpu[\"incremental_peak_gib\"]\n", + " ),\n", + " \"screened_cpu_to_gpu_memory_ratio\": finite_ratio(\n", + " screened[\"incremental_peak_gib\"], gpu[\"incremental_peak_gib\"]\n", + " ),\n", + " }\n", + " )\n", + "comparison_metrics_df = pd.DataFrame(comparison_rows)\n", + "\n", + "cross_run_rows: list[dict[str, Any]] = []\n", + "if len(DOCUMENTS) > 1 and not normalized_df.empty:\n", + " baseline_path, baseline_document = DOCUMENTS[0]\n", + " baseline_run_label = str(baseline_document.get(\"run_label\") or baseline_path.stem)\n", + " baseline = normalized_df[\n", + " normalized_df[\"run_label\"] == baseline_run_label\n", + " ].set_index([\"orientation_key\", \"mode\"])\n", + " for current_path, document in DOCUMENTS[1:]:\n", + " current_run_label = str(document.get(\"run_label\") or current_path.stem)\n", + " current = normalized_df[\n", + " normalized_df[\"run_label\"] == current_run_label\n", + " ].set_index([\"orientation_key\", \"mode\"])\n", + " for key in baseline.index.intersection(current.index):\n", + " baseline_row = baseline.loc[key]\n", + " current_row = current.loc[key]\n", + " cross_run_rows.append(\n", + " {\n", + " \"baseline_run_label\": baseline_run_label,\n", + " \"comparison_run_label\": current_run_label,\n", + " \"orientation_key\": key[0],\n", + " \"mode\": key[1],\n", + " \"runtime_speedup\": finite_ratio(\n", + " baseline_row[\"runtime_seconds\"],\n", + " current_row[\"runtime_seconds\"],\n", + " ),\n", + " \"memory_reduction\": finite_ratio(\n", + " baseline_row[\"incremental_peak_gib\"],\n", + " current_row[\"incremental_peak_gib\"],\n", + " ),\n", + " }\n", + " )\n", + "cross_run_df = pd.DataFrame(cross_run_rows)\n", + "\n", + "orientation_fit_rows: list[dict[str, Any]] = []\n", + "for (run_label, mode), group in normalized_df.groupby([\"run_label\", \"mode\"]):\n", + " for metric in (\"runtime_seconds\", \"incremental_peak_gib\"):\n", + " fit_data = group.dropna(subset=[\"azimuth_degrees\", \"polar_degrees\", metric])\n", + " if len(fit_data) < 5:\n", + " continue\n", + " azimuth = np.radians(fit_data[\"azimuth_degrees\"].to_numpy(float))\n", + " polar = np.radians(fit_data[\"polar_degrees\"].to_numpy(float))\n", + " design = np.column_stack(\n", + " [\n", + " np.ones(len(fit_data)),\n", + " np.sin(azimuth),\n", + " np.cos(azimuth),\n", + " np.sin(polar),\n", + " np.cos(polar),\n", + " ]\n", + " )\n", + " values = fit_data[metric].to_numpy(float)\n", + " coefficients, _, _, _ = np.linalg.lstsq(design, values, rcond=None)\n", + " residuals = values - design @ coefficients\n", + " total_variation = float(np.square(values - values.mean()).sum())\n", + " residual_variation = float(np.square(residuals).sum())\n", + " orientation_fit_rows.append(\n", + " {\n", + " \"run_label\": run_label,\n", + " \"mode\": mode,\n", + " \"metric\": metric,\n", + " \"intercept\": float(coefficients[0]),\n", + " \"sin_azimuth\": float(coefficients[1]),\n", + " \"cos_azimuth\": float(coefficients[2]),\n", + " \"sin_polar\": float(coefficients[3]),\n", + " \"cos_polar\": float(coefficients[4]),\n", + " \"r_squared\": (\n", + " 1.0 - residual_variation / total_variation\n", + " if total_variation > 0.0\n", + " else 1.0\n", + " ),\n", + " \"residual_rmse\": float(np.sqrt(np.mean(np.square(residuals)))),\n", + " }\n", + " )\n", + "orientation_fit_df = pd.DataFrame(orientation_fit_rows)\n", + "\n", + "print(\"Runtime statistics\")\n", + "display(runtime_statistics_df)\n", + "print(\"Memory statistics\")\n", + "display(memory_statistics_df)\n", + "print(\"Per-orientation speedup and reduction metrics\")\n", + "display(comparison_metrics_df)\n", + "print(\"Cross-run changes\")\n", + "display(cross_run_df)\n", + "print(\"Orientation-sensitivity fits\")\n", + "display(orientation_fit_df)" + ] + }, + { + "cell_type": "markdown", + "id": "bf4b225f", + "metadata": {}, + "source": [ + "## 10. Visualize Runtime Comparisons\n", + "\n", + "Runtime is shown by orientation index, as mode/revision distributions, and as 12-by-6 azimuth/polar heatmaps. Since AO count is fixed, route labels replace a screening-threshold marker.\n", + "\n", + "## 11. Visualize Memory Comparisons\n", + "\n", + "Incremental peak memory uses the same views, with GPU allocator baseline retained in the normalized table and exported metadata.\n", + "\n", + "## 12. Visualize Speedup and Memory Reduction\n", + "\n", + "Dense-to-screened and dense-to-GPU factors are rendered on the same angular grid. Values above one indicate faster execution or lower peak memory than CPU dense.\n", + "\n", + "## 13. Visualize Numerical Differences\n", + "\n", + "Signed fingerprint differences use CPU dense as zero. Change `FINGERPRINT_TO_PLOT` to inspect any of the six recorded fingerprint quantities." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ea347d87", + "metadata": {}, + "outputs": [], + "source": [ + "import re\n", + "\n", + "from matplotlib.figure import Figure\n", + "\n", + "LINE_STYLES = (\"-\", \"--\", \"-.\", \":\")\n", + "MARKERS = (\"o\", \"s\", \"^\", \"D\")\n", + "PLOT_STYLES = tuple(zip(LINE_STYLES, MARKERS, strict=True))\n", + "if len(DOCUMENTS) > len(PLOT_STYLES):\n", + " raise ValueError(f\"At most {len(PLOT_STYLES)} distinct run labels can be plotted\")\n", + "LABEL_PLOT_STYLES = {\n", + " str(document.get(\"run_label\") or path.stem): PLOT_STYLES[index]\n", + " for index, (path, document) in enumerate(DOCUMENTS)\n", + "}\n", + "\n", + "\n", + "def safe_filename(value: str) -> str:\n", + " return re.sub(r\"[^A-Za-z0-9_.-]+\", \"-\", value).strip(\"-.\") or \"result\"\n", + "\n", + "\n", + "def orientation_matrix(frame: pd.DataFrame, value_column: str) -> pd.DataFrame:\n", + " if frame.empty:\n", + " return pd.DataFrame()\n", + " return (\n", + " frame.pivot(\n", + " index=\"polar_degrees\",\n", + " columns=\"azimuth_degrees\",\n", + " values=value_column,\n", + " )\n", + " .sort_index()\n", + " .sort_index(axis=1)\n", + " )\n", + "\n", + "\n", + "def finite_value_range(values: Any) -> tuple[float, float] | None:\n", + " numeric = np.asarray(values, dtype=float)\n", + " finite = numeric[np.isfinite(numeric)]\n", + " if finite.size == 0:\n", + " return None\n", + " minimum = float(finite.min())\n", + " maximum = float(finite.max())\n", + " if minimum == maximum:\n", + " padding = max(abs(minimum) * 1e-9, np.finfo(float).eps)\n", + " return minimum - padding, maximum + padding\n", + " return minimum, maximum\n", + "\n", + "\n", + "def observed_range_errors(samples: Any, center: float) -> tuple[float, float]:\n", + " numeric = np.asarray(samples, dtype=float)\n", + " finite = numeric[np.isfinite(numeric)]\n", + " if finite.size == 0:\n", + " return 0.0, 0.0\n", + " return (\n", + " max(0.0, center - float(finite.min())),\n", + " max(0.0, float(finite.max()) - center),\n", + " )\n", + "\n", + "\n", + "def finish_figure(figure: Figure, output_path: Path | None = None) -> None:\n", + " if output_path is not None:\n", + " output_path.parent.mkdir(parents=True, exist_ok=True)\n", + " figure.savefig(output_path, dpi=180, bbox_inches=\"tight\")\n", + " plt.show()\n", + "\n", + "\n", + "def plot_orientation_lines(\n", + " value_column: str,\n", + " ylabel: str,\n", + " output_path: Path | None = None,\n", + " error_samples_column: str | None = None,\n", + ") -> None:\n", + " data = normalized_df.dropna(subset=[\"orientation_index\", value_column])\n", + " if data.empty:\n", + " print(f\"No successful values available for {value_column}\")\n", + " return\n", + " figure, axis = plt.subplots(figsize=(12, 6), constrained_layout=True)\n", + " for (run_label, mode), group in data.groupby([\"run_label\", \"mode\"]):\n", + " ordered = group.sort_values(\"orientation_index\")\n", + " x_values = ordered[\"orientation_index\"].to_numpy(float)\n", + " centers = ordered[value_column].to_numpy(float)\n", + " line_style, marker = LABEL_PLOT_STYLES[str(run_label)]\n", + " plot_options = {\n", + " \"color\": MODE_COLORS[mode],\n", + " \"linewidth\": 1.2,\n", + " \"alpha\": 0.85,\n", + " \"label\": f\"{run_label} {mode}\",\n", + " \"linestyle\": line_style,\n", + " \"marker\": marker,\n", + " \"markersize\": 3.5,\n", + " \"markevery\": max(1, len(ordered) // 12),\n", + " }\n", + " if error_samples_column is None:\n", + " axis.plot(x_values, centers, **plot_options)\n", + " else:\n", + " errors = np.asarray(\n", + " [\n", + " observed_range_errors(samples, center)\n", + " for samples, center in zip(\n", + " ordered[error_samples_column], centers, strict=True\n", + " )\n", + " ],\n", + " dtype=float,\n", + " ).T\n", + " axis.errorbar(\n", + " x_values,\n", + " centers,\n", + " yerr=errors,\n", + " capsize=2,\n", + " elinewidth=0.7,\n", + " **plot_options,\n", + " )\n", + " title = f\"{ylabel} across molecular orientations\"\n", + " if error_samples_column is not None:\n", + " title += \" (median and observed min-max)\"\n", + " axis.set(\n", + " title=title,\n", + " xlabel=\"Orientation index (azimuth-major, then polar)\",\n", + " ylabel=ylabel,\n", + " )\n", + " axis.grid(True, color=\"#D9D9D9\", linewidth=0.6)\n", + " axis.legend(fontsize=8, ncol=2)\n", + " finish_figure(figure, output_path)\n", + "\n", + "\n", + "def plot_mode_distribution(\n", + " value_column: str, ylabel: str, output_path: Path | None = None\n", + ") -> None:\n", + " data = normalized_df.dropna(subset=[value_column])\n", + " if data.empty:\n", + " print(f\"No successful values available for {value_column}\")\n", + " return\n", + " figure, axis = plt.subplots(figsize=(10, 6), constrained_layout=True)\n", + " sns.boxplot(\n", + " data=data,\n", + " x=\"mode\",\n", + " y=value_column,\n", + " hue=\"run_label\",\n", + " order=MODES,\n", + " showfliers=True,\n", + " ax=axis,\n", + " )\n", + " axis.set(title=f\"{ylabel} distribution by mode\", xlabel=\"Mode\", ylabel=ylabel)\n", + " axis.grid(True, axis=\"y\", color=\"#D9D9D9\", linewidth=0.6)\n", + " finish_figure(figure, output_path)\n", + "\n", + "\n", + "def plot_measurement_heatmaps(\n", + " value_column: str,\n", + " colorbar_label: str,\n", + " output_dir: Path | None = None,\n", + ") -> None:\n", + " if normalized_df[value_column].dropna().empty:\n", + " print(f\"No successful values available for {value_column}\")\n", + " return\n", + " for path, document in DOCUMENTS:\n", + " run_label = str(document.get(\"run_label\") or path.stem)\n", + " run_data = normalized_df[normalized_df[\"run_label\"] == run_label]\n", + " figure, axes = plt.subplots(\n", + " 1, len(MODES), figsize=(18, 4.8), constrained_layout=True\n", + " )\n", + " for axis, mode in zip(axes, MODES, strict=True):\n", + " matrix = orientation_matrix(\n", + " run_data[run_data[\"mode\"] == mode], value_column\n", + " )\n", + " value_range = finite_value_range(matrix)\n", + " if value_range is None:\n", + " axis.text(0.5, 0.5, \"No successful data\", ha=\"center\", va=\"center\")\n", + " axis.set_axis_off()\n", + " continue\n", + " minimum, maximum = value_range\n", + " sns.heatmap(\n", + " matrix,\n", + " mask=matrix.isna(),\n", + " cmap=\"viridis\",\n", + " vmin=minimum,\n", + " vmax=maximum,\n", + " cbar_kws={\"label\": colorbar_label},\n", + " ax=axis,\n", + " )\n", + " route_values = (\n", + " run_data.loc[run_data[\"mode\"] == mode, \"selected_route\"]\n", + " .dropna()\n", + " .unique()\n", + " )\n", + " route_label = \", \".join(str(value) for value in route_values) or \"no route\"\n", + " axis.set(\n", + " title=f\"{mode} ({route_label})\",\n", + " xlabel=\"Azimuth (degrees)\",\n", + " ylabel=\"Polar angle (degrees)\",\n", + " )\n", + " figure.suptitle(f\"{run_label} {colorbar_label} by orientation\")\n", + " output_path = (\n", + " output_dir\n", + " / f\"{safe_filename(run_label)}-{safe_filename(value_column)}-heatmap.png\"\n", + " if output_dir is not None\n", + " else None\n", + " )\n", + " finish_figure(figure, output_path)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "f41656f2", + "metadata": {}, + "outputs": [], + "source": [ + "RATIO_LABELS = {\n", + " \"dense_to_screened_runtime_speedup\": \"CPU dense / CPU screened runtime\",\n", + " \"dense_to_gpu_runtime_speedup\": \"CPU dense / GPU runtime\",\n", + " \"dense_to_screened_memory_reduction\": \"CPU dense / CPU screened peak\",\n", + " \"dense_to_gpu_memory_reduction\": \"CPU dense / GPU peak\",\n", + "}\n", + "\n", + "\n", + "def plot_ratio_heatmaps(\n", + " ratio_columns: tuple[str, ...],\n", + " title: str,\n", + " output_dir: Path | None = None,\n", + ") -> None:\n", + " if comparison_metrics_df.empty:\n", + " print(f\"No matched mode data available for {title}\")\n", + " return\n", + " for path, document in DOCUMENTS:\n", + " run_label = str(document.get(\"run_label\") or path.stem)\n", + " run_data = comparison_metrics_df[\n", + " comparison_metrics_df[\"run_label\"] == run_label\n", + " ]\n", + " figure, axes = plt.subplots(\n", + " 1, len(ratio_columns), figsize=(12, 4.8), constrained_layout=True\n", + " )\n", + " axes_array = np.atleast_1d(axes)\n", + " for axis, ratio_column in zip(axes_array, ratio_columns, strict=True):\n", + " matrix = orientation_matrix(run_data, ratio_column)\n", + " value_range = finite_value_range(matrix)\n", + " if value_range is None:\n", + " axis.text(0.5, 0.5, \"No successful data\", ha=\"center\", va=\"center\")\n", + " axis.set_axis_off()\n", + " continue\n", + " minimum, maximum = value_range\n", + " heatmap_options: dict[str, Any] = {\n", + " \"cmap\": \"RdYlGn\",\n", + " \"vmin\": minimum,\n", + " \"vmax\": maximum,\n", + " }\n", + " if minimum < 1.0 < maximum:\n", + " heatmap_options[\"center\"] = 1.0\n", + " sns.heatmap(\n", + " matrix,\n", + " mask=matrix.isna(),\n", + " cbar_kws={\"label\": \"Factor\"},\n", + " ax=axis,\n", + " **heatmap_options,\n", + " )\n", + " axis.set(\n", + " title=RATIO_LABELS[ratio_column],\n", + " xlabel=\"Azimuth (degrees)\",\n", + " ylabel=\"Polar angle (degrees)\",\n", + " )\n", + " figure.suptitle(f\"{run_label} {title}\")\n", + " output_path = (\n", + " output_dir / f\"{safe_filename(run_label)}-{safe_filename(title)}.png\"\n", + " if output_dir is not None\n", + " else None\n", + " )\n", + " finish_figure(figure, output_path)\n", + "\n", + "\n", + "def plot_fingerprint_heatmaps(fingerprint: str, output_dir: Path | None = None) -> None:\n", + " if fingerprint not in FINGERPRINT_LABELS:\n", + " raise ValueError(f\"Unknown fingerprint {fingerprint!r}\")\n", + " data = numerical_df[\n", + " (numerical_df[\"comparison\"] == \"mode_vs_cpu_dense\")\n", + " & (numerical_df[\"fingerprint\"] == fingerprint)\n", + " ]\n", + " if data.empty:\n", + " print(f\"No matched numerical data available for {fingerprint}\")\n", + " return\n", + " for run_label, run_data in data.groupby(\"run_label\"):\n", + " figure, axes = plt.subplots(1, 2, figsize=(12, 4.8), constrained_layout=True)\n", + " for axis, mode in zip(axes, (\"cpu_screened\", \"gpu\"), strict=True):\n", + " matrix = orientation_matrix(\n", + " run_data[run_data[\"mode\"] == mode], \"difference\"\n", + " )\n", + " value_range = finite_value_range(matrix)\n", + " if value_range is None:\n", + " axis.text(0.5, 0.5, \"No successful data\", ha=\"center\", va=\"center\")\n", + " axis.set_axis_off()\n", + " continue\n", + " minimum, maximum = value_range\n", + " heatmap_options = {\n", + " \"cmap\": \"coolwarm\",\n", + " \"vmin\": minimum,\n", + " \"vmax\": maximum,\n", + " }\n", + " if minimum < 0.0 < maximum:\n", + " heatmap_options[\"center\"] = 0.0\n", + " sns.heatmap(\n", + " matrix,\n", + " mask=matrix.isna(),\n", + " cbar_kws={\"label\": \"Mode - CPU dense\"},\n", + " ax=axis,\n", + " **heatmap_options,\n", + " )\n", + " axis.set(\n", + " title=mode,\n", + " xlabel=\"Azimuth (degrees)\",\n", + " ylabel=\"Polar angle (degrees)\",\n", + " )\n", + " figure.suptitle(f\"{run_label} {FINGERPRINT_LABELS[fingerprint]} difference\")\n", + " output_path = (\n", + " output_dir\n", + " / f\"{safe_filename(str(run_label))}-{safe_filename(fingerprint)}-difference.png\"\n", + " if output_dir is not None\n", + " else None\n", + " )\n", + " finish_figure(figure, output_path)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "942e3acc", + "metadata": {}, + "outputs": [], + "source": [ + "FINGERPRINT_TO_PLOT = \"xc_energy\"\n", + "\n", + "plot_orientation_lines(\n", + " \"runtime_seconds\",\n", + " \"Runtime (s)\",\n", + " error_samples_column=\"runtime_samples_seconds\",\n", + ")\n", + "plot_mode_distribution(\"runtime_seconds\", \"Runtime (s)\")\n", + "plot_measurement_heatmaps(\"runtime_seconds\", \"Runtime (s)\")\n", + "\n", + "plot_orientation_lines(\"incremental_peak_gib\", \"Incremental peak memory (GiB)\")\n", + "plot_mode_distribution(\"incremental_peak_gib\", \"Incremental peak memory (GiB)\")\n", + "plot_measurement_heatmaps(\"incremental_peak_gib\", \"Incremental peak memory (GiB)\")\n", + "\n", + "plot_ratio_heatmaps(\n", + " (\"dense_to_screened_runtime_speedup\", \"dense_to_gpu_runtime_speedup\"),\n", + " \"runtime speedup\",\n", + ")\n", + "plot_ratio_heatmaps(\n", + " (\"dense_to_screened_memory_reduction\", \"dense_to_gpu_memory_reduction\"),\n", + " \"memory reduction\",\n", + ")\n", + "plot_fingerprint_heatmaps(FINGERPRINT_TO_PLOT)" + ] + }, + { + "cell_type": "markdown", + "id": "be0bdf30", + "metadata": {}, + "source": [ + "## 14. Record Environment and Source Metadata\n", + "\n", + "This table keeps hardware, package versions, CUDA details, thread settings, scientific configuration, source commit, dirty state, implementation hash, and both runner hashes alongside every comparison.\n", + "\n", + "## 15. Export Comparison Results to JSON\n", + "\n", + "The comparison artifact contains JSON-safe configuration summaries, normalized rows, runtime and memory statistics, speedups, numerical differences, validation failures, angular fits, environment metadata, and source provenance.\n", + "\n", + "## 16. Save Tables and Figures\n", + "\n", + "The final cell writes CSV tables and deterministic PNG files below `benchmarks/results/rotation_comparison`. Re-running the cell refreshes the report artifacts from the currently selected input files." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "c530c585", + "metadata": {}, + "outputs": [], + "source": [ + "metadata_rows: list[dict[str, Any]] = []\n", + "validation_failure_counts = Counter(row[\"run_label\"] for row in validation_rows)\n", + "for path, document in DOCUMENTS:\n", + " environment = document.get(\"environment\", {})\n", + " packages = environment.get(\"packages\", {})\n", + " cuda = environment.get(\"cuda\", {})\n", + " configuration = document.get(\"configuration\", {})\n", + " source = document.get(\"source\", {})\n", + " hashes = document.get(\"runner_hashes\", {})\n", + " run_label = str(document.get(\"run_label\") or path.stem)\n", + " document_rows = normalized_df[normalized_df[\"run_label\"] == run_label]\n", + " implementation_hashes = sorted(\n", + " str(value) for value in document_rows[\"implementation_sha256\"].dropna().unique()\n", + " )\n", + " metadata_rows.append(\n", + " {\n", + " \"run_label\": run_label,\n", + " \"created_at\": document.get(\"created_at\"),\n", + " \"updated_at\": document.get(\"updated_at\"),\n", + " \"python\": environment.get(\"python\"),\n", + " \"python_executable\": environment.get(\"python_executable\"),\n", + " \"pyscf\": packages.get(\"pyscf\"),\n", + " \"skala\": packages.get(\"skala\"),\n", + " \"torch\": packages.get(\"torch\"),\n", + " \"cupy\": packages.get(\"cupy\"),\n", + " \"gpu4pyscf\": packages.get(\"gpu4pyscf\"),\n", + " \"memray\": packages.get(\"memray\"),\n", + " \"cuda_available\": cuda.get(\"available\"),\n", + " \"torch_cuda_version\": cuda.get(\"torch_cuda_version\"),\n", + " \"device_name\": cuda.get(\"device_name\"),\n", + " \"cpu_threads\": configuration.get(\"cpu_threads\"),\n", + " \"thread_environment\": environment.get(\"thread_environment\"),\n", + " \"basis\": configuration.get(\"basis\"),\n", + " \"functional\": configuration.get(\"functional\"),\n", + " \"grid_level\": configuration.get(\"grid_level\"),\n", + " \"grid_alignment\": configuration.get(\"grid_alignment\"),\n", + " \"max_memory_mb\": configuration.get(\"max_memory_mb\"),\n", + " \"orientation_count\": configuration.get(\"orientation_count\"),\n", + " \"commit\": source.get(\"commit\"),\n", + " \"branch\": source.get(\"branch\"),\n", + " \"dirty\": source.get(\"dirty\"),\n", + " \"implementation_hashes\": implementation_hashes,\n", + " \"worker_sha256\": hashes.get(\"worker_sha256\"),\n", + " \"rotation_runner_sha256\": hashes.get(\"rotation_runner_sha256\"),\n", + " \"validation_failure_count\": validation_failure_counts[run_label],\n", + " }\n", + " )\n", + "metadata_df = pd.DataFrame(metadata_rows)\n", + "display(metadata_df)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e35df038", + "metadata": {}, + "outputs": [], + "source": [ + "def records(frame: pd.DataFrame) -> list[dict[str, Any]]:\n", + " return frame.to_dict(orient=\"records\") if not frame.empty else []\n", + "\n", + "\n", + "def json_safe(value: Any) -> Any:\n", + " if isinstance(value, dict):\n", + " return {str(key): json_safe(item) for key, item in value.items()}\n", + " if isinstance(value, (list, tuple)):\n", + " return [json_safe(item) for item in value]\n", + " if isinstance(value, np.generic):\n", + " value = value.item()\n", + " if isinstance(value, float) and not np.isfinite(value):\n", + " return None\n", + " if value is pd.NA:\n", + " return None\n", + " return value\n", + "\n", + "\n", + "comparison_document = {\n", + " \"schema_version\": 1,\n", + " \"benchmark\": \"pyscf_ao_screening_rotation_comparison\",\n", + " \"generated_at\": datetime.now(UTC).isoformat(),\n", + " \"selected_run_labels\": [\n", + " str(document.get(\"run_label\") or path.stem) for path, document in DOCUMENTS\n", + " ],\n", + " \"configuration_summaries\": [\n", + " {\n", + " \"run_label\": str(document.get(\"run_label\") or path.stem),\n", + " \"configuration\": document.get(\"configuration\"),\n", + " }\n", + " for path, document in DOCUMENTS\n", + " ],\n", + " \"normalized_measurements\": records(normalized_df),\n", + " \"runtime_statistics\": records(runtime_statistics_df),\n", + " \"memory_statistics\": records(memory_statistics_df),\n", + " \"speedups_and_memory_reductions\": records(comparison_metrics_df),\n", + " \"cross_run_changes\": records(cross_run_df),\n", + " \"numerical_differences\": records(numerical_df),\n", + " \"numerical_summary\": records(numerical_summary_df),\n", + " \"orientation_sensitivity_fits\": records(orientation_fit_df),\n", + " \"validation_failures\": records(validation_df),\n", + " \"environment_metadata\": records(metadata_df),\n", + " \"source_provenance\": [\n", + " {\n", + " \"run_label\": str(document.get(\"run_label\") or path.stem),\n", + " \"source\": document.get(\"source\"),\n", + " \"runner_hashes\": document.get(\"runner_hashes\"),\n", + " }\n", + " for path, document in DOCUMENTS\n", + " ],\n", + "}\n", + "comparison_document = json_safe(comparison_document)\n", + "print(\n", + " f\"Prepared comparison JSON with \"\n", + " f\"{len(comparison_document['normalized_measurements'])} normalized rows\"\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9b9024b3", + "metadata": {}, + "outputs": [], + "source": [ + "ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)\n", + "FIGURE_DIR.mkdir(parents=True, exist_ok=True)\n", + "TABLE_DIR.mkdir(parents=True, exist_ok=True)\n", + "\n", + "with COMPARISON_JSON.open(\"w\", encoding=\"utf-8\") as stream:\n", + " json.dump(comparison_document, stream, indent=2, sort_keys=True, allow_nan=False)\n", + " stream.write(\"\\n\")\n", + "\n", + "tables = {\n", + " \"normalized-measurements.csv\": normalized_df,\n", + " \"validation-failures.csv\": validation_df,\n", + " \"status-summary.csv\": status_summary_df,\n", + " \"route-summary.csv\": route_summary_df,\n", + " \"runtime-statistics.csv\": runtime_statistics_df,\n", + " \"memory-statistics.csv\": memory_statistics_df,\n", + " \"speedups-and-memory-reductions.csv\": comparison_metrics_df,\n", + " \"cross-run-changes.csv\": cross_run_df,\n", + " \"numerical-differences.csv\": numerical_df,\n", + " \"numerical-summary.csv\": numerical_summary_df,\n", + " \"orientation-sensitivity-fits.csv\": orientation_fit_df,\n", + " \"environment-and-source-metadata.csv\": metadata_df,\n", + "}\n", + "for filename, table in tables.items():\n", + " table.to_csv(TABLE_DIR / filename, index=False)\n", + "\n", + "plot_orientation_lines(\n", + " \"runtime_seconds\",\n", + " \"Runtime (s)\",\n", + " FIGURE_DIR / \"runtime-by-orientation.png\",\n", + ")\n", + "plot_mode_distribution(\n", + " \"runtime_seconds\",\n", + " \"Runtime (s)\",\n", + " FIGURE_DIR / \"runtime-by-mode.png\",\n", + ")\n", + "plot_measurement_heatmaps(\"runtime_seconds\", \"Runtime (s)\", FIGURE_DIR)\n", + "plot_orientation_lines(\n", + " \"incremental_peak_gib\",\n", + " \"Incremental peak memory (GiB)\",\n", + " FIGURE_DIR / \"memory-by-orientation.png\",\n", + ")\n", + "plot_mode_distribution(\n", + " \"incremental_peak_gib\",\n", + " \"Incremental peak memory (GiB)\",\n", + " FIGURE_DIR / \"memory-by-mode.png\",\n", + ")\n", + "plot_measurement_heatmaps(\n", + " \"incremental_peak_gib\", \"Incremental peak memory (GiB)\", FIGURE_DIR\n", + ")\n", + "plot_ratio_heatmaps(\n", + " (\"dense_to_screened_runtime_speedup\", \"dense_to_gpu_runtime_speedup\"),\n", + " \"runtime speedup\",\n", + " FIGURE_DIR,\n", + ")\n", + "plot_ratio_heatmaps(\n", + " (\"dense_to_screened_memory_reduction\", \"dense_to_gpu_memory_reduction\"),\n", + " \"memory reduction\",\n", + " FIGURE_DIR,\n", + ")\n", + "for fingerprint in FINGERPRINT_LABELS:\n", + " plot_fingerprint_heatmaps(fingerprint, FIGURE_DIR)\n", + "\n", + "print(f\"Comparison JSON: {COMPARISON_JSON}\")\n", + "print(f\"Tables: {TABLE_DIR}\")\n", + "print(f\"Figures: {FIGURE_DIR}\")" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/benchmarks/run_pyscf_ao_screening_benchmark.py b/benchmarks/run_pyscf_ao_screening_benchmark.py new file mode 100644 index 00000000..381a47f0 --- /dev/null +++ b/benchmarks/run_pyscf_ao_screening_benchmark.py @@ -0,0 +1,1130 @@ +"""Run isolated Skala PySCF and GPU4PySCF AO-screening benchmarks.""" + +from __future__ import annotations + +import argparse +import hashlib +import importlib.metadata +import inspect +import json +import math +import os +import platform +import re +import socket +import subprocess +import sys +import tempfile +import time +import traceback +from contextlib import AbstractContextManager, nullcontext +from dataclasses import asdict, dataclass +from datetime import UTC, datetime +from itertools import pairwise +from pathlib import Path +from typing import Any, cast +from unittest.mock import patch + +FULL_CARBON_COUNTS = (2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12) +EXPECTED_AO_COUNTS = ( + 294, + 411, + 528, + 645, + 762, + 879, + 996, + 1113, + 1230, + 1347, + 1464, +) +MODES = ("cpu", "cpu_dense", "gpu") +MEASUREMENTS = ("runtime", "memory") +TERMINAL_STATUSES = {"ok", "timeout", "oom", "error", "skipped_after_resource_failure"} +WORKER_RESULT_PREFIX = "SKALA_BENCHMARK_RESULT=" +THREAD_ENVIRONMENT_VARIABLES = ( + "OMP_NUM_THREADS", + "MKL_NUM_THREADS", + "OPENBLAS_NUM_THREADS", + "NUMEXPR_NUM_THREADS", +) + +Vector = tuple[float, float, float] +Atom = tuple[str, float, float, float] + +CARBON_CARBON_BOND_ANGSTROM = 1.54 +CARBON_HYDROGEN_BOND_ANGSTROM = 1.09 +CARBON_BOND_ANGLE_DEGREES = 112.0 +COORDINATE_PRECISION = 12 +GEOMETRY_PARAMETERS = { + "version": "zigzag-alkane-v1", + "carbon_carbon_bond_angstrom": CARBON_CARBON_BOND_ANGSTROM, + "carbon_hydrogen_bond_angstrom": CARBON_HYDROGEN_BOND_ANGSTROM, + "carbon_bond_angle_degrees": CARBON_BOND_ANGLE_DEGREES, + "hydrogen_dot_product": -1.0 / 3.0, + "coordinate_precision": COORDINATE_PRECISION, +} +EXPECTED_AOS_BY_CARBON: dict[int, int] = dict( + zip(FULL_CARBON_COUNTS, EXPECTED_AO_COUNTS, strict=True) +) + + +def find_repository_root(start: Path) -> Path: + for candidate in (start.resolve(), *start.resolve().parents): + if (candidate / "pyproject.toml").is_file() and ( + candidate / "src" / "skala" + ).is_dir(): + return candidate + raise FileNotFoundError(f"Could not find the Skala repository above {start}") + + +RUNNER_ROOT = find_repository_root(Path(__file__).resolve()) + + +@dataclass(frozen=True) +class BenchmarkConfig: + source_root: Path + results_dir: Path + run_label: str + functional: str = "skala-1.1" + basis: str = "def2-qzvpp" + grid_level: int = 1 + grid_alignment: int = 1 + max_memory_mb: int = 2000 + cpu_threads: int = 4 + runtime_repetitions: int = 3 + worker_timeout_seconds: int = 30 * 60 + smoke_run: bool = False + + @property + def carbon_counts(self) -> tuple[int, ...]: + return FULL_CARBON_COUNTS[:1] if self.smoke_run else FULL_CARBON_COUNTS + + @property + def worker_thread_environment(self) -> dict[str, str]: + thread_count = str(self.cpu_threads) + return {name: thread_count for name in THREAD_ENVIRONMENT_VARIABLES} + + def as_json(self) -> dict[str, Any]: + data = asdict(self) + data["source_root"] = str(self.source_root) + data["results_dir"] = str(self.results_dir) + data["carbon_counts"] = list(self.carbon_counts) + data["full_carbon_counts"] = list(FULL_CARBON_COUNTS) + data["expected_ao_counts"] = list(EXPECTED_AO_COUNTS) + data["worker_thread_environment"] = self.worker_thread_environment + return data + + +def vector_add(left: Vector, right: Vector) -> Vector: + return tuple(a + b for a, b in zip(left, right, strict=True)) # type: ignore[return-value] + + +def vector_subtract(left: Vector, right: Vector) -> Vector: + return tuple(a - b for a, b in zip(left, right, strict=True)) # type: ignore[return-value] + + +def vector_scale(scale: float, vector: Vector) -> Vector: + return tuple(scale * value for value in vector) # type: ignore[return-value] + + +def vector_dot(left: Vector, right: Vector) -> float: + return sum(a * b for a, b in zip(left, right, strict=True)) + + +def vector_cross(left: Vector, right: Vector) -> Vector: + return ( + left[1] * right[2] - left[2] * right[1], + left[2] * right[0] - left[0] * right[2], + left[0] * right[1] - left[1] * right[0], + ) + + +def vector_normalize(vector: Vector) -> Vector: + norm = math.sqrt(vector_dot(vector, vector)) + if norm == 0.0: + raise ValueError("Cannot normalize a zero vector") + return vector_scale(1.0 / norm, vector) + + +def carbon_backbone(carbon_count: int) -> tuple[Vector, ...]: + if carbon_count < 2: + raise ValueError("The benchmark requires at least two carbon atoms") + half_turn = math.radians((180.0 - CARBON_BOND_ANGLE_DEGREES) / 2.0) + positions: list[Vector] = [(0.0, 0.0, 0.0)] + for bond_index in range(carbon_count - 1): + angle = half_turn if bond_index % 2 == 0 else -half_turn + direction = (math.cos(angle), math.sin(angle), 0.0) + positions.append( + vector_add( + positions[-1], vector_scale(CARBON_CARBON_BOND_ANGSTROM, direction) + ) + ) + + center = tuple( + sum(position[axis] for position in positions) / carbon_count + for axis in range(3) + ) + return tuple(vector_subtract(position, center) for position in positions) # type: ignore[arg-type] + + +def terminal_hydrogen_directions( + carbon: Vector, neighbor: Vector +) -> tuple[Vector, ...]: + neighbor_direction = vector_normalize(vector_subtract(neighbor, carbon)) + perpendicular = (0.0, 0.0, 1.0) + second_perpendicular = vector_normalize( + vector_cross(neighbor_direction, perpendicular) + ) + radial_scale = math.sqrt(8.0 / 9.0) + directions = [] + for index in range(3): + phase = 2.0 * math.pi * index / 3.0 + radial = vector_add( + vector_scale(math.cos(phase), perpendicular), + vector_scale(math.sin(phase), second_perpendicular), + ) + directions.append( + vector_add( + vector_scale(-1.0 / 3.0, neighbor_direction), + vector_scale(radial_scale, radial), + ) + ) + return tuple(directions) + + +def internal_hydrogen_directions( + carbon: Vector, previous_carbon: Vector, next_carbon: Vector +) -> tuple[Vector, Vector]: + previous_direction = vector_normalize(vector_subtract(previous_carbon, carbon)) + next_direction = vector_normalize(vector_subtract(next_carbon, carbon)) + neighbor_dot = vector_dot(previous_direction, next_direction) + in_plane_scale = (-1.0 / 3.0) / (1.0 + neighbor_dot) + in_plane = vector_scale( + in_plane_scale, vector_add(previous_direction, next_direction) + ) + normal = vector_normalize(vector_cross(previous_direction, next_direction)) + normal_scale = math.sqrt(max(0.0, 1.0 - vector_dot(in_plane, in_plane))) + return ( + vector_add(in_plane, vector_scale(normal_scale, normal)), + vector_subtract(in_plane, vector_scale(normal_scale, normal)), + ) + + +def generate_alkane_atoms(carbon_count: int) -> tuple[Atom, ...]: + carbons = carbon_backbone(carbon_count) + atoms: list[Atom] = [("C", *position) for position in carbons] + for index, carbon in enumerate(carbons): + if index == 0: + directions = terminal_hydrogen_directions(carbon, carbons[1]) + elif index == carbon_count - 1: + directions = terminal_hydrogen_directions(carbon, carbons[-2]) + else: + directions = internal_hydrogen_directions( + carbon, carbons[index - 1], carbons[index + 1] + ) + atoms.extend( + ( + "H", + *vector_add( + carbon, vector_scale(CARBON_HYDROGEN_BOND_ANGSTROM, direction) + ), + ) + for direction in directions + ) + return tuple(atoms) + + +def atoms_to_pyscf(atoms: tuple[Atom, ...]) -> str: + return "\n".join( + f"{element} {x:.{COORDINATE_PRECISION}f} " + f"{y:.{COORDINATE_PRECISION}f} {z:.{COORDINATE_PRECISION}f}" + for element, x, y, z in atoms + ) + + +@dataclass(frozen=True) +class MoleculeSpec: + carbon_count: int + expected_aos: int + formula: str + atoms: tuple[Atom, ...] + + @property + def atom_text(self) -> str: + return atoms_to_pyscf(self.atoms) + + @property + def coordinate_sha256(self) -> str: + return hashlib.sha256(self.atom_text.encode()).hexdigest() + + def as_json(self) -> dict[str, Any]: + return { + "carbon_count": self.carbon_count, + "expected_aos": self.expected_aos, + "formula": self.formula, + "atoms": [ + {"element": element, "xyz_angstrom": [x, y, z]} + for element, x, y, z in self.atoms + ], + "coordinate_sha256": self.coordinate_sha256, + } + + +def make_molecule_spec(carbon_count: int) -> MoleculeSpec: + hydrogen_count = 2 * carbon_count + 2 + return MoleculeSpec( + carbon_count=carbon_count, + expected_aos=EXPECTED_AOS_BY_CARBON[carbon_count], + formula=f"C{carbon_count}H{hydrogen_count}", + atoms=generate_alkane_atoms(carbon_count), + ) + + +FULL_MOLECULE_LADDER = tuple(make_molecule_spec(count) for count in FULL_CARBON_COUNTS) + + +def package_version(distribution: str) -> str | None: + try: + return importlib.metadata.version(distribution) + except importlib.metadata.PackageNotFoundError: + return None + + +def verify_skala_import(source_root: Path) -> str: + import skala + + imported_path = Path(skala.__file__).resolve() + expected_root = (source_root / "src").resolve() + try: + imported_path.relative_to(expected_root) + except ValueError as error: + raise RuntimeError( + f"Imported Skala from {imported_path}, expected a module below {expected_root}" + ) from error + return str(imported_path) + + +def collect_environment(payload: dict[str, Any]) -> dict[str, Any]: + import pyscf + import torch + + source_root = Path(payload["source_root"]).resolve() + imported_skala = verify_skala_import(source_root) + cuda_available = torch.cuda.is_available() + gpu_name = torch.cuda.get_device_name(0) if cuda_available else None + cupy_version = package_version("cupy-cuda12x") or package_version("cupy") + torch_cuda_version = getattr(getattr(torch, "version", None), "cuda", None) + return { + "python": sys.version, + "python_executable": sys.executable, + "platform": platform.platform(), + "hostname": socket.gethostname(), + "processor": platform.processor(), + "logical_cpu_count": os.cpu_count(), + "source_root": str(source_root), + "imported_skala": imported_skala, + "packages": { + "skala": package_version("skala"), + "pyscf": pyscf.__version__, + "gpu4pyscf": package_version("gpu4pyscf-cuda12x") + or package_version("gpu4pyscf"), + "torch": torch.__version__, + "cupy": cupy_version, + "memray": package_version("memray"), + }, + "cuda": { + "available": cuda_available, + "torch_cuda_version": torch_cuda_version, + "device_name": gpu_name, + "device_count": torch.cuda.device_count() if cuda_available else 0, + }, + "thread_environment": { + name: os.environ.get(name) for name in THREAD_ENVIRONMENT_VARIABLES + }, + } + + +def find_route_controller(numint: Any) -> tuple[Any, Any, Any]: + candidates = (numint, getattr(numint, "integrator", None)) + control_symbols = {"_should_screen_aos", "_functional_supports_atom_chunking"} + for candidate in candidates: + if candidate is None: + continue + route_callable = inspect.unwrap(type(candidate).__call__) + referenced_names = set(route_callable.__code__.co_names) + if referenced_names & control_symbols: + route_module = inspect.getmodule(route_callable) + if route_module is None: + raise RuntimeError( + f"Cannot identify the module defining {route_callable.__qualname__}" + ) + return candidate, route_callable, route_module + raise RuntimeError("Cannot find the Skala route-selection implementation") + + +def force_dense_route(numint: Any) -> AbstractContextManager[Any]: + route_owner, route_callable, route_module = find_route_controller(numint) + referenced_names = set(route_callable.__code__.co_names) + if "_should_screen_aos" in referenced_names and hasattr( + route_module, "_should_screen_aos" + ): + from tests.utils import patch_ao_screening + + return patch_ao_screening(False, module=route_module) + if "_functional_supports_atom_chunking" in referenced_names and hasattr( + type(route_owner), "_functional_supports_atom_chunking" + ): + return patch.object( + type(route_owner), + "_functional_supports_atom_chunking", + return_value=False, + ) + raise RuntimeError( + "Cannot force dense evaluation through " + f"{route_module.__name__}.{route_callable.__qualname__}" + ) + + +def route_metadata(numint: Any, mol: Any, forced_dense: bool) -> dict[str, Any]: + from pyscf.dft import numint as pyscf_numint + + route_owner, route_callable, route_module = find_route_controller(numint) + routing_source = inspect.getsource(route_callable) + routing_sha256 = hashlib.sha256(routing_source.encode()).hexdigest() + referenced_names = set(route_callable.__code__.co_names) + if "_should_screen_aos" in referenced_names and hasattr( + route_module, "_should_screen_aos" + ): + route_decision_callable = route_module._should_screen_aos + route_decision = bool(route_decision_callable(mol)) + supports_screened_evaluation = bool( + numint.feature_spec.supports_screened_evaluation + ) + route_selector = "ao_threshold" + elif "_functional_supports_atom_chunking" in referenced_names and hasattr( + type(route_owner), "_functional_supports_atom_chunking" + ): + route_decision_callable = route_owner._functional_supports_atom_chunking + route_decision = bool(route_decision_callable()) + supports_screened_evaluation = route_decision + route_selector = "functional_capability" + else: + raise RuntimeError("Unrecognized Skala route-selection API") + route_decision_source = inspect.getsource(route_decision_callable) + if "_should_screen_aos" in routing_source and ( + "_global_screened_features" in routing_source + or "_integrate_screened" in routing_source + ): + implementation = "threshold_gated_global_ao_screening" + elif "chunked_features" in routing_source: + implementation = "legacy_atom_chunking" + else: + implementation = "unclassified" + + switch_size = int(pyscf_numint.SWITCH_SIZE) + if forced_dense or not supports_screened_evaluation: + selected_route = "dense" + elif implementation == "threshold_gated_global_ao_screening": + selected_route = "global_ao_screening" if route_decision else "dense" + elif implementation == "legacy_atom_chunking": + selected_route = "atom_chunking" + else: + selected_route = "unknown" + return { + "request": "forced_dense" if forced_dense else "natural", + "implementation": implementation, + "implementation_sha256": routing_sha256, + "implementation_target": ( + f"{route_module.__name__}.{route_callable.__qualname__}" + ), + "route_selector": route_selector, + "route_decision": { + "target": ( + f"{route_decision_callable.__module__}." + f"{route_decision_callable.__qualname__}" + ), + "source": route_decision_source, + "source_sha256": hashlib.sha256(route_decision_source.encode()).hexdigest(), + "result": route_decision, + }, + "functional_supports_screened_evaluation": supports_screened_evaluation, + "pyscf_switch_size": switch_size, + "selected_route": selected_route, + } + + +def build_case(payload: dict[str, Any]) -> dict[str, Any]: + import numpy as np + import torch + from pyscf import dft, gto, lib + + source_root = Path(payload["source_root"]).resolve() + imported_skala = verify_skala_import(source_root) + thread_count = int(payload["cpu_threads"]) + lib.num_threads(thread_count) + torch.set_num_threads(thread_count) + + molecule = payload["molecule"] + mol = gto.M( + atom=molecule["atom_text"], + basis=payload["basis"], + charge=0, + spin=0, + unit="Angstrom", + cart=False, + verbose=0, + ) + initial_dm = dft.RKS(mol).get_init_guess() + backend = payload["backend"] + if backend == "cpu": + from skala.pyscf import SkalaKS as CpuSkalaKS + + ks = CpuSkalaKS(mol, xc=payload["functional"], with_dftd3=False) + dm = initial_dm + + def synchronize() -> None: + return None + + to_numpy = np.asarray + elif backend == "gpu": + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is not available") + import cupy + + from skala.gpu4pyscf import SkalaKS as GpuSkalaKS + + ks = GpuSkalaKS(mol, xc=payload["functional"], with_dftd3=False) + dm = cupy.asarray(initial_dm) + synchronize = torch.cuda.synchronize + to_numpy = cupy.asnumpy + else: + raise ValueError(f"Unknown backend: {backend}") + + ks.grids.level = int(payload["grid_level"]) + ks.grids.alignment = int(payload["grid_alignment"]) + ks.grids.build(sort_grids=False) + grid_weights = ks.grids.weights + if grid_weights is None: + raise RuntimeError("Grid construction did not produce weights") + numint = ks._numint + forced_dense = bool(payload["forced_dense"]) + route = route_metadata(numint, mol, forced_dense) + system = { + "formula": molecule["formula"], + "carbon_count": int(molecule["carbon_count"]), + "electron_count": int(mol.nelectron), + "actual_aos": int(mol.nao_nr()), + "grid_points": int(cast(Any, grid_weights).size), + "coordinate_sha256": molecule["coordinate_sha256"], + "imported_skala": imported_skala, + } + return { + "mol": mol, + "grids": ks.grids, + "dm": dm, + "numint": numint, + "backend": backend, + "synchronize": synchronize, + "to_numpy": to_numpy, + "route": route, + "system": system, + } + + +def fingerprint(result: tuple[Any, Any, Any], to_numpy: Any) -> dict[str, float]: + import numpy as np + + electron_integral, xc_energy, vxc = result + matrix = np.asarray(to_numpy(vxc), dtype=np.float64) + return { + "electron_integral": float(electron_integral), + "xc_energy": float(xc_energy), + "vxc_sum": float(matrix.sum()), + "vxc_trace": float(np.trace(matrix)), + "vxc_frobenius_norm": float(np.linalg.norm(matrix)), + "vxc_max_abs": float(np.max(np.abs(matrix))), + } + + +def run_measurement(payload: dict[str, Any]) -> dict[str, Any]: + import torch + + case = build_case(payload) + numint = case["numint"] + dense_route_override: AbstractContextManager[Any] + if not payload["forced_dense"]: + dense_route_override = nullcontext() + else: + dense_route_override = force_dense_route(numint) + + def evaluate() -> tuple[Any, Any, Any]: + return numint.nr_rks( + case["mol"], + case["grids"], + None, + case["dm"], + max_memory=int(payload["max_memory_mb"]), + ) + + result: tuple[Any, Any, Any] | None = None + with dense_route_override: + measurement = payload["measurement"] + if measurement == "runtime": + case["synchronize"]() + started = time.perf_counter() + result = evaluate() + case["synchronize"]() + elapsed_seconds = time.perf_counter() - started + measurement_data = {"runtime_seconds": elapsed_seconds} + elif measurement == "memory" and case["backend"] == "cpu": + import memray + + with tempfile.TemporaryDirectory() as temp_dir: + profile_path = Path(temp_dir) / "allocations.bin" + with memray.Tracker(profile_path): + result = evaluate() + peak_bytes = int(memray.FileReader(profile_path).metadata.peak_memory) + measurement_data = {"incremental_peak_bytes": peak_bytes} + elif measurement == "memory" and case["backend"] == "gpu": + case["synchronize"]() + torch.cuda.empty_cache() + case["synchronize"]() + baseline_bytes = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + result = evaluate() + case["synchronize"]() + peak_bytes = max(0, torch.cuda.max_memory_allocated() - baseline_bytes) + measurement_data = { + "incremental_peak_bytes": int(peak_bytes), + "allocator_baseline_bytes": int(baseline_bytes), + } + else: + raise ValueError(f"Unknown measurement: {measurement}") + + assert result is not None + return { + "status": "ok", + "measurement": payload["measurement"], + "mode": payload["mode"], + "route": case["route"], + "system": case["system"], + "fingerprint": fingerprint(result, case["to_numpy"]), + **measurement_data, + } + + +def classify_exception(error: Exception) -> str: + message = f"{type(error).__name__}: {error}".lower() + if ( + isinstance(error, MemoryError) + or "out of memory" in message + or "bad alloc" in message + ): + return "oom" + return "error" + + +def emit_worker_record(record: dict[str, Any]) -> None: + print(WORKER_RESULT_PREFIX + json.dumps(record, sort_keys=True), flush=True) + + +def worker_main() -> None: + payload = json.load(sys.stdin) + try: + if payload["operation"] == "environment": + emit_worker_record( + {"status": "ok", "environment": collect_environment(payload)} + ) + elif payload["operation"] == "measure": + emit_worker_record(run_measurement(payload)) + else: + raise ValueError(f"Unknown operation: {payload['operation']}") + except Exception as error: # noqa: BLE001 - serialize worker failures for the parent + emit_worker_record( + { + "status": classify_exception(error), + "error_type": type(error).__name__, + "error": str(error), + "traceback": traceback.format_exc()[-12000:], + "measurement": payload.get("measurement"), + "mode": payload.get("mode"), + } + ) + + +def utc_now() -> str: + return datetime.now(UTC).isoformat() + + +def git_output(source_root: Path, *arguments: str) -> str: + completed = subprocess.run( + ["git", "-C", str(source_root), *arguments], + check=True, + capture_output=True, + text=True, + ) + return completed.stdout.strip() + + +def source_metadata(source_root: Path) -> dict[str, Any]: + return { + "root": str(source_root), + "commit": git_output(source_root, "rev-parse", "HEAD"), + "branch": git_output(source_root, "branch", "--show-current") or None, + "dirty": bool(git_output(source_root, "status", "--porcelain")), + } + + +def runner_sha256() -> str: + return hashlib.sha256(Path(__file__).read_bytes()).hexdigest() + + +def worker_environment(config: BenchmarkConfig) -> dict[str, str]: + environment = os.environ.copy() + source_python_path = str(config.source_root / "src") + existing_python_path = environment.get("PYTHONPATH") + python_paths = [source_python_path, str(RUNNER_ROOT)] + if existing_python_path: + python_paths.append(existing_python_path) + environment["PYTHONPATH"] = os.pathsep.join(python_paths) + environment.update(config.worker_thread_environment) + return environment + + +def execute_worker(payload: dict[str, Any], config: BenchmarkConfig) -> dict[str, Any]: + try: + completed = subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--worker"], + input=json.dumps(payload), + cwd=config.source_root, + env=worker_environment(config), + capture_output=True, + text=True, + timeout=config.worker_timeout_seconds, + check=False, + ) + except subprocess.TimeoutExpired as error: + return { + "status": "timeout", + "measurement": payload.get("measurement"), + "mode": payload.get("mode"), + "error": f"Worker exceeded {config.worker_timeout_seconds} seconds", + "stdout_tail": (error.stdout or "")[-4000:], + "stderr_tail": (error.stderr or "")[-4000:], + } + + marker_lines = [ + line.removeprefix(WORKER_RESULT_PREFIX) + for line in completed.stdout.splitlines() + if line.startswith(WORKER_RESULT_PREFIX) + ] + if marker_lines: + record = json.loads(marker_lines[-1]) + if record["status"] != "ok": + record["stderr_tail"] = completed.stderr[-4000:] + return record + + combined_output = f"{completed.stdout}\n{completed.stderr}".lower() + status = ( + "oom" + if completed.returncode in {-9, 137} or "out of memory" in combined_output + else "error" + ) + return { + "status": status, + "measurement": payload.get("measurement"), + "mode": payload.get("mode"), + "error": f"Worker exited with code {completed.returncode} without a result record", + "stdout_tail": completed.stdout[-4000:], + "stderr_tail": completed.stderr[-4000:], + } + + +def execute_measurement( + payload: dict[str, Any], config: BenchmarkConfig +) -> dict[str, Any]: + if payload["measurement"] != "runtime": + return execute_worker(payload, config) + + runtime_samples: list[float] = [] + first_result: dict[str, Any] | None = None + for sample_index in range(config.runtime_repetitions): + result = execute_worker(payload, config) + if result["status"] != "ok": + result["runtime_samples_seconds"] = runtime_samples + result["failed_runtime_sample_index"] = sample_index + return result + if first_result is None: + first_result = result + else: + for key in ("route", "system"): + if result.get(key) != first_result.get(key): + raise ValueError( + f"Runtime worker {key} changed between isolated samples" + ) + runtime_samples.append(float(result["runtime_seconds"])) + + assert first_result is not None + first_result.pop("runtime_seconds") + first_result["runtime_samples_seconds"] = runtime_samples + return first_result + + +def atomic_write_json(path: Path, document: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary_path: Path | None = None + try: + with tempfile.NamedTemporaryFile( + "w", + encoding="utf-8", + dir=path.parent, + prefix=f".{path.name}.", + delete=False, + ) as stream: + json.dump(document, stream, indent=2, sort_keys=True, allow_nan=False) + stream.write("\n") + stream.flush() + os.fsync(stream.fileno()) + temporary_path = Path(stream.name) + temporary_path.replace(path) + finally: + if temporary_path is not None and temporary_path.exists(): + temporary_path.unlink() + + +def worker_payload( + config: BenchmarkConfig, + molecule: MoleculeSpec, + mode: str, + measurement: str, +) -> dict[str, Any]: + backend = "gpu" if mode.startswith("gpu") else "cpu" + return { + "operation": "measure", + "source_root": str(config.source_root), + "functional": config.functional, + "basis": config.basis, + "grid_level": config.grid_level, + "grid_alignment": config.grid_alignment, + "max_memory_mb": config.max_memory_mb, + "cpu_threads": config.cpu_threads, + "backend": backend, + "forced_dense": mode.endswith("_dense"), + "mode": mode, + "measurement": measurement, + "molecule": { + "carbon_count": molecule.carbon_count, + "formula": molecule.formula, + "atom_text": molecule.atom_text, + "coordinate_sha256": molecule.coordinate_sha256, + }, + } + + +def new_result_document( + config: BenchmarkConfig, + molecules: tuple[MoleculeSpec, ...], + source: dict[str, Any], + environment: dict[str, Any], +) -> dict[str, Any]: + created_at = utc_now() + return { + "schema_version": 1, + "created_at": created_at, + "updated_at": created_at, + "run_label": config.run_label, + "source": source, + "environment": environment, + "configuration": config.as_json(), + "geometry": GEOMETRY_PARAMETERS, + "worker_sha256": runner_sha256(), + "molecules": { + molecule.formula: { + **molecule.as_json(), + "observed": None, + "modes": {mode: {} for mode in MODES}, + } + for molecule in molecules + }, + } + + +def validate_resume_document( + document: dict[str, Any], config: BenchmarkConfig, source: dict[str, Any] +) -> None: + if document.get("schema_version") != 1: + raise ValueError("Cannot resume a result file with a different schema version") + if document.get("worker_sha256") != runner_sha256(): + raise ValueError( + "Cannot resume results created by a different runner implementation" + ) + if document.get("source", {}).get("commit") != source["commit"]: + raise ValueError("Cannot resume results from a different Git commit") + if document.get("configuration") != config.as_json(): + raise ValueError( + "Cannot resume results created with a different benchmark configuration" + ) + + +def merge_worker_result( + molecule_record: dict[str, Any], mode: str, measurement: str, result: dict[str, Any] +) -> None: + result = dict(result) + system = result.pop("system", None) + route = result.pop("route", None) + if system is not None: + observed = molecule_record.get("observed") + if observed is not None and observed != system: + raise ValueError( + f"Worker system metadata changed for {molecule_record['formula']}" + ) + molecule_record["observed"] = system + mode_record = molecule_record["modes"][mode] + if route is not None: + existing_route = mode_record.get("route") + if existing_route is not None and existing_route != route: + raise ValueError( + f"Worker route metadata changed for {molecule_record['formula']} {mode}" + ) + mode_record["route"] = route + if measurement == "runtime" and "runtime_seconds" in result: + result["runtime_samples_seconds"] = [result.pop("runtime_seconds")] + mode_record[measurement] = result + + +def result_path(config: BenchmarkConfig, source: dict[str, Any]) -> Path: + safe_label = re.sub(r"[^A-Za-z0-9_.-]+", "-", config.run_label).strip("-.") + if not safe_label: + raise ValueError("The run label must contain a filename-safe character") + return ( + config.results_dir + / f"skala-pyscf-ao-screening-{safe_label}-{source['commit'][:12]}.json" + ) + + +def run_worker_preflight(config: BenchmarkConfig) -> dict[str, Any]: + result = execute_worker( + {"operation": "environment", "source_root": str(config.source_root)}, config + ) + if result["status"] != "ok": + raise RuntimeError(f"Benchmark preflight failed: {result}") + environment = result["environment"] + imported_path = Path(environment["imported_skala"]) + imported_path.relative_to(config.source_root / "src") + required_packages = ("skala", "pyscf", "torch", "memray") + missing = [ + name for name in required_packages if not environment["packages"].get(name) + ] + if missing: + raise RuntimeError(f"Worker environment is missing packages: {missing}") + return environment + + +def atom_distance(left: Atom, right: Atom) -> float: + return math.dist(left[1:], right[1:]) + + +def validate_molecule_ladder(config: BenchmarkConfig) -> None: + from pyscf import gto + + assert len(FULL_MOLECULE_LADDER) == len(FULL_CARBON_COUNTS) + assert len( + {molecule.coordinate_sha256 for molecule in FULL_MOLECULE_LADDER} + ) == len(FULL_CARBON_COUNTS) + for molecule, expected_aos in zip( + FULL_MOLECULE_LADDER, EXPECTED_AO_COUNTS, strict=True + ): + carbon_count = molecule.carbon_count + hydrogen_count = 2 * carbon_count + 2 + assert molecule.formula == f"C{carbon_count}H{hydrogen_count}" + assert len(molecule.atoms) == carbon_count + hydrogen_count + assert molecule.atoms == generate_alkane_atoms(carbon_count) + + carbons = molecule.atoms[:carbon_count] + hydrogens = molecule.atoms[carbon_count:] + for left, right in pairwise(carbons): + assert math.isclose( + atom_distance(left, right), + CARBON_CARBON_BOND_ANGSTROM, + abs_tol=1e-12, + ) + for hydrogen in hydrogens: + nearest_carbon = min(atom_distance(hydrogen, carbon) for carbon in carbons) + assert math.isclose( + nearest_carbon, + CARBON_HYDROGEN_BOND_ANGSTROM, + abs_tol=1e-12, + ) + + mol = gto.M( + atom=molecule.atom_text, + basis=config.basis, + charge=0, + spin=0, + unit="Angstrom", + cart=False, + verbose=0, + ) + assert mol.nao_nr() == expected_aos == molecule.expected_aos + assert mol.nelectron % 2 == 0 + + +def run_benchmark( + config: BenchmarkConfig, + molecules: tuple[MoleculeSpec, ...], + environment: dict[str, Any], +) -> Path: + source = source_metadata(config.source_root) + output_path = result_path(config, source) + if output_path.exists(): + document = json.loads(output_path.read_text(encoding="utf-8")) + validate_resume_document(document, config, source) + else: + document = new_result_document(config, molecules, source, environment) + atomic_write_json(output_path, document) + + cuda_available = bool(document["environment"]["cuda"]["available"]) + for mode in MODES: + backend = "gpu" if mode.startswith("gpu") else "cpu" + for measurement in MEASUREMENTS: + blocked_by: dict[str, Any] | None = None + for molecule in molecules: + molecule_record = document["molecules"][molecule.formula] + existing = molecule_record["modes"][mode].get(measurement) + if existing and existing.get("status") in TERMINAL_STATUSES: + if existing["status"] in {"oom", "timeout"}: + blocked_by = { + "formula": molecule.formula, + "status": existing["status"], + } + continue + + result: dict[str, Any] + if backend == "gpu" and not cuda_available: + result = { + "status": "error", + "mode": mode, + "measurement": measurement, + "error": "CUDA is not available in the worker environment", + } + elif blocked_by is not None: + result = { + "status": "skipped_after_resource_failure", + "mode": mode, + "measurement": measurement, + "blocked_by": blocked_by, + } + else: + result = execute_measurement( + worker_payload(config, molecule, mode, measurement), config + ) + + merge_worker_result(molecule_record, mode, measurement, result) + document["updated_at"] = utc_now() + atomic_write_json(output_path, document) + if result["status"] in {"oom", "timeout"}: + blocked_by = { + "formula": molecule.formula, + "status": result["status"], + } + print( + f"{mode:9s} {measurement:7s} {molecule.formula:8s} " + f"{result['status']}" + ) + return output_path + + +def parse_arguments(argv: list[str]) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Benchmark one Skala XC/Vxc evaluation on CPU and GPU.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + parser.add_argument( + "--label", + required=True, + help="Result label, normally the revision name such as 'mr' or 'main'.", + ) + parser.add_argument( + "--source-root", + type=Path, + default=RUNNER_ROOT, + help="Skala checkout whose src/skala package is benchmarked.", + ) + parser.add_argument( + "--results-dir", + type=Path, + default=RUNNER_ROOT / "benchmarks" / "results", + help="Directory for commit-labelled JSON output.", + ) + parser.add_argument("--functional", default="skala-1.1") + parser.add_argument("--basis", default="def2-qzvpp") + parser.add_argument("--grid-level", type=int, default=1) + parser.add_argument("--max-memory-mb", type=int, default=2000) + parser.add_argument("--threads", type=int, default=4) + parser.add_argument("--runtime-repetitions", type=int, default=3) + parser.add_argument("--timeout-minutes", type=float, default=30.0) + parser.add_argument( + "--smoke", + action="store_true", + help="Run only C2H6 instead of the full 11-molecule ladder.", + ) + parser.add_argument( + "--preflight-only", + action="store_true", + help="Validate geometry, dependencies, source import, and CUDA without measurements.", + ) + return parser.parse_args(argv) + + +def config_from_arguments(arguments: argparse.Namespace) -> BenchmarkConfig: + if arguments.timeout_minutes <= 0: + raise ValueError("--timeout-minutes must be positive") + if arguments.threads <= 0: + raise ValueError("--threads must be positive") + if arguments.runtime_repetitions <= 0: + raise ValueError("--runtime-repetitions must be positive") + source_root = arguments.source_root.expanduser().resolve() + if not (source_root / "src" / "skala").is_dir(): + raise FileNotFoundError(f"No src/skala package below {source_root}") + return BenchmarkConfig( + source_root=source_root, + results_dir=arguments.results_dir.expanduser().resolve(), + run_label=arguments.label, + functional=arguments.functional, + basis=arguments.basis, + grid_level=arguments.grid_level, + max_memory_mb=arguments.max_memory_mb, + cpu_threads=arguments.threads, + runtime_repetitions=arguments.runtime_repetitions, + worker_timeout_seconds=round(arguments.timeout_minutes * 60), + smoke_run=arguments.smoke, + ) + + +def main(argv: list[str] | None = None) -> int: + arguments = parse_arguments(sys.argv[1:] if argv is None else argv) + config = config_from_arguments(arguments) + validate_molecule_ladder(config) + environment = run_worker_preflight(config) + print("Benchmark configuration:") + print(json.dumps(config.as_json(), indent=2, sort_keys=True)) + print(f"Skala import: {environment['imported_skala']}") + print(f"Python: {environment['python_executable']}") + print(f"CUDA: {environment['cuda']}") + if arguments.preflight_only: + print("Preflight passed.") + return 0 + + molecules = FULL_MOLECULE_LADDER[:1] if config.smoke_run else FULL_MOLECULE_LADDER + output_path = run_benchmark(config, molecules, environment) + print(f"Results written to {output_path}") + return 0 + + +if __name__ == "__main__": + if sys.argv[1:] == ["--worker"]: + worker_main() + else: + raise SystemExit(main()) diff --git a/benchmarks/run_pyscf_ao_screening_rotation_benchmark.py b/benchmarks/run_pyscf_ao_screening_rotation_benchmark.py new file mode 100644 index 00000000..953c5f16 --- /dev/null +++ b/benchmarks/run_pyscf_ao_screening_rotation_benchmark.py @@ -0,0 +1,409 @@ +"""Benchmark AO screening across rotations of one approximately 900-AO molecule. + +The default grid applies ``Rz(azimuth) @ Ry(polar)`` to C7H16 (879 AOs with +def2-qzvpp). Azimuth runs from 0 through 330 degrees and polar angle runs from +0 through 150 degrees, both in 30-degree steps, for 72 orientations per mode. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import math +import re +import sys +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Any + +import run_pyscf_ao_screening_benchmark as benchmark + +MODES = ("gpu", "cpu_dense", "cpu_screened") +MEASUREMENTS = ("runtime", "memory") +BASE_MOLECULE = benchmark.make_molecule_spec(7) + + +@dataclass(frozen=True) +class Orientation: + azimuth_degrees: int + polar_degrees: int + + @property + def key(self) -> str: + return f"azimuth_{self.azimuth_degrees:03d}_polar_{self.polar_degrees:03d}" + + def as_json(self) -> dict[str, int]: + return { + "azimuth_degrees": self.azimuth_degrees, + "polar_degrees": self.polar_degrees, + } + + +@dataclass(frozen=True) +class RotationBenchmarkConfig(benchmark.BenchmarkConfig): + azimuth_step_degrees: int = 30 + polar_step_degrees: int = 30 + + @property + def azimuth_angles(self) -> tuple[int, ...]: + return tuple(range(0, 360, self.azimuth_step_degrees)) + + @property + def polar_angles(self) -> tuple[int, ...]: + return tuple(range(0, 180, self.polar_step_degrees)) + + @property + def full_orientations(self) -> tuple[Orientation, ...]: + return tuple( + Orientation(azimuth, polar) + for azimuth in self.azimuth_angles + for polar in self.polar_angles + ) + + @property + def orientations(self) -> tuple[Orientation, ...]: + orientations = self.full_orientations + return orientations[:1] if self.smoke_run else orientations + + def as_json(self) -> dict[str, Any]: + data = asdict(self) + data["source_root"] = str(self.source_root) + data["results_dir"] = str(self.results_dir) + data["azimuth_angles_degrees"] = list(self.azimuth_angles) + data["polar_angles_degrees"] = list(self.polar_angles) + data["orientation_count"] = len(self.orientations) + data["full_orientation_count"] = len(self.full_orientations) + data["modes"] = list(MODES) + data["measurements"] = list(MEASUREMENTS) + data["worker_thread_environment"] = self.worker_thread_environment + return data + + +def rotate_atoms( + atoms: tuple[benchmark.Atom, ...], orientation: Orientation +) -> tuple[benchmark.Atom, ...]: + azimuth = math.radians(orientation.azimuth_degrees) + polar = math.radians(orientation.polar_degrees) + cos_azimuth = math.cos(azimuth) + sin_azimuth = math.sin(azimuth) + cos_polar = math.cos(polar) + sin_polar = math.sin(polar) + + rotated: list[benchmark.Atom] = [] + for element, x, y, z in atoms: + polar_x = cos_polar * x + sin_polar * z + polar_z = -sin_polar * x + cos_polar * z + rotated.append( + ( + element, + cos_azimuth * polar_x - sin_azimuth * y, + sin_azimuth * polar_x + cos_azimuth * y, + polar_z, + ) + ) + return tuple(rotated) + + +def rotated_molecule(orientation: Orientation) -> benchmark.MoleculeSpec: + return benchmark.MoleculeSpec( + carbon_count=BASE_MOLECULE.carbon_count, + expected_aos=BASE_MOLECULE.expected_aos, + formula=BASE_MOLECULE.formula, + atoms=rotate_atoms(BASE_MOLECULE.atoms, orientation), + ) + + +def runner_hashes() -> dict[str, str]: + return { + "rotation_runner_sha256": hashlib.sha256( + Path(__file__).read_bytes() + ).hexdigest(), + "worker_sha256": benchmark.runner_sha256(), + } + + +def validate_rotation_grid(config: RotationBenchmarkConfig) -> None: + from pyscf import gto + + molecules = tuple(rotated_molecule(item) for item in config.full_orientations) + coordinate_hashes = {molecule.coordinate_sha256 for molecule in molecules} + if len(coordinate_hashes) != len(molecules): + raise ValueError("The rotation grid produced duplicate coordinate sets") + + for molecule in molecules: + for original, rotated in zip(BASE_MOLECULE.atoms, molecule.atoms, strict=True): + if original[0] != rotated[0] or not math.isclose( + math.dist((0.0, 0.0, 0.0), original[1:]), + math.dist((0.0, 0.0, 0.0), rotated[1:]), + abs_tol=1e-12, + ): + raise ValueError("A rotation changed the molecular geometry") + + mol = gto.M( + atom=BASE_MOLECULE.atom_text, + basis=config.basis, + charge=0, + spin=0, + unit="Angstrom", + cart=False, + verbose=0, + ) + actual_aos = int(mol.nao_nr()) + if actual_aos != BASE_MOLECULE.expected_aos: + raise ValueError( + f"Expected {BASE_MOLECULE.expected_aos} AOs for {BASE_MOLECULE.formula} " + f"with {config.basis}, got {actual_aos}" + ) + + +def run_worker_preflight(config: RotationBenchmarkConfig) -> dict[str, Any]: + result = benchmark.execute_worker( + {"operation": "environment", "source_root": str(config.source_root)}, config + ) + if result["status"] != "ok": + raise RuntimeError(f"Benchmark preflight failed: {result}") + environment = result["environment"] + imported_path = Path(environment["imported_skala"]) + imported_path.relative_to(config.source_root / "src") + required_packages = ("skala", "pyscf", "torch", "memray") + missing = [ + name for name in required_packages if not environment["packages"].get(name) + ] + if missing: + raise RuntimeError(f"Worker environment is missing packages: {missing}") + return environment + + +def result_path(config: RotationBenchmarkConfig, source: dict[str, Any]) -> Path: + safe_label = re.sub(r"[^A-Za-z0-9_.-]+", "-", config.run_label).strip("-.") + if not safe_label: + raise ValueError("The run label must contain a filename-safe character") + return ( + config.results_dir + / f"skala-pyscf-ao-screening-rotations-{safe_label}-{source['commit'][:12]}.json" + ) + + +def new_result_document( + config: RotationBenchmarkConfig, + source: dict[str, Any], + environment: dict[str, Any], +) -> dict[str, Any]: + created_at = benchmark.utc_now() + return { + "schema_version": 1, + "benchmark": "pyscf_ao_screening_rotations", + "created_at": created_at, + "updated_at": created_at, + "run_label": config.run_label, + "source": source, + "environment": environment, + "configuration": config.as_json(), + "geometry": { + **benchmark.GEOMETRY_PARAMETERS, + "base_molecule": BASE_MOLECULE.as_json(), + "rotation_convention": "active Cartesian rotation Rz(azimuth) @ Ry(polar)", + }, + "runner_hashes": runner_hashes(), + "orientations": { + orientation.key: { + "index": index, + **orientation.as_json(), + "coordinate_sha256": rotated_molecule(orientation).coordinate_sha256, + "observed": None, + "modes": {mode: {} for mode in MODES}, + } + for index, orientation in enumerate(config.orientations) + }, + } + + +def validate_resume_document( + document: dict[str, Any], + config: RotationBenchmarkConfig, + source: dict[str, Any], +) -> None: + if document.get("schema_version") != 1: + raise ValueError("Cannot resume a result file with a different schema version") + if document.get("runner_hashes") != runner_hashes(): + raise ValueError( + "Cannot resume results created by different runner implementations" + ) + if document.get("source", {}).get("commit") != source["commit"]: + raise ValueError("Cannot resume results from a different Git commit") + if document.get("configuration") != config.as_json(): + raise ValueError( + "Cannot resume results created with a different benchmark configuration" + ) + + +def run_benchmark(config: RotationBenchmarkConfig, environment: dict[str, Any]) -> Path: + source = benchmark.source_metadata(config.source_root) + output_path = result_path(config, source) + if output_path.exists(): + document = json.loads(output_path.read_text(encoding="utf-8")) + validate_resume_document(document, config, source) + else: + document = new_result_document(config, source, environment) + benchmark.atomic_write_json(output_path, document) + + cuda_available = bool(document["environment"]["cuda"]["available"]) + for mode in MODES: + for measurement in MEASUREMENTS: + blocked_by: dict[str, Any] | None = None + for orientation in config.orientations: + orientation_record = document["orientations"][orientation.key] + existing = orientation_record["modes"][mode].get(measurement) + if existing and existing.get("status") in benchmark.TERMINAL_STATUSES: + if existing["status"] in {"oom", "timeout"}: + blocked_by = { + "orientation": orientation.key, + "status": existing["status"], + } + continue + + result: dict[str, Any] + if mode == "gpu" and not cuda_available: + result = { + "status": "error", + "mode": mode, + "measurement": measurement, + "error": "CUDA is not available in the worker environment", + } + elif blocked_by is not None: + result = { + "status": "skipped_after_resource_failure", + "mode": mode, + "measurement": measurement, + "blocked_by": blocked_by, + } + else: + molecule = rotated_molecule(orientation) + payload = benchmark.worker_payload( + config, molecule, mode, measurement + ) + payload["orientation"] = orientation.as_json() + result = benchmark.execute_measurement(payload, config) + + benchmark.merge_worker_result( + orientation_record, mode, measurement, result + ) + document["updated_at"] = benchmark.utc_now() + benchmark.atomic_write_json(output_path, document) + if result["status"] in {"oom", "timeout"}: + blocked_by = { + "orientation": orientation.key, + "status": result["status"], + } + print( + f"{mode:12s} {measurement:7s} {orientation.key} {result['status']}" + ) + return output_path + + +def parse_arguments(argv: list[str]) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=( + "Benchmark Skala XC/Vxc evaluation for 72 rotations of one 879-AO " + "molecule on GPU, dense CPU, and screened CPU." + ), + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + parser.add_argument( + "--label", + required=True, + help="Result label, normally the revision name such as 'mr' or 'main'.", + ) + parser.add_argument( + "--source-root", + type=Path, + default=benchmark.RUNNER_ROOT, + help="Skala checkout whose src/skala package is benchmarked.", + ) + parser.add_argument( + "--results-dir", + type=Path, + default=benchmark.RUNNER_ROOT / "benchmarks" / "results", + help="Directory for commit-labelled JSON output.", + ) + parser.add_argument("--functional", default="skala-1.1") + parser.add_argument("--basis", default="def2-qzvpp") + parser.add_argument("--grid-level", type=int, default=1) + parser.add_argument("--max-memory-mb", type=int, default=2000) + parser.add_argument("--threads", type=int, default=4) + parser.add_argument("--runtime-repetitions", type=int, default=3) + parser.add_argument("--timeout-minutes", type=float, default=30.0) + parser.add_argument("--azimuth-step-degrees", type=int, default=30) + parser.add_argument("--polar-step-degrees", type=int, default=30) + parser.add_argument( + "--smoke", + action="store_true", + help="Run only the unrotated orientation for each mode.", + ) + parser.add_argument( + "--preflight-only", + action="store_true", + help="Validate rotations, dependencies, source import, and CUDA without measurements.", + ) + return parser.parse_args(argv) + + +def config_from_arguments( + arguments: argparse.Namespace, +) -> RotationBenchmarkConfig: + if arguments.timeout_minutes <= 0: + raise ValueError("--timeout-minutes must be positive") + if arguments.threads <= 0: + raise ValueError("--threads must be positive") + if arguments.runtime_repetitions <= 0: + raise ValueError("--runtime-repetitions must be positive") + if arguments.azimuth_step_degrees <= 0 or 360 % arguments.azimuth_step_degrees: + raise ValueError("--azimuth-step-degrees must be a positive divisor of 360") + if arguments.polar_step_degrees <= 0 or 180 % arguments.polar_step_degrees: + raise ValueError("--polar-step-degrees must be a positive divisor of 180") + source_root = arguments.source_root.expanduser().resolve() + if not (source_root / "src" / "skala").is_dir(): + raise FileNotFoundError(f"No src/skala package below {source_root}") + return RotationBenchmarkConfig( + source_root=source_root, + results_dir=arguments.results_dir.expanduser().resolve(), + run_label=arguments.label, + functional=arguments.functional, + basis=arguments.basis, + grid_level=arguments.grid_level, + max_memory_mb=arguments.max_memory_mb, + cpu_threads=arguments.threads, + runtime_repetitions=arguments.runtime_repetitions, + worker_timeout_seconds=round(arguments.timeout_minutes * 60), + smoke_run=arguments.smoke, + azimuth_step_degrees=arguments.azimuth_step_degrees, + polar_step_degrees=arguments.polar_step_degrees, + ) + + +def main(argv: list[str] | None = None) -> int: + arguments = parse_arguments(sys.argv[1:] if argv is None else argv) + config = config_from_arguments(arguments) + validate_rotation_grid(config) + environment = run_worker_preflight(config) + print("Benchmark configuration:") + print(json.dumps(config.as_json(), indent=2, sort_keys=True)) + print( + f"Molecule: {BASE_MOLECULE.formula}, " + f"{BASE_MOLECULE.expected_aos} AOs with {config.basis}" + ) + print(f"Skala import: {environment['imported_skala']}") + print(f"Python: {environment['python_executable']}") + print(f"CUDA: {environment['cuda']}") + if arguments.preflight_only: + print("Preflight passed.") + return 0 + + output_path = run_benchmark(config, environment) + print(f"Results written to {output_path}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmarks/vxc_accuracy_grid_grouping.ipynb b/benchmarks/vxc_accuracy_grid_grouping.ipynb new file mode 100644 index 00000000..bf9c8485 --- /dev/null +++ b/benchmarks/vxc_accuracy_grid_grouping.ipynb @@ -0,0 +1,703 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "201eea25", + "metadata": {}, + "source": [ + "# GPU $V_{xc}$ accuracy versus grid grouping\n", + "\n", + "This notebook uses the carbon-chain/def2-QZVPP stress case from `test_gpu_screened_skala_matches_cpu_on_carbon_chain` to isolate how grouping grid points into GPU4PySCF screening blocks affects Skala's integrated $V_{xc}$.\n", + "\n", + "Grid levels 1 and 2 are evaluated independently, each against a dense CPU calculation on the identical grid and density matrix. Every screened candidate executes GPU4PySCF's real CUDA mask construction and AO evaluation with its installed $10^{-10}$ AO threshold and 4096-point block size." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3ce0ae4d", + "metadata": {}, + "outputs": [], + "source": [ + "from dataclasses import dataclass\n", + "from typing import Any\n", + "from unittest.mock import patch\n", + "\n", + "import cupy\n", + "import numpy as np\n", + "import torch\n", + "from pyscf import dft, gto\n", + "from tests.utils import patch_ao_screening\n", + "\n", + "from skala.functional import load_functional\n", + "from skala.functional.base import ExcFunctionalBase\n", + "from skala.pyscf import features as features_module\n", + "from skala.pyscf.backend import dft_gpu\n", + "from skala.pyscf.features import _spatial_grid_permutations\n", + "from skala.pyscf.numint import SkalaNumInt\n", + "\n", + "np.set_printoptions(precision=4, suppress=True)\n", + "\n", + "\n", + "@dataclass(frozen=True)\n", + "class Evaluation:\n", + " electron_count: float\n", + " xc_energy: float\n", + " vxc: np.ndarray\n", + " active_ao_counts: np.ndarray\n", + "\n", + "\n", + "@dataclass(frozen=True)\n", + "class GridExperiment:\n", + " level: int\n", + " coords: np.ndarrays\n", + " dense_reference: Evaluation\n", + " groupings: dict[str, list[np.ndarray]]\n", + " rows: list[dict[str, object]]\n", + "\n", + "\n", + "CARBON_CHAIN = \"\"\"\n", + "C 0.0 0.0 0.0\n", + "C 1.4 0.0 0.0\n", + "C 2.8 0.0 0.0\n", + "C 4.2 0.0 0.0\n", + "C 5.6 0.0 0.0\n", + "C 7.0 0.0 0.0\n", + "\"\"\"\n", + "GRID_LEVELS = (1, 2)\n", + "\n", + "assert torch.cuda.is_available()\n", + "assert dft_gpu is not None\n", + "mol = gto.M(atom=CARBON_CHAIN, basis=\"def2-qzvpp\", verbose=0)\n", + "dm = dft.RKS(mol).get_init_guess()\n", + "\n", + "cpu_functional = load_functional(\"skala-1.1\", device=torch.device(\"cpu\"))\n", + "gpu_functional = load_functional(\"skala-1.1\", device=torch.device(\"cuda:0\"))\n", + "assert isinstance(cpu_functional, ExcFunctionalBase)\n", + "assert isinstance(gpu_functional, ExcFunctionalBase)\n", + "cpu_numint = SkalaNumInt(cpu_functional, device=torch.device(\"cpu\"))\n", + "gpu_numint = SkalaNumInt(gpu_functional, device=torch.device(\"cuda:0\"))\n", + "GPU_BLOCK_SIZE = int(dft_gpu.numint.MIN_BLK_SIZE)\n", + "\n", + "print(f\"Atoms / AOs / shells: {mol.natm} / {mol.nao_nr()} / {mol.nbas}\")\n", + "print(f\"Grid levels: {GRID_LEVELS}\")\n", + "print(f\"GPU4PySCF AO threshold: {dft_gpu.numint.AO_THRESHOLD:.1e}\")\n", + "print(f\"GPU screening block size: {GPU_BLOCK_SIZE}\")" + ] + }, + { + "cell_type": "markdown", + "id": "c5410a08", + "metadata": {}, + "source": [ + "## Screening permutations\n", + "\n", + "Each algorithm partitions the complete grid required by that setting: level 1 has 31,080 points in 8 physical GPU groups, while level 2 has 67,248 points in 17 groups. Every group contains at most 4096 points.\n", + "\n", + "The atom-major case preserves PySCF's original grid order and divides the complete sequence into consecutive blocks. The production spatial case recursively partitions the complete coordinate set, choosing split directions from the spatial extent. The mixed cases exchange points between complete spatial groups while preserving every group size and leaving the final partial group intact. In every case, each grid point appears exactly once.\n", + "\n", + "Even spatially grouped 4096-point blocks can overlap many functions in a diffuse def2-QZVPP basis. `Active AO fraction` reports the grid-point-weighted fraction of AOs retained by the actual masks. `DM matmul proxy` weights the squared fraction, matching the leading $n_{\\mathrm{active}}^2 n_{\\mathrm{grid}}$ scaling of Skala's density-feature matrix multiplications; neither column is a measured runtime." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "fa8cb8c1", + "metadata": {}, + "outputs": [], + "source": [ + "def validate_partition(groups: list[np.ndarray], point_count: int) -> None:\n", + " if not groups or any(group.size == 0 for group in groups):\n", + " raise ValueError(\"Groups must be non-empty.\")\n", + " flattened = np.concatenate(groups)\n", + " if flattened.size != point_count or not np.array_equal(\n", + " np.sort(flattened), np.arange(point_count, dtype=np.int64)\n", + " ):\n", + " raise ValueError(\"Groups must partition every grid point exactly once.\")\n", + "\n", + "\n", + "def groups_from_permutation(\n", + " permutation: np.ndarray, block_size: int\n", + ") -> list[np.ndarray]:\n", + " groups = [\n", + " permutation[start : start + block_size].copy()\n", + " for start in range(0, permutation.size, block_size)\n", + " ]\n", + " validate_partition(groups, permutation.size)\n", + " return groups\n", + "\n", + "\n", + "def mix_spatial_groups(\n", + " groups: list[np.ndarray], mixing_fraction: float, point_count: int\n", + ") -> list[np.ndarray]:\n", + " if not 0 <= mixing_fraction <= 1:\n", + " raise ValueError(\"mixing_fraction must be between zero and one.\")\n", + "\n", + " complete = [group for group in groups if group.size == GPU_BLOCK_SIZE]\n", + " remainders = [group.copy() for group in groups if group.size != GPU_BLOCK_SIZE]\n", + " if mixing_fraction == 0 or len(complete) < 2:\n", + " return [group.copy() for group in groups]\n", + "\n", + " source = np.stack(complete)\n", + " mixed = source.copy()\n", + " mixed_columns = round(mixing_fraction * GPU_BLOCK_SIZE)\n", + " columns = np.floor(\n", + " np.arange(mixed_columns) * GPU_BLOCK_SIZE / mixed_columns\n", + " ).astype(np.int64)\n", + " for column_index, column in enumerate(columns):\n", + " shift = 1 + column_index % (source.shape[0] - 1)\n", + " mixed[:, column] = np.roll(source[:, column], shift)\n", + "\n", + " result = [row.copy() for row in mixed] + remainders\n", + " validate_partition(result, point_count)\n", + " return result\n", + "\n", + "\n", + "def build_matching_grids(level: int) -> tuple[Any, Any, np.ndarray]:\n", + " cpu_grids = dft.Grids(mol)\n", + " cpu_grids.level = level\n", + " cpu_grids.alignment = 1\n", + " cpu_grids.build(sort_grids=False)\n", + " assert cpu_grids.coords is not None and cpu_grids.weights is not None\n", + "\n", + " gpu_grids = dft_gpu.Grids(mol)\n", + " gpu_grids.level = level\n", + " gpu_grids.alignment = 1\n", + " gpu_grids.build(sort_grids=False)\n", + "\n", + " coords = np.asarray(cpu_grids.coords)\n", + " np.testing.assert_allclose(\n", + " coords, cupy.asnumpy(gpu_grids.coords), rtol=0.0, atol=0.0\n", + " )\n", + " np.testing.assert_allclose(\n", + " cpu_grids.weights,\n", + " cupy.asnumpy(gpu_grids.weights),\n", + " rtol=1e-12,\n", + " atol=1e-12,\n", + " )\n", + " return cpu_grids, gpu_grids, coords\n", + "\n", + "\n", + "def dense_cpu_reference(cpu_grids: Any) -> Evaluation:\n", + " with patch_ao_screening(False):\n", + " electron_count, xc_energy, vxc = cpu_numint.nr_rks(mol, cpu_grids, None, dm)\n", + " return Evaluation(\n", + " electron_count=float(electron_count),\n", + " xc_energy=float(xc_energy),\n", + " vxc=np.asarray(vxc),\n", + " active_ao_counts=np.asarray([mol.nao_nr()], dtype=np.int64),\n", + " )\n", + "\n", + "\n", + "def fresh_gpu_grids(template: Any) -> Any:\n", + " case_grids = dft_gpu.Grids(mol)\n", + " case_grids.level = template.level\n", + " case_grids.alignment = template.alignment\n", + " case_grids.coords = template.coords\n", + " case_grids.weights = template.weights\n", + " case_grids._non0ao_idx = None\n", + " return case_grids\n", + "\n", + "\n", + "def evaluate_gpu_permutation(\n", + " permutation: np.ndarray, gpu_grid_template: Any, point_count: int\n", + ") -> Evaluation:\n", + " inverse = np.empty_like(permutation)\n", + " inverse[permutation] = np.arange(point_count, dtype=np.int64)\n", + " case_grids = fresh_gpu_grids(gpu_grid_template)\n", + " with (\n", + " patch.object(\n", + " features_module,\n", + " \"_spatial_grid_permutations\",\n", + " return_value=(permutation, inverse),\n", + " ),\n", + " patch_ao_screening(True),\n", + " ):\n", + " electron_count, xc_energy, vxc = gpu_numint.nr_rks(\n", + " mol, case_grids, None, cupy.asarray(dm)\n", + " )\n", + "\n", + " prepared_grids, cached_forward, _ = features_module._prepare_spatially_sorted_grids(\n", + " mol, case_grids, GPU_BLOCK_SIZE, gpu=True\n", + " )\n", + " assert np.array_equal(cached_forward, permutation)\n", + " active_ao_counts = np.asarray(\n", + " [entry[1].size for entry in prepared_grids.get_non0ao_idx()],\n", + " dtype=np.int64,\n", + " )\n", + " return Evaluation(\n", + " electron_count=float(electron_count),\n", + " xc_energy=float(xc_energy),\n", + " vxc=cupy.asnumpy(vxc),\n", + " active_ao_counts=active_ao_counts,\n", + " )\n", + "\n", + "\n", + "def vxc_errors(\n", + " candidate: np.ndarray, dense_reference: Evaluation\n", + ") -> tuple[float, float]:\n", + " difference = candidate - dense_reference.vxc\n", + " return (\n", + " float(np.max(np.abs(difference))),\n", + " float(np.linalg.norm(difference) / np.linalg.norm(dense_reference.vxc)),\n", + " )\n", + "\n", + "\n", + "def normalized_within_group_radius(\n", + " groups: list[np.ndarray], coords: np.ndarray\n", + ") -> float:\n", + " global_center = coords.mean(axis=0)\n", + " global_rms = np.sqrt(np.mean(np.sum(np.square(coords - global_center), axis=1)))\n", + " within_sum = 0.0\n", + " for group in groups:\n", + " group_coords = coords[group]\n", + " center = group_coords.mean(axis=0)\n", + " within_sum += float(np.sum(np.square(group_coords - center)))\n", + " return float(np.sqrt(within_sum / coords.shape[0]) / global_rms)\n", + "\n", + "\n", + "def mean_maximum_bbox_iou(groups: list[np.ndarray], coords: np.ndarray) -> float:\n", + " if len(groups) == 1:\n", + " return 0.0\n", + " minimums = np.asarray([coords[group].min(axis=0) for group in groups])\n", + " maximums = np.asarray([coords[group].max(axis=0) for group in groups])\n", + " volumes = np.prod(np.maximum(maximums - minimums, 0.0), axis=1)\n", + " maximum_ious = []\n", + " for index in range(len(groups)):\n", + " intersection_extent = np.maximum(\n", + " np.minimum(maximums[index], maximums)\n", + " - np.maximum(minimums[index], minimums),\n", + " 0.0,\n", + " )\n", + " intersection = np.prod(intersection_extent, axis=1)\n", + " union = volumes[index] + volumes - intersection\n", + " iou = np.divide(\n", + " intersection,\n", + " union,\n", + " out=np.zeros_like(intersection),\n", + " where=union > 0,\n", + " )\n", + " iou[index] = 0.0\n", + " maximum_ious.append(float(iou.max()))\n", + " return float(np.mean(maximum_ious))\n", + "\n", + "\n", + "def summarize_case(\n", + " level: int,\n", + " name: str,\n", + " groups: list[np.ndarray],\n", + " evaluation: Evaluation,\n", + " coords: np.ndarray,\n", + " dense_reference: Evaluation,\n", + ") -> dict[str, object]:\n", + " point_count = coords.shape[0]\n", + " validate_partition(groups, point_count)\n", + " sizes = np.asarray([group.size for group in groups], dtype=np.int64)\n", + " active_aos = evaluation.active_ao_counts\n", + " assert sizes.size == active_aos.size\n", + " maximum_error, relative_error = vxc_errors(evaluation.vxc, dense_reference)\n", + " return {\n", + " \"level\": level,\n", + " \"grid_points\": point_count,\n", + " \"case\": name,\n", + " \"groups\": len(groups),\n", + " \"occupancy\": f\"{sizes.min()}/{np.median(sizes):.0f}/{sizes.max()}\",\n", + " \"active_aos\": f\"{active_aos.min()}/{np.median(active_aos):.0f}/{active_aos.max()}\",\n", + " \"radius\": normalized_within_group_radius(groups, coords),\n", + " \"bbox_iou\": mean_maximum_bbox_iou(groups, coords),\n", + " \"active_ao_fraction\": float(\n", + " np.sum(sizes * active_aos) / (point_count * mol.nao_nr())\n", + " ),\n", + " \"dm_matmul_proxy\": float(\n", + " np.sum(sizes * np.square(active_aos)) / (point_count * mol.nao_nr() ** 2)\n", + " ),\n", + " \"max_vxc_error\": maximum_error,\n", + " \"relative_vxc_error\": relative_error,\n", + " \"electron_error\": abs(\n", + " evaluation.electron_count - dense_reference.electron_count\n", + " ),\n", + " \"energy_error\": abs(evaluation.xc_energy - dense_reference.xc_energy),\n", + " }\n", + "\n", + "\n", + "def build_groupings(coords: np.ndarray) -> dict[str, list[np.ndarray]]:\n", + " point_count = coords.shape[0]\n", + " atom_major = groups_from_permutation(\n", + " np.arange(point_count, dtype=np.int64), GPU_BLOCK_SIZE\n", + " )\n", + " spatial_forward, _ = _spatial_grid_permutations(coords, GPU_BLOCK_SIZE)\n", + " spatial = groups_from_permutation(spatial_forward, GPU_BLOCK_SIZE)\n", + " return {\n", + " \"GPU4PySCF atom-major blocks\": atom_major,\n", + " \"GPU4PySCF spatial blocks\": spatial,\n", + " \"GPU4PySCF spatial, mix 0.500\": mix_spatial_groups(spatial, 0.5, point_count),\n", + " \"GPU4PySCF spatial, mix 1.000\": mix_spatial_groups(spatial, 1.0, point_count),\n", + " }\n", + "\n", + "\n", + "def run_grid_level(level: int) -> GridExperiment:\n", + " cpu_grids, gpu_grid_template, coords = build_matching_grids(level)\n", + " point_count = coords.shape[0]\n", + " dense_reference = dense_cpu_reference(cpu_grids)\n", + " groupings = build_groupings(coords)\n", + " all_points = [np.arange(point_count, dtype=np.int64)]\n", + " case_data = [(\"Dense CPU reference\", all_points, dense_reference)]\n", + "\n", + " print(f\"Level {level}: {point_count:,} points\")\n", + " for name, groups in groupings.items():\n", + " print(f\" Evaluating {name}...\")\n", + " evaluation = evaluate_gpu_permutation(\n", + " np.concatenate(groups), gpu_grid_template, point_count\n", + " )\n", + " assert np.isfinite(evaluation.electron_count)\n", + " assert np.isfinite(evaluation.xc_energy)\n", + " assert np.isfinite(evaluation.vxc).all()\n", + " assert np.allclose(evaluation.vxc, evaluation.vxc.T, rtol=1e-10, atol=1e-11)\n", + " case_data.append((name, groups, evaluation))\n", + "\n", + " rows = [\n", + " summarize_case(level, name, groups, evaluation, coords, dense_reference)\n", + " for name, groups, evaluation in case_data\n", + " ]\n", + " atom_major_error = rows[1][\"max_vxc_error\"]\n", + " spatial_error = rows[2][\"max_vxc_error\"]\n", + " assert isinstance(atom_major_error, float)\n", + " assert isinstance(spatial_error, float)\n", + " print(\n", + " \" Spatial grouping changes max |dVxc| by \"\n", + " f\"{atom_major_error / spatial_error:.2f}x.\"\n", + " )\n", + " return GridExperiment(level, coords, dense_reference, groupings, rows)\n", + "\n", + "\n", + "def render_results(rows: list[dict[str, object]]) -> str:\n", + " columns = (\n", + " (\"level\", \"Grid level\"),\n", + " (\"grid_points\", \"Grid points\"),\n", + " (\"case\", \"Case\"),\n", + " (\"groups\", \"Groups\"),\n", + " (\"occupancy\", \"Points min/med/max\"),\n", + " (\"active_aos\", \"Active AOs min/med/max\"),\n", + " (\"radius\", \"RMS radius\"),\n", + " (\"bbox_iou\", \"BBox IoU\"),\n", + " (\"active_ao_fraction\", \"Active AO fraction\"),\n", + " (\"dm_matmul_proxy\", \"DM matmul proxy\"),\n", + " (\"max_vxc_error\", \"max |dVxc|\"),\n", + " (\"relative_vxc_error\", \"rel. Frobenius\"),\n", + " (\"electron_error\", \"|dN|\"),\n", + " (\"energy_error\", \"|dExc|\"),\n", + " )\n", + " parts = [\n", + " '',\n", + " \"\",\n", + " ]\n", + " parts.extend(\n", + " f''\n", + " for _, label in columns\n", + " )\n", + " parts.append(\"\")\n", + " for row in rows:\n", + " parts.append(\"\")\n", + " for key, _ in columns:\n", + " value = row[key]\n", + " text = f\"{value:.3e}\" if isinstance(value, float) else str(value)\n", + " parts.append(\n", + " f''\n", + " )\n", + " parts.append(\"\")\n", + " parts.append(\"
{label}
{text}
\")\n", + " return \"\".join(parts)\n", + "\n", + "\n", + "class HTMLTable(str):\n", + " def _repr_html_(self) -> str:\n", + " return str(self)" + ] + }, + { + "cell_type": "markdown", + "id": "9aa97ac5", + "metadata": {}, + "source": [ + "## Results\n", + "\n", + "Each grid level has its own dense CPU reference and independently constructed grouping permutations. All screened rows use actual GPU4PySCF masks. Lower active-AO metrics mean more aggressive screening; lower error means closer agreement with that level's dense reference." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9ab116d3", + "metadata": {}, + "outputs": [], + "source": [ + "experiments = [run_grid_level(level) for level in GRID_LEVELS]\n", + "results = [row for experiment in experiments for row in experiment.rows]\n", + "\n", + "HTMLTable(render_results(results))" + ] + }, + { + "cell_type": "markdown", + "id": "4e6acd97", + "metadata": {}, + "source": [ + "## Fixed grid-group slices\n", + "\n", + "For each grid level, the figure shows three fixed slabs centered at $z=-1$, $0$, and $+1$ bohr relative to the molecular $x$-$y$ plane. Each slab includes points satisfying $|z-z_0|\\leq 0.25$ bohr. The outermost 1% of each grid, ranked by three-dimensional distance to the nearest carbon nucleus, is omitted to keep the molecular region legible.\n", + "\n", + "The selected points are accumulated in shared $x$-$y$ bins. Each occupied bin takes the color of its most frequent screening group; the color is blended toward white according to that group's fraction of points in the bin. Pure color means complete local agreement, while a pale bin contains a stronger mixture of groups. Black crosses mark the carbon nuclei." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e47fa055", + "metadata": {}, + "outputs": [], + "source": [ + "import matplotlib.pyplot as plt\n", + "from matplotlib.colors import BoundaryNorm, ListedColormap\n", + "\n", + "\n", + "def group_labels(groups: list[np.ndarray], point_count: int) -> np.ndarray:\n", + " labels = np.empty(point_count, dtype=np.int64)\n", + " for group_index, group in enumerate(groups):\n", + " labels[group] = group_index\n", + " return labels\n", + "\n", + "\n", + "def dominant_group_image(\n", + " coords: np.ndarray,\n", + " labels: np.ndarray,\n", + " selected: np.ndarray,\n", + " x_edges: np.ndarray,\n", + " y_edges: np.ndarray,\n", + " group_colors: np.ndarray,\n", + ") -> np.ndarray:\n", + " x_bin_count = x_edges.size - 1\n", + " y_bin_count = y_edges.size - 1\n", + " selected_coords = coords[selected]\n", + " x_bins = np.searchsorted(x_edges, selected_coords[:, 0], side=\"right\") - 1\n", + " y_bins = np.searchsorted(y_edges, selected_coords[:, 1], side=\"right\") - 1\n", + " x_bins = np.clip(x_bins, 0, x_bin_count - 1)\n", + " y_bins = np.clip(y_bins, 0, y_bin_count - 1)\n", + " flat_bins = y_bins * x_bin_count + x_bins\n", + " combined = labels[selected] * (x_bin_count * y_bin_count) + flat_bins\n", + " counts = np.bincount(\n", + " combined,\n", + " minlength=group_colors.shape[0] * x_bin_count * y_bin_count,\n", + " ).reshape(group_colors.shape[0], y_bin_count, x_bin_count)\n", + " assert int(counts.sum()) == int(selected.sum())\n", + "\n", + " totals = counts.sum(axis=0)\n", + " dominant_groups = counts.argmax(axis=0)\n", + " dominant_counts = counts.max(axis=0)\n", + " populated = totals > 0\n", + " dominant_fraction = np.divide(\n", + " dominant_counts,\n", + " totals,\n", + " out=np.zeros_like(dominant_counts, dtype=float),\n", + " where=populated,\n", + " )\n", + "\n", + " image = np.ones((*totals.shape, 4), dtype=float)\n", + " dominant_colors = group_colors[dominant_groups]\n", + " image[populated, :3] = 1.0 - dominant_fraction[populated, None] * (\n", + " 1.0 - dominant_colors[populated]\n", + " )\n", + " return image\n", + "\n", + "\n", + "def rgb_to_lab(rgb: np.ndarray) -> np.ndarray:\n", + " linear = np.where(\n", + " rgb <= 0.04045,\n", + " rgb / 12.92,\n", + " ((rgb + 0.055) / 1.055) ** 2.4,\n", + " )\n", + " transform = np.asarray(\n", + " [\n", + " [0.4124564, 0.3575761, 0.1804375],\n", + " [0.2126729, 0.7151522, 0.0721750],\n", + " [0.0193339, 0.1191920, 0.9503041],\n", + " ]\n", + " )\n", + " xyz = linear @ transform.T\n", + " xyz /= np.asarray([0.95047, 1.0, 1.08883])\n", + " delta = 6 / 29\n", + " transformed = np.where(\n", + " xyz > delta**3,\n", + " np.cbrt(xyz),\n", + " xyz / (3 * delta**2) + 4 / 29,\n", + " )\n", + " return np.column_stack(\n", + " (\n", + " 116 * transformed[:, 1] - 16,\n", + " 500 * (transformed[:, 0] - transformed[:, 1]),\n", + " 200 * (transformed[:, 1] - transformed[:, 2]),\n", + " )\n", + " )\n", + "\n", + "\n", + "def distinct_group_colors(count: int) -> np.ndarray:\n", + " levels = np.linspace(0.0, 1.0, 11)\n", + " candidates = np.stack(\n", + " np.meshgrid(levels, levels, levels, indexing=\"ij\"), axis=-1\n", + " ).reshape(-1, 3)\n", + " candidate_lab = rgb_to_lab(candidates)\n", + " chroma = np.linalg.norm(candidate_lab[:, 1:], axis=1)\n", + " keep = (candidate_lab[:, 0] >= 35) & (candidate_lab[:, 0] <= 75) & (chroma >= 35)\n", + " candidates = candidates[keep]\n", + " candidate_lab = candidate_lab[keep]\n", + "\n", + " seed = np.argmin(np.linalg.norm(candidates - np.asarray([0.0, 0.3, 0.8]), axis=1))\n", + " selected = [int(seed)]\n", + " minimum_distance = np.linalg.norm(candidate_lab - candidate_lab[seed], axis=1)\n", + " for _ in range(1, count):\n", + " index = int(np.argmax(minimum_distance))\n", + " selected.append(index)\n", + " distance = np.linalg.norm(candidate_lab - candidate_lab[index], axis=1)\n", + " minimum_distance = np.minimum(minimum_distance, distance)\n", + " return candidates[selected]\n", + "\n", + "\n", + "atom_coords = mol.atom_coords()\n", + "slice_centers = (-1.0, 0.0, 1.0)\n", + "slice_half_width = 0.25\n", + "retained_by_level = {}\n", + "for experiment in experiments:\n", + " nearest_atom_distance = np.linalg.norm(\n", + " experiment.coords[:, None, :] - atom_coords[None, :, :], axis=2\n", + " ).min(axis=1)\n", + " removed_count = round(0.01 * experiment.coords.shape[0])\n", + " retained = np.ones(experiment.coords.shape[0], dtype=bool)\n", + " outside_order = np.argsort(nearest_atom_distance, kind=\"stable\")\n", + " retained[outside_order[-removed_count:]] = False\n", + " assert retained.sum() == experiment.coords.shape[0] - removed_count\n", + " retained_by_level[experiment.level] = retained\n", + "\n", + "trimmed_xy = np.concatenate(\n", + " [\n", + " experiment.coords[retained_by_level[experiment.level], :2]\n", + " for experiment in experiments\n", + " ]\n", + ")\n", + "x_min, y_min = trimmed_xy.min(axis=0)\n", + "x_max, y_max = trimmed_xy.max(axis=0)\n", + "x_bin_count = 120\n", + "bin_width = (x_max - x_min) / x_bin_count\n", + "y_bin_count = max(1, int(np.ceil((y_max - y_min) / bin_width)))\n", + "y_center = 0.5 * (y_min + y_max)\n", + "x_edges = np.linspace(x_min, x_max, x_bin_count + 1)\n", + "y_edges = np.linspace(\n", + " y_center - 0.5 * y_bin_count * bin_width,\n", + " y_center + 0.5 * y_bin_count * bin_width,\n", + " y_bin_count + 1,\n", + ")\n", + "\n", + "for experiment in experiments:\n", + " coords = experiment.coords\n", + " retained = retained_by_level[experiment.level]\n", + " group_count = max(len(groups) for groups in experiment.groupings.values())\n", + " group_colors = distinct_group_colors(group_count)\n", + " palette = ListedColormap(group_colors)\n", + " norm = BoundaryNorm(np.arange(group_count + 1) - 0.5, palette.N)\n", + " labels_by_name = {\n", + " name: group_labels(groups, coords.shape[0])\n", + " for name, groups in experiment.groupings.items()\n", + " }\n", + "\n", + " figure, axes = plt.subplots(\n", + " len(labels_by_name),\n", + " len(slice_centers),\n", + " figsize=(15, 13),\n", + " sharex=True,\n", + " sharey=True,\n", + " constrained_layout=True,\n", + " squeeze=False,\n", + " )\n", + " for row_index, (name, labels) in enumerate(labels_by_name.items()):\n", + " for column_index, height in enumerate(slice_centers):\n", + " axis = axes[row_index, column_index]\n", + " selected = retained & (np.abs(coords[:, 2] - height) <= slice_half_width)\n", + " image = dominant_group_image(\n", + " coords,\n", + " labels,\n", + " selected,\n", + " x_edges,\n", + " y_edges,\n", + " group_colors,\n", + " )\n", + " axis.imshow(\n", + " image,\n", + " origin=\"lower\",\n", + " extent=(x_edges[0], x_edges[-1], y_edges[0], y_edges[-1]),\n", + " interpolation=\"nearest\",\n", + " aspect=\"equal\",\n", + " )\n", + " axis.scatter(\n", + " atom_coords[:, 0],\n", + " atom_coords[:, 1],\n", + " marker=\"x\",\n", + " c=\"black\",\n", + " s=24,\n", + " linewidths=1.0,\n", + " zorder=3,\n", + " )\n", + " axis.text(\n", + " 0.98,\n", + " 0.96,\n", + " f\"{selected.sum():,} points\",\n", + " ha=\"right\",\n", + " va=\"top\",\n", + " transform=axis.transAxes,\n", + " fontsize=8,\n", + " )\n", + " if row_index == 0:\n", + " axis.set_title(\n", + " f\"z = {height:+.1f} +/- {slice_half_width:.2f} bohr\",\n", + " fontsize=10,\n", + " )\n", + " if column_index == 0:\n", + " axis.set_ylabel(f\"{name}\\ny (bohr)\", fontsize=9)\n", + " if row_index == len(labels_by_name) - 1:\n", + " axis.set_xlabel(\"x (bohr)\")\n", + "\n", + " colorbar = figure.colorbar(\n", + " plt.cm.ScalarMappable(norm=norm, cmap=palette),\n", + " ax=axes,\n", + " ticks=np.arange(group_count),\n", + " shrink=0.82,\n", + " pad=0.02,\n", + " )\n", + " colorbar.ax.set_yticklabels(np.arange(1, group_count + 1))\n", + " colorbar.set_label(\"Dominant screening group\")\n", + " figure.suptitle(\n", + " f\"Level {experiment.level}: dominant GPU screening groups \"\n", + " f\"({coords.shape[0]:,} grid points)\"\n", + " )\n", + " plt.show()" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} From 5114d3b0ddc158755cae11d2e3c50bc0524a80a4 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Sat, 8 Aug 2026 21:34:23 +0200 Subject: [PATCH 31/39] improve model chunking --- .../pyscf_ao_screening_performance.ipynb | 55 +++++- .../run_pyscf_ao_screening_benchmark.py | 2 +- src/skala/pyscf/evaluation.py | 4 +- src/skala/pyscf/memory_estimators.py | 167 ++++++++---------- src/skala/pyscf/model_chunking.py | 165 ++++++++++------- src/skala/pyscf/xc_integrator.py | 14 +- tests/test_ao_screening.py | 36 +++- tests/test_evaluation.py | 2 +- tests/test_memory_estimators.py | 61 +++++-- tests/test_model_chunking.py | 113 ++++++++++++ 10 files changed, 413 insertions(+), 206 deletions(-) create mode 100644 tests/test_model_chunking.py diff --git a/benchmarks/pyscf_ao_screening_performance.ipynb b/benchmarks/pyscf_ao_screening_performance.ipynb index 9a6a39a8..6dcb60a7 100644 --- a/benchmarks/pyscf_ao_screening_performance.ipynb +++ b/benchmarks/pyscf_ao_screening_performance.ipynb @@ -28,7 +28,7 @@ "source": [ "## Select Result Files\n", "\n", - "By default, every matching result in `benchmarks/results` is loaded. Replace `SELECTED_RESULT_FILES` with an explicit list when comparing only particular labels or commits." + "By default, every compatible molecule-benchmark result in `benchmarks/results` is loaded. Results with other schemas, such as rotation comparisons, are reported and ignored. Replace `SELECTED_RESULT_FILES` with an explicit list when comparing only particular labels or commits." ] }, { @@ -71,13 +71,30 @@ " raise FileNotFoundError(f\"Could not find the Skala repository above {start}\")\n", "\n", "\n", + "def is_molecule_benchmark_result(path: Path) -> bool:\n", + " document = json.loads(path.read_text(encoding=\"utf-8\"))\n", + " return isinstance(document.get(\"molecules\"), dict)\n", + "\n", + "\n", "REPOSITORY_ROOT = find_repository_root(Path.cwd())\n", "RESULTS_DIR = REPOSITORY_ROOT / \"benchmarks\" / \"results\"\n", - "SELECTED_RESULT_FILES = sorted(RESULTS_DIR.glob(\"skala-pyscf-ao-screening-*.json\"))\n", + "CANDIDATE_RESULT_FILES = sorted(RESULTS_DIR.glob(\"skala-pyscf-ao-screening-*.json\"))\n", + "SELECTED_RESULT_FILES = [\n", + " path\n", + " for path in CANDIDATE_RESULT_FILES\n", + " if is_molecule_benchmark_result(path) and \"screening-screening\" not in path.name\n", + "]\n", + "IGNORED_RESULT_FILES = [\n", + " path for path in CANDIDATE_RESULT_FILES if path not in SELECTED_RESULT_FILES\n", + "]\n", "\n", "print(f\"Selected {len(SELECTED_RESULT_FILES)} result file(s) from {RESULTS_DIR}\")\n", "for result_file in SELECTED_RESULT_FILES:\n", - " print(f\" {result_file.name}\")" + " print(f\" {result_file.name}\")\n", + "if IGNORED_RESULT_FILES:\n", + " print(\"Ignored incompatible result file(s):\")\n", + " for result_file in IGNORED_RESULT_FILES:\n", + " print(f\" {result_file.name}\")" ] }, { @@ -271,6 +288,9 @@ "}\n", "REVISION_LINESTYLES = (\"-\", \"--\", \":\", \"-.\")\n", "RESULT_MARKERS = (\"o\", \"s\", \"^\", \"D\", \"v\", \"P\", \"X\")\n", + "MARKER_SIZE = 5\n", + "LEGEND_MARKER_SCALE = 1.8\n", + "LEGEND_HANDLE_LENGTH = 3.0\n", "\n", "\n", "def measurement_samples(mode_record: dict[str, Any], measurement: str) -> list[float]:\n", @@ -362,7 +382,14 @@ " if logarithmic:\n", " axis.set_yscale(\"log\")\n", " axis.grid(True, which=\"both\", color=\"#D9D9D9\", linewidth=0.6)\n", - " axis.legend(fontsize=8)\n", + " axis.legend(\n", + " fontsize=8,\n", + " loc=\"upper left\",\n", + " bbox_to_anchor=(1.02, 1.0),\n", + " borderaxespad=0.0,\n", + " markerscale=LEGEND_MARKER_SCALE,\n", + " handlelength=LEGEND_HANDLE_LENGTH,\n", + " )\n", "\n", "\n", "def plot_measurement(documents: list[dict[str, Any]], measurement: str) -> None:\n", @@ -383,7 +410,7 @@ " \"color\": MODE_COLORS[mode],\n", " \"linestyle\": line_style,\n", " \"marker\": marker,\n", - " \"markersize\": 4,\n", + " \"markersize\": MARKER_SIZE,\n", " \"label\": endpoint_label(curve_label, x_values),\n", " }\n", " if measurement == \"runtime\":\n", @@ -425,7 +452,7 @@ " \"color\": MODE_COLORS[\"cpu\"],\n", " \"linestyle\": line_style,\n", " \"marker\": marker,\n", - " \"markersize\": 4,\n", + " \"markersize\": MARKER_SIZE,\n", " \"label\": f\"{document['run_label']} cpu\",\n", " }\n", " if measurement == \"runtime\":\n", @@ -651,7 +678,7 @@ " color=MODE_COLORS[production_mode],\n", " linestyle=line_style,\n", " marker=marker,\n", - " markersize=4,\n", + " markersize=MARKER_SIZE,\n", " label=endpoint_label(curve_label, x_values),\n", " )\n", " fingerprint_label = FINGERPRINT_LABELS[fingerprint_key]\n", @@ -697,7 +724,7 @@ " document_index % len(REVISION_LINESTYLES)\n", " ],\n", " marker=RESULT_MARKERS[document_index % len(RESULT_MARKERS)],\n", - " markersize=4,\n", + " markersize=MARKER_SIZE,\n", " label=endpoint_label(curve_label, x_values),\n", " )\n", " axis.axhline(0.0, color=\"#777777\", linewidth=0.8, linestyle=\":\")\n", @@ -721,7 +748,17 @@ " else:\n", " axis.ticklabel_format(axis=\"y\", style=\"sci\", scilimits=(0, 0))\n", " axis.grid(True, which=\"both\", color=\"#D9D9D9\", linewidth=0.6)\n", - " axis.legend(fontsize=7)\n", + " handles, labels = axes.flat[0].get_legend_handles_labels()\n", + " figure.legend(\n", + " handles,\n", + " labels,\n", + " fontsize=8,\n", + " loc=\"center left\",\n", + " bbox_to_anchor=(1.01, 0.5),\n", + " borderaxespad=0.0,\n", + " markerscale=LEGEND_MARKER_SCALE,\n", + " handlelength=LEGEND_HANDLE_LENGTH,\n", + " )\n", " figure.suptitle(\n", " f\"CPU-dense reference differences from {reference_document['run_label']}\"\n", " )\n", diff --git a/benchmarks/run_pyscf_ao_screening_benchmark.py b/benchmarks/run_pyscf_ao_screening_benchmark.py index 381a47f0..749e9ef4 100644 --- a/benchmarks/run_pyscf_ao_screening_benchmark.py +++ b/benchmarks/run_pyscf_ao_screening_benchmark.py @@ -400,7 +400,7 @@ def route_metadata(numint: Any, mol: Any, forced_dense: bool) -> dict[str, Any]: route_decision_callable = route_module._should_screen_aos route_decision = bool(route_decision_callable(mol)) supports_screened_evaluation = bool( - numint.feature_spec.supports_screened_evaluation + route_owner.feature_spec.supports_spatial_decomposition ) route_selector = "ao_threshold" elif "_functional_supports_atom_chunking" in referenced_names and hasattr( diff --git a/src/skala/pyscf/evaluation.py b/src/skala/pyscf/evaluation.py index 042fd915..3c1e7ca6 100644 --- a/src/skala/pyscf/evaluation.py +++ b/src/skala/pyscf/evaluation.py @@ -92,8 +92,8 @@ def requires_atomic_layout(self) -> bool: return bool(self.names & _ATOMIC_LAYOUT_FEATURES) @property - def supports_screened_evaluation(self) -> bool: - """Return whether atom-aligned screened evaluation is supported.""" + def supports_spatial_decomposition(self) -> bool: + """Return whether spatial decomposition is supported.""" return Feature.ATOMIC_GRID_SIZES in self.names diff --git a/src/skala/pyscf/memory_estimators.py b/src/skala/pyscf/memory_estimators.py index e13dfb6a..4794d5ff 100644 --- a/src/skala/pyscf/memory_estimators.py +++ b/src/skala/pyscf/memory_estimators.py @@ -1,39 +1,47 @@ -"""Memory estimators for chunked calculations for Skala 1.1.""" +"""Memory estimators for screened calculations with Skala 1.1.""" import torch +_MODEL_ELEMENTS_PER_GRID_POINT = { + 0: 5830, + 1: 6680, + 2: 24230, +} +_GLOBAL_DENSE_BYTES_PER_AO_SQUARED = { + 0: 36.8, + 1: 37.0, + 2: 9.0, +} -def estimate_max_gridpoint_chunk_size( + +def estimate_max_model_atoms_per_chunk( dm: torch.Tensor, - deriv: int, + atomic_grid_sizes: torch.Tensor, + nfeatures: int, max_memory_in_mb: int | None = None, safety_fraction: float = 0.8, func_deriv: int = 1, - reserved_memory_in_bytes: int = 0, -) -> int: - """Heuristically limit grid points per atom-aligned model evaluation. +) -> dict[int, int]: + """Estimate an atom limit for every homogeneous atomic-grid-size group. - The dominant per-chunk allocation is the atomic-orbital matrix evaluated by - ``evaluate_full_grid`` (shape ``(ncomp, nao, n)`` in float64, with no AO screening), - together with the ``c0``/``ci`` products formed inside the feature function and - retained by autograd for the backward pass. Peak memory is therefore modelled - as affine in the number of grid points ``n`` (see - :func:`linear_peak_memory_model`):: + AO evaluation is completed globally before model chunking starts. Its AO-sized + terms therefore do not scale with each model chunk. Once chunks contain only + equal-sized atomic grids of size ``g``, the model's padded point count equals + its real point count, and its chunk-local peak for ``a`` atoms is modelled as:: - peak_bytes ~= bytes_per_point * n + fixed_overhead + chunk_bytes ~= model_bytes_per_point * g * a - The returned chunk size is the largest ``n`` whose predicted peak fits within - ``safety_fraction`` of the available memory. + For an explicit memory budget, globally live allocations are estimated from + ``atomic_grid_sizes`` and subtracted once. A probed CUDA free-memory value + already excludes allocations currently resident on the device, so the global + footprint is not subtracted a second time in that case. Args: - dm: Density matrix; only its device and trailing dimension are used. - ``dm.shape[-1]`` is taken as ``nao`` and ``dm.device`` selects how - available memory is probed. - deriv: Derivative order of the requested AO features (e.g. ``1`` for - MGGA), which sets the AO component count ``ncomp``. - max_memory_in_mb: Memory budget in **megabytes (MB)** to use on the device on which the density matrix is located. When ``None`` the - budget is probed automatically: free device memory on CUDA, available - physical RAM on CPU. + dm: Density matrix; its device selects how available memory is probed. + atomic_grid_sizes: Number of grid points belonging to each atom. + nfeatures: Number of globally stored raw features per grid point. + max_memory_in_mb: Memory budget in megabytes (MB). When ``None``, free + device memory is probed automatically on CUDA. safety_fraction: Fraction of the budget the predicted peak is allowed to occupy (``0 < safety_fraction <= 1``). Headroom for allocator fragmentation and transient buffers. @@ -41,15 +49,10 @@ def estimate_max_gridpoint_chunk_size( (``exc_only``), ``1`` first order (``__call__``/``V_xc``), ``2`` second order (``gen_response``/Hessian-vector product). Selects the calibrated coefficients. - reserved_memory_in_bytes: Memory already committed to allocations whose - size does not depend on the model chunk, such as global raw-feature - and cotangent buffers. - Returns: - Maximum number of grid points per chunk whose predicted peak memory fits - within ``safety_fraction`` of the budget. May be non-positive when the - ``fixed_overhead`` alone exceeds the budget; callers are expected to - clamp it to at least the largest atomic grid size. + Mapping from each distinct atomic grid size to the maximum number of atoms + of that size per model chunk. Values may be non-positive when the global + footprint exceeds the budget; callers are expected to clamp them to one. Raises: ValueError: If ``safety_fraction`` is outside ``(0, 1]``, or if @@ -59,6 +62,8 @@ def estimate_max_gridpoint_chunk_size( """ if not 0 < safety_fraction <= 1: raise ValueError("safety_fraction must be greater than 0 and at most 1") + if atomic_grid_sizes.numel() == 0 or torch.any(atomic_grid_sizes <= 0): + raise ValueError("atomic_grid_sizes must contain positive values") if max_memory_in_mb is None: match dm.device.type: @@ -77,18 +82,28 @@ def estimate_max_gridpoint_chunk_size( raise ValueError( f"Unsupported device type: {dm.device.type} for memory estimation. Supply max_memory_in_mb explicitly." ) + available_memory = int(free_bytes * safety_fraction) else: free_bytes = int(max_memory_in_mb * 1000**2) - free_bytes = int(free_bytes * safety_fraction) - reserved_memory_in_bytes + available_memory = int(free_bytes * safety_fraction) + available_memory -= estimate_global_screened_buffer_memory( + dm, nfeatures, atomic_grid_sizes, func_deriv + ) + + bytes_per_point = estimate_model_memory_per_grid_point(func_deriv) + return { + grid_size: available_memory // (grid_size * bytes_per_point) + for grid_size in map(int, torch.unique(atomic_grid_sizes).tolist()) + } - bytes_per_point, fixed_overhead = linear_peak_memory_model( - nao=dm.shape[-1], - deriv=deriv, - func_deriv=func_deriv, - ) - chunk_size = int((free_bytes - fixed_overhead) / bytes_per_point) - return chunk_size +def estimate_model_memory_per_grid_point(func_deriv: int) -> int: + """Return calibrated chunk-local Skala memory per homogeneous grid point.""" + try: + elements_per_point = _MODEL_ELEMENTS_PER_GRID_POINT[func_deriv] + except KeyError as error: + raise ValueError("Invalid func_deriv value") from error + return 8 * elements_per_point def estimate_global_raw_feature_buffer_memory( @@ -127,64 +142,20 @@ def estimate_global_raw_feature_buffer_memory( return buffer_count * batch_size * nfeatures * ngrids * 8 -def linear_peak_memory_model( - nao: int, - deriv: int, +def estimate_global_screened_buffer_memory( + dm: torch.Tensor, + nfeatures: int, + atomic_grid_sizes: torch.Tensor, func_deriv: int, -) -> tuple[float, float]: - """ - Return the coefficients of the linear model for peak memory usage in the number of grid points:: - - bytes ~= bytes_per_point * n + fixed_overhead - - Both terms are quadratic in ``nao`` and calibrated *per code path*:: - - bytes_per_point = 8 * (C_AO2*nao^2 + (ncomp + C_LIN)*nao + C_NET) - fixed_overhead = C_FIX * nao^2 - - Details: - The four coefficients are fitted directly for skala-1.1 to the empirical sweep (9999 - measured chunks, nao 38-4452, deriv=1 / MGGA) and then scaled by a single - per-path safety margin so the worst observed meas/pred ratio is 0.90 with - zero breaches. This replaces the earlier single ``autograd_factor`` that - multiplied *both* the nao^2 and the network-constant terms: the data show - the nao^2 (AO-retention) coefficient is almost path-independent (0.0065 / - 0.0067 / 0.0104 B), while only the network-activation constant scales - strongly across energy/first/second order. Decoupling them removes the - ~3x over-padding the old model carried on the second-order path. - """ - # Number of AO components for the requested derivative order. - ncomp = (deriv + 1) * (deriv + 2) * (deriv + 3) // 6 - - # Per-path calibrated coefficients, keyed by func_deriv (0=energy/exc_only, - # 1=first order/__call__, 2=second order/gen_response). Each tuple is - # (C_AO2, C_LIN, C_NET, C_FIX) in float64 elements (C_FIX already in bytes): - # * C_AO2 - nao^2 coefficient of the per-point cost (autograd-retained AO - # intermediates); barely grows with path. - # * C_LIN - retained AO columns beyond the raw ncomp matrix (c0/ci + grads). - # * C_NET - nao-independent enhancement-network activation elements/point; - # this is where the second-order double-backward graph shows up. - # * C_FIX - quadratic coefficient of the dense nao x nao buffers (dm0, dm1, - # hvp_total, Vxc accumulator, get_j), in bytes. - # Fitted on the cc-pVQZ/5Z/6Z + PAH sweep (coronene/cc-pV6Z reaches nao=4452) - # then scaled to worst-case ratio 0.90; tested max meas/pred 0.900, 0 breaches. - match func_deriv: - case 0: - C_AO2, C_LIN, C_NET, C_FIX = 1.07e-3, 5.2, 5830.0, 36.8 - case 1: - C_AO2, C_LIN, C_NET, C_FIX = 1.10e-3, 4.8, 6680.0, 37.0 - case 2: - C_AO2, C_LIN, C_NET, C_FIX = 1.80e-3, 1.7, 24230.0, 9.0 - case _: - raise ValueError("Invalid func_deriv value") - - elems_per_point = ( - C_AO2 * nao * nao # autograd-retained AO intermediates (nao^2) - + (ncomp + C_LIN) * nao # AO matrix + retained feature function memory - + C_NET # network activations (path-dependent) +) -> int: + """Estimate globally live buffers whose lifetimes overlap model chunks.""" + try: + dense_bytes_per_ao_squared = _GLOBAL_DENSE_BYTES_PER_AO_SQUARED[func_deriv] + except KeyError as error: + raise ValueError("Invalid func_deriv value") from error + + raw_feature_bytes = estimate_global_raw_feature_buffer_memory( + dm, nfeatures, int(atomic_grid_sizes.sum().item()), func_deriv ) - bytes_per_point = 8.0 * elems_per_point # float64 - # Dense nao x nao buffers; autograd-independent and already conservative. - fixed_overhead = C_FIX * nao * nao - - return bytes_per_point, fixed_overhead + dense_buffer_bytes = int(dense_bytes_per_ao_squared * dm.shape[-1] ** 2) + return raw_feature_bytes + dense_buffer_bytes diff --git a/src/skala/pyscf/model_chunking.py b/src/skala/pyscf/model_chunking.py index 6f08c88b..581bc720 100644 --- a/src/skala/pyscf/model_chunking.py +++ b/src/skala/pyscf/model_chunking.py @@ -21,12 +21,21 @@ from skala.pyscf.backend import Grid from skala.pyscf.features import get_grid_features from skala.pyscf.memory_estimators import ( - estimate_global_raw_feature_buffer_memory, - estimate_max_gridpoint_chunk_size, + estimate_max_model_atoms_per_chunk, ) LOG = logging.getLogger(__name__) +_GRID_POINT_FEATURES = ( + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, +) +_ATOM_FEATURES = ( + Feature.COARSE_0_ATOMIC_COORDS, + Feature.ATOMIC_GRID_SIZES, +) + class AtomGridChunk(NamedTuple): """Matching atom and grid slices for one model evaluation chunk.""" @@ -36,54 +45,76 @@ class AtomGridChunk(NamedTuple): def _make_atom_grid_chunks( - atomic_grid_sizes: Tensor, max_model_grid_points: int + atomic_grid_sizes: Tensor, max_atoms_per_grid_size: Mapping[int, int] ) -> list[AtomGridChunk]: - """Build atom-aligned slices up to the requested model grid chunk size.""" - if max_model_grid_points < atomic_grid_sizes.max().item(): - raise ValueError( - "max_model_grid_points must be at least the maximum atomic grid size" - ) + """Pack equal-sized atomic grids up to each size group's atom limit.""" + if any(max_atoms < 1 for max_atoms in max_atoms_per_grid_size.values()): + raise ValueError("max_atoms_per_grid_size values must be positive") chunks: list[AtomGridChunk] = [] + grid_sizes = [int(size) for size in atomic_grid_sizes.tolist()] atom_start = 0 grid_start = 0 - chunk_size = 0 - - for atom_index, atom_grid_size in enumerate(atomic_grid_sizes): - chunk_size += atom_grid_size.item() - if chunk_size > max_model_grid_points: + while atom_start < len(grid_sizes): + atom_grid_size = grid_sizes[atom_start] + group_stop = atom_start + 1 + while group_stop < len(grid_sizes) and grid_sizes[group_stop] == atom_grid_size: + group_stop += 1 + + atoms_per_chunk = max_atoms_per_grid_size[atom_grid_size] + for chunk_atom_start in range(atom_start, group_stop, atoms_per_chunk): + chunk_atom_stop = min(chunk_atom_start + atoms_per_chunk, group_stop) + chunk_grid_size = (chunk_atom_stop - chunk_atom_start) * atom_grid_size chunks.append( AtomGridChunk( - atom_slice=slice(atom_start, atom_index), - grid_slice=slice( - grid_start, grid_start + chunk_size - atom_grid_size.item() - ), + atom_slice=slice(chunk_atom_start, chunk_atom_stop), + grid_slice=slice(grid_start, grid_start + chunk_grid_size), ) ) - atom_start = atom_index - grid_start += chunk_size - atom_grid_size.item() - chunk_size = atom_grid_size.item() - - if chunk_size > 0: - chunks.append( - AtomGridChunk( - atom_slice=slice(atom_start, len(atomic_grid_sizes)), - grid_slice=slice(grid_start, grid_start + chunk_size), - ) - ) + grid_start += chunk_grid_size + atom_start = group_stop LOG.debug( - "Generated %d model chunks of grid sizes: %s", + "Generated %d homogeneous model chunks of grid sizes: %s", len(chunks), [chunk.grid_slice.stop - chunk.grid_slice.start for chunk in chunks], ) return chunks +class AtomGridOrder(NamedTuple): + """Atom and grid-point indices in ascending atomic-grid-size order.""" + + atom_indices: Tensor + grid_indices: Tensor + + +def _make_atom_grid_order(atomic_grid_sizes: Tensor) -> AtomGridOrder: + """Build a stable atom ordering and its matching complete grid-block ordering.""" + atom_indices = torch.argsort(atomic_grid_sizes, stable=True) + sorted_sizes = atomic_grid_sizes.index_select(0, atom_indices) + total_grid_points = int(atomic_grid_sizes.sum().item()) + + original_starts = atomic_grid_sizes.cumsum(0) - atomic_grid_sizes + sorted_starts = sorted_sizes.cumsum(0) - sorted_sizes + point_atom_indices = torch.repeat_interleave( + atom_indices, sorted_sizes, output_size=total_grid_points + ) + point_sorted_starts = torch.repeat_interleave( + sorted_starts, sorted_sizes, output_size=total_grid_points + ) + grid_indices = ( + original_starts.index_select(0, point_atom_indices) + + torch.arange(total_grid_points, device=atomic_grid_sizes.device) + - point_sorted_starts + ) + return AtomGridOrder(atom_indices=atom_indices, grid_indices=grid_indices) + + class ModelFeatureChunk(NamedTuple): """Chunk-local raw features and the corresponding model input dictionary.""" - grid_slice: slice + grid_indices: Tensor raw_features: Tensor model_features: FeatureMap @@ -96,36 +127,33 @@ class ModelFeatureChunker: grid_features: Mapping[Feature, Tensor] feature_function: feature_math.MGGAFeatureFunction chunk_layouts: Sequence[AtomGridChunk] + atom_order: Tensor + grid_order: Tensor is_spin_polarized: bool def __iter__(self) -> Iterator[ModelFeatureChunk]: """Yield detached raw features paired with atom-aligned model inputs.""" feature_spec = self.feature_function.feature_spec for layout in self.chunk_layouts: + atom_indices = self.atom_order[layout.atom_slice] + grid_indices = self.grid_order[layout.grid_slice] raw_features = ( - self.atom_major_raw_features[..., layout.grid_slice] + self.atom_major_raw_features.index_select(-1, grid_indices) .detach() .requires_grad_() ) model_features: FeatureMap = {} - for feature_name in ( - Feature.GRID_COORDS, - Feature.GRID_WEIGHTS, - Feature.ATOMIC_GRID_WEIGHTS, - ): + for feature_name in _GRID_POINT_FEATURES: if feature_spec.requests(feature_name): - model_features[feature_name] = self.grid_features[feature_name][ - layout.grid_slice - ] - - for feature_name in ( - Feature.COARSE_0_ATOMIC_COORDS, - Feature.ATOMIC_GRID_SIZES, - ): + model_features[feature_name] = self.grid_features[ + feature_name + ].index_select(0, grid_indices) + + for feature_name in _ATOM_FEATURES: if feature_spec.requests(feature_name): - model_features[feature_name] = self.grid_features[feature_name][ - layout.atom_slice - ] + model_features[feature_name] = self.grid_features[ + feature_name + ].index_select(0, atom_indices) if feature_spec.requests(Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE): max_size = int(model_features[Feature.ATOMIC_GRID_SIZES].max().item()) @@ -143,7 +171,7 @@ def __iter__(self) -> Iterator[ModelFeatureChunk]: feature, not self.is_spin_polarized, 2 ) yield ModelFeatureChunk( - grid_slice=layout.grid_slice, + grid_indices=grid_indices, raw_features=raw_features, model_features=model_features, ) @@ -161,41 +189,44 @@ def prepare_model_feature_chunks( ) -> ModelFeatureChunker: """Prepare memory-sized, atom-aligned chunks for functional model evaluation.""" feature_spec = feature_function.feature_spec - if not feature_spec.supports_screened_evaluation: + if not feature_spec.supports_spatial_decomposition: raise ValueError( f"Atom-aligned model chunking requires {Feature.ATOMIC_GRID_SIZES.value!r}." ) grid_features = get_grid_features(mol, dm, grids, feature_spec) - max_model_grid_points = estimate_max_gridpoint_chunk_size( + atomic_grid_sizes = grid_features[Feature.ATOMIC_GRID_SIZES] + atom_grid_order = _make_atom_grid_order(atomic_grid_sizes) + sorted_atomic_grid_sizes = atomic_grid_sizes.index_select( + 0, atom_grid_order.atom_indices + ) + + max_atoms_per_grid_size = estimate_max_model_atoms_per_chunk( dm=dm, - deriv=feature_function.deriv, + atomic_grid_sizes=sorted_atomic_grid_sizes, + nfeatures=feature_function.nfeats, max_memory_in_mb=max_memory_in_mb, safety_fraction=safety_fraction, func_deriv=deriv_order, - reserved_memory_in_bytes=estimate_global_raw_feature_buffer_memory( - dm, - feature_function.nfeats, - atom_major_raw_features.shape[-1], - deriv_order, - ), ) - max_atom_grid = int(grid_features[Feature.ATOMIC_GRID_SIZES].max().item()) - if max_model_grid_points < max_atom_grid: - LOG.warning( - "Adjusted model chunk size %d to match the largest atomic grid %d. " - "Hope for no OOM.", - max_model_grid_points, - max_atom_grid, - ) - max_model_grid_points = max_atom_grid + for grid_size, max_atoms in max_atoms_per_grid_size.items(): + if max_atoms < 1: + LOG.warning( + "Adjusted model chunk capacity for atomic grid size %d from %d " + "to one atom. Hope for no OOM.", + grid_size, + max_atoms, + ) + max_atoms_per_grid_size[grid_size] = 1 return ModelFeatureChunker( atom_major_raw_features=atom_major_raw_features, grid_features=grid_features, feature_function=feature_function, chunk_layouts=_make_atom_grid_chunks( - grid_features[Feature.ATOMIC_GRID_SIZES], max_model_grid_points + sorted_atomic_grid_sizes, max_atoms_per_grid_size ), + atom_order=atom_grid_order.atom_indices, + grid_order=atom_grid_order.grid_indices, is_spin_polarized=dm.ndim == 3, ) diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index ad2ea8fa..8f6c64f4 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -84,7 +84,7 @@ def __call__( ) -> XCResult: """Evaluate electron count, XC energy, and XC potential.""" self._validate_device(dm) - if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): + if self.feature_spec.supports_spatial_decomposition and _should_screen_aos(mol): return self._integrate_screened(mol, grids, dm, max_memory) return self._integrate_dense(mol, grids, dm, max_memory) @@ -98,7 +98,7 @@ def gen_response( ) -> Callable[[Tensor], Tensor]: """Build an XC-only Hessian-vector product callable.""" self._validate_device(dm0) - if self.feature_spec.supports_screened_evaluation and _should_screen_aos(mol): + if self.feature_spec.supports_spatial_decomposition and _should_screen_aos(mol): return self._gen_response_screened( mol, grids, @@ -196,7 +196,9 @@ def _integrate_screened( local_raw_features, torch.ones_like(energy_chunk), ) - atom_major_cotangent[..., chunk.grid_slice] = local_cotangent.detach() + atom_major_cotangent.index_copy_( + -1, chunk.grid_indices, local_cotangent.detach() + ) electron_count += ( (mol_features[Feature.DENSITY] * mol_features[Feature.GRID_WEIGHTS]) .sum(dim=-1) @@ -308,12 +310,12 @@ def hessian_vector_product(dm1: Tensor) -> Tensor: (local_hessian_action,) = torch.autograd.grad( local_gradient, local_raw_features, - atom_major_tangent[..., chunk.grid_slice], + atom_major_tangent.index_select(-1, chunk.grid_indices), ) else: local_hessian_action = torch.zeros_like(local_raw_features) - atom_major_hessian_action[..., chunk.grid_slice] = ( - local_hessian_action.detach() + atom_major_hessian_action.index_copy_( + -1, chunk.grid_indices, local_hessian_action.detach() ) del ( energy_chunk, diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 288d6f56..ac6de546 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -4,6 +4,7 @@ import numpy as np import pytest import torch +from benchmarks import run_pyscf_ao_screening_benchmark as benchmark_runner from pyscf import dft, gto from utils import QuadraticFunctional, patch_ao_screening @@ -127,6 +128,17 @@ def test_patch_ao_screening_restores_previous_decision(carbon: gto.Mole) -> None assert xc_integrator_module._should_screen_aos is original_decision +def test_benchmark_route_metadata_uses_integrator_feature_spec( + carbon: gto.Mole, +) -> None: + metadata = benchmark_runner.route_metadata( + SkalaNumInt(QuadraticFunctional()), carbon, forced_dense=False + ) + + assert metadata["implementation_target"].endswith("XCIntegrator.__call__") + assert metadata["functional_supports_screened_evaluation"] is True + + def test_active_cpu_ao_indices(carbon: gto.Mole) -> None: ao_loc = carbon.ao_loc_nr() screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) @@ -428,7 +440,7 @@ def __init__(self, raw_features: torch.Tensor) -> None: def __iter__(self) -> Iterator[ModelFeatureChunk]: raw_features = self.raw_features.detach().requires_grad_() yield ModelFeatureChunk( - grid_slice=slice(0, 1), + grid_indices=torch.tensor([0]), raw_features=raw_features, model_features={ Feature.ATOMIC_GRID_SIZES: torch.tensor([1]), @@ -858,8 +870,24 @@ def test_cpu_rks_uks_dense_screened_equivalence( ) +def test_cpu_quadratic_dense_screened_equivalence_heteronuclear() -> None: + mol = gto.M(atom="H 0 0 0; F 0 0 0.92", basis="sto-3g", spin=0, verbose=0) + grids = _minimal_atom_grid(mol) + numint = SkalaNumInt(QuadraticFunctional()) + dm = dft.RKS(mol).get_init_guess() + + with patch_ao_screening(False): + dense = numint.nr_rks(mol, grids, None, dm) + + with patch_ao_screening(True): + screened = numint.nr_rks(mol, grids, None, dm) + + for dense_value, screened_value in zip(dense, screened, strict=True): + assert np.allclose(dense_value, screened_value, rtol=1e-10, atol=1e-11) + + def test_cpu_response_dense_screened_equivalence() -> None: - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", spin=0, verbose=0) + mol = gto.M(atom="H 0 0 0; F 0 0 0.92", basis="sto-3g", spin=0, verbose=0) grids = _minimal_atom_grid(mol) ks = FakeKS(mol, grids) numint = SkalaNumInt(QuadraticFunctional()) @@ -891,8 +919,8 @@ def test_screened_ao_traversals_are_independent_of_model_chunking( atom_grid_size = grids.weights.size // mol.natm monkeypatch.setattr( model_chunking_module, - "estimate_max_gridpoint_chunk_size", - lambda *args, **kwargs: atom_grid_size, + "estimate_max_model_atoms_per_chunk", + lambda *args, **kwargs: {atom_grid_size: 1}, ) forward_calls = 0 diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py index 3d2bfe2f..93d0fe53 100644 --- a/tests/test_evaluation.py +++ b/tests/test_evaluation.py @@ -57,7 +57,7 @@ def test_feature_spec_derives_atomic_layout_requirements( assert spec.names == frozenset({feature}) assert spec.requires_atomic_layout - assert spec.supports_screened_evaluation is supports_screened_evaluation + assert spec.supports_spatial_decomposition is supports_screened_evaluation def test_evaluation_policy_defaults_and_is_immutable() -> None: diff --git a/tests/test_memory_estimators.py b/tests/test_memory_estimators.py index ec5b6200..e0a92655 100644 --- a/tests/test_memory_estimators.py +++ b/tests/test_memory_estimators.py @@ -3,8 +3,9 @@ from skala.pyscf.memory_estimators import ( estimate_global_raw_feature_buffer_memory, - estimate_max_gridpoint_chunk_size, - linear_peak_memory_model, + estimate_global_screened_buffer_memory, + estimate_max_model_atoms_per_chunk, + estimate_model_memory_per_grid_point, ) @@ -34,27 +35,50 @@ def test_global_raw_feature_buffer_memory_rejects_unsupported_order() -> None: ) -def test_reserved_memory_reduces_grid_chunk_size() -> None: +def test_global_screened_buffer_memory_uses_atomic_grid_sizes() -> None: dm = torch.eye(10, dtype=torch.float64) - bytes_per_point, _ = linear_peak_memory_model(nao=10, deriv=1, func_deriv=1) - base_chunk_size = estimate_max_gridpoint_chunk_size( - dm, - deriv=1, - max_memory_in_mb=100, - safety_fraction=1.0, - func_deriv=1, + atomic_grid_sizes = torch.tensor([10, 10, 20]) + + actual = estimate_global_screened_buffer_memory( + dm, nfeatures=5, atomic_grid_sizes=atomic_grid_sizes, func_deriv=1 ) - reserved_points = 123 - reserved_chunk_size = estimate_max_gridpoint_chunk_size( + + raw_feature_bytes = 4 * 5 * 40 * 8 + dense_buffer_bytes = int(37.0 * 10**2) + assert actual == raw_feature_bytes + dense_buffer_bytes + + +def test_model_atom_limits_are_estimated_per_atomic_grid_size() -> None: + dm = torch.eye(10, dtype=torch.float64) + atomic_grid_sizes = torch.tensor([10, 10, 20]) + + actual = estimate_max_model_atoms_per_chunk( dm, - deriv=1, - max_memory_in_mb=100, + atomic_grid_sizes=atomic_grid_sizes, + nfeatures=5, + max_memory_in_mb=10, safety_fraction=1.0, func_deriv=1, - reserved_memory_in_bytes=int(bytes_per_point * reserved_points), ) - assert base_chunk_size - reserved_chunk_size == reserved_points + available_memory = 10_000_000 - estimate_global_screened_buffer_memory( + dm, 5, atomic_grid_sizes, 1 + ) + bytes_per_point = estimate_model_memory_per_grid_point(1) + assert actual == { + 10: available_memory // (10 * bytes_per_point), + 20: available_memory // (20 * bytes_per_point), + } + + +@pytest.mark.parametrize( + ("func_deriv", "elements_per_point"), + [(0, 5830), (1, 6680), (2, 24230)], +) +def test_model_memory_per_grid_point_depends_on_functional_derivative( + func_deriv: int, elements_per_point: int +) -> None: + assert estimate_model_memory_per_grid_point(func_deriv) == 8 * elements_per_point @pytest.mark.parametrize("safety_fraction", [-0.1, 0.0, 1.1]) @@ -64,9 +88,10 @@ def test_model_grid_point_limit_rejects_invalid_safety_fraction( with pytest.raises( ValueError, match="safety_fraction must be greater than 0 and at most 1" ): - estimate_max_gridpoint_chunk_size( + estimate_max_model_atoms_per_chunk( torch.eye(2, dtype=torch.float64), - deriv=1, + atomic_grid_sizes=torch.tensor([10]), + nfeatures=5, max_memory_in_mb=100, safety_fraction=safety_fraction, ) diff --git a/tests/test_model_chunking.py b/tests/test_model_chunking.py new file mode 100644 index 00000000..6b024ded --- /dev/null +++ b/tests/test_model_chunking.py @@ -0,0 +1,113 @@ +# SPDX-License-Identifier: MIT + +from typing import cast + +import pytest +import torch +from pyscf import gto + +from skala.features import Feature, FeatureMap +from skala.pyscf import model_chunking +from skala.pyscf.backend import Grid +from skala.pyscf.evaluation import FeatureSpec +from skala.pyscf.feature_math import MGGAFeatureFunction + + +def test_prepare_model_feature_chunks_sorts_complete_atomic_grids( + monkeypatch: pytest.MonkeyPatch, +) -> None: + atomic_grid_sizes = torch.tensor([3, 1, 2, 1]) + point_ids = torch.arange(7, dtype=torch.float64) + atom_ids = torch.arange(4, dtype=torch.float64) + grid_features: FeatureMap = { + Feature.GRID_COORDS: point_ids[:, None].expand(-1, 3), + Feature.GRID_WEIGHTS: point_ids + 10, + Feature.ATOMIC_GRID_WEIGHTS: point_ids + 20, + Feature.COARSE_0_ATOMIC_COORDS: atom_ids[:, None].expand(-1, 3), + Feature.ATOMIC_GRID_SIZES: atomic_grid_sizes, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE: torch.zeros(3, 0, dtype=torch.long), + } + + def fake_get_grid_features(*args: object, **kwargs: object) -> FeatureMap: + return grid_features + + def fake_estimate_max_model_atoms_per_chunk( + **kwargs: object, + ) -> dict[int, int]: + return {1: 2, 2: 2, 3: 1} + + monkeypatch.setattr(model_chunking, "get_grid_features", fake_get_grid_features) + monkeypatch.setattr( + model_chunking, + "estimate_max_model_atoms_per_chunk", + fake_estimate_max_model_atoms_per_chunk, + ) + + feature_spec = FeatureSpec( + { + Feature.DENSITY, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + Feature.ATOMIC_GRID_WEIGHTS, + Feature.COARSE_0_ATOMIC_COORDS, + Feature.ATOMIC_GRID_SIZES, + Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE, + } + ) + raw_features = point_ids.reshape(1, -1) + chunker = model_chunking.prepare_model_feature_chunks( + mol=cast(gto.Mole, object()), + dm=torch.eye(1, dtype=torch.float64), + grids=cast(Grid, object()), + atom_major_raw_features=raw_features, + feature_function=MGGAFeatureFunction(feature_spec), + deriv_order=1, + ) + + expected_grid_order = torch.tensor([3, 6, 4, 5, 0, 1, 2]) + assert torch.equal(chunker.atom_order, torch.tensor([1, 3, 2, 0])) + assert torch.equal(chunker.grid_order, expected_grid_order) + assert torch.equal( + chunker.grid_features[Feature.ATOMIC_GRID_SIZES], atomic_grid_sizes + ) + assert torch.equal(chunker.atom_major_raw_features.flatten(), point_ids) + + chunks = list(chunker) + assert len(chunks) == 3 + assert torch.equal(chunks[0].grid_indices, torch.tensor([3, 6])) + assert torch.equal(chunks[0].raw_features.flatten(), torch.tensor([3.0, 6.0])) + assert torch.equal( + chunks[0].model_features[Feature.ATOMIC_GRID_SIZES], torch.tensor([1, 1]) + ) + assert torch.equal( + chunks[0].model_features[Feature.COARSE_0_ATOMIC_COORDS][:, 0], + torch.tensor([1.0, 3.0]), + ) + assert torch.equal(chunks[1].grid_indices, torch.tensor([4, 5])) + assert torch.equal(chunks[2].grid_indices, torch.tensor([0, 1, 2])) + + +def test_atom_grid_chunks_pack_equal_sizes_up_to_cap() -> None: + chunks = model_chunking._make_atom_grid_chunks( + torch.tensor([2, 2, 2, 2, 2]), max_atoms_per_grid_size={2: 2} + ) + + assert chunks == [ + model_chunking.AtomGridChunk(slice(0, 2), slice(0, 4)), + model_chunking.AtomGridChunk(slice(2, 4), slice(4, 8)), + model_chunking.AtomGridChunk(slice(4, 5), slice(8, 10)), + ] + + +def test_atom_grid_chunks_apply_limits_per_grid_size() -> None: + chunks = model_chunking._make_atom_grid_chunks( + torch.tensor([1, 1, 1, 2, 2, 2]), + max_atoms_per_grid_size={1: 3, 2: 1}, + ) + + assert chunks == [ + model_chunking.AtomGridChunk(slice(0, 3), slice(0, 3)), + model_chunking.AtomGridChunk(slice(3, 4), slice(3, 5)), + model_chunking.AtomGridChunk(slice(4, 5), slice(5, 7)), + model_chunking.AtomGridChunk(slice(5, 6), slice(7, 9)), + ] From 26fa246c03196fe809a67cc6ba8b6dab22ea69f3 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Sat, 8 Aug 2026 23:22:33 +0200 Subject: [PATCH 32/39] use analytic vjp for backpropagation of ml model outputs to density matrix --- src/skala/pyscf/ao_evaluation.py | 55 ++++++++-------- src/skala/pyscf/feature_math.py | 56 +++++++++++++++- tests/test_ao_screening.py | 106 +++++++++++++++++++++++++------ 3 files changed, 165 insertions(+), 52 deletions(-) diff --git a/src/skala/pyscf/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py index 2cdaf1c8..0c51d79a 100644 --- a/src/skala/pyscf/ao_evaluation.py +++ b/src/skala/pyscf/ao_evaluation.py @@ -2,7 +2,7 @@ """Blockwise atomic-orbital feature evaluation and custom autograd.""" -from collections.abc import Callable, Iterator +from collections.abc import Iterator from typing import NamedTuple, Protocol, TypeAlias, cast import numpy as np @@ -32,7 +32,7 @@ class _ChunkEvalForwardContext(Protocol): dm: Tensor mol: gto.Mole grids: Grid - feature_function: feature_math.FeatureFunction + feature_function: feature_math.LinearFeature blksize: int | None compile_feature_function: bool gpu: bool @@ -43,7 +43,7 @@ class _ChunkEvalBackwardContext(Protocol): dm: Tensor mol: gto.Mole grids: Grid - feature_function: feature_math.FeatureFunction + feature_function: feature_math.LinearFeature blksize: int | None compile_feature_function: bool gpu: bool @@ -94,29 +94,26 @@ def add_active_ao_submatrix(self, matrix: Tensor, block_result: Tensor) -> None: def _evaluate_feature_block( - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, block: _AOBlock, - active_dm_submatrix: Tensor, + active_dm_submatrix: Tensor | None, compile_feature_function: bool, feature_cotangent: Tensor | None = None, ) -> Tensor: """Evaluate one active-AO feature block or its feature-space VJP.""" - - def evaluate_features(dm: Tensor) -> Tensor: - return feature_function(dm, block.ao_values) - - evaluation_function: Callable[[Tensor], Tensor] = evaluate_features if feature_cotangent is not None: local_cotangent = feature_cotangent[..., block.grid_slice] + if compile_feature_function: + return torch.compile(feature_function.vjp)(block.ao_values, local_cotangent) + return feature_function.vjp(block.ao_values, local_cotangent) - def evaluate_vjp(dm: Tensor) -> Tensor: - return torch.func.vjp(evaluate_features, dm)[1](local_cotangent)[0] - - evaluation_function = evaluate_vjp - + if active_dm_submatrix is None: + raise ValueError("Feature evaluation requires a density matrix.") if compile_feature_function: - return torch.compile(evaluation_function)(active_dm_submatrix) - return evaluation_function(active_dm_submatrix) + return torch.compile(feature_function.forward)( + active_dm_submatrix, block.ao_values + ) + return feature_function(active_dm_submatrix, block.ao_values) class _CPUAOBlockLoop: @@ -148,7 +145,7 @@ def __init__( dm: Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, blksize: int | None, ) -> None: self.dm = dm @@ -236,7 +233,7 @@ def __init__( dm: Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, blksize: int | None, ) -> None: check_gpu_imports_were_successful() @@ -287,7 +284,7 @@ def _make_ao_block_loop( dm: Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, blksize: int | None, gpu: bool, ) -> _CPUAOBlockLoop | _GPUAOBlockLoop: @@ -303,7 +300,7 @@ def setup_context( Tensor, gto.Mole, Grid, - feature_math.FeatureFunction, + feature_math.LinearFeature, int | None, bool, bool, @@ -333,7 +330,7 @@ def forward( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, blksize: int | None, compile_feature_function: bool, gpu: bool, @@ -448,7 +445,7 @@ def setup_context( torch.Tensor, gto.Mole, Grid, - feature_math.FeatureFunction, + feature_math.LinearFeature, int | None, bool, bool, @@ -474,22 +471,20 @@ def forward( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, blksize: int | None, compile_feature_function: bool, gpu: bool, feature_cotangent: torch.Tensor, ) -> torch.Tensor: block_loop = _make_ao_block_loop(dm, mol, grids, feature_function, blksize, gpu) - dm_ordered = block_loop.order_aos(dm) out = torch.zeros_like(dm) for block in block_loop: - active_dm_submatrix = block.select_active_ao_submatrix(dm_ordered) block_result = _evaluate_feature_block( feature_function, block, - active_dm_submatrix, + None, compile_feature_function, feature_cotangent, ) @@ -539,7 +534,7 @@ def evaluate_full_grid( dm: torch.Tensor, mol: gto.Mole, coords: Array, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, compile_feature_function: bool = False, gpu: bool = False, ) -> torch.Tensor: @@ -562,7 +557,7 @@ def evaluate_full_grid( def _resolve_ao_block_size( mol: gto.Mole, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, block_size: int | None, max_memory: int, gpu: bool, @@ -592,7 +587,7 @@ def auto_chunk( dm: torch.Tensor, mol: gto.Mole, grids: Grid, - feature_function: feature_math.FeatureFunction, + feature_function: feature_math.LinearFeature, block_size: int | None = None, max_memory: int = 2000, gpu: bool = False, diff --git a/src/skala/pyscf/feature_math.py b/src/skala/pyscf/feature_math.py index f16d278c..b7177dfd 100644 --- a/src/skala/pyscf/feature_math.py +++ b/src/skala/pyscf/feature_math.py @@ -20,8 +20,8 @@ def maybe_expand_and_divide( return feature -class FeatureFunction(nn.Module, ABC): - """Base class for raw features evaluated from density and AO tensors.""" +class LinearFeature(nn.Module, ABC): + """Linear raw-feature map from a density matrix and fixed AO values.""" deriv: int nfeats: int @@ -29,11 +29,15 @@ class FeatureFunction(nn.Module, ABC): @abstractmethod def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: ... + @abstractmethod + def vjp(self, ao: torch.Tensor, cotangent: torch.Tensor) -> torch.Tensor: + """Apply the adjoint feature map to a feature-space cotangent.""" + @abstractmethod def to_dict(self, features: torch.Tensor) -> FeatureMap: ... -class MGGAFeatureFunction(FeatureFunction): +class MGGAFeatureFunction(LinearFeature): """Evaluate the requested linear meta-GGA density features.""" def __init__(self, feature_spec: FeatureSpec): @@ -120,3 +124,49 @@ def forward(self, dm: torch.Tensor, ao: torch.Tensor) -> torch.Tensor: if len(dm.shape) == 2: return features.reshape((self.nfeats, -1)) return features.reshape((*dm.shape[:-2], self.nfeats, -1)) + + def vjp(self, ao: torch.Tensor, cotangent: torch.Tensor) -> torch.Tensor: + """Apply the analytic adjoint of the linear MGGA feature map.""" + batch_shape = cotangent.shape[:-2] + ngrids = cotangent.shape[-1] + weights = cotangent.reshape(-1, self.nfeats, ngrids) + phi = ao if self.deriv == 0 else ao[0] + nao = phi.shape[-2] + + if self.deriv == 0: + result = (weights[:, 0, None, :] * phi) @ phi.transpose(-1, -2) + return result.reshape(*batch_shape, nao, nao) + + left = weights.new_zeros((weights.shape[0], nao, ngrids)) + feature_index = 0 + if self.feature_spec.with_density: + left += weights[:, feature_index, None, :] * phi + feature_index += 1 + + if self.feature_spec.with_grad: + left += 2 * torch.einsum( + "bcg,cig->big", + weights[:, feature_index : feature_index + 3], + ao[1:4], + ) + feature_index += 3 + + derivative_weight = weights.new_zeros((weights.shape[0], ngrids)) + if self.feature_spec.with_kin: + derivative_weight += 0.5 * weights[:, feature_index] + feature_index += 1 + + if self.feature_spec.with_lapl: + laplacian_weight = weights[:, feature_index] + derivative_weight += 2 * laplacian_weight + for component in (4, 7, 9): + left.addcmul_(laplacian_weight[:, None, :], ao[component], value=2) + + result = left @ phi.transpose(-1, -2) + if self.feature_spec.with_kin or self.feature_spec.with_lapl: + for component in range(1, 4): + weighted_derivative = derivative_weight[:, None, :] * ao[component] + result += weighted_derivative @ ao[component].transpose(-1, -2) + + result = 0.5 * (result + result.transpose(-1, -2)) + return result.reshape(*batch_shape, nao, nao) diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index ac6de546..b76cfa1a 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -1,10 +1,10 @@ from collections.abc import Callable, Iterator +from itertools import combinations from typing import Any import numpy as np import pytest import torch -from benchmarks import run_pyscf_ao_screening_benchmark as benchmark_runner from pyscf import dft, gto from utils import QuadraticFunctional, patch_ao_screening @@ -32,6 +32,13 @@ prepare_spatial_grid_layout, ) +_MGGA_FEATURES = (Feature.DENSITY, Feature.GRAD, Feature.KIN, Feature.LAPL) +_MGGA_FEATURE_COMBINATIONS = [ + combination + for size in range(1, len(_MGGA_FEATURES) + 1) + for combination in combinations(_MGGA_FEATURES, size) +] + @pytest.fixture def carbon() -> gto.Mole: @@ -105,6 +112,84 @@ def first_jvp(value: torch.Tensor) -> torch.Tensor: torch.testing.assert_close(second_jvp, torch.zeros_like(second_jvp)) +@pytest.mark.parametrize( + "feature_names", + _MGGA_FEATURE_COMBINATIONS, +) +@pytest.mark.parametrize("spin_channels", [None, 2]) +def test_mgga_analytic_vjp_matches_autograd( + feature_names: tuple[Feature, ...], spin_channels: int | None +) -> None: + feature_function = MGGAFeatureFunction(FeatureSpec(feature_names)) + ncomp = ( + (feature_function.deriv + 1) + * (feature_function.deriv + 2) + * (feature_function.deriv + 3) + // 6 + ) + generator = torch.Generator().manual_seed(0) + ao = torch.randn((ncomp, 3, 5), dtype=torch.float64, generator=generator) + if feature_function.deriv == 0: + ao = ao[0] + dm_shape = (3, 3) if spin_channels is None else (spin_channels, 3, 3) + dm = torch.randn(dm_shape, dtype=torch.float64, generator=generator) + + features, pullback = torch.func.vjp(lambda value: feature_function(value, ao), dm) + cotangent = torch.randn(features.shape, dtype=features.dtype, generator=generator) + + expected = pullback(cotangent)[0] + actual = feature_function.vjp(ao, cotangent) + + torch.testing.assert_close(actual, expected) + + +def test_feature_block_compiled_vjp_matches_eager( + monkeypatch: pytest.MonkeyPatch, +) -> None: + feature_function = MGGAFeatureFunction(FeatureSpec(_MGGA_FEATURES)) + generator = torch.Generator().manual_seed(0) + block = _AOBlock( + ao_values=torch.randn((10, 3, 5), dtype=torch.float64, generator=generator), + active_ao_indices=None, + grid_slice=slice(1, 6), + ) + dm = torch.randn((3, 3), dtype=torch.float64, generator=generator) + cotangent = torch.randn( + (feature_function.nfeats, 7), dtype=torch.float64, generator=generator + ) + eager_forward = _evaluate_feature_block( + feature_function, block, dm, compile_feature_function=False + ) + eager_vjp = _evaluate_feature_block( + feature_function, + block, + None, + compile_feature_function=False, + feature_cotangent=cotangent, + ) + + compile_function = torch.compile + monkeypatch.setattr( + torch, + "compile", + lambda function: compile_function(function, backend="eager"), + ) + + compiled_forward = _evaluate_feature_block( + feature_function, block, dm, compile_feature_function=True + ) + compiled_vjp = _evaluate_feature_block( + feature_function, + block, + None, + compile_feature_function=True, + feature_cotangent=cotangent, + ) + + torch.testing.assert_close(compiled_forward, eager_forward) + torch.testing.assert_close(compiled_vjp, eager_vjp) + + @pytest.mark.parametrize("feature_names", [[], [Feature.GRID_WEIGHTS]]) def test_mgga_requires_at_least_one_ao_derived_feature( feature_names: list[Feature], @@ -128,17 +213,6 @@ def test_patch_ao_screening_restores_previous_decision(carbon: gto.Mole) -> None assert xc_integrator_module._should_screen_aos is original_decision -def test_benchmark_route_metadata_uses_integrator_feature_spec( - carbon: gto.Mole, -) -> None: - metadata = benchmark_runner.route_metadata( - SkalaNumInt(QuadraticFunctional()), carbon, forced_dense=False - ) - - assert metadata["implementation_target"].endswith("XCIntegrator.__call__") - assert metadata["functional_supports_screened_evaluation"] is True - - def test_active_cpu_ao_indices(carbon: gto.Mole) -> None: ao_loc = carbon.ao_loc_nr() screen_index = np.zeros((2, carbon.nbas), dtype=np.uint8) @@ -557,16 +631,10 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: active_ao_indices=torch.tensor([0, 2]), grid_slice=slice(1, 3), ) - dm_ordered = torch.tensor( - [[2.0, 0.1, 0.2], [0.1, 1.0, 0.3], [0.2, 0.3, 3.0]], - dtype=torch.float64, - ) tangent_ordered = torch.tensor( [[0.5, 1.0, -0.2], [1.0, 0.4, 0.3], [-0.2, 0.3, 0.7]], dtype=torch.float64, ) - active_dm_submatrix = block.select_active_ao_submatrix(dm_ordered) - feature_jvp = _evaluate_feature_block( feature_function, block, @@ -582,7 +650,7 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: feature_vjp = _evaluate_feature_block( feature_function, block, - active_dm_submatrix, + None, compile_feature_function=False, feature_cotangent=full_grid_cotangent, ) From af88a1c96ed11eb46446733b8ce69573f5e398b4 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Sun, 9 Aug 2026 00:23:26 +0200 Subject: [PATCH 33/39] small optimisation --- src/skala/pyscf/feature_math.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/src/skala/pyscf/feature_math.py b/src/skala/pyscf/feature_math.py index b7177dfd..b6a2f66b 100644 --- a/src/skala/pyscf/feature_math.py +++ b/src/skala/pyscf/feature_math.py @@ -144,11 +144,12 @@ def vjp(self, ao: torch.Tensor, cotangent: torch.Tensor) -> torch.Tensor: feature_index += 1 if self.feature_spec.with_grad: - left += 2 * torch.einsum( - "bcg,cig->big", - weights[:, feature_index : feature_index + 3], - ao[1:4], - ) + for component in range(3): + left.addcmul_( + weights[:, feature_index + component, None, :], + ao[component + 1], + value=2, + ) feature_index += 3 derivative_weight = weights.new_zeros((weights.shape[0], ngrids)) From 59b020bf609094e66bdd3edc7968fb342558d192 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 06:13:05 +0200 Subject: [PATCH 34/39] use newer pyscf version in test --- .github/workflows/test.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index d7d41f89..6f402740 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -78,7 +78,7 @@ jobs: shell: micromamba-shell {0} profiling: - name: "Profiling (Python=3.12 & PySCF=2.9)" + name: "Profiling (Python=3.12 & PySCF=2.13.1)" runs-on: ubuntu-latest needs: - lint @@ -96,7 +96,7 @@ jobs: cache-downloads: true create-args: >- python=3.12 - pyscf=2.9 + pyscf=2.13.1 - name: Install package in development mode run: | From 01a42db21a9e2e3b6dc78d6a1fd76be1c643d3b3 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 06:29:35 +0200 Subject: [PATCH 35/39] pin versions --- environment-cpu.yml | 1 + environment-gpu.yml | 1 + 2 files changed, 2 insertions(+) diff --git a/environment-cpu.yml b/environment-cpu.yml index aa4f4353..2b3d9715 100644 --- a/environment-cpu.yml +++ b/environment-cpu.yml @@ -8,6 +8,7 @@ dependencies: - dftd3-python - e3nn - h5py + - hdf5 >=2.1,<2.2 - numpy <2.5 - opt_einsum_fx - pyscf >=2.8,<2.14 diff --git a/environment-gpu.yml b/environment-gpu.yml index bda42cf4..bf875ce9 100644 --- a/environment-gpu.yml +++ b/environment-gpu.yml @@ -8,6 +8,7 @@ dependencies: - dftd3-python - e3nn - h5py + - hdf5 >=2.1,<2.2 - numpy - opt_einsum_fx - pyscf >=2.8,<2.14 From 130bced2565430a384f3c7a9aa15d180e2ec6961 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 07:32:15 +0200 Subject: [PATCH 36/39] raise error bound --- tests/test_ao_screening_benchmark.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index da0e6048..40d60b3a 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -361,9 +361,10 @@ def test_screened_and_dense_values_agree( assert isinstance(cpu_functional, ExcFunctionalBase) dense = _make_benchmark_case(benchmark_spec, cpu_functional, "cpu").run() - scalar_rtol = 2e-10 if case.backend == "cpu" else 1e-8 - density_close = np.allclose(dense[0], screened[0], rtol=scalar_rtol, atol=1e-11) - energy_close = np.isclose(dense[1], screened[1], rtol=scalar_rtol, atol=1e-10) + density_rtol = 2e-10 if case.backend == "cpu" else 1e-8 + energy_rtol = 5e-10 if case.backend == "cpu" else 1e-8 + density_close = np.allclose(dense[0], screened[0], rtol=density_rtol, atol=1e-11) + energy_close = np.isclose(dense[1], screened[1], rtol=energy_rtol, atol=1e-10) dense_vxc = ( dense[2] if isinstance(dense[2], np.ndarray) From 0308eb8851e11d6e578fc707b2aea7bf499010b4 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 16:47:39 +0200 Subject: [PATCH 37/39] Make linear model simpler and add test --- src/skala/pyscf/ao_evaluation.py | 232 ++++++++++----------------- src/skala/pyscf/screening.py | 5 +- src/skala/pyscf/xc_integrator.py | 3 - tests/test_ao_screening.py | 24 +-- tests/test_gpu4pyscf_ao_screening.py | 4 +- tests/test_xc_integrator.py | 67 +++++++- 6 files changed, 171 insertions(+), 164 deletions(-) diff --git a/src/skala/pyscf/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py index 0c51d79a..5bb171cd 100644 --- a/src/skala/pyscf/ao_evaluation.py +++ b/src/skala/pyscf/ao_evaluation.py @@ -12,7 +12,6 @@ from torch.autograd import Function from torch.autograd.function import FunctionCtx from torch.utils.dlpack import from_dlpack -from typing_extensions import Unpack from skala.features import FeatureMap from skala.pyscf import feature_math @@ -28,25 +27,14 @@ _AOIndices: TypeAlias = np.ndarray[tuple[int], np.dtype[np.intp]] -class _ChunkEvalForwardContext(Protocol): - dm: Tensor +class _ChunkEvalContext(Protocol): mol: gto.Mole grids: Grid feature_function: feature_math.LinearFeature blksize: int | None compile_feature_function: bool - gpu: bool - vectors_jvp: tuple[Tensor, ...] - - -class _ChunkEvalBackwardContext(Protocol): - dm: Tensor - mol: gto.Mole - grids: Grid - feature_function: feature_math.LinearFeature - blksize: int | None - compile_feature_function: bool - gpu: bool + spin_shape: torch.Size + output_device: torch.device def _active_cpu_ao_indices(mol: gto.Mole, screen_index: _ScreenIndex) -> _AOIndices: @@ -142,13 +130,11 @@ class _CPUAOBlockLoop: def __init__( self, - dm: Tensor, mol: gto.Mole, grids: Grid, feature_function: feature_math.LinearFeature, blksize: int | None, ) -> None: - self.dm = dm self.mol = mol assert isinstance(grids, dft.Grids) self.grids = grids @@ -230,21 +216,20 @@ class _GPUAOBlockLoop: def __init__( self, - dm: Tensor, + device: torch.device, mol: gto.Mole, grids: Grid, feature_function: feature_math.LinearFeature, blksize: int | None, ) -> None: check_gpu_imports_were_successful() - self.dm = dm self.mol = mol self.grids = grids self.feature_function = feature_function self.blksize = blksize self.numint = dft_gpu.numint.NumInt().build(mol, grids.coords) self.numint.grid_blksize = blksize - self.sort_idx = torch.as_tensor(self.numint.gdftopt._ao_idx, device=dm.device) + self.sort_idx = torch.as_tensor(self.numint.gdftopt._ao_idx, device=device) self.unsort_idx = torch.argsort(self.sort_idx) def order_aos(self, matrix: Tensor) -> Tensor: @@ -281,15 +266,15 @@ def __iter__(self) -> Iterator[_AOBlock]: def _make_ao_block_loop( - dm: Tensor, + device: torch.device, mol: gto.Mole, grids: Grid, feature_function: feature_math.LinearFeature, blksize: int | None, - gpu: bool, ) -> _CPUAOBlockLoop | _GPUAOBlockLoop: - loop_type = _GPUAOBlockLoop if gpu else _CPUAOBlockLoop - return loop_type(dm, mol, grids, feature_function, blksize) + if device.type == "cuda": + return _GPUAOBlockLoop(device, mol, grids, feature_function, blksize) + return _CPUAOBlockLoop(mol, grids, feature_function, blksize) class ChunkEvalForward(Function): @@ -303,27 +288,20 @@ def setup_context( feature_math.LinearFeature, int | None, bool, - bool, - # The starred spelling requires Python 3.11. - Unpack[tuple[Tensor, ...]], # noqa: UP044 ], output: torch.Tensor, ) -> None: - if len(inputs) < 7: - raise ValueError("ChunkEvalForward requires seven fixed inputs.") - context = cast(_ChunkEvalForwardContext, ctx) + context = cast(_ChunkEvalContext, ctx) ( - context.dm, + dm, context.mol, context.grids, context.feature_function, context.blksize, context.compile_feature_function, - context.gpu, - *vectors_jvp, ) = inputs - context.vectors_jvp = tuple(vectors_jvp) - ctx.save_for_backward(context.dm) + context.spin_shape = dm.shape[:-2] + context.output_device = output.device @staticmethod def forward( @@ -333,11 +311,11 @@ def forward( feature_function: feature_math.LinearFeature, blksize: int | None, compile_feature_function: bool, - gpu: bool, - *vectors_jvp: torch.Tensor, ) -> torch.Tensor: ngrids = grids.weights.size - block_loop = _make_ao_block_loop(dm, mol, grids, feature_function, blksize, gpu) + block_loop = _make_ao_block_loop( + dm.device, mol, grids, feature_function, blksize + ) features = torch.zeros( *dm.shape[:-2], @@ -346,12 +324,7 @@ def forward( device=dm.device, dtype=dm.dtype, ) - # Raw AO features are linear in dm, so derivatives above first order vanish. - if len(vectors_jvp) > 1: - return features - - evaluation_dm = vectors_jvp[0] if vectors_jvp else dm - evaluation_dm_ordered = block_loop.order_aos(evaluation_dm) + evaluation_dm_ordered = block_loop.order_aos(dm) for block in block_loop: active_dm_submatrix = block.select_active_ao_submatrix( evaluation_dm_ordered @@ -366,75 +339,44 @@ def forward( return features @staticmethod - def jvp( - ctx: _ChunkEvalForwardContext, *grad_inputs: torch.Tensor | None - ) -> torch.Tensor: - if len(ctx.vectors_jvp) > 1: + def jvp(ctx: _ChunkEvalContext, *grad_inputs: torch.Tensor | None) -> torch.Tensor: + dm_tangent = grad_inputs[0] + if dm_tangent is None: return torch.zeros( - *ctx.dm.shape[:-2], + *ctx.spin_shape, ctx.feature_function.nfeats, ctx.grids.weights.size, - device=ctx.dm.device, - dtype=ctx.dm.dtype, + device=ctx.output_device, + dtype=torch.float64, ) - vector_tangent = grad_inputs[7] if ctx.vectors_jvp else grad_inputs[0] - if vector_tangent is None: - return torch.zeros( - *ctx.dm.shape[:-2], - ctx.feature_function.nfeats, - ctx.grids.weights.size, - device=ctx.dm.device, - dtype=ctx.dm.dtype, - ) - return ChunkEvalForward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - vector_tangent, + return cast( + Tensor, + ChunkEvalForward.apply( + dm_tangent, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ), ) @staticmethod def backward( - ctx: _ChunkEvalForwardContext, *grad_outputs: torch.Tensor + ctx: _ChunkEvalContext, *grad_outputs: torch.Tensor ) -> tuple[torch.Tensor | None, ...]: feature_cotangent = grad_outputs[0] - if ctx.vectors_jvp: - dm_grad = ctx.dm * 0 - else: - dm_grad = ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent, - ) - grads: list[Tensor | None] = [dm_grad] - grads += [None] * 6 - - for vector in ctx.vectors_jvp: - if len(ctx.vectors_jvp) == 1: - vector_grad = ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent, - ) - else: - vector_grad = vector * 0 - grads.append(vector_grad) - - return tuple(grads) + dm_cotangent = ChunkEvalBackward.apply( + feature_cotangent, + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, + ) + # PyTorch expects one gradient slot per forward input; the remaining + # arguments are AO-evaluation metadata and are not differentiable. + return dm_cotangent, None, None, None, None, None class ChunkEvalBackward(Function): @@ -448,38 +390,36 @@ def setup_context( feature_math.LinearFeature, int | None, bool, - bool, - torch.Tensor, ], output: torch.Tensor, ) -> None: - context = cast(_ChunkEvalBackwardContext, ctx) + context = cast(_ChunkEvalContext, ctx) ( - context.dm, + feature_cotangent, context.mol, context.grids, context.feature_function, context.blksize, context.compile_feature_function, - context.gpu, - _feature_cotangent, ) = inputs - ctx.save_for_backward(context.dm) + context.spin_shape = feature_cotangent.shape[:-2] + context.output_device = output.device @staticmethod def forward( - dm: torch.Tensor, + feature_cotangent: torch.Tensor, mol: gto.Mole, grids: Grid, feature_function: feature_math.LinearFeature, blksize: int | None, compile_feature_function: bool, - gpu: bool, - feature_cotangent: torch.Tensor, ) -> torch.Tensor: - block_loop = _make_ao_block_loop(dm, mol, grids, feature_function, blksize, gpu) + block_loop = _make_ao_block_loop( + feature_cotangent.device, mol, grids, feature_function, blksize + ) - out = torch.zeros_like(dm) + nao = mol.nao_nr() + out = feature_cotangent.new_zeros(*feature_cotangent.shape[:-2], nao, nao) for block in block_loop: block_result = _evaluate_feature_block( feature_function, @@ -492,42 +432,44 @@ def forward( return block_loop.restore_ao_order(out) @staticmethod - def jvp( - ctx: _ChunkEvalBackwardContext, *grad_inputs: torch.Tensor | None - ) -> torch.Tensor: - feature_cotangent_tangent = grad_inputs[7] + def jvp(ctx: _ChunkEvalContext, *grad_inputs: torch.Tensor | None) -> torch.Tensor: + feature_cotangent_tangent = grad_inputs[0] if feature_cotangent_tangent is None: - return torch.zeros_like(ctx.dm) - return ChunkEvalBackward.apply( - ctx.dm, - ctx.mol, - ctx.grids, - ctx.feature_function, - ctx.blksize, - ctx.compile_feature_function, - ctx.gpu, - feature_cotangent_tangent, - ) - - @staticmethod - def backward( - ctx: _ChunkEvalBackwardContext, *grad_outputs: torch.Tensor - ) -> tuple[torch.Tensor | None, ...]: - grads: list[Tensor | None] = [ctx.dm * 0] - grads += [None] * 6 - grads.append( - ChunkEvalForward.apply( - ctx.dm, + nao = ctx.mol.nao_nr() + return torch.zeros( + *ctx.spin_shape, + nao, + nao, + device=ctx.output_device, + dtype=torch.float64, + ) + return cast( + Tensor, + ChunkEvalBackward.apply( + feature_cotangent_tangent, ctx.mol, ctx.grids, ctx.feature_function, ctx.blksize, ctx.compile_feature_function, - ctx.gpu, - grad_outputs[0], - ) + ), + ) + + @staticmethod + def backward( + ctx: _ChunkEvalContext, *grad_outputs: torch.Tensor + ) -> tuple[torch.Tensor | None, ...]: + feature_cotangent_grad = ChunkEvalForward.apply( + grad_outputs[0], + ctx.mol, + ctx.grids, + ctx.feature_function, + ctx.blksize, + ctx.compile_feature_function, ) - return tuple(grads) + # PyTorch expects one gradient slot per forward input; the remaining + # arguments are AO-evaluation metadata and are not differentiable. + return feature_cotangent_grad, None, None, None, None, None def evaluate_full_grid( @@ -604,6 +546,6 @@ def auto_chunk( features = evaluate_full_grid(dm.double(), mol, grids.coords, feature_function) else: features = ChunkEvalForward.apply( - dm.double(), mol, grids, feature_function, blksize, False, gpu + dm.double(), mol, grids, feature_function, blksize, False ) return feature_function.to_dict(features) diff --git a/src/skala/pyscf/screening.py b/src/skala/pyscf/screening.py index 1359231e..0946faf1 100644 --- a/src/skala/pyscf/screening.py +++ b/src/skala/pyscf/screening.py @@ -174,7 +174,6 @@ def prepare_spatial_grid_layout( def screened_feature_jvp( - dm: Tensor, dm_tangent: Tensor, mol: gto.Mole, spatial_grid_layout: SpatialGridLayout, @@ -185,14 +184,12 @@ def screened_feature_jvp( sorted_tangent = cast( Tensor, ao_evaluation.ChunkEvalForward.apply( - dm, + dm_tangent, mol, spatial_grid_layout.sorted_grids, feature_function, spatial_grid_layout.block_size, compile_feature_function, - dm.device.type == "cuda", - dm_tangent, ), ) return sorted_tangent.index_select( diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index 8f6c64f4..f9dc5806 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -170,7 +170,6 @@ def _integrate_screened( feature_function, spatial_grid_layout.block_size, False, - dm.device.type == "cuda", ), ) atom_major_raw_features = sorted_raw_features.index_select( @@ -270,7 +269,6 @@ def _gen_response_screened( feature_function, spatial_grid_layout.block_size, False, - dm0.device.type == "cuda", ), ) atom_major_raw_features = sorted_raw_features.index_select( @@ -289,7 +287,6 @@ def _gen_response_screened( def hessian_vector_product(dm1: Tensor) -> Tensor: atom_major_tangent = screened_feature_jvp( - dm0, dm1, mol, spatial_grid_layout, diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index b76cfa1a..70ba3557 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -539,7 +539,6 @@ def fake_chunk_eval_forward( return dm.sum().reshape(1, 1) def fake_screened_feature_jvp( - dm: torch.Tensor, dm_tangent: torch.Tensor, mol: gto.Mole, spatial_grid_layout: object, @@ -662,18 +661,20 @@ def test_feature_block_helper_localizes_derivative_vectors() -> None: def test_chunk_eval_transforms_follow_linear_operator(carbon: gto.Mole) -> None: - """Check first and second JVPs and the feature-cotangent adjoint JVP.""" + """Check spin-resolved first and second JVPs and the adjoint JVP.""" grids = _minimal_atom_grid(carbon) feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) - dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + identity = torch.eye(carbon.nao_nr(), dtype=torch.float64) + dm = torch.stack((identity, 2 * identity)) tangent = torch.arange(1, dm.numel() + 1, dtype=dm.dtype).reshape(dm.shape) def evaluate(value: torch.Tensor) -> torch.Tensor: return ChunkEvalForward.apply( # type: ignore[no-untyped-call] - value, carbon, grids, feature_function, None, False, False + value, carbon, grids, feature_function, None, False ) features, feature_tangent = torch.func.jvp(evaluate, (dm,), (tangent,)) + assert features.shape[:1] == dm.shape[:-2] torch.testing.assert_close(feature_tangent, evaluate(tangent)) def first_jvp(value: torch.Tensor) -> torch.Tensor: @@ -689,9 +690,15 @@ def first_jvp(value: torch.Tensor) -> torch.Tensor: def apply_adjoint(value: torch.Tensor) -> torch.Tensor: return ChunkEvalBackward.apply( # type: ignore[no-untyped-call] - dm, carbon, grids, feature_function, None, False, False, value + value, carbon, grids, feature_function, None, False ) + dm_cotangent = apply_adjoint(feature_cotangent) + assert dm_cotangent.shape == dm.shape + torch.testing.assert_close( + torch.sum(features * feature_cotangent), + torch.sum(dm * dm_cotangent), + ) _, adjoint_tangent = torch.func.jvp( apply_adjoint, (feature_cotangent,), @@ -740,7 +747,7 @@ def block_loop( torch.arange(1, carbon.nao_nr() + 1, dtype=torch.float64) ).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm, carbon, grids, feature_function, block_size, False, False + dm, carbon, grids, feature_function, block_size, False ) expected_blocks = [] @@ -808,9 +815,8 @@ def block_loop( monkeypatch.setattr(dft.numint, "NumInt", FakeNumInt) feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY])) - dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) - blocks = list(_CPUAOBlockLoop(dm, carbon, grids, feature_function, block_size)) + blocks = list(_CPUAOBlockLoop(carbon, grids, feature_function, block_size)) assert len(blocks) == 2 assert blocks[0].active_ao_indices is None @@ -846,7 +852,7 @@ def block_loop( dm = torch.eye(carbon.nao_nr(), dtype=torch.float64).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm, carbon, grids, feature_function, ngrids, False, False + dm, carbon, grids, feature_function, ngrids, False ) (vxc,) = torch.autograd.grad(features.square().sum(), dm, create_graph=True) (hvp,) = torch.autograd.grad(vxc, dm, torch.ones_like(dm)) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index e3619d3f..9856ca81 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -323,7 +323,7 @@ def test_gpu_empty_ao_block_matches_dense_reference() -> None: ) dm = torch.eye(mol.nao_nr(), dtype=torch.float64, device="cuda", requires_grad=True) screened = ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm, mol, grids, feature_function, block_size, False, True + dm, mol, grids, feature_function, block_size, False ) dense = evaluate_full_grid(dm, mol, coords, feature_function, gpu=True) @@ -378,7 +378,7 @@ def block_loop( torch.arange(1, mol.nao_nr() + 1, dtype=torch.float64, device="cuda") ).requires_grad_() features = ChunkEvalForward.apply( # type: ignore[no-untyped-call] - dm, mol, grids, feature_function, None, False, True + dm, mol, grids, feature_function, None, False ) sort_idx_t = torch.as_tensor(sort_idx, device="cuda") diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py index 558038c1..6f8c85dd 100644 --- a/tests/test_xc_integrator.py +++ b/tests/test_xc_integrator.py @@ -1,13 +1,78 @@ import pytest import torch from pyscf import dft, gto -from utils import QuadraticFunctional +from utils import QuadraticFunctional, patch_ao_screening from skala.features import Feature, FeatureMap from skala.pyscf import xc_integrator as xc_integrator_module from skala.pyscf.xc_integrator import XCIntegrator, XCResult +def test_screened_xc_derivatives_match_finite_differences() -> None: + """Validate the screened first- and second-order XC derivatives numerically. + + The symmetric density matrix is varied along one symmetric direction as + ``D(t) = D + t P``. The centered energy slope is compared with the analytic + directional derivative ````, checking that the potential returned by + ``XCIntegrator`` is the derivative of the XC energy with respect to the density + matrix. + + The same two perturbed integrations also give a centered derivative of Vxc. Its + ``(0, 0)`` element is compared with the corresponding element of the analytic + Hessian action ``H(D) P`` returned by ``gen_response``. Checking one component + exercises the second-order path without constructing or finite-differencing the + full density-matrix Hessian. + + AO screening is forced so both comparisons cover the custom linear AO autograd + operators. Density, gradient, and kinetic features exercise the meta-GGA paths. + Because those raw features are linear in D and ``QuadraticFunctional`` is + quadratic in the features, the energy is quadratic in D and Vxc is linear; the + centered differences are therefore exact apart from floating-point roundoff. + """ + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + grids = dft.Grids(mol) + grids.level = 0 + grids.alignment = 1 + grids.build(sort_grids=False) + functional = QuadraticFunctional( + [ + Feature.ATOMIC_GRID_SIZES, + Feature.DENSITY, + Feature.GRAD, + Feature.KIN, + Feature.GRID_WEIGHTS, + ] + ) + integrator = XCIntegrator(functional) + dm = torch.tensor([[1.0, 0.2], [0.2, 0.8]], dtype=torch.float64) + direction = torch.tensor([[0.3, -0.2], [-0.2, 0.1]], dtype=dm.dtype) + step = 1e-4 + + with patch_ao_screening(True): + result = integrator(mol, grids, dm) + response = integrator.gen_response(mol, grids, dm.clone()) + plus = integrator(mol, grids, dm + step * direction) + minus = integrator(mol, grids, dm - step * direction) + hessian_action = response(direction) + + energy_slope = (plus.energy - minus.energy) / (2 * step) + potential_directional_derivative = torch.sum(result.potential * direction) + torch.testing.assert_close( + energy_slope, + potential_directional_derivative, + rtol=1e-9, + atol=1e-9, + ) + + potential_slope = (plus.potential[0, 0] - minus.potential[0, 0]) / (2 * step) + torch.testing.assert_close( + potential_slope, + hessian_action[0, 0], + rtol=1e-9, + atol=1e-9, + ) + + def test_xc_integrator_returns_tensors_and_xc_only_response( monkeypatch: pytest.MonkeyPatch, ) -> None: From c56a14a353ba30662ae8fc6cf188b563abf1f94d Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 17:21:04 +0200 Subject: [PATCH 38/39] require special skala grids --- src/skala/gpu4pyscf/dft.py | 51 ++++++----------------- src/skala/gpu4pyscf/grids.py | 57 ++++++++++++++++++++++---- src/skala/pyscf/dft.py | 61 +++++++--------------------- src/skala/pyscf/grids.py | 46 ++++++++++++++++++--- src/skala/pyscf/numint.py | 6 ++- src/skala/pyscf/xc_integrator.py | 42 ++++++++++++++----- tests/test_ao_screening.py | 52 +++++++++++++++++++----- tests/test_ao_screening_benchmark.py | 4 +- tests/test_gpu4pyscf_ao_screening.py | 38 ++++++++++++++++- tests/test_gpu4pyscf_classes.py | 24 +++++++---- tests/test_pyscf_classes.py | 54 +++++++++--------------- tests/test_xc_integrator.py | 17 +++++++- 12 files changed, 284 insertions(+), 168 deletions(-) diff --git a/src/skala/gpu4pyscf/dft.py b/src/skala/gpu4pyscf/dft.py index 0ce237b9..88f5dfaa 100644 --- a/src/skala/gpu4pyscf/dft.py +++ b/src/skala/gpu4pyscf/dft.py @@ -64,8 +64,7 @@ from skala.functional.base import ExcFunctionalBase from skala.gpu4pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient -from skala.gpu4pyscf.grids import UnsortableGrids -from skala.pyscf.dft import _build_grids_unsorted, _needs_unsorted_grids +from skala.gpu4pyscf.grids import SkalaGrids from skala.pyscf.numint import SkalaNumInt from skala.pyscf.utils import pyscf_version_newer_than_2_10 @@ -76,10 +75,10 @@ class SkalaRKS(dft.rks.RKS): # type: ignore[misc] with_dftd3: DFTD3Dispersion | None = None """DFT-D3 dispersion correction.""" - grids: dft.gen_grid.Grids + grids: SkalaGrids """Grids object""" - cphf_grids: dft.gen_grid.Grids + cphf_grids: SkalaGrids """Grids object for CPHF""" def __init__( @@ -94,13 +93,10 @@ def __init__( DFTD3Dispersion(mol, d3) if with_dftd3 and d3 is not None else None ) - self._needs_unsorted = _needs_unsorted_grids(xc) - if self._needs_unsorted: - self.grids = UnsortableGrids(mol)(level=self.grids.level) - self.cphf_grids = UnsortableGrids(mol)( - prune=self.cphf_grids.prune, atom_grid=self.cphf_grids.atom_grid - ) - _build_grids_unsorted(self.grids, mol) + self.grids = SkalaGrids(mol)(level=self.grids.level) + self.cphf_grids = SkalaGrids(mol)( + prune=self.cphf_grids.prune, atom_grid=self.cphf_grids.atom_grid + ) def energy_nuc(self) -> float: enuc = float(super().energy_nuc()) @@ -159,10 +155,10 @@ class SkalaUKS(dft.uks.UKS): # type: ignore[misc] with_dftd3: DFTD3Dispersion | None = None """DFT-D3 dispersion correction.""" - grids: dft.gen_grid.Grids + grids: SkalaGrids """Grids object""" - cphf_grids: dft.gen_grid.Grids + cphf_grids: SkalaGrids """Grids object for CPHF""" def __init__( @@ -177,13 +173,10 @@ def __init__( DFTD3Dispersion(mol, d3) if with_dftd3 and d3 is not None else None ) - self._needs_unsorted = _needs_unsorted_grids(xc) - if self._needs_unsorted: - self.grids = UnsortableGrids(mol)(level=self.grids.level) - self.cphf_grids = UnsortableGrids(mol)( - prune=self.cphf_grids.prune, atom_grid=self.cphf_grids.atom_grid - ) - _build_grids_unsorted(self.grids, mol) + self.grids = SkalaGrids(mol)(level=self.grids.level) + self.cphf_grids = SkalaGrids(mol)( + prune=self.cphf_grids.prune, atom_grid=self.cphf_grids.atom_grid + ) def energy_nuc(self) -> float: enuc = float(super().energy_nuc()) @@ -234,21 +227,3 @@ def density_fit( ks.Gradients = lambda: SkalaUKSGradient(ks) ks.nuc_grad_method = ks.Gradients return cast(SkalaUKS, ks) - - -# GPU4PySCF does not have a initialize_grids method, but a module level function that is called by the RKS and UKS classes. -# We need to monkeypatch this function to ensure that grids are initialized as unsorted when needed. -original_initialize_grids = dft.rks.initialize_grids - - -def initialize_grids( - ks: dft.rks.KohnShamDFT, mol: gto.Mole | None = None, dm: Any = None -) -> dft.rks.KohnShamDFT: - if getattr(ks, "_needs_unsorted", False) and ks.grids.coords is None: - _build_grids_unsorted(ks.grids, mol or ks.mol) - return ks - - return original_initialize_grids(ks, mol, dm) - - -dft.rks.initialize_grids = initialize_grids diff --git a/src/skala/gpu4pyscf/grids.py b/src/skala/gpu4pyscf/grids.py index a0ff6136..b09c9320 100644 --- a/src/skala/gpu4pyscf/grids.py +++ b/src/skala/gpu4pyscf/grids.py @@ -1,19 +1,62 @@ # SPDX-License-Identifier: MIT from logging import getLogger -from typing import Any +from typing import TYPE_CHECKING, Any from gpu4pyscf.dft import gen_grid from pyscf import gto +if TYPE_CHECKING: + from skala.pyscf.screening import SpatialGridLayout + LOG = getLogger(__name__) -class UnsortableGrids(gen_grid.Grids): # type: ignore +class SkalaGrids(gen_grid.Grids): # type: ignore + """GPU4PySCF grids with atom-major ordering and Skala layout caching.""" + + _spatial_grid_layout: "SpatialGridLayout | None" + _initializing: bool + + def __init__(self, mol: gto.Mole | None = None) -> None: + super().__setattr__("_initializing", True) + super().__init__(mol) + super().__setattr__("alignment", 1) + super().__setattr__("_initializing", False) + + def __setattr__(self, key: str, value: Any) -> None: + if ( + key == "alignment" + and value != 1 + and not getattr(self, "_initializing", False) + ): + raise ValueError(f"SkalaGrids alignment must be 1, got {value}") + if key in {"coords", "weights", "cutoff"}: + super().__setattr__("_spatial_grid_layout", None) + super().__setattr__(key, value) + def build( - self, mol: gto.Mole | None = None, with_non0tab: bool = False, **kwargs: Any - ) -> "UnsortableGrids": - sort_grids = kwargs.pop("sort_grids", None) - if sort_grids: + self, + mol: gto.Mole | None = None, + with_non0tab: bool = False, + sort_grids: bool = True, + sort_grids_of_each_atom: bool = False, + **kwargs: Any, + ) -> "SkalaGrids": + if sort_grids or sort_grids_of_each_atom: LOG.debug("sorted grids not supported, forcing unsorted grids") - return super().build(mol, with_non0tab, sort_grids=False, **kwargs) + return super().build( + mol, + with_non0tab, + sort_grids=False, + sort_grids_of_each_atom=False, + **kwargs, + ) + + def get_cached_spatial_grid_layout(self) -> "SpatialGridLayout | None": + """Return the spatial layout cached for the current grid state.""" + return getattr(self, "_spatial_grid_layout", None) + + def cache_spatial_grid_layout(self, layout: "SpatialGridLayout") -> None: + """Cache a spatial layout until layout-defining grid state changes.""" + self._spatial_grid_layout = layout diff --git a/src/skala/pyscf/dft.py b/src/skala/pyscf/dft.py index 499bf6f8..7a5927e7 100644 --- a/src/skala/pyscf/dft.py +++ b/src/skala/pyscf/dft.py @@ -50,7 +50,6 @@ """ -import logging import warnings from collections.abc import Callable from typing import Any, cast @@ -62,46 +61,18 @@ from pyscf.df import df_jk from skala.functional.base import ExcFunctionalBase -from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient -from skala.pyscf.grids import UnsortableGrids +from skala.pyscf.grids import SkalaGrids from skala.pyscf.numint import SkalaNumInt from skala.pyscf.utils import pyscf_version_newer_than_2_10 -logger = logging.getLogger(__name__) - - -def _needs_unsorted_grids(func: ExcFunctionalBase) -> bool: - """Return True when the functional needs per-atom grid ordering.""" - return FeatureSpec(func.features).requires_atomic_layout - - -def _build_grids_unsorted( - grids: dft.gen_grid.Grids, mol: gto.Mole -) -> dft.gen_grid.Grids: - """Build grids without sorting, preserving per-atom ordering. - - Also disables grid alignment padding, which would introduce extra - zero-weight grid points that are not accounted for in the per-atom - grid size decomposition used by the Skala functional. - """ - if grids.alignment != 1: - logger.debug( - "Overriding grids.alignment from %d to 1. " - "The Skala functional requires unsorted, unpadded grids.", - grids.alignment, - ) - grids.alignment = 1 - grids.build(mol, sort_grids=False) - return grids - class SkalaRKS(dft.rks.RKS): # type: ignore[misc] """Restricted Kohn-Sham method with support for Skala functional.""" xc: str - grids: dft.gen_grid.Grids + grids: SkalaGrids """Numerical integration grids.""" with_dftd3: DFTD3Dispersion | None = None @@ -118,23 +89,22 @@ def __init__( super().__init__(mol, xc="custom") self._keys.add("with_dftd3") self._numint = SkalaNumInt(xc, device=device or torch.device("cpu")) - self._needs_unsorted = _needs_unsorted_grids(xc) d3 = xc.get_d3_settings() self.with_dftd3 = ( DFTD3Dispersion(mol, d3) if with_dftd3 and d3 is not None else None ) - if self._needs_unsorted: - self.grids = UnsortableGrids(mol)(level=self.grids.level) - _build_grids_unsorted(self.grids, mol) + self.grids = SkalaGrids(mol)(level=self.grids.level) def initialize_grids( self, mol: gto.Mole | None = None, dm: np.ndarray | None = None ) -> "SkalaRKS": - # Ensure grids stay unsorted even if user changed grid settings after __init__ - if self._needs_unsorted and self.grids.coords is None: - _build_grids_unsorted(self.grids, mol or self.mol) + if not isinstance(self.grids, SkalaGrids): + raise TypeError( + "SkalaRKS requires skala.pyscf.grids.SkalaGrids, got " + f"{type(self.grids).__module__}.{type(self.grids).__name__}" + ) return super().initialize_grids(mol or self.mol, dm) def energy_nuc(self) -> float: @@ -190,7 +160,7 @@ class SkalaUKS(dft.uks.UKS): # type: ignore[misc] xc: str - grids: dft.gen_grid.Grids + grids: SkalaGrids """Numerical integration grids.""" with_dftd3: DFTD3Dispersion | None = None @@ -207,23 +177,22 @@ def __init__( super().__init__(mol, xc="custom") self._keys.add("with_dftd3") self._numint = SkalaNumInt(xc, device=device or torch.device("cpu")) - self._needs_unsorted = _needs_unsorted_grids(xc) d3 = xc.get_d3_settings() self.with_dftd3 = ( DFTD3Dispersion(mol, d3) if with_dftd3 and d3 is not None else None ) - if self._needs_unsorted: - self.grids = UnsortableGrids(mol)(level=self.grids.level) - _build_grids_unsorted(self.grids, mol) + self.grids = SkalaGrids(mol)(level=self.grids.level) def initialize_grids( self, mol: gto.Mole | None = None, dm: np.ndarray | None = None ) -> "SkalaUKS": - # Ensure grids stay unsorted even if user changed grid settings after __init__ - if self._needs_unsorted and self.grids.coords is None: - _build_grids_unsorted(self.grids, mol or self.mol) + if not isinstance(self.grids, SkalaGrids): + raise TypeError( + "SkalaUKS requires skala.pyscf.grids.SkalaGrids, got " + f"{type(self.grids).__module__}.{type(self.grids).__name__}" + ) return super().initialize_grids(mol or self.mol, dm) def energy_nuc(self) -> float: diff --git a/src/skala/pyscf/grids.py b/src/skala/pyscf/grids.py index 8d129ea1..bb752a9a 100644 --- a/src/skala/pyscf/grids.py +++ b/src/skala/pyscf/grids.py @@ -1,19 +1,55 @@ # SPDX-License-Identifier: MIT from logging import getLogger -from typing import Any +from typing import TYPE_CHECKING, Any from pyscf import gto from pyscf.dft import gen_grid +if TYPE_CHECKING: + from skala.pyscf.screening import SpatialGridLayout + LOG = getLogger(__name__) -class UnsortableGrids(gen_grid.Grids): # type: ignore +class SkalaGrids(gen_grid.Grids): # type: ignore + """PySCF grids with atom-major ordering and Skala layout caching.""" + + _spatial_grid_layout: "SpatialGridLayout | None" + _initializing: bool + + def __init__(self, mol: gto.Mole | None = None) -> None: + super().__setattr__("_initializing", True) + super().__init__(mol) + super().__setattr__("alignment", 1) + super().__setattr__("_initializing", False) + + def __setattr__(self, key: str, value: Any) -> None: + if ( + key == "alignment" + and value != 1 + and not getattr(self, "_initializing", False) + ): + raise ValueError(f"SkalaGrids alignment must be 1, got {value}") + if key in {"coords", "weights", "cutoff"}: + super().__setattr__("_spatial_grid_layout", None) + super().__setattr__(key, value) + def build( - self, mol: gto.Mole | None = None, with_non0tab: bool = False, **kwargs: Any - ) -> "UnsortableGrids": - sort_grids = kwargs.pop("sort_grids", None) + self, + mol: gto.Mole | None = None, + with_non0tab: bool = False, + sort_grids: bool = True, + **kwargs: Any, + ) -> "SkalaGrids": if sort_grids: LOG.debug("sorted grids not supported, forcing unsorted grids") return super().build(mol, with_non0tab, sort_grids=False, **kwargs) + + def get_cached_spatial_grid_layout(self) -> "SpatialGridLayout | None": + """Return the spatial layout cached for the current grid state.""" + return getattr(self, "_spatial_grid_layout", None) + + def cache_spatial_grid_layout(self, layout: "SpatialGridLayout") -> None: + """Cache a spatial layout until layout-defining grid state changes.""" + self._spatial_grid_layout = layout diff --git a/src/skala/pyscf/numint.py b/src/skala/pyscf/numint.py index 12ef9bcb..dbee4cb9 100644 --- a/src/skala/pyscf/numint.py +++ b/src/skala/pyscf/numint.py @@ -94,13 +94,15 @@ class SkalaNumInt(PySCFNumInt[Array]): ------- >>> from pyscf import gto, dft >>> from skala.functional import load_functional + >>> from skala.pyscf.grids import SkalaGrids >>> from skala.pyscf.numint import SkalaNumInt >>> >>> mol = gto.M(atom="H 0 0 0; H 0 0 1", basis="def2-svp", verbose=0) >>> ks = dft.KS(mol) >>> ks._numint = SkalaNumInt(load_functional("skala-1.1")) - >>> ks.grids.build(mol, sort_grids=False) # DOCTEST: Ellipsis - + >>> ks.grids = SkalaGrids(mol) + >>> ks.grids.build(mol) # DOCTEST: Ellipsis + >>> energy = ks.kernel() >>> print(energy) # DOCTEST: Ellipsis -1.1425799... diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py index f9dc5806..a7e223db 100644 --- a/src/skala/pyscf/xc_integrator.py +++ b/src/skala/pyscf/xc_integrator.py @@ -3,7 +3,7 @@ """Tensor-level exchange-correlation integration.""" from collections.abc import Callable -from typing import NamedTuple, cast +from typing import NamedTuple, Protocol, cast import torch from pyscf import gto @@ -16,6 +16,7 @@ from skala.pyscf.backend import Grid, check_gpu_imports_were_successful from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec from skala.pyscf.features import generate_features +from skala.pyscf.grids import SkalaGrids as PySCFSkalaGrids from skala.pyscf.model_chunking import prepare_model_feature_chunks from skala.pyscf.screening import ( CPU_AO_SCREENING_BLOCK_SIZE, @@ -25,6 +26,12 @@ ) +class _SpatialGridCache(Protocol): + def get_cached_spatial_grid_layout(self) -> SpatialGridLayout | None: ... + + def cache_spatial_grid_layout(self, layout: SpatialGridLayout) -> None: ... + + def _should_screen_aos(mol: gto.Mole) -> bool: """Return whether PySCF's sparse-contraction crossover is exceeded.""" # we use a smaller threshold because for MetaGGAs the AO evaluation is more expensive @@ -84,6 +91,7 @@ def __call__( ) -> XCResult: """Evaluate electron count, XC energy, and XC potential.""" self._validate_device(dm) + self._require_skala_grids(grids) if self.feature_spec.supports_spatial_decomposition and _should_screen_aos(mol): return self._integrate_screened(mol, grids, dm, max_memory) return self._integrate_dense(mol, grids, dm, max_memory) @@ -98,6 +106,7 @@ def gen_response( ) -> Callable[[Tensor], Tensor]: """Build an XC-only Hessian-vector product callable.""" self._validate_device(dm0) + self._require_skala_grids(grids) if self.feature_spec.supports_spatial_decomposition and _should_screen_aos(mol): return self._gen_response_screened( mol, @@ -118,16 +127,30 @@ def _validate_device(self, dm: Tensor) -> None: f"Density matrix device {dm.device} does not match functional device {self.device}" ) + def _require_skala_grids(self, grids: Grid) -> _SpatialGridCache: + if self.device.type == "cuda": + check_gpu_imports_were_successful() + from skala.gpu4pyscf.grids import SkalaGrids as GPU4PySCFSkalaGrids + + expected_type = GPU4PySCFSkalaGrids + else: + expected_type = PySCFSkalaGrids + + if not isinstance(grids, expected_type): + raise TypeError( + f"{self.device.type.upper()} Skala XC evaluation requires " + f"{expected_type.__module__}.{expected_type.__name__}, got " + f"{type(grids).__module__}.{type(grids).__name__}" + ) + return cast(_SpatialGridCache, grids) + def _get_spatial_grid_layout( self, mol: gto.Mole, grids: Grid, ) -> SpatialGridLayout: - grid_state = vars(grids) - spatial_grid_layout = cast( - SpatialGridLayout | None, - grid_state.get("_skala_spatial_grid_layout"), - ) + grid_cache = self._require_skala_grids(grids) + spatial_grid_layout = grid_cache.get_cached_spatial_grid_layout() if spatial_grid_layout is not None: return spatial_grid_layout @@ -140,12 +163,9 @@ def _get_spatial_grid_layout( block_size = CPU_AO_SCREENING_BLOCK_SIZE spatial_grid_layout = prepare_spatial_grid_layout( - mol, - grids, - block_size, - self.device, + mol, grids, block_size, self.device ) - grid_state["_skala_spatial_grid_layout"] = spatial_grid_layout + grid_cache.cache_spatial_grid_layout(spatial_grid_layout) return spatial_grid_layout def _integrate_screened( diff --git a/tests/test_ao_screening.py b/tests/test_ao_screening.py index 70ba3557..1c0fba21 100644 --- a/tests/test_ao_screening.py +++ b/tests/test_ao_screening.py @@ -24,6 +24,7 @@ ) from skala.pyscf.evaluation import FeatureSpec from skala.pyscf.feature_math import MGGAFeatureFunction +from skala.pyscf.grids import SkalaGrids from skala.pyscf.model_chunking import ModelFeatureChunk from skala.pyscf.numint import SkalaNumInt from skala.pyscf.screening import ( @@ -31,6 +32,7 @@ _decompose_grid_into_spatial_blocks, prepare_spatial_grid_layout, ) +from skala.pyscf.xc_integrator import XCIntegrator _MGGA_FEATURES = (Feature.DENSITY, Feature.GRAD, Feature.KIN, Feature.LAPL) _MGGA_FEATURE_COMBINATIONS = [ @@ -395,17 +397,17 @@ def fake_make_screen_index( assert screen_index_calls == 1 assert decomposition_block_sizes == [2] assert screened_molecules == [carbon] - assert not hasattr(grids, "_skala_spatial_grid_layout") + assert not hasattr(grids, "_spatial_grid_layout") def test_grid_reuses_spatial_grid_layout_across_numints( carbon: gto.Mole, monkeypatch: pytest.MonkeyPatch ) -> None: """Cache one spatial layout on each grid independently of the NumInt instance.""" - grids = dft.Grids(carbon) + grids = SkalaGrids(carbon) grids.coords = np.arange(18, dtype=np.float64).reshape(6, 3) grids.weights = np.arange(6, dtype=np.float64) - other_grids = dft.Grids(carbon) + other_grids = SkalaGrids(carbon) other_grids.coords = grids.coords.copy() other_grids.weights = grids.weights.copy() layouts: list[SpatialGridLayout] = [] @@ -435,7 +437,7 @@ def fake_prepare_spatial_grid_layout( layout = numint.integrator._get_spatial_grid_layout(carbon, grids) assert other_numint.integrator._get_spatial_grid_layout(carbon, grids) is layout - assert vars(grids)["_skala_spatial_grid_layout"] is layout + assert grids.get_cached_spatial_grid_layout() is layout assert len(layouts) == 1 numint.reset() @@ -443,7 +445,7 @@ def fake_prepare_spatial_grid_layout( other_layout = numint.integrator._get_spatial_grid_layout(carbon, other_grids) assert other_layout is not layout - assert vars(other_grids)["_skala_spatial_grid_layout"] is other_layout + assert other_grids.get_cached_spatial_grid_layout() is other_layout assert len(layouts) == 2 @@ -585,7 +587,7 @@ def fake_prepare_model_feature_chunks( ) numint = SkalaNumInt(QuadraticFunctional()) dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) - grids = dft.Grids(carbon) + grids = SkalaGrids(carbon) grids.weights = np.ones(1) ks = FakeKS(carbon, grids) @@ -865,13 +867,43 @@ def block_loop( assert torch.count_nonzero(hvp) == 0 -def _minimal_atom_grid(mol: gto.Mole) -> dft.Grids: - grids = dft.Grids(mol) +def _minimal_atom_grid(mol: gto.Mole) -> SkalaGrids: + grids = SkalaGrids(mol) grids.level = 0 grids.alignment = 1 return grids.build(sort_grids=False) +def test_atom_major_features_require_skala_grids(carbon: gto.Mole) -> None: + integrator = XCIntegrator(QuadraticFunctional()) + grids = dft.Grids(carbon) + dm = torch.eye(carbon.nao_nr(), dtype=torch.float64) + + with pytest.raises(TypeError, match=r"requires .*\.SkalaGrids"): + integrator(carbon, grids, dm) + with pytest.raises(TypeError, match=r"requires .*\.SkalaGrids"): + integrator.gen_response(carbon, grids, dm) + + +def test_skala_grids_invalidate_spatial_layout(carbon: gto.Mole) -> None: + integrator = XCIntegrator(QuadraticFunctional()) + grids = _minimal_atom_grid(carbon) + + layout = integrator._get_spatial_grid_layout(carbon, grids) + assert integrator._get_spatial_grid_layout(carbon, grids) is layout + + grids.reset() + assert grids.get_cached_spatial_grid_layout() is None + grids.level = 0 + grids.alignment = 1 + grids.build(sort_grids=False) + rebuilt_layout = integrator._get_spatial_grid_layout(carbon, grids) + assert rebuilt_layout is not layout + + grids.cutoff /= 10 + assert grids.get_cached_spatial_grid_layout() is None + + def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> None: numint = SkalaNumInt(QuadraticFunctional()) grids = _minimal_atom_grid(carbon) @@ -881,10 +913,10 @@ def test_numint_reset_does_not_clear_grid_spatial_layout(carbon: gto.Mole) -> No block_size=dft.gen_grid.BLKSIZE, device=torch.device("cpu"), ) - vars(grids)["_skala_spatial_grid_layout"] = spatial_grid_layout + grids.cache_spatial_grid_layout(spatial_grid_layout) assert numint.reset() is numint - assert vars(grids)["_skala_spatial_grid_layout"] is spatial_grid_layout + assert grids.get_cached_spatial_grid_layout() is spatial_grid_layout @pytest.mark.parametrize( diff --git a/tests/test_ao_screening_benchmark.py b/tests/test_ao_screening_benchmark.py index 40d60b3a..5b51d707 100644 --- a/tests/test_ao_screening_benchmark.py +++ b/tests/test_ao_screening_benchmark.py @@ -25,6 +25,7 @@ from skala.functional import load_functional from skala.functional.base import ExcFunctionalBase +from skala.pyscf.grids import SkalaGrids from skala.pyscf.numint import SkalaNumInt THREAD_COUNT = 4 @@ -142,8 +143,9 @@ def _make_benchmark_case( initial_dm = dft.RKS(mol).get_init_guess() if backend == "cpu": - grids = dft.Grids(mol) + grids = SkalaGrids(mol) grids.level = 1 + grids.alignment = 1 grids.build(sort_grids=False) dm: Any = initial_dm numint: Any = SkalaNumInt(functional) diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py index 9856ca81..73b24c51 100644 --- a/tests/test_gpu4pyscf_ao_screening.py +++ b/tests/test_gpu4pyscf_ao_screening.py @@ -28,6 +28,7 @@ from skala.features import Feature # noqa: E402 from skala.functional.base import ExcFunctionalBase # noqa: E402 from skala.gpu4pyscf import SkalaKS # noqa: E402 +from skala.gpu4pyscf.grids import SkalaGrids as GPU4PySCFSkalaGrids # noqa: E402 from skala.pyscf.ao_evaluation import ( # noqa: E402 ChunkEvalForward, evaluate_full_grid, @@ -35,8 +36,10 @@ from skala.pyscf.backend import dft_gpu # noqa: E402 from skala.pyscf.evaluation import FeatureSpec # noqa: E402 from skala.pyscf.feature_math import MGGAFeatureFunction # noqa: E402 +from skala.pyscf.grids import SkalaGrids as PySCFSkalaGrids # noqa: E402 from skala.pyscf.numint import SkalaNumInt # noqa: E402 from skala.pyscf.screening import prepare_spatial_grid_layout # noqa: E402 +from skala.pyscf.xc_integrator import XCIntegrator # noqa: E402 CARBON_CHAIN = """ C 0.0 0.0 0.0 @@ -84,6 +87,37 @@ def test_prepare_spatially_sorted_gpu_grids() -> None: assert layout.inverse_permutation.device.type == "cuda" +def test_gpu_atom_major_features_require_skala_grids() -> None: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + grids = dft_gpu.Grids(mol) + integrator = XCIntegrator(QuadraticFunctional(), device=torch.device("cuda:0")) + dm = torch.eye(mol.nao_nr(), dtype=torch.float64, device="cuda:0") + + with pytest.raises(TypeError, match=r"requires .*\.SkalaGrids"): + integrator(mol, grids, dm) + with pytest.raises(TypeError, match=r"requires .*\.SkalaGrids"): + integrator.gen_response(mol, grids, dm) + + +def test_gpu_skala_grids_invalidate_spatial_layout() -> None: + mol = gto.M(atom="H 0 0 0", basis="sto-3g", spin=1, verbose=0) + grids = GPU4PySCFSkalaGrids(mol) + grids.level = 0 + grids.alignment = 1 + grids.build() + integrator = XCIntegrator(QuadraticFunctional(), device=torch.device("cuda:0")) + + layout = integrator._get_spatial_grid_layout(mol, grids) + assert integrator._get_spatial_grid_layout(mol, grids) is layout + + grids.reset() + assert grids.get_cached_spatial_grid_layout() is None + grids.level = 0 + grids.alignment = 1 + grids.build() + assert integrator._get_spatial_grid_layout(mol, grids) is not layout + + @pytest.mark.parametrize( ("atom", "spin", "integration_method_name"), [ @@ -226,11 +260,11 @@ def test_gpu_screened_skala_matches_cpu_on_carbon_chain( sensitive to omitted derivative contributions. """ mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0) - cpu_grids = dft.Grids(mol) + cpu_grids = PySCFSkalaGrids(mol) cpu_grids.level = 1 cpu_grids.alignment = 1 cpu_grids.build(sort_grids=False) - gpu_grids = dft_gpu.Grids(mol) + gpu_grids = GPU4PySCFSkalaGrids(mol) gpu_grids.level = 1 gpu_grids.alignment = 1 gpu_grids.build(sort_grids=False) diff --git a/tests/test_gpu4pyscf_classes.py b/tests/test_gpu4pyscf_classes.py index a46f7cbd..ad874d8a 100644 --- a/tests/test_gpu4pyscf_classes.py +++ b/tests/test_gpu4pyscf_classes.py @@ -16,7 +16,7 @@ from skala.gpu4pyscf import SkalaKS # noqa: E402 from skala.gpu4pyscf.dft import SkalaRKS, SkalaUKS # noqa: E402 from skala.gpu4pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient # noqa: E402 -from skala.gpu4pyscf.grids import UnsortableGrids # noqa: E402 +from skala.gpu4pyscf.grids import SkalaGrids # noqa: E402 @pytest.fixture(params=["skala-1.0", "skala-1.1"]) @@ -73,8 +73,7 @@ def test_skala_class( assert ks.xc == "custom" assert isinstance(ks, SkalaRKS if mol.spin == 0 else SkalaUKS) assert ks.with_dftd3 is not None if with_dftd3 else ks.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(ks.grids, UnsortableGrids) + assert isinstance(ks.grids, SkalaGrids) ks_scanner = ks.as_scanner() assert isinstance(ks_scanner, SkalaRKS if mol.spin == 0 else SkalaUKS) @@ -87,17 +86,24 @@ def test_skala_class( grad = ks.nuc_grad_method() assert isinstance(grad, SkalaRKSGradient if mol.spin == 0 else SkalaUKSGradient) assert grad.with_dftd3 is not None if with_dftd3 else grad.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(grad.grids, UnsortableGrids) + assert isinstance(grad.grids, SkalaGrids) grad = ks.Gradients() assert isinstance(grad, SkalaRKSGradient if mol.spin == 0 else SkalaUKSGradient) assert grad.with_dftd3 is not None if with_dftd3 else grad.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(grad.grids, UnsortableGrids) + assert isinstance(grad.grids, SkalaGrids) ks = grad.base assert isinstance(ks, SkalaRKS if mol.spin == 0 else SkalaUKS) assert ks.with_dftd3 is not None if with_dftd3 else ks.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(ks.grids, UnsortableGrids) + assert isinstance(ks.grids, SkalaGrids) + + +def test_skala_grids_require_unit_alignment() -> None: + mol = gto.M(atom="H", basis="sto-3g", spin=1, verbose=0) + grids = SkalaGrids(mol) + + assert grids.alignment == 1 + grids.alignment = 1 + with pytest.raises(ValueError, match="alignment must be 1"): + grids.alignment = 256 diff --git a/tests/test_pyscf_classes.py b/tests/test_pyscf_classes.py index 8b6c6756..020fbc5c 100644 --- a/tests/test_pyscf_classes.py +++ b/tests/test_pyscf_classes.py @@ -1,13 +1,13 @@ from collections.abc import Callable import pytest -from pyscf import gto +from pyscf import dft, gto from skala.functional.base import ExcFunctionalBase from skala.pyscf import SkalaKS from skala.pyscf.dft import SkalaRKS, SkalaUKS from skala.pyscf.gradients import SkalaRKSGradient, SkalaUKSGradient -from skala.pyscf.grids import UnsortableGrids +from skala.pyscf.grids import SkalaGrids @pytest.fixture(params=["skala-1.0", "skala-1.1"]) @@ -64,8 +64,7 @@ def test_skala_class( assert ks.xc == "custom" assert isinstance(ks, SkalaRKS if mol.spin == 0 else SkalaUKS) assert ks.with_dftd3 is not None if with_dftd3 else ks.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(ks.grids, UnsortableGrids) + assert isinstance(ks.grids, SkalaGrids) ks_scanner = ks.as_scanner() assert isinstance(ks_scanner, SkalaRKS if mol.spin == 0 else SkalaUKS) @@ -78,20 +77,17 @@ def test_skala_class( grad = ks.nuc_grad_method() assert isinstance(grad, SkalaRKSGradient if mol.spin == 0 else SkalaUKSGradient) assert grad.with_dftd3 is not None if with_dftd3 else grad.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(grad.grids, UnsortableGrids) + assert isinstance(grad.grids, SkalaGrids) grad = ks.Gradients() assert isinstance(grad, SkalaRKSGradient if mol.spin == 0 else SkalaUKSGradient) assert grad.with_dftd3 is not None if with_dftd3 else grad.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(grad.grids, UnsortableGrids) + assert isinstance(grad.grids, SkalaGrids) ks = grad.base assert isinstance(ks, SkalaRKS if mol.spin == 0 else SkalaUKS) assert ks.with_dftd3 is not None if with_dftd3 else ks.with_dftd3 is None - if ks._needs_unsorted: - assert isinstance(ks.grids, UnsortableGrids) + assert isinstance(ks.grids, SkalaGrids) def test_skala_class_with_dftd3_and_native_functional_raises() -> None: @@ -111,34 +107,22 @@ def test_skala_class_with_native_functional_and_no_dftd3_is_allowed() -> None: assert not isinstance(ks, (SkalaRKS, SkalaUKS)) -def test_grid_alignment_mismatch_raises( - load_functional_cached: Callable[..., ExcFunctionalBase | str], +def test_initialize_grids_rejects_non_skala_grids( + skala_xc: ExcFunctionalBase, ) -> None: - """generate_features raises ValueError when grid has alignment padding.""" - from unittest.mock import patch - - import torch - - from skala.pyscf.features import generate_features - mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) - func = load_functional_cached("skala-1.1") - assert not isinstance(func, str) - - def _build_grids_keep_padding(grids: gto.Mole, mol: gto.Mole) -> gto.Mole: - """Build grids WITHOUT disabling alignment, so padding is preserved.""" - grids.build(mol, sort_grids=False) - return grids + ks = SkalaRKS(mol, xc=skala_xc) + ks.grids = dft.gen_grid.Grids(mol) - with patch("skala.pyscf.dft._build_grids_unsorted", _build_grids_keep_padding): - ks = SkalaKS(mol, xc=func, with_dftd3=False) + with pytest.raises(TypeError, match="SkalaRKS requires .*SkalaGrids"): + ks.initialize_grids() - # The default PySCF alignment is 8, so grids may have padding. - # Force alignment to something large to guarantee a mismatch. - ks.grids.alignment = 128 - ks.grids.build(mol, sort_grids=False) - dm = torch.from_numpy(ks.get_init_guess()) +def test_skala_grids_require_unit_alignment() -> None: + mol = gto.M(atom="H", basis="sto-3g", spin=1, verbose=0) + grids = SkalaGrids(mol) - with pytest.raises(ValueError, match="Grid size mismatch"): - generate_features(mol, dm, ks.grids, set(func.features)) + assert grids.alignment == 1 + grids.alignment = 1 + with pytest.raises(ValueError, match="alignment must be 1"): + grids.alignment = 8 diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py index 6f8c85dd..b427ef8e 100644 --- a/tests/test_xc_integrator.py +++ b/tests/test_xc_integrator.py @@ -5,6 +5,7 @@ from skala.features import Feature, FeatureMap from skala.pyscf import xc_integrator as xc_integrator_module +from skala.pyscf.grids import SkalaGrids from skala.pyscf.xc_integrator import XCIntegrator, XCResult @@ -30,7 +31,7 @@ def test_screened_xc_derivatives_match_finite_differences() -> None: centered differences are therefore exact apart from floating-point roundoff. """ mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) - grids = dft.Grids(mol) + grids = SkalaGrids(mol) grids.level = 0 grids.alignment = 1 grids.build(sort_grids=False) @@ -84,7 +85,7 @@ def test_xc_integrator_returns_tensors_and_xc_only_response( the Coulomb response that the higher-level NumInt wrapper adds. """ mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) - grids = dft.Grids(mol) + grids = SkalaGrids(mol) def fake_generate_features( mol: gto.Mole, @@ -115,3 +116,15 @@ def fake_generate_features( torch.testing.assert_close(result.energy, dm.new_tensor(128.0)) torch.testing.assert_close(result.potential, torch.full_like(dm, 32.0)) torch.testing.assert_close(response(torch.ones_like(dm)), torch.full_like(dm, 16.0)) + + +def test_xc_integrator_requires_skala_grids_for_density() -> None: + mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + grids = dft.Grids(mol) + integrator = XCIntegrator(QuadraticFunctional([Feature.DENSITY])) + dm = torch.eye(mol.nao_nr(), dtype=torch.float64) + + with pytest.raises(TypeError, match=r"XC evaluation requires .*\.SkalaGrids"): + integrator(mol, grids, dm) + with pytest.raises(TypeError, match=r"XC evaluation requires .*\.SkalaGrids"): + integrator.gen_response(mol, grids, dm) From eddd046a8a7f1e9b0ff564b1481ce84a34381ab5 Mon Sep 17 00:00:00 2001 From: jenswehner Date: Mon, 10 Aug 2026 21:19:35 +0200 Subject: [PATCH 39/39] add pyscf 2.9 fixes --- src/skala/pyscf/dft.py | 2 ++ tests/test_pyscf_classes.py | 12 ++++++++++++ 2 files changed, 14 insertions(+) diff --git a/src/skala/pyscf/dft.py b/src/skala/pyscf/dft.py index 7a5927e7..aedc1bb8 100644 --- a/src/skala/pyscf/dft.py +++ b/src/skala/pyscf/dft.py @@ -89,6 +89,7 @@ def __init__( super().__init__(mol, xc="custom") self._keys.add("with_dftd3") self._numint = SkalaNumInt(xc, device=device or torch.device("cpu")) + self.small_rho_cutoff = 0 # pyscf 2.9 default is 1e-7 d3 = xc.get_d3_settings() self.with_dftd3 = ( @@ -177,6 +178,7 @@ def __init__( super().__init__(mol, xc="custom") self._keys.add("with_dftd3") self._numint = SkalaNumInt(xc, device=device or torch.device("cpu")) + self.small_rho_cutoff = 0 # pyscf 2.9 default is 1e-7 d3 = xc.get_d3_settings() self.with_dftd3 = ( diff --git a/tests/test_pyscf_classes.py b/tests/test_pyscf_classes.py index 020fbc5c..b3979a02 100644 --- a/tests/test_pyscf_classes.py +++ b/tests/test_pyscf_classes.py @@ -126,3 +126,15 @@ def test_skala_grids_require_unit_alignment() -> None: grids.alignment = 1 with pytest.raises(ValueError, match="alignment must be 1"): grids.alignment = 8 + + +def test_skala_classes_disable_density_grid_pruning( + monkeypatch: pytest.MonkeyPatch, + skala_xc: ExcFunctionalBase, +) -> None: + monkeypatch.setattr(dft.rks.KohnShamDFT, "small_rho_cutoff", 1e-7) + rks_mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0) + uks_mol = gto.M(atom="H", basis="sto-3g", spin=1, verbose=0) + + assert SkalaRKS(rks_mol, xc=skala_xc).small_rho_cutoff == 0 + assert SkalaUKS(uks_mol, xc=skala_xc).small_rho_cutoff == 0