diff --git a/.github/actions/setup-pixi/action.yml b/.github/actions/setup-pixi/action.yml new file mode 100644 index 00000000..35494b81 --- /dev/null +++ b/.github/actions/setup-pixi/action.yml @@ -0,0 +1,34 @@ +name: Setup Pixi +description: Install Pixi and activate a locked project environment. + +inputs: + environment: + description: Pixi environment to install and activate. + required: true + working-directory: + description: Directory containing the Pixi project. + required: false + default: '' + pixi-bin-path: + description: Optional path at which to install the Pixi binary. + required: false + default: '' + post-cleanup: + description: Whether to remove Pixi and its environment after the job. + required: false + default: 'false' + +runs: + using: composite + steps: + - name: Setup Pixi + uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + with: + pixi-version: v0.78.0 + working-directory: ${{ inputs.working-directory }} + pixi-bin-path: ${{ inputs.pixi-bin-path }} + environments: ${{ inputs.environment }} + activate-environment: true + locked: true + cache: false + post-cleanup: ${{ inputs.post-cleanup }} \ No newline at end of file diff --git a/.github/workflows/benchmark-test.yml b/.github/workflows/benchmark-test.yml index 1ca715bf..0dcf3205 100644 --- a/.github/workflows/benchmark-test.yml +++ b/.github/workflows/benchmark-test.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &benchmark-test-paths - '.github/workflows/benchmark-test.yml' + - '.github/actions/setup-pixi/**' - 'benchmark/**' - 'website/_ext/benchmark_report.py' - 'pixi.lock' @@ -33,13 +34,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: default - activate-environment: true - locked: true - cache: false + environment: default - name: Run benchmark unit tests run: pytest -v benchmark/tests/ diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index d302aa47..36d7b8c5 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &docs-paths - '.github/workflows/docs.yml' + - '.github/actions/setup-pixi/**' - 'benchmark/pyproject.toml' - 'pixi.lock' - 'pixi.toml' @@ -38,13 +39,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: docs - activate-environment: true - locked: true - cache: false + environment: docs - name: Build documentation run: | diff --git a/.github/workflows/examples.yml b/.github/workflows/examples.yml index da7f84ea..cf192fb8 100644 --- a/.github/workflows/examples.yml +++ b/.github/workflows/examples.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &gauxc-paths - '.github/workflows/examples.yml' + - '.github/actions/setup-pixi/**' - 'gauxc/**' - 'pixi.lock' - 'pixi.toml' @@ -27,13 +28,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: assets - activate-environment: true - locked: true - cache: false + environment: assets - name: Download checkpoint run: >- @@ -86,14 +83,10 @@ jobs: uses: ./skala/.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./skala/.github/actions/setup-pixi with: - pixi-version: v0.75.0 working-directory: skala - environments: ${{ matrix.environment }} - activate-environment: true - locked: true - cache: false + environment: ${{ matrix.environment }} - name: Checkout development GauXC if: ${{ matrix.source == 'development' }} diff --git a/.github/workflows/gauxc-test.yml b/.github/workflows/gauxc-test.yml index cde33362..9dc85ec8 100644 --- a/.github/workflows/gauxc-test.yml +++ b/.github/workflows/gauxc-test.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &gauxc-test-paths - '.github/workflows/gauxc-test.yml' + - '.github/actions/setup-pixi/**' - 'gauxc/**' - 'pixi.lock' - 'pixi.toml' @@ -29,13 +30,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: default - activate-environment: true - locked: true - cache: false + environment: default - name: Run GauXC unit tests run: pytest -v gauxc/tests/ \ No newline at end of file diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index b97f764b..80b2d370 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -17,13 +17,9 @@ jobs: - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 # v4 - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: lint - activate-environment: true - locked: true - cache: false + environment: lint - name: Run pre-commit hooks run: pre-commit run --all-files diff --git a/.github/workflows/model-benchmark.yml b/.github/workflows/model-benchmark.yml index ccc881ed..68414b54 100644 --- a/.github/workflows/model-benchmark.yml +++ b/.github/workflows/model-benchmark.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &model-benchmark-paths - '.github/workflows/model-benchmark.yml' + - '.github/actions/setup-pixi/**' - 'conftest.py' - 'pixi.lock' - 'pixi.toml' @@ -35,14 +36,10 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 + environment: default pixi-bin-path: ${{ runner.temp }}/bin/pixi - environments: default - activate-environment: true - locked: true - cache: false post-cleanup: true - name: Run CPU model benchmarks @@ -63,14 +60,10 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 + environment: gpu-cuda12-torch213 pixi-bin-path: ${{ runner.temp }}/bin/pixi - environments: gpu-cuda12-torch213 - activate-environment: true - locked: true - cache: false post-cleanup: true - name: Verify CUDA is available diff --git a/.github/workflows/model-examples.yml b/.github/workflows/model-examples.yml index e0e513a1..5b87f863 100644 --- a/.github/workflows/model-examples.yml +++ b/.github/workflows/model-examples.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &model-example-paths - '.github/workflows/model-examples.yml' + - '.github/actions/setup-pixi/**' - 'model/examples/**' - 'model/src/**' - 'pixi.lock' @@ -29,13 +30,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: assets - activate-environment: true - locked: true - cache: false + environment: assets - name: Download checkpoint run: >- @@ -56,13 +53,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: default - activate-environment: true - locked: true - cache: false + environment: default - name: Generate features run: >- @@ -88,13 +81,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: cpp-integration - activate-environment: true - locked: true - cache: false + environment: cpp-integration - name: Configure and build project run: | @@ -131,13 +120,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: ftorch - activate-environment: true - locked: true - cache: false + environment: ftorch - name: Configure, build, and install project run: | diff --git a/.github/workflows/model-test.yml b/.github/workflows/model-test.yml index a712118c..b67fa70e 100644 --- a/.github/workflows/model-test.yml +++ b/.github/workflows/model-test.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &model-test-paths - '.github/workflows/model-test.yml' + - '.github/actions/setup-pixi/**' - 'model/**' - 'pixi.lock' - 'pixi.toml' @@ -30,13 +31,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: default - activate-environment: true - locked: true - cache: false + environment: default - name: Run model unit tests run: pytest -v model/tests/test_model.py model/tests/test_utils.py \ No newline at end of file diff --git a/.github/workflows/pypi.yml b/.github/workflows/pypi.yml index 8c4e2161..05ba35a4 100644 --- a/.github/workflows/pypi.yml +++ b/.github/workflows/pypi.yml @@ -23,13 +23,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: release - activate-environment: true - locked: true - cache: false + environment: release - name: Build a binary wheel and a source tarball run: python tools/build_release.py ${{ matrix.package }} @@ -52,13 +48,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: release - activate-environment: true - locked: true - cache: false + environment: release - name: Build conda artifact run: pixi publish --path skala --target-dir dist-conda diff --git a/.github/workflows/skala-test.yml b/.github/workflows/skala-test.yml index c2ff1e72..5fe50e35 100644 --- a/.github/workflows/skala-test.yml +++ b/.github/workflows/skala-test.yml @@ -5,6 +5,7 @@ on: branches: [main] paths: &skala-test-paths - '.github/workflows/skala-test.yml' + - '.github/actions/setup-pixi/**' - 'pixi.lock' - 'pixi.toml' - 'pyproject.toml' @@ -48,13 +49,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: ${{ matrix.environment }} - activate-environment: true - locked: true - cache: false + environment: ${{ matrix.environment }} - name: Run tests with coverage run: >- @@ -88,14 +85,10 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 + environment: ${{ matrix.environment }} pixi-bin-path: ${{ runner.temp }}/bin/pixi - environments: ${{ matrix.environment }} - activate-environment: true - locked: true - cache: false post-cleanup: true - name: Verify CUDA is available @@ -122,13 +115,9 @@ jobs: uses: ./.github/actions/cpu-diagnostics - name: Setup Pixi - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 + uses: ./.github/actions/setup-pixi with: - pixi-version: v0.75.0 - environments: test-py312-pyscf214-torch213 - activate-environment: true - locked: true - cache: false + environment: test-py312-pyscf214-torch213 - name: Run profiling tests run: pytest -v -m profiling skala/tests/ \ No newline at end of file diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 73664794..9eece202 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -16,7 +16,7 @@ or contact [opencode@microsoft.com](mailto:opencode@microsoft.com) with any addi ## Development setup -Install Pixi 0.75, then create the default locked development environment from the repository root: +Install Pixi 0.78, then create the default locked development environment from the repository root: ```bash pixi install --locked -e default diff --git a/pixi.lock b/pixi.lock index 5733fe21..dd83ec99 100644 --- a/pixi.lock +++ b/pixi.lock @@ -6,18 +6,7 @@ platforms: - __linux=4.18 - __glibc=2.28 - __archspec=0=x86_64 -- name: linux-aarch64 - virtual-packages: - - __unix=0=0 - - __linux=4.18 - - __glibc=2.28 - - __archspec=0=aarch64 -- name: osx-arm64 - virtual-packages: - - __unix=0=0 - - __osx=13.0 - - __archspec=0=m1 -- name: p1 +- name: linux-64-cuda12 subdir: linux-64 virtual-packages: - __cuda=12 @@ -25,7 +14,7 @@ platforms: - __linux=4.18 - __glibc=2.28 - __archspec=0=x86_64 -- name: p2 +- name: linux-64-cuda13 subdir: linux-64 virtual-packages: - __cuda=13 @@ -33,6 +22,17 @@ platforms: - __linux=4.18 - __glibc=2.28 - __archspec=0=x86_64 +- name: linux-aarch64 + virtual-packages: + - __unix=0=0 + - __linux=4.18 + - __glibc=2.28 + - __archspec=0=aarch64 +- name: osx-arm64 + virtual-packages: + - __unix=0=0 + - __osx=13.0 + - __archspec=0=m1 environments: assets: channels: @@ -270,7 +270,7 @@ environments: - pypi: https://files.pythonhosted.org/packages/40/2f/10a2e9ad81bd5c2479a255d18c4f02c734a3a21bf4b36a236ba615919592/pyscf_dispersion-1.5.0-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl - pypi: https://files.pythonhosted.org/packages/88/30/8e48717aff32f11fd78f981192d0b567971eb285ea048414b7dac2a211cf/pyscf-2.14.0-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl - pypi: https://files.pythonhosted.org/packages/9e/e9/1a19e42cd43cc1365e127db6aae85e1c671da1d9a5d746f4d34a50edb577/h5py-3.16.0-cp312-cp312-manylinux_2_28_x86_64.whl - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/binutils_impl_linux-64-2.46.1-default_hfdba357_102.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/brotli-1.2.0-h505cf86_3.conda @@ -1618,7 +1618,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tzdata-2026c-h151e31d_0.conda - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/binutils_impl_linux-64-2.46.1-default_hfdba357_102.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/bzip2-1.0.8-hda65f42_10.conda @@ -1735,7 +1735,7 @@ environments: channels: - url: https://conda.anaconda.org/conda-forge/ packages: - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-h4610da3_2.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.15-h6bbde05_1.conda @@ -2023,7 +2023,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tzdata-2026c-h151e31d_0.conda - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-h4610da3_2.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.15-h6bbde05_1.conda @@ -2272,7 +2272,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tzdata-2026c-h151e31d_0.conda - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-h4610da3_2.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.15-h6bbde05_1.conda @@ -2507,7 +2507,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tzdata-2026c-h151e31d_0.conda - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-h4610da3_2.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.15-h6bbde05_1.conda @@ -2728,7 +2728,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tzdata-2026c-h151e31d_0.conda - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-h4610da3_2.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.15-h6bbde05_1.conda @@ -2844,7 +2844,7 @@ environments: indexes: - https://pypi.org/simple packages: - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-hb7a77c6_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.14-h2aa3ae6_4.conda @@ -3111,7 +3111,7 @@ environments: indexes: - https://pypi.org/simple packages: - p1: + linux-64-cuda12: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-hb7a77c6_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.14-h2aa3ae6_4.conda @@ -3378,7 +3378,7 @@ environments: indexes: - https://pypi.org/simple packages: - p2: + linux-64-cuda13: - conda: https://conda.anaconda.org/conda-forge/linux-64/_openmp_mutex-4.5-7_kmp_llvm.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-auth-0.10.4-hb7a77c6_1.conda - conda: https://conda.anaconda.org/conda-forge/linux-64/aws-c-cal-0.9.14-h2aa3ae6_4.conda diff --git a/pixi.toml b/pixi.toml index 728c55f2..a4c99fe8 100644 --- a/pixi.toml +++ b/pixi.toml @@ -10,7 +10,7 @@ platforms = [ { name = "linux-64-cuda12", platform = "linux-64", cuda = "12" }, { name = "linux-64-cuda13", platform = "linux-64", cuda = "13" }, ] -requires-pixi = ">=0.75,<0.76" +requires-pixi = ">=0.78,<0.79" preview = ["pixi-build"] [workspace.conda-pypi-map] diff --git a/skala/src/skala/gpu4pyscf/gradients.py b/skala/src/skala/gpu4pyscf/gradients.py index 38ff3224..17830bb3 100644 --- a/skala/src/skala/gpu4pyscf/gradients.py +++ b/skala/src/skala/gpu4pyscf/gradients.py @@ -20,6 +20,11 @@ from skala.dispersion import DFTD3Dispersion from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase +from skala.pyscf.gradient_core import ( + contract_ao_derivative_block, + feature_derivatives, + grid_derivative_block, +) LOG = logging.getLogger(__name__) @@ -91,37 +96,15 @@ def veff_and_expl_nuc_grad( 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] + nuc_feats = {feat: mol_feats[feat] for feat in nuc_grad_feats} other_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: - exc_mol_feats = ( - dict(zip(nuc_feat_names, nuc_feat_tensors, strict=True)) | other_feats - ) - return functional.get_exc(exc_mol_feats) - - # torch.func.vjp wraps primals in functional tensors that may not expose - # backing storage, which TorchScript traced models can reject. - diff_nuc_feat_tensors = [ - feat.detach().requires_grad_(True) for feat in nuc_feat_tensors - ] - if len(diff_nuc_feat_tensors) > 0: - exc = exc_feat_func(*diff_nuc_feat_tensors) - dExc_tuple = torch.autograd.grad( - exc, - tuple(diff_nuc_feat_tensors), - create_graph=False, - retain_graph=False, - allow_unused=False, - ) - else: - dExc_tuple = () - dExc: FeatureMap = {} - for i in range(len(dExc_tuple)): - dExc[nuc_feat_names[i]] = dExc_tuple[i].detach() + def exc_feat_func(differentiable_features: FeatureMap) -> torch.Tensor: + return functional.get_exc(differentiable_features | other_feats) + + dExc = feature_derivatives(exc_feat_func, nuc_feats) LOG.debug("autograd gradients for nuclear features done") @@ -143,94 +126,17 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: if ao_deriv == 0: ao = ao[None, ...] atm_end = atm_start + weight.shape[0] + dExc_atm = grid_derivative_block(dExc, atm_start, atm_end) # Calculate the contribution to veff for this atomic grid - veff_atm = torch.zeros((2, 3, nao, nao), dtype=rdm1.dtype, device=rdm1.device) - - if Feature.DENSITY in nuc_grad_feats: - veff_atm += torch.einsum( - "si, xip, iq -> sxpq", - dExc[Feature.DENSITY][:, atm_start:atm_end], - ao[1:4], - ao[0], - ) - - 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] - ) - # XX, XY, XZ = 4, 5, 6 - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[4], ao[0] - ) - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[5], ao[0] - ) - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[6], ao[0] - ) - # YX, YY, YZ = 5, 7, 8 - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[5], ao[0] - ) - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[7], ao[0] - ) - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[8], ao[0] - ) - # ZX, ZY, ZZ = 6, 8, 9 - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[6], ao[0] - ) - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[8], ao[0] - ) - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[9], ao[0] - ) - - 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 - ) - veff_atm[:, 0] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[5], ao[2]) / 2 - ) - veff_atm[:, 0] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[6], ao[3]) / 2 - ) - # YX, YY, YZ = 5, 7, 8 - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[5], ao[1]) / 2 - ) - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[7], ao[2]) / 2 - ) - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[8], ao[3]) / 2 - ) - # ZX, ZY, ZZ = 6, 8, 9 - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[6], ao[1]) / 2 - ) - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[8], ao[2]) / 2 - ) - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[9], ao[3]) / 2 - ) + veff_atm = contract_ao_derivative_block(ao, dExc_atm) - if Feature.GRID_COORDS in nuc_grad_feats: + if Feature.GRID_COORDS in dExc_atm: # also add the explicit grid coordinate dependence - nuc_grad[atm_id] += dExc[Feature.GRID_COORDS][atm_start:atm_end].sum(dim=0) + nuc_grad[atm_id] += dExc_atm[Feature.GRID_COORDS].sum(dim=0) - if Feature.GRID_WEIGHTS in nuc_grad_feats: - Exc_dgw = dExc[Feature.GRID_WEIGHTS][atm_start:atm_end] + if Feature.GRID_WEIGHTS in dExc_atm: + Exc_dgw = dExc_atm[Feature.GRID_WEIGHTS] 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, diff --git a/skala/src/skala/pyscf/gradient_core.py b/skala/src/skala/pyscf/gradient_core.py new file mode 100644 index 00000000..48ea2a22 --- /dev/null +++ b/skala/src/skala/pyscf/gradient_core.py @@ -0,0 +1,218 @@ +# SPDX-License-Identifier: MIT + +"""Backend-independent PyTorch operations for PySCF nuclear gradients.""" + +from collections.abc import Callable, Iterable, Iterator +from dataclasses import dataclass +from enum import IntEnum +from types import EllipsisType +from typing import ClassVar, TypeAlias + +import torch + +from skala.features import Feature, FeatureMap + + +class _Direction(IntEnum): + """Cartesian direction indices used by feature and potential tensors.""" + + X = 0 + Y = 1 + Z = 2 + + +@dataclass(frozen=True) +class _PackedAO: + """Views into PySCF's packed AO derivative dimension.""" + + tensor: torch.Tensor + + # PySCF's ``eval_ao(..., deriv=2)`` order is + # value, x, y, z, xx, xy, xz, yy, yz, zz. Mixed derivatives are reused + # across the symmetric off-diagonal entries of the Cartesian Hessian. + _HESSIAN_COMPONENTS: ClassVar[tuple[tuple[int, int, int], ...]] = ( + (4, 5, 6), + (5, 7, 8), + (6, 8, 9), + ) + + @property + def value(self) -> torch.Tensor: + """AO values with shape ``(npoints, nao)``.""" + return self.tensor[0] + + @property + def gradient(self) -> torch.Tensor: + """AO gradients with shape ``(3, npoints, nao)``.""" + return self.tensor[1:4] + + def hessian(self) -> Iterator[tuple[_Direction, _Direction, torch.Tensor]]: + """Yield both Cartesian directions and each AO Hessian component.""" + for ao_direction in _Direction: + for feature_direction in _Direction: + component = self._HESSIAN_COMPONENTS[ao_direction][feature_direction] + yield ao_direction, feature_direction, self.tensor[component] + + +_FeatureSlice: TypeAlias = tuple[slice] | tuple[EllipsisType, slice] + + +def _disconnected_features_error(features: Iterable[Feature]) -> RuntimeError: + feature_names = ", ".join(sorted(feature.value for feature in features)) + return RuntimeError( + f"XC energy is disconnected from requested features: {feature_names}" + ) + + +def feature_derivatives( + exc_func: Callable[[FeatureMap], torch.Tensor], features: FeatureMap +) -> FeatureMap: + """Differentiate a scalar XC energy with respect to molecular features. + + Ordinary autograd tensors are used instead of ``torch.func.vjp`` functional + tensors because traced TorchScript models may require accessible backing + storage. + + Args: + exc_func: Callable accepting the differentiable features and returning + scalar XC energy. + features: Molecular features to differentiate, keyed by feature name. + + Returns: + XC energy derivatives keyed by feature name. + + Raises: + RuntimeError: If the XC energy is disconnected from a requested feature. + + """ + if not features: + return {} + + differentiable_features = { + feature: tensor.detach().requires_grad_(True) + for feature, tensor in features.items() + } + exc = exc_func(differentiable_features) + if not exc.requires_grad: + raise _disconnected_features_error(differentiable_features) + + gradients = torch.autograd.grad( + exc, + tuple(differentiable_features.values()), + create_graph=False, + retain_graph=False, + allow_unused=True, + ) + derivatives: FeatureMap = { + feature: gradient.detach() + for feature, gradient in zip(differentiable_features, gradients, strict=True) + if gradient is not None + } + if len(derivatives) != len(differentiable_features): + disconnected_features = differentiable_features.keys() - derivatives.keys() + raise _disconnected_features_error(disconnected_features) + return derivatives + + +def grid_derivative_block( + derivatives: FeatureMap, grid_start: int, grid_end: int +) -> FeatureMap: + """Select a slice of the feature derivatives from `grid_start` to `grid_end`. + + Density-like features store the grid-point dimension last, whereas grid + coordinates and weights store it first. Non-grid features, such as atomic + coordinates, are intentionally omitted. + + Args: + derivatives: XC energy derivatives keyed by molecular feature. + grid_start: Start of this block in the full molecular feature tensors. + grid_end: End of this block in the full molecular feature tensors. + + Returns: + Grid-resolved derivatives restricted to ``[grid_start:grid_end]``. + """ + grid_slice = slice(grid_start, grid_end) + feature_slices: dict[Feature, _FeatureSlice] = { + Feature.GRID_COORDS: (grid_slice,), + Feature.GRID_WEIGHTS: (grid_slice,), + Feature.DENSITY: (..., grid_slice), + Feature.GRAD: (..., grid_slice), + Feature.KIN: (..., grid_slice), + } + return { + feature: derivative[feature_slice] + for feature, derivative in derivatives.items() + if (feature_slice := feature_slices.get(feature)) is not None + } + + +def contract_ao_derivative_block( + ao: torch.Tensor, feature_derivatives: FeatureMap +) -> torch.Tensor: + """Contract XC feature derivatives with one-sided spatial AO derivatives. + + For each Cartesian direction, form one grid block's contribution to the + spatial derivative of the AO-basis XC potential. The AO associated with + the first matrix index is differentiated, while the AO associated with the + second index is held fixed. Consequently, the returned matrices are not + generally symmetric in their AO indices. + + Density, density-gradient, and kinetic-energy-density contributions are + included when their corresponding feature derivatives are present. + + Args: + ao: Packed component-major AO values with shape + ``(ncomponents, npoints, nao)``. Components follow PySCF's ordering: + ``value, x, y, z, xx, xy, xz, yy, yz, zz``. + feature_derivatives: XC energy derivatives with respect to features on + this grid block. Density and kinetic derivatives have shape + ``(2, npoints)``; gradient derivatives have shape + ``(2, 3, npoints)``. + + Returns: + One-sided AO-potential derivatives with shape ``(2, 3, nao, nao)``, + indexed by spin, Cartesian direction, and the two AO indices. The result + has the same device and dtype as ``ao``. + """ + packed_ao = _PackedAO(ao) + nao = ao.shape[-1] + potential_derivatives = ao.new_zeros((2, 3, nao, nao)) + + if Feature.DENSITY in feature_derivatives: + potential_derivatives += torch.einsum( + "si, xip, iq -> sxpq", + feature_derivatives[Feature.DENSITY], + packed_ao.gradient, + packed_ao.value, + ) + + if Feature.GRAD in feature_derivatives: + exc_dgrad = feature_derivatives[Feature.GRAD] + potential_derivatives += torch.einsum( + "syi, xip, yiq -> sxpq", + exc_dgrad, + packed_ao.gradient, + packed_ao.gradient, + ) + for ao_direction, feature_direction, ao_hessian in packed_ao.hessian(): + potential_derivatives[:, ao_direction] += torch.einsum( + "si, ip, iq -> spq", + exc_dgrad[:, feature_direction], + ao_hessian, + packed_ao.value, + ) + + if Feature.KIN in feature_derivatives: + exc_dkin = feature_derivatives[Feature.KIN] + for ao_direction, feature_direction, ao_hessian in packed_ao.hessian(): + potential_derivatives[:, ao_direction] += ( + torch.einsum( + "si, ip, iq -> spq", + exc_dkin, + ao_hessian, + packed_ao.gradient[feature_direction], + ) + / 2 + ) + + return potential_derivatives diff --git a/skala/src/skala/pyscf/gradients.py b/skala/src/skala/pyscf/gradients.py index c7af698b..ed2d405f 100644 --- a/skala/src/skala/pyscf/gradients.py +++ b/skala/src/skala/pyscf/gradients.py @@ -17,6 +17,11 @@ from skala.dispersion import DFTD3Dispersion from skala.features import Feature, FeatureMap from skala.functional.base import ExcFunctionalBase +from skala.pyscf.gradient_core import ( + contract_ao_derivative_block, + feature_derivatives, + grid_derivative_block, +) LOG = logging.getLogger(__name__) @@ -86,25 +91,17 @@ def veff_and_expl_nuc_grad( 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] + nuc_feats = {feat: mol_feats[feat] for feat in nuc_grad_feats} other_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: - exc_mol_feats = ( - dict(zip(nuc_feat_names, nuc_feat_tensors, strict=True)) | other_feats - ) - return functional.get_exc(exc_mol_feats) + def exc_feat_func(differentiable_features: FeatureMap) -> torch.Tensor: + return functional.get_exc(differentiable_features | other_feats) - dExc_func = torch.func.vjp(exc_feat_func, *nuc_feat_tensors)[1] - dExc_tuple = dExc_func(torch.tensor(1.0, dtype=rdm1.dtype)) - dExc: FeatureMap = {} - for i in range(len(dExc_tuple)): - dExc[nuc_feat_names[i]] = dExc_tuple[i].detach() + dExc = feature_derivatives(exc_feat_func, nuc_feats) - LOG.debug("torch.func.vjp done") + LOG.debug("autograd gradients for nuclear features done") nao = rdm1.shape[-1] veff = torch.zeros((2, 3, nao, nao), dtype=rdm1.dtype) @@ -121,94 +118,17 @@ def exc_feat_func(*nuc_feat_tensors: torch.Tensor) -> torch.Tensor: if ao_deriv == 0: ao = ao[None, ...] atm_end = atm_start + weight.shape[0] + dExc_atm = grid_derivative_block(dExc, atm_start, atm_end) # Calculate the contribution to veff for this atomic grid - veff_atm = torch.zeros((2, 3, nao, nao), dtype=rdm1.dtype) - - if Feature.DENSITY in nuc_grad_feats: - veff_atm += torch.einsum( - "si, xip, iq -> sxpq", - dExc[Feature.DENSITY][:, atm_start:atm_end], - ao[1:4], - ao[0], - ) - - 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] - ) - # XX, XY, XZ = 4, 5, 6 - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[4], ao[0] - ) - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[5], ao[0] - ) - veff_atm[:, 0] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[6], ao[0] - ) - # YX, YY, YZ = 5, 7, 8 - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[5], ao[0] - ) - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[7], ao[0] - ) - veff_atm[:, 1] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[8], ao[0] - ) - # ZX, ZY, ZZ = 6, 8, 9 - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 0], ao[6], ao[0] - ) - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 1], ao[8], ao[0] - ) - veff_atm[:, 2] += torch.einsum( - "si, ip, iq -> spq", Exc_dgrad_atm[:, 2], ao[9], ao[0] - ) - - 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 - ) - veff_atm[:, 0] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[5], ao[2]) / 2 - ) - veff_atm[:, 0] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[6], ao[3]) / 2 - ) - # YX, YY, YZ = 5, 7, 8 - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[5], ao[1]) / 2 - ) - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[7], ao[2]) / 2 - ) - veff_atm[:, 1] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[8], ao[3]) / 2 - ) - # ZX, ZY, ZZ = 6, 8, 9 - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[6], ao[1]) / 2 - ) - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[8], ao[2]) / 2 - ) - veff_atm[:, 2] += ( - torch.einsum("si, ip, iq -> spq", Exc_dkin_atm, ao[9], ao[3]) / 2 - ) + veff_atm = contract_ao_derivative_block(ao, dExc_atm) - if Feature.GRID_COORDS in nuc_grad_feats: + if Feature.GRID_COORDS in dExc_atm: # also add the explicit grid coordinate dependence - nuc_grad[atm_id] += dExc[Feature.GRID_COORDS][atm_start:atm_end].sum(dim=0) + nuc_grad[atm_id] += dExc_atm[Feature.GRID_COORDS].sum(dim=0) - if Feature.GRID_WEIGHTS in nuc_grad_feats: - Exc_dgw = dExc[Feature.GRID_WEIGHTS][atm_start:atm_end] + if Feature.GRID_WEIGHTS in dExc_atm: + Exc_dgw = dExc_atm[Feature.GRID_WEIGHTS] 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, diff --git a/skala/tests/test_gradient_core.py b/skala/tests/test_gradient_core.py new file mode 100644 index 00000000..214d217b --- /dev/null +++ b/skala/tests/test_gradient_core.py @@ -0,0 +1,111 @@ +# SPDX-License-Identifier: MIT + +"""Tests for backend-independent nuclear-gradient operations.""" + +import pytest +import torch +from skala.features import Feature +from skala.pyscf.gradient_core import ( + contract_ao_derivative_block, + feature_derivatives, + grid_derivative_block, +) + + +def test_feature_derivatives() -> None: + density = torch.tensor([1.0, 2.0], dtype=torch.float64) + weights = torch.tensor([3.0, 4.0], dtype=torch.float64) + + derivatives = feature_derivatives( + lambda features: ( + features[Feature.DENSITY].square() * features[Feature.GRID_WEIGHTS] + ).sum(), + {Feature.DENSITY: density, Feature.GRID_WEIGHTS: weights}, + ) + + torch.testing.assert_close(derivatives[Feature.DENSITY], 2 * density * weights) + torch.testing.assert_close(derivatives[Feature.GRID_WEIGHTS], density.square()) + assert feature_derivatives(lambda _: torch.tensor(0.0), {}) == {} + + +def test_feature_derivatives_rejects_disconnected_feature() -> None: + used = torch.tensor(2.0) + unused = torch.tensor(3.0) + + with pytest.raises( + RuntimeError, + match="XC energy is disconnected from requested features: grid_weights", + ): + feature_derivatives( + lambda features: features[Feature.DENSITY].square(), + {Feature.DENSITY: used, Feature.GRID_WEIGHTS: unused}, + ) + + +def test_feature_derivatives_rejects_constant_energy() -> None: + feature = torch.tensor(2.0) + + with pytest.raises( + RuntimeError, + match="XC energy is disconnected from requested features: density", + ): + feature_derivatives(lambda _: torch.tensor(1.0), {Feature.DENSITY: feature}) + + +def test_grid_derivative_block_slices_each_grid_dimension() -> None: + derivatives = { + Feature.DENSITY: torch.arange(8).reshape(2, 4), + Feature.GRAD: torch.arange(24).reshape(2, 3, 4), + Feature.GRID_COORDS: torch.arange(12).reshape(4, 3), + Feature.GRID_WEIGHTS: torch.arange(4), + Feature.COARSE_0_ATOMIC_COORDS: torch.arange(6).reshape(2, 3), + } + + block = grid_derivative_block(derivatives, grid_start=1, grid_end=3) + + assert set(block) == { + Feature.DENSITY, + Feature.GRAD, + Feature.GRID_COORDS, + Feature.GRID_WEIGHTS, + } + torch.testing.assert_close( + block[Feature.DENSITY], derivatives[Feature.DENSITY][..., 1:3] + ) + torch.testing.assert_close(block[Feature.GRAD], derivatives[Feature.GRAD][..., 1:3]) + torch.testing.assert_close( + block[Feature.GRID_COORDS], derivatives[Feature.GRID_COORDS][1:3] + ) + torch.testing.assert_close( + block[Feature.GRID_WEIGHTS], derivatives[Feature.GRID_WEIGHTS][1:3] + ) + + +def test_contract_ao_derivative_block_matches_reference() -> None: + ao = torch.tensor( + [2.0, 3.0, 5.0, 7.0, 11.0, 13.0, 17.0, 19.0, 23.0, 29.0], + dtype=torch.float64, + ).reshape(10, 1, 1) + derivatives = { + Feature.DENSITY: torch.tensor([[0.0, 31.0], [0.0, 37.0]], dtype=torch.float64), + Feature.GRAD: torch.tensor( + [ + [[0.0, 41.0], [0.0, 43.0], [0.0, 47.0]], + [[0.0, 53.0], [0.0, 59.0], [0.0, 61.0]], + ], + dtype=torch.float64, + ), + Feature.KIN: torch.tensor([[0.0, 67.0], [0.0, 71.0]], dtype=torch.float64), + } + actual = contract_ao_derivative_block(ao, grid_derivative_block(derivatives, 1, 2)) + expected = torch.tensor( + [ + [[[13074.5]], [[18389.5]], [[23562.5]]], + [[[15342.5]], [[21673.5]], [[27838.5]]], + ], + dtype=torch.float64, + ) + + torch.testing.assert_close(actual, expected) + assert actual.dtype == ao.dtype + assert actual.device == ao.device