diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml
index a20d0b8d..6f402740 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.13.1)"
+ 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.13.1
+
+ - 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/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..6dcb60a7
--- /dev/null
+++ b/benchmarks/pyscf_ao_screening_performance.ipynb
@@ -0,0 +1,820 @@
+{
+ "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 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."
+ ]
+ },
+ {
+ "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",
+ "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",
+ "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}\")\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}\")"
+ ]
+ },
+ {
+ "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",
+ "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",
+ " 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(\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",
+ " 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\": MARKER_SIZE,\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\": MARKER_SIZE,\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=MARKER_SIZE,\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=MARKER_SIZE,\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",
+ " 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",
+ " 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..749e9ef4
--- /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(
+ route_owner.feature_spec.supports_spatial_decomposition
+ )
+ 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'| {label} | '\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'| {text} | '\n",
+ " )\n",
+ " parts.append(\"
\")\n",
+ " parts.append(\"
\")\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
+}
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 0d415357..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
@@ -15,11 +16,14 @@ dependencies:
- pytorch * cpu_*
- qcelemental
# Testing and development
+ - memray
- pre-commit
- pytest
+ - pytest-benchmark
- pytest-cov
- pytest-randomly
+ - pytest-timeout
- ruff
- - mypy
+ - mypy >=2
- pip:
- huggingface_hub
diff --git a/environment-gpu.yml b/environment-gpu.yml
index 2abb6ad1..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
@@ -20,10 +21,12 @@ dependencies:
# Testing and development
- pre-commit
- pytest
+ - pytest-benchmark
- pytest-cov
- 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/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/pyproject.toml b/pyproject.toml
index bfb8fc26..0d389609 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -31,10 +31,13 @@ dependencies = [
optional-dependencies.dev = [
"pre-commit",
- "mypy",
+ "mypy>=2",
+ "memray",
"pytest",
+ "pytest-benchmark",
"pytest-cov",
"pytest-randomly",
+ "pytest-timeout",
]
optional-dependencies.doc = [
"ipywidgets",
@@ -59,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.screening"]
+disable_error_code = ["no-untyped-call"]
[tool.ruff]
target-version = "py311"
@@ -94,6 +96,13 @@ detect-same-package = true
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 = [
"error",
# PyTorch 2.11 deprecated `torch.jit.load`; Skala's pretrained checkpoints
@@ -116,6 +125,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/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/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"]
diff --git a/src/skala/functional/__init__.py b/src/skala/functional/__init__.py
index 36c64bc3..b31efae7 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",
]
@@ -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/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..f25ad1cd 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[Feature] = 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 1edcb9c8..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
@@ -125,7 +126,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 +135,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/model.py b/src/skala/functional/model.py
index 531709bc..d3c45092 100644
--- a/src/skala/functional/model.py
+++ b/src/skala/functional/model.py
@@ -9,13 +9,14 @@
"""
import math
-from typing import Any, cast
+from typing import Any, ClassVar, cast
import torch
from e3nn import o3
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,
)
@@ -57,7 +56,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:
@@ -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:
@@ -912,7 +910,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/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/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/__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/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/gradients.py b/src/skala/gpu4pyscf/gradients.py
index 54aa23ef..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,13 +88,13 @@ 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
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:
@@ -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/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/ao_evaluation.py b/src/skala/pyscf/ao_evaluation.py
new file mode 100644
index 00000000..5bb171cd
--- /dev/null
+++ b/src/skala/pyscf/ao_evaluation.py
@@ -0,0 +1,551 @@
+# SPDX-License-Identifier: MIT
+
+"""Blockwise atomic-orbital feature evaluation and custom autograd."""
+
+from collections.abc import Iterator
+from typing import NamedTuple, Protocol, TypeAlias, 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 torch.utils.dlpack import from_dlpack
+
+from skala.features import FeatureMap
+from skala.pyscf import feature_math
+from skala.pyscf.backend import (
+ Array,
+ Grid,
+ check_gpu_imports_were_successful,
+ dft_gpu,
+ 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 _ChunkEvalContext(Protocol):
+ mol: gto.Mole
+ grids: Grid
+ feature_function: feature_math.LinearFeature
+ blksize: int | None
+ compile_feature_function: bool
+ spin_shape: torch.Size
+ output_device: torch.device
+
+
+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
+ ``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)))
+
+
+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.
+ ``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_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_ao_indices[:, None], self.active_ao_indices[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_ao_indices[:, None], self.active_ao_indices[None, :]
+ ] += block_result
+
+
+def _evaluate_feature_block(
+ feature_function: feature_math.LinearFeature,
+ block: _AOBlock,
+ 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."""
+ 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)
+
+ if active_dm_submatrix is None:
+ raise ValueError("Feature evaluation requires a density matrix.")
+ if compile_feature_function:
+ return torch.compile(feature_function.forward)(
+ active_dm_submatrix, block.ao_values
+ )
+ return feature_function(active_dm_submatrix, block.ao_values)
+
+
+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.
+ """
+
+ def __init__(
+ self,
+ mol: gto.Mole,
+ grids: Grid,
+ feature_function: feature_math.LinearFeature,
+ blksize: int | None,
+ ) -> None:
+ self.mol = mol
+ assert isinstance(grids, dft.Grids)
+ self.grids = grids
+ self.feature_function = feature_function
+ self.blksize = blksize
+ self.numint = dft.numint.NumInt()
+
+ def order_aos(self, matrix: Tensor) -> Tensor:
+ return matrix
+
+ def restore_ao_order(self, matrix: Tensor) -> Tensor:
+ return matrix
+
+ def _active_ao_indices(
+ self,
+ non0tab: _ScreenIndex,
+ 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.from_numpy(_active_cpu_ao_indices(self.mol, block_non0tab))
+
+ def __iter__(self) -> Iterator[_AOBlock]:
+ non0tab = self.grids.non0tab
+
+ end = 0
+ 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=non0tab,
+ ):
+ start, end = end, end + block_weights.size
+ ao_values = torch.from_numpy(backend_ao_values).transpose(-1, -2)
+ 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
+ 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,
+ device: torch.device,
+ mol: gto.Mole,
+ grids: Grid,
+ feature_function: feature_math.LinearFeature,
+ blksize: int | None,
+ ) -> None:
+ check_gpu_imports_were_successful()
+ 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=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(
+ from_dlpack(backend_ao_values),
+ from_dlpack(active_ao_indices),
+ slice(start, end),
+ )
+
+
+def _make_ao_block_loop(
+ device: torch.device,
+ mol: gto.Mole,
+ grids: Grid,
+ feature_function: feature_math.LinearFeature,
+ blksize: int | None,
+) -> _CPUAOBlockLoop | _GPUAOBlockLoop:
+ if device.type == "cuda":
+ return _GPUAOBlockLoop(device, mol, grids, feature_function, blksize)
+ return _CPUAOBlockLoop(mol, grids, feature_function, blksize)
+
+
+class ChunkEvalForward(Function):
+ @staticmethod
+ def setup_context(
+ ctx: FunctionCtx,
+ inputs: tuple[
+ Tensor,
+ gto.Mole,
+ Grid,
+ feature_math.LinearFeature,
+ int | None,
+ bool,
+ ],
+ output: torch.Tensor,
+ ) -> None:
+ context = cast(_ChunkEvalContext, ctx)
+ (
+ dm,
+ context.mol,
+ context.grids,
+ context.feature_function,
+ context.blksize,
+ context.compile_feature_function,
+ ) = inputs
+ context.spin_shape = dm.shape[:-2]
+ context.output_device = output.device
+
+ @staticmethod
+ def forward(
+ dm: torch.Tensor,
+ mol: gto.Mole,
+ grids: Grid,
+ feature_function: feature_math.LinearFeature,
+ blksize: int | None,
+ compile_feature_function: bool,
+ ) -> torch.Tensor:
+ ngrids = grids.weights.size
+ block_loop = _make_ao_block_loop(
+ dm.device, mol, grids, feature_function, blksize
+ )
+
+ features = torch.zeros(
+ *dm.shape[:-2],
+ feature_function.nfeats,
+ ngrids,
+ device=dm.device,
+ dtype=dm.dtype,
+ )
+ evaluation_dm_ordered = block_loop.order_aos(dm)
+ for block in block_loop:
+ active_dm_submatrix = block.select_active_ao_submatrix(
+ evaluation_dm_ordered
+ )
+ temp_feature = _evaluate_feature_block(
+ feature_function,
+ block,
+ active_dm_submatrix,
+ compile_feature_function,
+ )
+ features[..., block.grid_slice] = temp_feature
+ return features
+
+ @staticmethod
+ 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.spin_shape,
+ ctx.feature_function.nfeats,
+ ctx.grids.weights.size,
+ device=ctx.output_device,
+ dtype=torch.float64,
+ )
+ return cast(
+ Tensor,
+ ChunkEvalForward.apply(
+ dm_tangent,
+ ctx.mol,
+ ctx.grids,
+ ctx.feature_function,
+ ctx.blksize,
+ ctx.compile_feature_function,
+ ),
+ )
+
+ @staticmethod
+ def backward(
+ ctx: _ChunkEvalContext, *grad_outputs: torch.Tensor
+ ) -> tuple[torch.Tensor | None, ...]:
+ feature_cotangent = grad_outputs[0]
+ 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):
+ @staticmethod
+ def setup_context(
+ ctx: FunctionCtx,
+ inputs: tuple[
+ torch.Tensor,
+ gto.Mole,
+ Grid,
+ feature_math.LinearFeature,
+ int | None,
+ bool,
+ ],
+ output: torch.Tensor,
+ ) -> None:
+ context = cast(_ChunkEvalContext, ctx)
+ (
+ feature_cotangent,
+ context.mol,
+ context.grids,
+ context.feature_function,
+ context.blksize,
+ context.compile_feature_function,
+ ) = inputs
+ context.spin_shape = feature_cotangent.shape[:-2]
+ context.output_device = output.device
+
+ @staticmethod
+ def forward(
+ feature_cotangent: torch.Tensor,
+ mol: gto.Mole,
+ grids: Grid,
+ feature_function: feature_math.LinearFeature,
+ blksize: int | None,
+ compile_feature_function: bool,
+ ) -> torch.Tensor:
+ block_loop = _make_ao_block_loop(
+ feature_cotangent.device, mol, grids, feature_function, blksize
+ )
+
+ 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,
+ block,
+ None,
+ compile_feature_function,
+ feature_cotangent,
+ )
+ block.add_active_ao_submatrix(out, block_result)
+ return block_loop.restore_ao_order(out)
+
+ @staticmethod
+ def jvp(ctx: _ChunkEvalContext, *grad_inputs: torch.Tensor | None) -> torch.Tensor:
+ feature_cotangent_tangent = grad_inputs[0]
+ if feature_cotangent_tangent is None:
+ 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,
+ ),
+ )
+
+ @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,
+ )
+ # 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(
+ dm: torch.Tensor,
+ mol: gto.Mole,
+ coords: Array,
+ feature_function: feature_math.LinearFeature,
+ 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.LinearFeature,
+ 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 None:
+ return None
+ raise ValueError("Setting custom block size is not supported on GPU.")
+
+ 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.LinearFeature,
+ block_size: int | None = None,
+ max_memory: int = 2000,
+ gpu: bool = False,
+) -> FeatureMap:
+ """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 = evaluate_full_grid(dm.double(), mol, grids.coords, feature_function)
+ else:
+ features = ChunkEvalForward.apply(
+ dm.double(), mol, grids, feature_function, blksize, False
+ )
+ return feature_function.to_dict(features)
diff --git a/src/skala/pyscf/dft.py b/src/skala/pyscf/dft.py
index 62786922..aedc1bb8 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.features import _ATOMIC_GRID_FEATURES
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 bool(set(func.features) & _ATOMIC_GRID_FEATURES)
-
-
-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,23 @@ 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)
+ self.small_rho_cutoff = 0 # pyscf 2.9 default is 1e-7
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 +161,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 +178,23 @@ 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)
+ self.small_rho_cutoff = 0 # pyscf 2.9 default is 1e-7
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/evaluation.py b/src/skala/pyscf/evaluation.py
new file mode 100644
index 00000000..3c1e7ca6
--- /dev/null
+++ b/src/skala/pyscf/evaluation.py
@@ -0,0 +1,105 @@
+# SPDX-License-Identifier: MIT
+
+"""Feature requirements and numerical-evaluation policy."""
+
+from collections.abc import Iterable
+from dataclasses import dataclass
+
+from skala.features import Feature
+
+_AO_FEATURES = frozenset(
+ {
+ Feature.DENSITY,
+ Feature.GRAD,
+ Feature.KIN,
+ Feature.LAPL,
+ }
+)
+_ATOMIC_LAYOUT_FEATURES = frozenset(
+ {
+ Feature.ATOMIC_GRID_WEIGHTS,
+ Feature.ATOMIC_GRID_SIZES,
+ Feature.ATOMIC_GRID_SIZE_BOUND_SHAPE,
+ }
+)
+
+
+class FeatureSpec:
+ """Normalized feature names and their evaluation requirements."""
+
+ 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 __hash__(self) -> int:
+ return hash(self.names)
+
+ 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(Feature.DENSITY)
+
+ @property
+ def with_grad(self) -> bool:
+ """Return whether the density gradient is requested."""
+ return self.requests(Feature.GRAD)
+
+ @property
+ def with_kin(self) -> bool:
+ """Return whether kinetic-energy density is requested."""
+ return self.requests(Feature.KIN)
+
+ @property
+ def with_lapl(self) -> bool:
+ """Return whether the density Laplacian is requested."""
+ return self.requests(Feature.LAPL)
+
+ @property
+ 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:
+ """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 Feature.LAPL in self.names:
+ return 2
+ if self.names & {Feature.GRAD, Feature.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_spatial_decomposition(self) -> bool:
+ """Return whether spatial decomposition is supported."""
+ return Feature.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/feature_math.py b/src/skala/pyscf/feature_math.py
new file mode 100644
index 00000000..b6a2f66b
--- /dev/null
+++ b/src/skala/pyscf/feature_math.py
@@ -0,0 +1,173 @@
+# SPDX-License-Identifier: MIT
+
+"""Raw density-feature mathematics and model formatting."""
+
+from abc import ABC, abstractmethod
+
+import torch
+from torch import nn
+
+from skala.features import Feature, FeatureMap
+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 LinearFeature(nn.Module, ABC):
+ """Linear raw-feature map from a density matrix and fixed AO values."""
+
+ deriv: int
+ nfeats: int
+
+ @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(LinearFeature):
+ """Evaluate the requested linear meta-GGA density features."""
+
+ def __init__(self, feature_spec: FeatureSpec):
+ super().__init__()
+
+ 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) -> FeatureMap:
+ """Convert a packed feature tensor to its named feature tensors."""
+ feature_index = 0
+ feature_dict: FeatureMap = {}
+ if self.feature_spec.with_density:
+ feature_dict[Feature.DENSITY] = features[..., feature_index, :]
+ feature_index += 1
+ if self.feature_spec.with_grad:
+ feature_dict[Feature.GRAD] = features[
+ ..., feature_index : feature_index + 3, :
+ ]
+ feature_index += 3
+ if self.feature_spec.with_kin:
+ feature_dict[Feature.KIN] = features[..., feature_index, :]
+ feature_index += 1
+ if self.feature_spec.with_lapl:
+ feature_dict[Feature.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))
+
+ 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:
+ 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))
+ 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/src/skala/pyscf/features.py b/src/skala/pyscf/features.py
index 2c5bfc11..751c7dfd 100644
--- a/src/skala/pyscf/features.py
+++ b/src/skala/pyscf/features.py
@@ -4,219 +4,35 @@
Methods for generating and manipulating density features.
"""
-import logging
-from abc import ABC, abstractmethod
-from collections.abc import Callable, Iterator
-from copy import copy
-
import numpy as np
import torch
-from pyscf import dft, gto
-from torch import Tensor, nn
-from torch.autograd import Function
-from torch.autograd.function import FunctionCtx
-
-from skala.pyscf.backend import (
- Array,
- Grid,
- check_gpu_imports_were_successful,
- dft_gpu,
- from_numpy_or_cupy,
-)
-from skala.pyscf.memory_estimators import estimate_max_grid_chunk_size
-
-LOG = logging.getLogger(__name__)
-
-DEFAULT_FEATURES = ["density", "kin", "grad", "grid_coords", "grid_weights"]
+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 = [
+ Feature.DENSITY,
+ Feature.KIN,
+ Feature.GRAD,
+ Feature.GRID_COORDS,
+ Feature.GRID_WEIGHTS,
+]
DEFAULT_FEATURES_SET = set(DEFAULT_FEATURES)
-# Features that require per-atom grid decomposition.
-_ATOMIC_GRID_FEATURES = {
- "atomic_grid_weights",
- "atomic_grid_sizes",
- "atomic_grid_size_bound_shape",
-}
-
-
-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 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,
-) -> 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.
-
- 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 dm is a 3D tensor, then we have a spin-polarized system
- with_spin = True if len(dm.shape) == 3 else False
-
- 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
- feat_tensor = non_chunk(
- dm.double(),
- mol,
- grids.coords[grid_slice],
- ff,
- compile_feature_function=compile_feature_function,
- gpu=dm.device.type == "cuda",
- )
-
- for k, v in ff.to_dict(feat_tensor).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]]:
- """
- 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(
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.
@@ -243,42 +59,31 @@ 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 = True if len(dm.shape) == 3 else False
+ 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, 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:
- mgga_features = auto_chunk(
+ if feature_spec.requires_ao_evaluation:
+ mgga_features = ao_evaluation.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,
+ feature_math.MGGAFeatureFunction(feature_spec),
+ block_size=evaluation_policy.ao_block_size,
max_memory=max_memory,
- fix_block_size=chunk_size is None,
gpu=gpu,
)
for feature in mgga_features:
- mol_features[feature] = maybe_expand_and_divide(
- mgga_features[feature], not with_spin, 2
+ mol_features[feature] = feature_math.maybe_expand_and_divide(
+ mgga_features[feature], not is_spin_polarized, 2
)
return mol_features
@@ -288,26 +93,26 @@ def get_grid_features(
mol: gto.Mole,
dm: Tensor,
grids: Grid,
- requested_features: set[str],
-) -> dict[str, Tensor]:
- grid_features = {}
+ feature_spec: FeatureSpec,
+) -> FeatureMap:
+ grid_features: FeatureMap = {}
- if "grid_coords" in requested_features:
- 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 "grid_weights" in requested_features:
- 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 "coarse_0_atomic_coords" in requested_features:
- 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
)
- 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
)
@@ -323,776 +128,23 @@ def get_grid_features(
f"Set grids.alignment = 1 before building grids to disable padding."
)
- if "atomic_grid_sizes" in requested_features:
- 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 "atomic_grid_size_bound_shape" in requested_features:
+ 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 "atomic_grid_weights" in requested_features:
+ 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
)
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,
-) -> 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 for chaining of derivatives.
- """
-
- def partial_feature_function(dm: torch.Tensor) -> torch.Tensor:
- return feature_function(dm, ao)
-
- 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,
-) -> 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 allow for chaining of derivatives over primals."""
-
- def reduced_vjp(primals: torch.Tensor) -> torch.Tensor:
- return torch.func.vjp(func, primals)[1](tangents)[0]
-
- return reduced_vjp
-
-
-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: ...
-
- @abstractmethod
- 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
- with_ked_var: bool
- with_ked_det: bool
-
- def __init__(
- self,
- with_density: bool = True,
- with_grad: bool = True,
- with_kin: bool = True,
- with_lapl: bool = False,
- with_ked_var: bool = False,
- with_ked_det: bool = False,
- ):
- super().__init__()
-
- self.with_density = with_density
- 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:
- 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
- )
-
- 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
- feature_dict: dict[str, torch.Tensor] = {}
- if self.with_density:
- feature_dict["density"] = features[..., feature_index, :]
- feature_index += 1
- if self.with_grad:
- feature_dict["grad"] = features[..., feature_index : feature_index + 3, :]
- feature_index += 3
- if self.with_kin:
- feature_dict["kin"] = features[..., feature_index, :]
- feature_index += 1
- 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])
- # 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.with_density:
- features[..., feat_idx, :] = torch.sum(c0 * ao[0][None, :, :], dim=-2)
- feat_idx += 1
-
- if self.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) and not with_Q:
- 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:
- feat_idx += 1
- if self.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:
- # 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 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:
- return features.reshape((*dm.shape[:-2], self.nfeats, -1))
-
-
-class ChunkEvalForward(Function):
- @staticmethod
- def setup_context(
- ctx: FunctionCtx,
- inputs: tuple[
- torch.Tensor,
- gto.Mole,
- Grid,
- FeatureFunction,
- int,
- 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: FeatureFunction,
- blksize: int,
- compile_feature_function: bool,
- gpu: bool,
- *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 if not gpu else 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())
-
- features = torch.zeros(
- *dm.shape[:-2],
- feature_function.nfeats,
- ngrids,
- device=dm.device,
- dtype=dm.dtype,
- )
- 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
- # 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, :]]
-
- # 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
- ),
- )
- for v_sorted in vectors_jvp_sorted:
- partial_func = partial_jvp_function_over_tangents(
- partial_func,
- v_sorted[..., mask[:, None], mask[None, :]],
- )
-
- # Compute feature (or its jvp) for this block with masked dm
- if compile_feature_function:
- temp_feature = torch.compile(partial_func)(masked_dm)
- else:
- temp_feature = partial_func(masked_dm)
-
- features[..., start:end] = temp_feature
- return features
-
- @staticmethod
- def jvp(ctx: FunctionCtx, grad_input: torch.Tensor) -> torch.Tensor:
- # Chain rule for the jvp
- return ChunkEvalForward.apply(
- ctx.dm,
- ctx.mol,
- ctx.grids,
- ctx.feature_function,
- ctx.blksize,
- ctx.compile_feature_function,
- ctx.gpu,
- *ctx.vectors_jvp,
- grad_input,
- )
-
- @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(
- 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,
- )
- ]
-
- # 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
-
- # 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(
- 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 :],
- )
- )
-
- return tuple(grads)
-
-
-class ChunkEvalBackward(Function):
- @staticmethod
- def setup_context(
- ctx: FunctionCtx,
- inputs: tuple[
- torch.Tensor,
- gto.Mole,
- Grid,
- FeatureFunction,
- list[str],
- int,
- bool,
- bool,
- torch.Tensor,
- ],
- output: tuple[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,
- ) = inputs
- ctx.save_for_backward(ctx.dm)
-
- @staticmethod
- def forward(
- dm: torch.Tensor,
- mol: gto.Mole,
- grids: Grid,
- feature_function: FeatureFunction,
- derivative_types: list[str],
- blksize: int,
- compile_feature_function: bool,
- 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,
- }
- 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())
-
- 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
-
- # Mask to only include the relevant AOs
- if mask is None or not gpu:
- mask = torch.arange(mol.nao_nr(), device=dm.device)
- else:
- mask = from_numpy_or_cupy(mask, device=dm.device, dtype=torch.long)
-
- # 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
- ),
- )
- for derivative_type, vector, v_sorted in zip(
- derivative_types, vectors, vectors_sorted, strict=True
- ):
- if derivative_type == "jvp":
- partial_func = partial_jvp_function_over_tangents(
- partial_func,
- 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, :]],
- )
- elif derivative_type == "first_vjp":
- partial_func = partial_vjp_function_over_tangents(
- partial_func, vector[..., start:end]
- )
- else:
- raise ValueError(
- f"Unknown derivative {derivative_type} (must be one of 'jvp', 'vjp', 'first_vjp')"
- )
- if compile_feature_function:
- out[..., mask[:, None], mask[None, :]] += torch.compile(partial_func)(
- dm_sorted[..., mask[:, None], mask[None, :]]
- )
- else:
- out[..., mask[:, None], mask[None, :]] += partial_func(
- dm_sorted[..., mask[:, None], mask[None, :]]
- )
- return out[..., unsort_idx, :][..., unsort_idx]
-
- @staticmethod
- def jvp(ctx: FunctionCtx, *grad_input: torch.Tensor) -> torch.Tensor:
- # Chain rule for the jvp
- 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,
- )
-
- @staticmethod
- def backward(
- ctx: FunctionCtx, *grad_outputs: torch.Tensor
- ) -> tuple[torch.Tensor | None, ...]:
- # Chain rule for the vjp
-
- # Gradient corresponding to dm
- grads = [
- ChunkEvalBackward.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,
- )
- ]
- # 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)
-
-
-def non_chunk(
- dm: torch.Tensor,
- mol: gto.Mole,
- coords: Array,
- feature_function: 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: 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/src/skala/pyscf/gradients.py b/src/skala/pyscf/gradients.py
index 02c0a97b..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,13 +83,13 @@ 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
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:
@@ -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/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/memory_estimators.py b/src/skala/pyscf/memory_estimators.py
index 6df15ea0..4794d5ff 100644
--- a/src/skala/pyscf/memory_estimators.py
+++ b/src/skala/pyscf/memory_estimators.py
@@ -1,38 +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_grid_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,
-) -> int:
- """Heuristically pick a grid chunk size for :func:`chunked_features`.
+) -> 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
- ``non_chunk`` (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.
@@ -40,18 +49,22 @@ 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.
-
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 ``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 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:
case "cuda":
@@ -69,78 +82,80 @@ def estimate_max_grid_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)
+ available_memory = int(free_bytes * safety_fraction)
+ available_memory -= estimate_global_screened_buffer_memory(
+ dm, nfeatures, atomic_grid_sizes, func_deriv
+ )
- 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)
+ 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())
+ }
- 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 linear_peak_memory_model(
- nao: int,
- deriv: int,
+
+def estimate_global_raw_feature_buffer_memory(
+ dm: torch.Tensor,
+ nfeatures: int,
+ ngrids: int,
func_deriv: int,
-) -> tuple[float, float]:
- """
- Return the coefficients of the linear model for peak memory usage in the number of grid points::
+) -> int:
+ """Estimate full-grid raw-feature storage for global screened evaluation.
- bytes ~= bytes_per_point * n + fixed_overhead
+ 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.
- Both terms are quadratic in ``nao`` and calibrated *per code path*::
+ 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.
- bytes_per_point = 8 * (C_AO2*nao^2 + (ncomp + C_LIN)*nao + C_NET)
- fixed_overhead = C_FIX * nao^2
+ Returns:
+ Estimated bytes occupied by global raw-feature buffers.
- 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.
+ Raises:
+ ValueError: If ``func_deriv`` is not first or second order.
"""
- # 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
+ buffer_count = 4
case 2:
- C_AO2, C_LIN, C_NET, C_FIX = 1.80e-3, 1.7, 24230.0, 9.0
+ buffer_count = 5
case _:
- raise ValueError("Invalid func_deriv value")
+ 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
- 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)
- )
- 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
+def estimate_global_screened_buffer_memory(
+ dm: torch.Tensor,
+ nfeatures: int,
+ atomic_grid_sizes: torch.Tensor,
+ func_deriv: int,
+) -> 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
+ )
+ 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
new file mode 100644
index 00000000..581bc720
--- /dev/null
+++ b/src/skala/pyscf/model_chunking.py
@@ -0,0 +1,232 @@
+# 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, 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_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."""
+
+ atom_slice: slice
+ grid_slice: slice
+
+
+def _make_atom_grid_chunks(
+ atomic_grid_sizes: Tensor, max_atoms_per_grid_size: Mapping[int, int]
+) -> list[AtomGridChunk]:
+ """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
+ 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(chunk_atom_start, chunk_atom_stop),
+ grid_slice=slice(grid_start, grid_start + chunk_grid_size),
+ )
+ )
+ grid_start += chunk_grid_size
+ atom_start = group_stop
+
+ LOG.debug(
+ "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_indices: Tensor
+ raw_features: Tensor
+ model_features: FeatureMap
+
+
+@dataclass(frozen=True)
+class ModelFeatureChunker:
+ """Reusable atom-aligned partition of raw and model features."""
+
+ atom_major_raw_features: Tensor
+ 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.index_select(-1, grid_indices)
+ .detach()
+ .requires_grad_()
+ )
+ model_features: FeatureMap = {}
+ for feature_name in _GRID_POINT_FEATURES:
+ if feature_spec.requests(feature_name):
+ 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
+ ].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())
+ model_features[Feature.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.is_spin_polarized, 2
+ )
+ yield ModelFeatureChunk(
+ grid_indices=grid_indices,
+ 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,
+ 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_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)
+ 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,
+ 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,
+ )
+ 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(
+ 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/numint.py b/src/skala/pyscf/numint.py
index 558f4d22..dbee4cb9 100644
--- a/src/skala/pyscf/numint.py
+++ b/src/skala/pyscf/numint.py
@@ -1,10 +1,10 @@
# 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 import gto
from torch import Tensor
from skala.functional.base import ExcFunctionalBase
@@ -12,12 +12,11 @@
KS,
Array,
Grid,
- check_gpu_imports_were_successful,
from_numpy_or_cupy,
to_cupy,
to_numpy,
)
-from skala.pyscf.features import chunked_features, generate_features
+from skala.pyscf.xc_integrator import XCIntegrator
class LibXCSpec(Protocol):
@@ -95,48 +94,48 @@ 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...
"""
- device: torch.device
-
def __init__(
self,
functional: ExcFunctionalBase,
chunk_size: int | None = None,
device: torch.device | None = None,
):
- if device is None:
- self.device = torch.get_default_device()
- else:
- self.device = device
-
- if self.device.type == "cuda":
- check_gpu_imports_were_successful()
-
- self.func = functional.to(device=self.device)
- self.chunk_size = chunk_size
-
- 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 to_backend(self, x: Tensor | list[Tensor]) -> Array | list[Array]:
+ self.integrator = XCIntegrator(functional, chunk_size=chunk_size, device=device)
+
+ @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
+
+ def _from_backend(self, x: Array) -> Tensor:
+ return from_numpy_or_cupy(x, device=self.device)
+
+ @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)
@@ -151,21 +150,18 @@ def get_rho(
max_memory: int = 2000,
verbose: int = 0,
) -> Array:
- mol_features = generate_features(
+ density = self.integrator.density(
mol,
- self.from_backend(dm),
+ self._from_backend(dm),
grids,
- features={"density"},
- chunk_size=self.chunk_size,
max_memory=max_memory,
- gpu=self.device.type == "cuda",
)
- return self.to_backend(mol_features["density"].sum(0)) # type: ignore
+ return self._to_backend(density)
def __call__(
self,
mol: gto.Mole,
- grids: dft.Grids,
+ grids: Grid,
xc_code: str | None,
dm: Tensor,
second_order: bool = False,
@@ -178,74 +174,19 @@ def __call__(
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.
+ 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 self.device != dm.device:
- raise ValueError(
- f"Density matrix device {dm.device} does not match functional device {self.device}"
+ if second_order:
+ raise NotImplementedError(
+ "Direct second-order evaluation is not supported; use gen_response()."
)
- if self._functional_supports_atom_chunking():
- 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)
- V_xc = torch.zeros_like(dm)
- for mol_features in chunked_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
- ):
- E_xc_chunk = self.func.get_exc(mol_features)
- (V_xc_chunk,) = torch.autograd.grad(
- E_xc_chunk,
- dm,
- torch.ones_like(E_xc_chunk),
- )
- 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
-
- 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",
- )
- 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
+ result = self.integrator(mol, grids, dm, max_memory=max_memory)
+ return result.electron_count, result.energy, result.potential
def nr_rks(
self,
@@ -258,9 +199,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) # type: ignore
+ return N.sum().item(), E_xc.item(), self._to_backend(V_xc)
def nr_uks(
self,
@@ -273,9 +214,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) # type: ignore
+ return self._to_backend(N), E_xc.item(), self._to_backend(V_xc)
class libxc:
__version__ = None
@@ -298,6 +239,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:
@@ -310,75 +252,24 @@ def gen_response(
if "with_j" in kwargs:
assert kwargs["with_j"]
- dm0 = self.from_backend(ks.make_rdm1(mo_coeff, mo_occ))
-
- if self._functional_supports_atom_chunking():
- dm0 = dm0.requires_grad_()
-
- 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(
- 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.
- ):
- 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)
- if ks.mol.spin == 0:
- v1 += vj
- else:
- v1 += vj[0] + vj[1]
- return v1
-
- 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]
-
- 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)
+ dm0 = self._from_backend(ks.make_rdm1(mo_coeff, mo_occ))
+ xc_response = self.integrator.gen_response(
+ ks.mol,
+ ks.grids,
+ dm0,
+ max_memory=ks.max_memory,
+ safety_fraction=kwargs.get("safety_fraction"),
+ )
- if ks.mol.spin == 0:
- v1 += vj
- else:
- v1 += vj[0] + vj[1]
+ def hessian_vector_product(dm1: Array) -> Array:
+ v1 = self._to_backend(xc_response(self._from_backend(dm1)))
+ vj = ks.get_j(ks.mol, dm1, hermi=1)
- return v1
+ if ks.mol.spin == 0:
+ v1 += vj
+ else:
+ v1 += vj[0] + vj[1]
- return hessian_vector_product
+ return v1
- def _functional_supports_atom_chunking(self) -> bool:
- return "atomic_grid_sizes" in self.func.features
+ return hessian_vector_product
diff --git a/src/skala/pyscf/screening.py b/src/skala/pyscf/screening.py
new file mode 100644
index 00000000..0946faf1
--- /dev/null
+++ b/src/skala/pyscf/screening.py
@@ -0,0 +1,197 @@
+# 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 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 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 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]]
+
+
+@dataclass(frozen=True)
+class SpatialGridLayout:
+ """Evaluation-ready spatial ordering derived from an atom-major grid."""
+
+ block_size: int
+ sorted_grids: Grid
+ forward_permutation: Tensor
+ inverse_permutation: Tensor
+
+
+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 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)``.
+ 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 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
+ positions = np.lexsort((indices, split_projections(indices)))
+ 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_spatial_grid_layout(
+ mol: gto.Mole,
+ grids: Grid,
+ block_size: int,
+ 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.
+ device: Torch device used for permutation tensors.
+
+ Returns:
+ 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.")
+
+ gpu = device.type == "cuda"
+ 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)
+ 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,
+ )
+ 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),
+ )
+
+
+def screened_feature_jvp(
+ dm_tangent: Tensor,
+ mol: gto.Mole,
+ spatial_grid_layout: SpatialGridLayout,
+ feature_function: feature_math.MGGAFeatureFunction,
+ compile_feature_function: bool = False,
+) -> Tensor:
+ """Apply the raw-feature Jacobian and restore atom-major grid order."""
+ sorted_tangent = cast(
+ Tensor,
+ ao_evaluation.ChunkEvalForward.apply(
+ dm_tangent,
+ mol,
+ spatial_grid_layout.sorted_grids,
+ feature_function,
+ spatial_grid_layout.block_size,
+ compile_feature_function,
+ ),
+ )
+ return sorted_tangent.index_select(
+ -1, spatial_grid_layout.inverse_permutation
+ ).detach()
diff --git a/src/skala/pyscf/xc_integrator.py b/src/skala/pyscf/xc_integrator.py
new file mode 100644
index 00000000..a7e223db
--- /dev/null
+++ b/src/skala/pyscf/xc_integrator.py
@@ -0,0 +1,383 @@
+# SPDX-License-Identifier: MIT
+
+"""Tensor-level exchange-correlation integration."""
+
+from collections.abc import Callable
+from typing import NamedTuple, Protocol, 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
+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,
+ SpatialGridLayout,
+ prepare_spatial_grid_layout,
+ screened_feature_jvp,
+)
+
+
+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
+ return 2 * mol.nao_nr() > pyscf_numint.SWITCH_SIZE
+
+
+class XCResult(NamedTuple):
+ """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={Feature.DENSITY},
+ chunk_size=self.evaluation_policy.ao_block_size,
+ max_memory=max_memory,
+ gpu=self.device.type == "cuda",
+ )
+ return mol_features[Feature.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)
+ 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)
+
+ 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)
+ self._require_skala_grids(grids)
+ if self.feature_spec.supports_spatial_decomposition 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 _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_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
+
+ 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_cache.cache_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_()
+ 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_eval,
+ mol,
+ spatial_grid_layout.sorted_grids,
+ feature_function,
+ spatial_grid_layout.block_size,
+ False,
+ ),
+ )
+ 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,
+ deriv_order=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.index_copy_(
+ -1, chunk.grid_indices, local_cotangent.detach()
+ )
+ electron_count += (
+ (mol_features[Feature.DENSITY] * mol_features[Feature.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) | {Feature.DENSITY, Feature.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[Feature.DENSITY] * mol_features[Feature.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,
+ ),
+ )
+ 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,
+ deriv_order=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(
+ 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.index_select(-1, chunk.grid_indices),
+ )
+ else:
+ local_hessian_action = torch.zeros_like(local_raw_features)
+ atom_major_hessian_action.index_copy_(
+ -1, chunk.grid_indices, 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
new file mode 100644
index 00000000..1c0fba21
--- /dev/null
+++ b/tests/test_ao_screening.py
@@ -0,0 +1,1065 @@
+from collections.abc import Callable, Iterator
+from itertools import combinations
+from typing import Any
+
+import numpy as np
+import pytest
+import torch
+from pyscf import dft, gto
+from utils import QuadraticFunctional, 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
+from skala.pyscf import xc_integrator as xc_integrator_module
+from skala.pyscf.ao_evaluation import (
+ ChunkEvalBackward,
+ ChunkEvalForward,
+ _active_cpu_ao_indices,
+ _AOBlock,
+ _CPUAOBlockLoop,
+ _evaluate_feature_block,
+ _resolve_ao_block_size,
+)
+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 (
+ SpatialGridLayout,
+ _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 = [
+ combination
+ for size in range(1, len(_MGGA_FEATURES) + 1)
+ for combination in combinations(_MGGA_FEATURES, size)
+]
+
+
+@pytest.fixture
+def carbon() -> gto.Mole:
+ return gto.M(atom="C 0 0 0", basis="sto-3g", spin=2, verbose=0)
+
+
+@pytest.mark.parametrize(
+ ("feature_names", "expected_deriv", "expected_nfeats"),
+ [
+ ({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[Feature],
+ expected_deriv: int,
+ expected_nfeats: int,
+) -> 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_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:
+ 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 feature_function.feature_spec is feature_spec
+ assert features.shape == (expected_nfeats, 3)
+ 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))
+
+
+@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],
+) -> None:
+ with pytest.raises(
+ ValueError, match="At least one AO-derived feature must be selected"
+ ):
+ MGGAFeatureFunction(FeatureSpec(feature_names))
+
+
+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 xc_integrator_module._should_screen_aos is original_decision
+
+
+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
+ 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_ao_indices(carbon, screen_index), expected)
+
+ empty = _active_cpu_ao_indices(carbon, np.zeros_like(screen_index))
+ assert empty.dtype == np.int64
+ assert empty.size == 0
+
+
+def test_resolve_ao_block_size_modes(carbon: gto.Mole) -> None:
+ 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.
+ 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_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 = _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)
+ 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_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)
+ coords = np.column_stack(
+ (100.0 * labels + offsets, np.zeros(labels.size), np.zeros(labels.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])
+
+
+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:
+ """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)
+ 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
+ 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
+ return forward, inverse
+
+ monkeypatch.setattr(
+ screening_module,
+ "_decompose_grid_into_spatial_blocks",
+ 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
+ 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)
+
+ 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)
+ 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
+ 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 partition_calls == 1
+ assert screen_index_calls == 1
+ assert decomposition_block_sizes == [2]
+ assert screened_molecules == [carbon]
+ 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 = SkalaGrids(carbon)
+ grids.coords = np.arange(18, dtype=np.float64).reshape(6, 3)
+ grids.weights = np.arange(6, dtype=np.float64)
+ other_grids = SkalaGrids(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(
+ xc_integrator_module,
+ "prepare_spatial_grid_layout",
+ fake_prepare_spatial_grid_layout,
+ )
+ 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
+ assert grids.get_cached_spatial_grid_layout() is layout
+ assert len(layouts) == 1
+
+ numint.reset()
+ assert numint.integrator._get_spatial_grid_layout(carbon, grids) is layout
+
+ other_layout = numint.integrator._get_spatial_grid_layout(carbon, other_grids)
+ assert other_layout is not layout
+ assert other_grids.get_cached_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
+ 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)
+
+
+def test_call_rejects_second_order_evaluation(carbon: gto.Mole) -> None:
+ numint = SkalaNumInt(QuadraticFunctional())
+
+ 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(
+ carbon: gto.Mole,
+ monkeypatch: pytest.MonkeyPatch,
+ expected: bool,
+ response_safety_fraction: float | None,
+) -> None:
+ routes: list[str] = []
+ safety_fractions: list[float] = []
+
+ def fake_generate_features(
+ mol: gto.Mole,
+ dm: torch.Tensor,
+ grids: object,
+ features: set[Feature] | None = None,
+ **kwargs: object,
+ ) -> FeatureMap:
+ routes.append("dense")
+ density = dm.square().sum().reshape(1).expand(2, 1) / 2
+ return {
+ Feature.ATOMIC_GRID_SIZES: torch.tensor([1]),
+ Feature.DENSITY: density,
+ Feature.GRID_WEIGHTS: torch.ones(1, dtype=dm.dtype),
+ }
+
+ class FakeSpatialGridLayout:
+ block_size = 1
+ forward_permutation = torch.tensor([0])
+ inverse_permutation = torch.tensor([0])
+
+ def __init__(self, sorted_grids: object) -> None:
+ self.sorted_grids = sorted_grids
+
+ 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_indices=torch.tensor([0]),
+ raw_features=raw_features,
+ model_features={
+ 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),
+ },
+ )
+
+ def fake_prepare_spatial_grid_layout(
+ mol: gto.Mole,
+ grids: object,
+ 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 dm.sum().reshape(1, 1)
+
+ def fake_screened_feature_jvp(
+ 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,
+ dm: torch.Tensor,
+ grids: object,
+ atom_major_raw_features: torch.Tensor,
+ feature_function: MGGAFeatureFunction,
+ deriv_order: int,
+ **kwargs: object,
+ ) -> FakeModelFeatureChunks:
+ safety_fraction = kwargs["safety_fraction"]
+ assert isinstance(safety_fraction, float)
+ safety_fractions.append(safety_fraction)
+ return FakeModelFeatureChunks(atom_major_raw_features)
+
+ monkeypatch.setattr(
+ xc_integrator_module, "generate_features", fake_generate_features
+ )
+ monkeypatch.setattr(
+ xc_integrator_module,
+ "prepare_spatial_grid_layout",
+ fake_prepare_spatial_grid_layout,
+ )
+ monkeypatch.setattr(
+ ChunkEvalForward,
+ "apply",
+ staticmethod(fake_chunk_eval_forward),
+ )
+ monkeypatch.setattr(
+ xc_integrator_module,
+ "prepare_model_feature_chunks",
+ fake_prepare_model_feature_chunks,
+ )
+ monkeypatch.setattr(
+ xc_integrator_module,
+ "screened_feature_jvp",
+ fake_screened_feature_jvp,
+ )
+ numint = SkalaNumInt(QuadraticFunctional())
+ dm = torch.eye(carbon.nao_nr(), dtype=torch.float64)
+ grids = SkalaGrids(carbon)
+ grids.weights = np.ones(1)
+
+ ks = FakeKS(carbon, grids)
+ response_kwargs = (
+ {}
+ if response_safety_fraction is None
+ else {"safety_fraction": response_safety_fraction}
+ )
+ 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"
+ 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_feature_block_helper_localizes_derivative_vectors() -> None:
+ """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),
+ active_ao_indices=torch.tensor([0, 2]),
+ grid_slice=slice(1, 3),
+ )
+ 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,
+ )
+ feature_jvp = _evaluate_feature_block(
+ feature_function,
+ block,
+ block.select_active_ao_submatrix(tangent_ordered),
+ compile_feature_function=False,
+ )
+ 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,
+ None,
+ 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_values, block.ao_values
+ )
+ torch.testing.assert_close(feature_vjp, expected_vjp)
+
+
+def test_chunk_eval_transforms_follow_linear_operator(carbon: gto.Mole) -> None:
+ """Check spin-resolved first and second JVPs and the adjoint JVP."""
+ grids = _minimal_atom_grid(carbon)
+ feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY]))
+ 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
+ )
+
+ 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:
+ 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]
+ 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,),
+ (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:
+ 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)
+
+ ao = np.arange(ngrids * carbon.nao_nr(), dtype=np.float64).reshape(
+ ngrids, carbon.nao_nr()
+ )
+ 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_ao_indices = _active_cpu_ao_indices(carbon, screen_index)
+
+ class FakeNumInt:
+ def block_loop(
+ self, *args: object, **kwargs: object
+ ) -> Iterator[tuple[np.ndarray, None, np.ndarray, np.ndarray]]:
+ assert kwargs["non0tab"] is screen_index
+ assert "strict_grid_order" not in kwargs
+ 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)
+
+ feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY]))
+ 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, block_size, False
+ )
+
+ 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_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_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:
+ """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)
+ 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([Feature.DENSITY]))
+
+ blocks = list(_CPUAOBlockLoop(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(
+ 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))
+ 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, 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([Feature.DENSITY]))
+ 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
+ )
+ (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) -> 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)
+ spatial_grid_layout = prepare_spatial_grid_layout(
+ carbon,
+ grids,
+ block_size=dft.gen_grid.BLKSIZE,
+ device=torch.device("cpu"),
+ )
+ grids.cache_spatial_grid_layout(spatial_grid_layout)
+
+ assert numint.reset() is numint
+ assert grids.get_cached_spatial_grid_layout() is spatial_grid_layout
+
+
+@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],
+ 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:
+ 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)
+ numint = SkalaNumInt(functional)
+ grids = _minimal_atom_grid(mol)
+ dm = mean_field.get_init_guess()
+
+ with patch_ao_screening(False):
+ dense = integration_method(numint, mol, grids, None, dm)
+
+ with patch_ao_screening(True):
+ screened = integration_method(numint, mol, grids, None, dm)
+
+ assert np.allclose(
+ dense[result_index], screened[result_index], rtol=rtol, atol=atol
+ )
+
+
+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; F 0 0 0.92", basis="sto-3g", spin=0, verbose=0)
+ grids = _minimal_atom_grid(mol)
+ ks = FakeKS(mol, grids)
+ 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(
+ mol.nao_nr(), mol.nao_nr()
+ )
+ dm1 += dm1.T
+
+ with patch_ao_screening(False):
+ dense_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
+ )
+
+
+@pytest.mark.parametrize("func_deriv", [1, 2])
+def test_screened_ao_traversals_are_independent_of_model_chunking(
+ 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(
+ model_chunking_module,
+ "estimate_max_model_atoms_per_chunk",
+ lambda *args, **kwargs: {atom_grid_size: 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) # type: ignore[no-untyped-call]
+
+ def counting_backward_apply(*args: object) -> torch.Tensor:
+ nonlocal backward_calls
+ backward_calls += 1
+ return original_backward_apply(*args) # type: ignore[no-untyped-call]
+
+ monkeypatch.setattr(ChunkEvalForward, "apply", counting_forward_apply)
+ monkeypatch.setattr(ChunkEvalBackward, "apply", counting_backward_apply)
+ functional = QuadraticFunctional()
+ numint = SkalaNumInt(functional)
+
+ 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
new file mode 100644
index 00000000..5b51d707
--- /dev/null
+++ b/tests/test_ao_screening_benchmark.py
@@ -0,0 +1,471 @@
+"""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
+import tempfile
+import traceback
+from collections.abc import Callable, Iterator
+from multiprocessing.connection import Connection
+from pathlib import Path
+from typing import Any, NamedTuple, cast
+
+import numpy as np
+import pytest
+import torch
+from pyscf import dft, gto, lib
+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.grids import SkalaGrids
+from skala.pyscf.numint import SkalaNumInt
+
+THREAD_COUNT = 4
+MAX_MEMORY_MB = 2000
+MEMORY_WORKER_TIMEOUT_SECONDS = 240
+
+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
+
+
+DeviceResult = tuple[float, float, object]
+
+
+class BenchmarkCase(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"),
+ pytest.param(BenchmarkSpec("tetracene", TETRACENE), id="tetracene"),
+]
+
+
+def _make_benchmark_case(
+ spec: BenchmarkSpec, functional: ExcFunctionalBase, backend: str
+) -> BenchmarkCase:
+ mol = gto.M(atom=spec.atoms, basis="def2-qzvpp", verbose=0)
+ initial_dm = dft.RKS(mol).get_init_guess()
+
+ if backend == "cpu":
+ grids = SkalaGrids(mol)
+ grids.level = 1
+ grids.alignment = 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 = numint.nr_rks(
+ mol,
+ grids,
+ None,
+ dm,
+ max_memory=MAX_MEMORY_MB,
+ )
+ synchronize()
+ return cast(DeviceResult, result)
+
+ return BenchmarkCase(backend, mol, run, synchronize)
+
+
+@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=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:
+ functional = load_functional_cached("skala-1.1")
+ assert isinstance(functional, ExcFunctionalBase)
+ return _make_benchmark_case(benchmark_spec, functional, "cpu")
+
+
+@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],
+) -> Iterator[BenchmarkCase]:
+ backend = cast(str, request.param)
+ if backend == "cpu":
+ functional = load_functional_cached("skala-1.1")
+ assert isinstance(functional, ExcFunctionalBase)
+ 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)
+ with patch_ao_screening(True):
+ yield case
+
+
+@pytest.fixture
+def screened_case(benchmark_case: BenchmarkCase) -> Iterator[BenchmarkCase]:
+ with patch_ao_screening(True):
+ yield benchmark_case
+
+
+@pytest.fixture
+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:
+ case.synchronize()
+ pedantic = cast(Callable[..., object], benchmark.pedantic)
+ pedantic(
+ case.run,
+ rounds=1,
+ iterations=2,
+ )
+
+
+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")
+
+ torch.cuda.synchronize()
+ torch.cuda.empty_cache()
+ baseline_bytes = torch.cuda.memory_allocated()
+ torch.cuda.reset_peak_memory_stats()
+ with patch_ao_screening(screened):
+ case.run()
+ 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, "cpu")
+
+ with tempfile.TemporaryDirectory() as tmpdir:
+ profile_path = Path(tmpdir) / "allocations.bin"
+ with (
+ patch_ao_screening(screened),
+ memray.Tracker(profile_path),
+ ):
+ case.run()
+ 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(
+ device_benchmark_case: BenchmarkCase,
+ benchmark_spec: BenchmarkSpec,
+ load_functional_cached: Callable[..., ExcFunctionalBase | str],
+) -> None:
+ case = device_benchmark_case
+ screened = case.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()
+
+ 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)
+ 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_vxc
+ )
+ 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}, "
+ 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_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_decision(
+ benchmark: BenchmarkFixture, dense_case: BenchmarkCase
+) -> None:
+ _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: BenchmarkCase,
+) -> None:
+ _benchmark_device_xc(benchmark, device_benchmark_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_ao_screening(
+ device_benchmark_case: BenchmarkCase,
+) -> None:
+ device_benchmark_case.run()
+
+
+@pytest.mark.profiling
+def test_profile_without_ao_screening_by_patching_decision(
+ device_benchmark_case: BenchmarkCase,
+) -> None:
+ with patch_ao_screening(False):
+ device_benchmark_case.run()
diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py
new file mode 100644
index 00000000..93d0fe53
--- /dev/null
+++ b/tests/test_evaluation.py
@@ -0,0 +1,69 @@
+from dataclasses import FrozenInstanceError
+
+import pytest
+
+from skala.features import Feature
+from skala.pyscf.evaluation import EvaluationPolicy, FeatureSpec
+
+
+def test_feature_name_parses_model_metadata_string() -> None:
+ assert Feature("density") is Feature.DENSITY
+
+
+@pytest.mark.parametrize(
+ ("features", "expected_order"),
+ [
+ ([], 0),
+ ([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[Feature], expected_order: int
+) -> None:
+ spec = FeatureSpec(features)
+
+ assert spec.requires_ao_evaluation is bool(features)
+ assert spec.ao_derivative_order == expected_order
+ 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"),
+ [
+ (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: Feature, supports_screened_evaluation: bool
+) -> None:
+ spec = FeatureSpec([feature, feature])
+
+ assert spec.names == frozenset({feature})
+ assert spec.requires_atomic_layout
+ assert spec.supports_spatial_decomposition 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]
diff --git a/tests/test_gpu4pyscf_ao_screening.py b/tests/test_gpu4pyscf_ao_screening.py
new file mode 100644
index 00000000..73b24c51
--- /dev/null
+++ b/tests/test_gpu4pyscf_ao_screening.py
@@ -0,0 +1,441 @@
+from collections.abc import Callable, Iterator
+from types import SimpleNamespace
+
+import numpy as np
+import pytest
+import torch
+from pyscf import dft, gto
+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.",
+ 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 utils import QuadraticFunctional, patch_ao_screening # 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.gpu4pyscf.grids import SkalaGrids as GPU4PySCFSkalaGrids # noqa: E402
+from skala.pyscf.ao_evaluation import ( # noqa: E402
+ ChunkEvalForward,
+ evaluate_full_grid,
+)
+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
+C 1.4 0.0 0.0
+C 2.8 0.0 0.0
+C 4.2 0.0 0.0
+"""
+
+
+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
+
+ 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
+ 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)
+ )
+ assert layout.forward_permutation.device.type == "cuda"
+ 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"),
+ [
+ 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],
+ atom: str,
+ spin: int,
+ integration_method_name: str,
+) -> None:
+ 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)
+ 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()
+ integrate = getattr(ks._numint, integration_method_name)
+
+ with patch_ao_screening(False):
+ dense = integrate(mol, ks.grids, None, dm)
+
+ with patch_ao_screening(True):
+ screened = integrate(mol, ks.grids, None, dm)
+
+ cupy.testing.assert_allclose(dense[0], screened[0], rtol=1e-9)
+ assert np.isclose(dense[1], screened[1], rtol=1e-9)
+ cupy.testing.assert_allclose(dense[2], screened[2], rtol=1e-8, atol=2e-9)
+
+
+@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)
+ 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)
+
+ with patch_ao_screening(True):
+ screened_response = ks._numint.gen_response(mo_coeff, mo_occ, ks=ks)
+
+ cupy.testing.assert_allclose(
+ screened_response(dm1),
+ dense_response(dm1),
+ rtol=1e-9,
+ atol=1e-10,
+ )
+
+
+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=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)
+ 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())
+
+ 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)
+
+ cupy.testing.assert_allclose(
+ screened_response(dm1),
+ dense_response(dm1),
+ rtol=1e-9,
+ atol=1e-9,
+ )
+
+
+def test_gpu_screened_skala_matches_cpu_on_carbon_chain(
+ load_functional_cached: Callable[..., ExcFunctionalBase | str],
+) -> None:
+ """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. 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.
+ """
+ mol = gto.M(atom=CARBON_CHAIN, basis="def2-qzvpp", verbose=0)
+ cpu_grids = PySCFSkalaGrids(mol)
+ cpu_grids.level = 1
+ cpu_grids.alignment = 1
+ cpu_grids.build(sort_grids=False)
+ gpu_grids = GPU4PySCFSkalaGrids(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,
+ )
+ 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)
+
+ with patch_ao_screening(False):
+ cpu_result = SkalaNumInt(cpu_functional, device=torch.device("cpu")).nr_rks(
+ mol, cpu_grids, None, 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
+ 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-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}; "
+ 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_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(
+ 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
+ )
+ 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)
+ (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)
+ 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 * block_size, dtype=cupy.float64).reshape(
+ active_sorted_aos.size, block_size
+ )
+ 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]]:
+ 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:],
+ )
+
+ monkeypatch.setattr(dft_gpu.numint, "NumInt", FakeGpuNumInt)
+
+ feature_function = MGGAFeatureFunction(FeatureSpec([Feature.DENSITY]))
+ 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
+ )
+
+ 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.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()
+ (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
diff --git a/tests/test_gpu4pyscf_classes.py b/tests/test_gpu4pyscf_classes.py
index 7aa28865..ad874d8a 100644
--- a/tests/test_gpu4pyscf_classes.py
+++ b/tests/test_gpu4pyscf_classes.py
@@ -4,17 +4,19 @@
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.",
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 SkalaGrids # noqa: E402
@pytest.fixture(params=["skala-1.0", "skala-1.1"])
@@ -71,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)
@@ -85,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_gpu4pyscf_gradients.py b/tests/test_gpu4pyscf_gradients.py
index 03e3c0cd..edf64cb9 100644
--- a/tests/test_gpu4pyscf_gradients.py
+++ b/tests/test_gpu4pyscf_gradients.py
@@ -2,6 +2,9 @@
import pytest
import torch
+from torch.utils.dlpack import from_dlpack
+
+pytestmark = pytest.mark.gpu
if not torch.cuda.is_available():
pytest.skip(
@@ -16,21 +19,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 _ridders import num_grad_ridders # noqa: E402
+from gpu4pyscf import dft, scf # noqa: E402
+from pyscf import gto # noqa: E402
+from test_pyscf_gradients import FULL_GRAD_REF # noqa: E402
+from utils import patch_ao_screening # noqa: E402
-from skala.functional.base import ExcFunctionalBase
-from skala.gpu4pyscf import SkalaKS
-from skala.gpu4pyscf.gradients import (
+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
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.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:
@@ -102,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)
@@ -110,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)
@@ -137,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)
@@ -160,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
@@ -178,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)
@@ -193,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
)
@@ -212,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
@@ -246,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)
@@ -267,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()
@@ -307,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(
@@ -327,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
@@ -361,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(
@@ -439,6 +466,67 @@ def test_full_grad(
)
+def test_nuclear_gradient_cpu_gpu_dense_screened_agree(
+ 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,
+ )
+ 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)
+
+ 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."""
@@ -448,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()
@@ -465,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)
@@ -476,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
new file mode 100644
index 00000000..e0a92655
--- /dev/null
+++ b/tests/test_memory_estimators.py
@@ -0,0 +1,97 @@
+import pytest
+import torch
+
+from skala.pyscf.memory_estimators import (
+ estimate_global_raw_feature_buffer_memory,
+ estimate_global_screened_buffer_memory,
+ estimate_max_model_atoms_per_chunk,
+ estimate_model_memory_per_grid_point,
+)
+
+
+@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_global_screened_buffer_memory_uses_atomic_grid_sizes() -> None:
+ dm = torch.eye(10, dtype=torch.float64)
+ 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
+ )
+
+ 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,
+ atomic_grid_sizes=atomic_grid_sizes,
+ nfeatures=5,
+ max_memory_in_mb=10,
+ safety_fraction=1.0,
+ func_deriv=1,
+ )
+
+ 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])
+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_model_atoms_per_chunk(
+ torch.eye(2, dtype=torch.float64),
+ atomic_grid_sizes=torch.tensor([10]),
+ nfeatures=5,
+ max_memory_in_mb=100,
+ safety_fraction=safety_fraction,
+ )
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_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)),
+ ]
diff --git a/tests/test_pyscf_classes.py b/tests/test_pyscf_classes.py
index 8b6c6756..b3979a02 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,34 @@ 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
+ mol = gto.M(atom="H 0 0 0; H 0 0 0.74", basis="sto-3g", verbose=0)
+ ks = SkalaRKS(mol, xc=skala_xc)
+ ks.grids = dft.gen_grid.Grids(mol)
- import torch
+ with pytest.raises(TypeError, match="SkalaRKS requires .*SkalaGrids"):
+ ks.initialize_grids()
- 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 test_skala_grids_require_unit_alignment() -> None:
+ mol = gto.M(atom="H", basis="sto-3g", spin=1, verbose=0)
+ grids = SkalaGrids(mol)
- 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
+ assert grids.alignment == 1
+ grids.alignment = 1
+ with pytest.raises(ValueError, match="alignment must be 1"):
+ grids.alignment = 8
- with patch("skala.pyscf.dft._build_grids_unsorted", _build_grids_keep_padding):
- ks = SkalaKS(mol, xc=func, with_dftd3=False)
- # 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_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)
- with pytest.raises(ValueError, match="Grid size mismatch"):
- generate_features(mol, dm, ks.grids, set(func.features))
+ assert SkalaRKS(rks_mol, xc=skala_xc).small_rho_cutoff == 0
+ assert SkalaUKS(uks_mol, xc=skala_xc).small_rho_cutoff == 0
diff --git a/tests/test_pyscf_gradients.py b/tests/test_pyscf_gradients.py
index 4f086b53..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)
@@ -556,7 +563,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)
diff --git a/tests/test_xc_integrator.py b/tests/test_xc_integrator.py
new file mode 100644
index 00000000..b427ef8e
--- /dev/null
+++ b/tests/test_xc_integrator.py
@@ -0,0 +1,130 @@
+import pytest
+import torch
+from pyscf import dft, gto
+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.grids import SkalaGrids
+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 = SkalaGrids(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:
+ """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 = SkalaGrids(mol)
+
+ def fake_generate_features(
+ mol: gto.Mole,
+ dm: torch.Tensor,
+ grids: object,
+ features: set[Feature],
+ **kwargs: object,
+ ) -> FeatureMap:
+ assert features == {Feature.DENSITY, Feature.GRID_WEIGHTS}
+ return {
+ Feature.DENSITY: dm.sum().reshape(1),
+ Feature.GRID_WEIGHTS: torch.tensor([2.0], dtype=dm.dtype),
+ }
+
+ monkeypatch.setattr(
+ xc_integrator_module,
+ "generate_features",
+ fake_generate_features,
+ )
+ integrator = XCIntegrator(QuadraticFunctional([Feature.DENSITY]))
+ 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))
+
+
+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)
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