From a084bf3f58ea4a6845adf77b2da5f9671ab0b87c Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 17:37:15 +0000 Subject: [PATCH 01/16] Compile the GPU kernels ahead of time as CUDA and drop Triton The EPG, pooled and PERK kernels are C++ over a small tile runtime (_tile.hpp), compiled by nvcc into blochsim._gpu, one file per kernel with 256- and 1024-thread bounded variants, linking the CUDA runtime statically. CMake builds it wherever it finds nvcc (BLOCHSIM_CUDA overrides), and the x86-64 manylinux wheel carries it. The same source compiled for the host (blochsim._gpu_host) replaces Triton's interpreter in the suite and is held to the C++ kernels. Triton, its modules, the interpreted marker and the cache-pruning script are gone. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H --- .github/ISSUE_TEMPLATE/bug_report.yml | 2 +- .github/SECURITY.md | 2 +- .github/skills/add-a-sequence/SKILL.md | 2 +- .github/skills/build-and-test/SKILL.md | 37 +- .github/workflows/test.yml | 17 +- .github/workflows/wheels.yml | 6 +- .pre-commit-config.yaml | 2 + CHANGELOG.md | 9 + CLAUDE.md | 62 +- CMakeLists.txt | 62 +- docs/developer_guide.md | 40 +- docs/explanation_figures.py | 2 +- docs/explanations/implementation.md | 9 +- docs/misc/related.md | 2 +- docs/user_guide.md | 19 +- pyproject.toml | 24 +- scripts/check_wheel.py | 15 +- scripts/prune_triton_cache.py | 142 - src/blochsim/_epg_kernels.hpp | 12222 ++++++++++ src/blochsim/_gpu.cu | 99 + src/blochsim/_gpu_host.cpp | 66 + src/blochsim/_gpu_kernel.cu.in | 20 + src/blochsim/_gpu_launch.py | 106 + src/blochsim/_kernels.hpp | 842 + src/blochsim/_launch.hpp | 117 + src/blochsim/_perk_kernels.hpp | 113 + src/blochsim/_pools_kernels.hpp | 1851 ++ src/blochsim/_tile.hpp | 1120 + src/blochsim/estimators/_dictionary.py | 2 +- src/blochsim/estimators/_perk.py | 9 +- src/blochsim/estimators/_perk_gpu.py | 113 + src/blochsim/estimators/_perk_triton.py | 332 - src/blochsim/sequence/_accelerators.py | 19 +- src/blochsim/sequence/_epg_gpu.py | 1893 ++ src/blochsim/sequence/_epg_triton.py | 18718 ---------------- src/blochsim/sequence/_parameters.py | 13 +- src/blochsim/sequence/_pools_gpu.py | 525 + src/blochsim/sequence/_pools_triton.py | 2226 -- tests/estimators/test_perk_kernel.py | 52 +- tests/sequence/test_both_pools.py | 82 +- tests/sequence/test_cuda_parity.py | 14 +- tests/sequence/test_dynamic_transmit.py | 2 +- tests/sequence/test_feature_gates.py | 12 +- tests/sequence/test_host_feature_mask.py | 3 +- tests/sequence/test_host_kernels.py | 66 + tests/sequence/test_interpreted.py | 101 - tests/sequence/test_many_pools_host.py | 42 + tests/sequence/test_many_pools_interpreted.py | 79 - tests/sequence/test_parameters.py | 2 +- tests/sequence/test_subspace_streams.py | 6 +- .../utils/{interpreted.py => host_kernels.py} | 232 +- .../{interpreted_pools.py => host_pools.py} | 37 +- 52 files changed, 19568 insertions(+), 22022 deletions(-) delete mode 100644 scripts/prune_triton_cache.py create mode 100644 src/blochsim/_epg_kernels.hpp create mode 100644 src/blochsim/_gpu.cu create mode 100644 src/blochsim/_gpu_host.cpp create mode 100644 src/blochsim/_gpu_kernel.cu.in create mode 100644 src/blochsim/_gpu_launch.py create mode 100644 src/blochsim/_kernels.hpp create mode 100644 src/blochsim/_launch.hpp create mode 100644 src/blochsim/_perk_kernels.hpp create mode 100644 src/blochsim/_pools_kernels.hpp create mode 100644 src/blochsim/_tile.hpp create mode 100644 src/blochsim/estimators/_perk_gpu.py delete mode 100644 src/blochsim/estimators/_perk_triton.py create mode 100644 src/blochsim/sequence/_epg_gpu.py delete mode 100644 src/blochsim/sequence/_epg_triton.py create mode 100644 src/blochsim/sequence/_pools_gpu.py delete mode 100644 src/blochsim/sequence/_pools_triton.py create mode 100644 tests/sequence/test_host_kernels.py delete mode 100644 tests/sequence/test_interpreted.py create mode 100644 tests/sequence/test_many_pools_host.py delete mode 100644 tests/sequence/test_many_pools_interpreted.py rename tests/utils/{interpreted.py => host_kernels.py} (81%) rename tests/utils/{interpreted_pools.py => host_pools.py} (86%) diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml index 8807346b..a623cbdf 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yml +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -63,7 +63,7 @@ body: multiple: true options: - CPU - - CUDA (Triton kernels) + - CUDA - Both validations: required: true diff --git a/.github/SECURITY.md b/.github/SECURITY.md index a3483bc1..48a394d0 100644 --- a/.github/SECURITY.md +++ b/.github/SECURITY.md @@ -29,7 +29,7 @@ a signal model are all Python that runs in your process, so a malicious In scope is anything that turns *data* into execution or into memory corruption -- a sequence description, a pulse waveform, a phantom or a dictionary read from a file, or values passed to the simulator, reaching the -C++ and Triton kernels. Those kernels index raw pointers, so an out-of-bounds +C++ and CUDA kernels. Those kernels index raw pointers, so an out-of-bounds read or write reachable from ordinary arguments is a vulnerability and not merely a bug. diff --git a/.github/skills/add-a-sequence/SKILL.md b/.github/skills/add-a-sequence/SKILL.md index db4c4bca..e3119d90 100644 --- a/.github/skills/add-a-sequence/SKILL.md +++ b/.github/skills/add-a-sequence/SKILL.md @@ -63,7 +63,7 @@ evaluates. Adding an effect is one line in a model's property declaration, plus the term itself in **both** kernels. The shared parameter ABI is `src/blochsim/sequence/_parameters.py`, read by the -Python dispatch, the C++ extension and the Triton kernels. A parameter added +Python dispatch, the C++ extension and the GPU kernels. A parameter added there is added in all three or in none. ## What the change has to arrive with diff --git a/.github/skills/build-and-test/SKILL.md b/.github/skills/build-and-test/SKILL.md index 3f851c91..6eb5e8a5 100644 --- a/.github/skills/build-and-test/SKILL.md +++ b/.github/skills/build-and-test/SKILL.md @@ -1,6 +1,6 @@ --- name: build-and-test -description: Compile the C++ kernels and run the BlochSim suite, including the Triton paths on a machine with no GPU. Use when asked to build, install, test, or reproduce a failure in blochsim. +description: Compile the C++ kernels and run the BlochSim suite, including the GPU kernels on a machine with no GPU. Use when asked to build, install, test, or reproduce a failure in blochsim. --- # Build and test BlochSim @@ -16,7 +16,8 @@ run alone, `pip install -e ".[test]"` is enough and much smaller. An editable install puts the Python sources on the path, so an edit under `src/blochsim` takes effect on the next import. The two C++ extensions are -compiled artifacts and do not: re-run the install after editing any `.cpp`. +compiled artifacts and do not: re-run the install after editing any `.cpp`, +`.hpp` or `.cu`. **Read the exit status, not the output.** A failed compile leaves the previously built `.so` importable, so the suite runs green against a kernel @@ -41,32 +42,14 @@ against closed forms — the shift, the RF rotation, relaxation, diffusion, flow spoiling, the two-pool and three-pool longitudinal steps — while `sequence/`, `model/`, `estimators/`, `recon/` and `optim/` cover the layers above. -## The Triton paths, without a GPU +## The GPU kernels, without a GPU -The tests carrying the `interpreted` marker are deselected by `addopts`. They -run a Triton kernel through Triton's CPU interpreter — no card, no compile, -about a minute each — and are how the GPU plumbing is verified anywhere: - -```sh -pytest tests/ -m interpreted -TRITON_INTERPRET=1 python your_script.py # the same trick, by hand -``` - -Triton has no wheel for macOS or Windows. On those platforms the default -deselection is not an optimisation, it is the only thing that runs. - -## Before you time anything - -Kernel compiles dominate a cold GPU run, not the arithmetic. A suite that takes -minutes on a card is mostly Triton compiling one specialization per feature -combination it meets; the second run of the same suite is a different number -entirely. Run the whole suite at natural boundaries rather than after every -edit. - -The cache keys on the *source* of `_epg_triton.py`, not on what it means, so a -formatting pass or a deleted dead local costs the same full recompile as a new -`tl.constexpr`. If the suite suddenly takes an hour where it took minutes, -check whether that file changed before looking for a performance regression. +The install compiles the GPU kernels for the host as `blochsim._gpu_host`, and +for the card as `blochsim._gpu` wherever CMake finds `nvcc` +(`--config-settings=cmake.define.BLOCHSIM_CUDA=ON` insists on it). +`tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` run the +host build against the C++ kernels, which is how the GPU kernels are verified +on any machine; the CUDA tests skip themselves without a card. ## Style diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index aaec7d27..48e16c1e 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -48,14 +48,6 @@ jobs: if: runner.os == 'Linux' run: python -m pip install torch --index-url https://download.pytorch.org/whl/cpu - # Triton reaches a runner as a dependency of the CUDA build of PyTorch, - # which the CPU index above deliberately leaves out, and it exists for - # Linux alone. Installing it is what lets the kernel-gate tests run; on - # the other two platforms they skip themselves. - - name: Install Triton (Linux only) - if: runner.os == 'Linux' - run: python -m pip install triton - # Not editable and not from a wheel: the C++ kernels are compiled here, # on this runner, with this compiler, which is the half of the package a # pure-Python job would never exercise. CMake and Ninja arrive as @@ -69,13 +61,12 @@ jobs: import importlib.util, pathlib spec = importlib.util.find_spec("blochsim") root = pathlib.Path(next(iter(spec.submodule_search_locations))) - for kernel in sorted(root.glob("_*_cpu*")): + for kernel in sorted(root.glob("_*_cpu*")) + sorted(root.glob("_gpu*")): print(kernel.name, kernel.stat().st_size, "bytes") - # What a runner with no card runs: the C++ kernels and every layer above - # them. The GPU tests skip themselves on ``torch.cuda.is_available()``, - # and the ones that reach for Triton's CPU interpreter carry the - # ``interpreted`` marker, which ``addopts`` deselects. + # What a runner with no card runs: the C++ kernels, the GPU kernels + # compiled for the host, and every layer above them. The GPU tests skip + # themselves on ``torch.cuda.is_available()``. - name: Run the tests run: pytest tests/ -n auto diff --git a/.github/workflows/wheels.yml b/.github/workflows/wheels.yml index 6ca1ae4f..50bda729 100644 --- a/.github/workflows/wheels.yml +++ b/.github/workflows/wheels.yml @@ -12,6 +12,8 @@ on: - setup.py - "src/blochsim/*.cpp" - "src/blochsim/*.hpp" + - "src/blochsim/*.cu" + - "src/blochsim/*.cu.in" - scripts/check_wheel.py - .github/workflows/wheels.yml workflow_dispatch: @@ -52,7 +54,9 @@ jobs: wheels: name: ${{ matrix.label }} runs-on: ${{ matrix.os }} - timeout-minutes: 60 + # The x86-64 manylinux wheel compiles the GPU kernels for every listed + # architecture. + timeout-minutes: 150 strategy: fail-fast: false matrix: diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 2df1e55e..c3dbf6a2 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -12,6 +12,8 @@ repos: hooks: - id: check-added-large-files args: [--maxkb=512] + # The EPG kernels for the card, in every mode. + exclude: ^src/blochsim/_epg_kernels\.hpp$ - id: check-case-conflict - id: check-merge-conflict - id: check-symlinks diff --git a/CHANGELOG.md b/CHANGELOG.md index 5f2a5c2f..210d8a77 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,15 @@ - **The package is `blochsim`.** The distribution, the import name and the repository are `blochsim`; `import torchsim` becomes `import blochsim`, with the same modules and names beneath it. +- **The GPU kernels are CUDA compiled ahead of time, and Triton is not used.** + The EPG, pooled and PERK kernels are C++ (`_epg_kernels.hpp`, + `_pools_kernels.hpp`, `_perk_kernels.hpp`) compiled by `nvcc` into + `blochsim._gpu`, which links the CUDA runtime statically; the x86-64 + manylinux wheel carries it, and a source build compiles it wherever CMake + finds `nvcc`. No kernel is compiled at the first call. The same kernels are + compiled for the host as `blochsim._gpu_host`, which the suite holds to the + C++ kernels; the `interpreted` marker is gone. A launch is at most 1024 + threads, so an EPG run of more than 1024 state orders on a card is refused. ### Added diff --git a/CLAUDE.md b/CLAUDE.md index 6b6972f0..c8014256 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -16,12 +16,14 @@ broken checkout (Windows without `core.symlinks`), not two documents. | `src/blochsim/sequence/` | The description an acquisition is assembled from — events, operators, builders — and the dispatch that turns one into a kernel launch. | | `src/blochsim/model/` | What a signal model *is*: the physics, the simulator that orders its events, and the binding that resolves a protocol's structure once and rebinds its values per call. | | `src/blochsim/_epg_cpu.cpp`, `_perk_cpu.cpp` | The CPU kernels. Every path the GPU has exists here too, and the two agree to float32 round-off. | -| `src/blochsim/sequence/_epg_triton.py`, `estimators/_perk_triton.py` | The GPU kernels. | +| `src/blochsim/_epg_kernels.hpp`, `_pools_kernels.hpp`, `_perk_kernels.hpp` | The GPU kernels, written over the tiles of `_tile.hpp`. `_kernels.hpp` is the table a launch reads: each kernel's parameters, and which of them size the block. | +| `src/blochsim/_gpu.cu`, `_gpu_kernel.cu.in`, `_gpu_host.cpp`, `_gpu_launch.py` | The same kernels compiled ahead of time for the card (`_gpu`), and for the host (`_gpu_host`), which is how the suite checks them without one; and the launcher both answer to. | +| `src/blochsim/sequence/_epg_gpu.py`, `_pools_gpu.py`, `estimators/_perk_gpu.py` | What a GPU launch is given: the tiling and the arguments. | | `src/blochsim/simulators/`, `estimators/`, `recon/`, `optim/` | The sequences that ship, and what is built on top of them. | | `tests/`, `examples/`, `docs/` | Mirrored by subpackage, executed by the gallery, built by Sphinx. | The shared parameter ABI — read by the Python dispatch, the C++ extension and -the Triton kernels alike — is `src/blochsim/sequence/_parameters.py`. A +the GPU kernels alike — is `src/blochsim/sequence/_parameters.py`. A parameter added there is added in all three places or in none. ## Commands @@ -30,7 +32,6 @@ parameter added there is added in all three places or in none. pip install -e ".[dev]" # the whole toolchain, and it compiles the kernels pytest tests/ # the suite, with coverage pytest tests/ -n auto # across cores -pytest tests/ -m interpreted # the Triton paths, through Triton's CPU interpreter pre-commit install # once per clone, or no hook runs on commit pre-commit run --all-files # exactly what CI's Lint job runs bash scripts/build_docs.sh # HTML into docs/build/html, examples executed @@ -57,40 +58,19 @@ pip install -e ".[dev]" ; echo "exit: $?" python -c "import blochsim._epg_cpu as k; print(k.__file__)" ``` -**Kernel compiles dominate a cold GPU run**, not the arithmetic. A suite that -takes minutes on a card is mostly Triton compiling one specialization per -feature combination it meets. Run the whole suite at natural boundaries, not -after every edit. - -**Triton's cache keys on the source of the kernel file, not on its meaning.** -Reformatting `_epg_triton.py`, or deleting a dead local from it, invalidates -every specialization exactly as a new `tl.constexpr` would: the next full run -on a card goes from minutes to the better part of an hour, almost all of it one -MLIR pass. Know that before letting a formatter touch that file, and say which -number you are quoting. - -**Most of what the cache holds is never loaded.** Each specialization is -written out as its `.ttir`, `.ttgir`, `.llir` and `.ptx` beside the `.cubin` -and the metadata a launch actually reads, and on a kernel this size the -intermediate IR is the bulk of the bytes. `TRITON_STORE_BINARY_ONLY=1` writes -the binary and its metadata alone; every compilation stage still runs, so -neither the generated code nor the set of specializations changes. -`TRITON_DISABLE_LINE_INFO=1` drops the source locations embedded in both, -which costs line attribution under `ncu`. Set `TRITON_CACHE_DIR` to somewhere -you keep and the compile is paid once per edit of the kernel file rather than -once per checkout. - -Neither knob stops the `.source` file, which carries a location per operation -and is most of an entry on a kernel this size; nothing reads it at launch. -Nothing evicts either, and the cache keys on the text of the kernel file, so -every edit orphans every entry made before it. -`scripts/prune_triton_cache.py` reports both and, with `--apply`, drops them: -the IR is free to drop, the stale entries recompile if you ask for them again. - -**The `interpreted` marker is deselected by default.** Those tests run a Triton -kernel through Triton's CPU interpreter — no GPU, no compile, about a minute -each. `TRITON_INTERPRET=1` does the same thing by hand for a script. It is how -the GPU plumbing is verified on a machine with no card. +**The GPU kernels are compiled where a CUDA compiler is found.** CMake looks +for `nvcc`, and `BLOCHSIM_CUDA=ON` or `OFF` (`--config-settings=cmake.define.BLOCHSIM_CUDA=ON`) +overrides what it finds; `CMAKE_CUDA_ARCHITECTURES` picks the cards. A kernel +compiles in a file of its own, twice -- bounded to 256 threads and to 1024 -- +so a build is minutes of `nvcc` spread over as many cores as Ninja is given. +The same kernels compiled for the host (`_gpu_host`) run one program at a +time over host tensors; `tests/sequence/test_host_kernels.py` and +`test_many_pools_host.py` hold them to the C++ kernels, which is how the GPU +path is verified on a machine with no card. + +**A block is at most 1024 threads.** A tile is an element per thread, so an +EPG launch with more than 1024 state orders, or a pooled one whose orders by +pools pass 1024, is refused with the kernel's name rather than launched. **`--cov` is on by default** through `addopts`, so a bare `pytest` writes `coverage.xml`. It is ignored, not tracked. @@ -129,7 +109,7 @@ written that way and each states its invariant in its module docstring. A test that only compares BlochSim to BlochSim proves the two agree, which was never in doubt. -Whatever you change in one kernel, change in the other. The C++ and Triton +Whatever you change in one kernel, change in the other. The C++ and GPU implementations are held to each other to float32 round-off, and a path that exists on one side and not the other is a bug in whichever side is missing it. @@ -207,12 +187,14 @@ beside the list the version switcher reads. ## Packaging The build is `scikit-build-core` driving `CMakeLists.txt`; `setup.py` is a shim -for tools that still shell out to it and configures nothing. Both kernels are +for tools that still shell out to it and configures nothing. The kernels are plain CPython extensions against the **stable ABI** from 3.10 on: they call no PyTorch API and link no PyTorch library, which is why one `cp310-abi3` wheel per platform serves every supported interpreter and why that wheel is a couple of megabytes rather than the size of libtorch. Keep it that way — a `#include -` in either `.cpp` ends all of that. +` in any of them ends all of that. The GPU module links the CUDA +runtime statically, so a machine needs the driver and nothing else, and the +x86-64 manylinux wheel is the one built with it. Wheels are built by cibuildwheel and published to PyPI by trusted publishing on a `v*.*.*` tag. `scripts/check_wheel.py` loads each compiled kernel by path, diff --git a/CMakeLists.txt b/CMakeLists.txt index 5d80ec2c..57eff498 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -2,7 +2,7 @@ cmake_minimum_required(VERSION 3.26) project(blochsim LANGUAGES CXX) -# Both kernels are plain CPython extensions against the stable ABI: they call +# The kernels are plain CPython extensions against the stable ABI: they call # no PyTorch API and link no PyTorch library, which is why one wheel per # platform serves every supported interpreter and why that wheel is a couple of # megabytes rather than the size of libtorch. @@ -56,3 +56,63 @@ endfunction() blochsim_add_kernel(_epg_cpu src/blochsim/_epg_cpu.cpp) blochsim_add_kernel(_perk_cpu src/blochsim/_perk_cpu.cpp) + +# The GPU kernels compiled for the host, one program at a time. Nothing in the +# package dispatches to them; they are how the suite checks the GPU kernels on +# a machine with no card. +option(BLOCHSIM_HOST_KERNELS "Compile the GPU kernels for the host, for the tests" ON) +if(BLOCHSIM_HOST_KERNELS) + blochsim_add_kernel(_gpu_host src/blochsim/_gpu_host.cpp) +endif() + +# The GPU kernels compiled ahead of time for the card, wherever a CUDA +# compiler is found unless BLOCHSIM_CUDA says otherwise. +include(CheckLanguage) +check_language(CUDA) +if(CMAKE_CUDA_COMPILER) + set(_blochsim_cuda_default ON) +else() + set(_blochsim_cuda_default OFF) +endif() +option(BLOCHSIM_CUDA "Compile the GPU kernels for the card" ${_blochsim_cuda_default}) + +if(BLOCHSIM_CUDA) + # Real code for each generation and PTX for the newest, which the driver + # compiles for a card newer than any listed. + if(NOT DEFINED CMAKE_CUDA_ARCHITECTURES) + set(CMAKE_CUDA_ARCHITECTURES 75-real 80-real 86-real 89-real 90) + endif() + enable_language(CUDA) + find_package(CUDAToolkit REQUIRED) + set(CMAKE_CUDA_STANDARD 17) + set(CMAKE_CUDA_STANDARD_REQUIRED ON) + + # One file per kernel, so they compile in parallel. + set(_blochsim_kernels + _three_pool_table_jvp_kernel _three_pool_table_kernel + _epg_vjp_kernel _epg_vjp_jvp_kernel _epg_real_vjp_jvp_kernel + _epg_real_vjp_kernel _epg_real_kernel _epg_real_jvp_kernel + _epg_kernel _epg_jvp_kernel + _pooled_kernel _pooled_adjoint_kernel + _regress_kernel _regress_vjp_kernel + ) + set(_blochsim_gpu_sources src/blochsim/_gpu.cu) + foreach(NAME IN LISTS _blochsim_kernels) + set(_source "${CMAKE_CURRENT_BINARY_DIR}/kernels/${NAME}.cu") + configure_file(src/blochsim/_gpu_kernel.cu.in "${_source}" @ONLY) + list(APPEND _blochsim_gpu_sources "${_source}") + endforeach() + + python_add_library(_gpu MODULE + USE_SABI ${BLOCHSIM_ABI3_VERSION} + WITH_SOABI + ${_blochsim_gpu_sources} + ) + target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim") + # The runtime is linked in, so a machine needs the driver and nothing else. + target_link_libraries(_gpu PRIVATE CUDA::cudart_static) + # Warning 221 is a double constant that rounds to zero in float, which the + # kernels write deliberately as the floor of a clamp. + target_compile_options(_gpu PRIVATE $<$:-diag-suppress=221>) + install(TARGETS _gpu DESTINATION blochsim) +endif() diff --git a/docs/developer_guide.md b/docs/developer_guide.md index ae089677..0914dea3 100644 --- a/docs/developer_guide.md +++ b/docs/developer_guide.md @@ -52,10 +52,11 @@ the compiler is on the path. **Git**, and a fork of https://github.com/pulserver/blochsim if you intend to open a pull request. -**An NVIDIA card, if you want to touch the GPU kernels.** They are Triton, -which arrives with the Linux CUDA wheels of PyTorch. You can develop and test -most of the Triton path without a card -- see {ref}`dev-tests` -- but only a -card runs it for real. +**The CUDA toolkit and an NVIDIA card, if you want to touch the GPU kernels.** +They are C++ compiled ahead of time by `nvcc`, which CMake uses wherever it +finds it. The same kernels are also compiled for the host, so you can develop +and test them without a card -- see {ref}`dev-tests` -- but only a card runs +them for real. ## Installing for development @@ -99,10 +100,11 @@ python -c "import blochsim._epg_cpu as k; print(k.__file__)" `src/blochsim/sequence/` : The description an acquisition is assembled from -- events, operators, - builders -- and the dispatch that turns one into a kernel launch. The - Triton kernels are `_epg_triton.py`; the shared parameter ABI, which the - Python dispatch, the C++ extension and the Triton kernels all read, is - `_parameters.py`. + builders -- and the dispatch that turns one into a kernel launch: + `_epg_gpu.py` and `_pools_gpu.py` for the GPU kernels, whose source is + `src/blochsim/_epg_kernels.hpp` and `_pools_kernels.hpp`. The shared + parameter ABI, which the Python dispatch, the C++ extension and the GPU + kernels all read, is `_parameters.py`. `src/blochsim/model/` : What a signal model is: the physics, the simulator that orders its events, @@ -249,23 +251,11 @@ diffusion, flow, spoiling, the two-pool and three-pool longitudinal steps -- while `sequence/`, `model/`, `estimators/`, `recon/` and `optim/` cover the layers above. -Two things to know before you time a run: - -**The `interpreted` marker is deselected by default.** Those tests run a -Triton kernel through Triton's CPU interpreter -- no GPU, no compile, and -about a minute each. That is how the GPU plumbing is verified on a machine -with no card: - -```sh -pytest tests/ -m interpreted -TRITON_INTERPRET=1 python your_script.py # the same trick, by hand -``` - -**Kernel compiles dominate a cold GPU run**, not the arithmetic. A suite that -takes minutes on a card is mostly Triton compiling one specialization per -feature combination it meets; the second run of the same suite is a different -number entirely. Run the whole suite at natural boundaries rather than after -every edit. +**The GPU kernels run without a card.** The install compiles them for the +host as well, one program at a time over host tensors, and +`tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` hold that +build to the C++ kernels. That is how the GPU kernels are verified on a +machine with no card; the tests that need one skip themselves. When you change physics, add the test that pins it against something outside BlochSim: a closed form, a published figure, or an isochromat summation you diff --git a/docs/explanation_figures.py b/docs/explanation_figures.py index 9d7da181..4a62e54d 100644 --- a/docs/explanation_figures.py +++ b/docs/explanation_figures.py @@ -674,7 +674,7 @@ def pipeline(): axis.text( boundary + 0.15, 1.35, - "C++ or Triton, once per run", + "C++ or CUDA, once per run", color=LONGITUDINAL, fontsize=11, ) diff --git a/docs/explanations/implementation.md b/docs/explanations/implementation.md index a3b68025..b378d2fa 100644 --- a/docs/explanations/implementation.md +++ b/docs/explanations/implementation.md @@ -7,7 +7,7 @@ Offline it builds a description from `layout()`; scanner-driven use starts from an incoming description and applies the simulator's handlers. - The resulting event stream is packed once and executed by fused CPU or - Triton kernels, one program per voxel. + CUDA kernels, one program per voxel. - Tissue Jacobians use forward mode; sequence optimization uses reverse mode; execution/offload policy is shared by every simulator. ``` @@ -117,10 +117,9 @@ is the sequence; what it is parallel in is the voxels. There are two implementations of that program and they compute the same thing. The CPU one is a threaded C++ extension, with a lane-vectorized path that runs -eight trains at once where the arithmetic allows it. The GPU one is written in -Triton and compiles a specialization per feature combination it meets -- which -is why the *first* call on a card can take tens of seconds while the second -takes milliseconds. Every mode exists on both sides: forward, forward-mode, +eight trains at once where the arithmetic allows it. The GPU one is CUDA, +compiled ahead of time, with a block of threads per program. Every mode exists +on both sides: forward, forward-mode, adjoint, forward-over-reverse, the real-subspace specialization, and the pool models. diff --git a/docs/misc/related.md b/docs/misc/related.md index 5aacb04c..ff717dba 100644 --- a/docs/misc/related.md +++ b/docs/misc/related.md @@ -46,7 +46,7 @@ and a signal-model simulator has no reason to carry a coil. | [EpyG](https://github.com/brennerd11/EpyG) | EPG | Python | CPU | -- | Python | | [mri-sim-py](https://github.com/utcsilab/mri-sim-py.epg) | EPG | Python, PyTorch | CPU, CUDA | Automatic, reverse mode | Python | | [snapMRF](https://github.com/dongwang881107/snapMRF) | EPG, with matching | CUDA C | CUDA | -- | Command line | -| **BlochSim** | EPG, and closed forms | Python, PyTorch, C++ and Triton kernels | CPU threads, CUDA, several cards | Automatic: forward, reverse, and forward over reverse | Python, or a description | +| **BlochSim** | EPG, and closed forms | Python, PyTorch, C++ and CUDA kernels | CPU threads, CUDA, several cards | Automatic: forward, reverse, and forward over reverse | Python, or a description | ## Where each one is the better tool diff --git a/docs/user_guide.md b/docs/user_guide.md index 463999ce..ea4da7b7 100644 --- a/docs/user_guide.md +++ b/docs/user_guide.md @@ -109,9 +109,9 @@ pick the one your driver supports, and install from that index -- pip install torch --index-url https://download.pytorch.org/whl/cu128 ``` -The GPU kernels are written in Triton, which comes with the Linux CUDA -wheels; you do not install it separately and you do not need the CUDA -toolkit, only a driver new enough for the build you picked. +The GPU kernels are compiled into the Linux x86-64 wheel of BlochSim; you do +not need the CUDA toolkit, only a driver new enough for the build you +picked. Built from source, they are compiled wherever CMake finds `nvcc`. ::: :::{tab-item} Apple silicon @@ -124,8 +124,8 @@ pip install torch ``` The simulation runs on the CPU kernels. There is no Metal path: the -fused state machine exists as C++ and as Triton, and Triton has no -Apple backend. +fused state machine exists as C++ for the CPU and as CUDA for NVIDIA +cards. ::: :::: @@ -164,11 +164,11 @@ from blochsim.simulators import FSESimulator acquisition = FSESimulator( ESP=5.0, TR=3000.0, - T1=torch.tensor([830.0, 1330.0, 4000.0]), # ms + T1=torch.tensor([830.0, 1330.0, 4000.0]), # ms T2=torch.tensor([80.0, 110.0, 2000.0]), ) signal = acquisition.simulate(flip=torch.full((48,), 60.0)) -print(signal.shape) # torch.Size([3, 48]) +print(signal.shape) # torch.Size([3, 48]) ``` On a card, hand it tissue that already lives there and the whole run follows: @@ -178,9 +178,8 @@ acquisition = acquisition.to("cuda") signal = acquisition.simulate(flip=torch.full((48,), 60.0, device="cuda")) ``` -The first GPU call pays for compiling the Triton kernel it needs -- tens of -seconds, once per kernel per machine, cached afterwards. A first call that -seems to hang is almost always that compile. +The kernels are compiled ahead of time, so the first GPU call costs what +any other does, apart from PyTorch initialising CUDA. ## Your first simulation diff --git a/pyproject.toml b/pyproject.toml index 10c0508c..60d8d342 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -125,7 +125,7 @@ wheel.packages = ["src/blochsim"] # against the CPython stable ABI from 3.10 on, so the tag is cp310-abi3. wheel.py-api = "cp310" # What ships is the compiled kernel, not the source it was compiled from. -wheel.exclude = ["*.cpp", "*.hpp"] +wheel.exclude = ["*.cpp", "*.hpp", "*.cu", "*.cu.in"] metadata.version.provider = "scikit_build_core.metadata.setuptools_scm" sdist.include = ["src/blochsim/_version.py"] @@ -143,12 +143,27 @@ test-command = "python {project}/scripts/check_wheel.py" # publishes no musllinux wheel for that install to resolve. The workflow loads # the musl kernels in an Alpine container instead, where nothing resolves. test-skip = "*musllinux*" +# The GPU kernels compiled for the host are for the suite, not for a user. +config-settings = {"cmake.define.BLOCHSIM_HOST_KERNELS" = "OFF"} [tool.cibuildwheel.linux] # scikit-build-core installs cmake and ninja as build-time wheels inside the # manylinux container; the image's compiler is what builds the kernels. archs = ["auto64"] +# The GPU kernels, on the one platform PyTorch publishes CUDA builds for that +# a manylinux image can compile: the toolkit's compiler and its static +# runtime, nothing the wheel carries beyond the module itself. +[[tool.cibuildwheel.overrides]] +select = "*-manylinux_x86_64" +before-all = [ + "dnf install -y dnf-plugins-core", + "dnf config-manager --add-repo https://developer.download.nvidia.com/compute/cuda/repos/rhel8/x86_64/cuda-rhel8.repo", + "dnf install -y cuda-nvcc-12-8 cuda-cudart-devel-12-8", +] +environment = {PATH = "/usr/local/cuda-12.8/bin:$PATH", CUDACXX = "/usr/local/cuda-12.8/bin/nvcc"} +config-settings = {"cmake.define.BLOCHSIM_HOST_KERNELS" = "OFF", "cmake.define.BLOCHSIM_CUDA" = "ON"} + [tool.cibuildwheel.macos] # An Apple silicon runner builds its own wheel and cross-compiles the Intel # one beside it; cibuildwheel skips the test for the architecture the runner @@ -170,10 +185,6 @@ addopts = [ "--cov=blochsim", "--cov-report=term-missing", "--cov-report=xml", - "-m", "not interpreted" -] -markers = [ - "interpreted: runs a kernel in Triton's CPU interpreter; a minute each" ] # Formatting and linting are both ruff: ``ruff format`` is the formatter and the @@ -225,9 +236,6 @@ ignore = [ "examples/**" = ["ANN", "D", "E402", "I001", "B018"] "docs/**" = ["ANN", "D1"] "scripts/**" = ["ANN", "D1"] -# A Triton kernel's parameters are compile-time constants and tensors the -# compiler types itself; an annotation on one says nothing Python will read. -"src/blochsim/**/*_triton.py" = ["ANN"] [tool.ruff.lint.isort] known-first-party = ["blochsim"] diff --git a/scripts/check_wheel.py b/scripts/check_wheel.py index 48ef33ce..4a7926c1 100644 --- a/scripts/check_wheel.py +++ b/scripts/check_wheel.py @@ -18,6 +18,9 @@ #: A kernel calls no PyTorch API and links no PyTorch library, so anything #: approaching this size means something was bundled that should not have been. LARGEST_REASONABLE_BYTES = 16 * 1024 * 1024 +#: The GPU kernels carry machine code for each architecture they were compiled +#: for, and the CUDA runtime linked in. +LARGEST_REASONABLE_GPU_BYTES = 96 * 1024 * 1024 def kernel_directory() -> pathlib.Path: @@ -35,6 +38,13 @@ def main() -> int: kernels = [path for path in kernels if path.suffix in {".so", ".pyd", ".dylib"}] if len(kernels) != 2: raise SystemExit(f"expected two kernels in {root}, found {kernels}") + # Present where the wheel was built with a CUDA compiler. Loading it needs + # no card: nothing reaches the driver until a launch. + kernels += [ + path + for path in sorted(root.glob("_gpu.*")) + if path.suffix in {".so", ".pyd", ".dylib"} + ] for path in kernels: name = path.name.split(".")[0] @@ -46,7 +56,10 @@ def main() -> int: size = path.stat().st_size print(f"{path.name}: {size} bytes, {len(dir(module))} attributes") - if size > LARGEST_REASONABLE_BYTES: + largest = ( + LARGEST_REASONABLE_GPU_BYTES if name == "_gpu" else LARGEST_REASONABLE_BYTES + ) + if size > largest: raise SystemExit(f"{path.name} is {size} bytes; something was bundled") print(f"ok on {sys.implementation.name} {'.'.join(map(str, sys.version_info[:3]))}") diff --git a/scripts/prune_triton_cache.py b/scripts/prune_triton_cache.py deleted file mode 100644 index f0a1162f..00000000 --- a/scripts/prune_triton_cache.py +++ /dev/null @@ -1,142 +0,0 @@ -"""Reclaim the disk a Triton cache grows into, without losing a compilation. - -Two things make the cache grow without bound. Each entry keeps the MLIR the -kernel was compiled from -- a `.source` file carrying a source location per -operation, which on a kernel the size of the EPG state machine is most of the -entry and is never read again: a launch reads the cubin and its metadata. -And nothing evicts: Triton keys an entry on the text of the file the kernel -was written in, so every edit to ``_epg_triton.py`` orphans every entry made -before it, and the orphans stay. - -So there are two prunes here. Dropping the IR is free -- the entries still -answer, and nothing recompiles. Dropping entries by age is not: an entry still -in use recompiles the next time it is asked for, which is the ordinary cold -cost of that kernel. - -Nothing is deleted without ``--apply``, so the first two commands below only -say what the third would do. - - python scripts/prune_triton_cache.py - python scripts/prune_triton_cache.py --drop-ir --older-than 30 - python scripts/prune_triton_cache.py --drop-ir --older-than 30 --apply - -``--cache`` names the directory; without it, ``TRITON_CACHE_DIR`` and then -``~/.triton/cache``. -""" - -from __future__ import annotations - -import argparse -import os -import pathlib -import shutil -import time - -#: What a launch reads. Everything else in an entry is there for a human. -LAUNCHED_FROM = {".cubin", ".hsaco", ".json", ".so"} - - -def cache_directory(named: str | None) -> pathlib.Path: - """The cache to work on, from the argument, the environment, or the default.""" - if named: - return pathlib.Path(named).expanduser() - from_environment = os.environ.get("TRITON_CACHE_DIR") - if from_environment: - return pathlib.Path(from_environment).expanduser() - return pathlib.Path.home() / ".triton" / "cache" - - -def entries(cache: pathlib.Path) -> list[pathlib.Path]: - """The per-kernel directories, each named by its compilation's hash.""" - return sorted(path for path in cache.iterdir() if path.is_dir()) - - -def bytes_under(path: pathlib.Path) -> int: - """Total size of everything below ``path``.""" - return sum(file.stat().st_size for file in path.rglob("*") if file.is_file()) - - -def intermediate_ir(entry: pathlib.Path) -> list[pathlib.Path]: - """The files in ``entry`` that no launch reads.""" - return [ - file - for file in entry.iterdir() - if file.is_file() and file.suffix not in LAUNCHED_FROM - ] - - -def touched(entry: pathlib.Path) -> float: - """When anything in ``entry`` was last read or written.""" - return max( - (file.stat().st_atime for file in entry.rglob("*") if file.is_file()), - default=entry.stat().st_atime, - ) - - -def main() -> None: - """Report what the cache holds, and drop what was asked for.""" - cli = argparse.ArgumentParser(description=__doc__.splitlines()[0]) - cli.add_argument("--cache", default="", help="the cache directory") - cli.add_argument( - "--drop-ir", - action="store_true", - help="delete the MLIR no launch reads", - ) - cli.add_argument( - "--older-than", - type=float, - default=0.0, - help="entries untouched for this many days, which recompile if asked for again", - ) - cli.add_argument( - "--apply", - action="store_true", - help="delete what the others select; without it nothing is touched", - ) - arguments = cli.parse_args() - - cache = cache_directory(arguments.cache) - if not cache.is_dir(): - print(f"no cache at {cache}") - return - - held = entries(cache) - total = sum(bytes_under(entry) for entry in held) - print(f"{cache}: {len(held)} entries, {total / 2**30:.2f} GiB") - - stale = [] - if arguments.older_than > 0: - cutoff = time.time() - arguments.older_than * 86400 - stale = [entry for entry in held if touched(entry) < cutoff] - - ir = {entry: intermediate_ir(entry) for entry in held if entry not in set(stale)} - ir_bytes = sum(file.stat().st_size for files in ir.values() for file in files) - stale_bytes = sum(bytes_under(entry) for entry in stale) - - print(f" IR no launch reads: {ir_bytes / 2**30:.2f} GiB") - if arguments.older_than > 0: - print( - f" entries untouched for {arguments.older_than:g} days: " - f"{len(stale)}, {stale_bytes / 2**30:.2f} GiB" - ) - - if not arguments.apply: - print("nothing deleted; pass --apply to act on the above") - return - - reclaimed = 0 - for entry in stale: - reclaimed += bytes_under(entry) - shutil.rmtree(entry) - if arguments.drop_ir: - for files in ir.values(): - for file in files: - reclaimed += file.stat().st_size - file.unlink() - - remaining = sum(bytes_under(entry) for entry in entries(cache)) - print(f"reclaimed {reclaimed / 2**30:.2f} GiB, {remaining / 2**30:.2f} GiB left") - - -if __name__ == "__main__": - main() diff --git a/src/blochsim/_epg_kernels.hpp b/src/blochsim/_epg_kernels.hpp new file mode 100644 index 00000000..395558de --- /dev/null +++ b/src/blochsim/_epg_kernels.hpp @@ -0,0 +1,12222 @@ +// The EPG kernels for the card: forward, forward-mode, adjoint and +// forward-over-reverse, complex and real, and the three-pool tables. Written +// over the tiles of _tile.hpp and included by _kernels.hpp inside ``epg``. + +// Diffusion damping and its directional derivative, per state order. +template +BSK_HD auto _damping_jvp(const T0& rate, const T1& rate_tangent, const T2& dt, const T3& dt_tangent, const T4& order) { + auto b_factor = (rate * dt); + auto b_tangent = ((rate_tangent * dt) + (rate * dt_tangent)); + auto squared = (order * order); + auto transverse_weight = ((squared + order) + 0.3333333333333333f); + auto damp_z = bsk::exp(((-b_factor) * squared)); + auto damp_t = bsk::exp(((-b_factor) * transverse_weight)); + return bsk::make_tup(damp_z, (damp_z * ((-b_tangent) * squared)), damp_t, (damp_t * ((-b_tangent) * transverse_weight))); +} + +// Two dual complex numbers added. +template +BSK_HD auto _dual_add(const T0& x, const T1& y) { + return bsk::make_tup((bsk::get<0>(x) + bsk::get<0>(y)), (bsk::get<1>(x) + bsk::get<1>(y)), (bsk::get<2>(x) + bsk::get<2>(y)), (bsk::get<3>(x) + bsk::get<3>(y))); +} + +template +BSK_HD auto _complex_mul(const T0& a_real, const T1& a_imag, const T2& b_real, const T3& b_imag) { + return bsk::make_tup(((a_real * b_real) - (a_imag * b_imag)), ((a_real * b_imag) + (a_imag * b_real))); +} + +// Product of two dual complex numbers. +template +BSK_HD auto _dual_mul(const T0& a_vr, const T1& a_vi, const T2& a_tr, const T3& a_ti, const T4& b_vr, const T5& b_vi, const T6& b_tr, const T7& b_ti) { + auto t0_ = _complex_mul(a_vr, a_vi, b_vr, b_vi); + auto value_real = bsk::get<0>(t0_); + auto value_imag = bsk::get<1>(t0_); + auto t1_ = _complex_mul(a_tr, a_ti, b_vr, b_vi); + auto left_real = bsk::get<0>(t1_); + auto left_imag = bsk::get<1>(t1_); + auto t2_ = _complex_mul(a_vr, a_vi, b_tr, b_ti); + auto right_real = bsk::get<0>(t2_); + auto right_imag = bsk::get<1>(t2_); + return bsk::make_tup(value_real, value_imag, (left_real + right_real), (left_imag + right_imag)); +} + +// ``conj(entry * spin)`` against a cotangent, entry and spin both dual. +// +// One row of a real mixing operator carried through the per-order turn, which +// is what a longitudinal cotangent walks back through. +template +BSK_HD auto _dual_back(const T0& entry, const T1& d_entry, const T2& spin_vr, const T3& spin_vi, const T4& spin_tr, const T5& spin_ti, const T6& br, const T7& bi, const T8& tr, const T9& ti) { + return _dual_mul((entry * spin_vr), (-(entry * spin_vi)), ((d_entry * spin_vr) + (entry * spin_tr)), (-((d_entry * spin_vi) + (entry * spin_ti))), br, bi, tr, ti); +} + +// A dual complex number's conjugate, both halves. +template +BSK_HD auto _dual_conj(const T0& z) { + return bsk::make_tup(bsk::get<0>(z), (-bsk::get<1>(z)), bsk::get<2>(z), (-bsk::get<3>(z))); +} + +// ``exp(i * angle)`` for a real dual angle. +template +BSK_HD auto _dual_polar(const T0& angle_value, const T1& angle_tangent) { + auto cosine = bsk::cos(angle_value); + auto sine = bsk::sin(angle_value); + return bsk::make_tup(cosine, sine, ((-sine) * angle_tangent), (cosine * angle_tangent)); +} + +// Two dual complex numbers multiplied. +template +BSK_HD auto _dual_product(const T0& x, const T1& y) { + return _dual_mul(bsk::get<0>(x), bsk::get<1>(x), bsk::get<2>(x), bsk::get<3>(x), bsk::get<0>(y), bsk::get<1>(y), bsk::get<2>(y), bsk::get<3>(y)); +} + +// ``real_part(conj(a) * b)``, the contraction an adjoint asks for. +template +BSK_HD auto _dual_real_conj_mul(const T0& a_vr, const T1& a_vi, const T2& a_tr, const T3& a_ti, const T4& b_vr, const T5& b_vi, const T6& b_tr, const T7& b_ti) { + auto value = ((a_vr * b_vr) + (a_vi * b_vi)); + auto tangent = ((((a_tr * b_vr) + (a_ti * b_vi)) + (a_vr * b_tr)) + (a_vi * b_ti)); + return bsk::make_tup(value, tangent); +} + +// A real dual number times a complex one. +template +BSK_HD auto _dual_scale(const T0& scale_value, const T1& scale_tangent, const T2& vr, const T3& vi, const T4& tr, const T5& ti) { + return bsk::make_tup((scale_value * vr), (scale_value * vi), ((scale_tangent * vr) + (scale_value * tr)), ((scale_tangent * vi) + (scale_value * ti))); +} + +// One dual complex number less another. +template +BSK_HD auto _dual_subtract(const T0& x, const T1& y) { + return bsk::make_tup((bsk::get<0>(x) - bsk::get<0>(y)), (bsk::get<1>(x) - bsk::get<1>(y)), (bsk::get<2>(x) - bsk::get<2>(y)), (bsk::get<3>(x) - bsk::get<3>(y))); +} + +template +BSK_HD auto _dual_times_i(const T0& vr, const T1& vi, const T2& tr, const T3& ti) { + return bsk::make_tup((-vi), vr, (-ti), tr); +} + +// The rotation a pulse performs at this voxel, read rather than read off. +// +// A tabulated pair covers a shape's every pulse because a static array +// reaches the rotation through one complex scalar; this one is integrated per +// pulse per voxel, so there is nothing to interpolate and the read is four +// floats. The row runs per train and per event, as the flip does. +template +BSK_HD auto _dynamic_pair_at(const T0& pairs, const T1& pair_index, const T2& event_base, const T3& event, const T4& atom, const T5& atom_count, const T6& mask) { + auto row = bsk::cast(bsk::ld(((pair_index + event_base) + event))); + auto entry = (pairs + (((row * atom_count) + atom) * 4)); + return bsk::make_tup(bsk::ld((entry + 0), mask, 1.0f), bsk::ld((entry + 1), mask, 0.0f), bsk::ld((entry + 2), mask, 0.0f), bsk::ld((entry + 3), mask, 0.0f)); +} + +// The rotation and the direction along it, with the phase applied. +// +// Shaped exactly as :func:`_profiled_pair_dual` returns, so the spinor +// operator and its adjoint read one from the other without knowing which +// they were handed. A pass that follows no direction holds the rotation +// still, and ``directed`` keeps the read for one out of the kernel. +template +BSK_HD auto _dynamic_pair_dual_at(const T0& pairs, const T1& pair_direction, const T2& pair_index, const T3& event_base, const T4& event, const T5& atom, const T6& atom_count, const T7& mask, const T8& phi_value, const T9& phi_tangent, const T10& directed) { + bsk::tup | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>> moved{}; + auto held = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, mask); + auto still = (bsk::get<0>(held) * 0.0f); + moved = bsk::make_tup(still, still, still, still); + if (bsk::truth(directed)) { + moved = _dynamic_pair_at(pair_direction, pair_index, event_base, event, atom, atom_count, mask); + } + auto a = bsk::make_tup(bsk::get<0>(held), bsk::get<1>(held), bsk::get<0>(moved), bsk::get<1>(moved)); + auto b = bsk::make_tup(bsk::get<2>(held), bsk::get<3>(held), bsk::get<2>(moved), bsk::get<3>(moved)); + auto turn = _dual_polar((-phi_value), (-phi_tangent)); + return bsk::make_tup(a, _dual_product(b, turn)); +} + +// One event's entry of a buffer carrying a row per train. +// +// ``duration``, ``flip`` and ``phase`` are indexed by the train and the event +// and never by the atom, so where there is one train the address is the same +// for every lane of the program and the value can be read once rather than +// once per element of the tile. +template +BSK_HD auto _event_value(const T0& values, const T1& event_base, const T2& event, const T3& active_atom, const T4& single_train) { + using Ret = bsk::tile_t | 0, 2)>; + if (bsk::truth(single_train)) { + // Spread over the program's problems, which a jitted helper has to do + // for itself: both arms of the branch have to hand back the one shape. + return bsk::convert((bsk::ld((values + event)) + bsk::zeros_like(bsk::cast(event_base)))); + } + return bsk::convert(bsk::ld(((values + event_base) + event), active_atom, 0.0f)); +} + +// Phase each dephasing order turns through over one interval. +// +// ``rate`` already carries the sequence's gradient geometry, so it is the +// winding per unit order per second: a longitudinal state at order l turns +// through ``l * rate * dt``. The transverse states sit half an order further +// along the gradient, which is where the extra half turn comes from. Order +// zero is left alone while longitudinal, so the recovery term is unaffected. +template +BSK_HD auto _flow(const T0& rate, const T1& dt, const T2& order) { + auto turn = (rate * dt); + return bsk::make_tup(((-order) * turn), ((-(order + 0.5f)) * turn)); +} + +// The lineshape, its slope and its curvature, from the same cubic. +// +// The table covers the magnitude, so the slope changes sign with the offset +// and the curvature does not: an even function's second derivative is even. +template +BSK_HD auto _lineshape_at_curve(const T0& lineshape, const T1& offset_hz, const T2& bins, const T3& step) { + auto last = (bins - 1); + auto magnitude = bsk::truediv(bsk::abs(offset_hz), step); + auto scaled = bsk::minimum(magnitude, (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto base = (bsk::cast(lower) * 2); + auto near = bsk::ld((lineshape + base)); + auto near_slope = bsk::ld(((lineshape + base) + 1)); + auto far = bsk::ld(((lineshape + base) + 2)); + auto far_slope = bsk::ld(((lineshape + base) + 3)); + auto value = (((((((2.0f * u3) - (3.0f * u2)) + 1.0f) * near) + ((((u3 - (2.0f * u2)) + u) * step) * near_slope)) + (((-2.0f * u3) + (3.0f * u2)) * far)) + (((u3 - u2) * step) * far_slope)); + auto direction = bsk::where((offset_hz < 0.0f), -1.0f, 1.0f); + auto slope = (direction * (((bsk::truediv((((6.0f * u2) - (6.0f * u)) * near), step) + ((((3.0f * u2) - (4.0f * u)) + 1.0f) * near_slope)) + bsk::truediv((((-6.0f * u2) + (6.0f * u)) * far), step)) + (((3.0f * u2) - (2.0f * u)) * far_slope))); + auto curve = (((bsk::truediv((((12.0f * u) - 6.0f) * near), (step * step)) + bsk::truediv((((6.0f * u) - 4.0f) * near_slope), step)) + bsk::truediv((((-12.0f * u) + 6.0f) * far), (step * step))) + bsk::truediv((((6.0f * u) - 2.0f) * far_slope), step)); + auto beyond = (magnitude > last); + return bsk::make_tup(value, bsk::where(beyond, 0.0f, slope), bsk::where(beyond, 0.0f, curve)); +} + +// The lineshape and its derivative in the *signed* offset. +// +// The table covers the magnitude, so the slope changes sign with the offset; +// past the last knot the read is constant and the slope is zero. +template +BSK_HD auto _lineshape_at_slope(const T0& lineshape, const T1& offset_hz, const T2& bins, const T3& step) { + auto last = (bins - 1); + auto magnitude = bsk::truediv(bsk::abs(offset_hz), step); + auto scaled = bsk::minimum(magnitude, (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto base = (bsk::cast(lower) * 2); + auto near = bsk::ld((lineshape + base)); + auto near_slope = bsk::ld(((lineshape + base) + 1)); + auto far = bsk::ld(((lineshape + base) + 2)); + auto far_slope = bsk::ld(((lineshape + base) + 3)); + auto value = (((((((2.0f * u3) - (3.0f * u2)) + 1.0f) * near) + ((((u3 - (2.0f * u2)) + u) * step) * near_slope)) + (((-2.0f * u3) + (3.0f * u2)) * far)) + (((u3 - u2) * step) * far_slope)); + auto direction = bsk::where((offset_hz < 0.0f), -1.0f, 1.0f); + auto slope = (direction * (((bsk::truediv((((6.0f * u2) - (6.0f * u)) * near), step) + ((((3.0f * u2) - (4.0f * u)) + 1.0f) * near_slope)) + bsk::truediv((((-6.0f * u2) + (6.0f * u)) * far), step)) + (((3.0f * u2) - (2.0f * u)) * far_slope))); + return bsk::make_tup(value, bsk::where((magnitude > last), 0.0f, slope)); +} + +// The pair, its slope and its curvature in the flip angle. +// +// The second-order pass differentiates the read twice, and a Hermite segment +// is a cubic, so all three come from the same four knot values. Returned in +// threes per component: value, slope, curvature. +template +BSK_HD auto _profile_pair_curve(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + auto g00 = bsk::truediv(((6.0f * u2) - (6.0f * u)), step); + auto g10 = (((3.0f * u2) - (4.0f * u)) + 1.0f); + auto g01 = bsk::truediv(((6.0f * u) - (6.0f * u2)), step); + auto g11 = ((3.0f * u2) - (2.0f * u)); + auto c00 = bsk::truediv(((12.0f * u) - 6.0f), (step * step)); + auto c10 = bsk::truediv(((6.0f * u) - 4.0f), step); + auto c01 = bsk::truediv((6.0f - (12.0f * u)), (step * step)); + auto c11 = bsk::truediv(((6.0f * u) - 2.0f), step); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto near = [&](int c) { return bsk::ld(((profile + base) + c)); }; + auto near_slope = [&](int c) { return bsk::ld((((profile + base) + 4) + c)); }; + auto far = [&](int c) { return bsk::ld((((profile + base) + 8) + c)); }; + auto far_slope = [&](int c) { return bsk::ld((((profile + base) + 12) + c)); }; + auto value = [&](int c) { + return ((((h00 * near(c)) + (h10 * near_slope(c))) + (h01 * far(c))) + (h11 * far_slope(c))); + }; + auto slope = [&](int c) { + return ((((g00 * near(c)) + (g10 * near_slope(c))) + (g01 * far(c))) + (g11 * far_slope(c))); + }; + auto curve = [&](int c) { + return ((((c00 * near(c)) + (c10 * near_slope(c))) + (c01 * far(c))) + (c11 * far_slope(c))); + }; + return bsk::make_tup(value(0), slope(0), curve(0), value(1), slope(1), curve(1), + value(2), slope(2), curve(2), value(3), slope(3), curve(3)); +} + +// The pair a shaped pulse turns through, and its slope, as duals. +// +// The flip angle carries the tangent into the table, so the pair's tangent is +// the stored slope and the slope's own tangent is the segment's curvature. +// The RF phase turns the axis once the pair is out, and so reaches ``b``. +template +BSK_HD auto _profiled_pair_dual(const T0& profile, const T1& row, const T2& alpha_value, const T3& alpha_tangent, const T4& phi_value, const T5& phi_tangent, const T6& bins, const T7& step) { + auto read = _profile_pair_curve(profile, row, alpha_value, bins, step); + auto a = bsk::make_tup(bsk::get<0>(read), bsk::get<3>(read), (bsk::get<1>(read) * alpha_tangent), (bsk::get<4>(read) * alpha_tangent)); + auto slope_a = bsk::make_tup(bsk::get<1>(read), bsk::get<4>(read), (bsk::get<2>(read) * alpha_tangent), (bsk::get<5>(read) * alpha_tangent)); + auto b = bsk::make_tup(bsk::get<6>(read), bsk::get<9>(read), (bsk::get<7>(read) * alpha_tangent), (bsk::get<10>(read) * alpha_tangent)); + auto slope_b = bsk::make_tup(bsk::get<7>(read), bsk::get<10>(read), (bsk::get<8>(read) * alpha_tangent), (bsk::get<11>(read) * alpha_tangent)); + auto turn = _dual_polar((-phi_value), (-phi_tangent)); + return bsk::make_tup(a, _dual_product(b, turn), slope_a, _dual_product(slope_b, turn)); +} + +// One row of the rotation applied to the states, values only. +template +BSK_HD auto _dual_row(const T0& first, const T1& second, const T2& third, const T3& fp_r, const T4& fp_i, const T5& fm_r, const T6& fm_i, const T7& z_r, const T8& z_i) { + auto real = ((((((bsk::get<0>(first) * fp_r) - (bsk::get<1>(first) * fp_i)) + (bsk::get<0>(second) * fm_r)) - (bsk::get<1>(second) * fm_i)) + (bsk::get<0>(third) * z_r)) - (bsk::get<1>(third) * z_i)); + auto imag = ((((((bsk::get<0>(first) * fp_i) + (bsk::get<1>(first) * fp_r)) + (bsk::get<0>(second) * fm_i)) + (bsk::get<1>(second) * fm_r)) + (bsk::get<0>(third) * z_i)) + (bsk::get<1>(third) * z_r)); + return bsk::make_tup(real, imag); +} + +// The rotation's nine coefficients and their tangents. +// +// Every entry is a product of two factors drawn from the pair and its +// conjugate, so five products carry all nine: ``a^2``, ``b^2``, ``a b``, +// ``a conj(b)`` and the norm difference. +template +BSK_HD auto _spinor_coefficients(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& dar, const T5& dai, const T6& dbr, const T7& dbi) { + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto daa_r = (2.0f * ((ar * dar) - (ai * dai))); + auto daa_i = (2.0f * ((dar * ai) + (ar * dai))); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto dbb_r = (2.0f * ((br * dbr) - (bi * dbi))); + auto dbb_i = (2.0f * ((dbr * bi) + (br * dbi))); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto dab_r = ((((dar * br) + (ar * dbr)) - (dai * bi)) - (ai * dbi)); + auto dab_i = ((((dar * bi) + (ar * dbi)) + (dai * br)) + (ai * dbr)); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto dcross_r = ((((dar * br) + (ar * dbr)) + (dai * bi)) + (ai * dbi)); + auto dcross_i = ((((dar * bi) + (ar * dbi)) - (dai * br)) - (ai * dbr)); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + auto dt22 = (2.0f * ((((ar * dar) + (ai * dai)) - (br * dbr)) - (bi * dbi))); + return bsk::make_tup(bsk::make_tup(aa_r, (-aa_i), daa_r, (-daa_i)), bsk::make_tup((-bb_r), bb_i, (-dbb_r), dbb_i), bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i), (-2.0f * dab_r), (2.0f * dab_i)), bsk::make_tup((-bb_r), (-bb_i), (-dbb_r), (-dbb_i)), bsk::make_tup(aa_r, aa_i, daa_r, daa_i), bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i), (-2.0f * dab_r), (-2.0f * dab_i)), bsk::make_tup(cross_r, cross_i, dcross_r, dcross_i), bsk::make_tup(cross_r, (-cross_i), dcross_r, (-dcross_i)), bsk::make_tup(t22, (0.0f * t22), dt22, (0.0f * dt22))); +} + +// The same row built from the coefficients' tangents instead. +template +BSK_HD auto _tangent_row(const T0& first, const T1& second, const T2& third, const T3& fp_r, const T4& fp_i, const T5& fm_r, const T6& fm_i, const T7& z_r, const T8& z_i) { + auto real = ((((((bsk::get<2>(first) * fp_r) - (bsk::get<3>(first) * fp_i)) + (bsk::get<2>(second) * fm_r)) - (bsk::get<3>(second) * fm_i)) + (bsk::get<2>(third) * z_r)) - (bsk::get<3>(third) * z_i)); + auto imag = ((((((bsk::get<2>(first) * fp_i) + (bsk::get<3>(first) * fp_r)) + (bsk::get<2>(second) * fm_i)) + (bsk::get<3>(second) * fm_r)) + (bsk::get<2>(third) * z_i)) + (bsk::get<3>(third) * z_r)); + return bsk::make_tup(real, imag); +} + +// The spinor rotation carrying a forward-mode tangent. +// +// Both the states and the pair naming the rotation move, so the tangent is +// ``T dx + dT x``. +template +BSK_HD auto _rotate_spinor_dual(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& dar, const T5& dai, const T6& dbr, const T7& dbi, const T8& fp_r, const T9& fp_i, const T10& fm_r, const T11& fm_i, const T12& z_r, const T13& z_i, const T14& dfp_r, const T15& dfp_i, const T16& dfm_r, const T17& dfm_i, const T18& dz_r, const T19& dz_i) { + auto t0_ = _spinor_coefficients(ar, ai, br, bi, dar, dai, dbr, dbi); + auto t00 = bsk::get<0>(t0_); + auto t01 = bsk::get<1>(t0_); + auto t02 = bsk::get<2>(t0_); + auto t10 = bsk::get<3>(t0_); + auto t11 = bsk::get<4>(t0_); + auto t12 = bsk::get<5>(t0_); + auto t20 = bsk::get<6>(t0_); + auto t21 = bsk::get<7>(t0_); + auto t22 = bsk::get<8>(t0_); + auto t1_ = _dual_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_pr = bsk::get<0>(t1_); + auto out_pi = bsk::get<1>(t1_); + auto t2_ = _dual_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_mr = bsk::get<0>(t2_); + auto out_mi = bsk::get<1>(t2_); + auto t3_ = _dual_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_zr = bsk::get<0>(t3_); + auto out_zi = bsk::get<1>(t3_); + auto t4_ = _dual_row(t00, t01, t02, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dpr = bsk::get<0>(t4_); + auto dpi = bsk::get<1>(t4_); + auto t5_ = _dual_row(t10, t11, t12, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dmr = bsk::get<0>(t5_); + auto dmi = bsk::get<1>(t5_); + auto t6_ = _dual_row(t20, t21, t22, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dzr = bsk::get<0>(t6_); + auto dzi = bsk::get<1>(t6_); + auto t7_ = _tangent_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tpr = bsk::get<0>(t7_); + auto tpi = bsk::get<1>(t7_); + auto t8_ = _tangent_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tmr = bsk::get<0>(t8_); + auto tmi = bsk::get<1>(t8_); + auto t9_ = _tangent_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tzr = bsk::get<0>(t9_); + auto tzi = bsk::get<1>(t9_); + return bsk::make_tup(out_pr, out_pi, out_mr, out_mi, out_zr, out_zi, (dpr + tpr), (dpi + tpi), (dmr + tmr), (dmi + tmi), (dzr + tzr), (dzi + tzi)); +} + +// Seven of the nine rotation coefficients; the rest follow by symmetry. +// +// ``t11`` repeats ``t00`` and ``t10`` is the conjugate of ``t01``, so the +// caller derives those. Feeding ``(cos, sin)`` gives the rotation itself and +// ``(sin, cos)`` rearranged gives its derivative in the flip angle, which is +// why this is one routine rather than two. +template +BSK_HD auto _rotation_block(const T0& a_value, const T1& a_tangent, const T2& b_value, const T3& b_tangent, const T4& c_value, const T5& c_tangent, const T6& d_value, const T7& d_tangent, const T8& p1r, const T9& p1i, const T10& p1tr, const T11& p1ti, const T12& p2r, const T13& p2i, const T14& p2tr, const T15& p2ti, const T16& pcr, const T17& pci, const T18& pctr, const T19& pcti) { + auto t00 = bsk::make_tup(a_value, (0.0f * a_value), a_tangent, (0.0f * a_tangent)); + auto t01 = _dual_scale(b_value, b_tangent, p2r, p2i, p2tr, p2ti); + auto t02 = _dual_mul((0.0f * c_value), (-c_value), (0.0f * c_tangent), (-c_tangent), p1r, p1i, p1tr, p1ti); + auto t12 = _dual_mul((0.0f * c_value), c_value, (0.0f * c_tangent), c_tangent, pcr, pci, pctr, pcti); + auto t20 = _dual_mul((0.0f * c_value), (-0.5f * c_value), (0.0f * c_tangent), (-0.5f * c_tangent), pcr, pci, pctr, pcti); + auto t21 = _dual_mul((0.0f * c_value), (0.5f * c_value), (0.0f * c_tangent), (0.5f * c_tangent), p1r, p1i, p1tr, p1ti); + auto t22 = bsk::make_tup(d_value, (0.0f * d_value), d_tangent, (0.0f * d_tangent)); + return bsk::make_tup(t00, t01, t02, t12, t20, t21, t22); +} + +// ``values`` moved one order down: ``result[k] = values[k + 1]``. +// +// The top order has no neighbour to read, so it reads itself and the caller +// masks it away. +template +BSK_HD auto _down(const T0& values, const T1& state) { + return bsk::gather_x(values, bsk::minimum(state + 1, bsk::width_x() - 1)); +} + +// ``values`` moved one configuration order up: ``result[k] = values[k - 1]``. +// +// Order zero is left to the caller, which fills it from the sequence's own +// boundary condition rather than from a neighbour. +template +BSK_HD auto _up(const T0& values, const T1& state) { + return bsk::gather_x(values, bsk::maximum(state - 1, 0)); +} + +template +BSK_HD auto _shift(const T0& fplus_real, const T1& fplus_imag, const T2& fminus_real, const T3& fminus_imag, const T4& state, const T5& state_mask, const T6& state_count) { + bsk::tile_t | 0, 3)> plus_imag{}; + bsk::tile_t | 0, 3)> plus_real{}; + auto keep_up = bsk::band((state > 0), state_mask); + auto keep_down = bsk::band(((state + 1) < state_count), state_mask); + plus_real = bsk::where(keep_up, _up(fplus_real, state), 0.0f); + plus_imag = bsk::where(keep_up, _up(fplus_imag, state), 0.0f); + auto minus_real = bsk::where(keep_down, _down(fminus_real, state), 0.0f); + auto minus_imag = bsk::where(keep_down, _down(fminus_imag, state), 0.0f); + plus_real = bsk::where((state == 0), minus_real, plus_real); + plus_imag = bsk::where((state == 0), (-minus_imag), plus_imag); + return bsk::make_tup(plus_real, plus_imag, minus_real, minus_imag); +} + +// Order zero of ``values``, spread across every order. +template +BSK_HD auto _first(const T0& values, const T1& state) { + return bsk::gather_x(values, state * 0); +} + +// Transpose of ``_shift``. +// +// The conjugate refill at order zero sends the incoming plus adjoint back +// onto minus, conjugated, at the index the minus shift moves it to. +template +BSK_HD auto _shift_adjoint(const T0& plus_bar_real, const T1& plus_bar_imag, const T2& minus_bar_real, const T3& minus_bar_imag, const T4& state, const T5& state_mask, const T6& state_count) { + bsk::tile_t | 0, 3)> shifted_mi{}; + bsk::tile_t | 0, 3)> shifted_mr{}; + auto carry_real = bsk::where(state_mask, _first(plus_bar_real, state), 0.0f); + auto carry_imag = (-bsk::where(state_mask, _first(plus_bar_imag, state), 0.0f)); + auto forward = bsk::band(((state + 1) < state_count), state_mask); + auto backward = bsk::band((state > 0), state_mask); + auto shifted_pr = bsk::where(forward, _down(plus_bar_real, state), 0.0f); + auto shifted_pi = bsk::where(forward, _down(plus_bar_imag, state), 0.0f); + shifted_mr = bsk::where(backward, _up(minus_bar_real, state), 0.0f); + shifted_mi = bsk::where(backward, _up(minus_bar_imag, state), 0.0f); + shifted_mr = bsk::where((state == 1), (shifted_mr + carry_real), shifted_mr); + shifted_mi = bsk::where((state == 1), (shifted_mi + carry_imag), shifted_mi); + return bsk::make_tup(shifted_pr, shifted_pi, shifted_mr, shifted_mi); +} + +// Four dual complex numbers added. +template +BSK_HD auto _dual_sum(const T0& first, const T1& second, const T2& third, const T3& fourth) { + return bsk::make_tup((((bsk::get<0>(first) + bsk::get<0>(second)) + bsk::get<0>(third)) + bsk::get<0>(fourth)), (((bsk::get<1>(first) + bsk::get<1>(second)) + bsk::get<1>(third)) + bsk::get<1>(fourth)), (((bsk::get<2>(first) + bsk::get<2>(second)) + bsk::get<2>(third)) + bsk::get<2>(fourth)), (((bsk::get<3>(first) + bsk::get<3>(second)) + bsk::get<3>(third)) + bsk::get<3>(fourth))); +} + +// A dual complex number scaled by a real constant. +template +BSK_HD auto _dual_weigh(const T0& z, const T1& factor) { + return bsk::make_tup((factor * bsk::get<0>(z)), (factor * bsk::get<1>(z)), (factor * bsk::get<2>(z)), (factor * bsk::get<3>(z))); +} + +// The spinor rotation's adjoint, on dual numbers. +// +// Returns the cotangent on the Cayley-Klein pair and the three state +// cotangents sent back through the conjugate transpose. Every entry of the +// matrix is a product of two factors drawn from the pair and its conjugate, +// so the pair's two Wirtinger halves are linear in the outer product of the +// seed with the state the rotation acted on -- which is why this is a closed +// form rather than a differentiated matrix. +template +BSK_HD auto _spinor_adjoint_dual(const T0& a, const T1& b, const T2& sp, const T3& sm, const T4& rz, const T5& pb, const T6& mb, const T7& zb) { + auto t0_ = _spinor_coefficients(bsk::get<0>(a), bsk::get<1>(a), bsk::get<0>(b), bsk::get<1>(b), bsk::get<2>(a), bsk::get<3>(a), bsk::get<2>(b), bsk::get<3>(b)); + auto t00 = bsk::get<0>(t0_); + auto t01 = bsk::get<1>(t0_); + auto t02 = bsk::get<2>(t0_); + auto t10 = bsk::get<3>(t0_); + auto t11 = bsk::get<4>(t0_); + auto t12 = bsk::get<5>(t0_); + auto t20 = bsk::get<6>(t0_); + auto t21 = bsk::get<7>(t0_); + auto t22 = bsk::get<8>(t0_); + auto conj_pb = _dual_conj(pb); + auto conj_mb = _dual_conj(mb); + auto conj_zb = _dual_conj(zb); + auto m00 = _dual_product(conj_pb, sp); + auto m01 = _dual_product(conj_pb, sm); + auto m02 = _dual_product(conj_pb, rz); + auto m10 = _dual_product(conj_mb, sp); + auto m11 = _dual_product(conj_mb, sm); + auto m12 = _dual_product(conj_mb, rz); + auto m20 = _dual_product(conj_zb, sp); + auto m21 = _dual_product(conj_zb, sm); + auto m22 = _dual_product(conj_zb, rz); + auto conj_a = _dual_conj(a); + auto conj_b = _dual_conj(b); + auto holding_conj_a = _dual_sum(_dual_weigh(_dual_product(a, m11), 2.0f), _dual_weigh(_dual_product(b, m12), -2.0f), _dual_product(conj_b, m21), _dual_product(conj_a, m22)); + auto holding_a = _dual_sum(_dual_weigh(_dual_product(conj_a, m00), 2.0f), _dual_weigh(_dual_product(conj_b, m02), -2.0f), _dual_product(b, m20), _dual_product(a, m22)); + auto holding_conj_b = _dual_sum(_dual_weigh(_dual_product(b, m10), -2.0f), _dual_weigh(_dual_product(a, m12), -2.0f), _dual_product(conj_a, m20), _dual_weigh(_dual_product(conj_b, m22), -1.0f)); + auto holding_b = _dual_sum(_dual_weigh(_dual_product(conj_b, m01), -2.0f), _dual_weigh(_dual_product(conj_a, m02), -2.0f), _dual_product(a, m21), _dual_weigh(_dual_product(b, m22), -1.0f)); + auto zero = _dual_weigh(m00, 0.0f); + auto grad_a = _dual_sum(_dual_conj(holding_conj_a), holding_a, zero, zero); + auto grad_b = _dual_sum(_dual_conj(holding_conj_b), holding_b, zero, zero); + auto next_pb = _dual_sum(_dual_product(_dual_conj(t00), pb), _dual_product(_dual_conj(t10), mb), _dual_product(_dual_conj(t20), zb), zero); + auto next_mb = _dual_sum(_dual_product(_dual_conj(t01), pb), _dual_product(_dual_conj(t11), mb), _dual_product(_dual_conj(t21), zb), zero); + auto next_zb = _dual_sum(_dual_product(_dual_conj(t02), pb), _dual_product(_dual_conj(t12), mb), _dual_product(_dual_conj(t22), zb), zero); + return bsk::make_tup(grad_a, grad_b, next_pb, next_mb, next_zb); +} + +// Send the cotangent on one pulse's rotation to its row. +// +// Summed over the dephasing orders first: the pair multiplies every one of +// them, so what reaches the row is the sum. The value plane is the adjoint +// and the tangent plane its own derivative, which is the split every other +// gradient here takes. +template +BSK_HD auto _store_pair_cotangent(const T0& grad_value, const T1& grad_tangent, const T2& pair_index, const T3& event_base, const T4& event, const T5& atom, const T6& atom_count, const T7& turning, const T8& mask, const T9& state_mask, const T10& grad_a, const T11& grad_b) { + auto row = bsk::cast(bsk::ld(((pair_index + event_base) + event))); + auto entry = (((row * atom_count) + atom) * 4); + // The block is padded to a power of two, and the orders past the last one + // carry whatever the sweep left there -- so the sum is taken over the + // orders that exist rather than over the block. + auto keep = bsk::band(turning, state_mask); + bsk::atomic_add(((grad_value + entry) + 0), bsk::sum_x(bsk::where(keep, bsk::get<0>(grad_a), 0.0f)), mask); + bsk::atomic_add(((grad_value + entry) + 1), bsk::sum_x(bsk::where(keep, bsk::get<1>(grad_a), 0.0f)), mask); + bsk::atomic_add(((grad_value + entry) + 2), bsk::sum_x(bsk::where(keep, bsk::get<0>(grad_b), 0.0f)), mask); + bsk::atomic_add(((grad_value + entry) + 3), bsk::sum_x(bsk::where(keep, bsk::get<1>(grad_b), 0.0f)), mask); + bsk::atomic_add(((grad_tangent + entry) + 0), bsk::sum_x(bsk::where(keep, bsk::get<2>(grad_a), 0.0f)), mask); + bsk::atomic_add(((grad_tangent + entry) + 1), bsk::sum_x(bsk::where(keep, bsk::get<3>(grad_a), 0.0f)), mask); + bsk::atomic_add(((grad_tangent + entry) + 2), bsk::sum_x(bsk::where(keep, bsk::get<2>(grad_b), 0.0f)), mask); + bsk::atomic_add(((grad_tangent + entry) + 3), bsk::sum_x(bsk::where(keep, bsk::get<3>(grad_b), 0.0f)), mask); +} + +// Which row of the stacked tables this pulse reads. +// +// Its own shape's block of ``locations`` rows, then the voxel's place along +// the slice. +template +BSK_HD auto _table_row(const T0& profile_index, const T1& event, const T2& location, const T3& locations) { + return ((bsk::cast(bsk::ld((profile_index + event))) * locations) + location); +} + +// The bare three-pool operator, assembled from its shared pieces. +// +// Both branches are formed and one is chosen: a ``where`` evaluates each +// side, so the divisor each of them carries is guarded whether or not it +// is the side taken. +// +// In double, and before any attenuation -- what a reverse sweep reads, +// and what :func:`_three_pool_weigh_jvp` turns into an interval's step. +template +BSK_HD auto _three_pool_assemble_jvp(const T0& free, const T1& d_free, const T2& pool_b, const T3& d_pool_b, const T4& pool_c, const T5& d_pool_c, const T6& a00, const T7& d_a00, const T8& a01, const T9& d_a01, const T10& a02, const T11& d_a02, const T12& a10, const T13& d_a10, const T14& a11, const T15& d_a11, const T16& a20, const T17& d_a20, const T18& a22, const T19& d_a22, const T20& s00, const T21& d_s00, const T22& s11, const T23& d_s11, const T24& s22, const T25& d_s22, const T26& minors, const T27& d_minors, const T28& sum_flat, const T29& sum_linear, const T30& sum_square, const T31& d_sum_flat, const T32& d_sum_linear, const T33& d_sum_square, const T34& lift, const T35& d_lift, const T36& low, const T37& middle, const T38& d_low, const T39& d_middle, const T40& leading, const T41& d_leading, const T42& first, const T43& d_first, const T44& second, const T45& d_second, const T46& determinant, const T47& d_determinant, const T48& high, const T49& d_high, const T50& radius, const T51& d_radius, const T52& cube, const T53& raw, const T54& d_raw, const T55& argument, const T56& inside_limit, const T57& angle, const T58& d_angle, const T59& centre, const T60& d_centre, const T61& trailing, const T62& d_trailing, const T63& guarded, const T64& d_guarded, const T65& q00, const T66& d_q00, const T67& q01, const T68& d_q01, const T69& q02, const T70& d_q02, const T71& q10, const T72& d_q10, const T73& q11, const T74& d_q11, const T75& q12, const T76& d_q12, const T77& q20, const T78& d_q20, const T79& q21, const T80& d_q21, const T81& q22, const T82& d_q22, const T83& narrow) { + bsk::tile_t | 0, 3)> def_00{}; + bsk::tile_t | 0, 3)> def_01{}; + bsk::tile_t | 0, 3)> def_02{}; + bsk::tile_t | 0, 3)> def_10{}; + bsk::tile_t | 0, 3)> def_11{}; + bsk::tile_t | 0, 3)> def_12{}; + bsk::tile_t | 0, 3)> def_20{}; + bsk::tile_t | 0, 3)> def_21{}; + bsk::tile_t | 0, 3)> def_22{}; + bsk::tile_t | 0, 3)> dif_00{}; + bsk::tile_t | 0, 3)> dif_01{}; + bsk::tile_t | 0, 3)> dif_02{}; + bsk::tile_t | 0, 3)> dif_10{}; + bsk::tile_t | 0, 3)> dif_11{}; + bsk::tile_t | 0, 3)> dif_12{}; + bsk::tile_t | 0, 3)> dif_20{}; + bsk::tile_t | 0, 3)> dif_21{}; + bsk::tile_t | 0, 3)> dif_22{}; + auto c00 = (lift * ((sum_flat + (sum_linear * s00)) + (sum_square * q00))); + auto d_c00 = ((d_lift * ((sum_flat + (sum_linear * s00)) + (sum_square * q00))) + (lift * ((((d_sum_flat + (d_sum_linear * s00)) + (sum_linear * d_s00)) + (d_sum_square * q00)) + (sum_square * d_q00)))); + auto c01 = (lift * ((sum_linear * a01) + (sum_square * q01))); + auto d_c01 = ((d_lift * ((sum_linear * a01) + (sum_square * q01))) + (lift * ((((d_sum_linear * a01) + (sum_linear * d_a01)) + (d_sum_square * q01)) + (sum_square * d_q01)))); + auto c02 = (lift * ((sum_linear * a02) + (sum_square * q02))); + auto d_c02 = ((d_lift * ((sum_linear * a02) + (sum_square * q02))) + (lift * ((((d_sum_linear * a02) + (sum_linear * d_a02)) + (d_sum_square * q02)) + (sum_square * d_q02)))); + auto c10 = (lift * ((sum_linear * a10) + (sum_square * q10))); + auto d_c10 = ((d_lift * ((sum_linear * a10) + (sum_square * q10))) + (lift * ((((d_sum_linear * a10) + (sum_linear * d_a10)) + (d_sum_square * q10)) + (sum_square * d_q10)))); + auto c11 = (lift * ((sum_flat + (sum_linear * s11)) + (sum_square * q11))); + auto d_c11 = ((d_lift * ((sum_flat + (sum_linear * s11)) + (sum_square * q11))) + (lift * ((((d_sum_flat + (d_sum_linear * s11)) + (sum_linear * d_s11)) + (d_sum_square * q11)) + (sum_square * d_q11)))); + auto c12 = (lift * (sum_square * q12)); + auto d_c12 = ((d_lift * (sum_square * q12)) + (lift * ((d_sum_square * q12) + (sum_square * d_q12)))); + auto c20 = (lift * ((sum_linear * a20) + (sum_square * q20))); + auto d_c20 = ((d_lift * ((sum_linear * a20) + (sum_square * q20))) + (lift * ((((d_sum_linear * a20) + (sum_linear * d_a20)) + (d_sum_square * q20)) + (sum_square * d_q20)))); + auto c21 = (lift * (sum_square * q21)); + auto d_c21 = ((d_lift * (sum_square * q21)) + (lift * ((d_sum_square * q21) + (sum_square * d_q21)))); + auto c22 = (lift * ((sum_flat + (sum_linear * s22)) + (sum_square * q22))); + auto d_c22 = ((d_lift * ((sum_flat + (sum_linear * s22)) + (sum_square * q22))) + (lift * ((((d_sum_flat + (d_sum_linear * s22)) + (sum_linear * d_s22)) + (d_sum_square * q22)) + (sum_square * d_q22)))); + // --- the Newton form's two factors, for the eigenvalue branch --- + auto m00 = (a00 - low); + auto d_m00 = (d_a00 - d_low); + auto m11 = (a11 - low); + auto d_m11 = (d_a11 - d_low); + auto m22 = (a22 - low); + auto d_m22 = (d_a22 - d_low); + auto n00 = (a00 - middle); + auto d_n00 = (d_a00 - d_middle); + auto n11 = (a11 - middle); + auto d_n11 = (d_a11 - d_middle); + auto n22 = (a22 - middle); + auto d_n22 = (d_a22 - d_middle); + auto p00 = (((m00 * n00) + (a01 * a10)) + (a02 * a20)); + auto d_p00 = ((((((d_m00 * n00) + (m00 * d_n00)) + (d_a01 * a10)) + (a01 * d_a10)) + (d_a02 * a20)) + (a02 * d_a20)); + auto p01 = (a01 * (m00 + n11)); + auto d_p01 = ((d_a01 * (m00 + n11)) + (a01 * (d_m00 + d_n11))); + auto p02 = (a02 * (m00 + n22)); + auto d_p02 = ((d_a02 * (m00 + n22)) + (a02 * (d_m00 + d_n22))); + auto p10 = (a10 * (n00 + m11)); + auto d_p10 = ((d_a10 * (n00 + m11)) + (a10 * (d_n00 + d_m11))); + auto p11 = ((a10 * a01) + (m11 * n11)); + auto d_p11 = ((((d_a10 * a01) + (a10 * d_a01)) + (d_m11 * n11)) + (m11 * d_n11)); + auto p12 = (a10 * a02); + auto d_p12 = ((d_a10 * a02) + (a10 * d_a02)); + auto p20 = (a20 * (n00 + m22)); + auto d_p20 = ((d_a20 * (n00 + m22)) + (a20 * (d_n00 + d_m22))); + auto p21 = (a20 * a01); + auto d_p21 = ((d_a20 * a01) + (a20 * d_a01)); + auto p22 = ((a20 * a02) + (m22 * n22)); + auto d_p22 = ((((d_a20 * a02) + (a20 * d_a02)) + (d_m22 * n22)) + (m22 * d_n22)); + auto e00 = ((leading + (first * m00)) + (second * p00)); + auto d_e00 = ((((d_leading + (d_first * m00)) + (first * d_m00)) + (d_second * p00)) + (second * d_p00)); + auto e01 = ((first * a01) + (second * p01)); + auto d_e01 = ((((d_first * a01) + (first * d_a01)) + (d_second * p01)) + (second * d_p01)); + auto e02 = ((first * a02) + (second * p02)); + auto d_e02 = ((((d_first * a02) + (first * d_a02)) + (d_second * p02)) + (second * d_p02)); + auto e10 = ((first * a10) + (second * p10)); + auto d_e10 = ((((d_first * a10) + (first * d_a10)) + (d_second * p10)) + (second * d_p10)); + auto e11 = ((leading + (first * m11)) + (second * p11)); + auto d_e11 = ((((d_leading + (d_first * m11)) + (first * d_m11)) + (d_second * p11)) + (second * d_p11)); + auto e12 = (second * p12); + auto d_e12 = ((d_second * p12) + (second * d_p12)); + auto e20 = ((first * a20) + (second * p20)); + auto d_e20 = ((((d_first * a20) + (first * d_a20)) + (d_second * p20)) + (second * d_p20)); + auto e21 = (second * p21); + auto d_e21 = ((d_second * p21) + (second * d_p21)); + auto e22 = ((leading + (first * m22)) + (second * p22)); + auto d_e22 = ((((d_leading + (d_first * m22)) + (first * d_m22)) + (d_second * p22)) + (second * d_p22)); + if (bsk::truth(narrow)) { + // The caller has bounded the spread, so the roots are unreachable and + // everything that leads to them goes with this select. + auto t0_ = bsk::make_tup(c00, d_c00); + def_00 = bsk::get<0>(t0_); + dif_00 = bsk::get<1>(t0_); + auto t1_ = bsk::make_tup(c01, d_c01); + def_01 = bsk::get<0>(t1_); + dif_01 = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup(c02, d_c02); + def_02 = bsk::get<0>(t2_); + dif_02 = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup(c10, d_c10); + def_10 = bsk::get<0>(t3_); + dif_10 = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(c11, d_c11); + def_11 = bsk::get<0>(t4_); + dif_11 = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup(c12, d_c12); + def_12 = bsk::get<0>(t5_); + dif_12 = bsk::get<1>(t5_); + auto t6_ = bsk::make_tup(c20, d_c20); + def_20 = bsk::get<0>(t6_); + dif_20 = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(c21, d_c21); + def_21 = bsk::get<0>(t7_); + dif_21 = bsk::get<1>(t7_); + auto t8_ = bsk::make_tup(c22, d_c22); + def_22 = bsk::get<0>(t8_); + dif_22 = bsk::get<1>(t8_); + } else { + auto close = ((-2.0f * minors) < 1.0f); + def_00 = bsk::where(close, c00, e00); + dif_00 = bsk::where(close, d_c00, d_e00); + def_01 = bsk::where(close, c01, e01); + dif_01 = bsk::where(close, d_c01, d_e01); + def_02 = bsk::where(close, c02, e02); + dif_02 = bsk::where(close, d_c02, d_e02); + def_10 = bsk::where(close, c10, e10); + dif_10 = bsk::where(close, d_c10, d_e10); + def_11 = bsk::where(close, c11, e11); + dif_11 = bsk::where(close, d_c11, d_e11); + def_12 = bsk::where(close, c12, e12); + dif_12 = bsk::where(close, d_c12, d_e12); + def_20 = bsk::where(close, c20, e20); + dif_20 = bsk::where(close, d_c20, d_e20); + def_21 = bsk::where(close, c21, e21); + dif_21 = bsk::where(close, d_c21, d_e21); + def_22 = bsk::where(close, c22, e22); + dif_22 = bsk::where(close, d_c22, d_e22); + } + return bsk::make_tup(def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22); +} + +// Read one interval's three-pool operator and a direction through it. +// +// The row carries the tissue's share of the direction; the interval's own is +// ``A1 C d_dt``, formed here because ``d_dt`` belongs to the event and the +// row is shared. Returns the nine entries and three recoveries with their +// tangents, in the order :func:`_three_pool_step_jvp` returns them. +template +BSK_HD auto _three_pool_from_table_jvp(const T0& table, const T1& row, const T2& atom, const T3& voxel_count, const T4& mask, const T5& r1_free, const T6& r1_pool_b, const T7& r1_bound, const T8& exchange_b, const T9& exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& d_dt, const T15& attenuation, const T16& d_attenuation) { + auto free = ((1.0f - fraction_b) - fraction_c); + auto d_free = ((-d_fraction_b) - d_fraction_c); + auto a00 = ((((-exchange_b) * fraction_b) - (exchange_c * fraction_c)) - r1_free); + auto a01 = (exchange_b * free); + auto a02 = (exchange_c * free); + auto a10 = (exchange_b * fraction_b); + auto a11 = (((-exchange_b) * free) - r1_pool_b); + auto a20 = (exchange_c * fraction_c); + auto a22 = (((-exchange_c) * free) - r1_bound); + auto base = ((table + (row * (18 * voxel_count))) + atom); + auto c00 = bsk::ld((base + (0 * voxel_count)), mask, 0.0f); + auto c01 = bsk::ld((base + (1 * voxel_count)), mask, 0.0f); + auto c02 = bsk::ld((base + (2 * voxel_count)), mask, 0.0f); + auto c10 = bsk::ld((base + (3 * voxel_count)), mask, 0.0f); + auto c11 = bsk::ld((base + (4 * voxel_count)), mask, 0.0f); + auto c12 = bsk::ld((base + (5 * voxel_count)), mask, 0.0f); + auto c20 = bsk::ld((base + (6 * voxel_count)), mask, 0.0f); + auto c21 = bsk::ld((base + (7 * voxel_count)), mask, 0.0f); + auto c22 = bsk::ld((base + (8 * voxel_count)), mask, 0.0f); + // The row's tangent, plus what the event's own interval direction adds. + auto t00 = (bsk::ld((base + (9 * voxel_count)), mask, 0.0f) + (d_dt * (((a00 * c00) + (a01 * c10)) + (a02 * c20)))); + auto t01 = (bsk::ld((base + (10 * voxel_count)), mask, 0.0f) + (d_dt * (((a00 * c01) + (a01 * c11)) + (a02 * c21)))); + auto t02 = (bsk::ld((base + (11 * voxel_count)), mask, 0.0f) + (d_dt * (((a00 * c02) + (a01 * c12)) + (a02 * c22)))); + auto t10 = (bsk::ld((base + (12 * voxel_count)), mask, 0.0f) + (d_dt * ((a10 * c00) + (a11 * c10)))); + auto t11 = (bsk::ld((base + (13 * voxel_count)), mask, 0.0f) + (d_dt * ((a10 * c01) + (a11 * c11)))); + auto t12 = (bsk::ld((base + (14 * voxel_count)), mask, 0.0f) + (d_dt * ((a10 * c02) + (a11 * c12)))); + auto t20 = (bsk::ld((base + (15 * voxel_count)), mask, 0.0f) + (d_dt * ((a20 * c00) + (a22 * c20)))); + auto t21 = (bsk::ld((base + (16 * voxel_count)), mask, 0.0f) + (d_dt * ((a20 * c01) + (a22 * c21)))); + auto t22 = (bsk::ld((base + (17 * voxel_count)), mask, 0.0f) + (d_dt * ((a20 * c02) + (a22 * c22)))); + auto e00 = (attenuation * c00); + auto e01 = (attenuation * c01); + auto e02 = (attenuation * c02); + auto e10 = (attenuation * c10); + auto e11 = (attenuation * c11); + auto e12 = (attenuation * c12); + auto e20 = (attenuation * c20); + auto e21 = (attenuation * c21); + auto e22 = (attenuation * c22); + auto f00 = ((d_attenuation * c00) + (attenuation * t00)); + auto f01 = ((d_attenuation * c01) + (attenuation * t01)); + auto f02 = ((d_attenuation * c02) + (attenuation * t02)); + auto f10 = ((d_attenuation * c10) + (attenuation * t10)); + auto f11 = ((d_attenuation * c11) + (attenuation * t11)); + auto f12 = ((d_attenuation * c12) + (attenuation * t12)); + auto f20 = ((d_attenuation * c20) + (attenuation * t20)); + auto f21 = ((d_attenuation * c21) + (attenuation * t21)); + auto f22 = ((d_attenuation * c22) + (attenuation * t22)); + // The equilibrium the recoveries are taken against moves with the + // fractions, so it carries a direction of its own. + auto grow_free = (free - (((e00 * free) + (e01 * fraction_b)) + (e02 * fraction_c))); + auto grow_pool_b = (fraction_b - (((e10 * free) + (e11 * fraction_b)) + (e12 * fraction_c))); + auto grow_bound = (fraction_c - (((e20 * free) + (e21 * fraction_b)) + (e22 * fraction_c))); + auto d_grow_free = (d_free - ((((((f00 * free) + (f01 * fraction_b)) + (f02 * fraction_c)) + (e00 * d_free)) + (e01 * d_fraction_b)) + (e02 * d_fraction_c))); + auto d_grow_pool_b = (d_fraction_b - ((((((f10 * free) + (f11 * fraction_b)) + (f12 * fraction_c)) + (e10 * d_free)) + (e11 * d_fraction_b)) + (e12 * d_fraction_c))); + auto d_grow_bound = (d_fraction_c - ((((((f20 * free) + (f21 * fraction_b)) + (f22 * fraction_c)) + (e20 * d_free)) + (e21 * d_fraction_b)) + (e22 * d_fraction_c))); + return bsk::make_tup(e00, e01, e02, e10, e11, e12, e20, e21, e22, grow_free, grow_pool_b, grow_bound, f00, f01, f02, f10, f11, f12, f20, f21, f22, d_grow_free, d_grow_pool_b, d_grow_bound); +} + +// Nine cotangents against nine entries, less what the recoveries take. +// +// A recovery is ``m - E m`` with ``m`` the equilibrium +// ``(free, fraction_b, fraction_c)``, so it differentiates through the same +// nine entries with the equilibrium contracted out of them. +template +BSK_HD auto _three_pool_contract(const T0& x00, const T1& x01, const T2& x02, const T3& x10, const T4& x11, const T5& x12, const T6& x20, const T7& x21, const T8& x22, const T9& e11, const T10& e12, const T11& e13, const T12& e21, const T13& e22, const T14& e23, const T15& e31, const T16& e32, const T17& e33, const T18& rec_free, const T19& rec_pool_b, const T20& rec_bound, const T21& free, const T22& fraction_b, const T23& fraction_c) { + return ((((((((((((e11 * x00) + (e12 * x01)) + (e13 * x02)) + (e21 * x10)) + (e22 * x11)) + (e23 * x12)) + (e31 * x20)) + (e32 * x21)) + (e33 * x22)) - (rec_free * (((x00 * free) + (x01 * fraction_b)) + (x02 * fraction_c)))) - (rec_pool_b * (((x10 * free) + (x11 * fraction_b)) + (x12 * fraction_c)))) - (rec_bound * (((x20 * free) + (x21 * fraction_b)) + (x22 * fraction_c)))); +} + +// The interval and the attenuation, from a dual pair's cotangents. +// +// ``dE/d(dt)`` is ``A1 E``, and the direction that quantity carries follows +// from the same generator: with ``C_dot == C_row + A1 C d_dt``, the +// derivative in the interval is ``A1_dot C + A1 C_row + A1 A1 C d_dt``. So +// three products of the generator against the tabulated row serve what the +// eigenvalues would otherwise be re-formed for, and these two quantities are +// the only ones that stay per event. +// +// Returns the interval and attenuation gradients, value then tangent, in the +// order :func:`_three_pool_step_adjoint_jvp` returns them. +template +BSK_HD auto _three_pool_interval_adjoint_jvp(const T0& table, const T1& row, const T2& atom, const T3& voxel_count, const T4& mask, const T5& r1_free, const T6& d_r1_free, const T7& r1_pool_b, const T8& d_r1_pool_b, const T9& r1_bound, const T10& d_r1_bound, const T11& exchange_b, const T12& d_exchange_b, const T13& exchange_c, const T14& d_exchange_c, const T15& fraction_b, const T16& d_fraction_b, const T17& fraction_c, const T18& d_fraction_c, const T19& d_dt, const T20& attenuation, const T21& d_attenuation, const T22& b11, const T23& b12, const T24& b13, const T25& b21, const T26& b22, const T27& b23, const T28& b31, const T29& b32, const T30& b33, const T31& bfree, const T32& bpool_b, const T33& bbound, const T34& t11, const T35& t12, const T36& t13, const T37& t21, const T38& t22, const T39& t23, const T40& t31, const T41& t32, const T42& t33, const T43& tfree, const T44& tpool_b, const T45& tbound) { + auto free = ((1.0f - fraction_b) - fraction_c); + auto d_free = ((-d_fraction_b) - d_fraction_c); + auto a00 = ((((-exchange_b) * fraction_b) - (exchange_c * fraction_c)) - r1_free); + auto a01 = (exchange_b * free); + auto a02 = (exchange_c * free); + auto a10 = (exchange_b * fraction_b); + auto a11 = (((-exchange_b) * free) - r1_pool_b); + auto a20 = (exchange_c * fraction_c); + auto a22 = (((-exchange_c) * free) - r1_bound); + auto da00 = ((((((-d_exchange_b) * fraction_b) - (exchange_b * d_fraction_b)) - (d_exchange_c * fraction_c)) - (exchange_c * d_fraction_c)) - d_r1_free); + auto da01 = ((d_exchange_b * free) + (exchange_b * d_free)); + auto da02 = ((d_exchange_c * free) + (exchange_c * d_free)); + auto da10 = ((d_exchange_b * fraction_b) + (exchange_b * d_fraction_b)); + auto da11 = ((((-d_exchange_b) * free) - (exchange_b * d_free)) - d_r1_pool_b); + auto da20 = ((d_exchange_c * fraction_c) + (exchange_c * d_fraction_c)); + auto da22 = ((((-d_exchange_c) * free) - (exchange_c * d_free)) - d_r1_bound); + auto base = ((table + (row * (18 * voxel_count))) + atom); + auto c00 = bsk::ld((base + (0 * voxel_count)), mask, 0.0f); + auto c01 = bsk::ld((base + (1 * voxel_count)), mask, 0.0f); + auto c02 = bsk::ld((base + (2 * voxel_count)), mask, 0.0f); + auto c10 = bsk::ld((base + (3 * voxel_count)), mask, 0.0f); + auto c11 = bsk::ld((base + (4 * voxel_count)), mask, 0.0f); + auto c12 = bsk::ld((base + (5 * voxel_count)), mask, 0.0f); + auto c20 = bsk::ld((base + (6 * voxel_count)), mask, 0.0f); + auto c21 = bsk::ld((base + (7 * voxel_count)), mask, 0.0f); + auto c22 = bsk::ld((base + (8 * voxel_count)), mask, 0.0f); + auto r00 = bsk::ld((base + (9 * voxel_count)), mask, 0.0f); + auto r01 = bsk::ld((base + (10 * voxel_count)), mask, 0.0f); + auto r02 = bsk::ld((base + (11 * voxel_count)), mask, 0.0f); + auto r10 = bsk::ld((base + (12 * voxel_count)), mask, 0.0f); + auto r11 = bsk::ld((base + (13 * voxel_count)), mask, 0.0f); + auto r12 = bsk::ld((base + (14 * voxel_count)), mask, 0.0f); + auto r20 = bsk::ld((base + (15 * voxel_count)), mask, 0.0f); + auto r21 = bsk::ld((base + (16 * voxel_count)), mask, 0.0f); + auto r22 = bsk::ld((base + (17 * voxel_count)), mask, 0.0f); + // P = A1 C, Q = A1_dot C + A1 C_row, S = A1 P. + auto p00 = (((a00 * c00) + (a01 * c10)) + (a02 * c20)); + auto p01 = (((a00 * c01) + (a01 * c11)) + (a02 * c21)); + auto p02 = (((a00 * c02) + (a01 * c12)) + (a02 * c22)); + auto p10 = ((a10 * c00) + (a11 * c10)); + auto p11 = ((a10 * c01) + (a11 * c11)); + auto p12 = ((a10 * c02) + (a11 * c12)); + auto p20 = ((a20 * c00) + (a22 * c20)); + auto p21 = ((a20 * c01) + (a22 * c21)); + auto p22 = ((a20 * c02) + (a22 * c22)); + auto q00 = ((((((da00 * c00) + (da01 * c10)) + (da02 * c20)) + (a00 * r00)) + (a01 * r10)) + (a02 * r20)); + auto q01 = ((((((da00 * c01) + (da01 * c11)) + (da02 * c21)) + (a00 * r01)) + (a01 * r11)) + (a02 * r21)); + auto q02 = ((((((da00 * c02) + (da01 * c12)) + (da02 * c22)) + (a00 * r02)) + (a01 * r12)) + (a02 * r22)); + auto q10 = ((((da10 * c00) + (da11 * c10)) + (a10 * r00)) + (a11 * r10)); + auto q11 = ((((da10 * c01) + (da11 * c11)) + (a10 * r01)) + (a11 * r11)); + auto q12 = ((((da10 * c02) + (da11 * c12)) + (a10 * r02)) + (a11 * r12)); + auto q20 = ((((da20 * c00) + (da22 * c20)) + (a20 * r00)) + (a22 * r20)); + auto q21 = ((((da20 * c01) + (da22 * c21)) + (a20 * r01)) + (a22 * r21)); + auto q22 = ((((da20 * c02) + (da22 * c22)) + (a20 * r02)) + (a22 * r22)); + auto s00 = (((a00 * p00) + (a01 * p10)) + (a02 * p20)); + auto s01 = (((a00 * p01) + (a01 * p11)) + (a02 * p21)); + auto s02 = (((a00 * p02) + (a01 * p12)) + (a02 * p22)); + auto s10 = ((a10 * p00) + (a11 * p10)); + auto s11 = ((a10 * p01) + (a11 * p11)); + auto s12 = ((a10 * p02) + (a11 * p12)); + auto s20 = ((a20 * p00) + (a22 * p20)); + auto s21 = ((a20 * p01) + (a22 * p21)); + auto s22 = ((a20 * p02) + (a22 * p22)); + // The direction the tabulated operator carries, and the interval's own + // share of it. + auto d00 = (r00 + (p00 * d_dt)); + auto d01 = (r01 + (p01 * d_dt)); + auto d02 = (r02 + (p02 * d_dt)); + auto d10 = (r10 + (p10 * d_dt)); + auto d11 = (r11 + (p11 * d_dt)); + auto d12 = (r12 + (p12 * d_dt)); + auto d20 = (r20 + (p20 * d_dt)); + auto d21 = (r21 + (p21 * d_dt)); + auto d22 = (r22 + (p22 * d_dt)); + // dE/d(dt) with the attenuation held, and the direction that carries. + auto g00 = (attenuation * p00); + auto g01 = (attenuation * p01); + auto g02 = (attenuation * p02); + auto g10 = (attenuation * p10); + auto g11 = (attenuation * p11); + auto g12 = (attenuation * p12); + auto g20 = (attenuation * p20); + auto g21 = (attenuation * p21); + auto g22 = (attenuation * p22); + auto w00 = ((d_attenuation * p00) + (attenuation * (q00 + (s00 * d_dt)))); + auto w01 = ((d_attenuation * p01) + (attenuation * (q01 + (s01 * d_dt)))); + auto w02 = ((d_attenuation * p02) + (attenuation * (q02 + (s02 * d_dt)))); + auto w10 = ((d_attenuation * p10) + (attenuation * (q10 + (s10 * d_dt)))); + auto w11 = ((d_attenuation * p11) + (attenuation * (q11 + (s11 * d_dt)))); + auto w12 = ((d_attenuation * p12) + (attenuation * (q12 + (s12 * d_dt)))); + auto w20 = ((d_attenuation * p20) + (attenuation * (q20 + (s20 * d_dt)))); + auto w21 = ((d_attenuation * p21) + (attenuation * (q21 + (s21 * d_dt)))); + auto w22 = ((d_attenuation * p22) + (attenuation * (q22 + (s22 * d_dt)))); + // This kernel carries every quantity as a dual pair, so the tangent + // returned beside a gradient is that gradient's own directional + // derivative -- not the gradient with respect to the direction. + auto grad_dt_v = _three_pool_contract(g00, g01, g02, g10, g11, g12, g20, g21, g22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c); + auto grad_dt_t = ((_three_pool_contract(g00, g01, g02, g10, g11, g12, g20, g21, g22, t11, t12, t13, t21, t22, t23, t31, t32, t33, tfree, tpool_b, tbound, free, fraction_b, fraction_c) + _three_pool_contract(w00, w01, w02, w10, w11, w12, w20, w21, w22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c)) - (((bfree * (((g00 * d_free) + (g01 * d_fraction_b)) + (g02 * d_fraction_c))) + (bpool_b * (((g10 * d_free) + (g11 * d_fraction_b)) + (g12 * d_fraction_c)))) + (bbound * (((g20 * d_free) + (g21 * d_fraction_b)) + (g22 * d_fraction_c))))); + auto grad_att_v = _three_pool_contract(c00, c01, c02, c10, c11, c12, c20, c21, c22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c); + auto grad_att_t = ((_three_pool_contract(c00, c01, c02, c10, c11, c12, c20, c21, c22, t11, t12, t13, t21, t22, t23, t31, t32, t33, tfree, tpool_b, tbound, free, fraction_b, fraction_c) + _three_pool_contract(d00, d01, d02, d10, d11, d12, d20, d21, d22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c)) - (((bfree * (((c00 * d_free) + (c01 * d_fraction_b)) + (c02 * d_fraction_c))) + (bpool_b * (((c10 * d_free) + (c11 * d_fraction_b)) + (c12 * d_fraction_c)))) + (bbound * (((c20 * d_free) + (c21 * d_fraction_b)) + (c22 * d_fraction_c))))); + return bsk::make_tup(grad_dt_v, grad_att_v, grad_dt_t, grad_att_t); +} + +// ``[a, b] exp`` and its directional derivative. +// +// The derivative of a divided difference is the next one along, +// ``d/da [a,b] = [a,a,b]``, which near the coalescence is again a series in +// the gap's square rather than a quotient that vanishes over a vanishing +// denominator. +template +BSK_HD auto _exp_difference_jvp(const T0& lower, const T1& d_lower, const T2& upper, const T3& d_upper, const T4& low_exp, const T5& high_exp) { + auto half = (0.5f * (upper - lower)); + auto d_half = (0.5f * (d_upper - d_lower)); + auto near = (bsk::abs(half) < 0.0001f); + auto square = (half * half); + auto d_square = ((2.0f * half) * d_half); + // exp(mid) * sinh(half)/half, both factors expanded about zero. + auto lift = (low_exp * ((1.0f + half) + (0.5f * square))); + auto d_lift = (low_exp * (((d_lower * ((1.0f + half) + (0.5f * square))) + d_half) + (half * d_half))); + auto sinch = ((1.0f + bsk::truediv(square, 6.0f)) + bsk::truediv((square * square), 120.0f)); + auto d_sinch = (bsk::truediv(d_square, 6.0f) + bsk::truediv(((2.0f * square) * d_square), 120.0f)); + auto series = (lift * sinch); + auto d_series = ((d_lift * sinch) + (lift * d_sinch)); + auto gap = bsk::where(near, 1.0f, (upper - lower)); + auto d_gap = bsk::where(near, 0.0f, (d_upper - d_lower)); + auto quotient = bsk::truediv((high_exp - low_exp), gap); + auto d_quotient = bsk::truediv((((high_exp * d_upper) - (low_exp * d_lower)) - (quotient * d_gap)), gap); + return bsk::make_tup(bsk::where(near, series, quotient), bsk::where(near, d_series, d_quotient)); +} + +// The three-pool operator's shared front half, as duals. +// +// The generator, its two invariants, the series coefficients and the three +// roots with the divided differences between them -- everything both the +// operator and its reverse sweep are assembled from, computed once in double +// so the two cannot drift apart. +template +BSK_HD auto _three_pool_pieces_jvp_in_precision(const T0& r1_free, const T1& d_r1_free, const T2& r1_pool_b, const T3& d_r1_pool_b, const T4& r1_bound, const T5& d_r1_bound, const T6& exchange_b, const T7& d_exchange_b, const T8& exchange_c, const T9& d_exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& dt, const T15& d_dt, const T16& narrow) { + bsk::tile_t | 0, 3)> d_flat{}; + bsk::tile_t | 0, 3)> d_linear{}; + bsk::tile_t | 0, 3)> d_square{}; + bsk::tile_t | 0, 3)> d_sum_flat{}; + bsk::tile_t | 0, 3)> d_sum_linear{}; + bsk::tile_t | 0, 3)> d_sum_square{}; + float factorial{}; + bsk::tile_t | 0, 3)> flat{}; + bsk::tile_t | 0, 3)> linear{}; + bsk::tile_t | 0, 3)> square{}; + bsk::tile_t | 0, 3)> sum_flat{}; + bsk::tile_t | 0, 3)> sum_linear{}; + bsk::tile_t | 0, 3)> sum_square{}; + auto step = bsk::cast(dt); + auto d_step = bsk::cast(d_dt); + auto free = bsk::cast(((Work(1.0) - fraction_b) - fraction_c)); + auto d_free = bsk::cast(((-d_fraction_b) - d_fraction_c)); + auto pool_b = bsk::cast(fraction_b); + auto d_pool_b = bsk::cast(d_fraction_b); + auto pool_c = bsk::cast(fraction_c); + auto d_pool_c = bsk::cast(d_fraction_c); + auto rate_b = bsk::cast(exchange_b); + auto d_rate_b = bsk::cast(d_exchange_b); + auto rate_c = bsk::cast(exchange_c); + auto d_rate_c = bsk::cast(d_exchange_c); + auto kab = (rate_b * pool_b); + auto d_kab = ((d_rate_b * pool_b) + (rate_b * d_pool_b)); + auto kba = (rate_b * free); + auto d_kba = ((d_rate_b * free) + (rate_b * d_free)); + auto kac = (rate_c * pool_c); + auto d_kac = ((d_rate_c * pool_c) + (rate_c * d_pool_c)); + auto kca = (rate_c * free); + auto d_kca = ((d_rate_c * free) + (rate_c * d_free)); + auto row_a = (((-kab) - kac) - bsk::cast(r1_free)); + auto d_row_a = (((-d_kab) - d_kac) - bsk::cast(d_r1_free)); + auto row_b = ((-kba) - bsk::cast(r1_pool_b)); + auto d_row_b = ((-d_kba) - bsk::cast(d_r1_pool_b)); + auto row_c = ((-kca) - bsk::cast(r1_bound)); + auto d_row_c = ((-d_kca) - bsk::cast(d_r1_bound)); + auto a00 = (row_a * step); + auto d_a00 = ((d_row_a * step) + (row_a * d_step)); + auto a01 = (kba * step); + auto d_a01 = ((d_kba * step) + (kba * d_step)); + auto a02 = (kca * step); + auto d_a02 = ((d_kca * step) + (kca * d_step)); + auto a10 = (kab * step); + auto d_a10 = ((d_kab * step) + (kab * d_step)); + auto a11 = (row_b * step); + auto d_a11 = ((d_row_b * step) + (row_b * d_step)); + auto a20 = (kac * step); + auto d_a20 = ((d_kac * step) + (kac * d_step)); + auto a22 = (row_c * step); + auto d_a22 = ((d_row_c * step) + (row_c * d_step)); + auto third = bsk::truediv(((a00 + a11) + a22), Work(3.0)); + auto d_third = bsk::truediv(((d_a00 + d_a11) + d_a22), Work(3.0)); + auto s00 = (a00 - third); + auto d_s00 = (d_a00 - d_third); + auto s11 = (a11 - third); + auto d_s11 = (d_a11 - d_third); + auto s22 = (a22 - third); + auto d_s22 = (d_a22 - d_third); + auto minors = (((((s00 * s11) - (a01 * a10)) + (s00 * s22)) - (a02 * a20)) + (s11 * s22)); + auto d_minors = ((((((((((d_s00 * s11) + (s00 * d_s11)) - (d_a01 * a10)) - (a01 * d_a10)) + (d_s00 * s22)) + (s00 * d_s22)) - (d_a02 * a20)) - (a02 * d_a20)) + (d_s11 * s22)) + (s11 * d_s22)); + auto determinant = ((((s00 * s11) * s22) - (a01 * (a10 * s22))) + (a02 * ((-s11) * a20))); + auto d_determinant = ((((((((((d_s00 * s11) * s22) + ((s00 * d_s11) * s22)) + ((s00 * s11) * d_s22)) - ((d_a01 * a10) * s22)) - ((a01 * d_a10) * s22)) - ((a01 * a10) * d_s22)) - ((d_a02 * s11) * a20)) - ((a02 * d_s11) * a20)) - ((a02 * s11) * d_a20)); + // --- close together: the series reduced modulo x^3 + minors x - det --- + flat = (Work(1.0) + (Work(0.0) * third)); + linear = (Work(0.0) * third); + square = (Work(0.0) * third); + d_flat = (Work(0.0) * third); + d_linear = (Work(0.0) * third); + d_square = (Work(0.0) * third); + sum_flat = flat; + sum_linear = linear; + sum_square = square; + d_sum_flat = d_flat; + d_sum_linear = d_linear; + d_sum_square = d_square; + factorial = Work(1.0); + #pragma unroll + for (std::int64_t order = 1; order < 16; order += 1) { + auto next_flat = (square * determinant); + auto d_next_flat = ((d_square * determinant) + (square * d_determinant)); + auto next_linear = (flat - (square * minors)); + auto d_next_linear = ((d_flat - (d_square * minors)) - (square * d_minors)); + auto next_square = linear; + auto d_next_square = d_linear; + flat = next_flat; + linear = next_linear; + square = next_square; + d_flat = d_next_flat; + d_linear = d_next_linear; + d_square = d_next_square; + factorial = (factorial * order); + auto weight = bsk::truediv(Work(1.0), factorial); + sum_flat = (sum_flat + (weight * flat)); + sum_linear = (sum_linear + (weight * linear)); + sum_square = (sum_square + (weight * square)); + d_sum_flat = (d_sum_flat + (weight * d_flat)); + d_sum_linear = (d_sum_linear + (weight * d_linear)); + d_sum_square = (d_sum_square + (weight * d_square)); + } + auto lift = bsk::exp(third); + auto d_lift = (lift * d_third); + // --- far apart: the Newton form at the three roots --- + auto inside = ((-minors) * Work(0.3333333333333333)); + auto d_inside = ((-d_minors) * Work(0.3333333333333333)); + auto radius = bsk::sqrt(bsk::maximum(inside, Work(1e-300))); + auto d_radius = bsk::where((inside > Work(0.0)), bsk::truediv((Work(0.5) * d_inside), radius), Work(0.0)); + auto cube = ((radius * radius) * radius); + auto raw = bsk::truediv((Work(0.5) * determinant), cube); + auto d_raw = bsk::truediv(((Work(0.5) * d_determinant) - ((((raw * Work(3.0)) * radius) * radius) * d_radius)), cube); + auto inside_limit = bsk::band((raw > Work(-0.9999999999999999)), (raw < Work(0.9999999999999999))); + auto argument = bsk::minimum(bsk::maximum(raw, Work(-0.9999999999999999)), Work(0.9999999999999999)); + auto d_argument = bsk::where(inside_limit, d_raw, Work(0.0)); + auto angle = bsk::truediv(bsk::acos(argument), Work(3.0)); + auto d_angle = bsk::truediv((-d_argument), (Work(3.0) * bsk::sqrt(bsk::maximum((Work(1.0) - (argument * argument)), Work(1e-300))))); + auto root_a = (((Work(2.0) * radius) * bsk::cos(angle)) + third); + auto d_root_a = ((((Work(2.0) * d_radius) * bsk::cos(angle)) - (((Work(2.0) * radius) * bsk::sin(angle)) * d_angle)) + d_third); + auto root_b = (((Work(2.0) * radius) * bsk::cos((angle - Work(2.0943951023931957)))) + third); + auto d_root_b = ((((Work(2.0) * d_radius) * bsk::cos((angle - Work(2.0943951023931957)))) - (((Work(2.0) * radius) * bsk::sin((angle - Work(2.0943951023931957)))) * d_angle)) + d_third); + auto root_c = (((Work(2.0) * radius) * bsk::cos((angle - Work(4.188790204786391)))) + third); + auto d_root_c = ((((Work(2.0) * d_radius) * bsk::cos((angle - Work(4.188790204786391)))) - (((Work(2.0) * radius) * bsk::sin((angle - Work(4.188790204786391)))) * d_angle)) + d_third); + // Sorting is a permutation, so the tangents follow their own values. + auto low = bsk::minimum(bsk::minimum(root_a, root_b), root_c); + auto high = bsk::maximum(bsk::maximum(root_a, root_b), root_c); + auto middle = bsk::maximum(bsk::minimum(root_a, root_b), bsk::minimum(bsk::maximum(root_a, root_b), root_c)); + auto d_low = bsk::where((root_a == low), d_root_a, bsk::where((root_b == low), d_root_b, d_root_c)); + auto d_high = bsk::where((root_a == high), d_root_a, bsk::where((root_b == high), d_root_b, d_root_c)); + auto d_middle = bsk::where((root_a == middle), d_root_a, bsk::where((root_b == middle), d_root_b, d_root_c)); + auto leading = bsk::exp(low); + auto d_leading = (leading * d_low); + auto centre = bsk::exp(middle); + auto d_centre = (centre * d_middle); + auto trailing = bsk::exp(high); + auto d_trailing = (trailing * d_high); + auto t0_ = _exp_difference_jvp(low, d_low, middle, d_middle, leading, centre); + auto first = bsk::get<0>(t0_); + auto d_first = bsk::get<1>(t0_); + auto t1_ = _exp_difference_jvp(middle, d_middle, high, d_high, centre, trailing); + auto upper = bsk::get<0>(t1_); + auto d_upper = bsk::get<1>(t1_); + auto span = (high - low); + auto d_span = (d_high - d_low); + auto guarded = bsk::where((span > Work(0.0)), span, Work(1.0)); + auto d_guarded = bsk::where((span > Work(0.0)), d_span, Work(0.0)); + auto second = bsk::truediv((upper - first), guarded); + auto d_second = bsk::truediv(((d_upper - d_first) - (second * d_guarded)), guarded); + // --- the shifted generator squared, for the series branch --- + auto q00 = (((s00 * s00) + (a01 * a10)) + (a02 * a20)); + auto d_q00 = ((((((Work(2.0) * s00) * d_s00) + (d_a01 * a10)) + (a01 * d_a10)) + (d_a02 * a20)) + (a02 * d_a20)); + auto q01 = (a01 * (s00 + s11)); + auto d_q01 = ((d_a01 * (s00 + s11)) + (a01 * (d_s00 + d_s11))); + auto q02 = (a02 * (s00 + s22)); + auto d_q02 = ((d_a02 * (s00 + s22)) + (a02 * (d_s00 + d_s22))); + auto q10 = (a10 * (s00 + s11)); + auto d_q10 = ((d_a10 * (s00 + s11)) + (a10 * (d_s00 + d_s11))); + auto q11 = ((a10 * a01) + (s11 * s11)); + auto d_q11 = (((d_a10 * a01) + (a10 * d_a01)) + ((Work(2.0) * s11) * d_s11)); + auto q12 = (a10 * a02); + auto d_q12 = ((d_a10 * a02) + (a10 * d_a02)); + auto q20 = (a20 * (s00 + s22)); + auto d_q20 = ((d_a20 * (s00 + s22)) + (a20 * (d_s00 + d_s22))); + auto q21 = (a20 * a01); + auto d_q21 = ((d_a20 * a01) + (a20 * d_a01)); + auto q22 = ((a20 * a02) + (s22 * s22)); + auto d_q22 = (((d_a20 * a02) + (a20 * d_a02)) + ((Work(2.0) * s22) * d_s22)); + return bsk::make_tup(free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, a00, d_a00, a01, d_a01, a02, d_a02, a10, d_a10, a11, d_a11, a20, d_a20, a22, d_a22, s00, d_s00, s11, d_s11, s22, d_s22, minors, d_minors, sum_flat, sum_linear, sum_square, d_sum_flat, d_sum_linear, d_sum_square, lift, d_lift, low, middle, d_low, d_middle, leading, d_leading, first, d_first, second, d_second, determinant, d_determinant, high, d_high, radius, d_radius, cube, raw, d_raw, argument, inside_limit, angle, d_angle, centre, d_centre, trailing, d_trailing, guarded, d_guarded, q00, d_q00, q01, d_q01, q02, d_q02, q10, d_q10, q11, d_q11, q12, d_q12, q20, d_q20, q21, d_q21, q22, d_q22); +} + +template +BSK_HD auto _three_pool_pieces_jvp(const T0& r1_free, const T1& d_r1_free, const T2& r1_pool_b, const T3& d_r1_pool_b, const T4& r1_bound, const T5& d_r1_bound, const T6& exchange_b, const T7& d_exchange_b, const T8& exchange_c, const T9& d_exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& dt, const T15& d_dt, const T16& narrow) { + using R = decltype(_three_pool_pieces_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, narrow)); + if (bsk::truth(narrow)) { + return bsk::convert(_three_pool_pieces_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, narrow)); + } + return _three_pool_pieces_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, narrow); +} + +// The reverse of :func:`_exp_difference`, onto both points, on a direction. +// +// Near the coalescence the slope comes from the same series the value does, +// because the difference quotient's own derivative is a cancellation divided +// by a small number twice over. +template +BSK_HD auto _exp_difference_adjoint_jvp(const T0& lower, const T1& d_lower, const T2& upper, const T3& d_upper, const T4& exp_lower, const T5& d_exp_lower, const T6& exp_upper, const T7& d_exp_upper, const T8& seed, const T9& d_seed) { + auto half = (0.5f * (upper - lower)); + auto d_half = (0.5f * (d_upper - d_lower)); + auto near = (bsk::abs(half) < 0.0001f); + auto poly = ((1.0f + half) + ((0.5f * half) * half)); + auto d_poly = (d_half + (half * d_half)); + auto even = (1.0f + bsk::truediv((half * half), 6.0f)); + auto d_even = bsk::truediv((half * d_half), 3.0f); + auto slope = (((1.0f + half) * even) + ((poly * half) * 0.3333333333333333f)); + auto d_slope = (((d_half * even) + ((1.0f + half) * d_even)) + (((d_poly * half) + (poly * d_half)) * 0.3333333333333333f)); + auto series = ((exp_lower * poly) * even); + auto d_series = (((d_exp_lower * poly) * even) + (exp_lower * ((d_poly * even) + (poly * d_even)))); + auto swing = ((0.5f * exp_lower) * slope); + auto d_swing = (0.5f * ((d_exp_lower * slope) + (exp_lower * d_slope))); + auto gap = bsk::where(near, 1.0f, (upper - lower)); + auto d_gap = bsk::where(near, 0.0f, (d_upper - d_lower)); + auto value = bsk::truediv((exp_upper - exp_lower), gap); + auto d_value = bsk::truediv(((d_exp_upper - d_exp_lower) - (value * d_gap)), gap); + auto far_lower = bsk::truediv((value - exp_lower), gap); + auto d_far_lower = bsk::truediv(((d_value - d_exp_lower) - (far_lower * d_gap)), gap); + auto far_upper = bsk::truediv((exp_upper - value), gap); + auto d_far_upper = bsk::truediv(((d_exp_upper - d_value) - (far_upper * d_gap)), gap); + auto to_lower = bsk::where(near, (series - swing), far_lower); + auto d_to_lower = bsk::where(near, (d_series - d_swing), d_far_lower); + auto to_upper = bsk::where(near, swing, far_upper); + auto d_to_upper = bsk::where(near, d_swing, d_far_upper); + return bsk::make_tup((seed * to_lower), ((d_seed * to_lower) + (seed * d_to_lower)), (seed * to_upper), ((d_seed * to_upper) + (seed * d_to_upper))); +} + +// The reverse sweep of :func:`_three_pool_step`, carried on a direction. +// +// Reads the pieces and the bare operator the replay already formed, so an +// interval's transcendentals are taken once for the pass rather than once +// for each direction through it, and in the same double. +// +// Both branches are swept, each by the algebra its own forward used, and the +// choice between them is made on the cotangents rather than on the way in -- +// a ``where`` evaluates both sides, so each side's divisors are guarded. +// +// The series branch is a polynomial in the two invariants alone, so its +// reverse is reached by carrying the recurrence's sensitivity to those two +// forward beside it, which needs no history of the sixteen terms. +// +// Returned as the gradients w.r.t. ``(r1_free, r1_pool_b, r1_bound, +// exchange_b, exchange_c, fraction_b, fraction_c, dt, attenuation)`` and +// then their nine tangents. +template +BSK_HD auto _three_pool_step_adjoint_jvp_in_precision(const T0& r1_free, const T1& d_r1_free, const T2& r1_pool_b, const T3& d_r1_pool_b, const T4& r1_bound, const T5& d_r1_bound, const T6& exchange_b, const T7& d_exchange_b, const T8& exchange_c, const T9& d_exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& dt, const T15& d_dt, const T16& attenuation, const T17& d_attenuation, const T18& bar_e00, const T19& d_bar_e00, const T20& bar_e01, const T21& d_bar_e01, const T22& bar_e02, const T23& d_bar_e02, const T24& bar_e10, const T25& d_bar_e10, const T26& bar_e11, const T27& d_bar_e11, const T28& bar_e12, const T29& d_bar_e12, const T30& bar_e20, const T31& d_bar_e20, const T32& bar_e21, const T33& d_bar_e21, const T34& bar_e22, const T35& d_bar_e22, const T36& bar_grow_free, const T37& d_bar_grow_free, const T38& bar_grow_pool_b, const T39& d_bar_grow_pool_b, const T40& bar_grow_bound, const T41& d_bar_grow_bound, const T42& free, const T43& d_free, const T44& pool_b, const T45& d_pool_b, const T46& pool_c, const T47& d_pool_c, const T48& a00, const T49& d_a00, const T50& a01, const T51& d_a01, const T52& a02, const T53& d_a02, const T54& a10, const T55& d_a10, const T56& a11, const T57& d_a11, const T58& a20, const T59& d_a20, const T60& a22, const T61& d_a22, const T62& s00, const T63& d_s00, const T64& s11, const T65& d_s11, const T66& s22, const T67& d_s22, const T68& minors, const T69& d_minors, const T70& sum_flat, const T71& sum_linear, const T72& sum_square, const T73& d_sum_flat, const T74& d_sum_linear, const T75& d_sum_square, const T76& lift, const T77& d_lift, const T78& low, const T79& middle, const T80& d_low, const T81& d_middle, const T82& leading, const T83& d_leading, const T84& first, const T85& d_first, const T86& second, const T87& d_second, const T88& determinant, const T89& d_determinant, const T90& high, const T91& d_high, const T92& radius, const T93& d_radius, const T94& cube, const T95& raw, const T96& d_raw, const T97& argument, const T98& inside_limit, const T99& angle, const T100& d_angle, const T101& centre, const T102& d_centre, const T103& trailing, const T104& d_trailing, const T105& guarded, const T106& d_guarded, const T107& q00, const T108& d_q00, const T109& q01, const T110& d_q01, const T111& q02, const T112& d_q02, const T113& q10, const T114& d_q10, const T115& q11, const T116& d_q11, const T117& q12, const T118& d_q12, const T119& q20, const T120& d_q20, const T121& q21, const T122& d_q21, const T123& q22, const T124& d_q22, const T125& def_00, const T126& dif_00, const T127& def_01, const T128& dif_01, const T129& def_02, const T130& dif_02, const T131& def_10, const T132& dif_10, const T133& def_11, const T134& dif_11, const T135& def_12, const T136& dif_12, const T137& def_20, const T138& dif_20, const T139& def_21, const T140& dif_21, const T141& def_22, const T142& dif_22, const T143& narrow) { + bsk::tile_t | 0, 2)> bar_a00{}; + bsk::tile_t | 0, 2)> bar_a01{}; + bsk::tile_t | 0, 2)> bar_a02{}; + bsk::tile_t | 0, 2)> bar_a10{}; + bsk::tile_t | 0, 2)> bar_a11{}; + bsk::tile_t | 0, 2)> bar_a20{}; + bsk::tile_t | 0, 2)> bar_a22{}; + bsk::tile_t | 0, 2)> bar_determinant{}; + bsk::tile_t | 0, 2)> bar_first{}; + bsk::tile_t | 0, 2)> bar_high{}; + bsk::tile_t | 0, 2)> bar_low{}; + bsk::tile_t | 0, 2)> bar_middle{}; + bsk::tile_t | 0, 2)> bar_minors{}; + bsk::tile_t | 0, 2)> bar_radius{}; + bsk::tile_t | 0, 2)> bar_third{}; + bsk::tile_t | 0, 2)> d_bar_a00{}; + bsk::tile_t | 0, 2)> d_bar_a01{}; + bsk::tile_t | 0, 2)> d_bar_a02{}; + bsk::tile_t | 0, 2)> d_bar_a10{}; + bsk::tile_t | 0, 2)> d_bar_a11{}; + bsk::tile_t | 0, 2)> d_bar_a20{}; + bsk::tile_t | 0, 2)> d_bar_a22{}; + bsk::tile_t | 0, 2)> d_bar_determinant{}; + bsk::tile_t | 0, 2)> d_bar_first{}; + bsk::tile_t | 0, 2)> d_bar_high{}; + bsk::tile_t | 0, 2)> d_bar_low{}; + bsk::tile_t | 0, 2)> d_bar_middle{}; + bsk::tile_t | 0, 2)> d_bar_minors{}; + bsk::tile_t | 0, 2)> d_bar_radius{}; + bsk::tile_t | 0, 2)> d_bar_third{}; + bsk::tile_t | 0, 2)> d_flat{}; + bsk::tile_t | 0, 2)> d_fu{}; + bsk::tile_t | 0, 2)> d_fv{}; + bsk::tile_t | 0, 2)> d_linear{}; + bsk::tile_t | 0, 2)> d_lu{}; + bsk::tile_t | 0, 2)> d_lv{}; + bsk::tile_t | 0, 2)> d_slope_u_flat{}; + bsk::tile_t | 0, 2)> d_slope_u_linear{}; + bsk::tile_t | 0, 2)> d_slope_u_square{}; + bsk::tile_t | 0, 2)> d_slope_v_flat{}; + bsk::tile_t | 0, 2)> d_slope_v_linear{}; + bsk::tile_t | 0, 2)> d_slope_v_square{}; + bsk::tile_t | 0, 2)> d_square{}; + bsk::tile_t | 0, 2)> d_su{}; + bsk::tile_t | 0, 2)> d_sv{}; + bsk::tile_t | 0, 2)> d_turn_series{}; + float factorial{}; + bsk::tile_t | 0, 2)> flat{}; + bsk::tile_t | 0, 2)> fu{}; + bsk::tile_t | 0, 2)> fv{}; + bsk::tile_t | 0, 2)> linear{}; + bsk::tile_t | 0, 2)> lu{}; + bsk::tile_t | 0, 2)> lv{}; + bsk::tile_t | 0, 2)> slope_u_flat{}; + bsk::tile_t | 0, 2)> slope_u_linear{}; + bsk::tile_t | 0, 2)> slope_u_square{}; + bsk::tile_t | 0, 2)> slope_v_flat{}; + bsk::tile_t | 0, 2)> slope_v_linear{}; + bsk::tile_t | 0, 2)> slope_v_square{}; + bsk::tile_t | 0, 2)> square{}; + bsk::tile_t | 0, 2)> su{}; + bsk::tile_t | 0, 2)> sv{}; + bsk::tile_t | 0, 2)> turn_series{}; + // --- the recovery and the attenuation, which both branches share --- + auto damp = bsk::cast(attenuation); + auto d_damp = bsk::cast(d_attenuation); + auto r0 = bsk::cast(bar_grow_free); + auto d_r0 = bsk::cast(d_bar_grow_free); + auto r1 = bsk::cast(bar_grow_pool_b); + auto d_r1 = bsk::cast(d_bar_grow_pool_b); + auto r2 = bsk::cast(bar_grow_bound); + auto d_r2 = bsk::cast(d_bar_grow_bound); + auto y00 = (bsk::cast(bar_e00) - (r0 * free)); + auto d_y00 = ((bsk::cast(d_bar_e00) - (d_r0 * free)) - (r0 * d_free)); + auto y01 = (bsk::cast(bar_e01) - (r0 * pool_b)); + auto d_y01 = ((bsk::cast(d_bar_e01) - (d_r0 * pool_b)) - (r0 * d_pool_b)); + auto y02 = (bsk::cast(bar_e02) - (r0 * pool_c)); + auto d_y02 = ((bsk::cast(d_bar_e02) - (d_r0 * pool_c)) - (r0 * d_pool_c)); + auto y10 = (bsk::cast(bar_e10) - (r1 * free)); + auto d_y10 = ((bsk::cast(d_bar_e10) - (d_r1 * free)) - (r1 * d_free)); + auto y11 = (bsk::cast(bar_e11) - (r1 * pool_b)); + auto d_y11 = ((bsk::cast(d_bar_e11) - (d_r1 * pool_b)) - (r1 * d_pool_b)); + auto y12 = (bsk::cast(bar_e12) - (r1 * pool_c)); + auto d_y12 = ((bsk::cast(d_bar_e12) - (d_r1 * pool_c)) - (r1 * d_pool_c)); + auto y20 = (bsk::cast(bar_e20) - (r2 * free)); + auto d_y20 = ((bsk::cast(d_bar_e20) - (d_r2 * free)) - (r2 * d_free)); + auto y21 = (bsk::cast(bar_e21) - (r2 * pool_b)); + auto d_y21 = ((bsk::cast(d_bar_e21) - (d_r2 * pool_b)) - (r2 * d_pool_b)); + auto y22 = (bsk::cast(bar_e22) - (r2 * pool_c)); + auto d_y22 = ((bsk::cast(d_bar_e22) - (d_r2 * pool_c)) - (r2 * d_pool_c)); + // The bare operator is read rather than the attenuation divided back out + // of the weighed one -- a washed-out interval leaves nothing to divide by. + auto bar_damp = (((((((((y00 * def_00) + (y01 * def_01)) + (y02 * def_02)) + (y10 * def_10)) + (y11 * def_11)) + (y12 * def_12)) + (y20 * def_20)) + (y21 * def_21)) + (y22 * def_22)); + auto d_bar_damp = ((((((((((((((((((d_y00 * def_00) + (y00 * dif_00)) + (d_y01 * def_01)) + (y01 * dif_01)) + (d_y02 * def_02)) + (y02 * dif_02)) + (d_y10 * def_10)) + (y10 * dif_10)) + (d_y11 * def_11)) + (y11 * dif_11)) + (d_y12 * def_12)) + (y12 * dif_12)) + (d_y20 * def_20)) + (y20 * dif_20)) + (d_y21 * def_21)) + (y21 * dif_21)) + (d_y22 * def_22)) + (y22 * dif_22)); + auto column_free = (((r0 * def_00) + (r1 * def_10)) + (r2 * def_20)); + auto d_column_free = ((((((d_r0 * def_00) + (r0 * dif_00)) + (d_r1 * def_10)) + (r1 * dif_10)) + (d_r2 * def_20)) + (r2 * dif_20)); + auto column_pool_b = (((r0 * def_01) + (r1 * def_11)) + (r2 * def_21)); + auto d_column_pool_b = ((((((d_r0 * def_01) + (r0 * dif_01)) + (d_r1 * def_11)) + (r1 * dif_11)) + (d_r2 * def_21)) + (r2 * dif_21)); + auto column_bound = (((r0 * def_02) + (r1 * def_12)) + (r2 * def_22)); + auto d_column_bound = ((((((d_r0 * def_02) + (r0 * dif_02)) + (d_r1 * def_12)) + (r1 * dif_12)) + (d_r2 * def_22)) + (r2 * dif_22)); + auto bar_free = (r0 - (damp * column_free)); + auto d_bar_free = ((d_r0 - (d_damp * column_free)) - (damp * d_column_free)); + auto bar_pool_b = (r1 - (damp * column_pool_b)); + auto d_bar_pool_b = ((d_r1 - (d_damp * column_pool_b)) - (damp * d_column_pool_b)); + auto bar_pool_c = (r2 - (damp * column_bound)); + auto d_bar_pool_c = ((d_r2 - (d_damp * column_bound)) - (damp * d_column_bound)); + auto o00 = (damp * y00); + auto d_o00 = ((d_damp * y00) + (damp * d_y00)); + auto o01 = (damp * y01); + auto d_o01 = ((d_damp * y01) + (damp * d_y01)); + auto o02 = (damp * y02); + auto d_o02 = ((d_damp * y02) + (damp * d_y02)); + auto o10 = (damp * y10); + auto d_o10 = ((d_damp * y10) + (damp * d_y10)); + auto o11 = (damp * y11); + auto d_o11 = ((d_damp * y11) + (damp * d_y11)); + auto o12 = (damp * y12); + auto d_o12 = ((d_damp * y12) + (damp * d_y12)); + auto o20 = (damp * y20); + auto d_o20 = ((d_damp * y20) + (damp * d_y20)); + auto o21 = (damp * y21); + auto d_o21 = ((d_damp * y21) + (damp * d_y21)); + auto o22 = (damp * y22); + auto d_o22 = ((d_damp * y22) + (damp * d_y22)); + // --- close together: the series in the two invariants, run backwards --- + auto scale00 = (o00 * lift); + auto d_scale00 = ((d_o00 * lift) + (o00 * d_lift)); + auto scale01 = (o01 * lift); + auto d_scale01 = ((d_o01 * lift) + (o01 * d_lift)); + auto scale02 = (o02 * lift); + auto d_scale02 = ((d_o02 * lift) + (o02 * d_lift)); + auto scale10 = (o10 * lift); + auto d_scale10 = ((d_o10 * lift) + (o10 * d_lift)); + auto scale11 = (o11 * lift); + auto d_scale11 = ((d_o11 * lift) + (o11 * d_lift)); + auto scale12 = (o12 * lift); + auto d_scale12 = ((d_o12 * lift) + (o12 * d_lift)); + auto scale20 = (o20 * lift); + auto d_scale20 = ((d_o20 * lift) + (o20 * d_lift)); + auto scale21 = (o21 * lift); + auto d_scale21 = ((d_o21 * lift) + (o21 * d_lift)); + auto scale22 = (o22 * lift); + auto d_scale22 = ((d_o22 * lift) + (o22 * d_lift)); + auto bar_flat = ((scale00 + scale11) + scale22); + auto d_bar_flat = ((d_scale00 + d_scale11) + d_scale22); + auto bar_linear = (((((((scale00 * s00) + (scale01 * a01)) + (scale02 * a02)) + (scale10 * a10)) + (scale11 * s11)) + (scale20 * a20)) + (scale22 * s22)); + auto d_bar_linear = ((((((((((((((d_scale00 * s00) + (scale00 * d_s00)) + (d_scale01 * a01)) + (scale01 * d_a01)) + (d_scale02 * a02)) + (scale02 * d_a02)) + (d_scale10 * a10)) + (scale10 * d_a10)) + (d_scale11 * s11)) + (scale11 * d_s11)) + (d_scale20 * a20)) + (scale20 * d_a20)) + (d_scale22 * s22)) + (scale22 * d_s22)); + auto bar_square = (((((((((scale00 * q00) + (scale01 * q01)) + (scale02 * q02)) + (scale10 * q10)) + (scale11 * q11)) + (scale12 * q12)) + (scale20 * q20)) + (scale21 * q21)) + (scale22 * q22)); + auto d_bar_square = ((((((((((((((((((d_scale00 * q00) + (scale00 * d_q00)) + (d_scale01 * q01)) + (scale01 * d_q01)) + (d_scale02 * q02)) + (scale02 * d_q02)) + (d_scale10 * q10)) + (scale10 * d_q10)) + (d_scale11 * q11)) + (scale11 * d_q11)) + (d_scale12 * q12)) + (scale12 * d_q12)) + (d_scale20 * q20)) + (scale20 * d_q20)) + (d_scale21 * q21)) + (scale21 * d_q21)) + (d_scale22 * q22)) + (scale22 * d_q22)); + // ``lift`` multiplies the whole bracket, so the shift it carries picks up + // the bracket back again -- which is what the three sums contract to. + turn_series = (((sum_flat * bar_flat) + (sum_linear * bar_linear)) + (sum_square * bar_square)); + d_turn_series = ((((((d_sum_flat * bar_flat) + (sum_flat * d_bar_flat)) + (d_sum_linear * bar_linear)) + (sum_linear * d_bar_linear)) + (d_sum_square * bar_square)) + (sum_square * d_bar_square)); + auto g00 = (sum_square * scale00); + auto d_g00 = ((d_sum_square * scale00) + (sum_square * d_scale00)); + auto g01 = (sum_square * scale01); + auto d_g01 = ((d_sum_square * scale01) + (sum_square * d_scale01)); + auto g02 = (sum_square * scale02); + auto d_g02 = ((d_sum_square * scale02) + (sum_square * d_scale02)); + auto g10 = (sum_square * scale10); + auto d_g10 = ((d_sum_square * scale10) + (sum_square * d_scale10)); + auto g11 = (sum_square * scale11); + auto d_g11 = ((d_sum_square * scale11) + (sum_square * d_scale11)); + auto g12 = (sum_square * scale12); + auto d_g12 = ((d_sum_square * scale12) + (sum_square * d_scale12)); + auto g20 = (sum_square * scale20); + auto d_g20 = ((d_sum_square * scale20) + (sum_square * d_scale20)); + auto g21 = (sum_square * scale21); + auto d_g21 = ((d_sum_square * scale21) + (sum_square * d_scale21)); + auto g22 = (sum_square * scale22); + auto d_g22 = ((d_sum_square * scale22) + (sum_square * d_scale22)); + // The square's reverse, ``g @ shifted^T + shifted^T @ g``. + auto v00 = (((((((g00 * s00) + (g01 * a01)) + (g02 * a02)) + (s00 * g00)) + (a10 * g10)) + (a20 * g20)) + (sum_linear * scale00)); + auto d_v00 = ((((((((((((((d_g00 * s00) + (g00 * d_s00)) + (d_g01 * a01)) + (g01 * d_a01)) + (d_g02 * a02)) + (g02 * d_a02)) + (d_s00 * g00)) + (s00 * d_g00)) + (d_a10 * g10)) + (a10 * d_g10)) + (d_a20 * g20)) + (a20 * d_g20)) + (d_sum_linear * scale00)) + (sum_linear * d_scale00)); + auto v01 = ((((((g00 * a10) + (g01 * s11)) + (s00 * g01)) + (a10 * g11)) + (a20 * g21)) + (sum_linear * scale01)); + auto d_v01 = ((((((((((((d_g00 * a10) + (g00 * d_a10)) + (d_g01 * s11)) + (g01 * d_s11)) + (d_s00 * g01)) + (s00 * d_g01)) + (d_a10 * g11)) + (a10 * d_g11)) + (d_a20 * g21)) + (a20 * d_g21)) + (d_sum_linear * scale01)) + (sum_linear * d_scale01)); + auto v02 = ((((((g00 * a20) + (g02 * s22)) + (s00 * g02)) + (a10 * g12)) + (a20 * g22)) + (sum_linear * scale02)); + auto d_v02 = ((((((((((((d_g00 * a20) + (g00 * d_a20)) + (d_g02 * s22)) + (g02 * d_s22)) + (d_s00 * g02)) + (s00 * d_g02)) + (d_a10 * g12)) + (a10 * d_g12)) + (d_a20 * g22)) + (a20 * d_g22)) + (d_sum_linear * scale02)) + (sum_linear * d_scale02)); + auto v10 = ((((((g10 * s00) + (g11 * a01)) + (g12 * a02)) + (a01 * g00)) + (s11 * g10)) + (sum_linear * scale10)); + auto d_v10 = ((((((((((((d_g10 * s00) + (g10 * d_s00)) + (d_g11 * a01)) + (g11 * d_a01)) + (d_g12 * a02)) + (g12 * d_a02)) + (d_a01 * g00)) + (a01 * d_g00)) + (d_s11 * g10)) + (s11 * d_g10)) + (d_sum_linear * scale10)) + (sum_linear * d_scale10)); + auto v11 = (((((g10 * a10) + (g11 * s11)) + (a01 * g01)) + (s11 * g11)) + (sum_linear * scale11)); + auto d_v11 = ((((((((((d_g10 * a10) + (g10 * d_a10)) + (d_g11 * s11)) + (g11 * d_s11)) + (d_a01 * g01)) + (a01 * d_g01)) + (d_s11 * g11)) + (s11 * d_g11)) + (d_sum_linear * scale11)) + (sum_linear * d_scale11)); + auto v20 = ((((((g20 * s00) + (g21 * a01)) + (g22 * a02)) + (a02 * g00)) + (s22 * g20)) + (sum_linear * scale20)); + auto d_v20 = ((((((((((((d_g20 * s00) + (g20 * d_s00)) + (d_g21 * a01)) + (g21 * d_a01)) + (d_g22 * a02)) + (g22 * d_a02)) + (d_a02 * g00)) + (a02 * d_g00)) + (d_s22 * g20)) + (s22 * d_g20)) + (d_sum_linear * scale20)) + (sum_linear * d_scale20)); + auto v22 = (((((g20 * a20) + (g22 * s22)) + (a02 * g02)) + (s22 * g22)) + (sum_linear * scale22)); + auto d_v22 = ((((((((((d_g20 * a20) + (g20 * d_a20)) + (d_g22 * s22)) + (g22 * d_s22)) + (d_a02 * g02)) + (a02 * d_g02)) + (d_s22 * g22)) + (s22 * d_g22)) + (d_sum_linear * scale22)) + (sum_linear * d_scale22)); + turn_series = (turn_series - ((v00 + v11) + v22)); + d_turn_series = (d_turn_series - ((d_v00 + d_v11) + d_v22)); + // The recurrence's own sensitivity to the two invariants, carried forward + // beside it: two numbers reach the whole series, so their derivatives are + // cheaper to push forward than the sixteen terms are to keep. + flat = (Work(1.0) + (Work(0.0) * a00)); + linear = (Work(0.0) * a00); + square = (Work(0.0) * a00); + d_flat = (Work(0.0) * a00); + d_linear = (Work(0.0) * a00); + d_square = (Work(0.0) * a00); + fu = (Work(0.0) * a00); + lu = (Work(0.0) * a00); + su = (Work(0.0) * a00); + d_fu = (Work(0.0) * a00); + d_lu = (Work(0.0) * a00); + d_su = (Work(0.0) * a00); + fv = (Work(0.0) * a00); + lv = (Work(0.0) * a00); + sv = (Work(0.0) * a00); + d_fv = (Work(0.0) * a00); + d_lv = (Work(0.0) * a00); + d_sv = (Work(0.0) * a00); + slope_u_flat = (Work(0.0) * a00); + slope_u_linear = (Work(0.0) * a00); + slope_u_square = (Work(0.0) * a00); + d_slope_u_flat = (Work(0.0) * a00); + d_slope_u_linear = (Work(0.0) * a00); + d_slope_u_square = (Work(0.0) * a00); + slope_v_flat = (Work(0.0) * a00); + slope_v_linear = (Work(0.0) * a00); + slope_v_square = (Work(0.0) * a00); + d_slope_v_flat = (Work(0.0) * a00); + d_slope_v_linear = (Work(0.0) * a00); + d_slope_v_square = (Work(0.0) * a00); + factorial = Work(1.0); + #pragma unroll + for (std::int64_t order = 1; order < 16; order += 1) { + auto next_flat = (square * determinant); + auto d_next_flat = ((d_square * determinant) + (square * d_determinant)); + auto next_linear = (flat - (square * minors)); + auto d_next_linear = ((d_flat - (d_square * minors)) - (square * d_minors)); + auto next_square = linear; + auto d_next_square = d_linear; + auto next_fu = (su * determinant); + auto d_next_fu = ((d_su * determinant) + (su * d_determinant)); + auto next_lu = ((fu - (su * minors)) - square); + auto d_next_lu = (((d_fu - (d_su * minors)) - (su * d_minors)) - d_square); + auto next_su = lu; + auto d_next_su = d_lu; + auto next_fv = ((sv * determinant) + square); + auto d_next_fv = (((d_sv * determinant) + (sv * d_determinant)) + d_square); + auto next_lv = (fv - (sv * minors)); + auto d_next_lv = ((d_fv - (d_sv * minors)) - (sv * d_minors)); + auto next_sv = lv; + auto d_next_sv = d_lv; + flat = next_flat; + linear = next_linear; + square = next_square; + d_flat = d_next_flat; + d_linear = d_next_linear; + d_square = d_next_square; + fu = next_fu; + lu = next_lu; + su = next_su; + d_fu = d_next_fu; + d_lu = d_next_lu; + d_su = d_next_su; + fv = next_fv; + lv = next_lv; + sv = next_sv; + d_fv = d_next_fv; + d_lv = d_next_lv; + d_sv = d_next_sv; + factorial = (factorial * order); + auto weight = bsk::truediv(Work(1.0), factorial); + slope_u_flat = (slope_u_flat + (weight * fu)); + slope_u_linear = (slope_u_linear + (weight * lu)); + slope_u_square = (slope_u_square + (weight * su)); + d_slope_u_flat = (d_slope_u_flat + (weight * d_fu)); + d_slope_u_linear = (d_slope_u_linear + (weight * d_lu)); + d_slope_u_square = (d_slope_u_square + (weight * d_su)); + slope_v_flat = (slope_v_flat + (weight * fv)); + slope_v_linear = (slope_v_linear + (weight * lv)); + slope_v_square = (slope_v_square + (weight * sv)); + d_slope_v_flat = (d_slope_v_flat + (weight * d_fv)); + d_slope_v_linear = (d_slope_v_linear + (weight * d_lv)); + d_slope_v_square = (d_slope_v_square + (weight * d_sv)); + } + auto minors_series = (((bar_flat * slope_u_flat) + (bar_linear * slope_u_linear)) + (bar_square * slope_u_square)); + auto d_minors_series = ((((((d_bar_flat * slope_u_flat) + (bar_flat * d_slope_u_flat)) + (d_bar_linear * slope_u_linear)) + (bar_linear * d_slope_u_linear)) + (d_bar_square * slope_u_square)) + (bar_square * d_slope_u_square)); + auto determinant_series = (((bar_flat * slope_v_flat) + (bar_linear * slope_v_linear)) + (bar_square * slope_v_square)); + auto d_determinant_series = ((((((d_bar_flat * slope_v_flat) + (bar_flat * d_slope_v_flat)) + (d_bar_linear * slope_v_linear)) + (bar_linear * d_slope_v_linear)) + (d_bar_square * slope_v_square)) + (bar_square * d_slope_v_square)); + // --- far apart: back through the Newton form and the three roots --- + auto m00 = (a00 - low); + auto d_m00 = (d_a00 - d_low); + auto m11 = (a11 - low); + auto d_m11 = (d_a11 - d_low); + auto m22 = (a22 - low); + auto d_m22 = (d_a22 - d_low); + auto n00 = (a00 - middle); + auto d_n00 = (d_a00 - d_middle); + auto n11 = (a11 - middle); + auto d_n11 = (d_a11 - d_middle); + auto n22 = (a22 - middle); + auto d_n22 = (d_a22 - d_middle); + auto p00 = (((m00 * n00) + (a01 * a10)) + (a02 * a20)); + auto d_p00 = ((((((d_m00 * n00) + (m00 * d_n00)) + (d_a01 * a10)) + (a01 * d_a10)) + (d_a02 * a20)) + (a02 * d_a20)); + auto p01 = (a01 * (m00 + n11)); + auto d_p01 = ((d_a01 * (m00 + n11)) + (a01 * (d_m00 + d_n11))); + auto p02 = (a02 * (m00 + n22)); + auto d_p02 = ((d_a02 * (m00 + n22)) + (a02 * (d_m00 + d_n22))); + auto p10 = (a10 * (m11 + n00)); + auto d_p10 = ((d_a10 * (m11 + n00)) + (a10 * (d_m11 + d_n00))); + auto p11 = ((m11 * n11) + (a01 * a10)); + auto d_p11 = ((((d_m11 * n11) + (m11 * d_n11)) + (d_a01 * a10)) + (a01 * d_a10)); + auto p12 = (a10 * a02); + auto d_p12 = ((d_a10 * a02) + (a10 * d_a02)); + auto p20 = (a20 * (m22 + n00)); + auto d_p20 = ((d_a20 * (m22 + n00)) + (a20 * (d_m22 + d_n00))); + auto p21 = (a20 * a01); + auto d_p21 = ((d_a20 * a01) + (a20 * d_a01)); + auto p22 = ((m22 * n22) + (a02 * a20)); + auto d_p22 = ((((d_m22 * n22) + (m22 * d_n22)) + (d_a02 * a20)) + (a02 * d_a20)); + auto bar_leading = ((o00 + o11) + o22); + auto d_bar_leading = ((d_o00 + d_o11) + d_o22); + bar_first = (((((((o00 * m00) + (o01 * a01)) + (o02 * a02)) + (o10 * a10)) + (o11 * m11)) + (o20 * a20)) + (o22 * m22)); + d_bar_first = ((((((((((((((d_o00 * m00) + (o00 * d_m00)) + (d_o01 * a01)) + (o01 * d_a01)) + (d_o02 * a02)) + (o02 * d_a02)) + (d_o10 * a10)) + (o10 * d_a10)) + (d_o11 * m11)) + (o11 * d_m11)) + (d_o20 * a20)) + (o20 * d_a20)) + (d_o22 * m22)) + (o22 * d_m22)); + auto bar_second = (((((((((o00 * p00) + (o01 * p01)) + (o02 * p02)) + (o10 * p10)) + (o11 * p11)) + (o12 * p12)) + (o20 * p20)) + (o21 * p21)) + (o22 * p22)); + auto d_bar_second = ((((((((((((((((((d_o00 * p00) + (o00 * d_p00)) + (d_o01 * p01)) + (o01 * d_p01)) + (d_o02 * p02)) + (o02 * d_p02)) + (d_o10 * p10)) + (o10 * d_p10)) + (d_o11 * p11)) + (o11 * d_p11)) + (d_o12 * p12)) + (o12 * d_p12)) + (d_o20 * p20)) + (o20 * d_p20)) + (d_o21 * p21)) + (o21 * d_p21)) + (d_o22 * p22)) + (o22 * d_p22)); + auto z00 = (second * o00); + auto d_z00 = ((d_second * o00) + (second * d_o00)); + auto z01 = (second * o01); + auto d_z01 = ((d_second * o01) + (second * d_o01)); + auto z02 = (second * o02); + auto d_z02 = ((d_second * o02) + (second * d_o02)); + auto z10 = (second * o10); + auto d_z10 = ((d_second * o10) + (second * d_o10)); + auto z11 = (second * o11); + auto d_z11 = ((d_second * o11) + (second * d_o11)); + auto z12 = (second * o12); + auto d_z12 = ((d_second * o12) + (second * d_o12)); + auto z20 = (second * o20); + auto d_z20 = ((d_second * o20) + (second * d_o20)); + auto z21 = (second * o21); + auto d_z21 = ((d_second * o21) + (second * d_o21)); + auto z22 = (second * o22); + auto d_z22 = ((d_second * o22) + (second * d_o22)); + // ``z @ n^T``, the product's reverse onto the first factor. + auto u00 = (((z00 * n00) + (z01 * a01)) + (z02 * a02)); + auto d_u00 = ((((((d_z00 * n00) + (z00 * d_n00)) + (d_z01 * a01)) + (z01 * d_a01)) + (d_z02 * a02)) + (z02 * d_a02)); + auto u01 = ((z00 * a10) + (z01 * n11)); + auto d_u01 = ((((d_z00 * a10) + (z00 * d_a10)) + (d_z01 * n11)) + (z01 * d_n11)); + auto u02 = ((z00 * a20) + (z02 * n22)); + auto d_u02 = ((((d_z00 * a20) + (z00 * d_a20)) + (d_z02 * n22)) + (z02 * d_n22)); + auto u10 = (((z10 * n00) + (z11 * a01)) + (z12 * a02)); + auto d_u10 = ((((((d_z10 * n00) + (z10 * d_n00)) + (d_z11 * a01)) + (z11 * d_a01)) + (d_z12 * a02)) + (z12 * d_a02)); + auto u11 = ((z10 * a10) + (z11 * n11)); + auto d_u11 = ((((d_z10 * a10) + (z10 * d_a10)) + (d_z11 * n11)) + (z11 * d_n11)); + auto u20 = (((z20 * n00) + (z21 * a01)) + (z22 * a02)); + auto d_u20 = ((((((d_z20 * n00) + (z20 * d_n00)) + (d_z21 * a01)) + (z21 * d_a01)) + (d_z22 * a02)) + (z22 * d_a02)); + auto u22 = ((z20 * a20) + (z22 * n22)); + auto d_u22 = ((((d_z20 * a20) + (z20 * d_a20)) + (d_z22 * n22)) + (z22 * d_n22)); + // ``m^T @ z``, onto the second. + auto w00 = (((m00 * z00) + (a10 * z10)) + (a20 * z20)); + auto d_w00 = ((((((d_m00 * z00) + (m00 * d_z00)) + (d_a10 * z10)) + (a10 * d_z10)) + (d_a20 * z20)) + (a20 * d_z20)); + auto w01 = (((m00 * z01) + (a10 * z11)) + (a20 * z21)); + auto d_w01 = ((((((d_m00 * z01) + (m00 * d_z01)) + (d_a10 * z11)) + (a10 * d_z11)) + (d_a20 * z21)) + (a20 * d_z21)); + auto w02 = (((m00 * z02) + (a10 * z12)) + (a20 * z22)); + auto d_w02 = ((((((d_m00 * z02) + (m00 * d_z02)) + (d_a10 * z12)) + (a10 * d_z12)) + (d_a20 * z22)) + (a20 * d_z22)); + auto w10 = ((a01 * z00) + (m11 * z10)); + auto d_w10 = ((((d_a01 * z00) + (a01 * d_z00)) + (d_m11 * z10)) + (m11 * d_z10)); + auto w11 = ((a01 * z01) + (m11 * z11)); + auto d_w11 = ((((d_a01 * z01) + (a01 * d_z01)) + (d_m11 * z11)) + (m11 * d_z11)); + auto w20 = ((a02 * z00) + (m22 * z20)); + auto d_w20 = ((((d_a02 * z00) + (a02 * d_z00)) + (d_m22 * z20)) + (m22 * d_z20)); + auto w22 = ((a02 * z02) + (m22 * z22)); + auto d_w22 = ((((d_a02 * z02) + (a02 * d_z02)) + (d_m22 * z22)) + (m22 * d_z22)); + bar_low = (((bar_leading * leading) - (first * ((o00 + o11) + o22))) - ((u00 + u11) + u22)); + d_bar_low = (((((d_bar_leading * leading) + (bar_leading * d_leading)) - (d_first * ((o00 + o11) + o22))) - (first * ((d_o00 + d_o11) + d_o22))) - ((d_u00 + d_u11) + d_u22)); + bar_middle = (-((w00 + w11) + w22)); + d_bar_middle = (-((d_w00 + d_w11) + d_w22)); + bar_high = (Work(0.0) * a00); + d_bar_high = (Work(0.0) * a00); + auto span = (high - low); + auto positive = (span > Work(0.0)); + auto bar_upper = bsk::truediv(bar_second, guarded); + auto d_bar_upper = bsk::truediv((d_bar_second - (bar_upper * d_guarded)), guarded); + bar_first = (bar_first - bar_upper); + d_bar_first = (d_bar_first - d_bar_upper); + auto bar_span = bsk::where(positive, ((-bar_upper) * second), Work(0.0)); + auto d_bar_span = bsk::where(positive, (((-d_bar_upper) * second) - (bar_upper * d_second)), Work(0.0)); + bar_high = (bar_high + bar_span); + d_bar_high = (d_bar_high + d_bar_span); + bar_low = (bar_low - bar_span); + d_bar_low = (d_bar_low - d_bar_span); + auto t0_ = _exp_difference_adjoint_jvp(low, d_low, middle, d_middle, leading, d_leading, centre, d_centre, bar_first, d_bar_first); + auto from_first_low = bsk::get<0>(t0_); + auto d_from_first_low = bsk::get<1>(t0_); + auto from_first_middle = bsk::get<2>(t0_); + auto d_from_first_middle = bsk::get<3>(t0_); + auto t1_ = _exp_difference_adjoint_jvp(middle, d_middle, high, d_high, centre, d_centre, trailing, d_trailing, bar_upper, d_bar_upper); + auto from_upper_middle = bsk::get<0>(t1_); + auto d_from_upper_middle = bsk::get<1>(t1_); + auto from_upper_high = bsk::get<2>(t1_); + auto d_from_upper_high = bsk::get<3>(t1_); + bar_low = (bar_low + from_first_low); + d_bar_low = (d_bar_low + d_from_first_low); + bar_middle = ((bar_middle + from_first_middle) + from_upper_middle); + d_bar_middle = ((d_bar_middle + d_from_first_middle) + d_from_upper_middle); + bar_high = (bar_high + from_upper_high); + d_bar_high = (d_bar_high + d_from_upper_high); + // The three roots come off one angle a third of a turn apart, and the + // cosine puts them in a fixed order: the last turn is the lowest, the + // first the highest, whatever the angle is. + auto swing_low = (angle - Work(4.188790204786391)); + auto swing_middle = (angle - Work(2.0943951023931957)); + auto cos_low = bsk::cos(swing_low); + auto cos_middle = bsk::cos(swing_middle); + auto cos_high = bsk::cos(angle); + auto sin_low = bsk::sin(swing_low); + auto sin_middle = bsk::sin(swing_middle); + auto sin_high = bsk::sin(angle); + bar_radius = (Work(2.0) * (((cos_low * bar_low) + (cos_middle * bar_middle)) + (cos_high * bar_high))); + d_bar_radius = (Work(2.0) * ((((cos_low * d_bar_low) + (cos_middle * d_bar_middle)) + (cos_high * d_bar_high)) - (d_angle * (((sin_low * bar_low) + (sin_middle * bar_middle)) + (sin_high * bar_high))))); + auto swept = (((sin_low * bar_low) + (sin_middle * bar_middle)) + (sin_high * bar_high)); + auto d_swept = ((((sin_low * d_bar_low) + (sin_middle * d_bar_middle)) + (sin_high * d_bar_high)) + (d_angle * (((cos_low * bar_low) + (cos_middle * bar_middle)) + (cos_high * bar_high)))); + auto bar_angle = ((Work(-2.0) * radius) * swept); + auto d_bar_angle = (Work(-2.0) * ((d_radius * swept) + (radius * d_swept))); + auto turn_roots = ((bar_low + bar_middle) + bar_high); + auto d_turn_roots = ((d_bar_low + d_bar_middle) + d_bar_high); + // ``acos`` is clamped, and where it is the angle no longer moves with the + // cubic's argument -- which is what keeps a double root differentiable. + auto d_argument = bsk::where(inside_limit, d_raw, Work(0.0)); + auto inner = (Work(1.0) - (argument * argument)); + auto d_inner = ((Work(-2.0) * argument) * d_argument); + auto stem = bsk::sqrt(bsk::maximum(inner, Work(1e-300))); + auto tilt = bsk::truediv(Work(-1.0), (Work(3.0) * stem)); + auto d_tilt = bsk::truediv(d_inner, (((Work(6.0) * stem) * stem) * stem)); + auto bar_raw = bsk::where(inside_limit, (bar_angle * tilt), Work(0.0)); + auto d_bar_raw = bsk::where(inside_limit, ((d_bar_angle * tilt) + (bar_angle * d_tilt)), Work(0.0)); + auto safe_radius = bsk::where((radius > Work(1e-30)), radius, Work(1.0)); + auto d_safe_radius = bsk::where((radius > Work(1e-30)), d_radius, Work(0.0)); + auto safe_cube = ((safe_radius * safe_radius) * safe_radius); + auto d_safe_cube = (((Work(3.0) * safe_radius) * safe_radius) * d_safe_radius); + auto determinant_roots = bsk::truediv((Work(0.5) * bar_raw), safe_cube); + auto d_determinant_roots = bsk::truediv(((Work(0.5) * d_bar_raw) - (determinant_roots * d_safe_cube)), safe_cube); + auto pull = bsk::where(inside_limit, bsk::truediv(((Work(-3.0) * raw) * bar_raw), safe_radius), Work(0.0)); + auto d_pull = bsk::where(inside_limit, bsk::truediv(((Work(-3.0) * ((d_raw * bar_raw) + (raw * d_bar_raw))) - (pull * d_safe_radius)), safe_radius), Work(0.0)); + bar_radius = (bar_radius + pull); + d_bar_radius = (d_bar_radius + d_pull); + auto minors_roots = bsk::truediv((-bar_radius), (Work(6.0) * safe_radius)); + auto d_minors_roots = bsk::truediv(((-d_bar_radius) - ((minors_roots * Work(6.0)) * d_safe_radius)), (Work(6.0) * safe_radius)); + // --- the branch chosen on the cotangents, not on the way in --- + if (bsk::truth(narrow)) { + auto t2_ = bsk::make_tup(v00, d_v00); + bar_a00 = bsk::get<0>(t2_); + d_bar_a00 = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup(v01, d_v01); + bar_a01 = bsk::get<0>(t3_); + d_bar_a01 = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(v02, d_v02); + bar_a02 = bsk::get<0>(t4_); + d_bar_a02 = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup(v10, d_v10); + bar_a10 = bsk::get<0>(t5_); + d_bar_a10 = bsk::get<1>(t5_); + auto t6_ = bsk::make_tup(v11, d_v11); + bar_a11 = bsk::get<0>(t6_); + d_bar_a11 = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(v20, d_v20); + bar_a20 = bsk::get<0>(t7_); + d_bar_a20 = bsk::get<1>(t7_); + auto t8_ = bsk::make_tup(v22, d_v22); + bar_a22 = bsk::get<0>(t8_); + d_bar_a22 = bsk::get<1>(t8_); + auto t9_ = bsk::make_tup(turn_series, d_turn_series); + bar_third = bsk::get<0>(t9_); + d_bar_third = bsk::get<1>(t9_); + auto t10_ = bsk::make_tup(minors_series, d_minors_series); + bar_minors = bsk::get<0>(t10_); + d_bar_minors = bsk::get<1>(t10_); + auto t11_ = bsk::make_tup(determinant_series, d_determinant_series); + bar_determinant = bsk::get<0>(t11_); + d_bar_determinant = bsk::get<1>(t11_); + } else { + auto close = ((Work(-2.0) * minors) < Work(1.0)); + bar_a00 = bsk::where(close, v00, (((first * o00) + u00) + w00)); + d_bar_a00 = bsk::where(close, d_v00, ((((d_first * o00) + (first * d_o00)) + d_u00) + d_w00)); + bar_a01 = bsk::where(close, v01, (((first * o01) + u01) + w01)); + d_bar_a01 = bsk::where(close, d_v01, ((((d_first * o01) + (first * d_o01)) + d_u01) + d_w01)); + bar_a02 = bsk::where(close, v02, (((first * o02) + u02) + w02)); + d_bar_a02 = bsk::where(close, d_v02, ((((d_first * o02) + (first * d_o02)) + d_u02) + d_w02)); + bar_a10 = bsk::where(close, v10, (((first * o10) + u10) + w10)); + d_bar_a10 = bsk::where(close, d_v10, ((((d_first * o10) + (first * d_o10)) + d_u10) + d_w10)); + bar_a11 = bsk::where(close, v11, (((first * o11) + u11) + w11)); + d_bar_a11 = bsk::where(close, d_v11, ((((d_first * o11) + (first * d_o11)) + d_u11) + d_w11)); + bar_a20 = bsk::where(close, v20, (((first * o20) + u20) + w20)); + d_bar_a20 = bsk::where(close, d_v20, ((((d_first * o20) + (first * d_o20)) + d_u20) + d_w20)); + bar_a22 = bsk::where(close, v22, (((first * o22) + u22) + w22)); + d_bar_a22 = bsk::where(close, d_v22, ((((d_first * o22) + (first * d_o22)) + d_u22) + d_w22)); + bar_third = bsk::where(close, turn_series, turn_roots); + d_bar_third = bsk::where(close, d_turn_series, d_turn_roots); + bar_minors = bsk::where(close, minors_series, minors_roots); + d_bar_minors = bsk::where(close, d_minors_series, d_minors_roots); + bar_determinant = bsk::where(close, determinant_series, determinant_roots); + d_bar_determinant = bsk::where(close, d_determinant_series, d_determinant_roots); + } + // --- the two invariants back onto the shifted generator --- + auto cofactor00 = (s11 * s22); + auto d_cofactor00 = ((d_s11 * s22) + (s11 * d_s22)); + auto cofactor11 = ((s00 * s22) - (a02 * a20)); + auto d_cofactor11 = ((((d_s00 * s22) + (s00 * d_s22)) - (d_a02 * a20)) - (a02 * d_a20)); + auto cofactor22 = ((s00 * s11) - (a01 * a10)); + auto d_cofactor22 = ((((d_s00 * s11) + (s00 * d_s11)) - (d_a01 * a10)) - (a01 * d_a10)); + auto shift00 = ((bar_minors * (s11 + s22)) + (bar_determinant * cofactor00)); + auto d_shift00 = ((((d_bar_minors * (s11 + s22)) + (bar_minors * (d_s11 + d_s22))) + (d_bar_determinant * cofactor00)) + (bar_determinant * d_cofactor00)); + auto shift11 = ((bar_minors * (s00 + s22)) + (bar_determinant * cofactor11)); + auto d_shift11 = ((((d_bar_minors * (s00 + s22)) + (bar_minors * (d_s00 + d_s22))) + (d_bar_determinant * cofactor11)) + (bar_determinant * d_cofactor11)); + auto shift22 = ((bar_minors * (s00 + s11)) + (bar_determinant * cofactor22)); + auto d_shift22 = ((((d_bar_minors * (s00 + s11)) + (bar_minors * (d_s00 + d_s11))) + (d_bar_determinant * cofactor22)) + (bar_determinant * d_cofactor22)); + auto shift01 = ((-a10) * (bar_minors + (bar_determinant * s22))); + auto d_shift01 = (((-d_a10) * (bar_minors + (bar_determinant * s22))) - (a10 * ((d_bar_minors + (d_bar_determinant * s22)) + (bar_determinant * d_s22)))); + auto shift10 = ((-a01) * (bar_minors + (bar_determinant * s22))); + auto d_shift10 = (((-d_a01) * (bar_minors + (bar_determinant * s22))) - (a01 * ((d_bar_minors + (d_bar_determinant * s22)) + (bar_determinant * d_s22)))); + auto shift02 = ((-a20) * (bar_minors + (bar_determinant * s11))); + auto d_shift02 = (((-d_a20) * (bar_minors + (bar_determinant * s11))) - (a20 * ((d_bar_minors + (d_bar_determinant * s11)) + (bar_determinant * d_s11)))); + auto shift20 = ((-a02) * (bar_minors + (bar_determinant * s11))); + auto d_shift20 = (((-d_a02) * (bar_minors + (bar_determinant * s11))) - (a02 * ((d_bar_minors + (d_bar_determinant * s11)) + (bar_determinant * d_s11)))); + bar_a00 = (bar_a00 + shift00); + d_bar_a00 = (d_bar_a00 + d_shift00); + bar_a01 = (bar_a01 + shift01); + d_bar_a01 = (d_bar_a01 + d_shift01); + bar_a02 = (bar_a02 + shift02); + d_bar_a02 = (d_bar_a02 + d_shift02); + bar_a10 = (bar_a10 + shift10); + d_bar_a10 = (d_bar_a10 + d_shift10); + bar_a11 = (bar_a11 + shift11); + d_bar_a11 = (d_bar_a11 + d_shift11); + bar_a20 = (bar_a20 + shift20); + d_bar_a20 = (d_bar_a20 + d_shift20); + bar_a22 = (bar_a22 + shift22); + d_bar_a22 = (d_bar_a22 + d_shift22); + bar_third = (bar_third - ((shift00 + shift11) + shift22)); + d_bar_third = (d_bar_third - ((d_shift00 + d_shift11) + d_shift22)); + bar_a00 = (bar_a00 + bsk::truediv(bar_third, Work(3.0))); + d_bar_a00 = (d_bar_a00 + bsk::truediv(d_bar_third, Work(3.0))); + bar_a11 = (bar_a11 + bsk::truediv(bar_third, Work(3.0))); + d_bar_a11 = (d_bar_a11 + bsk::truediv(d_bar_third, Work(3.0))); + bar_a22 = (bar_a22 + bsk::truediv(bar_third, Work(3.0))); + d_bar_a22 = (d_bar_a22 + bsk::truediv(d_bar_third, Work(3.0))); + // --- the generator back onto the rates, the fractions and the interval --- + auto step = bsk::cast(dt); + auto d_step = bsk::cast(d_dt); + auto rate_b = bsk::cast(exchange_b); + auto d_rate_b = bsk::cast(d_exchange_b); + auto rate_c = bsk::cast(exchange_c); + auto d_rate_c = bsk::cast(d_exchange_c); + auto kab = (rate_b * pool_b); + auto d_kab = ((d_rate_b * pool_b) + (rate_b * d_pool_b)); + auto kba = (rate_b * free); + auto d_kba = ((d_rate_b * free) + (rate_b * d_free)); + auto kac = (rate_c * pool_c); + auto d_kac = ((d_rate_c * pool_c) + (rate_c * d_pool_c)); + auto kca = (rate_c * free); + auto d_kca = ((d_rate_c * free) + (rate_c * d_free)); + auto row_a = (((-kab) - kac) - bsk::cast(r1_free)); + auto d_row_a = (((-d_kab) - d_kac) - bsk::cast(d_r1_free)); + auto row_b = ((-kba) - bsk::cast(r1_pool_b)); + auto d_row_b = ((-d_kba) - bsk::cast(d_r1_pool_b)); + auto row_c = ((-kca) - bsk::cast(r1_bound)); + auto d_row_c = ((-d_kca) - bsk::cast(d_r1_bound)); + auto bar_step = (((((((row_a * bar_a00) + (kba * bar_a01)) + (kca * bar_a02)) + (kab * bar_a10)) + (row_b * bar_a11)) + (kac * bar_a20)) + (row_c * bar_a22)); + auto d_bar_step = ((((((((((((((d_row_a * bar_a00) + (row_a * d_bar_a00)) + (d_kba * bar_a01)) + (kba * d_bar_a01)) + (d_kca * bar_a02)) + (kca * d_bar_a02)) + (d_kab * bar_a10)) + (kab * d_bar_a10)) + (d_row_b * bar_a11)) + (row_b * d_bar_a11)) + (d_kac * bar_a20)) + (kac * d_bar_a20)) + (d_row_c * bar_a22)) + (row_c * d_bar_a22)); + auto bar_kab = (step * (bar_a10 - bar_a00)); + auto d_bar_kab = ((d_step * (bar_a10 - bar_a00)) + (step * (d_bar_a10 - d_bar_a00))); + auto bar_kba = (step * (bar_a01 - bar_a11)); + auto d_bar_kba = ((d_step * (bar_a01 - bar_a11)) + (step * (d_bar_a01 - d_bar_a11))); + auto bar_kac = (step * (bar_a20 - bar_a00)); + auto d_bar_kac = ((d_step * (bar_a20 - bar_a00)) + (step * (d_bar_a20 - d_bar_a00))); + auto bar_kca = (step * (bar_a02 - bar_a22)); + auto d_bar_kca = ((d_step * (bar_a02 - bar_a22)) + (step * (d_bar_a02 - d_bar_a22))); + auto whole_free = ((bar_free + (rate_b * bar_kba)) + (rate_c * bar_kca)); + auto d_whole_free = ((((d_bar_free + (d_rate_b * bar_kba)) + (rate_b * d_bar_kba)) + (d_rate_c * bar_kca)) + (rate_c * d_bar_kca)); + auto whole_pool_b = (bar_pool_b + (rate_b * bar_kab)); + auto d_whole_pool_b = ((d_bar_pool_b + (d_rate_b * bar_kab)) + (rate_b * d_bar_kab)); + auto whole_pool_c = (bar_pool_c + (rate_c * bar_kac)); + auto d_whole_pool_c = ((d_bar_pool_c + (d_rate_c * bar_kac)) + (rate_c * d_bar_kac)); + return bsk::make_tup(bsk::cast(((-step) * bar_a00)), bsk::cast(((-step) * bar_a11)), bsk::cast(((-step) * bar_a22)), bsk::cast(((pool_b * bar_kab) + (free * bar_kba))), bsk::cast(((pool_c * bar_kac) + (free * bar_kca))), bsk::cast((whole_pool_b - whole_free)), bsk::cast((whole_pool_c - whole_free)), bsk::cast(bar_step), bsk::cast(bar_damp), bsk::cast((((-d_step) * bar_a00) - (step * d_bar_a00))), bsk::cast((((-d_step) * bar_a11) - (step * d_bar_a11))), bsk::cast((((-d_step) * bar_a22) - (step * d_bar_a22))), bsk::cast(((((d_pool_b * bar_kab) + (pool_b * d_bar_kab)) + (d_free * bar_kba)) + (free * d_bar_kba))), bsk::cast(((((d_pool_c * bar_kac) + (pool_c * d_bar_kac)) + (d_free * bar_kca)) + (free * d_bar_kca))), bsk::cast((d_whole_pool_b - d_whole_free)), bsk::cast((d_whole_pool_c - d_whole_free)), bsk::cast(d_bar_step), bsk::cast(d_bar_damp)); +} + +template +BSK_HD auto _three_pool_step_adjoint_jvp(const T0& r1_free, const T1& d_r1_free, const T2& r1_pool_b, const T3& d_r1_pool_b, const T4& r1_bound, const T5& d_r1_bound, const T6& exchange_b, const T7& d_exchange_b, const T8& exchange_c, const T9& d_exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& dt, const T15& d_dt, const T16& attenuation, const T17& d_attenuation, const T18& bar_e00, const T19& d_bar_e00, const T20& bar_e01, const T21& d_bar_e01, const T22& bar_e02, const T23& d_bar_e02, const T24& bar_e10, const T25& d_bar_e10, const T26& bar_e11, const T27& d_bar_e11, const T28& bar_e12, const T29& d_bar_e12, const T30& bar_e20, const T31& d_bar_e20, const T32& bar_e21, const T33& d_bar_e21, const T34& bar_e22, const T35& d_bar_e22, const T36& bar_grow_free, const T37& d_bar_grow_free, const T38& bar_grow_pool_b, const T39& d_bar_grow_pool_b, const T40& bar_grow_bound, const T41& d_bar_grow_bound, const T42& free, const T43& d_free, const T44& pool_b, const T45& d_pool_b, const T46& pool_c, const T47& d_pool_c, const T48& a00, const T49& d_a00, const T50& a01, const T51& d_a01, const T52& a02, const T53& d_a02, const T54& a10, const T55& d_a10, const T56& a11, const T57& d_a11, const T58& a20, const T59& d_a20, const T60& a22, const T61& d_a22, const T62& s00, const T63& d_s00, const T64& s11, const T65& d_s11, const T66& s22, const T67& d_s22, const T68& minors, const T69& d_minors, const T70& sum_flat, const T71& sum_linear, const T72& sum_square, const T73& d_sum_flat, const T74& d_sum_linear, const T75& d_sum_square, const T76& lift, const T77& d_lift, const T78& low, const T79& middle, const T80& d_low, const T81& d_middle, const T82& leading, const T83& d_leading, const T84& first, const T85& d_first, const T86& second, const T87& d_second, const T88& determinant, const T89& d_determinant, const T90& high, const T91& d_high, const T92& radius, const T93& d_radius, const T94& cube, const T95& raw, const T96& d_raw, const T97& argument, const T98& inside_limit, const T99& angle, const T100& d_angle, const T101& centre, const T102& d_centre, const T103& trailing, const T104& d_trailing, const T105& guarded, const T106& d_guarded, const T107& q00, const T108& d_q00, const T109& q01, const T110& d_q01, const T111& q02, const T112& d_q02, const T113& q10, const T114& d_q10, const T115& q11, const T116& d_q11, const T117& q12, const T118& d_q12, const T119& q20, const T120& d_q20, const T121& q21, const T122& d_q21, const T123& q22, const T124& d_q22, const T125& def_00, const T126& dif_00, const T127& def_01, const T128& dif_01, const T129& def_02, const T130& dif_02, const T131& def_10, const T132& dif_10, const T133& def_11, const T134& dif_11, const T135& def_12, const T136& dif_12, const T137& def_20, const T138& dif_20, const T139& def_21, const T140& dif_21, const T141& def_22, const T142& dif_22, const T143& narrow) { + using R = decltype(_three_pool_step_adjoint_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, attenuation, d_attenuation, bar_e00, d_bar_e00, bar_e01, d_bar_e01, bar_e02, d_bar_e02, bar_e10, d_bar_e10, bar_e11, d_bar_e11, bar_e12, d_bar_e12, bar_e20, d_bar_e20, bar_e21, d_bar_e21, bar_e22, d_bar_e22, bar_grow_free, d_bar_grow_free, bar_grow_pool_b, d_bar_grow_pool_b, bar_grow_bound, d_bar_grow_bound, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, a00, d_a00, a01, d_a01, a02, d_a02, a10, d_a10, a11, d_a11, a20, d_a20, a22, d_a22, s00, d_s00, s11, d_s11, s22, d_s22, minors, d_minors, sum_flat, sum_linear, sum_square, d_sum_flat, d_sum_linear, d_sum_square, lift, d_lift, low, middle, d_low, d_middle, leading, d_leading, first, d_first, second, d_second, determinant, d_determinant, high, d_high, radius, d_radius, cube, raw, d_raw, argument, inside_limit, angle, d_angle, centre, d_centre, trailing, d_trailing, guarded, d_guarded, q00, d_q00, q01, d_q01, q02, d_q02, q10, d_q10, q11, d_q11, q12, d_q12, q20, d_q20, q21, d_q21, q22, d_q22, def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, narrow)); + if (bsk::truth(narrow)) { + return bsk::convert(_three_pool_step_adjoint_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, attenuation, d_attenuation, bar_e00, d_bar_e00, bar_e01, d_bar_e01, bar_e02, d_bar_e02, bar_e10, d_bar_e10, bar_e11, d_bar_e11, bar_e12, d_bar_e12, bar_e20, d_bar_e20, bar_e21, d_bar_e21, bar_e22, d_bar_e22, bar_grow_free, d_bar_grow_free, bar_grow_pool_b, d_bar_grow_pool_b, bar_grow_bound, d_bar_grow_bound, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, a00, d_a00, a01, d_a01, a02, d_a02, a10, d_a10, a11, d_a11, a20, d_a20, a22, d_a22, s00, d_s00, s11, d_s11, s22, d_s22, minors, d_minors, sum_flat, sum_linear, sum_square, d_sum_flat, d_sum_linear, d_sum_square, lift, d_lift, low, middle, d_low, d_middle, leading, d_leading, first, d_first, second, d_second, determinant, d_determinant, high, d_high, radius, d_radius, cube, raw, d_raw, argument, inside_limit, angle, d_angle, centre, d_centre, trailing, d_trailing, guarded, d_guarded, q00, d_q00, q01, d_q01, q02, d_q02, q10, d_q10, q11, d_q11, q12, d_q12, q20, d_q20, q21, d_q21, q22, d_q22, def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, narrow)); + } + return _three_pool_step_adjoint_jvp_in_precision(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, attenuation, d_attenuation, bar_e00, d_bar_e00, bar_e01, d_bar_e01, bar_e02, d_bar_e02, bar_e10, d_bar_e10, bar_e11, d_bar_e11, bar_e12, d_bar_e12, bar_e20, d_bar_e20, bar_e21, d_bar_e21, bar_e22, d_bar_e22, bar_grow_free, d_bar_grow_free, bar_grow_pool_b, d_bar_grow_pool_b, bar_grow_bound, d_bar_grow_bound, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, a00, d_a00, a01, d_a01, a02, d_a02, a10, d_a10, a11, d_a11, a20, d_a20, a22, d_a22, s00, d_s00, s11, d_s11, s22, d_s22, minors, d_minors, sum_flat, sum_linear, sum_square, d_sum_flat, d_sum_linear, d_sum_square, lift, d_lift, low, middle, d_low, d_middle, leading, d_leading, first, d_first, second, d_second, determinant, d_determinant, high, d_high, radius, d_radius, cube, raw, d_raw, argument, inside_limit, angle, d_angle, centre, d_centre, trailing, d_trailing, guarded, d_guarded, q00, d_q00, q01, d_q01, q02, d_q02, q10, d_q10, q11, d_q11, q12, d_q12, q20, d_q20, q21, d_q21, q22, d_q22, def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, narrow); +} + +// An interval's step, from the bare operator and what survives it. +// +// The recovery is ``(I - E) m0`` rather than a solve, which is what the +// equilibrium being a fixed point of the generator buys. +template +BSK_HD auto _three_pool_weigh_jvp_in_precision(const T0& def_00, const T1& dif_00, const T2& def_01, const T3& dif_01, const T4& def_02, const T5& dif_02, const T6& def_10, const T7& dif_10, const T8& def_11, const T9& dif_11, const T10& def_12, const T11& dif_12, const T12& def_20, const T13& dif_20, const T14& def_21, const T15& dif_21, const T16& def_22, const T17& dif_22, const T18& free, const T19& d_free, const T20& pool_b, const T21& d_pool_b, const T22& pool_c, const T23& d_pool_c, const T24& attenuation, const T25& d_attenuation, const T26& narrow) { + auto damp = bsk::cast(attenuation); + auto d_damp = bsk::cast(d_attenuation); + auto w00 = (damp * def_00); + auto dw00 = ((d_damp * def_00) + (damp * dif_00)); + auto w01 = (damp * def_01); + auto dw01 = ((d_damp * def_01) + (damp * dif_01)); + auto w02 = (damp * def_02); + auto dw02 = ((d_damp * def_02) + (damp * dif_02)); + auto w10 = (damp * def_10); + auto dw10 = ((d_damp * def_10) + (damp * dif_10)); + auto w11 = (damp * def_11); + auto dw11 = ((d_damp * def_11) + (damp * dif_11)); + auto w12 = (damp * def_12); + auto dw12 = ((d_damp * def_12) + (damp * dif_12)); + auto w20 = (damp * def_20); + auto dw20 = ((d_damp * def_20) + (damp * dif_20)); + auto w21 = (damp * def_21); + auto dw21 = ((d_damp * def_21) + (damp * dif_21)); + auto w22 = (damp * def_22); + auto dw22 = ((d_damp * def_22) + (damp * dif_22)); + auto grow_free = (free - (((w00 * free) + (w01 * pool_b)) + (w02 * pool_c))); + auto d_grow_free = (d_free - ((((((dw00 * free) + (w00 * d_free)) + (dw01 * pool_b)) + (w01 * d_pool_b)) + (dw02 * pool_c)) + (w02 * d_pool_c))); + auto grow_pool_b = (pool_b - (((w10 * free) + (w11 * pool_b)) + (w12 * pool_c))); + auto d_grow_pool_b = (d_pool_b - ((((((dw10 * free) + (w10 * d_free)) + (dw11 * pool_b)) + (w11 * d_pool_b)) + (dw12 * pool_c)) + (w12 * d_pool_c))); + auto grow_bound = (pool_c - (((w20 * free) + (w21 * pool_b)) + (w22 * pool_c))); + auto d_grow_bound = (d_pool_c - ((((((dw20 * free) + (w20 * d_free)) + (dw21 * pool_b)) + (w21 * d_pool_b)) + (dw22 * pool_c)) + (w22 * d_pool_c))); + return bsk::make_tup(w00, w01, w02, w10, w11, w12, w20, w21, w22, grow_free, grow_pool_b, grow_bound, dw00, dw01, dw02, dw10, dw11, dw12, dw20, dw21, dw22, d_grow_free, d_grow_pool_b, d_grow_bound); +} + +template +BSK_HD auto _three_pool_weigh_jvp(const T0& def_00, const T1& dif_00, const T2& def_01, const T3& dif_01, const T4& def_02, const T5& dif_02, const T6& def_10, const T7& dif_10, const T8& def_11, const T9& dif_11, const T10& def_12, const T11& dif_12, const T12& def_20, const T13& dif_20, const T14& def_21, const T15& dif_21, const T16& def_22, const T17& dif_22, const T18& free, const T19& d_free, const T20& pool_b, const T21& d_pool_b, const T22& pool_c, const T23& d_pool_c, const T24& attenuation, const T25& d_attenuation, const T26& narrow) { + using R = decltype(_three_pool_weigh_jvp_in_precision(def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, attenuation, d_attenuation, narrow)); + if (bsk::truth(narrow)) { + return bsk::convert(_three_pool_weigh_jvp_in_precision(def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, attenuation, d_attenuation, narrow)); + } + return _three_pool_weigh_jvp_in_precision(def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, attenuation, d_attenuation, narrow); +} + +// The three-pool longitudinal step and its directional derivative. +// +// The same closed form :func:`_three_pool_step` evaluates, carried +// alongside a tangent and in the same double precision. Returns the +// nine entries and three recoveries, then their twelve tangents. +template +BSK_HD auto _three_pool_step_jvp(const T0& r1_free, const T1& d_r1_free, const T2& r1_pool_b, const T3& d_r1_pool_b, const T4& r1_bound, const T5& d_r1_bound, const T6& exchange_b, const T7& d_exchange_b, const T8& exchange_c, const T9& d_exchange_c, const T10& fraction_b, const T11& d_fraction_b, const T12& fraction_c, const T13& d_fraction_c, const T14& dt, const T15& d_dt, const T16& attenuation, const T17& d_attenuation, const T18& narrow) { + auto t0_ = _three_pool_pieces_jvp(r1_free, d_r1_free, r1_pool_b, d_r1_pool_b, r1_bound, d_r1_bound, exchange_b, d_exchange_b, exchange_c, d_exchange_c, fraction_b, d_fraction_b, fraction_c, d_fraction_c, dt, d_dt, narrow); + auto free = bsk::get<0>(t0_); + auto d_free = bsk::get<1>(t0_); + auto pool_b = bsk::get<2>(t0_); + auto d_pool_b = bsk::get<3>(t0_); + auto pool_c = bsk::get<4>(t0_); + auto d_pool_c = bsk::get<5>(t0_); + auto a00 = bsk::get<6>(t0_); + auto d_a00 = bsk::get<7>(t0_); + auto a01 = bsk::get<8>(t0_); + auto d_a01 = bsk::get<9>(t0_); + auto a02 = bsk::get<10>(t0_); + auto d_a02 = bsk::get<11>(t0_); + auto a10 = bsk::get<12>(t0_); + auto d_a10 = bsk::get<13>(t0_); + auto a11 = bsk::get<14>(t0_); + auto d_a11 = bsk::get<15>(t0_); + auto a20 = bsk::get<16>(t0_); + auto d_a20 = bsk::get<17>(t0_); + auto a22 = bsk::get<18>(t0_); + auto d_a22 = bsk::get<19>(t0_); + auto s00 = bsk::get<20>(t0_); + auto d_s00 = bsk::get<21>(t0_); + auto s11 = bsk::get<22>(t0_); + auto d_s11 = bsk::get<23>(t0_); + auto s22 = bsk::get<24>(t0_); + auto d_s22 = bsk::get<25>(t0_); + auto minors = bsk::get<26>(t0_); + auto d_minors = bsk::get<27>(t0_); + auto sum_flat = bsk::get<28>(t0_); + auto sum_linear = bsk::get<29>(t0_); + auto sum_square = bsk::get<30>(t0_); + auto d_sum_flat = bsk::get<31>(t0_); + auto d_sum_linear = bsk::get<32>(t0_); + auto d_sum_square = bsk::get<33>(t0_); + auto lift = bsk::get<34>(t0_); + auto d_lift = bsk::get<35>(t0_); + auto low = bsk::get<36>(t0_); + auto middle = bsk::get<37>(t0_); + auto d_low = bsk::get<38>(t0_); + auto d_middle = bsk::get<39>(t0_); + auto leading = bsk::get<40>(t0_); + auto d_leading = bsk::get<41>(t0_); + auto first = bsk::get<42>(t0_); + auto d_first = bsk::get<43>(t0_); + auto second = bsk::get<44>(t0_); + auto d_second = bsk::get<45>(t0_); + auto determinant = bsk::get<46>(t0_); + auto d_determinant = bsk::get<47>(t0_); + auto high = bsk::get<48>(t0_); + auto d_high = bsk::get<49>(t0_); + auto radius = bsk::get<50>(t0_); + auto d_radius = bsk::get<51>(t0_); + auto cube = bsk::get<52>(t0_); + auto raw = bsk::get<53>(t0_); + auto d_raw = bsk::get<54>(t0_); + auto argument = bsk::get<55>(t0_); + auto inside_limit = bsk::get<56>(t0_); + auto angle = bsk::get<57>(t0_); + auto d_angle = bsk::get<58>(t0_); + auto centre = bsk::get<59>(t0_); + auto d_centre = bsk::get<60>(t0_); + auto trailing = bsk::get<61>(t0_); + auto d_trailing = bsk::get<62>(t0_); + auto guarded = bsk::get<63>(t0_); + auto d_guarded = bsk::get<64>(t0_); + auto q00 = bsk::get<65>(t0_); + auto d_q00 = bsk::get<66>(t0_); + auto q01 = bsk::get<67>(t0_); + auto d_q01 = bsk::get<68>(t0_); + auto q02 = bsk::get<69>(t0_); + auto d_q02 = bsk::get<70>(t0_); + auto q10 = bsk::get<71>(t0_); + auto d_q10 = bsk::get<72>(t0_); + auto q11 = bsk::get<73>(t0_); + auto d_q11 = bsk::get<74>(t0_); + auto q12 = bsk::get<75>(t0_); + auto d_q12 = bsk::get<76>(t0_); + auto q20 = bsk::get<77>(t0_); + auto d_q20 = bsk::get<78>(t0_); + auto q21 = bsk::get<79>(t0_); + auto d_q21 = bsk::get<80>(t0_); + auto q22 = bsk::get<81>(t0_); + auto d_q22 = bsk::get<82>(t0_); + auto t1_ = _three_pool_assemble_jvp(free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, a00, d_a00, a01, d_a01, a02, d_a02, a10, d_a10, a11, d_a11, a20, d_a20, a22, d_a22, s00, d_s00, s11, d_s11, s22, d_s22, minors, d_minors, sum_flat, sum_linear, sum_square, d_sum_flat, d_sum_linear, d_sum_square, lift, d_lift, low, middle, d_low, d_middle, leading, d_leading, first, d_first, second, d_second, determinant, d_determinant, high, d_high, radius, d_radius, cube, raw, d_raw, argument, inside_limit, angle, d_angle, centre, d_centre, trailing, d_trailing, guarded, d_guarded, q00, d_q00, q01, d_q01, q02, d_q02, q10, d_q10, q11, d_q11, q12, d_q12, q20, d_q20, q21, d_q21, q22, d_q22, narrow); + auto def_00 = bsk::get<0>(t1_); + auto dif_00 = bsk::get<1>(t1_); + auto def_01 = bsk::get<2>(t1_); + auto dif_01 = bsk::get<3>(t1_); + auto def_02 = bsk::get<4>(t1_); + auto dif_02 = bsk::get<5>(t1_); + auto def_10 = bsk::get<6>(t1_); + auto dif_10 = bsk::get<7>(t1_); + auto def_11 = bsk::get<8>(t1_); + auto dif_11 = bsk::get<9>(t1_); + auto def_12 = bsk::get<10>(t1_); + auto dif_12 = bsk::get<11>(t1_); + auto def_20 = bsk::get<12>(t1_); + auto dif_20 = bsk::get<13>(t1_); + auto def_21 = bsk::get<14>(t1_); + auto dif_21 = bsk::get<15>(t1_); + auto def_22 = bsk::get<16>(t1_); + auto dif_22 = bsk::get<17>(t1_); + auto t2_ = _three_pool_weigh_jvp(def_00, dif_00, def_01, dif_01, def_02, dif_02, def_10, dif_10, def_11, dif_11, def_12, dif_12, def_20, dif_20, def_21, dif_21, def_22, dif_22, free, d_free, pool_b, d_pool_b, pool_c, d_pool_c, attenuation, d_attenuation, narrow); + auto w00 = bsk::get<0>(t2_); + auto w01 = bsk::get<1>(t2_); + auto w02 = bsk::get<2>(t2_); + auto w10 = bsk::get<3>(t2_); + auto w11 = bsk::get<4>(t2_); + auto w12 = bsk::get<5>(t2_); + auto w20 = bsk::get<6>(t2_); + auto w21 = bsk::get<7>(t2_); + auto w22 = bsk::get<8>(t2_); + auto grow_free = bsk::get<9>(t2_); + auto grow_pool_b = bsk::get<10>(t2_); + auto grow_bound = bsk::get<11>(t2_); + auto dw00 = bsk::get<12>(t2_); + auto dw01 = bsk::get<13>(t2_); + auto dw02 = bsk::get<14>(t2_); + auto dw10 = bsk::get<15>(t2_); + auto dw11 = bsk::get<16>(t2_); + auto dw12 = bsk::get<17>(t2_); + auto dw20 = bsk::get<18>(t2_); + auto dw21 = bsk::get<19>(t2_); + auto dw22 = bsk::get<20>(t2_); + auto d_grow_free = bsk::get<21>(t2_); + auto d_grow_pool_b = bsk::get<22>(t2_); + auto d_grow_bound = bsk::get<23>(t2_); + return bsk::make_tup(bsk::cast(w00), bsk::cast(w01), bsk::cast(w02), bsk::cast(w10), bsk::cast(w11), bsk::cast(w12), bsk::cast(w20), bsk::cast(w21), bsk::cast(w22), bsk::cast(grow_free), bsk::cast(grow_pool_b), bsk::cast(grow_bound), bsk::cast(dw00), bsk::cast(dw01), bsk::cast(dw02), bsk::cast(dw10), bsk::cast(dw11), bsk::cast(dw12), bsk::cast(dw20), bsk::cast(dw21), bsk::cast(dw22), bsk::cast(d_grow_free), bsk::cast(d_grow_pool_b), bsk::cast(d_grow_bound)); +} + +// The reverse sweep of :func:`_two_pool_step`, carried on a direction. +// +// Recomputes the forward rather than carrying it across the event: the whole +// thing is a handful of transcendentals once per interval, against a state +// loop that runs per dephasing order. +// +// Where the discriminant is small the value is still formed from the two +// eigenvalues -- a sum, which loses nothing -- but the derivative is taken +// from the series, because ``d cosh(d)/d(d^2)`` reached through +// ``(e^{t+d} - e^{t-d})/2d`` is a cancellation divided by a small number. +// +// Returned as the gradients w.r.t. ``(r1_free, r1_bound, exchange, bound, +// dt, attenuation)`` then their six tangents. +template +BSK_HD auto _two_pool_step_adjoint_jvp(const T0& r1_free, const T1& d_r1_free, const T2& r1_bound, const T3& d_r1_bound, const T4& exchange, const T5& d_exchange, const T6& bound, const T7& d_bound, const T8& dt, const T9& d_dt, const T10& attenuation, const T11& d_attenuation, const T12& bar_e11, const T13& d_bar_e11, const T14& bar_e12, const T15& d_bar_e12, const T16& bar_e21, const T17& d_bar_e21, const T18& bar_e22, const T19& d_bar_e22, const T20& bar_free, const T21& d_bar_free, const T22& bar_bound, const T23& d_bar_bound) { + bsk::tile_t | 0, 2)> back_bound{}; + bsk::tile_t | 0, 2)> back_free{}; + bsk::tile_t | 0, 2)> bar_half_gap{}; + bsk::tile_t | 0, 2)> bar_l12{}; + bsk::tile_t | 0, 2)> bar_l21{}; + bsk::tile_t | 0, 2)> d_back_bound{}; + bsk::tile_t | 0, 2)> d_back_free{}; + bsk::tile_t | 0, 2)> d_bar_half_gap{}; + bsk::tile_t | 0, 2)> d_bar_l12{}; + bsk::tile_t | 0, 2)> d_bar_l21{}; + auto free = (1.0f - bound); + auto d_free = (-d_bound); + auto kab = (exchange * bound); + auto d_kab = ((d_exchange * bound) + (exchange * d_bound)); + auto kba = (exchange * free); + auto d_kba = ((d_exchange * free) + (exchange * d_free)); + auto l11 = (((-kab) - r1_free) * dt); + auto d_l11 = ((((-d_kab) - d_r1_free) * dt) + (((-kab) - r1_free) * d_dt)); + auto l12 = (kba * dt); + auto d_l12 = ((d_kba * dt) + (kba * d_dt)); + auto l21 = (kab * dt); + auto d_l21 = ((d_kab * dt) + (kab * d_dt)); + auto l22 = (((-kba) - r1_bound) * dt); + auto d_l22 = ((((-d_kba) - d_r1_bound) * dt) + (((-kba) - r1_bound) * d_dt)); + auto half_trace = (0.5f * (l11 + l22)); + auto d_half_trace = (0.5f * (d_l11 + d_l22)); + auto half_gap = (0.5f * (l11 - l22)); + auto d_half_gap = (0.5f * (d_l11 - d_l22)); + auto square = ((half_gap * half_gap) + (l12 * l21)); + auto d_square = ((((2.0f * half_gap) * d_half_gap) + (d_l12 * l21)) + (l12 * d_l21)); + auto turning = (square > 1e-12f); + auto root = bsk::sqrt(bsk::maximum(square, 0.0f)); + auto guarded = bsk::where(turning, root, 1.0f); + auto d_root = bsk::where(turning, bsk::truediv((0.5f * d_square), guarded), 0.0f); + auto upper = bsk::exp((half_trace + root)); + auto d_upper = (upper * (d_half_trace + d_root)); + auto lower = bsk::exp((half_trace - root)); + auto d_lower = (lower * (d_half_trace - d_root)); + auto cosine = (0.5f * (upper + lower)); + auto d_cosine = (0.5f * (d_upper + d_lower)); + auto plain = bsk::exp(half_trace); + auto d_plain = (plain * d_half_trace); + auto poly = ((1.0f + bsk::truediv(square, 6.0f)) + bsk::truediv((square * square), 120.0f)); + auto d_poly = (bsk::truediv(d_square, 6.0f) + bsk::truediv((square * d_square), 60.0f)); + auto scale = bsk::where(turning, bsk::truediv((0.5f * (upper - lower)), guarded), (plain * poly)); + auto d_scale = bsk::where(turning, (bsk::truediv((0.5f * (d_upper - d_lower)), guarded) - bsk::truediv(((0.5f * (upper - lower)) * d_root), (guarded * guarded))), ((d_plain * poly) + (plain * d_poly))); + auto bare11 = (cosine + (scale * half_gap)); + auto d_bare11 = ((d_cosine + (d_scale * half_gap)) + (scale * d_half_gap)); + auto bare12 = (scale * l12); + auto d_bare12 = ((d_scale * l12) + (scale * d_l12)); + auto bare21 = (scale * l21); + auto d_bare21 = ((d_scale * l21) + (scale * d_l21)); + auto bare22 = (cosine - (scale * half_gap)); + auto d_bare22 = ((d_cosine - (d_scale * half_gap)) - (scale * d_half_gap)); + // The recovery reaches the operator's four entries and the two fractions. + auto carried11 = (bar_e11 - (bar_free * free)); + auto d_carried11 = (d_bar_e11 - ((d_bar_free * free) + (bar_free * d_free))); + auto carried12 = (bar_e12 - (bar_free * bound)); + auto d_carried12 = (d_bar_e12 - ((d_bar_free * bound) + (bar_free * d_bound))); + auto carried21 = (bar_e21 - (bar_bound * free)); + auto d_carried21 = (d_bar_e21 - ((d_bar_bound * free) + (bar_bound * d_free))); + auto carried22 = (bar_e22 - (bar_bound * bound)); + auto d_carried22 = (d_bar_e22 - ((d_bar_bound * bound) + (bar_bound * d_bound))); + auto e11 = (attenuation * bare11); + auto d_e11 = ((d_attenuation * bare11) + (attenuation * d_bare11)); + auto e12 = (attenuation * bare12); + auto d_e12 = ((d_attenuation * bare12) + (attenuation * d_bare12)); + auto e21 = (attenuation * bare21); + auto d_e21 = ((d_attenuation * bare21) + (attenuation * d_bare21)); + auto e22 = (attenuation * bare22); + auto d_e22 = ((d_attenuation * bare22) + (attenuation * d_bare22)); + back_free = ((bar_free * (1.0f - e11)) - (bar_bound * e21)); + d_back_free = (((d_bar_free * (1.0f - e11)) - (bar_free * d_e11)) - ((d_bar_bound * e21) + (bar_bound * d_e21))); + back_bound = ((bar_bound * (1.0f - e22)) - (bar_free * e12)); + d_back_bound = (((d_bar_bound * (1.0f - e22)) - (bar_bound * d_e22)) - ((d_bar_free * e12) + (bar_free * d_e12))); + auto back_attenuation = ((((carried11 * bare11) + (carried12 * bare12)) + (carried21 * bare21)) + (carried22 * bare22)); + auto d_back_attenuation = ((((((((d_carried11 * bare11) + (carried11 * d_bare11)) + (d_carried12 * bare12)) + (carried12 * d_bare12)) + (d_carried21 * bare21)) + (carried21 * d_bare21)) + (d_carried22 * bare22)) + (carried22 * d_bare22)); + auto scaled11 = (attenuation * carried11); + auto d_scaled11 = ((d_attenuation * carried11) + (attenuation * d_carried11)); + auto scaled12 = (attenuation * carried12); + auto d_scaled12 = ((d_attenuation * carried12) + (attenuation * d_carried12)); + auto scaled21 = (attenuation * carried21); + auto d_scaled21 = ((d_attenuation * carried21) + (attenuation * d_carried21)); + auto scaled22 = (attenuation * carried22); + auto d_scaled22 = ((d_attenuation * carried22) + (attenuation * d_carried22)); + auto bar_cosine = (scaled11 + scaled22); + auto d_bar_cosine = (d_scaled11 + d_scaled22); + auto gap = (scaled11 - scaled22); + auto d_gap = (d_scaled11 - d_scaled22); + auto bar_scale = (((gap * half_gap) + (scaled12 * l12)) + (scaled21 * l21)); + auto d_bar_scale = ((((((d_gap * half_gap) + (gap * d_half_gap)) + (d_scaled12 * l12)) + (scaled12 * d_l12)) + (d_scaled21 * l21)) + (scaled21 * d_l21)); + bar_half_gap = (scale * gap); + d_bar_half_gap = ((d_scale * gap) + (scale * d_gap)); + bar_l12 = (scale * scaled12); + d_bar_l12 = ((d_scale * scaled12) + (scale * d_scaled12)); + bar_l21 = (scale * scaled21); + d_bar_l21 = ((d_scale * scaled21) + (scale * d_scaled21)); + auto series_trace = ((bar_cosine * cosine) + (bar_scale * scale)); + auto d_series_trace = ((((d_bar_cosine * cosine) + (bar_cosine * d_cosine)) + (d_bar_scale * scale)) + (bar_scale * d_scale)); + auto cosine_poly = (0.5f + bsk::truediv(square, 12.0f)); + auto d_cosine_poly = bsk::truediv(d_square, 12.0f); + auto scale_poly = (0.16666666666666666f + bsk::truediv(square, 60.0f)); + auto d_scale_poly = bsk::truediv(d_square, 60.0f); + auto series_square = (plain * ((bar_cosine * cosine_poly) + (bar_scale * scale_poly))); + auto d_series_square = ((d_plain * ((bar_cosine * cosine_poly) + (bar_scale * scale_poly))) + (plain * ((((d_bar_cosine * cosine_poly) + (bar_cosine * d_cosine_poly)) + (d_bar_scale * scale_poly)) + (bar_scale * d_scale_poly)))); + auto inverse = bsk::where(turning, bsk::truediv(1.0f, guarded), 0.0f); + auto d_inverse = bsk::where(turning, bsk::truediv((-d_root), (guarded * guarded)), 0.0f); + auto bar_upper = (0.5f * (bar_cosine + (bar_scale * inverse))); + auto d_bar_upper = (0.5f * ((d_bar_cosine + (d_bar_scale * inverse)) + (bar_scale * d_inverse))); + auto bar_lower = (0.5f * (bar_cosine - (bar_scale * inverse))); + auto d_bar_lower = (0.5f * ((d_bar_cosine - (d_bar_scale * inverse)) - (bar_scale * d_inverse))); + auto root_trace = ((bar_upper * upper) + (bar_lower * lower)); + auto d_root_trace = ((((d_bar_upper * upper) + (bar_upper * d_upper)) + (d_bar_lower * lower)) + (bar_lower * d_lower)); + auto bar_root = (((bar_upper * upper) - (bar_lower * lower)) - ((bar_scale * scale) * inverse)); + auto d_bar_root = (((((d_bar_upper * upper) + (bar_upper * d_upper)) - (d_bar_lower * lower)) - (bar_lower * d_lower)) - ((((d_bar_scale * scale) * inverse) + ((bar_scale * d_scale) * inverse)) + ((bar_scale * scale) * d_inverse))); + auto root_square = ((0.5f * bar_root) * inverse); + auto d_root_square = (0.5f * ((d_bar_root * inverse) + (bar_root * d_inverse))); + auto bar_half_trace = bsk::where(turning, root_trace, series_trace); + auto d_bar_half_trace = bsk::where(turning, d_root_trace, d_series_trace); + auto bar_square = bsk::where(turning, root_square, series_square); + auto d_bar_square = bsk::where(turning, d_root_square, d_series_square); + bar_half_gap = (bar_half_gap + ((2.0f * bar_square) * half_gap)); + d_bar_half_gap = (d_bar_half_gap + (2.0f * ((d_bar_square * half_gap) + (bar_square * d_half_gap)))); + bar_l12 = (bar_l12 + (bar_square * l21)); + d_bar_l12 = (d_bar_l12 + ((d_bar_square * l21) + (bar_square * d_l21))); + bar_l21 = (bar_l21 + (bar_square * l12)); + d_bar_l21 = (d_bar_l21 + ((d_bar_square * l12) + (bar_square * d_l12))); + auto bar_l11 = (0.5f * (bar_half_trace + bar_half_gap)); + auto d_bar_l11 = (0.5f * (d_bar_half_trace + d_bar_half_gap)); + auto bar_l22 = (0.5f * (bar_half_trace - bar_half_gap)); + auto d_bar_l22 = (0.5f * (d_bar_half_trace - d_bar_half_gap)); + auto bar_kab = ((bar_l21 - bar_l11) * dt); + auto d_bar_kab = (((d_bar_l21 - d_bar_l11) * dt) + ((bar_l21 - bar_l11) * d_dt)); + auto bar_kba = ((bar_l12 - bar_l22) * dt); + auto d_bar_kba = (((d_bar_l12 - d_bar_l22) * dt) + ((bar_l12 - bar_l22) * d_dt)); + auto back_dt = ((((bar_l11 * ((-kab) - r1_free)) + (bar_l12 * kba)) + (bar_l21 * kab)) + (bar_l22 * ((-kba) - r1_bound))); + auto d_back_dt = ((((((((d_bar_l11 * ((-kab) - r1_free)) + (bar_l11 * ((-d_kab) - d_r1_free))) + (d_bar_l12 * kba)) + (bar_l12 * d_kba)) + (d_bar_l21 * kab)) + (bar_l21 * d_kab)) + (d_bar_l22 * ((-kba) - r1_bound))) + (bar_l22 * ((-d_kba) - d_r1_bound))); + back_bound = (back_bound + (bar_kab * exchange)); + d_back_bound = (d_back_bound + ((d_bar_kab * exchange) + (bar_kab * d_exchange))); + back_free = (back_free + (bar_kba * exchange)); + d_back_free = (d_back_free + ((d_bar_kba * exchange) + (bar_kba * d_exchange))); + return bsk::make_tup(((-bar_l11) * dt), ((-bar_l22) * dt), ((bar_kab * bound) + (bar_kba * free)), (back_bound - back_free), back_dt, back_attenuation, (-((d_bar_l11 * dt) + (bar_l11 * d_dt))), (-((d_bar_l22 * dt) + (bar_l22 * d_dt))), ((((d_bar_kab * bound) + (bar_kab * d_bound)) + (d_bar_kba * free)) + (bar_kba * d_free)), (d_back_bound - d_back_free), d_back_dt, d_back_attenuation); +} + +// The two-pool operator and its directional derivative. +// +// The same closed form :func:`_two_pool_step` evaluates, carried alongside a +// tangent. Returned as the six outputs then their six tangents. +template +BSK_HD auto _two_pool_step_jvp(const T0& r1_free, const T1& d_r1_free, const T2& r1_bound, const T3& d_r1_bound, const T4& exchange, const T5& d_exchange, const T6& bound, const T7& d_bound, const T8& dt, const T9& d_dt, const T10& attenuation, const T11& d_attenuation) { + auto free = (1.0f - bound); + auto d_free = (-d_bound); + auto kab = (exchange * bound); + auto d_kab = ((d_exchange * bound) + (exchange * d_bound)); + auto kba = (exchange * free); + auto d_kba = ((d_exchange * free) + (exchange * d_free)); + auto l11 = (((-kab) - r1_free) * dt); + auto d_l11 = ((((-d_kab) - d_r1_free) * dt) + (((-kab) - r1_free) * d_dt)); + auto l12 = (kba * dt); + auto d_l12 = ((d_kba * dt) + (kba * d_dt)); + auto l21 = (kab * dt); + auto d_l21 = ((d_kab * dt) + (kab * d_dt)); + auto l22 = (((-kba) - r1_bound) * dt); + auto d_l22 = ((((-d_kba) - d_r1_bound) * dt) + (((-kba) - r1_bound) * d_dt)); + auto half_trace = (0.5f * (l11 + l22)); + auto d_half_trace = (0.5f * (d_l11 + d_l22)); + auto half_gap = (0.5f * (l11 - l22)); + auto d_half_gap = (0.5f * (d_l11 - d_l22)); + auto square = ((half_gap * half_gap) + (l12 * l21)); + auto d_square = ((((2.0f * half_gap) * d_half_gap) + (d_l12 * l21)) + (l12 * d_l21)); + auto turning = (square > 1e-12f); + auto root = bsk::sqrt(bsk::maximum(square, 0.0f)); + auto guarded = bsk::where(turning, root, 1.0f); + auto d_root = bsk::where(turning, bsk::truediv((0.5f * d_square), guarded), 0.0f); + auto upper = bsk::exp((half_trace + root)); + auto d_upper = (upper * (d_half_trace + d_root)); + auto lower = bsk::exp((half_trace - root)); + auto d_lower = (lower * (d_half_trace - d_root)); + auto cosine = (0.5f * (upper + lower)); + auto d_cosine = (0.5f * (d_upper + d_lower)); + // sinh(d)/d by series where the root has no derivative of its own. + auto plain = bsk::exp(half_trace); + auto d_plain = (plain * d_half_trace); + auto poly = ((1.0f + bsk::truediv(square, 6.0f)) + bsk::truediv((square * square), 120.0f)); + auto d_poly = (bsk::truediv(d_square, 6.0f) + bsk::truediv((square * d_square), 60.0f)); + auto scale = bsk::where(turning, bsk::truediv((0.5f * (upper - lower)), guarded), (plain * poly)); + auto d_scale = bsk::where(turning, (bsk::truediv((0.5f * (d_upper - d_lower)), guarded) - bsk::truediv(((0.5f * (upper - lower)) * d_root), (guarded * guarded))), ((d_plain * poly) + (plain * d_poly))); + auto e11 = (attenuation * (cosine + (scale * half_gap))); + auto d_e11 = ((d_attenuation * (cosine + (scale * half_gap))) + (attenuation * ((d_cosine + (d_scale * half_gap)) + (scale * d_half_gap)))); + auto e12 = ((attenuation * scale) * l12); + auto d_e12 = ((((d_attenuation * scale) * l12) + ((attenuation * d_scale) * l12)) + ((attenuation * scale) * d_l12)); + auto e21 = ((attenuation * scale) * l21); + auto d_e21 = ((((d_attenuation * scale) * l21) + ((attenuation * d_scale) * l21)) + ((attenuation * scale) * d_l21)); + auto e22 = (attenuation * (cosine - (scale * half_gap))); + auto d_e22 = ((d_attenuation * (cosine - (scale * half_gap))) + (attenuation * ((d_cosine - (d_scale * half_gap)) - (scale * d_half_gap)))); + auto grow_free = (free - ((e11 * free) + (e12 * bound))); + auto d_grow_free = (d_free - ((((d_e11 * free) + (e11 * d_free)) + (d_e12 * bound)) + (e12 * d_bound))); + auto grow_bound = (bound - ((e21 * free) + (e22 * bound))); + auto d_grow_bound = (d_bound - ((((d_e21 * free) + (e21 * d_free)) + (d_e22 * bound)) + (e22 * d_bound))); + return bsk::make_tup(e11, e12, e21, e22, grow_free, grow_bound, d_e11, d_e12, d_e21, d_e22, d_grow_free, d_grow_bound); +} + +// ``e^z`` for ``z`` carried as a pair of floats. +template +BSK_HD auto _complex_exp(const T0& real, const T1& imag) { + auto scale = bsk::exp(real); + return bsk::make_tup((scale * bsk::cos(imag)), (scale * bsk::sin(imag))); +} + +// ``e^z`` and its directional derivative, which is ``e^z`` times it. +template +BSK_HD auto _complex_exp_jvp(const T0& real, const T1& imag, const T2& d_real, const T3& d_imag) { + auto t0_ = _complex_exp(real, imag); + auto value_real = bsk::get<0>(t0_); + auto value_imag = bsk::get<1>(t0_); + return bsk::make_tup(value_real, value_imag, ((value_real * d_real) - (value_imag * d_imag)), ((value_real * d_imag) + (value_imag * d_real))); +} + +// A square root of a complex number carried as a pair of floats. +// +// Which of the two it is does not matter here: the only thing that reads it +// is even in it, so the branch cut the principal root carries is unreachable. +template +BSK_HD auto _complex_sqrt(const T0& real, const T1& imag) { + auto magnitude = bsk::sqrt(((real * real) + (imag * imag))); + auto root_real = bsk::sqrt(bsk::maximum((0.5f * (magnitude + real)), 0.0f)); + auto root_imag = bsk::sqrt(bsk::maximum((0.5f * (magnitude - real)), 0.0f)); + return bsk::make_tup(root_real, bsk::where((imag < 0.0f), (-root_imag), root_imag)); +} + +// A complex square root and its directional derivative. +// +// The derivative divides by twice the root, so a caller keeps the origin -- +// where the root has none -- on its series branch. +template +BSK_HD auto _complex_sqrt_jvp(const T0& real, const T1& imag, const T2& d_real, const T3& d_imag) { + auto t0_ = _complex_sqrt(real, imag); + auto root_real = bsk::get<0>(t0_); + auto root_imag = bsk::get<1>(t0_); + auto guard = (2.0f * ((root_real * root_real) + (root_imag * root_imag))); + auto live = (guard > 0.0f); + auto guarded = bsk::where(live, guard, 1.0f); + // dz / (2 w) == dz * conj(2 w) / |2 w|^2 + auto tangent_real = bsk::where(live, bsk::truediv(((d_real * root_real) + (d_imag * root_imag)), guarded), 0.0f); + auto tangent_imag = bsk::where(live, bsk::truediv(((d_imag * root_real) - (d_real * root_imag)), guarded), 0.0f); + return bsk::make_tup(root_real, root_imag, tangent_real, tangent_imag); +} + +// ``1/z`` for a dual complex number, and the tangent that goes with it. +template +BSK_HD auto _dual_reciprocal(const T0& z) { + auto norm = ((bsk::get<0>(z) * bsk::get<0>(z)) + (bsk::get<1>(z) * bsk::get<1>(z))); + auto guard = bsk::where((norm > 0.0f), norm, 1.0f); + auto value_real = bsk::truediv(bsk::get<0>(z), guard); + auto value_imag = bsk::truediv((-bsk::get<1>(z)), guard); + auto t0_ = _complex_mul(value_real, value_imag, value_real, value_imag); + auto square_real = bsk::get<0>(t0_); + auto square_imag = bsk::get<1>(t0_); + auto t1_ = _complex_mul(square_real, square_imag, bsk::get<2>(z), bsk::get<3>(z)); + auto tangent_real = bsk::get<0>(t1_); + auto tangent_imag = bsk::get<1>(t1_); + return bsk::make_tup(value_real, value_imag, (-tangent_real), (-tangent_imag)); +} + +// The reverse sweep of :func:`_two_pool_transverse_step_jvp`. +// +// Every step from the four generator entries to the four operator entries is +// holomorphic, so the sweep is the longitudinal one with complex numbers in +// place of real ones and no conjugates along the way. That holds because the +// cotangents arrive as row covectors -- ``bar_e`` is the number with ``dL = +// Re(bar_e de)`` -- and only where a complex intermediate meets one of the +// real inputs is a real part taken. +// +// Takes the four cotangents as dual complex quadruples and returns the seven +// real gradients, each as a value and a tangent. +template +BSK_HD auto _two_pool_transverse_adjoint_jvp(const T0& r2_free, const T1& d_r2_free, const T2& r2_bound, const T3& d_r2_bound, const T4& exchange, const T5& d_exchange, const T6& bound, const T7& d_bound, const T8& free, const T9& d_free, const T10& shift_hz, const T11& d_shift_hz, const T12& dt, const T13& d_dt, const T14& attenuation, const T15& d_attenuation, const T16& bar_e11, const T17& bar_e12, const T18& bar_e21, const T19& bar_e22) { + bsk::tup | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>> bar_half_gap{}; + bsk::tup | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>> bar_l12{}; + bsk::tup | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>, bsk::tile_t | 0, 2)>> bar_l21{}; + auto zero = (0.0f * dt); + auto kab = (exchange * bound); + auto d_kab = ((d_exchange * bound) + (exchange * d_bound)); + auto kba = (exchange * free); + auto d_kba = ((d_exchange * free) + (exchange * d_free)); + auto turn = -6.283185307179586f; + auto l11 = bsk::make_tup((((-kab) - r2_free) * dt), zero, ((((-d_kab) - d_r2_free) * dt) + (((-kab) - r2_free) * d_dt)), zero); + auto l12 = bsk::make_tup((kba * dt), zero, ((d_kba * dt) + (kba * d_dt)), zero); + auto l21 = bsk::make_tup((kab * dt), zero, ((d_kab * dt) + (kab * d_dt)), zero); + auto l22 = bsk::make_tup((((-kba) - r2_bound) * dt), (turn * (shift_hz * dt)), ((((-d_kba) - d_r2_bound) * dt) + (((-kba) - r2_bound) * d_dt)), (turn * ((d_shift_hz * dt) + (shift_hz * d_dt)))); + auto half_trace = _dual_weigh(_dual_add(l11, l22), 0.5f); + auto half_gap = _dual_weigh(_dual_subtract(l11, l22), 0.5f); + auto square = _dual_add(_dual_product(half_gap, half_gap), _dual_product(l12, l21)); + auto delta = _complex_sqrt_jvp(bsk::get<0>(square), bsk::get<1>(square), bsk::get<2>(square), bsk::get<3>(square)); + auto upper = [&](const auto& s0_) { return _complex_exp_jvp(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(_dual_add(half_trace, delta)); + auto lower = [&](const auto& s0_) { return _complex_exp_jvp(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(_dual_subtract(half_trace, delta)); + auto plain = _complex_exp_jvp(bsk::get<0>(half_trace), bsk::get<1>(half_trace), bsk::get<2>(half_trace), bsk::get<3>(half_trace)); + auto cosine = _dual_weigh(_dual_add(upper, lower), 0.5f); + auto turning = (((bsk::get<0>(square) * bsk::get<0>(square)) + (bsk::get<1>(square) * bsk::get<1>(square))) > 1e-24f); + // Off the branch the reciprocal is taken at one instead, so a discriminant + // at the origin never divides anything the series answer then discards. + auto guarded = bsk::make_tup(bsk::where(turning, bsk::get<0>(delta), 1.0f), bsk::where(turning, bsk::get<1>(delta), 0.0f), bsk::where(turning, bsk::get<2>(delta), 0.0f), bsk::where(turning, bsk::get<3>(delta), 0.0f)); + auto inverse = _dual_reciprocal(guarded); + auto divided = _dual_product(_dual_weigh(_dual_subtract(upper, lower), 0.5f), inverse); + auto square2 = _dual_product(square, square); + auto poly = bsk::make_tup(((1.0f + bsk::truediv(bsk::get<0>(square), 6.0f)) + bsk::truediv(bsk::get<0>(square2), 120.0f)), (bsk::truediv(bsk::get<1>(square), 6.0f) + bsk::truediv(bsk::get<1>(square2), 120.0f)), (bsk::truediv(bsk::get<2>(square), 6.0f) + bsk::truediv(bsk::get<2>(square2), 120.0f)), (bsk::truediv(bsk::get<3>(square), 6.0f) + bsk::truediv(bsk::get<3>(square2), 120.0f))); + auto series = _dual_product(plain, poly); + auto scale = bsk::make_tup(bsk::where(turning, bsk::get<0>(divided), bsk::get<0>(series)), bsk::where(turning, bsk::get<1>(divided), bsk::get<1>(series)), bsk::where(turning, bsk::get<2>(divided), bsk::get<2>(series)), bsk::where(turning, bsk::get<3>(divided), bsk::get<3>(series))); + auto off = _dual_product(scale, half_gap); + auto bare_11 = _dual_add(cosine, off); + auto bare_12 = _dual_product(scale, l12); + auto bare_21 = _dual_product(scale, l21); + auto bare_22 = _dual_subtract(cosine, off); + auto bar_attenuation = _dual_sum(_dual_product(bar_e11, bare_11), _dual_product(bar_e12, bare_12), _dual_product(bar_e21, bare_21), _dual_product(bar_e22, bare_22)); + auto scaled_11 = [&](const auto& s2_) { return _dual_scale(attenuation, d_attenuation, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_e11); + auto scaled_12 = [&](const auto& s2_) { return _dual_scale(attenuation, d_attenuation, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_e12); + auto scaled_21 = [&](const auto& s2_) { return _dual_scale(attenuation, d_attenuation, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_e21); + auto scaled_22 = [&](const auto& s2_) { return _dual_scale(attenuation, d_attenuation, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_e22); + auto diagonal = _dual_subtract(scaled_11, scaled_22); + auto bar_cosine = _dual_add(scaled_11, scaled_22); + auto bar_scale = _dual_add(_dual_product(diagonal, half_gap), _dual_add(_dual_product(scaled_12, l12), _dual_product(scaled_21, l21))); + bar_half_gap = _dual_product(scale, diagonal); + bar_l12 = _dual_product(scale, scaled_12); + bar_l21 = _dual_product(scale, scaled_21); + auto series_trace = _dual_add(_dual_product(bar_cosine, cosine), _dual_product(bar_scale, scale)); + auto series_square = _dual_product(plain, _dual_add(_dual_product(bar_cosine, bsk::make_tup((0.5f + bsk::truediv(bsk::get<0>(square), 12.0f)), bsk::truediv(bsk::get<1>(square), 12.0f), bsk::truediv(bsk::get<2>(square), 12.0f), bsk::truediv(bsk::get<3>(square), 12.0f))), _dual_product(bar_scale, bsk::make_tup((0.16666666666666666f + bsk::truediv(bsk::get<0>(square), 60.0f)), bsk::truediv(bsk::get<1>(square), 60.0f), bsk::truediv(bsk::get<2>(square), 60.0f), bsk::truediv(bsk::get<3>(square), 60.0f))))); + auto bar_upper = _dual_weigh(_dual_add(bar_cosine, _dual_product(bar_scale, inverse)), 0.5f); + auto bar_lower = _dual_weigh(_dual_subtract(bar_cosine, _dual_product(bar_scale, inverse)), 0.5f); + auto split_trace = _dual_add(_dual_product(bar_upper, upper), _dual_product(bar_lower, lower)); + auto bar_delta = _dual_subtract(_dual_subtract(_dual_product(bar_upper, upper), _dual_product(bar_lower, lower)), _dual_product(_dual_product(bar_scale, scale), inverse)); + auto split_square = _dual_weigh(_dual_product(bar_delta, inverse), 0.5f); + auto bar_half_trace = bsk::make_tup(bsk::where(turning, bsk::get<0>(split_trace), bsk::get<0>(series_trace)), bsk::where(turning, bsk::get<1>(split_trace), bsk::get<1>(series_trace)), bsk::where(turning, bsk::get<2>(split_trace), bsk::get<2>(series_trace)), bsk::where(turning, bsk::get<3>(split_trace), bsk::get<3>(series_trace))); + auto bar_square = bsk::make_tup(bsk::where(turning, bsk::get<0>(split_square), bsk::get<0>(series_square)), bsk::where(turning, bsk::get<1>(split_square), bsk::get<1>(series_square)), bsk::where(turning, bsk::get<2>(split_square), bsk::get<2>(series_square)), bsk::where(turning, bsk::get<3>(split_square), bsk::get<3>(series_square))); + bar_half_gap = _dual_add(bar_half_gap, _dual_weigh(_dual_product(bar_square, half_gap), 2.0f)); + bar_l12 = _dual_add(bar_l12, _dual_product(bar_square, l21)); + bar_l21 = _dual_add(bar_l21, _dual_product(bar_square, l12)); + auto bar_l11 = _dual_weigh(_dual_add(bar_half_trace, bar_half_gap), 0.5f); + auto bar_l22 = _dual_weigh(_dual_subtract(bar_half_trace, bar_half_gap), 0.5f); + auto bar_kab = [&](const auto& s2_) { return _dual_scale(dt, d_dt, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(_dual_subtract(bar_l21, bar_l11)); + auto bar_kba = [&](const auto& s2_) { return _dual_scale(dt, d_dt, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(_dual_subtract(bar_l12, bar_l22)); + auto slope_22 = bsk::make_tup(((-kba) - r2_bound), (turn * shift_hz), ((-d_kba) - d_r2_bound), (turn * d_shift_hz)); + auto bar_dt = _dual_sum([&](const auto& s2_) { return _dual_scale(((-kab) - r2_free), ((-d_kab) - d_r2_free), bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_l11), [&](const auto& s2_) { return _dual_scale(kba, d_kba, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_l12), [&](const auto& s2_) { return _dual_scale(kab, d_kab, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_l21), _dual_product(slope_22, bar_l22)); + auto r2_free_bar = [&](const auto& s2_) { return _dual_scale((-dt), (-d_dt), bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_l11); + auto r2_bound_bar = [&](const auto& s2_) { return _dual_scale((-dt), (-d_dt), bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_l22); + auto exchange_bar = _dual_add([&](const auto& s2_) { return _dual_scale(bound, d_bound, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_kab), [&](const auto& s2_) { return _dual_scale(free, d_free, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_kba)); + auto bound_bar = [&](const auto& s2_) { return _dual_scale(exchange, d_exchange, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_kab); + auto free_bar = [&](const auto& s2_) { return _dual_scale(exchange, d_exchange, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(bar_kba); + auto shift_bar = [&](const auto& s2_) { return _dual_scale((turn * dt), (turn * d_dt), bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }([&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(bar_l22)); + return bsk::make_tup(bsk::get<0>(r2_free_bar), bsk::get<2>(r2_free_bar), bsk::get<0>(r2_bound_bar), bsk::get<2>(r2_bound_bar), bsk::get<0>(exchange_bar), bsk::get<2>(exchange_bar), bsk::get<0>(bound_bar), bsk::get<2>(bound_bar), bsk::get<0>(free_bar), bsk::get<2>(free_bar), bsk::get<0>(shift_bar), bsk::get<2>(shift_bar), bsk::get<0>(bar_dt), bsk::get<2>(bar_dt), bsk::get<0>(bar_attenuation), bsk::get<2>(bar_attenuation)); +} + +// The transverse operator and its directional derivative. +// +// The same closed form :func:`_two_pool_transverse_step` evaluates, carried +// alongside a tangent. Returned as the four entries then their four tangents, +// each a pair of floats. +template +BSK_HD auto _two_pool_transverse_step_jvp(const T0& r2_free, const T1& d_r2_free, const T2& r2_bound, const T3& d_r2_bound, const T4& exchange, const T5& d_exchange, const T6& bound, const T7& d_bound, const T8& free, const T9& d_free, const T10& shift_hz, const T11& d_shift_hz, const T12& dt, const T13& d_dt, const T14& attenuation, const T15& d_attenuation) { + auto kab = (exchange * bound); + auto d_kab = ((d_exchange * bound) + (exchange * d_bound)); + auto kba = (exchange * free); + auto d_kba = ((d_exchange * free) + (exchange * d_free)); + auto l11 = (((-kab) - r2_free) * dt); + auto d_l11 = ((((-d_kab) - d_r2_free) * dt) + (((-kab) - r2_free) * d_dt)); + auto l12 = (kba * dt); + auto d_l12 = ((d_kba * dt) + (kba * d_dt)); + auto l21 = (kab * dt); + auto d_l21 = ((d_kab * dt) + (kab * d_dt)); + auto l22 = (((-kba) - r2_bound) * dt); + auto d_l22 = ((((-d_kba) - d_r2_bound) * dt) + (((-kba) - r2_bound) * d_dt)); + auto turn = -6.283185307179586f; + auto l22_imag = ((turn * shift_hz) * dt); + auto d_l22_imag = (turn * ((d_shift_hz * dt) + (shift_hz * d_dt))); + auto trace_real = (0.5f * (l11 + l22)); + auto d_trace_real = (0.5f * (d_l11 + d_l22)); + auto trace_imag = (0.5f * l22_imag); + auto d_trace_imag = (0.5f * d_l22_imag); + auto gap_real = (0.5f * (l11 - l22)); + auto d_gap_real = (0.5f * (d_l11 - d_l22)); + auto gap_imag = (-0.5f * l22_imag); + auto d_gap_imag = (-0.5f * d_l22_imag); + auto square_real = (((gap_real * gap_real) - (gap_imag * gap_imag)) + (l12 * l21)); + auto d_square_real = (((((2.0f * gap_real) * d_gap_real) - ((2.0f * gap_imag) * d_gap_imag)) + (d_l12 * l21)) + (l12 * d_l21)); + auto square_imag = ((2.0f * gap_real) * gap_imag); + auto d_square_imag = (2.0f * ((d_gap_real * gap_imag) + (gap_real * d_gap_imag))); + auto t0_ = _complex_sqrt_jvp(square_real, square_imag, d_square_real, d_square_imag); + auto root_real = bsk::get<0>(t0_); + auto root_imag = bsk::get<1>(t0_); + auto d_root_real = bsk::get<2>(t0_); + auto d_root_imag = bsk::get<3>(t0_); + auto t1_ = _complex_exp_jvp((trace_real + root_real), (trace_imag + root_imag), (d_trace_real + d_root_real), (d_trace_imag + d_root_imag)); + auto upper_real = bsk::get<0>(t1_); + auto upper_imag = bsk::get<1>(t1_); + auto d_upper_real = bsk::get<2>(t1_); + auto d_upper_imag = bsk::get<3>(t1_); + auto t2_ = _complex_exp_jvp((trace_real - root_real), (trace_imag - root_imag), (d_trace_real - d_root_real), (d_trace_imag - d_root_imag)); + auto lower_real = bsk::get<0>(t2_); + auto lower_imag = bsk::get<1>(t2_); + auto d_lower_real = bsk::get<2>(t2_); + auto d_lower_imag = bsk::get<3>(t2_); + auto cos_real = (0.5f * (upper_real + lower_real)); + auto cos_imag = (0.5f * (upper_imag + lower_imag)); + auto d_cos_real = (0.5f * (d_upper_real + d_lower_real)); + auto d_cos_imag = (0.5f * (d_upper_imag + d_lower_imag)); + auto turning = (((square_real * square_real) + (square_imag * square_imag)) > 1e-24f); + auto half_real = (0.5f * (upper_real - lower_real)); + auto half_imag = (0.5f * (upper_imag - lower_imag)); + auto d_half_real = (0.5f * (d_upper_real - d_lower_real)); + auto d_half_imag = (0.5f * (d_upper_imag - d_lower_imag)); + auto norm = ((root_real * root_real) + (root_imag * root_imag)); + auto guard = bsk::where(turning, norm, 1.0f); + auto d_norm = bsk::where(turning, (2.0f * ((root_real * d_root_real) + (root_imag * d_root_imag))), 0.0f); + // (a / w) with w complex: a * conj(w) / |w|^2, differentiated as a quotient. + auto top_real = ((half_real * root_real) + (half_imag * root_imag)); + auto top_imag = ((half_imag * root_real) - (half_real * root_imag)); + auto d_top_real = ((((d_half_real * root_real) + (half_real * d_root_real)) + (d_half_imag * root_imag)) + (half_imag * d_root_imag)); + auto d_top_imag = ((((d_half_imag * root_real) + (half_imag * d_root_real)) - (d_half_real * root_imag)) - (half_real * d_root_imag)); + auto divided_real = bsk::truediv(top_real, guard); + auto divided_imag = bsk::truediv(top_imag, guard); + auto d_divided_real = bsk::truediv((d_top_real - (divided_real * d_norm)), guard); + auto d_divided_imag = bsk::truediv((d_top_imag - (divided_imag * d_norm)), guard); + auto t3_ = _complex_exp_jvp(trace_real, trace_imag, d_trace_real, d_trace_imag); + auto plain_real = bsk::get<0>(t3_); + auto plain_imag = bsk::get<1>(t3_); + auto d_plain_real = bsk::get<2>(t3_); + auto d_plain_imag = bsk::get<3>(t3_); + auto square2_real = ((square_real * square_real) - (square_imag * square_imag)); + auto square2_imag = ((2.0f * square_real) * square_imag); + auto d_square2_real = (((2.0f * square_real) * d_square_real) - ((2.0f * square_imag) * d_square_imag)); + auto d_square2_imag = (2.0f * ((d_square_real * square_imag) + (square_real * d_square_imag))); + auto poly_real = ((1.0f + bsk::truediv(square_real, 6.0f)) + bsk::truediv(square2_real, 120.0f)); + auto poly_imag = (bsk::truediv(square_imag, 6.0f) + bsk::truediv(square2_imag, 120.0f)); + auto d_poly_real = (bsk::truediv(d_square_real, 6.0f) + bsk::truediv(d_square2_real, 120.0f)); + auto d_poly_imag = (bsk::truediv(d_square_imag, 6.0f) + bsk::truediv(d_square2_imag, 120.0f)); + auto series_real = ((plain_real * poly_real) - (plain_imag * poly_imag)); + auto series_imag = ((plain_real * poly_imag) + (plain_imag * poly_real)); + auto d_series_real = ((((d_plain_real * poly_real) + (plain_real * d_poly_real)) - (d_plain_imag * poly_imag)) - (plain_imag * d_poly_imag)); + auto d_series_imag = ((((d_plain_real * poly_imag) + (plain_real * d_poly_imag)) + (d_plain_imag * poly_real)) + (plain_imag * d_poly_real)); + auto scale_real = bsk::where(turning, divided_real, series_real); + auto scale_imag = bsk::where(turning, divided_imag, series_imag); + auto d_scale_real = bsk::where(turning, d_divided_real, d_series_real); + auto d_scale_imag = bsk::where(turning, d_divided_imag, d_series_imag); + auto off_real = ((scale_real * gap_real) - (scale_imag * gap_imag)); + auto off_imag = ((scale_real * gap_imag) + (scale_imag * gap_real)); + auto d_off_real = ((((d_scale_real * gap_real) + (scale_real * d_gap_real)) - (d_scale_imag * gap_imag)) - (scale_imag * d_gap_imag)); + auto d_off_imag = ((((d_scale_real * gap_imag) + (scale_real * d_gap_imag)) + (d_scale_imag * gap_real)) + (scale_imag * d_gap_real)); + auto e11_real = (cos_real + off_real); + auto e11_imag = (cos_imag + off_imag); + auto d_e11_real = (d_cos_real + d_off_real); + auto d_e11_imag = (d_cos_imag + d_off_imag); + auto e22_real = (cos_real - off_real); + auto e22_imag = (cos_imag - off_imag); + auto d_e22_real = (d_cos_real - d_off_real); + auto d_e22_imag = (d_cos_imag - d_off_imag); + return bsk::make_tup((attenuation * e11_real), (attenuation * e11_imag), ((attenuation * scale_real) * l12), ((attenuation * scale_imag) * l12), ((attenuation * scale_real) * l21), ((attenuation * scale_imag) * l21), (attenuation * e22_real), (attenuation * e22_imag), ((d_attenuation * e11_real) + (attenuation * d_e11_real)), ((d_attenuation * e11_imag) + (attenuation * d_e11_imag)), (((d_attenuation * scale_real) * l12) + (attenuation * ((d_scale_real * l12) + (scale_real * d_l12)))), (((d_attenuation * scale_imag) * l12) + (attenuation * ((d_scale_imag * l12) + (scale_imag * d_l12)))), (((d_attenuation * scale_real) * l21) + (attenuation * ((d_scale_real * l21) + (scale_real * d_l21)))), (((d_attenuation * scale_imag) * l21) + (attenuation * ((d_scale_imag * l21) + (scale_imag * d_l21)))), ((d_attenuation * e22_real) + (attenuation * d_e22_real)), ((d_attenuation * e22_imag) + (attenuation * d_e22_imag))); +} + +// The same fraction and its directional derivative. +template +BSK_HD auto _washout_jvp(const T0& rate, const T1& rate_tangent, const T2& dt, const T3& dt_tangent) { + auto fraction = (rate * dt); + auto live = (fraction < 1.0f); + return bsk::make_tup(bsk::where(live, (1.0f - fraction), 0.0f), bsk::where(live, (-((rate_tangent * dt) + (rate * dt_tangent))), 0.0f)); +} + +BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, float* grad_pair_value, float* grad_pair_tangent, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_b1_phase, float* dot_b0, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_velocity, float* dot_bound_fraction, float* dot_exchange_rate, float* dot_t1_bound, float* dot_pool_b_fraction, float* dot_pool_b_exchange, float* dot_t1_pool_b, float* dot_t2_pool_b, float* dot_pool_b_shift, float* dot_duration, float* dot_flip, float* dot_phase, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, std::int64_t row_count, float* grad_output_real, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_phase_value, float* grad_phase_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_vr, float* trajectory_vi, float* trajectory_tr, float* trajectory_ti, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t directed, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t recording, std::int64_t block_states, std::int64_t problems) { + bsk::tup, bsk::V, bsk::V, bsk::V> a11{}; + bsk::tup, bsk::V, bsk::V, bsk::V> a12{}; + bsk::tup, bsk::V, bsk::V, bsk::V> a21{}; + bsk::tup, bsk::V, bsk::V, bsk::V> a22{}; + bsk::V absorbed_tangent{}; + bsk::V absorbed_value{}; + bsk::tup, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V> across{}; + bsk::tup, bsk::V, bsk::V, bsk::V> add1{}; + bsk::tup, bsk::V, bsk::V, bsk::V> add2{}; + bsk::V alpha_b_t{}; + bsk::V alpha_b_v{}; + bsk::V alpha_t{}; + bsk::V alpha_tangent{}; + bsk::V alpha_v{}; + bsk::V alpha_value{}; + bsk::V angle_tangent{}; + bsk::V angle_value{}; + bsk::V ati{}; + bsk::V atom_b0{}; + bsk::V atom_b1{}; + bsk::V atom_b1_phase{}; + bsk::V atom_bound{}; + bsk::V atom_damping{}; + bsk::V atom_exchange{}; + bsk::V atom_flow{}; + bsk::V atom_free{}; + bsk::V atom_inv{}; + bsk::V atom_m0{}; + bsk::V atom_semisolid{}; + bsk::V atom_semisolid_exchange{}; + bsk::V atom_shift{}; + bsk::V atom_t1b{}; + bsk::V atom_t2b{}; + bsk::V atom_washout{}; + bsk::V atr{}; + bsk::V att_rate{}; + bsk::V att_span{}; + bsk::V attenuation_t{}; + bsk::V attenuation_v{}; + bsk::V avi{}; + bsk::V avr{}; + bsk::V back_att_t{}; + bsk::V back_att_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_bb{}; + bsk::V back_bound_t{}; + bsk::V back_bound_v{}; + bsk::V back_dt_t{}; + bsk::V back_dt_v{}; + bsk::V back_exch_t{}; + bsk::V back_exch_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_mb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_pb{}; + bsk::V back_r1_t{}; + bsk::V back_r1_v{}; + bsk::V back_r1b_t{}; + bsk::V back_r1b_v{}; + bsk::V back_r1c_t{}; + bsk::V back_r1c_v{}; + bsk::V back_semi_t{}; + bsk::V back_semi_v{}; + bsk::V back_sexch_t{}; + bsk::V back_sexch_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_ub{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_wb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_zb{}; + bsk::V bare1_tangent{}; + bsk::V bare1_value{}; + bsk::V bare2_tangent{}; + bsk::V bare2_value{}; + std::int32_t base_row{}; + bsk::V bbti{}; + bsk::V bbtr{}; + bsk::V bbvi{}; + bsk::V bbvr{}; + bsk::V bmti{}; + bsk::V bmtr{}; + bsk::V bmvi{}; + bsk::V bmvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> bound_bar{}; + bsk::tup, bsk::V, bsk::V, bsk::V> bound_part{}; + bsk::V bpti{}; + bsk::V bptr{}; + bsk::V bpvi{}; + bsk::V bpvr{}; + bsk::V bti{}; + bsk::V btr{}; + bsk::V bvi{}; + bsk::V bvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> carried{}; + bsk::V cbti{}; + bsk::V cbtr{}; + bsk::V cbvi{}; + bsk::V cbvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> conjugated{}; + bsk::V cos_tangent{}; + bsk::V cos_value{}; + bsk::V cot2_t{}; + bsk::V cot2_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> cross_in{}; + bsk::tup, bsk::V, bsk::V, bsk::V> cross_out{}; + bsk::V cti{}; + bsk::V ctr{}; + bsk::V cvi{}; + bsk::V cvr{}; + bsk::V d_b0{}; + bsk::V d_b1{}; + bsk::V d_b1_phase{}; + bsk::V d_boundf{}; + bsk::V d_damping{}; + bsk::V d_exchange{}; + bsk::V d_flow{}; + bsk::V d_free{}; + bsk::V d_grow_free{}; + bsk::V d_grow_pool_b{}; + bsk::V d_grow_semisolid{}; + bsk::V d_inv{}; + bsk::V d_m0{}; + bsk::V d_semisolid_exchange{}; + bsk::V d_semisolid_t1{}; + bsk::V d_semisolidf{}; + bsk::V d_shift{}; + bsk::V d_t11{}; + bsk::V d_t12{}; + bsk::V d_t13{}; + bsk::V d_t1b{}; + bsk::V d_t21{}; + bsk::V d_t22{}; + bsk::V d_t23{}; + bsk::V d_t2b{}; + bsk::V d_t31{}; + bsk::V d_t32{}; + bsk::V d_t33{}; + bsk::V d_turn{}; + bsk::V d_w11{}; + bsk::V d_w12{}; + bsk::V d_w13{}; + bsk::V d_w21{}; + bsk::V d_w22{}; + bsk::V d_w23{}; + bsk::V d_w31{}; + bsk::V d_w32{}; + bsk::V d_w33{}; + bsk::V d_washout{}; + bsk::V damp_pair_t{}; + bsk::V damp_pair_v{}; + bsk::V damp_t{}; + bsk::V damp_t_tangent{}; + bsk::V damp_z{}; + bsk::V damp_z_tangent{}; + bsk::tup, bsk::V, bsk::V, bsk::V> damped{}; + bsk::V de11{}; + bsk::V de12{}; + bsk::V de21{}; + bsk::V de22{}; + bsk::V direction{}; + bool do_shift{}; + bsk::V drec_b{}; + bsk::V drec_f{}; + bsk::V dry1_tangent{}; + bsk::V dry1_value{}; + bsk::V dry2_tangent{}; + bsk::V dry2_value{}; + bsk::V dt_tangent{}; + bsk::V dt_value{}; + bsk::V dturn_t{}; + bsk::V dturn_z{}; + bsk::V duration_t{}; + bsk::V duration_v{}; + bsk::V e11_t{}; + bsk::V e11_v{}; + bsk::V e12_t{}; + bsk::V e12_v{}; + bsk::V e1_tangent{}; + bsk::V e1_value{}; + bsk::V e21_t{}; + bsk::V e21_v{}; + bsk::V e22_t{}; + bsk::V e22_v{}; + bsk::V e2_tangent{}; + bsk::V e2_value{}; + std::int64_t event{}; + std::int32_t event_action{}; + bsk::V event_dot_flip{}; + bsk::V event_dot_phase{}; + bsk::V event_flip{}; + std::int32_t event_kind{}; + bsk::V event_phase{}; + float event_saturation{}; + bsk::tup, bsk::V, bsk::V, bsk::V> free_bar{}; + bsk::tup, bsk::V, bsk::V, bsk::V> free_minus{}; + bsk::tup, bsk::V, bsk::V, bsk::V> free_part{}; + bsk::tup, bsk::V, bsk::V, bsk::V> free_plus{}; + bsk::V g_b0t{}; + bsk::V g_b0v{}; + bsk::V g_b1pt{}; + bsk::V g_b1pv{}; + bsk::V g_b1t{}; + bsk::V g_b1v{}; + bsk::V g_boundt{}; + bsk::V g_boundv{}; + bsk::V g_difft{}; + bsk::V g_diffv{}; + bsk::V g_excht{}; + bsk::V g_exchv{}; + bsk::V g_flowt{}; + bsk::V g_flowv{}; + bsk::V g_invt{}; + bsk::V g_invv{}; + bsk::V g_m0t{}; + bsk::V g_m0v{}; + bsk::V g_semit{}; + bsk::V g_semiv{}; + bsk::V g_sexcht{}; + bsk::V g_sexchv{}; + bsk::V g_shiftt{}; + bsk::V g_shiftv{}; + bsk::V g_t1bt{}; + bsk::V g_t1bv{}; + bsk::V g_t1ct{}; + bsk::V g_t1cv{}; + bsk::V g_t1t{}; + bsk::V g_t1v{}; + bsk::V g_t2bt{}; + bsk::V g_t2bv{}; + bsk::V g_t2t{}; + bsk::V g_t2v{}; + bsk::V g_washt{}; + bsk::V g_washv{}; + bsk::V grad_alpha_t{}; + bsk::V grad_alpha_v{}; + bsk::V grad_angle_t{}; + bsk::V grad_angle_v{}; + bsk::V grad_e1_t{}; + bsk::V grad_e1_v{}; + bsk::V grad_e2_t{}; + bsk::V grad_e2_v{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; + bsk::V held{}; + bsk::tup, bsk::V, bsk::V, bsk::V> held_bar{}; + bsk::V held_semisolid{}; + bsk::tup, bsk::V, bsk::V, bsk::V> held_state{}; + bool invert{}; + bool is_inversion{}; + bool is_rf{}; + bsk::V iti{}; + bsk::V itr{}; + bsk::V ivi{}; + bsk::V ivr{}; + bsk::V long_damp_t{}; + bsk::V long_damp_v{}; + bsk::V lti{}; + bsk::V ltr{}; + bsk::V lvi{}; + bsk::V lvr{}; + bsk::V mbti{}; + bsk::V mbtr{}; + bsk::V mbvi{}; + bsk::V mbvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> mixed_bound{}; + bsk::tup, bsk::V, bsk::V, bsk::V> mixed_free{}; + bsk::tup, bsk::V, bsk::V, bsk::V> mixed_semisolid{}; + bsk::tup, bsk::V, bsk::V, bsk::V> mo{}; + bsk::V mti{}; + bsk::V mtr{}; + bsk::V mvi{}; + bsk::V mvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> n0{}; + bsk::tup, bsk::V, bsk::V, bsk::V> n1{}; + bsk::tup, bsk::V, bsk::V, bsk::V> n2{}; + bsk::tup, bsk::V, bsk::V, bsk::V> next_mb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> next_pb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> next_ub{}; + bsk::tup, bsk::V, bsk::V, bsk::V> next_wb{}; + bsk::V offset_value{}; + bsk::V one_att{}; + bsk::V other_t{}; + bsk::V other_v{}; + bsk::V oti{}; + bsk::V otr{}; + bsk::V ovi{}; + bsk::V ovr{}; + bsk::V p1i{}; + bsk::V p1r{}; + bsk::V p1ti{}; + bsk::V p1tr{}; + bsk::V p2i{}; + bsk::V p2r{}; + bsk::V p2ti{}; + bsk::V p2tr{}; + bsk::V part_t{}; + bsk::V part_v{}; + bsk::V pbti{}; + bsk::V pbtr{}; + bsk::V pbvi{}; + bsk::V pbvr{}; + bsk::V pe11{}; + bsk::V pe12{}; + bsk::V pe21{}; + bsk::V pe22{}; + bsk::V per_angle_t{}; + bsk::V per_angle_v{}; + bsk::V phi_b_t{}; + bsk::V phi_b_v{}; + bsk::V phi_t{}; + bsk::V phi_tangent{}; + bsk::V phi_v{}; + bsk::V phi_value{}; + bsk::tup, bsk::V, bsk::V, bsk::V> po{}; + bsk::tup, bsk::V, bsk::V, bsk::V> pool_minus{}; + bsk::tup, bsk::V, bsk::V, bsk::V> pool_plus{}; + bsk::V pool_row{}; + bsk::V power_tangent{}; + bsk::V power_value{}; + bool pre_shift{}; + bsk::V prec_b{}; + bsk::V prec_f{}; + bsk::V problem{}; + bsk::V pti{}; + bsk::V ptr{}; + bsk::V pvi{}; + bsk::V pvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> q0{}; + bsk::tup, bsk::V, bsk::V, bsk::V> q1{}; + bsk::tup, bsk::V, bsk::V, bsk::V> q2{}; + bsk::V qi{}; + bsk::V qr{}; + bsk::V qti{}; + bsk::V qtr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> r12{}; + bsk::V r1b_tangent{}; + bsk::V r1b_value{}; + bsk::V r1c_tangent{}; + bsk::V r1c_value{}; + bsk::tup, bsk::V, bsk::V, bsk::V> r21{}; + bsk::tup, bsk::V, bsk::V, bsk::V> r22{}; + bsk::V r2b_tangent{}; + bsk::V r2b_value{}; + bsk::V rbmti{}; + bsk::V rbmtr{}; + bsk::V rbmvi{}; + bsk::V rbmvr{}; + bsk::V rbpti{}; + bsk::V rbptr{}; + bsk::V rbpvi{}; + bsk::V rbpvr{}; + bsk::V rbti{}; + bsk::V rbtr{}; + bsk::V rbvi{}; + bsk::V rbvr{}; + bsk::V rcti{}; + bsk::V rctr{}; + bsk::V rcvi{}; + bsk::V rcvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> recorded{}; + bsk::V recovery_tangent{}; + bsk::V recovery_value{}; + bsk::V rmti{}; + bsk::V rmtr{}; + bsk::V rmvi{}; + bsk::V rmvr{}; + bool rotate{}; + std::int64_t row{}; + bsk::tup, bsk::V, bsk::V, bsk::V> row0{}; + bsk::V rpti{}; + bsk::V rptr{}; + bsk::V rpvi{}; + bsk::V rpvr{}; + bsk::V rzti{}; + bsk::V rztr{}; + bsk::V rzvi{}; + bsk::V rzvr{}; + bsk::V sat_alpha_t{}; + bsk::V sat_alpha_v{}; + bsk::V sat_b0_t{}; + bsk::V sat_b0_v{}; + bool saturating{}; + bsk::V sbmti{}; + bsk::V sbmtr{}; + bsk::V sbmvi{}; + bsk::V sbmvr{}; + bsk::V sbpti{}; + bsk::V sbptr{}; + bsk::V sbpvi{}; + bsk::V sbpvr{}; + bsk::V scale1_tangent{}; + bsk::V scale2_tangent{}; + bsk::V shape_slope{}; + bsk::V shape_tangent{}; + bsk::V shape_value{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_a{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_b{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_bb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_mb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_pb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_slope_a{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_slope_b{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_ub{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_wb{}; + bsk::tup, bsk::V, bsk::V, bsk::V> shaped_zb{}; + bsk::V sin_tangent{}; + bsk::V sin_value{}; + bsk::V slope1_t{}; + bsk::V slope1_v{}; + bsk::V slope1b_t{}; + bsk::V slope1b_v{}; + bsk::V slope1c_t{}; + bsk::V slope1c_v{}; + bsk::V slot{}; + bsk::tup, bsk::V, bsk::V, bsk::V> spin{}; + bool spoil{}; + bsk::V spread_t{}; + bsk::V spread_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> spun_bound{}; + bsk::tup, bsk::V, bsk::V, bsk::V> spun_free{}; + bsk::V spun_mti{}; + bsk::V spun_mtr{}; + bsk::V spun_mvi{}; + bsk::V spun_mvr{}; + bsk::V spun_pti{}; + bsk::V spun_ptr{}; + bsk::V spun_pvi{}; + bsk::V spun_pvr{}; + bsk::V spun_zti{}; + bsk::V spun_ztr{}; + bsk::V spun_zvi{}; + bsk::V spun_zvr{}; + bsk::V sti{}; + bsk::V str_{}; + bsk::V svi{}; + bsk::V svr{}; + bsk::V szi{}; + bsk::V szr{}; + bsk::V szti{}; + bsk::V sztr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> t00{}; + bsk::tup, bsk::V, bsk::V, bsk::V> t01{}; + bsk::tup, bsk::V, bsk::V, bsk::V> t02{}; + bsk::V t11{}; + bsk::V t12{}; + bsk::V t13{}; + bsk::tup, bsk::V, bsk::V, bsk::V> t20{}; + bsk::V t21{}; + bsk::V t22{}; + bsk::V t23{}; + bsk::V t31{}; + bsk::V t32{}; + bsk::V t33{}; + bsk::V three_a00{}; + bsk::V three_a01{}; + bsk::V three_a02{}; + bsk::V three_a10{}; + bsk::V three_a11{}; + bsk::V three_a20{}; + bsk::V three_a22{}; + bsk::V three_angle{}; + bsk::V three_argument{}; + bsk::V three_centre{}; + bsk::V three_cube{}; + bsk::V three_d_a00{}; + bsk::V three_d_a01{}; + bsk::V three_d_a02{}; + bsk::V three_d_a10{}; + bsk::V three_d_a11{}; + bsk::V three_d_a20{}; + bsk::V three_d_a22{}; + bsk::V three_d_angle{}; + bsk::V three_d_centre{}; + bsk::V three_d_determinant{}; + bsk::V three_d_first{}; + bsk::V three_d_free{}; + bsk::V three_d_guarded{}; + bsk::V three_d_high{}; + bsk::V three_d_leading{}; + bsk::V three_d_lift{}; + bsk::V three_d_low{}; + bsk::V three_d_middle{}; + bsk::V three_d_minors{}; + bsk::V three_d_pool_b{}; + bsk::V three_d_pool_c{}; + bsk::V three_d_q00{}; + bsk::V three_d_q01{}; + bsk::V three_d_q02{}; + bsk::V three_d_q10{}; + bsk::V three_d_q11{}; + bsk::V three_d_q12{}; + bsk::V three_d_q20{}; + bsk::V three_d_q21{}; + bsk::V three_d_q22{}; + bsk::V three_d_radius{}; + bsk::V three_d_raw{}; + bsk::V three_d_s00{}; + bsk::V three_d_s11{}; + bsk::V three_d_s22{}; + bsk::V three_d_second{}; + bsk::V three_d_sum_flat{}; + bsk::V three_d_sum_linear{}; + bsk::V three_d_sum_square{}; + bsk::V three_d_trailing{}; + bsk::V three_def_00{}; + bsk::V three_def_01{}; + bsk::V three_def_02{}; + bsk::V three_def_10{}; + bsk::V three_def_11{}; + bsk::V three_def_12{}; + bsk::V three_def_20{}; + bsk::V three_def_21{}; + bsk::V three_def_22{}; + bsk::V three_determinant{}; + bsk::V three_dif_00{}; + bsk::V three_dif_01{}; + bsk::V three_dif_02{}; + bsk::V three_dif_10{}; + bsk::V three_dif_11{}; + bsk::V three_dif_12{}; + bsk::V three_dif_20{}; + bsk::V three_dif_21{}; + bsk::V three_dif_22{}; + bsk::V three_first{}; + bsk::V three_free{}; + bsk::V three_guarded{}; + bsk::V three_high{}; + bsk::V three_inside_limit{}; + bsk::V three_leading{}; + bsk::V three_lift{}; + bsk::V three_low{}; + bsk::V three_middle{}; + bsk::V three_minors{}; + bsk::V three_pool_b{}; + bsk::V three_pool_c{}; + bsk::V three_q00{}; + bsk::V three_q01{}; + bsk::V three_q02{}; + bsk::V three_q10{}; + bsk::V three_q11{}; + bsk::V three_q12{}; + bsk::V three_q20{}; + bsk::V three_q21{}; + bsk::V three_q22{}; + bsk::V three_radius{}; + bsk::V three_raw{}; + bsk::V three_s00{}; + bsk::V three_s11{}; + bsk::V three_s22{}; + bsk::V three_second{}; + bsk::V three_sum_flat{}; + bsk::V three_sum_linear{}; + bsk::V three_sum_square{}; + bsk::V three_trailing{}; + bsk::V turn_t{}; + bsk::V turn_z{}; + bsk::V turned_mti{}; + bsk::V turned_mtr{}; + bsk::V turned_mvi{}; + bsk::V turned_mvr{}; + bsk::V turned_pti{}; + bsk::V turned_ptr{}; + bsk::V turned_pvi{}; + bsk::V turned_pvr{}; + bsk::V turned_zti{}; + bsk::V turned_ztr{}; + bsk::V turned_zvi{}; + bsk::V turned_zvr{}; + bsk::V two_pool_dt_t{}; + bsk::V two_pool_dt_v{}; + bsk::tup, bsk::V, bsk::V, bsk::V> u1{}; + bsk::tup, bsk::V, bsk::V, bsk::V> u2{}; + bsk::V ubti{}; + bsk::V ubtr{}; + bsk::V ubvi{}; + bsk::V ubvr{}; + bsk::V ui{}; + bsk::V ur{}; + bsk::V uti{}; + bsk::V utr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> w0{}; + bsk::tup, bsk::V, bsk::V, bsk::V> w1{}; + bsk::V w11{}; + bsk::V w12{}; + bsk::V w13{}; + bsk::tup, bsk::V, bsk::V, bsk::V> w2{}; + bsk::V w21{}; + bsk::V w22{}; + bsk::V w23{}; + bsk::V w31{}; + bsk::V w32{}; + bsk::V w33{}; + bsk::V wash_t{}; + bsk::V wash_v{}; + bsk::V wbti{}; + bsk::V wbtr{}; + bsk::V wbvi{}; + bsk::V wbvr{}; + bsk::V wound_t{}; + bsk::V wound_v{}; + bsk::V wout_tangent{}; + bsk::V wout_value{}; + bsk::V wti{}; + bsk::V wtr{}; + bsk::V wvi{}; + bsk::V wvr{}; + bsk::V xbmti{}; + bsk::V xbmtr{}; + bsk::V xbmvi{}; + bsk::V xbmvr{}; + bsk::V xbpti{}; + bsk::V xbptr{}; + bsk::V xbpvi{}; + bsk::V xbpvr{}; + bsk::V xbti{}; + bsk::V xbtr{}; + bsk::V xbvi{}; + bsk::V xbvr{}; + bsk::V xcti{}; + bsk::V xctr{}; + bsk::V xcvi{}; + bsk::V xcvr{}; + bsk::V yi{}; + bsk::V yr{}; + bsk::V yti{}; + bsk::V ytr{}; + bsk::V zangle_t{}; + bsk::V zangle_v{}; + bsk::V zbti{}; + bsk::V zbtr{}; + bsk::V zbvi{}; + bsk::V zbvr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> zo{}; + bsk::V zti{}; + bsk::V ztr{}; + bsk::V zvi{}; + bsk::V zvr{}; + problem = (problem_base + (bsk::program_id(0) * problems)); + problem = (problem + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < problem_end); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + // Voxels are spread over the slice voxel-major, so a voxel's place along + // the slice is its index modulo the profile's width. One pulse shape holds + // that many consecutive rows, and the event says which shape it drives. + auto location = bsk::mod(atom, locations); + auto local = (problem - problem_base); + // A second pool rides along as planes of its own: it enters an event as its + // own vector and the RF operator acts on it, so the reverse sweep cannot + // replay it from the free pool's. A semisolid pool adds one plane, a + // chemically exchanging one three, and the two together add four. + auto record_stride = (bsk::select(bsk::truth((pools == 3)), 7, bsk::select(bsk::truth((pools == 2)), 6, bsk::select(bsk::truth((pools == 1)), 4, 3))) * state_count); + auto trajectory = (((local * event_count) * record_stride) + state); + auto minus_plane = state_count; + auto long_plane = (2 * state_count); + auto bound_plane = (3 * state_count); + auto bplus_plane = (4 * state_count); + auto bminus_plane = (5 * state_count); + auto semisolid_plane = (6 * state_count); + auto empty = bsk::full(0); + pvr = empty; + pvi = empty; + ptr = empty; + pti = empty; + mvr = empty; + mvi = empty; + mtr = empty; + mti = empty; + bvr = empty; + bvi = empty; + btr = empty; + bti = empty; + bpvr = empty; + bpvi = empty; + bptr = empty; + bpti = empty; + bmvr = empty; + bmvi = empty; + bmtr = empty; + bmti = empty; + cvr = empty; + cvi = empty; + ctr = empty; + cti = empty; + atom_bound = 0.0f; + d_boundf = 0.0f; + atom_exchange = 0.0f; + d_exchange = 0.0f; + atom_t1b = 1.0f; + d_t1b = 0.0f; + r1b_value = 0.0f; + r1b_tangent = 0.0f; + atom_t2b = 1.0f; + d_t2b = 0.0f; + r2b_value = 0.0f; + r2b_tangent = 0.0f; + atom_shift = 0.0f; + d_shift = 0.0f; + atom_semisolid = 0.0f; + d_semisolidf = 0.0f; + atom_semisolid_exchange = 0.0f; + d_semisolid_exchange = 0.0f; + r1c_value = 0.0f; + r1c_tangent = 0.0f; + if (bsk::truth((pools == 1))) { + atom_bound = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + d_boundf = bsk::ld((dot_bound_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + d_exchange = bsk::ld((dot_exchange_rate + scalar_atom), active_atom, 0.0f); + atom_t1b = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + d_t1b = bsk::ld((dot_t1_bound + scalar_atom), active_atom, 0.0f); + r1b_value = bsk::truediv(1000.0f, atom_t1b); + r1b_tangent = bsk::truediv((-1000.0f * d_t1b), (atom_t1b * atom_t1b)); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + atom_bound = bsk::ld((pool_b_fraction + scalar_atom), active_atom, 0.0f); + d_boundf = bsk::ld((dot_pool_b_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((pool_b_exchange + scalar_atom), active_atom, 0.0f); + d_exchange = bsk::ld((dot_pool_b_exchange + scalar_atom), active_atom, 0.0f); + atom_t1b = bsk::ld((t1_pool_b + scalar_atom), active_atom, 1.0f); + d_t1b = bsk::ld((dot_t1_pool_b + scalar_atom), active_atom, 0.0f); + r1b_value = bsk::truediv(1000.0f, atom_t1b); + r1b_tangent = bsk::truediv((-1000.0f * d_t1b), (atom_t1b * atom_t1b)); + atom_t2b = bsk::ld((t2_pool_b + scalar_atom), active_atom, 1.0f); + d_t2b = bsk::ld((dot_t2_pool_b + scalar_atom), active_atom, 0.0f); + r2b_value = bsk::truediv(1000.0f, atom_t2b); + r2b_tangent = bsk::truediv((-1000.0f * d_t2b), (atom_t2b * atom_t2b)); + atom_shift = bsk::ld((pool_b_shift + scalar_atom), active_atom, 0.0f); + d_shift = bsk::ld((dot_pool_b_shift + scalar_atom), active_atom, 0.0f); + } + if (bsk::truth((pools == 3))) { + atom_semisolid = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + d_semisolidf = bsk::ld((dot_bound_fraction + scalar_atom), active_atom, 0.0f); + atom_semisolid_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + d_semisolid_exchange = bsk::ld((dot_exchange_rate + scalar_atom), active_atom, 0.0f); + held_semisolid = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + d_semisolid_t1 = bsk::ld((dot_t1_bound + scalar_atom), active_atom, 0.0f); + r1c_value = bsk::truediv(1000.0f, held_semisolid); + r1c_tangent = bsk::truediv((-1000.0f * d_semisolid_t1), (held_semisolid * held_semisolid)); + cvr = (empty + bsk::where((state == 0), (atom_semisolid + 0.0f), 0.0f)); + ctr = (empty + bsk::where((state == 0), (d_semisolidf + 0.0f), 0.0f)); + } + if (bsk::truth((pools > 0))) { + atom_free = ((1.0f - atom_bound) - atom_semisolid); + d_free = ((-d_boundf) - d_semisolidf); + zvr = (empty + bsk::where((state == 0), atom_free, 0.0f)); + ztr = (empty + bsk::where((state == 0), d_free, 0.0f)); + bvr = (empty + bsk::where((state == 0), (atom_bound + 0.0f), 0.0f)); + btr = (empty + bsk::where((state == 0), (d_boundf + 0.0f), 0.0f)); + } else { + atom_free = (1.0f + (0.0f * atom_bound)); + d_free = (0.0f * atom_bound); + zvr = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + ztr = empty; + } + zvi = empty; + zti = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_b1_phase = 0.0f; + atom_b0 = 0.0f; + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld((b1_phase + scalar_atom), active_atom, 0.0f); + atom_b0 = bsk::ld((b0 + scalar_atom), active_atom, 0.0f); + } + atom_inv = 1.0f; + if (bsk::truth(inverting)) { + atom_inv = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + auto d_t1 = bsk::ld((dot_t1 + atom), active_atom, 0.0f); + auto d_t2 = bsk::ld((dot_t2 + atom), active_atom, 0.0f); + d_m0 = 0.0f; + if (bsk::truth(density)) { + d_m0 = bsk::ld((dot_m0 + scalar_atom), active_atom, 0.0f); + } + d_b1 = 0.0f; + if (bsk::truth(transmit)) { + d_b1 = bsk::ld((dot_b1 + scalar_atom), active_atom, 0.0f); + } + d_b1_phase = 0.0f; + d_b0 = 0.0f; + if (bsk::truth(off_axis)) { + d_b1_phase = bsk::ld((dot_b1_phase + scalar_atom), active_atom, 0.0f); + d_b0 = bsk::ld((dot_b0 + scalar_atom), active_atom, 0.0f); + } + d_inv = 0.0f; + if (bsk::truth(inverting)) { + d_inv = bsk::ld((dot_inversion_efficiency + scalar_atom), active_atom, 0.0f); + } + atom_damping = 0.0f; + d_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + d_damping = bsk::ld((dot_diffusion + scalar_atom), active_atom, 0.0f); + } + atom_flow = 0.0f; + d_flow = 0.0f; + direction = 0.0f; + atom_washout = 0.0f; + d_washout = 0.0f; + if (bsk::truth(moving)) { + auto atom_velocity = bsk::ld((velocity + scalar_atom), active_atom, 0.0f); + auto d_velocity = bsk::ld((dot_velocity + scalar_atom), active_atom, 0.0f); + atom_flow = (atom_velocity * flow_scale); + d_flow = (d_velocity * flow_scale); + // |v| has no derivative at the origin, so a still voxel contributes + // none. + direction = (bsk::cast((atom_velocity > 0.0f)) - bsk::cast((atom_velocity < 0.0f))); + atom_washout = (bsk::abs(atom_velocity) * washout_scale); + d_washout = ((direction * d_velocity) * washout_scale); + } + auto order = bsk::cast(state); + auto longitudinal_weight = (order * order); + auto transverse_weight = ((longitudinal_weight + order) + 0.3333333333333333f); + auto r1_value = bsk::truediv(1000.0f, atom_t1); + auto r1_tangent = bsk::truediv((-1000.0f * d_t1), (atom_t1 * atom_t1)); + auto r2_value = bsk::truediv(1000.0f, atom_t2); + auto r2_tangent = bsk::truediv((-1000.0f * d_t2), (atom_t2 * atom_t2)); + auto event_base = (train * event_count); + // The forward half records the trajectory the reverse half walks back, + // and the two are launched separately: each compiles the sweep it is + // asked for and no more. + if (bsk::truth(recording)) { + for (std::int64_t event = 0; event < event_count; event += 1) { + slot = (trajectory + (event * record_stride)); + bsk::st((trajectory_vr + slot), pvr, state_mask); + bsk::st((trajectory_vi + slot), pvi, state_mask); + bsk::st((trajectory_tr + slot), ptr, state_mask); + bsk::st((trajectory_ti + slot), pti, state_mask); + bsk::st(((trajectory_vr + slot) + minus_plane), mvr, state_mask); + bsk::st(((trajectory_vi + slot) + minus_plane), mvi, state_mask); + bsk::st(((trajectory_tr + slot) + minus_plane), mtr, state_mask); + bsk::st(((trajectory_ti + slot) + minus_plane), mti, state_mask); + bsk::st(((trajectory_vr + slot) + long_plane), zvr, state_mask); + bsk::st(((trajectory_vi + slot) + long_plane), zvi, state_mask); + bsk::st(((trajectory_tr + slot) + long_plane), ztr, state_mask); + bsk::st(((trajectory_ti + slot) + long_plane), zti, state_mask); + if (bsk::truth((pools > 0))) { + bsk::st(((trajectory_vr + slot) + bound_plane), bvr, state_mask); + bsk::st(((trajectory_vi + slot) + bound_plane), bvi, state_mask); + bsk::st(((trajectory_tr + slot) + bound_plane), btr, state_mask); + bsk::st(((trajectory_ti + slot) + bound_plane), bti, state_mask); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + bsk::st(((trajectory_vr + slot) + bplus_plane), bpvr, state_mask); + bsk::st(((trajectory_vi + slot) + bplus_plane), bpvi, state_mask); + bsk::st(((trajectory_tr + slot) + bplus_plane), bptr, state_mask); + bsk::st(((trajectory_ti + slot) + bplus_plane), bpti, state_mask); + bsk::st(((trajectory_vr + slot) + bminus_plane), bmvr, state_mask); + bsk::st(((trajectory_vi + slot) + bminus_plane), bmvi, state_mask); + bsk::st(((trajectory_tr + slot) + bminus_plane), bmtr, state_mask); + bsk::st(((trajectory_ti + slot) + bminus_plane), bmti, state_mask); + } + if (bsk::truth((pools == 3))) { + bsk::st(((trajectory_vr + slot) + semisolid_plane), cvr, state_mask); + bsk::st(((trajectory_vi + slot) + semisolid_plane), cvi, state_mask); + bsk::st(((trajectory_tr + slot) + semisolid_plane), ctr, state_mask); + bsk::st(((trajectory_ti + slot) + semisolid_plane), cti, state_mask); + } + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + dt_tangent = _event_value(dot_duration, event_base, event, active_atom, single_train); + wout_value = 1.0f; + wout_tangent = 0.0f; + if (bsk::truth(moving)) { + auto t0_ = _washout_jvp(atom_washout, d_washout, dt_value, dt_tangent); + wout_value = bsk::get<0>(t0_); + wout_tangent = bsk::get<1>(t0_); + } + dry1_value = bsk::exp(((-r1_value) * dt_value)); + dry1_tangent = ((-dry1_value) * ((r1_value * dt_tangent) + (r1_tangent * dt_value))); + dry2_value = bsk::exp(((-r2_value) * dt_value)); + dry2_tangent = ((-dry2_value) * ((r2_value * dt_tangent) + (r2_tangent * dt_value))); + e1_value = (dry1_value * wout_value); + e1_tangent = ((dry1_tangent * wout_value) + (dry1_value * wout_tangent)); + e2_value = (dry2_value * wout_value); + e2_tangent = ((dry2_tangent * wout_value) + (dry2_value * wout_tangent)); + damp_z = 1.0f; + damp_z_tangent = 0.0f; + damp_t = 1.0f; + damp_t_tangent = 0.0f; + if (bsk::truth(diffusing)) { + auto t1_ = _damping_jvp(atom_damping, d_damping, dt_value, dt_tangent, order); + damp_z = bsk::get<0>(t1_); + damp_z_tangent = bsk::get<1>(t1_); + damp_t = bsk::get<2>(t1_); + damp_t_tangent = bsk::get<3>(t1_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + auto t2_ = bsk::make_tup((1.0f - e1_value), (-e1_tangent)); + recovery_value = bsk::get<0>(t2_); + recovery_tangent = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup(e1_value, e1_tangent); + bare1_value = bsk::get<0>(t3_); + bare1_tangent = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(e2_value, e2_tangent); + bare2_value = bsk::get<0>(t4_); + bare2_tangent = bsk::get<1>(t4_); + e1_tangent = ((e1_tangent * damp_z) + (bare1_value * damp_z_tangent)); + e1_value = (bare1_value * damp_z); + e2_tangent = ((e2_tangent * damp_t) + (bare2_value * damp_t_tangent)); + e2_value = (bare2_value * damp_t); + turn_t = 0.0f; + dturn_t = 0.0f; + auto t5_ = bsk::make_tup(1.0f, 0.0f, 0.0f, 0.0f); + szr = bsk::get<0>(t5_); + szi = bsk::get<1>(t5_); + sztr = bsk::get<2>(t5_); + szti = bsk::get<3>(t5_); + if (bsk::truth(moving)) { + auto t6_ = _flow(atom_flow, dt_value, order); + turn_z = bsk::get<0>(t6_); + turn_t = bsk::get<1>(t6_); + d_turn = ((d_flow * dt_value) + (atom_flow * dt_tangent)); + dturn_z = ((-order) * d_turn); + dturn_t = ((-(order + 0.5f)) * d_turn); + auto t7_ = _dual_polar(turn_z, dturn_z); + szr = bsk::get<0>(t7_); + szi = bsk::get<1>(t7_); + sztr = bsk::get<2>(t7_); + szti = bsk::get<3>(t7_); + } + auto t8_ = bsk::make_tup(1.0f, 0.0f, 0.0f, 0.0f); + qr = bsk::get<0>(t8_); + qi = bsk::get<1>(t8_); + qtr = bsk::get<2>(t8_); + qti = bsk::get<3>(t8_); + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + angle_value = ((-6.283185307179586f * (atom_b0 * dt_value)) + turn_t); + angle_tangent = ((-6.283185307179586f * ((d_b0 * dt_value) + (atom_b0 * dt_tangent))) + dturn_t); + auto t9_ = _dual_polar(angle_value, angle_tangent); + qr = bsk::get<0>(t9_); + qi = bsk::get<1>(t9_); + qtr = bsk::get<2>(t9_); + qti = bsk::get<3>(t9_); + } + auto t10_ = _dual_scale(e2_value, e2_tangent, qr, qi, qtr, qti); + ovr = bsk::get<0>(t10_); + ovi = bsk::get<1>(t10_); + otr = bsk::get<2>(t10_); + oti = bsk::get<3>(t10_); + auto t11_ = _dual_scale(e1_value, e1_tangent, szr, szi, sztr, szti); + lvr = bsk::get<0>(t11_); + lvi = bsk::get<1>(t11_); + ltr = bsk::get<2>(t11_); + lti = bsk::get<3>(t11_); + // The damping and the off-resonance turn both pools take; with an + // exchanging one the relaxation itself sits inside the operator instead + // of in the scalar the free pool alone multiplies by. + carried = _dual_scale(damp_t, damp_t_tangent, qr, qi, qtr, qti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + across = _two_pool_transverse_step_jvp(r2_value, r2_tangent, r2b_value, r2b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, atom_free, d_free, atom_shift, d_shift, dt_value, dt_tangent, wout_value, wout_tangent); + a11 = bsk::make_tup(bsk::get<0>(across), bsk::get<1>(across), bsk::get<8>(across), bsk::get<9>(across)); + a12 = bsk::make_tup(bsk::get<2>(across), bsk::get<3>(across), bsk::get<10>(across), bsk::get<11>(across)); + a21 = bsk::make_tup(bsk::get<4>(across), bsk::get<5>(across), bsk::get<12>(across), bsk::get<13>(across)); + a22 = bsk::make_tup(bsk::get<6>(across), bsk::get<7>(across), bsk::get<14>(across), bsk::get<15>(across)); + free_plus = bsk::make_tup(pvr, pvi, ptr, pti); + pool_plus = bsk::make_tup(bpvr, bpvi, bptr, bpti); + free_minus = bsk::make_tup(mvr, mvi, mtr, mti); + pool_minus = bsk::make_tup(bmvr, bmvi, bmtr, bmti); + conjugated = _dual_conj(carried); + // ``F-`` takes the conjugate of the operator entry by entry, not + // its transpose: it is the conjugate state following the conjugate + // map. + auto t12_ = _dual_product(_dual_add(_dual_product(a11, free_plus), _dual_product(a12, pool_plus)), carried); + pvr = bsk::get<0>(t12_); + pvi = bsk::get<1>(t12_); + ptr = bsk::get<2>(t12_); + pti = bsk::get<3>(t12_); + auto t13_ = _dual_product(_dual_add(_dual_product(a21, free_plus), _dual_product(a22, pool_plus)), carried); + bpvr = bsk::get<0>(t13_); + bpvi = bsk::get<1>(t13_); + bptr = bsk::get<2>(t13_); + bpti = bsk::get<3>(t13_); + auto t14_ = _dual_product(_dual_add(_dual_product(_dual_conj(a11), free_minus), _dual_product(_dual_conj(a12), pool_minus)), conjugated); + mvr = bsk::get<0>(t14_); + mvi = bsk::get<1>(t14_); + mtr = bsk::get<2>(t14_); + mti = bsk::get<3>(t14_); + auto t15_ = _dual_product(_dual_add(_dual_product(_dual_conj(a21), free_minus), _dual_product(_dual_conj(a22), pool_minus)), conjugated); + bmvr = bsk::get<0>(t15_); + bmvi = bsk::get<1>(t15_); + bmtr = bsk::get<2>(t15_); + bmti = bsk::get<3>(t15_); + } else { + auto t16_ = _dual_mul(ovr, ovi, otr, oti, pvr, pvi, ptr, pti); + pvr = bsk::get<0>(t16_); + pvi = bsk::get<1>(t16_); + ptr = bsk::get<2>(t16_); + pti = bsk::get<3>(t16_); + auto t17_ = _dual_mul(ovr, (-ovi), otr, (-oti), mvr, mvi, mtr, mti); + mvr = bsk::get<0>(t17_); + mvi = bsk::get<1>(t17_); + mtr = bsk::get<2>(t17_); + mti = bsk::get<3>(t17_); + } + if (bsk::truth((pools == 3))) { + // Three pools mix through a 3x3 formed in double, tangent and all: + // a direction through an operator this ill-conditioned needs the + // width as much as the value does. + if (bsk::truth(tabulated)) { + auto t18_ = _three_pool_from_table_jvp(pool_table, bsk::ld(((duration_row + event_base) + event), active_atom, 0), atom, atom_count, active_atom, r1_value, r1b_value, r1c_value, atom_exchange, atom_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_tangent, wout_value, wout_tangent); + t11 = bsk::get<0>(t18_); + t12 = bsk::get<1>(t18_); + t13 = bsk::get<2>(t18_); + t21 = bsk::get<3>(t18_); + t22 = bsk::get<4>(t18_); + t23 = bsk::get<5>(t18_); + t31 = bsk::get<6>(t18_); + t32 = bsk::get<7>(t18_); + t33 = bsk::get<8>(t18_); + grow_free = bsk::get<9>(t18_); + grow_pool_b = bsk::get<10>(t18_); + grow_semisolid = bsk::get<11>(t18_); + d_t11 = bsk::get<12>(t18_); + d_t12 = bsk::get<13>(t18_); + d_t13 = bsk::get<14>(t18_); + d_t21 = bsk::get<15>(t18_); + d_t22 = bsk::get<16>(t18_); + d_t23 = bsk::get<17>(t18_); + d_t31 = bsk::get<18>(t18_); + d_t32 = bsk::get<19>(t18_); + d_t33 = bsk::get<20>(t18_); + d_grow_free = bsk::get<21>(t18_); + d_grow_pool_b = bsk::get<22>(t18_); + d_grow_semisolid = bsk::get<23>(t18_); + } else { + auto t19_ = _three_pool_step_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_value, dt_tangent, wout_value, wout_tangent, narrow); + t11 = bsk::get<0>(t19_); + t12 = bsk::get<1>(t19_); + t13 = bsk::get<2>(t19_); + t21 = bsk::get<3>(t19_); + t22 = bsk::get<4>(t19_); + t23 = bsk::get<5>(t19_); + t31 = bsk::get<6>(t19_); + t32 = bsk::get<7>(t19_); + t33 = bsk::get<8>(t19_); + grow_free = bsk::get<9>(t19_); + grow_pool_b = bsk::get<10>(t19_); + grow_semisolid = bsk::get<11>(t19_); + d_t11 = bsk::get<12>(t19_); + d_t12 = bsk::get<13>(t19_); + d_t13 = bsk::get<14>(t19_); + d_t21 = bsk::get<15>(t19_); + d_t22 = bsk::get<16>(t19_); + d_t23 = bsk::get<17>(t19_); + d_t31 = bsk::get<18>(t19_); + d_t32 = bsk::get<19>(t19_); + d_t33 = bsk::get<20>(t19_); + d_grow_free = bsk::get<21>(t19_); + d_grow_pool_b = bsk::get<22>(t19_); + d_grow_semisolid = bsk::get<23>(t19_); + } + spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti); + auto was_free = bsk::make_tup(zvr, zvi, ztr, zti); + auto was_pool_b = bsk::make_tup(bvr, bvi, btr, bti); + auto was_semisolid = bsk::make_tup(cvr, cvi, ctr, cti); + mixed_free = _dual_add(_dual_add(_dual_scale(t11, d_t11, bsk::get<0>(was_free), bsk::get<1>(was_free), bsk::get<2>(was_free), bsk::get<3>(was_free)), _dual_scale(t12, d_t12, bsk::get<0>(was_pool_b), bsk::get<1>(was_pool_b), bsk::get<2>(was_pool_b), bsk::get<3>(was_pool_b))), _dual_scale(t13, d_t13, bsk::get<0>(was_semisolid), bsk::get<1>(was_semisolid), bsk::get<2>(was_semisolid), bsk::get<3>(was_semisolid))); + auto mixed_pool_b = _dual_add(_dual_add(_dual_scale(t21, d_t21, bsk::get<0>(was_free), bsk::get<1>(was_free), bsk::get<2>(was_free), bsk::get<3>(was_free)), _dual_scale(t22, d_t22, bsk::get<0>(was_pool_b), bsk::get<1>(was_pool_b), bsk::get<2>(was_pool_b), bsk::get<3>(was_pool_b))), _dual_scale(t23, d_t23, bsk::get<0>(was_semisolid), bsk::get<1>(was_semisolid), bsk::get<2>(was_semisolid), bsk::get<3>(was_semisolid))); + mixed_semisolid = _dual_add(_dual_add(_dual_scale(t31, d_t31, bsk::get<0>(was_free), bsk::get<1>(was_free), bsk::get<2>(was_free), bsk::get<3>(was_free)), _dual_scale(t32, d_t32, bsk::get<0>(was_pool_b), bsk::get<1>(was_pool_b), bsk::get<2>(was_pool_b), bsk::get<3>(was_pool_b))), _dual_scale(t33, d_t33, bsk::get<0>(was_semisolid), bsk::get<1>(was_semisolid), bsk::get<2>(was_semisolid), bsk::get<3>(was_semisolid))); + auto t20_ = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(mixed_free), bsk::get<1>(mixed_free), bsk::get<2>(mixed_free), bsk::get<3>(mixed_free)); + zvr = bsk::get<0>(t20_); + zvi = bsk::get<1>(t20_); + ztr = bsk::get<2>(t20_); + zti = bsk::get<3>(t20_); + auto t21_ = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(mixed_pool_b), bsk::get<1>(mixed_pool_b), bsk::get<2>(mixed_pool_b), bsk::get<3>(mixed_pool_b)); + bvr = bsk::get<0>(t21_); + bvi = bsk::get<1>(t21_); + btr = bsk::get<2>(t21_); + bti = bsk::get<3>(t21_); + auto t22_ = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(mixed_semisolid), bsk::get<1>(mixed_semisolid), bsk::get<2>(mixed_semisolid), bsk::get<3>(mixed_semisolid)); + cvr = bsk::get<0>(t22_); + cvi = bsk::get<1>(t22_); + ctr = bsk::get<2>(t22_); + cti = bsk::get<3>(t22_); + zvr = (zvr + bsk::where((state == 0), grow_free, 0.0f)); + ztr = (ztr + bsk::where((state == 0), d_grow_free, 0.0f)); + bvr = (bvr + bsk::where((state == 0), grow_pool_b, 0.0f)); + btr = (btr + bsk::where((state == 0), d_grow_pool_b, 0.0f)); + cvr = (cvr + bsk::where((state == 0), grow_semisolid, 0.0f)); + ctr = (ctr + bsk::where((state == 0), d_grow_semisolid, 0.0f)); + } else if (bsk::truth((pools > 0))) { + // The exchange operator is a property of the interval, not of a + // dephasing order, so it is formed once and the per-order damping + // and turn multiply it. + auto t23_ = _two_pool_step_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, dt_value, dt_tangent, wout_value, wout_tangent); + pe11 = bsk::get<0>(t23_); + pe12 = bsk::get<1>(t23_); + pe21 = bsk::get<2>(t23_); + pe22 = bsk::get<3>(t23_); + prec_f = bsk::get<4>(t23_); + prec_b = bsk::get<5>(t23_); + de11 = bsk::get<6>(t23_); + de12 = bsk::get<7>(t23_); + de21 = bsk::get<8>(t23_); + de22 = bsk::get<9>(t23_); + drec_f = bsk::get<10>(t23_); + drec_b = bsk::get<11>(t23_); + spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti); + free_part = _dual_scale(pe11, de11, zvr, zvi, ztr, zti); + cross_in = _dual_scale(pe12, de12, bvr, bvi, btr, bti); + cross_out = _dual_scale(pe21, de21, zvr, zvi, ztr, zti); + bound_part = _dual_scale(pe22, de22, bvr, bvi, btr, bti); + auto t24_ = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), (bsk::get<0>(free_part) + bsk::get<0>(cross_in)), (bsk::get<1>(free_part) + bsk::get<1>(cross_in)), (bsk::get<2>(free_part) + bsk::get<2>(cross_in)), (bsk::get<3>(free_part) + bsk::get<3>(cross_in))); + zvr = bsk::get<0>(t24_); + zvi = bsk::get<1>(t24_); + ztr = bsk::get<2>(t24_); + zti = bsk::get<3>(t24_); + auto t25_ = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), (bsk::get<0>(cross_out) + bsk::get<0>(bound_part)), (bsk::get<1>(cross_out) + bsk::get<1>(bound_part)), (bsk::get<2>(cross_out) + bsk::get<2>(bound_part)), (bsk::get<3>(cross_out) + bsk::get<3>(bound_part))); + bvr = bsk::get<0>(t25_); + bvi = bsk::get<1>(t25_); + btr = bsk::get<2>(t25_); + bti = bsk::get<3>(t25_); + zvr = (zvr + bsk::where((state == 0), prec_f, 0.0f)); + ztr = (ztr + bsk::where((state == 0), drec_f, 0.0f)); + bvr = (bvr + bsk::where((state == 0), prec_b, 0.0f)); + btr = (btr + bsk::where((state == 0), drec_b, 0.0f)); + } else { + auto t26_ = _dual_mul(lvr, lvi, ltr, lti, zvr, zvi, ztr, zti); + zvr = bsk::get<0>(t26_); + zvi = bsk::get<1>(t26_); + ztr = bsk::get<2>(t26_); + zti = bsk::get<3>(t26_); + zvr = (zvr + bsk::where((state == 0), recovery_value, 0.0f)); + ztr = (ztr + bsk::where((state == 0), recovery_tangent, 0.0f)); + } + event_action = bsk::cast(bsk::ld((action + event))); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t27_ = _shift(pvr, pvi, mvr, mvi, state, state_mask, state_count); + svr = bsk::get<0>(t27_); + svi = bsk::get<1>(t27_); + wvr = bsk::get<2>(t27_); + wvi = bsk::get<3>(t27_); + auto t28_ = _shift(ptr, pti, mtr, mti, state, state_mask, state_count); + str_ = bsk::get<0>(t28_); + sti = bsk::get<1>(t28_); + wtr = bsk::get<2>(t28_); + wti = bsk::get<3>(t28_); + pvr = bsk::where(pre_shift, svr, pvr); + pvi = bsk::where(pre_shift, svi, pvi); + ptr = bsk::where(pre_shift, str_, ptr); + pti = bsk::where(pre_shift, sti, pti); + mvr = bsk::where(pre_shift, wvr, mvr); + mvi = bsk::where(pre_shift, wvi, mvi); + mtr = bsk::where(pre_shift, wtr, mtr); + mti = bsk::where(pre_shift, wti, mti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t29_ = _shift(bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count); + svr = bsk::get<0>(t29_); + svi = bsk::get<1>(t29_); + wvr = bsk::get<2>(t29_); + wvi = bsk::get<3>(t29_); + auto t30_ = _shift(bptr, bpti, bmtr, bmti, state, state_mask, state_count); + str_ = bsk::get<0>(t30_); + sti = bsk::get<1>(t30_); + wtr = bsk::get<2>(t30_); + wti = bsk::get<3>(t30_); + bpvr = bsk::where(pre_shift, svr, bpvr); + bpvi = bsk::where(pre_shift, svi, bpvi); + bptr = bsk::where(pre_shift, str_, bptr); + bpti = bsk::where(pre_shift, sti, bpti); + bmvr = bsk::where(pre_shift, wvr, bmvr); + bmvi = bsk::where(pre_shift, wvi, bmvi); + bmtr = bsk::where(pre_shift, wtr, bmtr); + bmti = bsk::where(pre_shift, wti, bmti); + } + event_kind = bsk::ld((kind + event)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + auto t31_ = _dual_scale((-atom_inv), (-d_inv), zvr, zvi, ztr, zti); + ivr = bsk::get<0>(t31_); + ivi = bsk::get<1>(t31_); + itr = bsk::get<2>(t31_); + iti = bsk::get<3>(t31_); + zvr = bsk::where(invert, ivr, zvr); + zvi = bsk::where(invert, ivi, zvi); + ztr = bsk::where(invert, itr, ztr); + zti = bsk::where(invert, iti, zti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // A semisolid pool is saturated by an adiabatic sweep rather than + // turned over; a chemically exchanging one is free water and + // inverts like any other. + auto t32_ = _dual_scale((-atom_inv), (-d_inv), bvr, bvi, btr, bti); + ivr = bsk::get<0>(t32_); + ivi = bsk::get<1>(t32_); + itr = bsk::get<2>(t32_); + iti = bsk::get<3>(t32_); + bvr = bsk::where(invert, ivr, bvr); + bvi = bsk::where(invert, ivi, bvi); + btr = bsk::where(invert, itr, btr); + bti = bsk::where(invert, iti, bti); + } + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_dot_flip = _event_value(dot_flip, event_base, event, active_atom, single_train); + event_phase = _event_value(phase, event_base, event, active_atom, single_train); + event_dot_phase = _event_value(dot_phase, event_base, event, active_atom, single_train); + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse a row of its own. + if (bsk::truth(shimmed)) { + row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + } + d_b1 = bsk::ld(((dot_b1 + row) + atom), active_atom, 0.0f); + if (bsk::truth(off_axis)) { + d_b1_phase = bsk::ld(((dot_b1_phase + row) + atom), active_atom, 0.0f); + } + } + alpha_value = (event_flip * atom_b1); + alpha_tangent = ((event_dot_flip * atom_b1) + (event_flip * d_b1)); + phi_value = (event_phase + atom_b1_phase); + phi_tangent = (event_dot_phase + d_b1_phase); + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + // The semisolid pool absorbs the power the pulse deposits, so it + // reads the bare flip the transmit field gives the voxel -- not the + // slice-shaped rotation the free pool takes from the table. + offset_value = (bsk::ld((rf_frequency + event)) - atom_b0); + auto t33_ = _lineshape_at_slope(lineshape, offset_value, lineshape_bins, lineshape_step); + shape_value = bsk::get<0>(t33_); + shape_slope = bsk::get<1>(t33_); + shape_tangent = (shape_slope * (-d_b0)); + event_saturation = bsk::ld((saturation + event)); + power_value = ((event_saturation * alpha_value) * alpha_value); + power_tangent = (((event_saturation * 2.0f) * alpha_value) * alpha_tangent); + absorbed_value = bsk::exp((power_value * shape_value)); + absorbed_tangent = (absorbed_value * ((power_tangent * shape_value) + (power_value * shape_tangent))); + saturating = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((pools == 1))) { + auto sat_b = _dual_scale(absorbed_value, absorbed_tangent, bvr, bvi, btr, bti); + bvr = bsk::where(saturating, bsk::get<0>(sat_b), bvr); + bvi = bsk::where(saturating, bsk::get<1>(sat_b), bvi); + btr = bsk::where(saturating, bsk::get<2>(sat_b), btr); + bti = bsk::where(saturating, bsk::get<3>(sat_b), bti); + } else { + auto sat_c = _dual_scale(absorbed_value, absorbed_tangent, cvr, cvi, ctr, cti); + cvr = bsk::where(saturating, bsk::get<0>(sat_c), cvr); + cvi = bsk::where(saturating, bsk::get<1>(sat_c), cvi); + ctr = bsk::where(saturating, bsk::get<2>(sat_c), ctr); + cti = bsk::where(saturating, bsk::get<3>(sat_c), cti); + } + } + cos_value = bsk::cos(alpha_value); + sin_value = bsk::sin(alpha_value); + cos_tangent = ((-sin_value) * alpha_tangent); + sin_tangent = (cos_value * alpha_tangent); + auto t34_ = _dual_polar(phi_value, phi_tangent); + p1r = bsk::get<0>(t34_); + p1i = bsk::get<1>(t34_); + p1tr = bsk::get<2>(t34_); + p1ti = bsk::get<3>(t34_); + auto t35_ = _dual_mul(p1r, p1i, p1tr, p1ti, p1r, p1i, p1tr, p1ti); + p2r = bsk::get<0>(t35_); + p2i = bsk::get<1>(t35_); + p2tr = bsk::get<2>(t35_); + p2ti = bsk::get<3>(t35_); + auto t36_ = _rotation_block((0.5f * (1.0f + cos_value)), (0.5f * cos_tangent), (0.5f * (1.0f - cos_value)), (-0.5f * cos_tangent), sin_value, sin_tangent, cos_value, cos_tangent, p1r, p1i, p1tr, p1ti, p2r, p2i, p2tr, p2ti, p1r, (-p1i), p1tr, (-p1ti)); + t00 = bsk::get<0>(t36_); + t01 = bsk::get<1>(t36_); + t02 = bsk::get<2>(t36_); + r12 = bsk::get<3>(t36_); + t20 = bsk::get<4>(t36_); + r21 = bsk::get<5>(t36_); + r22 = bsk::get<6>(t36_); + auto a0 = _dual_mul(bsk::get<0>(t00), bsk::get<1>(t00), bsk::get<2>(t00), bsk::get<3>(t00), pvr, pvi, ptr, pti); + auto a1 = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), mvr, mvi, mtr, mti); + auto a2 = _dual_mul(bsk::get<0>(t02), bsk::get<1>(t02), bsk::get<2>(t02), bsk::get<3>(t02), zvr, zvi, ztr, zti); + auto b0_ = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), pvr, pvi, ptr, pti); + auto b1_ = _dual_mul(bsk::get<0>(t00), bsk::get<1>(t00), bsk::get<2>(t00), bsk::get<3>(t00), mvr, mvi, mtr, mti); + auto b2 = _dual_mul(bsk::get<0>(r12), bsk::get<1>(r12), bsk::get<2>(r12), bsk::get<3>(r12), zvr, zvi, ztr, zti); + auto c0 = _dual_mul(bsk::get<0>(t20), bsk::get<1>(t20), bsk::get<2>(t20), bsk::get<3>(t20), pvr, pvi, ptr, pti); + auto c1 = _dual_mul(bsk::get<0>(r21), bsk::get<1>(r21), bsk::get<2>(r21), bsk::get<3>(r21), mvr, mvi, mtr, mti); + auto c2 = _dual_mul(bsk::get<0>(r22), bsk::get<1>(r22), bsk::get<2>(r22), bsk::get<3>(r22), zvr, zvi, ztr, zti); + turned_pvr = ((bsk::get<0>(a0) + bsk::get<0>(a1)) + bsk::get<0>(a2)); + turned_pvi = ((bsk::get<1>(a0) + bsk::get<1>(a1)) + bsk::get<1>(a2)); + turned_ptr = ((bsk::get<2>(a0) + bsk::get<2>(a1)) + bsk::get<2>(a2)); + turned_pti = ((bsk::get<3>(a0) + bsk::get<3>(a1)) + bsk::get<3>(a2)); + turned_mvr = ((bsk::get<0>(b0_) + bsk::get<0>(b1_)) + bsk::get<0>(b2)); + turned_mvi = ((bsk::get<1>(b0_) + bsk::get<1>(b1_)) + bsk::get<1>(b2)); + turned_mtr = ((bsk::get<2>(b0_) + bsk::get<2>(b1_)) + bsk::get<2>(b2)); + turned_mti = ((bsk::get<3>(b0_) + bsk::get<3>(b1_)) + bsk::get<3>(b2)); + turned_zvr = ((bsk::get<0>(c0) + bsk::get<0>(c1)) + bsk::get<0>(c2)); + turned_zvi = ((bsk::get<1>(c0) + bsk::get<1>(c1)) + bsk::get<1>(c2)); + turned_ztr = ((bsk::get<2>(c0) + bsk::get<2>(c1)) + bsk::get<2>(c2)); + turned_zti = ((bsk::get<3>(c0) + bsk::get<3>(c1)) + bsk::get<3>(c2)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + if (bsk::truth(dynamic)) { + auto t37_ = _dynamic_pair_dual_at(pairs, pair_direction, pair_index, event_base, event, atom, atom_count, active_atom, phi_value, phi_tangent, directed); + shaped_a = bsk::get<0>(t37_); + shaped_b = bsk::get<1>(t37_); + } else { + auto t38_ = _profiled_pair_dual(profile, _table_row(profile_index, event, location, locations), alpha_value, alpha_tangent, phi_value, phi_tangent, profile_bins, profile_step); + shaped_a = bsk::get<0>(t38_); + shaped_b = bsk::get<1>(t38_); + } + auto t39_ = _rotate_spinor_dual(bsk::get<0>(shaped_a), bsk::get<1>(shaped_a), bsk::get<0>(shaped_b), bsk::get<1>(shaped_b), bsk::get<2>(shaped_a), bsk::get<3>(shaped_a), bsk::get<2>(shaped_b), bsk::get<3>(shaped_b), pvr, pvi, mvr, mvi, zvr, zvi, ptr, pti, mtr, mti, ztr, zti); + turned_pvr = bsk::get<0>(t39_); + turned_pvi = bsk::get<1>(t39_); + turned_mvr = bsk::get<2>(t39_); + turned_mvi = bsk::get<3>(t39_); + turned_zvr = bsk::get<4>(t39_); + turned_zvi = bsk::get<5>(t39_); + turned_ptr = bsk::get<6>(t39_); + turned_pti = bsk::get<7>(t39_); + turned_mtr = bsk::get<8>(t39_); + turned_mti = bsk::get<9>(t39_); + turned_ztr = bsk::get<10>(t39_); + turned_zti = bsk::get<11>(t39_); + } + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // The same pulse, the same rotation. A chemical shift moves where a + // pool precesses, not what a pulse does to it. + auto e0 = _dual_mul(bsk::get<0>(t00), bsk::get<1>(t00), bsk::get<2>(t00), bsk::get<3>(t00), bpvr, bpvi, bptr, bpti); + auto e1_ = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), bmvr, bmvi, bmtr, bmti); + auto e2_ = _dual_mul(bsk::get<0>(t02), bsk::get<1>(t02), bsk::get<2>(t02), bsk::get<3>(t02), bvr, bvi, btr, bti); + auto f0 = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), bpvr, bpvi, bptr, bpti); + auto f1 = _dual_mul(bsk::get<0>(t00), bsk::get<1>(t00), bsk::get<2>(t00), bsk::get<3>(t00), bmvr, bmvi, bmtr, bmti); + auto f2 = _dual_mul(bsk::get<0>(r12), bsk::get<1>(r12), bsk::get<2>(r12), bsk::get<3>(r12), bvr, bvi, btr, bti); + auto h0 = _dual_mul(bsk::get<0>(t20), bsk::get<1>(t20), bsk::get<2>(t20), bsk::get<3>(t20), bpvr, bpvi, bptr, bpti); + auto h1 = _dual_mul(bsk::get<0>(r21), bsk::get<1>(r21), bsk::get<2>(r21), bsk::get<3>(r21), bmvr, bmvi, bmtr, bmti); + auto h2 = _dual_mul(bsk::get<0>(r22), bsk::get<1>(r22), bsk::get<2>(r22), bsk::get<3>(r22), bvr, bvi, btr, bti); + spun_pvr = ((bsk::get<0>(e0) + bsk::get<0>(e1_)) + bsk::get<0>(e2_)); + spun_pvi = ((bsk::get<1>(e0) + bsk::get<1>(e1_)) + bsk::get<1>(e2_)); + spun_ptr = ((bsk::get<2>(e0) + bsk::get<2>(e1_)) + bsk::get<2>(e2_)); + spun_pti = ((bsk::get<3>(e0) + bsk::get<3>(e1_)) + bsk::get<3>(e2_)); + spun_mvr = ((bsk::get<0>(f0) + bsk::get<0>(f1)) + bsk::get<0>(f2)); + spun_mvi = ((bsk::get<1>(f0) + bsk::get<1>(f1)) + bsk::get<1>(f2)); + spun_mtr = ((bsk::get<2>(f0) + bsk::get<2>(f1)) + bsk::get<2>(f2)); + spun_mti = ((bsk::get<3>(f0) + bsk::get<3>(f1)) + bsk::get<3>(f2)); + spun_zvr = ((bsk::get<0>(h0) + bsk::get<0>(h1)) + bsk::get<0>(h2)); + spun_zvi = ((bsk::get<1>(h0) + bsk::get<1>(h1)) + bsk::get<1>(h2)); + spun_ztr = ((bsk::get<2>(h0) + bsk::get<2>(h1)) + bsk::get<2>(h2)); + spun_zti = ((bsk::get<3>(h0) + bsk::get<3>(h1)) + bsk::get<3>(h2)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + auto t40_ = _rotate_spinor_dual(bsk::get<0>(shaped_a), bsk::get<1>(shaped_a), bsk::get<0>(shaped_b), bsk::get<1>(shaped_b), bsk::get<2>(shaped_a), bsk::get<3>(shaped_a), bsk::get<2>(shaped_b), bsk::get<3>(shaped_b), bpvr, bpvi, bmvr, bmvi, bvr, bvi, bptr, bpti, bmtr, bmti, btr, bti); + spun_pvr = bsk::get<0>(t40_); + spun_pvi = bsk::get<1>(t40_); + spun_mvr = bsk::get<2>(t40_); + spun_mvi = bsk::get<3>(t40_); + spun_zvr = bsk::get<4>(t40_); + spun_zvi = bsk::get<5>(t40_); + spun_ptr = bsk::get<6>(t40_); + spun_pti = bsk::get<7>(t40_); + spun_mtr = bsk::get<8>(t40_); + spun_mti = bsk::get<9>(t40_); + spun_ztr = bsk::get<10>(t40_); + spun_zti = bsk::get<11>(t40_); + } + bpvr = bsk::where(rotate, spun_pvr, bpvr); + bpvi = bsk::where(rotate, spun_pvi, bpvi); + bptr = bsk::where(rotate, spun_ptr, bptr); + bpti = bsk::where(rotate, spun_pti, bpti); + bmvr = bsk::where(rotate, spun_mvr, bmvr); + bmvi = bsk::where(rotate, spun_mvi, bmvi); + bmtr = bsk::where(rotate, spun_mtr, bmtr); + bmti = bsk::where(rotate, spun_mti, bmti); + bvr = bsk::where(rotate, spun_zvr, bvr); + bvi = bsk::where(rotate, spun_zvi, bvi); + btr = bsk::where(rotate, spun_ztr, btr); + bti = bsk::where(rotate, spun_zti, bti); + } + pvr = bsk::where(rotate, turned_pvr, pvr); + pvi = bsk::where(rotate, turned_pvi, pvi); + ptr = bsk::where(rotate, turned_ptr, ptr); + pti = bsk::where(rotate, turned_pti, pti); + mvr = bsk::where(rotate, turned_mvr, mvr); + mvi = bsk::where(rotate, turned_mvi, mvi); + mtr = bsk::where(rotate, turned_mtr, mtr); + mti = bsk::where(rotate, turned_mti, mti); + zvr = bsk::where(rotate, turned_zvr, zvr); + zvi = bsk::where(rotate, turned_zvi, zvi); + ztr = bsk::where(rotate, turned_ztr, ztr); + zti = bsk::where(rotate, turned_zti, zti); + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + auto t41_ = _shift(pvr, pvi, mvr, mvi, state, state_mask, state_count); + svr = bsk::get<0>(t41_); + svi = bsk::get<1>(t41_); + wvr = bsk::get<2>(t41_); + wvi = bsk::get<3>(t41_); + auto t42_ = _shift(ptr, pti, mtr, mti, state, state_mask, state_count); + str_ = bsk::get<0>(t42_); + sti = bsk::get<1>(t42_); + wtr = bsk::get<2>(t42_); + wti = bsk::get<3>(t42_); + pvr = bsk::where(do_shift, svr, pvr); + pvi = bsk::where(do_shift, svi, pvi); + ptr = bsk::where(do_shift, str_, ptr); + pti = bsk::where(do_shift, sti, pti); + mvr = bsk::where(do_shift, wvr, mvr); + mvi = bsk::where(do_shift, wvi, mvi); + mtr = bsk::where(do_shift, wtr, mtr); + mti = bsk::where(do_shift, wti, mti); + spoil = (bsk::band(event_action, 8) != 0); + pvr = bsk::where(spoil, 0.0f, pvr); + pvi = bsk::where(spoil, 0.0f, pvi); + ptr = bsk::where(spoil, 0.0f, ptr); + pti = bsk::where(spoil, 0.0f, pti); + mvr = bsk::where(spoil, 0.0f, mvr); + mvi = bsk::where(spoil, 0.0f, mvi); + mtr = bsk::where(spoil, 0.0f, mtr); + mti = bsk::where(spoil, 0.0f, mti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t43_ = _shift(bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count); + svr = bsk::get<0>(t43_); + svi = bsk::get<1>(t43_); + wvr = bsk::get<2>(t43_); + wvi = bsk::get<3>(t43_); + auto t44_ = _shift(bptr, bpti, bmtr, bmti, state, state_mask, state_count); + str_ = bsk::get<0>(t44_); + sti = bsk::get<1>(t44_); + wtr = bsk::get<2>(t44_); + wti = bsk::get<3>(t44_); + bpvr = bsk::where(spoil, 0.0f, bsk::where(do_shift, svr, bpvr)); + bpvi = bsk::where(spoil, 0.0f, bsk::where(do_shift, svi, bpvi)); + bptr = bsk::where(spoil, 0.0f, bsk::where(do_shift, str_, bptr)); + bpti = bsk::where(spoil, 0.0f, bsk::where(do_shift, sti, bpti)); + bmvr = bsk::where(spoil, 0.0f, bsk::where(do_shift, wvr, bmvr)); + bmvi = bsk::where(spoil, 0.0f, bsk::where(do_shift, wvi, bmvi)); + bmtr = bsk::where(spoil, 0.0f, bsk::where(do_shift, wtr, bmtr)); + bmti = bsk::where(spoil, 0.0f, bsk::where(do_shift, wti, bmti)); + } + } + return; + } + // ---- reverse ---- + pbvr = empty; + pbvi = empty; + pbtr = empty; + pbti = empty; + mbvr = empty; + mbvi = empty; + mbtr = empty; + mbti = empty; + zbvr = empty; + zbvi = empty; + zbtr = empty; + zbti = empty; + bbvr = empty; + bbvi = empty; + bbtr = empty; + bbti = empty; + ubvr = empty; + ubvi = empty; + ubtr = empty; + ubti = empty; + wbvr = empty; + wbvi = empty; + wbtr = empty; + wbti = empty; + cbvr = empty; + cbvi = empty; + cbtr = empty; + cbti = empty; + auto zero = bsk::full(0); + g_boundv = zero; + g_boundt = zero; + g_exchv = zero; + g_excht = zero; + g_t1bv = zero; + g_t1bt = zero; + g_t2bv = zero; + g_t2bt = zero; + g_shiftv = zero; + g_shiftt = zero; + g_semiv = zero; + g_semit = zero; + g_sexchv = zero; + g_sexcht = zero; + g_t1cv = zero; + g_t1ct = zero; + g_diffv = zero; + g_difft = zero; + g_flowv = zero; + g_flowt = zero; + g_washv = zero; + g_washt = zero; + g_t1v = zero; + g_t1t = zero; + g_t2v = zero; + g_t2t = zero; + g_m0v = zero; + g_m0t = zero; + g_b1v = zero; + g_b1t = zero; + g_b1pv = zero; + g_b1pt = zero; + g_b0v = zero; + g_b0t = zero; + g_invv = zero; + g_invt = zero; + for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + event = ((event_count - 1) - reverse); + slot = (trajectory + (event * record_stride)); + auto xpvr = bsk::ld((trajectory_vr + slot), state_mask, 0.0f); + auto xpvi = bsk::ld((trajectory_vi + slot), state_mask, 0.0f); + auto xptr = bsk::ld((trajectory_tr + slot), state_mask, 0.0f); + auto xpti = bsk::ld((trajectory_ti + slot), state_mask, 0.0f); + auto xmvr = bsk::ld(((trajectory_vr + slot) + minus_plane), state_mask, 0.0f); + auto xmvi = bsk::ld(((trajectory_vi + slot) + minus_plane), state_mask, 0.0f); + auto xmtr = bsk::ld(((trajectory_tr + slot) + minus_plane), state_mask, 0.0f); + auto xmti = bsk::ld(((trajectory_ti + slot) + minus_plane), state_mask, 0.0f); + auto xzvr = bsk::ld(((trajectory_vr + slot) + long_plane), state_mask, 0.0f); + auto xzvi = bsk::ld(((trajectory_vi + slot) + long_plane), state_mask, 0.0f); + auto xztr = bsk::ld(((trajectory_tr + slot) + long_plane), state_mask, 0.0f); + auto xzti = bsk::ld(((trajectory_ti + slot) + long_plane), state_mask, 0.0f); + xbvr = empty; + xbvi = empty; + xbtr = empty; + xbti = empty; + xbpvr = empty; + xbpvi = empty; + xbptr = empty; + xbpti = empty; + xbmvr = empty; + xbmvi = empty; + xbmtr = empty; + xbmti = empty; + xcvr = empty; + xcvi = empty; + xctr = empty; + xcti = empty; + if (bsk::truth((pools > 0))) { + xbvr = bsk::ld(((trajectory_vr + slot) + bound_plane), state_mask, 0.0f); + xbvi = bsk::ld(((trajectory_vi + slot) + bound_plane), state_mask, 0.0f); + xbtr = bsk::ld(((trajectory_tr + slot) + bound_plane), state_mask, 0.0f); + xbti = bsk::ld(((trajectory_ti + slot) + bound_plane), state_mask, 0.0f); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + xbpvr = bsk::ld(((trajectory_vr + slot) + bplus_plane), state_mask, 0.0f); + xbpvi = bsk::ld(((trajectory_vi + slot) + bplus_plane), state_mask, 0.0f); + xbptr = bsk::ld(((trajectory_tr + slot) + bplus_plane), state_mask, 0.0f); + xbpti = bsk::ld(((trajectory_ti + slot) + bplus_plane), state_mask, 0.0f); + xbmvr = bsk::ld(((trajectory_vr + slot) + bminus_plane), state_mask, 0.0f); + xbmvi = bsk::ld(((trajectory_vi + slot) + bminus_plane), state_mask, 0.0f); + xbmtr = bsk::ld(((trajectory_tr + slot) + bminus_plane), state_mask, 0.0f); + xbmti = bsk::ld(((trajectory_ti + slot) + bminus_plane), state_mask, 0.0f); + } + if (bsk::truth((pools == 3))) { + xcvr = bsk::ld(((trajectory_vr + slot) + semisolid_plane), state_mask, 0.0f); + xcvi = bsk::ld(((trajectory_vi + slot) + semisolid_plane), state_mask, 0.0f); + xctr = bsk::ld(((trajectory_tr + slot) + semisolid_plane), state_mask, 0.0f); + xcti = bsk::ld(((trajectory_ti + slot) + semisolid_plane), state_mask, 0.0f); + } + event_action = bsk::cast(bsk::ld((action + event))); + event_kind = bsk::ld((kind + event)); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + dt_tangent = _event_value(dot_duration, event_base, event, active_atom, single_train); + wout_value = 1.0f; + wout_tangent = 0.0f; + if (bsk::truth(moving)) { + auto t45_ = _washout_jvp(atom_washout, d_washout, dt_value, dt_tangent); + wout_value = bsk::get<0>(t45_); + wout_tangent = bsk::get<1>(t45_); + } + dry1_value = bsk::exp(((-r1_value) * dt_value)); + dry1_tangent = ((-dry1_value) * ((r1_value * dt_tangent) + (r1_tangent * dt_value))); + dry2_value = bsk::exp(((-r2_value) * dt_value)); + dry2_tangent = ((-dry2_value) * ((r2_value * dt_tangent) + (r2_tangent * dt_value))); + e1_value = (dry1_value * wout_value); + e1_tangent = ((dry1_tangent * wout_value) + (dry1_value * wout_tangent)); + e2_value = (dry2_value * wout_value); + e2_tangent = ((dry2_tangent * wout_value) + (dry2_value * wout_tangent)); + damp_z = 1.0f; + damp_z_tangent = 0.0f; + damp_t = 1.0f; + damp_t_tangent = 0.0f; + if (bsk::truth(diffusing)) { + auto t46_ = _damping_jvp(atom_damping, d_damping, dt_value, dt_tangent, order); + damp_z = bsk::get<0>(t46_); + damp_z_tangent = bsk::get<1>(t46_); + damp_t = bsk::get<2>(t46_); + damp_t_tangent = bsk::get<3>(t46_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + auto t47_ = bsk::make_tup((1.0f - e1_value), (-e1_tangent)); + recovery_value = bsk::get<0>(t47_); + recovery_tangent = bsk::get<1>(t47_); + auto t48_ = bsk::make_tup(e1_value, e1_tangent); + bare1_value = bsk::get<0>(t48_); + bare1_tangent = bsk::get<1>(t48_); + auto t49_ = bsk::make_tup(e2_value, e2_tangent); + bare2_value = bsk::get<0>(t49_); + bare2_tangent = bsk::get<1>(t49_); + e1_tangent = ((e1_tangent * damp_z) + (bare1_value * damp_z_tangent)); + e1_value = (bare1_value * damp_z); + e2_tangent = ((e2_tangent * damp_t) + (bare2_value * damp_t_tangent)); + e2_value = (bare2_value * damp_t); + turn_t = 0.0f; + dturn_t = 0.0f; + auto t50_ = bsk::make_tup(1.0f, 0.0f, 0.0f, 0.0f); + szr = bsk::get<0>(t50_); + szi = bsk::get<1>(t50_); + sztr = bsk::get<2>(t50_); + szti = bsk::get<3>(t50_); + if (bsk::truth(moving)) { + auto t51_ = _flow(atom_flow, dt_value, order); + turn_z = bsk::get<0>(t51_); + turn_t = bsk::get<1>(t51_); + d_turn = ((d_flow * dt_value) + (atom_flow * dt_tangent)); + dturn_z = ((-order) * d_turn); + dturn_t = ((-(order + 0.5f)) * d_turn); + auto t52_ = _dual_polar(turn_z, dturn_z); + szr = bsk::get<0>(t52_); + szi = bsk::get<1>(t52_); + sztr = bsk::get<2>(t52_); + szti = bsk::get<3>(t52_); + } + auto t53_ = bsk::make_tup(1.0f, 0.0f, 0.0f, 0.0f); + qr = bsk::get<0>(t53_); + qi = bsk::get<1>(t53_); + qtr = bsk::get<2>(t53_); + qti = bsk::get<3>(t53_); + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + angle_value = ((-6.283185307179586f * (atom_b0 * dt_value)) + turn_t); + angle_tangent = ((-6.283185307179586f * ((d_b0 * dt_value) + (atom_b0 * dt_tangent))) + dturn_t); + auto t54_ = _dual_polar(angle_value, angle_tangent); + qr = bsk::get<0>(t54_); + qi = bsk::get<1>(t54_); + qtr = bsk::get<2>(t54_); + qti = bsk::get<3>(t54_); + } + auto t55_ = _dual_scale(e2_value, e2_tangent, qr, qi, qtr, qti); + ovr = bsk::get<0>(t55_); + ovi = bsk::get<1>(t55_); + otr = bsk::get<2>(t55_); + oti = bsk::get<3>(t55_); + auto t56_ = _dual_scale(e1_value, e1_tangent, szr, szi, sztr, szti); + lvr = bsk::get<0>(t56_); + lvi = bsk::get<1>(t56_); + ltr = bsk::get<2>(t56_); + lti = bsk::get<3>(t56_); + // Replay the intra-event stages from the recorded entry state. + carried = _dual_scale(damp_t, damp_t_tangent, qr, qi, qtr, qti); + rbpvr = empty; + rbpvi = empty; + rbptr = empty; + rbpti = empty; + rbmvr = empty; + rbmvi = empty; + rbmtr = empty; + rbmti = empty; + a11 = bsk::make_tup(empty, empty, empty, empty); + a12 = bsk::make_tup(empty, empty, empty, empty); + a21 = bsk::make_tup(empty, empty, empty, empty); + a22 = bsk::make_tup(empty, empty, empty, empty); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + across = _two_pool_transverse_step_jvp(r2_value, r2_tangent, r2b_value, r2b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, atom_free, d_free, atom_shift, d_shift, dt_value, dt_tangent, wout_value, wout_tangent); + a11 = bsk::make_tup(bsk::get<0>(across), bsk::get<1>(across), bsk::get<8>(across), bsk::get<9>(across)); + a12 = bsk::make_tup(bsk::get<2>(across), bsk::get<3>(across), bsk::get<10>(across), bsk::get<11>(across)); + a21 = bsk::make_tup(bsk::get<4>(across), bsk::get<5>(across), bsk::get<12>(across), bsk::get<13>(across)); + a22 = bsk::make_tup(bsk::get<6>(across), bsk::get<7>(across), bsk::get<14>(across), bsk::get<15>(across)); + free_plus = bsk::make_tup(xpvr, xpvi, xptr, xpti); + pool_plus = bsk::make_tup(xbpvr, xbpvi, xbptr, xbpti); + free_minus = bsk::make_tup(xmvr, xmvi, xmtr, xmti); + pool_minus = bsk::make_tup(xbmvr, xbmvi, xbmtr, xbmti); + conjugated = _dual_conj(carried); + auto t57_ = _dual_product(_dual_add(_dual_product(a11, free_plus), _dual_product(a12, pool_plus)), carried); + rpvr = bsk::get<0>(t57_); + rpvi = bsk::get<1>(t57_); + rptr = bsk::get<2>(t57_); + rpti = bsk::get<3>(t57_); + auto t58_ = _dual_product(_dual_add(_dual_product(a21, free_plus), _dual_product(a22, pool_plus)), carried); + rbpvr = bsk::get<0>(t58_); + rbpvi = bsk::get<1>(t58_); + rbptr = bsk::get<2>(t58_); + rbpti = bsk::get<3>(t58_); + auto t59_ = _dual_product(_dual_add(_dual_product(_dual_conj(a11), free_minus), _dual_product(_dual_conj(a12), pool_minus)), conjugated); + rmvr = bsk::get<0>(t59_); + rmvi = bsk::get<1>(t59_); + rmtr = bsk::get<2>(t59_); + rmti = bsk::get<3>(t59_); + auto t60_ = _dual_product(_dual_add(_dual_product(_dual_conj(a21), free_minus), _dual_product(_dual_conj(a22), pool_minus)), conjugated); + rbmvr = bsk::get<0>(t60_); + rbmvi = bsk::get<1>(t60_); + rbmtr = bsk::get<2>(t60_); + rbmti = bsk::get<3>(t60_); + } else { + auto t61_ = _dual_mul(ovr, ovi, otr, oti, xpvr, xpvi, xptr, xpti); + rpvr = bsk::get<0>(t61_); + rpvi = bsk::get<1>(t61_); + rptr = bsk::get<2>(t61_); + rpti = bsk::get<3>(t61_); + auto t62_ = _dual_mul(ovr, (-ovi), otr, (-oti), xmvr, xmvi, xmtr, xmti); + rmvr = bsk::get<0>(t62_); + rmvi = bsk::get<1>(t62_); + rmtr = bsk::get<2>(t62_); + rmti = bsk::get<3>(t62_); + } + rbvr = empty; + rbvi = empty; + rbtr = empty; + rbti = empty; + rcvr = empty; + rcvi = empty; + rctr = empty; + rcti = empty; + if (bsk::truth((pools == 3))) { + if (bsk::truth(tabulated)) { + // The walk back needs the operator and the direction + // through it, which the row already holds -- and pooling + // the cotangents took what the eigenvalues were formed + // for, so nothing here reads them. + pool_row = bsk::ld(((duration_row + event_base) + event), active_atom, 0); + auto t63_ = _three_pool_from_table_jvp(pool_table, pool_row, atom, atom_count, active_atom, r1_value, r1b_value, r1c_value, atom_exchange, atom_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_tangent, wout_value, wout_tangent); + w11 = bsk::get<0>(t63_); + w12 = bsk::get<1>(t63_); + w13 = bsk::get<2>(t63_); + w21 = bsk::get<3>(t63_); + w22 = bsk::get<4>(t63_); + w23 = bsk::get<5>(t63_); + w31 = bsk::get<6>(t63_); + w32 = bsk::get<7>(t63_); + w33 = bsk::get<8>(t63_); + grow_free = bsk::get<9>(t63_); + grow_pool_b = bsk::get<10>(t63_); + grow_semisolid = bsk::get<11>(t63_); + d_w11 = bsk::get<12>(t63_); + d_w12 = bsk::get<13>(t63_); + d_w13 = bsk::get<14>(t63_); + d_w21 = bsk::get<15>(t63_); + d_w22 = bsk::get<16>(t63_); + d_w23 = bsk::get<17>(t63_); + d_w31 = bsk::get<18>(t63_); + d_w32 = bsk::get<19>(t63_); + d_w33 = bsk::get<20>(t63_); + d_grow_free = bsk::get<21>(t63_); + d_grow_pool_b = bsk::get<22>(t63_); + d_grow_semisolid = bsk::get<23>(t63_); + } else { + auto t64_ = _three_pool_pieces_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_value, dt_tangent, narrow); + three_free = bsk::get<0>(t64_); + three_d_free = bsk::get<1>(t64_); + three_pool_b = bsk::get<2>(t64_); + three_d_pool_b = bsk::get<3>(t64_); + three_pool_c = bsk::get<4>(t64_); + three_d_pool_c = bsk::get<5>(t64_); + three_a00 = bsk::get<6>(t64_); + three_d_a00 = bsk::get<7>(t64_); + three_a01 = bsk::get<8>(t64_); + three_d_a01 = bsk::get<9>(t64_); + three_a02 = bsk::get<10>(t64_); + three_d_a02 = bsk::get<11>(t64_); + three_a10 = bsk::get<12>(t64_); + three_d_a10 = bsk::get<13>(t64_); + three_a11 = bsk::get<14>(t64_); + three_d_a11 = bsk::get<15>(t64_); + three_a20 = bsk::get<16>(t64_); + three_d_a20 = bsk::get<17>(t64_); + three_a22 = bsk::get<18>(t64_); + three_d_a22 = bsk::get<19>(t64_); + three_s00 = bsk::get<20>(t64_); + three_d_s00 = bsk::get<21>(t64_); + three_s11 = bsk::get<22>(t64_); + three_d_s11 = bsk::get<23>(t64_); + three_s22 = bsk::get<24>(t64_); + three_d_s22 = bsk::get<25>(t64_); + three_minors = bsk::get<26>(t64_); + three_d_minors = bsk::get<27>(t64_); + three_sum_flat = bsk::get<28>(t64_); + three_sum_linear = bsk::get<29>(t64_); + three_sum_square = bsk::get<30>(t64_); + three_d_sum_flat = bsk::get<31>(t64_); + three_d_sum_linear = bsk::get<32>(t64_); + three_d_sum_square = bsk::get<33>(t64_); + three_lift = bsk::get<34>(t64_); + three_d_lift = bsk::get<35>(t64_); + three_low = bsk::get<36>(t64_); + three_middle = bsk::get<37>(t64_); + three_d_low = bsk::get<38>(t64_); + three_d_middle = bsk::get<39>(t64_); + three_leading = bsk::get<40>(t64_); + three_d_leading = bsk::get<41>(t64_); + three_first = bsk::get<42>(t64_); + three_d_first = bsk::get<43>(t64_); + three_second = bsk::get<44>(t64_); + three_d_second = bsk::get<45>(t64_); + three_determinant = bsk::get<46>(t64_); + three_d_determinant = bsk::get<47>(t64_); + three_high = bsk::get<48>(t64_); + three_d_high = bsk::get<49>(t64_); + three_radius = bsk::get<50>(t64_); + three_d_radius = bsk::get<51>(t64_); + three_cube = bsk::get<52>(t64_); + three_raw = bsk::get<53>(t64_); + three_d_raw = bsk::get<54>(t64_); + three_argument = bsk::get<55>(t64_); + three_inside_limit = bsk::get<56>(t64_); + three_angle = bsk::get<57>(t64_); + three_d_angle = bsk::get<58>(t64_); + three_centre = bsk::get<59>(t64_); + three_d_centre = bsk::get<60>(t64_); + three_trailing = bsk::get<61>(t64_); + three_d_trailing = bsk::get<62>(t64_); + three_guarded = bsk::get<63>(t64_); + three_d_guarded = bsk::get<64>(t64_); + three_q00 = bsk::get<65>(t64_); + three_d_q00 = bsk::get<66>(t64_); + three_q01 = bsk::get<67>(t64_); + three_d_q01 = bsk::get<68>(t64_); + three_q02 = bsk::get<69>(t64_); + three_d_q02 = bsk::get<70>(t64_); + three_q10 = bsk::get<71>(t64_); + three_d_q10 = bsk::get<72>(t64_); + three_q11 = bsk::get<73>(t64_); + three_d_q11 = bsk::get<74>(t64_); + three_q12 = bsk::get<75>(t64_); + three_d_q12 = bsk::get<76>(t64_); + three_q20 = bsk::get<77>(t64_); + three_d_q20 = bsk::get<78>(t64_); + three_q21 = bsk::get<79>(t64_); + three_d_q21 = bsk::get<80>(t64_); + three_q22 = bsk::get<81>(t64_); + three_d_q22 = bsk::get<82>(t64_); + auto t65_ = _three_pool_assemble_jvp(three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, narrow); + three_def_00 = bsk::get<0>(t65_); + three_dif_00 = bsk::get<1>(t65_); + three_def_01 = bsk::get<2>(t65_); + three_dif_01 = bsk::get<3>(t65_); + three_def_02 = bsk::get<4>(t65_); + three_dif_02 = bsk::get<5>(t65_); + three_def_10 = bsk::get<6>(t65_); + three_dif_10 = bsk::get<7>(t65_); + three_def_11 = bsk::get<8>(t65_); + three_dif_11 = bsk::get<9>(t65_); + three_def_12 = bsk::get<10>(t65_); + three_dif_12 = bsk::get<11>(t65_); + three_def_20 = bsk::get<12>(t65_); + three_dif_20 = bsk::get<13>(t65_); + three_def_21 = bsk::get<14>(t65_); + three_dif_21 = bsk::get<15>(t65_); + three_def_22 = bsk::get<16>(t65_); + three_dif_22 = bsk::get<17>(t65_); + auto t66_ = _three_pool_weigh_jvp(three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, wout_value, wout_tangent, narrow); + w11 = bsk::get<0>(t66_); + w12 = bsk::get<1>(t66_); + w13 = bsk::get<2>(t66_); + w21 = bsk::get<3>(t66_); + w22 = bsk::get<4>(t66_); + w23 = bsk::get<5>(t66_); + w31 = bsk::get<6>(t66_); + w32 = bsk::get<7>(t66_); + w33 = bsk::get<8>(t66_); + grow_free = bsk::get<9>(t66_); + grow_pool_b = bsk::get<10>(t66_); + grow_semisolid = bsk::get<11>(t66_); + d_w11 = bsk::get<12>(t66_); + d_w12 = bsk::get<13>(t66_); + d_w13 = bsk::get<14>(t66_); + d_w21 = bsk::get<15>(t66_); + d_w22 = bsk::get<16>(t66_); + d_w23 = bsk::get<17>(t66_); + d_w31 = bsk::get<18>(t66_); + d_w32 = bsk::get<19>(t66_); + d_w33 = bsk::get<20>(t66_); + d_grow_free = bsk::get<21>(t66_); + d_grow_pool_b = bsk::get<22>(t66_); + d_grow_semisolid = bsk::get<23>(t66_); + // The operator is O(1) once formed, so the per-order loop below + // takes it at the width the states are carried in. + w11 = bsk::cast(w11); + w12 = bsk::cast(w12); + w13 = bsk::cast(w13); + w21 = bsk::cast(w21); + w22 = bsk::cast(w22); + w23 = bsk::cast(w23); + w31 = bsk::cast(w31); + w32 = bsk::cast(w32); + w33 = bsk::cast(w33); + grow_free = bsk::cast(grow_free); + grow_pool_b = bsk::cast(grow_pool_b); + grow_semisolid = bsk::cast(grow_semisolid); + d_w11 = bsk::cast(d_w11); + d_w12 = bsk::cast(d_w12); + d_w13 = bsk::cast(d_w13); + d_w21 = bsk::cast(d_w21); + d_w22 = bsk::cast(d_w22); + d_w23 = bsk::cast(d_w23); + d_w31 = bsk::cast(d_w31); + d_w32 = bsk::cast(d_w32); + d_w33 = bsk::cast(d_w33); + d_grow_free = bsk::cast(d_grow_free); + d_grow_pool_b = bsk::cast(d_grow_pool_b); + d_grow_semisolid = bsk::cast(d_grow_semisolid); + } + spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti); + mixed_free = _dual_add(_dual_add(_dual_scale(w11, d_w11, xzvr, xzvi, xztr, xzti), _dual_scale(w12, d_w12, xbvr, xbvi, xbtr, xbti)), _dual_scale(w13, d_w13, xcvr, xcvi, xctr, xcti)); + mixed_bound = _dual_add(_dual_add(_dual_scale(w21, d_w21, xzvr, xzvi, xztr, xzti), _dual_scale(w22, d_w22, xbvr, xbvi, xbtr, xbti)), _dual_scale(w23, d_w23, xcvr, xcvi, xctr, xcti)); + mixed_semisolid = _dual_add(_dual_add(_dual_scale(w31, d_w31, xzvr, xzvi, xztr, xzti), _dual_scale(w32, d_w32, xbvr, xbvi, xbtr, xbti)), _dual_scale(w33, d_w33, xcvr, xcvi, xctr, xcti)); + auto t67_ = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_free); + rzvr = bsk::get<0>(t67_); + rzvi = bsk::get<1>(t67_); + rztr = bsk::get<2>(t67_); + rzti = bsk::get<3>(t67_); + auto t68_ = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_bound); + rbvr = bsk::get<0>(t68_); + rbvi = bsk::get<1>(t68_); + rbtr = bsk::get<2>(t68_); + rbti = bsk::get<3>(t68_); + auto t69_ = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_semisolid); + rcvr = bsk::get<0>(t69_); + rcvi = bsk::get<1>(t69_); + rctr = bsk::get<2>(t69_); + rcti = bsk::get<3>(t69_); + rzvr = (rzvr + bsk::where((state == 0), grow_free, 0.0f)); + rztr = (rztr + bsk::where((state == 0), d_grow_free, 0.0f)); + rbvr = (rbvr + bsk::where((state == 0), grow_pool_b, 0.0f)); + rbtr = (rbtr + bsk::where((state == 0), d_grow_pool_b, 0.0f)); + rcvr = (rcvr + bsk::where((state == 0), grow_semisolid, 0.0f)); + rctr = (rctr + bsk::where((state == 0), d_grow_semisolid, 0.0f)); + } else if (bsk::truth((pools > 0))) { + auto t70_ = _two_pool_step_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, dt_value, dt_tangent, wout_value, wout_tangent); + pe11 = bsk::get<0>(t70_); + pe12 = bsk::get<1>(t70_); + pe21 = bsk::get<2>(t70_); + pe22 = bsk::get<3>(t70_); + prec_f = bsk::get<4>(t70_); + prec_b = bsk::get<5>(t70_); + de11 = bsk::get<6>(t70_); + de12 = bsk::get<7>(t70_); + de21 = bsk::get<8>(t70_); + de22 = bsk::get<9>(t70_); + drec_f = bsk::get<10>(t70_); + drec_b = bsk::get<11>(t70_); + spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti); + free_part = _dual_scale(pe11, de11, xzvr, xzvi, xztr, xzti); + cross_in = _dual_scale(pe12, de12, xbvr, xbvi, xbtr, xbti); + cross_out = _dual_scale(pe21, de21, xzvr, xzvi, xztr, xzti); + bound_part = _dual_scale(pe22, de22, xbvr, xbvi, xbtr, xbti); + mixed_free = bsk::make_tup((bsk::get<0>(free_part) + bsk::get<0>(cross_in)), (bsk::get<1>(free_part) + bsk::get<1>(cross_in)), (bsk::get<2>(free_part) + bsk::get<2>(cross_in)), (bsk::get<3>(free_part) + bsk::get<3>(cross_in))); + mixed_bound = bsk::make_tup((bsk::get<0>(cross_out) + bsk::get<0>(bound_part)), (bsk::get<1>(cross_out) + bsk::get<1>(bound_part)), (bsk::get<2>(cross_out) + bsk::get<2>(bound_part)), (bsk::get<3>(cross_out) + bsk::get<3>(bound_part))); + auto t71_ = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_free); + rzvr = bsk::get<0>(t71_); + rzvi = bsk::get<1>(t71_); + rztr = bsk::get<2>(t71_); + rzti = bsk::get<3>(t71_); + auto t72_ = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_bound); + rbvr = bsk::get<0>(t72_); + rbvi = bsk::get<1>(t72_); + rbtr = bsk::get<2>(t72_); + rbti = bsk::get<3>(t72_); + rzvr = (rzvr + bsk::where((state == 0), prec_f, 0.0f)); + rztr = (rztr + bsk::where((state == 0), drec_f, 0.0f)); + rbvr = (rbvr + bsk::where((state == 0), prec_b, 0.0f)); + rbtr = (rbtr + bsk::where((state == 0), drec_b, 0.0f)); + } else { + auto t73_ = _dual_mul(lvr, lvi, ltr, lti, xzvr, xzvi, xztr, xzti); + rzvr = bsk::get<0>(t73_); + rzvi = bsk::get<1>(t73_); + rztr = bsk::get<2>(t73_); + rzti = bsk::get<3>(t73_); + rzvr = (rzvr + bsk::where((state == 0), recovery_value, 0.0f)); + rztr = (rztr + bsk::where((state == 0), recovery_tangent, 0.0f)); + } + pre_shift = (bsk::band(event_action, 1) != 0); + auto t74_ = _shift(rpvr, rpvi, rmvr, rmvi, state, state_mask, state_count); + svr = bsk::get<0>(t74_); + svi = bsk::get<1>(t74_); + wvr = bsk::get<2>(t74_); + wvi = bsk::get<3>(t74_); + auto t75_ = _shift(rptr, rpti, rmtr, rmti, state, state_mask, state_count); + str_ = bsk::get<0>(t75_); + sti = bsk::get<1>(t75_); + wtr = bsk::get<2>(t75_); + wti = bsk::get<3>(t75_); + auto spvr = bsk::where(pre_shift, svr, rpvr); + auto spvi = bsk::where(pre_shift, svi, rpvi); + auto sptr = bsk::where(pre_shift, str_, rptr); + auto spti = bsk::where(pre_shift, sti, rpti); + auto smvr = bsk::where(pre_shift, wvr, rmvr); + auto smvi = bsk::where(pre_shift, wvi, rmvi); + auto smtr = bsk::where(pre_shift, wtr, rmtr); + auto smti = bsk::where(pre_shift, wti, rmti); + sbpvr = rbpvr; + sbpvi = rbpvi; + sbptr = rbptr; + sbpti = rbpti; + sbmvr = rbmvr; + sbmvi = rbmvi; + sbmtr = rbmtr; + sbmti = rbmti; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t76_ = _shift(rbpvr, rbpvi, rbmvr, rbmvi, state, state_mask, state_count); + svr = bsk::get<0>(t76_); + svi = bsk::get<1>(t76_); + wvr = bsk::get<2>(t76_); + wvi = bsk::get<3>(t76_); + auto t77_ = _shift(rbptr, rbpti, rbmtr, rbmti, state, state_mask, state_count); + str_ = bsk::get<0>(t77_); + sti = bsk::get<1>(t77_); + wtr = bsk::get<2>(t77_); + wti = bsk::get<3>(t77_); + sbpvr = bsk::where(pre_shift, svr, rbpvr); + sbpvi = bsk::where(pre_shift, svi, rbpvi); + sbptr = bsk::where(pre_shift, str_, rbptr); + sbpti = bsk::where(pre_shift, sti, rbpti); + sbmvr = bsk::where(pre_shift, wvr, rbmvr); + sbmvi = bsk::where(pre_shift, wvi, rbmvi); + sbmtr = bsk::where(pre_shift, wtr, rbmtr); + sbmti = bsk::where(pre_shift, wti, rbmti); + } + // Undo the trailing spoil or shift. + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + spoil = (bsk::band(event_action, 8) != 0); + auto t78_ = _shift_adjoint(pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count); + avr = bsk::get<0>(t78_); + avi = bsk::get<1>(t78_); + bvr = bsk::get<2>(t78_); + bvi = bsk::get<3>(t78_); + auto t79_ = _shift_adjoint(pbtr, pbti, mbtr, mbti, state, state_mask, state_count); + atr = bsk::get<0>(t79_); + ati = bsk::get<1>(t79_); + btr = bsk::get<2>(t79_); + bti = bsk::get<3>(t79_); + auto trailing = bsk::band(do_shift, bsk::bnot(spoil)); + pbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, avr, pbvr)); + pbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, avi, pbvi)); + pbtr = bsk::where(spoil, 0.0f, bsk::where(trailing, atr, pbtr)); + pbti = bsk::where(spoil, 0.0f, bsk::where(trailing, ati, pbti)); + mbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, bvr, mbvr)); + mbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, bvi, mbvi)); + mbtr = bsk::where(spoil, 0.0f, bsk::where(trailing, btr, mbtr)); + mbti = bsk::where(spoil, 0.0f, bsk::where(trailing, bti, mbti)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t80_ = _shift_adjoint(ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count); + avr = bsk::get<0>(t80_); + avi = bsk::get<1>(t80_); + bvr = bsk::get<2>(t80_); + bvi = bsk::get<3>(t80_); + auto t81_ = _shift_adjoint(ubtr, ubti, wbtr, wbti, state, state_mask, state_count); + atr = bsk::get<0>(t81_); + ati = bsk::get<1>(t81_); + btr = bsk::get<2>(t81_); + bti = bsk::get<3>(t81_); + ubvr = bsk::where(spoil, 0.0f, bsk::where(trailing, avr, ubvr)); + ubvi = bsk::where(spoil, 0.0f, bsk::where(trailing, avi, ubvi)); + ubtr = bsk::where(spoil, 0.0f, bsk::where(trailing, atr, ubtr)); + ubti = bsk::where(spoil, 0.0f, bsk::where(trailing, ati, ubti)); + wbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, bvr, wbvr)); + wbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, bvi, wbvi)); + wbtr = bsk::where(spoil, 0.0f, bsk::where(trailing, btr, wbtr)); + wbti = bsk::where(spoil, 0.0f, bsk::where(trailing, bti, wbti)); + } + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_dot_flip = _event_value(dot_flip, event_base, event, active_atom, single_train); + event_phase = _event_value(phase, event_base, event, active_atom, single_train); + event_dot_phase = _event_value(dot_phase, event_base, event, active_atom, single_train); + // ---- recorded sample ---- + auto record = bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)); + auto out_ = bsk::ld((output_index + event)); + auto seed_mask = bsk::band(bsk::band(active_atom, record), (out_ >= 0)); + auto seed_real = bsk::ld(((grad_output_real + (problem * output_count)) + out_), seed_mask, 0.0f); + auto seed_imag = bsk::ld(((grad_output_imag + (problem * output_count)) + out_), seed_mask, 0.0f); + auto t82_ = _dual_polar((-event_phase), (-event_dot_phase)); + auto dvr = bsk::get<0>(t82_); + auto dvi = bsk::get<1>(t82_); + auto dtr = bsk::get<2>(t82_); + auto dti = bsk::get<3>(t82_); + // A coil sees the whole voxel, so what it records is the sum over pools. + recorded = bsk::make_tup(spvr, spvi, sptr, spti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + recorded = _dual_add(recorded, bsk::make_tup(sbpvr, sbpvi, sbptr, sbpti)); + } + // grad_m0 = Re(conj(seed) * recorded * demodulation) + auto t83_ = [&](const auto& s0_) { return _dual_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), dvr, dvi, dtr, dti); }(recorded); + auto wr = bsk::get<0>(t83_); + auto wi = bsk::get<1>(t83_); + auto wtr_ = bsk::get<2>(t83_); + auto wti_ = bsk::get<3>(t83_); + auto t84_ = _dual_real_conj_mul(seed_real, seed_imag, (0.0f * seed_real), (0.0f * seed_imag), wr, wi, wtr_, wti_); + auto m0_value = bsk::get<0>(t84_); + auto m0_tangent = bsk::get<1>(t84_); + g_m0v = (g_m0v + bsk::sum_x(bsk::where((state == 0), m0_value, 0.0f))); + g_m0t = (g_m0t + bsk::sum_x(bsk::where((state == 0), m0_tangent, 0.0f))); + // grad_phase = Re(conj(seed) * m0 * recorded * (-i) * demodulation) + auto t85_ = [&](const auto& s2_) { return _dual_scale(atom_m0, d_m0, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(recorded); + yr = bsk::get<0>(t85_); + yi = bsk::get<1>(t85_); + ytr = bsk::get<2>(t85_); + yti = bsk::get<3>(t85_); + auto t86_ = _dual_times_i(yr, yi, ytr, yti); + yr = bsk::get<0>(t86_); + yi = bsk::get<1>(t86_); + ytr = bsk::get<2>(t86_); + yti = bsk::get<3>(t86_); + auto t87_ = bsk::make_tup((-yr), (-yi), (-ytr), (-yti)); + yr = bsk::get<0>(t87_); + yi = bsk::get<1>(t87_); + ytr = bsk::get<2>(t87_); + yti = bsk::get<3>(t87_); + auto t88_ = _dual_mul(yr, yi, ytr, yti, dvr, dvi, dtr, dti); + yr = bsk::get<0>(t88_); + yi = bsk::get<1>(t88_); + ytr = bsk::get<2>(t88_); + yti = bsk::get<3>(t88_); + auto t89_ = _dual_real_conj_mul(seed_real, seed_imag, (0.0f * seed_real), (0.0f * seed_imag), yr, yi, ytr, yti); + auto phase_value = bsk::get<0>(t89_); + auto phase_tangent = bsk::get<1>(t89_); + bsk::atomic_add(((grad_phase_value + event_base) + event), bsk::sum_x(bsk::where((state == 0), phase_value, 0.0f)), seed_mask); + bsk::atomic_add(((grad_phase_tangent + event_base) + event), bsk::sum_x(bsk::where((state == 0), phase_tangent, 0.0f)), seed_mask); + // fplus_bar[0] += conj(m0 * demodulation) * seed + auto t90_ = _dual_scale(atom_m0, d_m0, dvr, dvi, dtr, dti); + auto kr = bsk::get<0>(t90_); + auto ki = bsk::get<1>(t90_); + auto ktr = bsk::get<2>(t90_); + auto kti = bsk::get<3>(t90_); + auto t91_ = _dual_mul(kr, (-ki), ktr, (-kti), seed_real, seed_imag, (0.0f * seed_real), (0.0f * seed_imag)); + auto sr = bsk::get<0>(t91_); + auto si = bsk::get<1>(t91_); + auto stg_r = bsk::get<2>(t91_); + auto stg_i = bsk::get<3>(t91_); + pbvr = (pbvr + bsk::where((state == 0), sr, 0.0f)); + pbvi = (pbvi + bsk::where((state == 0), si, 0.0f)); + pbtr = (pbtr + bsk::where((state == 0), stg_r, 0.0f)); + pbti = (pbti + bsk::where((state == 0), stg_i, 0.0f)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + ubvr = (ubvr + bsk::where((state == 0), sr, 0.0f)); + ubvi = (ubvi + bsk::where((state == 0), si, 0.0f)); + ubtr = (ubtr + bsk::where((state == 0), stg_r, 0.0f)); + ubti = (ubti + bsk::where((state == 0), stg_i, 0.0f)); + } + // ---- RF adjoint ---- + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + auto t92_ = _dual_real_conj_mul(zbvr, zbvi, zbtr, zbti, (-rzvr), (-rzvi), (-rztr), (-rzti)); + auto inv_value = bsk::get<0>(t92_); + auto inv_tangent = bsk::get<1>(t92_); + g_invv = (g_invv + bsk::sum_x(bsk::where(invert, inv_value, 0.0f))); + g_invt = (g_invt + bsk::sum_x(bsk::where(invert, inv_tangent, 0.0f))); + auto t93_ = _dual_scale((-atom_inv), (-d_inv), zbvr, zbvi, zbtr, zbti); + ivr = bsk::get<0>(t93_); + ivi = bsk::get<1>(t93_); + itr = bsk::get<2>(t93_); + iti = bsk::get<3>(t93_); + zbvr = bsk::where(invert, ivr, zbvr); + zbvi = bsk::where(invert, ivi, zbvi); + zbtr = bsk::where(invert, itr, zbtr); + zbti = bsk::where(invert, iti, zbti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t94_ = _dual_real_conj_mul(bbvr, bbvi, bbtr, bbti, (-rbvr), (-rbvi), (-rbtr), (-rbti)); + auto pool_v = bsk::get<0>(t94_); + auto pool_t = bsk::get<1>(t94_); + g_invv = (g_invv + bsk::sum_x(bsk::where(invert, pool_v, 0.0f))); + g_invt = (g_invt + bsk::sum_x(bsk::where(invert, pool_t, 0.0f))); + auto t95_ = _dual_scale((-atom_inv), (-d_inv), bbvr, bbvi, bbtr, bbti); + ivr = bsk::get<0>(t95_); + ivi = bsk::get<1>(t95_); + itr = bsk::get<2>(t95_); + iti = bsk::get<3>(t95_); + bbvr = bsk::where(invert, ivr, bbvr); + bbvi = bsk::where(invert, ivi, bbvi); + bbtr = bsk::where(invert, itr, bbtr); + bbti = bsk::where(invert, iti, bbti); + } + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse a row of its own. + if (bsk::truth(shimmed)) { + row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + } + d_b1 = bsk::ld(((dot_b1 + row) + atom), active_atom, 0.0f); + if (bsk::truth(off_axis)) { + d_b1_phase = bsk::ld(((dot_b1_phase + row) + atom), active_atom, 0.0f); + } + } + alpha_value = (event_flip * atom_b1); + alpha_tangent = ((event_dot_flip * atom_b1) + (event_flip * d_b1)); + phi_value = (event_phase + atom_b1_phase); + phi_tangent = (event_dot_phase + d_b1_phase); + sat_alpha_v = zero; + sat_alpha_t = zero; + sat_b0_v = zero; + sat_b0_t = zero; + if (bsk::truth(broadened)) { + // The pulse scales every order of the bound pool by one real + // number, so its cotangent is a single sum over the states it + // multiplied. The lineshape's own slope is differentiated too, + // which is what the curvature the reader returns is for. + offset_value = (bsk::ld((rf_frequency + event)) - atom_b0); + auto t96_ = _lineshape_at_curve(lineshape, offset_value, lineshape_bins, lineshape_step); + shape_value = bsk::get<0>(t96_); + shape_slope = bsk::get<1>(t96_); + auto shape_curve = bsk::get<2>(t96_); + shape_tangent = (shape_slope * (-d_b0)); + auto slope_tangent = (shape_curve * (-d_b0)); + event_saturation = bsk::ld((saturation + event)); + power_value = ((event_saturation * alpha_value) * alpha_value); + power_tangent = (((event_saturation * 2.0f) * alpha_value) * alpha_tangent); + absorbed_value = bsk::exp((power_value * shape_value)); + absorbed_tangent = (absorbed_value * ((power_tangent * shape_value) + (power_value * shape_tangent))); + if (bsk::truth((pools == 1))) { + held_bar = bsk::make_tup(bbvr, bbvi, bbtr, bbti); + held_state = bsk::make_tup(rbvr, rbvi, rbtr, rbti); + } else { + held_bar = bsk::make_tup(cbvr, cbvi, cbtr, cbti); + held_state = bsk::make_tup(rcvr, rcvi, rctr, rcti); + } + auto t97_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(held_bar, held_state); + auto per_state_v = bsk::get<0>(t97_); + auto per_state_t = bsk::get<1>(t97_); + auto grad_absorbed_v = bsk::sum_x(per_state_v); + auto grad_absorbed_t = bsk::sum_x(per_state_t); + auto grad_exponent_v = (grad_absorbed_v * absorbed_value); + auto grad_exponent_t = ((grad_absorbed_t * absorbed_value) + (grad_absorbed_v * absorbed_tangent)); + auto twice = (event_saturation * 2.0f); + sat_alpha_v = (grad_exponent_v * ((twice * alpha_value) * shape_value)); + sat_alpha_t = ((grad_exponent_t * ((twice * alpha_value) * shape_value)) + ((grad_exponent_v * twice) * ((alpha_tangent * shape_value) + (alpha_value * shape_tangent)))); + // The lineshape is read at the pulse's offset from the voxel, so a + // step in the voxel's own off-resonance moves the read the other + // way. + sat_b0_v = ((-grad_exponent_v) * (power_value * shape_slope)); + sat_b0_t = (-((grad_exponent_t * (power_value * shape_slope)) + (grad_exponent_v * ((power_tangent * shape_slope) + (power_value * slope_tangent))))); + damped = [&](const auto& s2_) { return _dual_scale(absorbed_value, absorbed_tangent, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_)); }(held_bar); + saturating = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((pools == 1))) { + bbvr = bsk::where(saturating, bsk::get<0>(damped), bbvr); + bbvi = bsk::where(saturating, bsk::get<1>(damped), bbvi); + bbtr = bsk::where(saturating, bsk::get<2>(damped), bbtr); + bbti = bsk::where(saturating, bsk::get<3>(damped), bbti); + } else { + cbvr = bsk::where(saturating, bsk::get<0>(damped), cbvr); + cbvi = bsk::where(saturating, bsk::get<1>(damped), cbvi); + cbtr = bsk::where(saturating, bsk::get<2>(damped), cbtr); + cbti = bsk::where(saturating, bsk::get<3>(damped), cbti); + } + } + cos_value = bsk::cos(alpha_value); + sin_value = bsk::sin(alpha_value); + cos_tangent = ((-sin_value) * alpha_tangent); + sin_tangent = (cos_value * alpha_tangent); + auto t98_ = _dual_polar(phi_value, phi_tangent); + p1r = bsk::get<0>(t98_); + p1i = bsk::get<1>(t98_); + p1tr = bsk::get<2>(t98_); + p1ti = bsk::get<3>(t98_); + auto t99_ = _dual_mul(p1r, p1i, p1tr, p1ti, p1r, p1i, p1tr, p1ti); + p2r = bsk::get<0>(t99_); + p2i = bsk::get<1>(t99_); + p2tr = bsk::get<2>(t99_); + p2ti = bsk::get<3>(t99_); + auto t100_ = _rotation_block((0.5f * (1.0f + cos_value)), (0.5f * cos_tangent), (0.5f * (1.0f - cos_value)), (-0.5f * cos_tangent), sin_value, sin_tangent, cos_value, cos_tangent, p1r, p1i, p1tr, p1ti, p2r, p2i, p2tr, p2ti, p1r, (-p1i), p1tr, (-p1ti)); + t00 = bsk::get<0>(t100_); + t01 = bsk::get<1>(t100_); + t02 = bsk::get<2>(t100_); + r12 = bsk::get<3>(t100_); + t20 = bsk::get<4>(t100_); + r21 = bsk::get<5>(t100_); + r22 = bsk::get<6>(t100_); + // The flip angle reaches a shaped pulse's rotation through the slope + // stored beside it, or not at all when the rotation is read per voxel, + // so the operator's derivative in the flip is only built where the + // pulse is a flip and a phase. + alpha_v = empty; + alpha_t = empty; + phi_v = empty; + phi_t = empty; + alpha_b_v = empty; + alpha_b_t = empty; + phi_b_v = empty; + phi_b_t = empty; + if (bsk::truth((bsk::truth((!bsk::truth(profiled))) && bsk::truth((!bsk::truth(dynamic)))))) { + auto t101_ = _rotation_block((-0.5f * sin_value), (-0.5f * sin_tangent), (0.5f * sin_value), (0.5f * sin_tangent), cos_value, cos_tangent, (-sin_value), (-sin_tangent), p1r, p1i, p1tr, p1ti, p2r, p2i, p2tr, p2ti, p1r, (-p1i), p1tr, (-p1ti)); + auto d00 = bsk::get<0>(t101_); + auto d01 = bsk::get<1>(t101_); + auto d02 = bsk::get<2>(t101_); + auto d12 = bsk::get<3>(t101_); + auto d20 = bsk::get<4>(t101_); + auto d21 = bsk::get<5>(t101_); + auto d22 = bsk::get<6>(t101_); + // d/dalpha, contracted with the adjoint. + row0 = _dual_mul(bsk::get<0>(d00), bsk::get<1>(d00), bsk::get<2>(d00), bsk::get<3>(d00), spvr, spvi, sptr, spti); + add1 = _dual_mul(bsk::get<0>(d01), bsk::get<1>(d01), bsk::get<2>(d01), bsk::get<3>(d01), smvr, smvi, smtr, smti); + add2 = _dual_mul(bsk::get<0>(d02), bsk::get<1>(d02), bsk::get<2>(d02), bsk::get<3>(d02), rzvr, rzvi, rztr, rzti); + auto t102_ = _dual_real_conj_mul(pbvr, pbvi, pbtr, pbti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + alpha_v = bsk::get<0>(t102_); + alpha_t = bsk::get<1>(t102_); + row0 = _dual_mul(bsk::get<0>(d01), (-bsk::get<1>(d01)), bsk::get<2>(d01), (-bsk::get<3>(d01)), spvr, spvi, sptr, spti); + add1 = _dual_mul(bsk::get<0>(d00), bsk::get<1>(d00), bsk::get<2>(d00), bsk::get<3>(d00), smvr, smvi, smtr, smti); + add2 = _dual_mul(bsk::get<0>(d12), bsk::get<1>(d12), bsk::get<2>(d12), bsk::get<3>(d12), rzvr, rzvi, rztr, rzti); + auto t103_ = _dual_real_conj_mul(mbvr, mbvi, mbtr, mbti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + part_v = bsk::get<0>(t103_); + part_t = bsk::get<1>(t103_); + alpha_v = (alpha_v + part_v); + alpha_t = (alpha_t + part_t); + row0 = _dual_mul(bsk::get<0>(d20), bsk::get<1>(d20), bsk::get<2>(d20), bsk::get<3>(d20), spvr, spvi, sptr, spti); + add1 = _dual_mul(bsk::get<0>(d21), bsk::get<1>(d21), bsk::get<2>(d21), bsk::get<3>(d21), smvr, smvi, smtr, smti); + add2 = _dual_mul(bsk::get<0>(d22), bsk::get<1>(d22), bsk::get<2>(d22), bsk::get<3>(d22), rzvr, rzvi, rztr, rzti); + auto t104_ = _dual_real_conj_mul(zbvr, zbvi, zbtr, zbti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + part_v = bsk::get<0>(t104_); + part_t = bsk::get<1>(t104_); + alpha_v = (alpha_v + part_v); + alpha_t = (alpha_t + part_t); + // d/dphi, where only the phase factors carry the dependence. + u1 = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), smvr, smvi, smtr, smti); + u2 = _dual_mul(bsk::get<0>(t02), bsk::get<1>(t02), bsk::get<2>(t02), bsk::get<3>(t02), rzvr, rzvi, rztr, rzti); + auto t105_ = _dual_times_i(((2.0f * bsk::get<0>(u1)) + bsk::get<0>(u2)), ((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2)), ((2.0f * bsk::get<2>(u1)) + bsk::get<2>(u2)), ((2.0f * bsk::get<3>(u1)) + bsk::get<3>(u2))); + ur = bsk::get<0>(t105_); + ui = bsk::get<1>(t105_); + utr = bsk::get<2>(t105_); + uti = bsk::get<3>(t105_); + auto t106_ = _dual_real_conj_mul(pbvr, pbvi, pbtr, pbti, ur, ui, utr, uti); + phi_v = bsk::get<0>(t106_); + phi_t = bsk::get<1>(t106_); + u1 = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), spvr, spvi, sptr, spti); + u2 = _dual_mul(bsk::get<0>(r12), bsk::get<1>(r12), bsk::get<2>(r12), bsk::get<3>(r12), rzvr, rzvi, rztr, rzti); + auto t107_ = _dual_times_i(((-2.0f * bsk::get<0>(u1)) - bsk::get<0>(u2)), ((-2.0f * bsk::get<1>(u1)) - bsk::get<1>(u2)), ((-2.0f * bsk::get<2>(u1)) - bsk::get<2>(u2)), ((-2.0f * bsk::get<3>(u1)) - bsk::get<3>(u2))); + ur = bsk::get<0>(t107_); + ui = bsk::get<1>(t107_); + utr = bsk::get<2>(t107_); + uti = bsk::get<3>(t107_); + auto t108_ = _dual_real_conj_mul(mbvr, mbvi, mbtr, mbti, ur, ui, utr, uti); + part_v = bsk::get<0>(t108_); + part_t = bsk::get<1>(t108_); + phi_v = (phi_v + part_v); + phi_t = (phi_t + part_t); + u1 = _dual_mul(bsk::get<0>(t20), bsk::get<1>(t20), bsk::get<2>(t20), bsk::get<3>(t20), spvr, spvi, sptr, spti); + u2 = _dual_mul(bsk::get<0>(r21), bsk::get<1>(r21), bsk::get<2>(r21), bsk::get<3>(r21), smvr, smvi, smtr, smti); + auto t109_ = _dual_times_i((bsk::get<0>(u2) - bsk::get<0>(u1)), (bsk::get<1>(u2) - bsk::get<1>(u1)), (bsk::get<2>(u2) - bsk::get<2>(u1)), (bsk::get<3>(u2) - bsk::get<3>(u1))); + ur = bsk::get<0>(t109_); + ui = bsk::get<1>(t109_); + utr = bsk::get<2>(t109_); + uti = bsk::get<3>(t109_); + auto t110_ = _dual_real_conj_mul(zbvr, zbvi, zbtr, zbti, ur, ui, utr, uti); + part_v = bsk::get<0>(t110_); + part_t = bsk::get<1>(t110_); + phi_v = (phi_v + part_v); + phi_t = (phi_t + part_t); + // The same pulse turns the exchanging pool, so its cotangent adds to + // the flip and phase the free pool already left. + alpha_b_v = (0.0f * alpha_v); + alpha_b_t = (0.0f * alpha_v); + phi_b_v = (0.0f * alpha_v); + phi_b_t = (0.0f * alpha_v); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + row0 = _dual_mul(bsk::get<0>(d00), bsk::get<1>(d00), bsk::get<2>(d00), bsk::get<3>(d00), sbpvr, sbpvi, sbptr, sbpti); + add1 = _dual_mul(bsk::get<0>(d01), bsk::get<1>(d01), bsk::get<2>(d01), bsk::get<3>(d01), sbmvr, sbmvi, sbmtr, sbmti); + add2 = _dual_mul(bsk::get<0>(d02), bsk::get<1>(d02), bsk::get<2>(d02), bsk::get<3>(d02), rbvr, rbvi, rbtr, rbti); + auto t111_ = _dual_real_conj_mul(ubvr, ubvi, ubtr, ubti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + alpha_b_v = bsk::get<0>(t111_); + alpha_b_t = bsk::get<1>(t111_); + row0 = _dual_mul(bsk::get<0>(d01), (-bsk::get<1>(d01)), bsk::get<2>(d01), (-bsk::get<3>(d01)), sbpvr, sbpvi, sbptr, sbpti); + add1 = _dual_mul(bsk::get<0>(d00), bsk::get<1>(d00), bsk::get<2>(d00), bsk::get<3>(d00), sbmvr, sbmvi, sbmtr, sbmti); + add2 = _dual_mul(bsk::get<0>(d12), bsk::get<1>(d12), bsk::get<2>(d12), bsk::get<3>(d12), rbvr, rbvi, rbtr, rbti); + auto t112_ = _dual_real_conj_mul(wbvr, wbvi, wbtr, wbti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + part_v = bsk::get<0>(t112_); + part_t = bsk::get<1>(t112_); + alpha_b_v = (alpha_b_v + part_v); + alpha_b_t = (alpha_b_t + part_t); + row0 = _dual_mul(bsk::get<0>(d20), bsk::get<1>(d20), bsk::get<2>(d20), bsk::get<3>(d20), sbpvr, sbpvi, sbptr, sbpti); + add1 = _dual_mul(bsk::get<0>(d21), bsk::get<1>(d21), bsk::get<2>(d21), bsk::get<3>(d21), sbmvr, sbmvi, sbmtr, sbmti); + add2 = _dual_mul(bsk::get<0>(d22), bsk::get<1>(d22), bsk::get<2>(d22), bsk::get<3>(d22), rbvr, rbvi, rbtr, rbti); + auto t113_ = _dual_real_conj_mul(bbvr, bbvi, bbtr, bbti, ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)), ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)), ((bsk::get<2>(row0) + bsk::get<2>(add1)) + bsk::get<2>(add2)), ((bsk::get<3>(row0) + bsk::get<3>(add1)) + bsk::get<3>(add2))); + part_v = bsk::get<0>(t113_); + part_t = bsk::get<1>(t113_); + alpha_b_v = (alpha_b_v + part_v); + alpha_b_t = (alpha_b_t + part_t); + u1 = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), sbmvr, sbmvi, sbmtr, sbmti); + u2 = _dual_mul(bsk::get<0>(t02), bsk::get<1>(t02), bsk::get<2>(t02), bsk::get<3>(t02), rbvr, rbvi, rbtr, rbti); + auto t114_ = _dual_times_i(((2.0f * bsk::get<0>(u1)) + bsk::get<0>(u2)), ((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2)), ((2.0f * bsk::get<2>(u1)) + bsk::get<2>(u2)), ((2.0f * bsk::get<3>(u1)) + bsk::get<3>(u2))); + ur = bsk::get<0>(t114_); + ui = bsk::get<1>(t114_); + utr = bsk::get<2>(t114_); + uti = bsk::get<3>(t114_); + auto t115_ = _dual_real_conj_mul(ubvr, ubvi, ubtr, ubti, ur, ui, utr, uti); + phi_b_v = bsk::get<0>(t115_); + phi_b_t = bsk::get<1>(t115_); + u1 = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), sbpvr, sbpvi, sbptr, sbpti); + u2 = _dual_mul(bsk::get<0>(r12), bsk::get<1>(r12), bsk::get<2>(r12), bsk::get<3>(r12), rbvr, rbvi, rbtr, rbti); + auto t116_ = _dual_times_i(((-2.0f * bsk::get<0>(u1)) - bsk::get<0>(u2)), ((-2.0f * bsk::get<1>(u1)) - bsk::get<1>(u2)), ((-2.0f * bsk::get<2>(u1)) - bsk::get<2>(u2)), ((-2.0f * bsk::get<3>(u1)) - bsk::get<3>(u2))); + ur = bsk::get<0>(t116_); + ui = bsk::get<1>(t116_); + utr = bsk::get<2>(t116_); + uti = bsk::get<3>(t116_); + auto t117_ = _dual_real_conj_mul(wbvr, wbvi, wbtr, wbti, ur, ui, utr, uti); + part_v = bsk::get<0>(t117_); + part_t = bsk::get<1>(t117_); + phi_b_v = (phi_b_v + part_v); + phi_b_t = (phi_b_t + part_t); + u1 = _dual_mul(bsk::get<0>(t20), bsk::get<1>(t20), bsk::get<2>(t20), bsk::get<3>(t20), sbpvr, sbpvi, sbptr, sbpti); + u2 = _dual_mul(bsk::get<0>(r21), bsk::get<1>(r21), bsk::get<2>(r21), bsk::get<3>(r21), sbmvr, sbmvi, sbmtr, sbmti); + auto t118_ = _dual_times_i((bsk::get<0>(u2) - bsk::get<0>(u1)), (bsk::get<1>(u2) - bsk::get<1>(u1)), (bsk::get<2>(u2) - bsk::get<2>(u1)), (bsk::get<3>(u2) - bsk::get<3>(u1))); + ur = bsk::get<0>(t118_); + ui = bsk::get<1>(t118_); + utr = bsk::get<2>(t118_); + uti = bsk::get<3>(t118_); + auto t119_ = _dual_real_conj_mul(bbvr, bbvi, bbtr, bbti, ur, ui, utr, uti); + part_v = bsk::get<0>(t119_); + part_t = bsk::get<1>(t119_); + phi_b_v = (phi_b_v + part_v); + phi_b_t = (phi_b_t + part_t); + } + } + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + shaped_slope_a = bsk::make_tup(empty, empty, empty, empty); + shaped_slope_b = bsk::make_tup(empty, empty, empty, empty); + if (bsk::truth(dynamic)) { + auto t120_ = _dynamic_pair_dual_at(pairs, pair_direction, pair_index, event_base, event, atom, atom_count, active_atom, phi_value, phi_tangent, directed); + shaped_a = bsk::get<0>(t120_); + shaped_b = bsk::get<1>(t120_); + } else { + auto t121_ = _profiled_pair_dual(profile, _table_row(profile_index, event, location, locations), alpha_value, alpha_tangent, phi_value, phi_tangent, profile_bins, profile_step); + shaped_a = bsk::get<0>(t121_); + shaped_b = bsk::get<1>(t121_); + shaped_slope_a = bsk::get<2>(t121_); + shaped_slope_b = bsk::get<3>(t121_); + } + auto t122_ = _spinor_adjoint_dual(shaped_a, shaped_b, bsk::make_tup(spvr, spvi, sptr, spti), bsk::make_tup(smvr, smvi, smtr, smti), bsk::make_tup(rzvr, rzvi, rztr, rzti), bsk::make_tup(pbvr, pbvi, pbtr, pbti), bsk::make_tup(mbvr, mbvi, mbtr, mbti), bsk::make_tup(zbvr, zbvi, zbtr, zbti)); + auto grad_a = bsk::get<0>(t122_); + auto grad_b = bsk::get<1>(t122_); + shaped_pb = bsk::get<2>(t122_); + shaped_mb = bsk::get<3>(t122_); + shaped_zb = bsk::get<4>(t122_); + if (bsk::truth(dynamic)) { + // The flip is inside the pair rather than read against it, so + // it has no gradient here: the cotangent goes out on the + // rotation and whatever integrated it carries the rest. ``b`` + // was turned by the phase after the pair came out, so the + // cotangent turns back the other way. + alpha_v = empty; + alpha_t = empty; + auto back = _dual_product(grad_b, _dual_conj(_dual_polar((-phi_value), (-phi_tangent)))); + _store_pair_cotangent(grad_pair_value, grad_pair_tangent, pair_index, event_base, event, atom, atom_count, bsk::band(is_rf, bsk::bnot(is_inversion)), active_atom, state_mask, grad_a, back); + } else { + auto t123_ = _dual_real_conj_mul(bsk::get<0>(grad_a), bsk::get<1>(grad_a), bsk::get<2>(grad_a), bsk::get<3>(grad_a), bsk::get<0>(shaped_slope_a), bsk::get<1>(shaped_slope_a), bsk::get<2>(shaped_slope_a), bsk::get<3>(shaped_slope_a)); + alpha_v = bsk::get<0>(t123_); + alpha_t = bsk::get<1>(t123_); + auto t124_ = _dual_real_conj_mul(bsk::get<0>(grad_b), bsk::get<1>(grad_b), bsk::get<2>(grad_b), bsk::get<3>(grad_b), bsk::get<0>(shaped_slope_b), bsk::get<1>(shaped_slope_b), bsk::get<2>(shaped_slope_b), bsk::get<3>(shaped_slope_b)); + part_v = bsk::get<0>(t124_); + part_t = bsk::get<1>(t124_); + alpha_v = (alpha_v + part_v); + alpha_t = (alpha_t + part_t); + } + // d(b e^{-i phi})/dphi is -i times it, and nothing else moves. + auto t125_ = _dual_times_i(bsk::get<0>(shaped_b), bsk::get<1>(shaped_b), bsk::get<2>(shaped_b), bsk::get<3>(shaped_b)); + auto turn_r = bsk::get<0>(t125_); + auto turn_i = bsk::get<1>(t125_); + auto turn_tr = bsk::get<2>(t125_); + auto turn_ti = bsk::get<3>(t125_); + auto t126_ = _dual_real_conj_mul(bsk::get<0>(grad_b), bsk::get<1>(grad_b), bsk::get<2>(grad_b), bsk::get<3>(grad_b), (-turn_r), (-turn_i), (-turn_tr), (-turn_ti)); + phi_v = bsk::get<0>(t126_); + phi_t = bsk::get<1>(t126_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t127_ = _spinor_adjoint_dual(shaped_a, shaped_b, bsk::make_tup(sbpvr, sbpvi, sbptr, sbpti), bsk::make_tup(sbmvr, sbmvi, sbmtr, sbmti), bsk::make_tup(rbvr, rbvi, rbtr, rbti), bsk::make_tup(ubvr, ubvi, ubtr, ubti), bsk::make_tup(wbvr, wbvi, wbtr, wbti), bsk::make_tup(bbvr, bbvi, bbtr, bbti)); + auto pool_a = bsk::get<0>(t127_); + auto pool_b_pair = bsk::get<1>(t127_); + shaped_ub = bsk::get<2>(t127_); + shaped_wb = bsk::get<3>(t127_); + shaped_bb = bsk::get<4>(t127_); + if (bsk::truth(dynamic)) { + // The same pulse turned this pool, so its cotangent lands + // on the same row. + auto pool_back = _dual_product(pool_b_pair, _dual_conj(_dual_polar((-phi_value), (-phi_tangent)))); + _store_pair_cotangent(grad_pair_value, grad_pair_tangent, pair_index, event_base, event, atom, atom_count, bsk::band(is_rf, bsk::bnot(is_inversion)), active_atom, state_mask, pool_a, pool_back); + } else { + auto t128_ = _dual_real_conj_mul(bsk::get<0>(pool_a), bsk::get<1>(pool_a), bsk::get<2>(pool_a), bsk::get<3>(pool_a), bsk::get<0>(shaped_slope_a), bsk::get<1>(shaped_slope_a), bsk::get<2>(shaped_slope_a), bsk::get<3>(shaped_slope_a)); + alpha_b_v = bsk::get<0>(t128_); + alpha_b_t = bsk::get<1>(t128_); + auto t129_ = _dual_real_conj_mul(bsk::get<0>(pool_b_pair), bsk::get<1>(pool_b_pair), bsk::get<2>(pool_b_pair), bsk::get<3>(pool_b_pair), bsk::get<0>(shaped_slope_b), bsk::get<1>(shaped_slope_b), bsk::get<2>(shaped_slope_b), bsk::get<3>(shaped_slope_b)); + part_v = bsk::get<0>(t129_); + part_t = bsk::get<1>(t129_); + alpha_b_v = (alpha_b_v + part_v); + alpha_b_t = (alpha_b_t + part_t); + } + auto t130_ = _dual_real_conj_mul(bsk::get<0>(pool_b_pair), bsk::get<1>(pool_b_pair), bsk::get<2>(pool_b_pair), bsk::get<3>(pool_b_pair), (-turn_r), (-turn_i), (-turn_tr), (-turn_ti)); + phi_b_v = bsk::get<0>(t130_); + phi_b_t = bsk::get<1>(t130_); + } + } + alpha_v = (alpha_v + alpha_b_v); + alpha_t = (alpha_t + alpha_b_t); + phi_v = (phi_v + phi_b_v); + phi_t = (phi_t + phi_b_t); + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + grad_alpha_v = bsk::sum_x(bsk::where(rotate, alpha_v, 0.0f)); + grad_alpha_t = bsk::sum_x(bsk::where(rotate, alpha_t, 0.0f)); + auto grad_phi_v = bsk::sum_x(bsk::where(rotate, phi_v, 0.0f)); + auto grad_phi_t = bsk::sum_x(bsk::where(rotate, phi_t, 0.0f)); + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + auto turning = bsk::where(rotate, 1.0f, 0.0f); + grad_alpha_v = (grad_alpha_v + (sat_alpha_v * turning)); + grad_alpha_t = (grad_alpha_t + (sat_alpha_t * turning)); + g_b0v = (g_b0v + (sat_b0_v * turning)); + g_b0t = (g_b0t + (sat_b0_t * turning)); + } + // Conjugate transpose of the rotation. + n0 = _dual_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), bsk::get<2>(t00), (-bsk::get<3>(t00)), pbvr, pbvi, pbtr, pbti); + n1 = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), mbvr, mbvi, mbtr, mbti); + n2 = _dual_mul(bsk::get<0>(t20), (-bsk::get<1>(t20)), bsk::get<2>(t20), (-bsk::get<3>(t20)), zbvr, zbvi, zbtr, zbti); + q0 = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), pbvr, pbvi, pbtr, pbti); + q1 = _dual_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), bsk::get<2>(t00), (-bsk::get<3>(t00)), mbvr, mbvi, mbtr, mbti); + q2 = _dual_mul(bsk::get<0>(r21), (-bsk::get<1>(r21)), bsk::get<2>(r21), (-bsk::get<3>(r21)), zbvr, zbvi, zbtr, zbti); + w0 = _dual_mul(bsk::get<0>(t02), (-bsk::get<1>(t02)), bsk::get<2>(t02), (-bsk::get<3>(t02)), pbvr, pbvi, pbtr, pbti); + w1 = _dual_mul(bsk::get<0>(r12), (-bsk::get<1>(r12)), bsk::get<2>(r12), (-bsk::get<3>(r12)), mbvr, mbvi, mbtr, mbti); + w2 = _dual_mul(bsk::get<0>(r22), (-bsk::get<1>(r22)), bsk::get<2>(r22), (-bsk::get<3>(r22)), zbvr, zbvi, zbtr, zbti); + back_pb = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2)), ((bsk::get<2>(n0) + bsk::get<2>(n1)) + bsk::get<2>(n2)), ((bsk::get<3>(n0) + bsk::get<3>(n1)) + bsk::get<3>(n2))); + back_mb = bsk::make_tup(((bsk::get<0>(q0) + bsk::get<0>(q1)) + bsk::get<0>(q2)), ((bsk::get<1>(q0) + bsk::get<1>(q1)) + bsk::get<1>(q2)), ((bsk::get<2>(q0) + bsk::get<2>(q1)) + bsk::get<2>(q2)), ((bsk::get<3>(q0) + bsk::get<3>(q1)) + bsk::get<3>(q2))); + back_zb = bsk::make_tup(((bsk::get<0>(w0) + bsk::get<0>(w1)) + bsk::get<0>(w2)), ((bsk::get<1>(w0) + bsk::get<1>(w1)) + bsk::get<1>(w2)), ((bsk::get<2>(w0) + bsk::get<2>(w1)) + bsk::get<2>(w2)), ((bsk::get<3>(w0) + bsk::get<3>(w1)) + bsk::get<3>(w2))); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + n0 = _dual_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), bsk::get<2>(t00), (-bsk::get<3>(t00)), ubvr, ubvi, ubtr, ubti); + n1 = _dual_mul(bsk::get<0>(t01), bsk::get<1>(t01), bsk::get<2>(t01), bsk::get<3>(t01), wbvr, wbvi, wbtr, wbti); + n2 = _dual_mul(bsk::get<0>(t20), (-bsk::get<1>(t20)), bsk::get<2>(t20), (-bsk::get<3>(t20)), bbvr, bbvi, bbtr, bbti); + q0 = _dual_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bsk::get<2>(t01), (-bsk::get<3>(t01)), ubvr, ubvi, ubtr, ubti); + q1 = _dual_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), bsk::get<2>(t00), (-bsk::get<3>(t00)), wbvr, wbvi, wbtr, wbti); + q2 = _dual_mul(bsk::get<0>(r21), (-bsk::get<1>(r21)), bsk::get<2>(r21), (-bsk::get<3>(r21)), bbvr, bbvi, bbtr, bbti); + w0 = _dual_mul(bsk::get<0>(t02), (-bsk::get<1>(t02)), bsk::get<2>(t02), (-bsk::get<3>(t02)), ubvr, ubvi, ubtr, ubti); + w1 = _dual_mul(bsk::get<0>(r12), (-bsk::get<1>(r12)), bsk::get<2>(r12), (-bsk::get<3>(r12)), wbvr, wbvi, wbtr, wbti); + w2 = _dual_mul(bsk::get<0>(r22), (-bsk::get<1>(r22)), bsk::get<2>(r22), (-bsk::get<3>(r22)), bbvr, bbvi, bbtr, bbti); + back_ub = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2)), ((bsk::get<2>(n0) + bsk::get<2>(n1)) + bsk::get<2>(n2)), ((bsk::get<3>(n0) + bsk::get<3>(n1)) + bsk::get<3>(n2))); + back_wb = bsk::make_tup(((bsk::get<0>(q0) + bsk::get<0>(q1)) + bsk::get<0>(q2)), ((bsk::get<1>(q0) + bsk::get<1>(q1)) + bsk::get<1>(q2)), ((bsk::get<2>(q0) + bsk::get<2>(q1)) + bsk::get<2>(q2)), ((bsk::get<3>(q0) + bsk::get<3>(q1)) + bsk::get<3>(q2))); + back_bb = bsk::make_tup(((bsk::get<0>(w0) + bsk::get<0>(w1)) + bsk::get<0>(w2)), ((bsk::get<1>(w0) + bsk::get<1>(w1)) + bsk::get<1>(w2)), ((bsk::get<2>(w0) + bsk::get<2>(w1)) + bsk::get<2>(w2)), ((bsk::get<3>(w0) + bsk::get<3>(w1)) + bsk::get<3>(w2))); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + back_ub = shaped_ub; + back_wb = shaped_wb; + back_bb = shaped_bb; + } + ubvr = bsk::where(rotate, bsk::get<0>(back_ub), ubvr); + ubvi = bsk::where(rotate, bsk::get<1>(back_ub), ubvi); + ubtr = bsk::where(rotate, bsk::get<2>(back_ub), ubtr); + ubti = bsk::where(rotate, bsk::get<3>(back_ub), ubti); + wbvr = bsk::where(rotate, bsk::get<0>(back_wb), wbvr); + wbvi = bsk::where(rotate, bsk::get<1>(back_wb), wbvi); + wbtr = bsk::where(rotate, bsk::get<2>(back_wb), wbtr); + wbti = bsk::where(rotate, bsk::get<3>(back_wb), wbti); + bbvr = bsk::where(rotate, bsk::get<0>(back_bb), bbvr); + bbvi = bsk::where(rotate, bsk::get<1>(back_bb), bbvi); + bbtr = bsk::where(rotate, bsk::get<2>(back_bb), bbtr); + bbti = bsk::where(rotate, bsk::get<3>(back_bb), bbti); + } + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + back_pb = shaped_pb; + back_mb = shaped_mb; + back_zb = shaped_zb; + } + pbvr = bsk::where(rotate, bsk::get<0>(back_pb), pbvr); + pbvi = bsk::where(rotate, bsk::get<1>(back_pb), pbvi); + pbtr = bsk::where(rotate, bsk::get<2>(back_pb), pbtr); + pbti = bsk::where(rotate, bsk::get<3>(back_pb), pbti); + mbvr = bsk::where(rotate, bsk::get<0>(back_mb), mbvr); + mbvi = bsk::where(rotate, bsk::get<1>(back_mb), mbvi); + mbtr = bsk::where(rotate, bsk::get<2>(back_mb), mbtr); + mbti = bsk::where(rotate, bsk::get<3>(back_mb), mbti); + zbvr = bsk::where(rotate, bsk::get<0>(back_zb), zbvr); + zbvi = bsk::where(rotate, bsk::get<1>(back_zb), zbvi); + zbtr = bsk::where(rotate, bsk::get<2>(back_zb), zbtr); + zbti = bsk::where(rotate, bsk::get<3>(back_zb), zbti); + auto writes_flip = bsk::band(active_atom, rotate); + bsk::atomic_add(((grad_flip_value + event_base) + event), (grad_alpha_v * atom_b1), writes_flip); + bsk::atomic_add(((grad_flip_tangent + event_base) + event), ((grad_alpha_t * atom_b1) + (grad_alpha_v * d_b1)), writes_flip); + bsk::atomic_add(((grad_phase_value + event_base) + event), grad_phi_v, writes_flip); + bsk::atomic_add(((grad_phase_tangent + event_base) + event), grad_phi_t, writes_flip); + // A pulse's transmit gradient belongs to the shim it drives, so with + // several it lands in that shim's row here rather than in a register + // summed over the whole train. ``row`` is the offset of the row the + // replay above read. + if (bsk::truth(shimmed)) { + bsk::atomic_add((((grad_tissue_value + (3 * atom_count)) + row) + atom), (grad_alpha_v * event_flip), writes_flip); + bsk::atomic_add((((grad_tissue_tangent + (3 * atom_count)) + row) + atom), ((grad_alpha_t * event_flip) + (grad_alpha_v * event_dot_flip)), writes_flip); + bsk::atomic_add((((grad_tissue_value + (((4 + shim_rows) - 1) * atom_count)) + row) + atom), grad_phi_v, writes_flip); + bsk::atomic_add((((grad_tissue_tangent + (((4 + shim_rows) - 1) * atom_count)) + row) + atom), grad_phi_t, writes_flip); + } else { + g_b1v = (g_b1v + (grad_alpha_v * event_flip)); + g_b1t = (g_b1t + ((grad_alpha_t * event_flip) + (grad_alpha_v * event_dot_flip))); + g_b1pv = (g_b1pv + grad_phi_v); + g_b1pt = (g_b1pt + grad_phi_t); + } + auto t131_ = _shift_adjoint(pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count); + avr = bsk::get<0>(t131_); + avi = bsk::get<1>(t131_); + bvr = bsk::get<2>(t131_); + bvi = bsk::get<3>(t131_); + auto t132_ = _shift_adjoint(pbtr, pbti, mbtr, mbti, state, state_mask, state_count); + atr = bsk::get<0>(t132_); + ati = bsk::get<1>(t132_); + btr = bsk::get<2>(t132_); + bti = bsk::get<3>(t132_); + pbvr = bsk::where(pre_shift, avr, pbvr); + pbvi = bsk::where(pre_shift, avi, pbvi); + pbtr = bsk::where(pre_shift, atr, pbtr); + pbti = bsk::where(pre_shift, ati, pbti); + mbvr = bsk::where(pre_shift, bvr, mbvr); + mbvi = bsk::where(pre_shift, bvi, mbvi); + mbtr = bsk::where(pre_shift, btr, mbtr); + mbti = bsk::where(pre_shift, bti, mbti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t133_ = _shift_adjoint(ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count); + avr = bsk::get<0>(t133_); + avi = bsk::get<1>(t133_); + bvr = bsk::get<2>(t133_); + bvi = bsk::get<3>(t133_); + auto t134_ = _shift_adjoint(ubtr, ubti, wbtr, wbti, state, state_mask, state_count); + atr = bsk::get<0>(t134_); + ati = bsk::get<1>(t134_); + btr = bsk::get<2>(t134_); + bti = bsk::get<3>(t134_); + ubvr = bsk::where(pre_shift, avr, ubvr); + ubvi = bsk::where(pre_shift, avi, ubvi); + ubtr = bsk::where(pre_shift, atr, ubtr); + ubti = bsk::where(pre_shift, ati, ubti); + wbvr = bsk::where(pre_shift, bvr, wbvr); + wbvi = bsk::where(pre_shift, bvi, wbvi); + wbtr = bsk::where(pre_shift, btr, wbtr); + wbti = bsk::where(pre_shift, bti, wbti); + } + // ---- relaxation and off-resonance adjoint ---- + grad_e2_v = zero; + grad_e2_t = zero; + attenuation_v = zero; + attenuation_t = zero; + two_pool_dt_v = zero; + two_pool_dt_t = zero; + // The damping is homogeneous of degree one in every transverse state it + // acts on, so its gradient times the damping itself is the cotangent + // taken against the states the interval leaves. With one pool that is + // the same thing as the relaxation factor's own gradient, scaled. + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto plus_side = _dual_add(_dual_product(_dual_conj(bsk::make_tup(pbvr, pbvi, pbtr, pbti)), bsk::make_tup(rpvr, rpvi, rptr, rpti)), _dual_product(_dual_conj(bsk::make_tup(ubvr, ubvi, ubtr, ubti)), bsk::make_tup(rbpvr, rbpvi, rbptr, rbpti))); + auto minus_side = _dual_add(_dual_product(_dual_conj(bsk::make_tup(mbvr, mbvi, mbtr, mbti)), bsk::make_tup(rmvr, rmvi, rmtr, rmti)), _dual_product(_dual_conj(bsk::make_tup(wbvr, wbvi, wbtr, wbti)), bsk::make_tup(rbmvr, rbmvi, rbmtr, rbmti))); + damped = _dual_add(plus_side, minus_side); + auto wound = [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(_dual_subtract(plus_side, minus_side)); + cot2_v = bsk::get<0>(damped); + cot2_t = bsk::get<2>(damped); + per_angle_v = bsk::get<0>(wound); + per_angle_t = bsk::get<2>(wound); + } else { + auto pq = _dual_mul(qr, qi, qtr, qti, xpvr, xpvi, xptr, xpti); + auto mq = _dual_mul(qr, (-qi), qtr, (-qti), xmvr, xmvi, xmtr, xmti); + auto t135_ = _dual_real_conj_mul(pbvr, pbvi, pbtr, pbti, bsk::get<0>(pq), bsk::get<1>(pq), bsk::get<2>(pq), bsk::get<3>(pq)); + auto e2_v = bsk::get<0>(t135_); + auto e2_t = bsk::get<1>(t135_); + auto t136_ = _dual_real_conj_mul(mbvr, mbvi, mbtr, mbti, bsk::get<0>(mq), bsk::get<1>(mq), bsk::get<2>(mq), bsk::get<3>(mq)); + part_v = bsk::get<0>(t136_); + part_t = bsk::get<1>(t136_); + auto bare_cot_v = (e2_v + part_v); + auto bare_cot_t = (e2_t + part_t); + grad_e2_v = bsk::sum_x((bare_cot_v * damp_t)); + grad_e2_t = bsk::sum_x(((bare_cot_v * damp_t_tangent) + (bare_cot_t * damp_t))); + per_angle_v = empty; + per_angle_t = empty; + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + po = _dual_mul(ovr, ovi, otr, oti, xpvr, xpvi, xptr, xpti); + po = _dual_times_i(bsk::get<0>(po), bsk::get<1>(po), bsk::get<2>(po), bsk::get<3>(po)); + mo = _dual_mul(ovr, (-ovi), otr, (-oti), xmvr, xmvi, xmtr, xmti); + mo = _dual_times_i(bsk::get<0>(mo), bsk::get<1>(mo), bsk::get<2>(mo), bsk::get<3>(mo)); + auto t137_ = _dual_real_conj_mul(pbvr, pbvi, pbtr, pbti, bsk::get<0>(po), bsk::get<1>(po), bsk::get<2>(po), bsk::get<3>(po)); + auto angle_v = bsk::get<0>(t137_); + auto angle_t = bsk::get<1>(t137_); + auto t138_ = _dual_real_conj_mul(mbvr, mbvi, mbtr, mbti, bsk::get<0>(mo), bsk::get<1>(mo), bsk::get<2>(mo), bsk::get<3>(mo)); + part_v = bsk::get<0>(t138_); + part_t = bsk::get<1>(t138_); + per_angle_v = (angle_v - part_v); + per_angle_t = (angle_t - part_t); + } + cot2_v = ((bare_cot_v * bare2_value) * damp_t); + cot2_t = ((((bare_cot_t * bare2_value) * damp_t) + ((bare_cot_v * bare2_tangent) * damp_t)) + ((bare_cot_v * bare2_value) * damp_t_tangent)); + } + // A turn of the transverse states and the off-resonance angle are the + // same derivative; only the weight each order carries differs. + grad_angle_v = zero; + grad_angle_t = zero; + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + grad_angle_v = bsk::sum_x(per_angle_v); + grad_angle_t = bsk::sum_x(per_angle_t); + } + grad_e1_v = zero; + grad_e1_t = zero; + if (bsk::truth((pools == 3))) { + // The nine entries of the mixing operator and the three recoveries, + // summed over the orders that share them, then pushed back through + // the closed form once for the whole interval and in double. + free_bar = bsk::make_tup(zbvr, zbvi, zbtr, zbti); + bound_bar = bsk::make_tup(bbvr, bbvi, bbtr, bbti); + auto semi_bar = bsk::make_tup(cbvr, cbvi, cbtr, cbti); + spun_free = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), xzvr, xzvi, xztr, xzti); + spun_bound = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), xbvr, xbvi, xbtr, xbti); + auto spun_semi = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), xcvr, xcvi, xctr, xcti); + auto t139_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_free); + e11_v = bsk::get<0>(t139_); + e11_t = bsk::get<1>(t139_); + auto t140_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_bound); + e12_v = bsk::get<0>(t140_); + e12_t = bsk::get<1>(t140_); + auto t141_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_semi); + auto e13_v = bsk::get<0>(t141_); + auto e13_t = bsk::get<1>(t141_); + auto t142_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_free); + e21_v = bsk::get<0>(t142_); + e21_t = bsk::get<1>(t142_); + auto t143_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_bound); + e22_v = bsk::get<0>(t143_); + e22_t = bsk::get<1>(t143_); + auto t144_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_semi); + auto e23_v = bsk::get<0>(t144_); + auto e23_t = bsk::get<1>(t144_); + auto t145_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(semi_bar, spun_free); + auto e31_v = bsk::get<0>(t145_); + auto e31_t = bsk::get<1>(t145_); + auto t146_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(semi_bar, spun_bound); + auto e32_v = bsk::get<0>(t146_); + auto e32_t = bsk::get<1>(t146_); + auto t147_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(semi_bar, spun_semi); + auto e33_v = bsk::get<0>(t147_); + auto e33_t = bsk::get<1>(t147_); + if (bsk::truth(tabulated)) { + // Every gradient but the interval's own and the + // attenuation's is linear in these cotangents, so the + // events sharing a length pool them here and pay the + // closed form once each after the walk back. The tangent + // gradient carries a third term in the event's own + // interval direction, which pools as the value cotangents + // weighted by it. + auto bar11_v = bsk::sum_x(e11_v); + auto bar12_v = bsk::sum_x(e12_v); + auto bar13_v = bsk::sum_x(e13_v); + auto bar21_v = bsk::sum_x(e21_v); + auto bar22_v = bsk::sum_x(e22_v); + auto bar23_v = bsk::sum_x(e23_v); + auto bar31_v = bsk::sum_x(e31_v); + auto bar32_v = bsk::sum_x(e32_v); + auto bar33_v = bsk::sum_x(e33_v); + auto barfree_v = bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f)); + auto barpool_v = bsk::sum_x(bsk::where((state == 0), bbvr, 0.0f)); + auto barbound_v = bsk::sum_x(bsk::where((state == 0), cbvr, 0.0f)); + auto bar11_t = bsk::sum_x(e11_t); + auto bar12_t = bsk::sum_x(e12_t); + auto bar13_t = bsk::sum_x(e13_t); + auto bar21_t = bsk::sum_x(e21_t); + auto bar22_t = bsk::sum_x(e22_t); + auto bar23_t = bsk::sum_x(e23_t); + auto bar31_t = bsk::sum_x(e31_t); + auto bar32_t = bsk::sum_x(e32_t); + auto bar33_t = bsk::sum_x(e33_t); + auto barfree_t = bsk::sum_x(bsk::where((state == 0), zbtr, 0.0f)); + auto barpool_t = bsk::sum_x(bsk::where((state == 0), bbtr, 0.0f)); + auto barbound_t = bsk::sum_x(bsk::where((state == 0), cbtr, 0.0f)); + held = (pool_bars + (((local * row_count) + pool_row) * 36)); + bsk::st((held + 0), (bsk::ld((held + 0), active_atom, 0.0f) + bar11_v), active_atom); + bsk::st((held + 1), (bsk::ld((held + 1), active_atom, 0.0f) + bar12_v), active_atom); + bsk::st((held + 2), (bsk::ld((held + 2), active_atom, 0.0f) + bar13_v), active_atom); + bsk::st((held + 3), (bsk::ld((held + 3), active_atom, 0.0f) + bar21_v), active_atom); + bsk::st((held + 4), (bsk::ld((held + 4), active_atom, 0.0f) + bar22_v), active_atom); + bsk::st((held + 5), (bsk::ld((held + 5), active_atom, 0.0f) + bar23_v), active_atom); + bsk::st((held + 6), (bsk::ld((held + 6), active_atom, 0.0f) + bar31_v), active_atom); + bsk::st((held + 7), (bsk::ld((held + 7), active_atom, 0.0f) + bar32_v), active_atom); + bsk::st((held + 8), (bsk::ld((held + 8), active_atom, 0.0f) + bar33_v), active_atom); + bsk::st((held + 9), (bsk::ld((held + 9), active_atom, 0.0f) + barfree_v), active_atom); + bsk::st((held + 10), (bsk::ld((held + 10), active_atom, 0.0f) + barpool_v), active_atom); + bsk::st((held + 11), (bsk::ld((held + 11), active_atom, 0.0f) + barbound_v), active_atom); + bsk::st((held + 12), (bsk::ld((held + 12), active_atom, 0.0f) + bar11_t), active_atom); + bsk::st((held + 13), (bsk::ld((held + 13), active_atom, 0.0f) + bar12_t), active_atom); + bsk::st((held + 14), (bsk::ld((held + 14), active_atom, 0.0f) + bar13_t), active_atom); + bsk::st((held + 15), (bsk::ld((held + 15), active_atom, 0.0f) + bar21_t), active_atom); + bsk::st((held + 16), (bsk::ld((held + 16), active_atom, 0.0f) + bar22_t), active_atom); + bsk::st((held + 17), (bsk::ld((held + 17), active_atom, 0.0f) + bar23_t), active_atom); + bsk::st((held + 18), (bsk::ld((held + 18), active_atom, 0.0f) + bar31_t), active_atom); + bsk::st((held + 19), (bsk::ld((held + 19), active_atom, 0.0f) + bar32_t), active_atom); + bsk::st((held + 20), (bsk::ld((held + 20), active_atom, 0.0f) + bar33_t), active_atom); + bsk::st((held + 21), (bsk::ld((held + 21), active_atom, 0.0f) + barfree_t), active_atom); + bsk::st((held + 22), (bsk::ld((held + 22), active_atom, 0.0f) + barpool_t), active_atom); + bsk::st((held + 23), (bsk::ld((held + 23), active_atom, 0.0f) + barbound_t), active_atom); + bsk::st((held + 24), (bsk::ld((held + 24), active_atom, 0.0f) + (dt_tangent * bar11_v)), active_atom); + bsk::st((held + 25), (bsk::ld((held + 25), active_atom, 0.0f) + (dt_tangent * bar12_v)), active_atom); + bsk::st((held + 26), (bsk::ld((held + 26), active_atom, 0.0f) + (dt_tangent * bar13_v)), active_atom); + bsk::st((held + 27), (bsk::ld((held + 27), active_atom, 0.0f) + (dt_tangent * bar21_v)), active_atom); + bsk::st((held + 28), (bsk::ld((held + 28), active_atom, 0.0f) + (dt_tangent * bar22_v)), active_atom); + bsk::st((held + 29), (bsk::ld((held + 29), active_atom, 0.0f) + (dt_tangent * bar23_v)), active_atom); + bsk::st((held + 30), (bsk::ld((held + 30), active_atom, 0.0f) + (dt_tangent * bar31_v)), active_atom); + bsk::st((held + 31), (bsk::ld((held + 31), active_atom, 0.0f) + (dt_tangent * bar32_v)), active_atom); + bsk::st((held + 32), (bsk::ld((held + 32), active_atom, 0.0f) + (dt_tangent * bar33_v)), active_atom); + bsk::st((held + 33), (bsk::ld((held + 33), active_atom, 0.0f) + (dt_tangent * barfree_v)), active_atom); + bsk::st((held + 34), (bsk::ld((held + 34), active_atom, 0.0f) + (dt_tangent * barpool_v)), active_atom); + bsk::st((held + 35), (bsk::ld((held + 35), active_atom, 0.0f) + (dt_tangent * barbound_v)), active_atom); + auto t148_ = _three_pool_interval_adjoint_jvp(pool_table, pool_row, atom, atom_count, active_atom, r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_tangent, wout_value, wout_tangent, bar11_v, bar12_v, bar13_v, bar21_v, bar22_v, bar23_v, bar31_v, bar32_v, bar33_v, barfree_v, barpool_v, barbound_v, bar11_t, bar12_t, bar13_t, bar21_t, bar22_t, bar23_t, bar31_t, bar32_t, bar33_t, barfree_t, barpool_t, barbound_t); + back_dt_v = bsk::get<0>(t148_); + back_att_v = bsk::get<1>(t148_); + back_dt_t = bsk::get<2>(t148_); + back_att_t = bsk::get<3>(t148_); + attenuation_v = back_att_v; + attenuation_t = back_att_t; + two_pool_dt_v = back_dt_v; + two_pool_dt_t = back_dt_t; + } else { + auto t149_ = _three_pool_step_adjoint_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, dt_value, dt_tangent, wout_value, wout_tangent, bsk::sum_x(e11_v), bsk::sum_x(e11_t), bsk::sum_x(e12_v), bsk::sum_x(e12_t), bsk::sum_x(e13_v), bsk::sum_x(e13_t), bsk::sum_x(e21_v), bsk::sum_x(e21_t), bsk::sum_x(e22_v), bsk::sum_x(e22_t), bsk::sum_x(e23_v), bsk::sum_x(e23_t), bsk::sum_x(e31_v), bsk::sum_x(e31_t), bsk::sum_x(e32_v), bsk::sum_x(e32_t), bsk::sum_x(e33_v), bsk::sum_x(e33_t), bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f)), bsk::sum_x(bsk::where((state == 0), zbtr, 0.0f)), bsk::sum_x(bsk::where((state == 0), bbvr, 0.0f)), bsk::sum_x(bsk::where((state == 0), bbtr, 0.0f)), bsk::sum_x(bsk::where((state == 0), cbvr, 0.0f)), bsk::sum_x(bsk::where((state == 0), cbtr, 0.0f)), three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, narrow); + back_r1_v = bsk::get<0>(t149_); + back_r1b_v = bsk::get<1>(t149_); + back_r1c_v = bsk::get<2>(t149_); + back_exch_v = bsk::get<3>(t149_); + back_sexch_v = bsk::get<4>(t149_); + back_bound_v = bsk::get<5>(t149_); + back_semi_v = bsk::get<6>(t149_); + back_dt_v = bsk::get<7>(t149_); + back_att_v = bsk::get<8>(t149_); + back_r1_t = bsk::get<9>(t149_); + back_r1b_t = bsk::get<10>(t149_); + back_r1c_t = bsk::get<11>(t149_); + back_exch_t = bsk::get<12>(t149_); + back_sexch_t = bsk::get<13>(t149_); + back_bound_t = bsk::get<14>(t149_); + back_semi_t = bsk::get<15>(t149_); + back_dt_t = bsk::get<16>(t149_); + back_att_t = bsk::get<17>(t149_); + slope1_v = bsk::truediv(-1000.0f, (atom_t1 * atom_t1)); + slope1_t = bsk::truediv((2000.0f * d_t1), ((atom_t1 * atom_t1) * atom_t1)); + slope1b_v = bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)); + slope1b_t = bsk::truediv((2000.0f * d_t1b), ((atom_t1b * atom_t1b) * atom_t1b)); + slope1c_v = bsk::truediv(-1000.0f, (held_semisolid * held_semisolid)); + slope1c_t = bsk::truediv((2000.0f * d_semisolid_t1), ((held_semisolid * held_semisolid) * held_semisolid)); + g_t1v = (g_t1v + (back_r1_v * slope1_v)); + g_t1t = (g_t1t + ((back_r1_t * slope1_v) + (back_r1_v * slope1_t))); + g_t1bv = (g_t1bv + (back_r1b_v * slope1b_v)); + g_t1bt = (g_t1bt + ((back_r1b_t * slope1b_v) + (back_r1b_v * slope1b_t))); + g_t1cv = (g_t1cv + (back_r1c_v * slope1c_v)); + g_t1ct = (g_t1ct + ((back_r1c_t * slope1c_v) + (back_r1c_v * slope1c_t))); + g_exchv = (g_exchv + back_exch_v); + g_excht = (g_excht + back_exch_t); + g_sexchv = (g_sexchv + back_sexch_v); + g_sexcht = (g_sexcht + back_sexch_t); + g_boundv = (g_boundv + back_bound_v); + g_boundt = (g_boundt + back_bound_t); + g_semiv = (g_semiv + back_semi_v); + g_semit = (g_semit + back_semi_t); + attenuation_v = back_att_v; + attenuation_t = back_att_t; + two_pool_dt_v = back_dt_v; + two_pool_dt_t = back_dt_t; + } + auto turned_free = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_free); + auto turned_bound = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_bound); + auto turned_semi = [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_semisolid); + auto t150_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, turned_free); + damp_pair_v = bsk::get<0>(t150_); + damp_pair_t = bsk::get<1>(t150_); + auto t151_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, turned_bound); + other_v = bsk::get<0>(t151_); + other_t = bsk::get<1>(t151_); + auto t152_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(semi_bar, turned_semi); + auto stuck_v = bsk::get<0>(t152_); + auto stuck_t = bsk::get<1>(t152_); + long_damp_v = ((damp_pair_v + other_v) + stuck_v); + long_damp_t = ((damp_pair_t + other_t) + stuck_t); + auto t153_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(turned_free)); + zangle_v = bsk::get<0>(t153_); + zangle_t = bsk::get<1>(t153_); + auto t154_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(turned_bound)); + part_v = bsk::get<0>(t154_); + part_t = bsk::get<1>(t154_); + zangle_v = (zangle_v + part_v); + zangle_t = (zangle_t + part_t); + auto t155_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(semi_bar, [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }(turned_semi)); + part_v = bsk::get<0>(t155_); + part_t = bsk::get<1>(t155_); + zangle_v = (zangle_v + part_v); + zangle_t = (zangle_t + part_t); + auto col_free = _dual_add(_dual_add([&](const auto& s2_, const auto& s3_) { return _dual_back(w11, d_w11, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, free_bar), [&](const auto& s2_, const auto& s3_) { return _dual_back(w21, d_w21, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, bound_bar)), [&](const auto& s2_, const auto& s3_) { return _dual_back(w31, d_w31, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, semi_bar)); + auto col_bound = _dual_add(_dual_add([&](const auto& s2_, const auto& s3_) { return _dual_back(w12, d_w12, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, free_bar), [&](const auto& s2_, const auto& s3_) { return _dual_back(w22, d_w22, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, bound_bar)), [&](const auto& s2_, const auto& s3_) { return _dual_back(w32, d_w32, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, semi_bar)); + auto col_semi = _dual_add(_dual_add([&](const auto& s2_, const auto& s3_) { return _dual_back(w13, d_w13, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, free_bar), [&](const auto& s2_, const auto& s3_) { return _dual_back(w23, d_w23, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, bound_bar)), [&](const auto& s2_, const auto& s3_) { return _dual_back(w33, d_w33, bsk::get<0>(s2_), bsk::get<1>(s2_), bsk::get<2>(s2_), bsk::get<3>(s2_), bsk::get<0>(s3_), bsk::get<1>(s3_), bsk::get<2>(s3_), bsk::get<3>(s3_)); }(spin, semi_bar)); + auto t156_ = col_free; + zbvr = bsk::get<0>(t156_); + zbvi = bsk::get<1>(t156_); + zbtr = bsk::get<2>(t156_); + zbti = bsk::get<3>(t156_); + auto t157_ = col_bound; + bbvr = bsk::get<0>(t157_); + bbvi = bsk::get<1>(t157_); + bbtr = bsk::get<2>(t157_); + bbti = bsk::get<3>(t157_); + auto t158_ = col_semi; + cbvr = bsk::get<0>(t158_); + cbvi = bsk::get<1>(t158_); + cbtr = bsk::get<2>(t158_); + cbti = bsk::get<3>(t158_); + } else if (bsk::truth((pools > 0))) { + // The four entries of the exchange operator and the two recoveries, + // summed over the orders that share them, then pushed back through + // the closed form once for the whole interval. + free_bar = bsk::make_tup(zbvr, zbvi, zbtr, zbti); + bound_bar = bsk::make_tup(bbvr, bbvi, bbtr, bbti); + spun_free = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), xzvr, xzvi, xztr, xzti); + spun_bound = _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), xbvr, xbvi, xbtr, xbti); + auto t159_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_free); + e11_v = bsk::get<0>(t159_); + e11_t = bsk::get<1>(t159_); + auto t160_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_bound); + e12_v = bsk::get<0>(t160_); + e12_t = bsk::get<1>(t160_); + auto t161_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_free); + e21_v = bsk::get<0>(t161_); + e21_t = bsk::get<1>(t161_); + auto t162_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_bound); + e22_v = bsk::get<0>(t162_); + e22_t = bsk::get<1>(t162_); + auto bar_e11_v = bsk::sum_x(e11_v); + auto bar_e11_t = bsk::sum_x(e11_t); + auto bar_e12_v = bsk::sum_x(e12_v); + auto bar_e12_t = bsk::sum_x(e12_t); + auto bar_e21_v = bsk::sum_x(e21_v); + auto bar_e21_t = bsk::sum_x(e21_t); + auto bar_e22_v = bsk::sum_x(e22_v); + auto bar_e22_t = bsk::sum_x(e22_t); + auto rec_f_v = bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f)); + auto rec_f_t = bsk::sum_x(bsk::where((state == 0), zbtr, 0.0f)); + auto rec_b_v = bsk::sum_x(bsk::where((state == 0), bbvr, 0.0f)); + auto rec_b_t = bsk::sum_x(bsk::where((state == 0), bbtr, 0.0f)); + auto t163_ = _two_pool_step_adjoint_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, dt_value, dt_tangent, wout_value, wout_tangent, bar_e11_v, bar_e11_t, bar_e12_v, bar_e12_t, bar_e21_v, bar_e21_t, bar_e22_v, bar_e22_t, rec_f_v, rec_f_t, rec_b_v, rec_b_t); + back_r1_v = bsk::get<0>(t163_); + back_r1b_v = bsk::get<1>(t163_); + back_exch_v = bsk::get<2>(t163_); + back_bound_v = bsk::get<3>(t163_); + back_dt_v = bsk::get<4>(t163_); + back_att_v = bsk::get<5>(t163_); + back_r1_t = bsk::get<6>(t163_); + back_r1b_t = bsk::get<7>(t163_); + back_exch_t = bsk::get<8>(t163_); + back_bound_t = bsk::get<9>(t163_); + back_dt_t = bsk::get<10>(t163_); + back_att_t = bsk::get<11>(t163_); + // r1 = 1000/t1, so a rate gradient reaches the time through the + // square of it. + slope1_v = bsk::truediv(-1000.0f, (atom_t1 * atom_t1)); + slope1_t = bsk::truediv((2000.0f * d_t1), ((atom_t1 * atom_t1) * atom_t1)); + slope1b_v = bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)); + slope1b_t = bsk::truediv((2000.0f * d_t1b), ((atom_t1b * atom_t1b) * atom_t1b)); + g_t1v = (g_t1v + (back_r1_v * slope1_v)); + g_t1t = (g_t1t + ((back_r1_t * slope1_v) + (back_r1_v * slope1_t))); + g_t1bv = (g_t1bv + (back_r1b_v * slope1b_v)); + g_t1bt = (g_t1bt + ((back_r1b_t * slope1b_v) + (back_r1b_v * slope1b_t))); + g_exchv = (g_exchv + back_exch_v); + g_excht = (g_excht + back_exch_t); + g_boundv = (g_boundv + back_bound_v); + g_boundt = (g_boundt + back_bound_t); + attenuation_v = back_att_v; + attenuation_t = back_att_t; + two_pool_dt_v = back_dt_v; + two_pool_dt_t = back_dt_t; + // Both pools take the same per-order damping and turn, so each + // collects the cotangent of the mixture that reached it. + auto t164_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_free)); + damp_pair_v = bsk::get<0>(t164_); + damp_pair_t = bsk::get<1>(t164_); + auto t165_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, [&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_bound)); + other_v = bsk::get<0>(t165_); + other_t = bsk::get<1>(t165_); + long_damp_v = (damp_pair_v + other_v); + long_damp_t = (damp_pair_t + other_t); + auto spun_mix_free = [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }([&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_free)); + auto spun_mix_bound = [&](const auto& s0_) { return _dual_times_i(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_)); }([&](const auto& s4_) { return _dual_mul(bsk::get<0>(spin), bsk::get<1>(spin), bsk::get<2>(spin), bsk::get<3>(spin), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(mixed_bound)); + auto t166_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(free_bar, spun_mix_free); + zangle_v = bsk::get<0>(t166_); + zangle_t = bsk::get<1>(t166_); + auto t167_ = [&](const auto& s0_, const auto& s1_) { return _dual_real_conj_mul(bsk::get<0>(s0_), bsk::get<1>(s0_), bsk::get<2>(s0_), bsk::get<3>(s0_), bsk::get<0>(s1_), bsk::get<1>(s1_), bsk::get<2>(s1_), bsk::get<3>(s1_)); }(bound_bar, spun_mix_bound); + part_v = bsk::get<0>(t167_); + part_t = bsk::get<1>(t167_); + zangle_v = (zangle_v + part_v); + zangle_t = (zangle_t + part_t); + auto back_z = [&](const auto& s4_) { return _dual_mul((pe11 * bsk::get<0>(spin)), (-(pe11 * bsk::get<1>(spin))), ((de11 * bsk::get<0>(spin)) + (pe11 * bsk::get<2>(spin))), (-((de11 * bsk::get<1>(spin)) + (pe11 * bsk::get<3>(spin)))), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(free_bar); + auto cross_z = [&](const auto& s4_) { return _dual_mul((pe21 * bsk::get<0>(spin)), (-(pe21 * bsk::get<1>(spin))), ((de21 * bsk::get<0>(spin)) + (pe21 * bsk::get<2>(spin))), (-((de21 * bsk::get<1>(spin)) + (pe21 * bsk::get<3>(spin)))), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(bound_bar); + auto back_b = [&](const auto& s4_) { return _dual_mul((pe12 * bsk::get<0>(spin)), (-(pe12 * bsk::get<1>(spin))), ((de12 * bsk::get<0>(spin)) + (pe12 * bsk::get<2>(spin))), (-((de12 * bsk::get<1>(spin)) + (pe12 * bsk::get<3>(spin)))), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(free_bar); + auto cross_b = [&](const auto& s4_) { return _dual_mul((pe22 * bsk::get<0>(spin)), (-(pe22 * bsk::get<1>(spin))), ((de22 * bsk::get<0>(spin)) + (pe22 * bsk::get<2>(spin))), (-((de22 * bsk::get<1>(spin)) + (pe22 * bsk::get<3>(spin)))), bsk::get<0>(s4_), bsk::get<1>(s4_), bsk::get<2>(s4_), bsk::get<3>(s4_)); }(bound_bar); + auto next_zbvr = (bsk::get<0>(back_z) + bsk::get<0>(cross_z)); + auto next_zbvi = (bsk::get<1>(back_z) + bsk::get<1>(cross_z)); + auto next_zbtr = (bsk::get<2>(back_z) + bsk::get<2>(cross_z)); + auto next_zbti = (bsk::get<3>(back_z) + bsk::get<3>(cross_z)); + bbvr = (bsk::get<0>(back_b) + bsk::get<0>(cross_b)); + bbvi = (bsk::get<1>(back_b) + bsk::get<1>(cross_b)); + bbtr = (bsk::get<2>(back_b) + bsk::get<2>(cross_b)); + bbti = (bsk::get<3>(back_b) + bsk::get<3>(cross_b)); + zbvr = next_zbvr; + zbvi = next_zbvi; + zbtr = next_zbtr; + zbti = next_zbti; + } else { + auto spun = _dual_mul(szr, szi, sztr, szti, xzvr, xzvi, xztr, xzti); + auto t168_ = _dual_real_conj_mul(zbvr, zbvi, zbtr, zbti, bsk::get<0>(spun), bsk::get<1>(spun), bsk::get<2>(spun), bsk::get<3>(spun)); + auto e1_v = bsk::get<0>(t168_); + auto e1_t = bsk::get<1>(t168_); + grad_e1_v = bsk::sum_x((e1_v * damp_z)); + grad_e1_t = bsk::sum_x(((e1_v * damp_z_tangent) + (e1_t * damp_z))); + grad_e1_v = (grad_e1_v - bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f))); + grad_e1_t = (grad_e1_t - bsk::sum_x(bsk::where((state == 0), zbtr, 0.0f))); + long_damp_v = ((e1_v * bare1_value) * damp_z); + long_damp_t = ((((e1_t * bare1_value) * damp_z) + ((e1_v * bare1_tangent) * damp_z)) + ((e1_v * bare1_value) * damp_z_tangent)); + // The longitudinal states turn too, and by a whole order rather + // than the transverse half-order more. + zo = _dual_mul(lvr, lvi, ltr, lti, xzvr, xzvi, xztr, xzti); + zo = _dual_times_i(bsk::get<0>(zo), bsk::get<1>(zo), bsk::get<2>(zo), bsk::get<3>(zo)); + auto t169_ = _dual_real_conj_mul(zbvr, zbvi, zbtr, zbti, bsk::get<0>(zo), bsk::get<1>(zo), bsk::get<2>(zo), bsk::get<3>(zo)); + zangle_v = bsk::get<0>(t169_); + zangle_t = bsk::get<1>(t169_); + auto t170_ = _dual_mul(lvr, (-lvi), ltr, (-lti), zbvr, zbvi, zbtr, zbti); + zbvr = bsk::get<0>(t170_); + zbvi = bsk::get<1>(t170_); + zbtr = bsk::get<2>(t170_); + zbti = bsk::get<3>(t170_); + } + next_pb = bsk::make_tup(pbvr, pbvi, pbtr, pbti); + next_mb = bsk::make_tup(mbvr, mbvi, mbtr, mbti); + next_ub = bsk::make_tup(ubvr, ubvi, ubtr, ubti); + next_wb = bsk::make_tup(wbvr, wbvi, wbtr, wbti); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // The four entries of the transverse operator, summed over the + // orders that share them, then pushed back through the closed form + // once for the whole interval. ``F-`` follows the conjugate of the + // operator, so its cotangent lands on the entry itself rather than + // on the conjugate of it. + auto ap = bsk::make_tup(pbvr, pbvi, pbtr, pbti); + auto am = bsk::make_tup(mbvr, mbvi, mbtr, mbti); + auto aub = bsk::make_tup(ubvr, ubvi, ubtr, ubti); + auto awb = bsk::make_tup(wbvr, wbvi, wbtr, wbti); + auto fp = bsk::make_tup(xpvr, xpvi, xptr, xpti); + auto fm = bsk::make_tup(xmvr, xmvi, xmtr, xmti); + auto bp = bsk::make_tup(xbpvr, xbpvi, xbptr, xbpti); + auto bm = bsk::make_tup(xbmvr, xbmvi, xbmtr, xbmti); + auto term11 = _dual_product(_dual_add(_dual_product(_dual_conj(ap), fp), _dual_product(am, _dual_conj(fm))), carried); + auto term12 = _dual_product(_dual_add(_dual_product(_dual_conj(ap), bp), _dual_product(am, _dual_conj(bm))), carried); + auto term21 = _dual_product(_dual_add(_dual_product(_dual_conj(aub), fp), _dual_product(awb, _dual_conj(fm))), carried); + auto term22 = _dual_product(_dual_add(_dual_product(_dual_conj(aub), bp), _dual_product(awb, _dual_conj(bm))), carried); + auto bar11 = bsk::make_tup(bsk::sum_x(bsk::get<0>(term11)), bsk::sum_x(bsk::get<1>(term11)), bsk::sum_x(bsk::get<2>(term11)), bsk::sum_x(bsk::get<3>(term11))); + auto bar12 = bsk::make_tup(bsk::sum_x(bsk::get<0>(term12)), bsk::sum_x(bsk::get<1>(term12)), bsk::sum_x(bsk::get<2>(term12)), bsk::sum_x(bsk::get<3>(term12))); + auto bar21 = bsk::make_tup(bsk::sum_x(bsk::get<0>(term21)), bsk::sum_x(bsk::get<1>(term21)), bsk::sum_x(bsk::get<2>(term21)), bsk::sum_x(bsk::get<3>(term21))); + auto bar22 = bsk::make_tup(bsk::sum_x(bsk::get<0>(term22)), bsk::sum_x(bsk::get<1>(term22)), bsk::sum_x(bsk::get<2>(term22)), bsk::sum_x(bsk::get<3>(term22))); + auto t171_ = _two_pool_transverse_adjoint_jvp(r2_value, r2_tangent, r2b_value, r2b_tangent, atom_exchange, d_exchange, atom_bound, d_boundf, atom_free, d_free, atom_shift, d_shift, dt_value, dt_tangent, wout_value, wout_tangent, bar11, bar12, bar21, bar22); + auto back_r2_v = bsk::get<0>(t171_); + auto back_r2_t = bsk::get<1>(t171_); + auto back_r2b_v = bsk::get<2>(t171_); + auto back_r2b_t = bsk::get<3>(t171_); + auto back_xexch_v = bsk::get<4>(t171_); + auto back_xexch_t = bsk::get<5>(t171_); + auto back_xbound_v = bsk::get<6>(t171_); + auto back_xbound_t = bsk::get<7>(t171_); + auto back_xfree_v = bsk::get<8>(t171_); + auto back_xfree_t = bsk::get<9>(t171_); + auto back_shift_v = bsk::get<10>(t171_); + auto back_shift_t = bsk::get<11>(t171_); + auto back_xdt_v = bsk::get<12>(t171_); + auto back_xdt_t = bsk::get<13>(t171_); + auto back_xatt_v = bsk::get<14>(t171_); + auto back_xatt_t = bsk::get<15>(t171_); + auto slope2_v = bsk::truediv(-1000.0f, (atom_t2 * atom_t2)); + auto slope2_t = bsk::truediv((2000.0f * d_t2), ((atom_t2 * atom_t2) * atom_t2)); + auto slope2b_v = bsk::truediv(-1000.0f, (atom_t2b * atom_t2b)); + auto slope2b_t = bsk::truediv((2000.0f * d_t2b), ((atom_t2b * atom_t2b) * atom_t2b)); + g_t2v = (g_t2v + (back_r2_v * slope2_v)); + g_t2t = (g_t2t + ((back_r2_t * slope2_v) + (back_r2_v * slope2_t))); + g_t2bv = (g_t2bv + (back_r2b_v * slope2b_v)); + g_t2bt = (g_t2bt + ((back_r2b_t * slope2b_v) + (back_r2b_v * slope2b_t))); + g_exchv = (g_exchv + back_xexch_v); + g_excht = (g_excht + back_xexch_t); + // The free water is what both second pools leave, so a cotangent + // on it reaches each of their fractions turned over. + g_boundv = (g_boundv + (back_xbound_v - back_xfree_v)); + g_boundt = (g_boundt + (back_xbound_t - back_xfree_t)); + if (bsk::truth((pools == 3))) { + g_semiv = (g_semiv - back_xfree_v); + g_semit = (g_semit - back_xfree_t); + } + g_shiftv = (g_shiftv + back_shift_v); + g_shiftt = (g_shiftt + back_shift_t); + attenuation_v = (attenuation_v + back_xatt_v); + attenuation_t = (attenuation_t + back_xatt_t); + two_pool_dt_v = (two_pool_dt_v + back_xdt_v); + two_pool_dt_t = (two_pool_dt_t + back_xdt_t); + auto step11 = _dual_product(a11, carried); + auto step12 = _dual_product(a12, carried); + auto step21 = _dual_product(a21, carried); + auto step22 = _dual_product(a22, carried); + next_pb = _dual_add(_dual_product(_dual_conj(step11), ap), _dual_product(_dual_conj(step21), aub)); + next_ub = _dual_add(_dual_product(_dual_conj(step12), ap), _dual_product(_dual_conj(step22), aub)); + next_mb = _dual_add(_dual_product(step11, am), _dual_product(step21, awb)); + next_wb = _dual_add(_dual_product(step12, am), _dual_product(step22, awb)); + } + // The rate and the interval multiply every order's b-weight, so both + // take a weighted sum rather than one scalar. Order zero carries no + // longitudinal weight, which keeps recovery out of this. + spread_v = zero; + spread_t = zero; + if (bsk::truth(diffusing)) { + auto weighted_v = ((long_damp_v * longitudinal_weight) + (cot2_v * transverse_weight)); + auto weighted_t = ((long_damp_t * longitudinal_weight) + (cot2_t * transverse_weight)); + spread_v = bsk::sum_x(weighted_v); + spread_t = bsk::sum_x(weighted_t); + g_diffv = (g_diffv + ((-spread_v) * dt_value)); + g_difft = (g_difft + (-((spread_v * dt_tangent) + (spread_t * dt_value)))); + } + wound_v = zero; + wound_t = zero; + if (bsk::truth(moving)) { + wound_v = bsk::sum_x(((per_angle_v * (order + 0.5f)) + (zangle_v * order))); + wound_t = bsk::sum_x(((per_angle_t * (order + 0.5f)) + (zangle_t * order))); + g_flowv = (g_flowv + ((-wound_v) * dt_value)); + g_flowt = (g_flowt + (-((wound_v * dt_tangent) + (wound_t * dt_value)))); + } + // Washout scales both relaxation factors, so its gradient is the one + // they already carry, taken against the factors before that scaling. + // Past the clamp the interval has replaced the voxel outright and + // nothing further depends on the rate. + wash_v = zero; + wash_t = zero; + if (bsk::truth(moving)) { + auto live = bsk::cast(((atom_washout * dt_value) < 1.0f)); + wash_v = ((-live) * (((grad_e1_v * dry1_value) + (grad_e2_v * dry2_value)) + attenuation_v)); + wash_t = ((-live) * (((((grad_e1_v * dry1_tangent) + (grad_e1_t * dry1_value)) + (grad_e2_v * dry2_tangent)) + (grad_e2_t * dry2_value)) + attenuation_t)); + g_washv = (g_washv + (wash_v * dt_value)); + g_washt = (g_washt + ((wash_v * dt_tangent) + (wash_t * dt_value))); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t172_ = next_pb; + pbvr = bsk::get<0>(t172_); + pbvi = bsk::get<1>(t172_); + pbtr = bsk::get<2>(t172_); + pbti = bsk::get<3>(t172_); + auto t173_ = next_mb; + mbvr = bsk::get<0>(t173_); + mbvi = bsk::get<1>(t173_); + mbtr = bsk::get<2>(t173_); + mbti = bsk::get<3>(t173_); + auto t174_ = next_ub; + ubvr = bsk::get<0>(t174_); + ubvi = bsk::get<1>(t174_); + ubtr = bsk::get<2>(t174_); + ubti = bsk::get<3>(t174_); + auto t175_ = next_wb; + wbvr = bsk::get<0>(t175_); + wbvi = bsk::get<1>(t175_); + wbtr = bsk::get<2>(t175_); + wbti = bsk::get<3>(t175_); + } else { + auto t176_ = _dual_mul(ovr, (-ovi), otr, (-oti), pbvr, pbvi, pbtr, pbti); + pbvr = bsk::get<0>(t176_); + pbvi = bsk::get<1>(t176_); + pbtr = bsk::get<2>(t176_); + pbti = bsk::get<3>(t176_); + auto t177_ = _dual_mul(ovr, ovi, otr, oti, mbvr, mbvi, mbtr, mbti); + mbvr = bsk::get<0>(t177_); + mbvi = bsk::get<1>(t177_); + mbtr = bsk::get<2>(t177_); + mbti = bsk::get<3>(t177_); + } + auto inverse1_value = bsk::truediv(1000.0f, (atom_t1 * atom_t1)); + auto inverse1_tangent = bsk::truediv((-2000.0f * d_t1), ((atom_t1 * atom_t1) * atom_t1)); + auto inverse2_value = bsk::truediv(1000.0f, (atom_t2 * atom_t2)); + auto inverse2_tangent = bsk::truediv((-2000.0f * d_t2), ((atom_t2 * atom_t2) * atom_t2)); + auto scale1_value = ((bare1_value * dt_value) * inverse1_value); + scale1_tangent = ((bare1_tangent * dt_value) * inverse1_value); + scale1_tangent = (scale1_tangent + ((bare1_value * dt_tangent) * inverse1_value)); + scale1_tangent = (scale1_tangent + ((bare1_value * dt_value) * inverse1_tangent)); + auto scale2_value = ((bare2_value * dt_value) * inverse2_value); + scale2_tangent = ((bare2_tangent * dt_value) * inverse2_value); + scale2_tangent = (scale2_tangent + ((bare2_value * dt_tangent) * inverse2_value)); + scale2_tangent = (scale2_tangent + ((bare2_value * dt_value) * inverse2_tangent)); + g_t1v = (g_t1v + (grad_e1_v * scale1_value)); + g_t1t = (g_t1t + ((grad_e1_v * scale1_tangent) + (grad_e1_t * scale1_value))); + g_t2v = (g_t2v + (grad_e2_v * scale2_value)); + g_t2t = (g_t2t + ((grad_e2_v * scale2_tangent) + (grad_e2_t * scale2_value))); + auto turn = -6.283185307179586f; + g_b0v = (g_b0v + (grad_angle_v * (turn * dt_value))); + g_b0t = (g_b0t + ((grad_angle_v * (turn * dt_tangent)) + (grad_angle_t * (turn * dt_value)))); + auto decay1_value = (r1_value * bare1_value); + auto decay1_tangent = ((r1_value * bare1_tangent) + (r1_tangent * bare1_value)); + auto decay2_value = (r2_value * bare2_value); + auto decay2_tangent = ((r2_value * bare2_tangent) + (r2_tangent * bare2_value)); + duration_v = (((-grad_e1_v) * decay1_value) - (grad_e2_v * decay2_value)); + duration_v = (duration_v + ((grad_angle_v * (turn * atom_b0)) + two_pool_dt_v)); + duration_t = (-((grad_e1_v * decay1_tangent) + (grad_e1_t * decay1_value))); + duration_t = (duration_t - ((grad_e2_v * decay2_tangent) + (grad_e2_t * decay2_value))); + duration_t = (duration_t + ((grad_angle_v * (turn * d_b0)) + (grad_angle_t * (turn * atom_b0)))); + duration_t = (duration_t + two_pool_dt_t); + duration_v = (duration_v + (((-spread_v) * atom_damping) - (wound_v * atom_flow))); + duration_t = (duration_t + (-((spread_v * d_damping) + (spread_t * atom_damping)))); + duration_t = (duration_t + (-((wound_v * d_flow) + (wound_t * atom_flow)))); + duration_v = (duration_v + (wash_v * atom_washout)); + duration_t = (duration_t + ((wash_v * d_washout) + (wash_t * atom_washout))); + bsk::atomic_add(((grad_duration_value + event_base) + event), duration_v, active_atom); + bsk::atomic_add(((grad_duration_tangent + event_base) + event), duration_t, active_atom); + } + if (bsk::truth((bsk::truth((pools == 3)) && bsk::truth(tabulated)))) { + // One closed form per distinct length rather than one per event, + // run twice. The walk back pooled the cotangents the eigenvalues + // are pushed through and the closed form is linear in them, so the + // pieces of the sum are the sum of the pieces. A gradient's own + // direction depends on the interval as well, and a row is shared + // by events whose interval directions differ -- so the second pass + // takes that dependence alone, driven by the cotangents the walk + // back weighted by each event's direction and read at a unit one. + for (std::int64_t row = 0; row < row_count; row += 1) { + held = (pool_bars + (((local * row_count) + row) * 36)); + auto row_dt = (bsk::ld((pool_durations + row)) + zero); + auto nil = (0.0f * row_dt); + auto unit = (1.0f + nil); + one_att = unit; + att_rate = nil; + att_span = nil; + if (bsk::truth(moving)) { + auto t178_ = _washout_jvp(atom_washout, d_washout, row_dt, nil); + one_att = bsk::get<0>(t178_); + att_rate = bsk::get<1>(t178_); + auto t179_ = _washout_jvp(atom_washout, nil, row_dt, unit); + auto _held_att = bsk::get<0>(t179_); + att_span = bsk::get<1>(t179_); + } + auto t180_ = _three_pool_pieces_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, row_dt, nil, narrow); + three_free = bsk::get<0>(t180_); + three_d_free = bsk::get<1>(t180_); + three_pool_b = bsk::get<2>(t180_); + three_d_pool_b = bsk::get<3>(t180_); + three_pool_c = bsk::get<4>(t180_); + three_d_pool_c = bsk::get<5>(t180_); + three_a00 = bsk::get<6>(t180_); + three_d_a00 = bsk::get<7>(t180_); + three_a01 = bsk::get<8>(t180_); + three_d_a01 = bsk::get<9>(t180_); + three_a02 = bsk::get<10>(t180_); + three_d_a02 = bsk::get<11>(t180_); + three_a10 = bsk::get<12>(t180_); + three_d_a10 = bsk::get<13>(t180_); + three_a11 = bsk::get<14>(t180_); + three_d_a11 = bsk::get<15>(t180_); + three_a20 = bsk::get<16>(t180_); + three_d_a20 = bsk::get<17>(t180_); + three_a22 = bsk::get<18>(t180_); + three_d_a22 = bsk::get<19>(t180_); + three_s00 = bsk::get<20>(t180_); + three_d_s00 = bsk::get<21>(t180_); + three_s11 = bsk::get<22>(t180_); + three_d_s11 = bsk::get<23>(t180_); + three_s22 = bsk::get<24>(t180_); + three_d_s22 = bsk::get<25>(t180_); + three_minors = bsk::get<26>(t180_); + three_d_minors = bsk::get<27>(t180_); + three_sum_flat = bsk::get<28>(t180_); + three_sum_linear = bsk::get<29>(t180_); + three_sum_square = bsk::get<30>(t180_); + three_d_sum_flat = bsk::get<31>(t180_); + three_d_sum_linear = bsk::get<32>(t180_); + three_d_sum_square = bsk::get<33>(t180_); + three_lift = bsk::get<34>(t180_); + three_d_lift = bsk::get<35>(t180_); + three_low = bsk::get<36>(t180_); + three_middle = bsk::get<37>(t180_); + three_d_low = bsk::get<38>(t180_); + three_d_middle = bsk::get<39>(t180_); + three_leading = bsk::get<40>(t180_); + three_d_leading = bsk::get<41>(t180_); + three_first = bsk::get<42>(t180_); + three_d_first = bsk::get<43>(t180_); + three_second = bsk::get<44>(t180_); + three_d_second = bsk::get<45>(t180_); + three_determinant = bsk::get<46>(t180_); + three_d_determinant = bsk::get<47>(t180_); + three_high = bsk::get<48>(t180_); + three_d_high = bsk::get<49>(t180_); + three_radius = bsk::get<50>(t180_); + three_d_radius = bsk::get<51>(t180_); + three_cube = bsk::get<52>(t180_); + three_raw = bsk::get<53>(t180_); + three_d_raw = bsk::get<54>(t180_); + three_argument = bsk::get<55>(t180_); + three_inside_limit = bsk::get<56>(t180_); + three_angle = bsk::get<57>(t180_); + three_d_angle = bsk::get<58>(t180_); + three_centre = bsk::get<59>(t180_); + three_d_centre = bsk::get<60>(t180_); + three_trailing = bsk::get<61>(t180_); + three_d_trailing = bsk::get<62>(t180_); + three_guarded = bsk::get<63>(t180_); + three_d_guarded = bsk::get<64>(t180_); + three_q00 = bsk::get<65>(t180_); + three_d_q00 = bsk::get<66>(t180_); + three_q01 = bsk::get<67>(t180_); + three_d_q01 = bsk::get<68>(t180_); + three_q02 = bsk::get<69>(t180_); + three_d_q02 = bsk::get<70>(t180_); + three_q10 = bsk::get<71>(t180_); + three_d_q10 = bsk::get<72>(t180_); + three_q11 = bsk::get<73>(t180_); + three_d_q11 = bsk::get<74>(t180_); + three_q12 = bsk::get<75>(t180_); + three_d_q12 = bsk::get<76>(t180_); + three_q20 = bsk::get<77>(t180_); + three_d_q20 = bsk::get<78>(t180_); + three_q21 = bsk::get<79>(t180_); + three_d_q21 = bsk::get<80>(t180_); + three_q22 = bsk::get<81>(t180_); + three_d_q22 = bsk::get<82>(t180_); + auto t181_ = _three_pool_assemble_jvp(three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, narrow); + three_def_00 = bsk::get<0>(t181_); + three_dif_00 = bsk::get<1>(t181_); + three_def_01 = bsk::get<2>(t181_); + three_dif_01 = bsk::get<3>(t181_); + three_def_02 = bsk::get<4>(t181_); + three_dif_02 = bsk::get<5>(t181_); + three_def_10 = bsk::get<6>(t181_); + three_dif_10 = bsk::get<7>(t181_); + three_def_11 = bsk::get<8>(t181_); + three_dif_11 = bsk::get<9>(t181_); + three_def_12 = bsk::get<10>(t181_); + three_dif_12 = bsk::get<11>(t181_); + three_def_20 = bsk::get<12>(t181_); + three_dif_20 = bsk::get<13>(t181_); + three_def_21 = bsk::get<14>(t181_); + three_dif_21 = bsk::get<15>(t181_); + three_def_22 = bsk::get<16>(t181_); + three_dif_22 = bsk::get<17>(t181_); + auto t182_ = _three_pool_step_adjoint_jvp(r1_value, r1_tangent, r1b_value, r1b_tangent, r1c_value, r1c_tangent, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_boundf, atom_semisolid, d_semisolidf, row_dt, nil, one_att, att_rate, bsk::ld((held + 0), active_atom, 0.0f), bsk::ld((held + 12), active_atom, 0.0f), bsk::ld((held + 1), active_atom, 0.0f), bsk::ld((held + 13), active_atom, 0.0f), bsk::ld((held + 2), active_atom, 0.0f), bsk::ld((held + 14), active_atom, 0.0f), bsk::ld((held + 3), active_atom, 0.0f), bsk::ld((held + 15), active_atom, 0.0f), bsk::ld((held + 4), active_atom, 0.0f), bsk::ld((held + 16), active_atom, 0.0f), bsk::ld((held + 5), active_atom, 0.0f), bsk::ld((held + 17), active_atom, 0.0f), bsk::ld((held + 6), active_atom, 0.0f), bsk::ld((held + 18), active_atom, 0.0f), bsk::ld((held + 7), active_atom, 0.0f), bsk::ld((held + 19), active_atom, 0.0f), bsk::ld((held + 8), active_atom, 0.0f), bsk::ld((held + 20), active_atom, 0.0f), bsk::ld((held + 9), active_atom, 0.0f), bsk::ld((held + 21), active_atom, 0.0f), bsk::ld((held + 10), active_atom, 0.0f), bsk::ld((held + 22), active_atom, 0.0f), bsk::ld((held + 11), active_atom, 0.0f), bsk::ld((held + 23), active_atom, 0.0f), three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, narrow); + back_r1_v = bsk::get<0>(t182_); + back_r1b_v = bsk::get<1>(t182_); + back_r1c_v = bsk::get<2>(t182_); + back_exch_v = bsk::get<3>(t182_); + back_sexch_v = bsk::get<4>(t182_); + back_bound_v = bsk::get<5>(t182_); + back_semi_v = bsk::get<6>(t182_); + auto _row_dt_v = bsk::get<7>(t182_); + auto _row_att_v = bsk::get<8>(t182_); + back_r1_t = bsk::get<9>(t182_); + back_r1b_t = bsk::get<10>(t182_); + back_r1c_t = bsk::get<11>(t182_); + back_exch_t = bsk::get<12>(t182_); + back_sexch_t = bsk::get<13>(t182_); + back_bound_t = bsk::get<14>(t182_); + back_semi_t = bsk::get<15>(t182_); + auto _row_dt_t = bsk::get<16>(t182_); + auto _row_att_t = bsk::get<17>(t182_); + auto t183_ = _three_pool_pieces_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, row_dt, unit, narrow); + auto alt_free = bsk::get<0>(t183_); + auto alt_d_free = bsk::get<1>(t183_); + auto alt_pool_b = bsk::get<2>(t183_); + auto alt_d_pool_b = bsk::get<3>(t183_); + auto alt_pool_c = bsk::get<4>(t183_); + auto alt_d_pool_c = bsk::get<5>(t183_); + auto alt_a00 = bsk::get<6>(t183_); + auto alt_d_a00 = bsk::get<7>(t183_); + auto alt_a01 = bsk::get<8>(t183_); + auto alt_d_a01 = bsk::get<9>(t183_); + auto alt_a02 = bsk::get<10>(t183_); + auto alt_d_a02 = bsk::get<11>(t183_); + auto alt_a10 = bsk::get<12>(t183_); + auto alt_d_a10 = bsk::get<13>(t183_); + auto alt_a11 = bsk::get<14>(t183_); + auto alt_d_a11 = bsk::get<15>(t183_); + auto alt_a20 = bsk::get<16>(t183_); + auto alt_d_a20 = bsk::get<17>(t183_); + auto alt_a22 = bsk::get<18>(t183_); + auto alt_d_a22 = bsk::get<19>(t183_); + auto alt_s00 = bsk::get<20>(t183_); + auto alt_d_s00 = bsk::get<21>(t183_); + auto alt_s11 = bsk::get<22>(t183_); + auto alt_d_s11 = bsk::get<23>(t183_); + auto alt_s22 = bsk::get<24>(t183_); + auto alt_d_s22 = bsk::get<25>(t183_); + auto alt_minors = bsk::get<26>(t183_); + auto alt_d_minors = bsk::get<27>(t183_); + auto alt_sum_flat = bsk::get<28>(t183_); + auto alt_sum_linear = bsk::get<29>(t183_); + auto alt_sum_square = bsk::get<30>(t183_); + auto alt_d_sum_flat = bsk::get<31>(t183_); + auto alt_d_sum_linear = bsk::get<32>(t183_); + auto alt_d_sum_square = bsk::get<33>(t183_); + auto alt_lift = bsk::get<34>(t183_); + auto alt_d_lift = bsk::get<35>(t183_); + auto alt_low = bsk::get<36>(t183_); + auto alt_middle = bsk::get<37>(t183_); + auto alt_d_low = bsk::get<38>(t183_); + auto alt_d_middle = bsk::get<39>(t183_); + auto alt_leading = bsk::get<40>(t183_); + auto alt_d_leading = bsk::get<41>(t183_); + auto alt_first = bsk::get<42>(t183_); + auto alt_d_first = bsk::get<43>(t183_); + auto alt_second = bsk::get<44>(t183_); + auto alt_d_second = bsk::get<45>(t183_); + auto alt_determinant = bsk::get<46>(t183_); + auto alt_d_determinant = bsk::get<47>(t183_); + auto alt_high = bsk::get<48>(t183_); + auto alt_d_high = bsk::get<49>(t183_); + auto alt_radius = bsk::get<50>(t183_); + auto alt_d_radius = bsk::get<51>(t183_); + auto alt_cube = bsk::get<52>(t183_); + auto alt_raw = bsk::get<53>(t183_); + auto alt_d_raw = bsk::get<54>(t183_); + auto alt_argument = bsk::get<55>(t183_); + auto alt_inside_limit = bsk::get<56>(t183_); + auto alt_angle = bsk::get<57>(t183_); + auto alt_d_angle = bsk::get<58>(t183_); + auto alt_centre = bsk::get<59>(t183_); + auto alt_d_centre = bsk::get<60>(t183_); + auto alt_trailing = bsk::get<61>(t183_); + auto alt_d_trailing = bsk::get<62>(t183_); + auto alt_guarded = bsk::get<63>(t183_); + auto alt_d_guarded = bsk::get<64>(t183_); + auto alt_q00 = bsk::get<65>(t183_); + auto alt_d_q00 = bsk::get<66>(t183_); + auto alt_q01 = bsk::get<67>(t183_); + auto alt_d_q01 = bsk::get<68>(t183_); + auto alt_q02 = bsk::get<69>(t183_); + auto alt_d_q02 = bsk::get<70>(t183_); + auto alt_q10 = bsk::get<71>(t183_); + auto alt_d_q10 = bsk::get<72>(t183_); + auto alt_q11 = bsk::get<73>(t183_); + auto alt_d_q11 = bsk::get<74>(t183_); + auto alt_q12 = bsk::get<75>(t183_); + auto alt_d_q12 = bsk::get<76>(t183_); + auto alt_q20 = bsk::get<77>(t183_); + auto alt_d_q20 = bsk::get<78>(t183_); + auto alt_q21 = bsk::get<79>(t183_); + auto alt_d_q21 = bsk::get<80>(t183_); + auto alt_q22 = bsk::get<81>(t183_); + auto alt_d_q22 = bsk::get<82>(t183_); + auto t184_ = _three_pool_assemble_jvp(alt_free, alt_d_free, alt_pool_b, alt_d_pool_b, alt_pool_c, alt_d_pool_c, alt_a00, alt_d_a00, alt_a01, alt_d_a01, alt_a02, alt_d_a02, alt_a10, alt_d_a10, alt_a11, alt_d_a11, alt_a20, alt_d_a20, alt_a22, alt_d_a22, alt_s00, alt_d_s00, alt_s11, alt_d_s11, alt_s22, alt_d_s22, alt_minors, alt_d_minors, alt_sum_flat, alt_sum_linear, alt_sum_square, alt_d_sum_flat, alt_d_sum_linear, alt_d_sum_square, alt_lift, alt_d_lift, alt_low, alt_middle, alt_d_low, alt_d_middle, alt_leading, alt_d_leading, alt_first, alt_d_first, alt_second, alt_d_second, alt_determinant, alt_d_determinant, alt_high, alt_d_high, alt_radius, alt_d_radius, alt_cube, alt_raw, alt_d_raw, alt_argument, alt_inside_limit, alt_angle, alt_d_angle, alt_centre, alt_d_centre, alt_trailing, alt_d_trailing, alt_guarded, alt_d_guarded, alt_q00, alt_d_q00, alt_q01, alt_d_q01, alt_q02, alt_d_q02, alt_q10, alt_d_q10, alt_q11, alt_d_q11, alt_q12, alt_d_q12, alt_q20, alt_d_q20, alt_q21, alt_d_q21, alt_q22, alt_d_q22, narrow); + auto alt_def_00 = bsk::get<0>(t184_); + auto alt_dif_00 = bsk::get<1>(t184_); + auto alt_def_01 = bsk::get<2>(t184_); + auto alt_dif_01 = bsk::get<3>(t184_); + auto alt_def_02 = bsk::get<4>(t184_); + auto alt_dif_02 = bsk::get<5>(t184_); + auto alt_def_10 = bsk::get<6>(t184_); + auto alt_dif_10 = bsk::get<7>(t184_); + auto alt_def_11 = bsk::get<8>(t184_); + auto alt_dif_11 = bsk::get<9>(t184_); + auto alt_def_12 = bsk::get<10>(t184_); + auto alt_dif_12 = bsk::get<11>(t184_); + auto alt_def_20 = bsk::get<12>(t184_); + auto alt_dif_20 = bsk::get<13>(t184_); + auto alt_def_21 = bsk::get<14>(t184_); + auto alt_dif_21 = bsk::get<15>(t184_); + auto alt_def_22 = bsk::get<16>(t184_); + auto alt_dif_22 = bsk::get<17>(t184_); + auto t185_ = _three_pool_step_adjoint_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, row_dt, unit, one_att, att_span, bsk::ld((held + 24), active_atom, 0.0f), nil, bsk::ld((held + 25), active_atom, 0.0f), nil, bsk::ld((held + 26), active_atom, 0.0f), nil, bsk::ld((held + 27), active_atom, 0.0f), nil, bsk::ld((held + 28), active_atom, 0.0f), nil, bsk::ld((held + 29), active_atom, 0.0f), nil, bsk::ld((held + 30), active_atom, 0.0f), nil, bsk::ld((held + 31), active_atom, 0.0f), nil, bsk::ld((held + 32), active_atom, 0.0f), nil, bsk::ld((held + 33), active_atom, 0.0f), nil, bsk::ld((held + 34), active_atom, 0.0f), nil, bsk::ld((held + 35), active_atom, 0.0f), nil, alt_free, alt_d_free, alt_pool_b, alt_d_pool_b, alt_pool_c, alt_d_pool_c, alt_a00, alt_d_a00, alt_a01, alt_d_a01, alt_a02, alt_d_a02, alt_a10, alt_d_a10, alt_a11, alt_d_a11, alt_a20, alt_d_a20, alt_a22, alt_d_a22, alt_s00, alt_d_s00, alt_s11, alt_d_s11, alt_s22, alt_d_s22, alt_minors, alt_d_minors, alt_sum_flat, alt_sum_linear, alt_sum_square, alt_d_sum_flat, alt_d_sum_linear, alt_d_sum_square, alt_lift, alt_d_lift, alt_low, alt_middle, alt_d_low, alt_d_middle, alt_leading, alt_d_leading, alt_first, alt_d_first, alt_second, alt_d_second, alt_determinant, alt_d_determinant, alt_high, alt_d_high, alt_radius, alt_d_radius, alt_cube, alt_raw, alt_d_raw, alt_argument, alt_inside_limit, alt_angle, alt_d_angle, alt_centre, alt_d_centre, alt_trailing, alt_d_trailing, alt_guarded, alt_d_guarded, alt_q00, alt_d_q00, alt_q01, alt_d_q01, alt_q02, alt_d_q02, alt_q10, alt_d_q10, alt_q11, alt_d_q11, alt_q12, alt_d_q12, alt_q20, alt_d_q20, alt_q21, alt_d_q21, alt_q22, alt_d_q22, alt_def_00, alt_dif_00, alt_def_01, alt_dif_01, alt_def_02, alt_dif_02, alt_def_10, alt_dif_10, alt_def_11, alt_dif_11, alt_def_12, alt_dif_12, alt_def_20, alt_dif_20, alt_def_21, alt_dif_21, alt_def_22, alt_dif_22, narrow); + auto _span_r1_v = bsk::get<0>(t185_); + auto _span_r1b_v = bsk::get<1>(t185_); + auto _span_r1c_v = bsk::get<2>(t185_); + auto _span_exch_v = bsk::get<3>(t185_); + auto _span_sexch_v = bsk::get<4>(t185_); + auto _span_bound_v = bsk::get<5>(t185_); + auto _span_semi_v = bsk::get<6>(t185_); + auto _span_dt_v = bsk::get<7>(t185_); + auto _span_att_v = bsk::get<8>(t185_); + auto span_r1_t = bsk::get<9>(t185_); + auto span_r1b_t = bsk::get<10>(t185_); + auto span_r1c_t = bsk::get<11>(t185_); + auto span_exch_t = bsk::get<12>(t185_); + auto span_sexch_t = bsk::get<13>(t185_); + auto span_bound_t = bsk::get<14>(t185_); + auto span_semi_t = bsk::get<15>(t185_); + auto _span_dt_t = bsk::get<16>(t185_); + auto _span_att_t = bsk::get<17>(t185_); + slope1_v = bsk::truediv(-1000.0f, (atom_t1 * atom_t1)); + slope1_t = bsk::truediv((2000.0f * d_t1), ((atom_t1 * atom_t1) * atom_t1)); + slope1b_v = bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)); + slope1b_t = bsk::truediv((2000.0f * d_t1b), ((atom_t1b * atom_t1b) * atom_t1b)); + slope1c_v = bsk::truediv(-1000.0f, (held_semisolid * held_semisolid)); + slope1c_t = bsk::truediv((2000.0f * d_semisolid_t1), ((held_semisolid * held_semisolid) * held_semisolid)); + auto row_r1_t = (back_r1_t + span_r1_t); + auto row_r1b_t = (back_r1b_t + span_r1b_t); + auto row_r1c_t = (back_r1c_t + span_r1c_t); + g_t1v = (g_t1v + (back_r1_v * slope1_v)); + g_t1t = (g_t1t + ((row_r1_t * slope1_v) + (back_r1_v * slope1_t))); + g_t1bv = (g_t1bv + (back_r1b_v * slope1b_v)); + g_t1bt = (g_t1bt + ((row_r1b_t * slope1b_v) + (back_r1b_v * slope1b_t))); + g_t1cv = (g_t1cv + (back_r1c_v * slope1c_v)); + g_t1ct = (g_t1ct + ((row_r1c_t * slope1c_v) + (back_r1c_v * slope1c_t))); + g_exchv = (g_exchv + back_exch_v); + g_excht = (g_excht + (back_exch_t + span_exch_t)); + g_sexchv = (g_sexchv + back_sexch_v); + g_sexcht = (g_sexcht + (back_sexch_t + span_sexch_t)); + g_boundv = (g_boundv + back_bound_v); + g_boundt = (g_boundt + (back_bound_t + span_bound_t)); + g_semiv = (g_semiv + back_semi_v); + g_semit = (g_semit + (back_semi_t + span_semi_t)); + } + } + auto velocity_v = ((g_flowv * flow_scale) + ((g_washv * direction) * washout_scale)); + auto velocity_t = ((g_flowt * flow_scale) + ((g_washt * direction) * washout_scale)); + if (bsk::truth((pools > 0))) { + // The fraction also sets where each pool starts, which the walk back + // reaches last. + g_boundv = (g_boundv + bsk::sum_x(bsk::where((state == 0), (bbvr - zbvr), 0.0f))); + g_boundt = (g_boundt + bsk::sum_x(bsk::where((state == 0), (bbtr - zbtr), 0.0f))); + } + if (bsk::truth((pools == 3))) { + g_semiv = (g_semiv + bsk::sum_x(bsk::where((state == 0), (cbvr - zbvr), 0.0f))); + g_semit = (g_semit + bsk::sum_x(bsk::where((state == 0), (cbtr - zbtr), 0.0f))); + } + if (bsk::truth((pools == 1))) { + base_row = (9 + (2 * (shim_rows - 1))); + bsk::atomic_add(((grad_tissue_value + (base_row * atom_count)) + atom), g_boundv, active_atom); + bsk::atomic_add(((grad_tissue_tangent + (base_row * atom_count)) + atom), g_boundt, active_atom); + bsk::atomic_add(((grad_tissue_value + ((base_row + 1) * atom_count)) + atom), g_exchv, active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((base_row + 1) * atom_count)) + atom), g_excht, active_atom); + bsk::atomic_add(((grad_tissue_value + ((base_row + 2) * atom_count)) + atom), g_t1bv, active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((base_row + 2) * atom_count)) + atom), g_t1bt, active_atom); + } + if (bsk::truth((pools == 3))) { + auto semisolid_row = (9 + (2 * (shim_rows - 1))); + auto stuck = bsk::make_tup(g_semiv, g_sexchv, g_t1cv); + auto stuck_tangents = bsk::make_tup(g_semit, g_sexcht, g_t1ct); + bsk::static_for<0, 3, 1>([&](auto offset_c) { + constexpr std::int64_t offset = decltype(offset_c)::value; + bsk::atomic_add(((grad_tissue_value + ((semisolid_row + offset) * atom_count)) + atom), bsk::get(stuck), active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((semisolid_row + offset) * atom_count)) + atom), bsk::get(stuck_tangents), active_atom); + }); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + base_row = (12 + (2 * (shim_rows - 1))); + auto rows = bsk::make_tup(g_boundv, g_exchv, g_t1bv, g_t2bv, g_shiftv); + auto tangent_rows = bsk::make_tup(g_boundt, g_excht, g_t1bt, g_t2bt, g_shiftt); + bsk::static_for<0, 5, 1>([&](auto offset_c) { + constexpr std::int64_t offset = decltype(offset_c)::value; + bsk::atomic_add(((grad_tissue_value + ((base_row + offset) * atom_count)) + atom), bsk::get(rows), active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((base_row + offset) * atom_count)) + atom), bsk::get(tangent_rows), active_atom); + }); + } + auto values = bsk::make_tup(g_t1v, g_t2v, g_m0v, g_b1v, g_b1pv, g_b0v, g_invv, g_diffv, velocity_v); + auto tangents = bsk::make_tup(g_t1t, g_t2t, g_m0t, g_b1t, g_b1pt, g_b0t, g_invt, g_difft, velocity_t); + bsk::static_for<0, 9, 1>([&](auto parameter_c) { + constexpr std::int64_t parameter = decltype(parameter_c)::value; + // The transmit pair went to its shim's row above when there is more + // than one; the rest sit past whatever rows that pair took. + if (bsk::truth((bsk::truth((!bsk::truth(shimmed))) || bsk::truth((bsk::truth((parameter != 3)) && bsk::truth((parameter != 4))))))) { + auto plane = bsk::select(bsk::truth((parameter < 3)), parameter, (parameter + (2 * (shim_rows - 1)))); + bsk::atomic_add(((grad_tissue_value + (plane * atom_count)) + atom), bsk::get(values), active_atom); + bsk::atomic_add(((grad_tissue_tangent + (plane * atom_count)) + atom), bsk::get(tangents), active_atom); + } + }); +} + +// Longitudinal and transverse diffusion damping for one interval. +// +// ``rate`` already carries the sequence's gradient geometry, so an interval's +// b-factor is that rate times its duration. Order zero has no longitudinal +// weight, which is what keeps the recovery term undamped. +template +BSK_HD auto _damping(const T0& rate, const T1& dt, const T2& order) { + auto b_factor = (rate * dt); + auto squared = (order * order); + return bsk::make_tup(bsk::exp(((-b_factor) * squared)), bsk::exp(((-b_factor) * ((squared + order) + 0.3333333333333333f)))); +} + +// How well the bound pool absorbs a pulse this far off its resonance. +// +// Cubic Hermite between the two knots bracketing the offset, taken in +// magnitude because the lineshape is even, and clamped at the far end. Each +// knot is two floats -- the value then its slope -- so the two a read needs +// are four contiguous ones. +template +BSK_HD auto _lineshape_at(const T0& lineshape, const T1& offset_hz, const T2& bins, const T3& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::truediv(bsk::abs(offset_hz), step), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto base = (bsk::cast(lower) * 2); + auto near = bsk::ld((lineshape + base)); + auto near_slope = bsk::ld(((lineshape + base) + 1)); + auto far = bsk::ld(((lineshape + base) + 2)); + auto far_slope = bsk::ld(((lineshape + base) + 3)); + return (((((((2.0f * u3) - (3.0f * u2)) + 1.0f) * near) + ((((u3 - (2.0f * u2)) + u) * step) * near_slope)) + (((-2.0f * u3) + (3.0f * u2)) * far)) + (((u3 - u2) * step) * far_slope)); +} + +// The Cayley-Klein pair the transition table holds at this flip angle. +// +// Cubic Hermite between the two knots bracketing ``theta``, clamped at both +// ends: a cubic run off its grid leaves the unit circle. Each knot is eight +// floats -- the pair then its slope, real before imaginary -- so the two a +// read needs are sixteen contiguous ones. +template +BSK_HD auto _profile_pair(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto component = [&](int c) { + auto near = bsk::ld(((profile + base) + c)); + auto near_slope = bsk::ld((((profile + base) + 4) + c)); + auto far = bsk::ld((((profile + base) + 8) + c)); + auto far_slope = bsk::ld((((profile + base) + 12) + c)); + return ((((h00 * near) + (h10 * near_slope)) + (h01 * far)) + (h11 * far_slope)); + }; + return bsk::make_tup(component(0), component(1), component(2), component(3)); +} + +// One pool through a hard pulse named by its flip angle and phase. +// +// Pulled out of the kernel body so a second pool can take the same rotation: +// a chemical shift moves where a pool precesses, not what a pulse does to it. +template +BSK_HD auto _rotate_flip_phase(const T0& cosine, const T1& sine, const T2& cos_phi, const T3& sin_phi, const T4& cos_2phi, const T5& sin_2phi, const T6& fp_r, const T7& fp_i, const T8& fm_r, const T9& fm_i, const T10& z_r, const T11& z_i) { + auto cosine_half_sq = (0.5f * (1.0f + cosine)); + auto sine_half_sq = (0.5f * (1.0f - cosine)); + auto half_sine = (0.5f * sine); + // Every sum of products is one fused multiply-add in a fixed order, so a + // kernel compiled for fewer terms rounds the rotation as the full one does. + auto minus_2phi_r = bsk::fma(cos_2phi, fm_r, (-(sin_2phi * fm_i))); + auto minus_2phi_i = bsk::fma(sin_2phi, fm_r, (cos_2phi * fm_i)); + auto plus_2phi_r = bsk::fma(cos_2phi, fp_r, (sin_2phi * fp_i)); + auto plus_2phi_i = bsk::fma(cos_2phi, fp_i, (-(sin_2phi * fp_r))); + auto z_turn_a = bsk::fma(sin_phi, z_r, (cos_phi * z_i)); + auto z_turn_b = bsk::fma(sin_phi, z_i, (-(cos_phi * z_r))); + auto z_turn_c = bsk::fma(sin_phi, z_r, (-(cos_phi * z_i))); + auto z_turn_d = bsk::fma(cos_phi, z_r, (sin_phi * z_i)); + auto rotated_pr = bsk::fma(sine, z_turn_a, bsk::fma(sine_half_sq, minus_2phi_r, (cosine_half_sq * fp_r))); + auto rotated_pi = bsk::fma(sine, z_turn_b, bsk::fma(sine_half_sq, minus_2phi_i, (cosine_half_sq * fp_i))); + auto rotated_mr = bsk::fma(sine, z_turn_c, bsk::fma(cosine_half_sq, fm_r, (sine_half_sq * plus_2phi_r))); + auto rotated_mi = bsk::fma(sine, z_turn_d, bsk::fma(cosine_half_sq, fm_i, (sine_half_sq * plus_2phi_i))); + auto plus_turn_r = bsk::fma(sin_phi, fp_r, (-(cos_phi * fp_i))); + auto minus_turn_r = bsk::fma(sin_phi, fm_r, (cos_phi * fm_i)); + auto plus_turn_i = bsk::fma(cos_phi, fp_r, (sin_phi * fp_i)); + auto minus_turn_i = bsk::fma(cos_phi, fm_r, (-(sin_phi * fm_i))); + auto rotated_zr = bsk::fma(cosine, z_r, bsk::fma((-half_sine), minus_turn_r, ((-half_sine) * plus_turn_r))); + auto rotated_zi = bsk::fma(cosine, z_i, bsk::fma(half_sine, minus_turn_i, ((-half_sine) * plus_turn_i))); + return bsk::make_tup(rotated_pr, rotated_pi, rotated_mr, rotated_mi, rotated_zr, rotated_zi); +} + +// The rotation named by its Cayley-Klein pair, applied to the states. +// +// T = [ conj(a)^2 -conj(b)^2 -2 conj(a b) ] +// [ -b^2 a^2 -2 a b ] +// [ conj(a) b a conj(b) |a|^2-|b|^2 ] +template +BSK_HD auto _rotate_spinor(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& fp_r, const T5& fp_i, const T6& fm_r, const T7& fm_i, const T8& z_r, const T9& z_i) { + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto t0_ = bsk::make_tup(aa_r, (-aa_i)); + auto t00_r = bsk::get<0>(t0_); + auto t00_i = bsk::get<1>(t0_); + auto t1_ = bsk::make_tup((-bb_r), bb_i); + auto t01_r = bsk::get<0>(t1_); + auto t01_i = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i)); + auto t02_r = bsk::get<0>(t2_); + auto t02_i = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup((-bb_r), (-bb_i)); + auto t10_r = bsk::get<0>(t3_); + auto t10_i = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(aa_r, aa_i); + auto t11_r = bsk::get<0>(t4_); + auto t11_i = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i)); + auto t12_r = bsk::get<0>(t5_); + auto t12_i = bsk::get<1>(t5_); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto t6_ = bsk::make_tup(cross_r, cross_i); + auto t20_r = bsk::get<0>(t6_); + auto t20_i = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(cross_r, (-cross_i)); + auto t21_r = bsk::get<0>(t7_); + auto t21_i = bsk::get<1>(t7_); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + auto out_pr = ((((((t00_r * fp_r) - (t00_i * fp_i)) + (t01_r * fm_r)) - (t01_i * fm_i)) + (t02_r * z_r)) - (t02_i * z_i)); + auto out_pi = ((((((t00_r * fp_i) + (t00_i * fp_r)) + (t01_r * fm_i)) + (t01_i * fm_r)) + (t02_r * z_i)) + (t02_i * z_r)); + auto out_mr = ((((((t10_r * fp_r) - (t10_i * fp_i)) + (t11_r * fm_r)) - (t11_i * fm_i)) + (t12_r * z_r)) - (t12_i * z_i)); + auto out_mi = ((((((t10_r * fp_i) + (t10_i * fp_r)) + (t11_r * fm_i)) + (t11_i * fm_r)) + (t12_r * z_i)) + (t12_i * z_r)); + auto out_zr = (((((t20_r * fp_r) - (t20_i * fp_i)) + (t21_r * fm_r)) - (t21_i * fm_i)) + (t22 * z_r)); + auto out_zi = (((((t20_r * fp_i) + (t20_i * fp_r)) + (t21_r * fm_i)) + (t21_i * fm_r)) + (t22 * z_i)); + return bsk::make_tup(out_pr, out_pi, out_mr, out_mi, out_zr, out_zi); +} + +// The sine and cosine of ``x`` from one reduction by a quarter turn. +// +// Cody and Waite's three-part quarter turn and the single-precision Cephes +// polynomials on the eighth turn either side of zero, to about an ulp where +// ``|x|`` is a flip angle; one reduction serves both where two library calls +// would each make their own. +template +BSK_HD auto _sincos(const T0& x) { + bsk::tile_t | 0, 2)> cosine{}; + bsk::tile_t | 0, 2)> r{}; + bsk::tile_t | 0, 2)> sine{}; + auto quarter = bsk::rint((x * 0.6366197723675814f)); + r = bsk::fma((-quarter), 1.5703125f, x); + r = bsk::fma((-quarter), 0.0004837512969970703f, r); + r = bsk::fma((-quarter), 7.549789954891882e-08f, r); + auto r2 = (r * r); + sine = bsk::fma(-0.00019515295891f, r2, 0.0083321608736f); + sine = bsk::fma(sine, r2, -0.16666654611f); + sine = bsk::fma((r * r2), sine, r); + cosine = bsk::fma(2.443315711809948e-05f, r2, -0.001388731625493765f); + cosine = bsk::fma(cosine, r2, 0.04166664568298827f); + cosine = bsk::fma((r2 * r2), cosine, bsk::fma(-0.5f, r2, 1.0f)); + auto q = bsk::band(bsk::cast(quarter), 3); + auto s = bsk::where((q == 0), sine, bsk::where((q == 1), cosine, bsk::where((q == 2), (-sine), (-cosine)))); + auto c = bsk::where((q == 0), cosine, bsk::where((q == 1), (-sine), bsk::where((q == 2), (-cosine), sine))); + return bsk::make_tup(s, c); +} + +// What each pool recovers over the interval, beside the operator itself. +// +// Returns the nine entries and the three recoveries, narrowed to float32 +// once they are an operator. +template +BSK_HD auto _three_pool_recovery(const T0& e00, const T1& e01, const T2& e02, const T3& e10, const T4& e11, const T5& e12, const T6& e20, const T7& e21, const T8& e22, const T9& free, const T10& pool_b, const T11& pool_c) { + auto grow_free = (free - (((e00 * free) + (e01 * pool_b)) + (e02 * pool_c))); + auto grow_pool_b = (pool_b - (((e10 * free) + (e11 * pool_b)) + (e12 * pool_c))); + auto grow_bound = (pool_c - (((e20 * free) + (e21 * pool_b)) + (e22 * pool_c))); + return bsk::make_tup(bsk::cast(e00), bsk::cast(e01), bsk::cast(e02), bsk::cast(e10), bsk::cast(e11), bsk::cast(e12), bsk::cast(e20), bsk::cast(e21), bsk::cast(e22), bsk::cast(grow_free), bsk::cast(grow_pool_b), bsk::cast(grow_bound)); +} + +// Read one interval's three-pool operator, and what each pool recovers. +// +// The stored row is undamped, so the washout the event carries is applied +// here and the three recoveries follow from the damped entries -- which is +// what makes one row serve every event of the same length whatever its +// washout. +template +BSK_HD auto _three_pool_from_table(const T0& table, const T1& row, const T2& atom, const T3& voxel_count, const T4& mask, const T5& attenuation, const T6& free, const T7& pool_b, const T8& pool_c) { + auto base = ((table + (row * (9 * voxel_count))) + atom); + return _three_pool_recovery((attenuation * bsk::ld((base + (0 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (1 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (2 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (3 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (4 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (5 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (6 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (7 * voxel_count)), mask, 0.0f)), (attenuation * bsk::ld((base + (8 * voxel_count)), mask, 0.0f)), free, pool_b, pool_c); +} + +// ``[a, b] exp``, from exponentials the caller has already taken. +// +// Near the coalescence ``sinh(d)/d`` is even in the gap, so the series is a +// polynomial in its square; the exponential of the midpoint is reached from +// the lower one by a series too, because over a gap this small it is one. +template +BSK_HD auto _exp_difference(const T0& lower, const T1& upper, const T2& exp_lower, const T3& exp_upper) { + auto half = (0.5f * (upper - lower)); + auto near = (bsk::abs(half) < 0.0001f); + auto square = (half * half); + // exp(mid) * sinh(half)/half, with both factors expanded about zero. + auto series = ((exp_lower * ((1.0f + half) + (0.5f * square))) * (1.0f + bsk::truediv(square, 6.0f))); + auto gap = bsk::where(near, 1.0f, (upper - lower)); + return bsk::where(near, series, bsk::truediv((exp_upper - exp_lower), gap)); +} + +// ``expm((K - diag(R1)) t)`` for free water beside both second pools. +// +// Free water is pool a, the chemically exchanging pool b and the semisolid +// pool c; each second pool exchanges with the free water and not with the +// other. Returns the nine entries and the three recoveries, narrowed to +// float32 once they are an operator. +// +// Two branches, by how far apart the eigenvalues are. Where they are close +// the exponential's own series is reduced modulo the characteristic +// polynomial, which forms no root at all. Where they are far apart the +// interpolating polynomial is taken in Newton form at the three roots, each +// of which is non-positive, so a long interval cannot overflow. +// +// ``narrow`` says the caller has bounded the spread below +// :data:`blochsim.sequence._parameters.NARROW_SPREAD` for every voxel and +// every interval it will pass, so +// only the series can be reached. The roots then cost nothing, and the series +// holds the answer to float32 without being carried in double -- +// :func:`blochsim.sequence._parameters.narrow_three_pool` is what decides it. +template +BSK_HD auto _three_pool_step_in_precision(const T0& r1_free, const T1& r1_pool_b, const T2& r1_bound, const T3& exchange_b, const T4& exchange_c, const T5& fraction_b, const T6& fraction_c, const T7& dt, const T8& attenuation, const T9& narrow) { + using Ret = bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>>; + bsk::tile_t | 0, 3)> e00{}; + bsk::tile_t | 0, 3)> e01{}; + bsk::tile_t | 0, 3)> e02{}; + bsk::tile_t | 0, 3)> e10{}; + bsk::tile_t | 0, 3)> e11{}; + bsk::tile_t | 0, 3)> e12{}; + bsk::tile_t | 0, 3)> e20{}; + bsk::tile_t | 0, 3)> e21{}; + bsk::tile_t | 0, 3)> e22{}; + float factorial{}; + bsk::tile_t | 0, 3)> flat{}; + bsk::tile_t | 0, 3)> linear{}; + bsk::tile_t | 0, 3)> square{}; + bsk::tile_t | 0, 3)> sum_flat{}; + bsk::tile_t | 0, 3)> sum_linear{}; + bsk::tile_t | 0, 3)> sum_square{}; + auto terms = bsk::select(bsk::truth(narrow), 24, 16); + auto step = bsk::cast(dt); + auto free = bsk::cast(((Work(1.0) - fraction_b) - fraction_c)); + auto pool_b = bsk::cast(fraction_b); + auto pool_c = bsk::cast(fraction_c); + auto kab = (bsk::cast(exchange_b) * pool_b); + auto kba = (bsk::cast(exchange_b) * free); + auto kac = (bsk::cast(exchange_c) * pool_c); + auto kca = (bsk::cast(exchange_c) * free); + auto a00 = ((((-kab) - kac) - bsk::cast(r1_free)) * step); + auto a01 = (kba * step); + auto a02 = (kca * step); + auto a10 = (kab * step); + auto a11 = (((-kba) - bsk::cast(r1_pool_b)) * step); + auto a20 = (kac * step); + auto a22 = (((-kca) - bsk::cast(r1_bound)) * step); + auto third = bsk::truediv(((a00 + a11) + a22), Work(3.0)); + auto s00 = (a00 - third); + auto s11 = (a11 - third); + auto s22 = (a22 - third); + // The two second pools do not exchange, so the generator keeps a pair of + // structural zeros the products below are written around. + auto minors = (((((s00 * s11) - (a01 * a10)) + (s00 * s22)) - (a02 * a20)) + (s11 * s22)); + auto determinant = ((((s00 * s11) * s22) - (a01 * (a10 * s22))) + (a02 * ((-s11) * a20))); + // --- close together: the series reduced modulo x^3 + minors x - det --- + flat = (Work(1.0) + (Work(0.0) * third)); + linear = (Work(0.0) * third); + square = (Work(0.0) * third); + sum_flat = flat; + sum_linear = linear; + sum_square = square; + factorial = Work(1.0); + #pragma unroll + for (std::int64_t order = 1; order < terms; order += 1) { + auto next_flat = (square * determinant); + auto next_linear = (flat - (square * minors)); + auto next_square = linear; + flat = next_flat; + linear = next_linear; + square = next_square; + factorial = (factorial * order); + auto weight = bsk::truediv(Work(1.0), factorial); + sum_flat = (sum_flat + (weight * flat)); + sum_linear = (sum_linear + (weight * linear)); + sum_square = (sum_square + (weight * square)); + } + auto q00 = (((s00 * s00) + (a01 * a10)) + (a02 * a20)); + auto q01 = ((s00 * a01) + (a01 * s11)); + auto q02 = ((s00 * a02) + (a02 * s22)); + auto q10 = ((a10 * s00) + (s11 * a10)); + auto q11 = ((a10 * a01) + (s11 * s11)); + auto q12 = (a10 * a02); + auto q20 = ((a20 * s00) + (s22 * a20)); + auto q21 = (a20 * a01); + auto q22 = ((a20 * a02) + (s22 * s22)); + auto lift = bsk::exp(third); + auto c00 = (lift * ((sum_flat + (sum_linear * s00)) + (sum_square * q00))); + auto c01 = (lift * ((sum_linear * a01) + (sum_square * q01))); + auto c02 = (lift * ((sum_linear * a02) + (sum_square * q02))); + auto c10 = (lift * ((sum_linear * a10) + (sum_square * q10))); + auto c11 = (lift * ((sum_flat + (sum_linear * s11)) + (sum_square * q11))); + auto c12 = (lift * (sum_square * q12)); + auto c20 = (lift * ((sum_linear * a20) + (sum_square * q20))); + auto c21 = (lift * (sum_square * q21)); + auto c22 = (lift * ((sum_flat + (sum_linear * s22)) + (sum_square * q22))); + // --- far apart: the Newton form at the three roots --- + auto damp = bsk::cast(attenuation); + if (bsk::truth(narrow)) { + e00 = (damp * c00); + e01 = (damp * c01); + e02 = (damp * c02); + e10 = (damp * c10); + e11 = (damp * c11); + e12 = (damp * c12); + e20 = (damp * c20); + e21 = (damp * c21); + e22 = (damp * c22); + return bsk::convert(_three_pool_recovery(e00, e01, e02, e10, e11, e12, e20, e21, e22, free, pool_b, pool_c)); + } + auto radius = bsk::sqrt(bsk::maximum(((-minors) * Work(0.3333333333333333)), Work(1e-300))); + auto argument = bsk::minimum(bsk::maximum(bsk::truediv((Work(0.5) * determinant), ((radius * radius) * radius)), Work(-0.9999999999999999)), Work(0.9999999999999999)); + auto angle = bsk::truediv(bsk::acos(argument), Work(3.0)); + auto root_a = (((Work(2.0) * radius) * bsk::cos(angle)) + third); + auto root_b = (((Work(2.0) * radius) * bsk::cos((angle - Work(2.0943951023931957)))) + third); + auto root_c = (((Work(2.0) * radius) * bsk::cos((angle - Work(4.188790204786391)))) + third); + auto low = bsk::minimum(bsk::minimum(root_a, root_b), root_c); + auto high = bsk::maximum(bsk::maximum(root_a, root_b), root_c); + auto middle = bsk::maximum(bsk::minimum(root_a, root_b), bsk::minimum(bsk::maximum(root_a, root_b), root_c)); + // Three exponentials serve every divided difference between them. + auto leading = bsk::exp(low); + auto centre = bsk::exp(middle); + auto trailing = bsk::exp(high); + auto first = _exp_difference(low, middle, leading, centre); + auto span = (high - low); + auto second = bsk::truediv((_exp_difference(middle, high, centre, trailing) - first), bsk::where((span > Work(0.0)), span, Work(1.0))); + auto m00 = (a00 - low); + auto m11 = (a11 - low); + auto m22 = (a22 - low); + auto n00 = (a00 - middle); + auto n11 = (a11 - middle); + auto n22 = (a22 - middle); + auto p00 = (((m00 * n00) + (a01 * a10)) + (a02 * a20)); + auto p01 = ((m00 * a01) + (a01 * n11)); + auto p02 = ((m00 * a02) + (a02 * n22)); + auto p10 = ((a10 * n00) + (m11 * a10)); + auto p11 = ((a10 * a01) + (m11 * n11)); + auto p12 = (a10 * a02); + auto p20 = ((a20 * n00) + (m22 * a20)); + auto p21 = (a20 * a01); + auto p22 = ((a20 * a02) + (m22 * n22)); + auto d00 = ((leading + (first * m00)) + (second * p00)); + auto d01 = ((first * a01) + (second * p01)); + auto d02 = ((first * a02) + (second * p02)); + auto d10 = ((first * a10) + (second * p10)); + auto d11 = ((leading + (first * m11)) + (second * p11)); + auto d12 = (second * p12); + auto d20 = ((first * a20) + (second * p20)); + auto d21 = (second * p21); + auto d22 = ((leading + (first * m22)) + (second * p22)); + // The shifted roots sum to zero, so the sum of their squares is -2 * minors + // and none is larger than the root of that. + auto close = ((Work(-2.0) * minors) < Work(1.0)); + e00 = (damp * bsk::where(close, c00, d00)); + e01 = (damp * bsk::where(close, c01, d01)); + e02 = (damp * bsk::where(close, c02, d02)); + e10 = (damp * bsk::where(close, c10, d10)); + e11 = (damp * bsk::where(close, c11, d11)); + e12 = (damp * bsk::where(close, c12, d12)); + e20 = (damp * bsk::where(close, c20, d20)); + e21 = (damp * bsk::where(close, c21, d21)); + e22 = (damp * bsk::where(close, c22, d22)); + return bsk::convert(_three_pool_recovery(e00, e01, e02, e10, e11, e12, e20, e21, e22, free, pool_b, pool_c)); +} + +template +BSK_HD auto _three_pool_step(const T0& r1_free, const T1& r1_pool_b, const T2& r1_bound, const T3& exchange_b, const T4& exchange_c, const T5& fraction_b, const T6& fraction_c, const T7& dt, const T8& attenuation, const T9& narrow) { + using R = decltype(_three_pool_step_in_precision(r1_free, r1_pool_b, r1_bound, exchange_b, exchange_c, fraction_b, fraction_c, dt, attenuation, narrow)); + if (bsk::truth(narrow)) { + return bsk::convert(_three_pool_step_in_precision(r1_free, r1_pool_b, r1_bound, exchange_b, exchange_c, fraction_b, fraction_c, dt, attenuation, narrow)); + } + return _three_pool_step_in_precision(r1_free, r1_pool_b, r1_bound, exchange_b, exchange_c, fraction_b, fraction_c, dt, attenuation, narrow); +} + +// The two-pool longitudinal operator over one interval, and its recovery. +// +// ``expm((K - diag(R1)) t)`` in the exact 2x2 closed form. Its discriminant +// is a square plus a product of two non-negative rates, so the root is real +// and the branch a general exponential would need does not exist here. +// ``sinh(d)/d`` is taken by series near the origin, where the root has no +// derivative of its own. +// +// The equilibrium each pool relaxes toward is its own fraction, so the +// recovery is ``(I - E1) (1 - f, f)`` and needs no solve. Returned as +// ``(e11, e12, e21, e22, recovery_free, recovery_bound)``. +template +BSK_HD auto _two_pool_step(const T0& r1_free, const T1& r1_bound, const T2& exchange, const T3& bound, const T4& dt, const T5& attenuation) { + auto free = (1.0f - bound); + auto kab = (exchange * bound); + auto kba = (exchange * free); + auto l11 = (((-kab) - r1_free) * dt); + auto l12 = (kba * dt); + auto l21 = (kab * dt); + auto l22 = (((-kba) - r1_bound) * dt); + auto half_trace = (0.5f * (l11 + l22)); + auto half_gap = (0.5f * (l11 - l22)); + auto square = ((half_gap * half_gap) + (l12 * l21)); + // tau +/- d are the eigenvalues, both non-positive for a decaying system, + // so their exponentials are bounded by one. Formed that way rather than as + // e^tau cosh(d), which over a long interval is an underflow times an + // overflow. + auto root = bsk::sqrt(bsk::maximum(square, 0.0f)); + auto upper = bsk::exp((half_trace + root)); + auto lower = bsk::exp((half_trace - root)); + auto cosine = (0.5f * (upper + lower)); + auto turning = (square > 1e-12f); + auto guarded = bsk::where(turning, root, 1.0f); + auto scale = bsk::where(turning, bsk::truediv((0.5f * (upper - lower)), guarded), (bsk::exp(half_trace) * ((1.0f + bsk::truediv(square, 6.0f)) + bsk::truediv((square * square), 120.0f)))); + auto e11 = (attenuation * (cosine + (scale * half_gap))); + auto e12 = ((attenuation * scale) * l12); + auto e21 = ((attenuation * scale) * l21); + auto e22 = (attenuation * (cosine - (scale * half_gap))); + return bsk::make_tup(e11, e12, e21, e22, (free - ((e11 * free) + (e12 * bound))), (bound - ((e21 * free) + (e22 * bound)))); +} + +// The transverse operator of two chemically exchanging pools. +// +// ``expm((K - diag(R2) - 2 pi i diag(df)) t)``, in the closed form the +// longitudinal pair uses -- the numbers have become complex, the algebra has +// not. Returned as the four entries, each a pair of floats. +// +// A semisolid pool holds a share of the voxel without carrying any transverse +// magnetization, so it is absent from this 2x2 and present in ``free`` -- how +// much free water the exchange sees. +// +// There is no recovery term: transverse magnetization relaxes toward zero. +template +BSK_HD auto _two_pool_transverse_step(const T0& r2_free, const T1& r2_bound, const T2& exchange, const T3& bound, const T4& free, const T5& shift_hz, const T6& dt, const T7& attenuation) { + auto kab = (exchange * bound); + auto kba = (exchange * free); + auto l11 = (((-kab) - r2_free) * dt); + auto l12 = (kba * dt); + auto l21 = (kab * dt); + auto l22 = (((-kba) - r2_bound) * dt); + // Only pool b's offset appears: pool a sits at whatever off-resonance the + // free precession already carries the whole voxel through. + auto l22_imag = ((-6.283185307179586f * shift_hz) * dt); + auto trace_real = (0.5f * (l11 + l22)); + auto trace_imag = (0.5f * l22_imag); + auto gap_real = (0.5f * (l11 - l22)); + auto gap_imag = (-0.5f * l22_imag); + auto square_real = (((gap_real * gap_real) - (gap_imag * gap_imag)) + (l12 * l21)); + auto square_imag = ((2.0f * gap_real) * gap_imag); + auto t0_ = _complex_sqrt(square_real, square_imag); + auto root_real = bsk::get<0>(t0_); + auto root_imag = bsk::get<1>(t0_); + auto t1_ = _complex_exp((trace_real + root_real), (trace_imag + root_imag)); + auto upper_real = bsk::get<0>(t1_); + auto upper_imag = bsk::get<1>(t1_); + auto t2_ = _complex_exp((trace_real - root_real), (trace_imag - root_imag)); + auto lower_real = bsk::get<0>(t2_); + auto lower_imag = bsk::get<1>(t2_); + auto cos_real = (0.5f * (upper_real + lower_real)); + auto cos_imag = (0.5f * (upper_imag + lower_imag)); + // ``sinh(d)/d`` by series near the origin, where the root has no + // derivative of its own. + auto turning = (((square_real * square_real) + (square_imag * square_imag)) > 1e-24f); + auto half_real = (0.5f * (upper_real - lower_real)); + auto half_imag = (0.5f * (upper_imag - lower_imag)); + auto guard = bsk::where(turning, ((root_real * root_real) + (root_imag * root_imag)), 1.0f); + auto divided_real = bsk::truediv(((half_real * root_real) + (half_imag * root_imag)), guard); + auto divided_imag = bsk::truediv(((half_imag * root_real) - (half_real * root_imag)), guard); + auto t3_ = _complex_exp(trace_real, trace_imag); + auto plain_real = bsk::get<0>(t3_); + auto plain_imag = bsk::get<1>(t3_); + auto square2_real = ((square_real * square_real) - (square_imag * square_imag)); + auto square2_imag = ((2.0f * square_real) * square_imag); + auto poly_real = ((1.0f + bsk::truediv(square_real, 6.0f)) + bsk::truediv(square2_real, 120.0f)); + auto poly_imag = (bsk::truediv(square_imag, 6.0f) + bsk::truediv(square2_imag, 120.0f)); + auto series_real = ((plain_real * poly_real) - (plain_imag * poly_imag)); + auto series_imag = ((plain_real * poly_imag) + (plain_imag * poly_real)); + auto scale_real = bsk::where(turning, divided_real, series_real); + auto scale_imag = bsk::where(turning, divided_imag, series_imag); + auto off_real = ((scale_real * gap_real) - (scale_imag * gap_imag)); + auto off_imag = ((scale_real * gap_imag) + (scale_imag * gap_real)); + return bsk::make_tup((attenuation * (cos_real + off_real)), (attenuation * (cos_imag + off_imag)), ((attenuation * scale_real) * l12), ((attenuation * scale_imag) * l12), ((attenuation * scale_real) * l21), ((attenuation * scale_imag) * l21), (attenuation * (cos_real - off_real)), (attenuation * (cos_imag - off_imag))); +} + +// The fraction of a voxel's spins that stay put over one interval. +// +// Inflowing spins are taken to be fully relaxed and unexcited, which makes +// washout an affine map of the shape longitudinal recovery already has: +// +// wout * (Z * e1 + (1 - e1)) + win == Z * (e1 * wout) + (1 - e1 * wout) +// +// so scaling both relaxation factors by it carries the whole term. Clamped at +// one, past which the interval has replaced the voxel outright. +template +BSK_HD auto _washout(const T0& rate, const T1& dt) { + return (1.0f - bsk::minimum((rate * dt), 1.0f)); +} + +BSK_HD void _epg_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* bound_exchange, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, float* phase_cos, float* phase_sin, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V atom_b0{}; + bsk::V atom_b1{}; + bsk::V atom_b1_phase{}; + bsk::V atom_bound{}; + bsk::V atom_damping{}; + bsk::V atom_exchange{}; + bsk::V atom_flow{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V atom_r1_bound{}; + bsk::V atom_r1_semisolid{}; + bsk::V atom_r2_bound{}; + bsk::V atom_semisolid{}; + bsk::V atom_semisolid_exchange{}; + bsk::V atom_shift{}; + bsk::V atom_washout{}; + bsk::V b1_cos{}; + bsk::V b1_sin{}; + bsk::V b_rot_mi{}; + bsk::V b_rot_mr{}; + bsk::V b_rot_pi{}; + bsk::V b_rot_pr{}; + bsk::V b_rot_zi{}; + bsk::V b_rot_zr{}; + bsk::V bminus_imag{}; + bsk::V bminus_real{}; + bsk::V bound_imag{}; + bsk::V bound_real{}; + bsk::V bplus_imag{}; + bsk::V bplus_real{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bsk::V e1{}; + bsk::V e2{}; + bsk::V fminus_imag{}; + bsk::V fminus_real{}; + bsk::V fplus_imag{}; + bsk::V fplus_real{}; + bsk::V free_imag{}; + bsk::V free_real{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; + bsk::V held_imag{}; + bsk::V held_real{}; + bsk::V longitudinal_imag{}; + bsk::V longitudinal_real{}; + bsk::V off_cos{}; + bsk::V off_sin{}; + bsk::V old_real{}; + bsk::tup, bsk::V, bsk::V, bsk::V> pair{}; + bsk::V read_imag{}; + bsk::V read_real{}; + bsk::V rotated_mi{}; + bsk::V rotated_mr{}; + bsk::V rotated_pi{}; + bsk::V rotated_pr{}; + bsk::V rotated_zi{}; + bsk::V rotated_zr{}; + bsk::V semisolid_imag{}; + bsk::V semisolid_real{}; + bsk::V shaped_mi{}; + bsk::V shaped_mr{}; + bsk::V shaped_pi{}; + bsk::V shaped_pr{}; + bsk::V shaped_zi{}; + bsk::V shaped_zr{}; + bsk::V spun_bi{}; + bsk::V spun_br{}; + bsk::V t11{}; + bsk::V t12{}; + bsk::V t13{}; + bsk::V t21{}; + bsk::V t22{}; + bsk::V t23{}; + bsk::V t31{}; + bsk::V t32{}; + bsk::V t33{}; + bsk::V turn_cos{}; + bsk::V turn_sin{}; + bsk::V turn_t{}; + bsk::V turn_z{}; + bsk::V wout{}; + auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < (train_count * atom_count)); + // A partial block carries lanes with no problem behind them, and they must + // take no part in a reduction or a store. + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + // Voxels are spread over the slice voxel-major, so a voxel's place along + // the slice is its index modulo the profile's width. One pulse shape holds + // that many consecutive rows, and the event says which shape it drives. + auto location = bsk::mod(atom, locations); + auto empty = bsk::full(0); + fplus_real = empty; + fplus_imag = empty; + fminus_real = empty; + fminus_imag = empty; + // A second pool holds its own share of the equilibrium. The semisolid one + // carries longitudinal states alone -- nothing dephases it, so it reaches + // the higher orders only through exchange with the free pool's -- while the + // chemically exchanging one carries a transverse pair of its own. + // + // ``bound`` is whichever second pool the longitudinal step pairs the free + // water with -- the semisolid one when it is the only one, the exchanging + // one otherwise -- and ``semisolid`` is the third, which only a three-pool + // run carries. + atom_bound = 0.0f; + atom_exchange = 0.0f; + atom_r1_bound = 0.0f; + atom_r2_bound = 0.0f; + atom_shift = 0.0f; + atom_semisolid = 0.0f; + atom_semisolid_exchange = 0.0f; + atom_r1_semisolid = 0.0f; + if (bsk::truth((pools == 1))) { + atom_bound = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((bound_exchange + scalar_atom), active_atom, 0.0f); + atom_r1_bound = bsk::truediv(1000.0f, bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f)); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + atom_bound = bsk::ld((pool_b_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((pool_b_exchange + scalar_atom), active_atom, 0.0f); + atom_r1_bound = bsk::truediv(1000.0f, bsk::ld((t1_pool_b + scalar_atom), active_atom, 1.0f)); + atom_r2_bound = bsk::truediv(1000.0f, bsk::ld((t2_pool_b + scalar_atom), active_atom, 1.0f)); + atom_shift = bsk::ld((pool_b_shift + scalar_atom), active_atom, 0.0f); + } + if (bsk::truth((pools == 3))) { + atom_semisolid = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + atom_semisolid_exchange = bsk::ld((bound_exchange + scalar_atom), active_atom, 0.0f); + atom_r1_semisolid = bsk::truediv(1000.0f, bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f)); + } + // A semisolid pool holds a share of the voxel without carrying any + // transverse magnetization, so the 2x2 below is blind to it and the + // exchange inside that 2x2 is not. + auto atom_free = ((1.0f - atom_bound) - atom_semisolid); + longitudinal_real = (empty + bsk::where((state == 0), atom_free, 0.0f)); + longitudinal_imag = empty; + bound_real = (empty + bsk::where((state == 0), (atom_bound + 0.0f), 0.0f)); + bound_imag = empty; + semisolid_real = (empty + bsk::where((state == 0), (atom_semisolid + 0.0f), 0.0f)); + semisolid_imag = empty; + bplus_real = empty; + bplus_imag = empty; + bminus_real = empty; + bminus_imag = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_b1_phase = 0.0f; + atom_b0 = 0.0f; + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld((b1_phase + scalar_atom), active_atom, 0.0f); + atom_b0 = bsk::ld((b0 + scalar_atom), active_atom, 0.0f); + } + b1_cos = bsk::cos(atom_b1_phase); + b1_sin = bsk::sin(atom_b1_phase); + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + atom_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + } + atom_flow = 0.0f; + atom_washout = 0.0f; + if (bsk::truth(moving)) { + auto atom_velocity = bsk::ld((velocity + scalar_atom), active_atom, 0.0f); + atom_flow = (atom_velocity * flow_scale); + atom_washout = (bsk::abs(atom_velocity) * washout_scale); + } + auto order = bsk::cast(state); + auto event_base = (train * event_count); + for (std::int64_t event = 0; event < event_count; event += 1) { + auto dt = _event_value(duration, event_base, event, active_atom, single_train); + wout = 1.0f; + if (bsk::truth(moving)) { + wout = _washout(atom_washout, dt); + } + e1 = (bsk::exp(((-bsk::truediv(1000.0f, atom_t1)) * dt)) * wout); + e2 = (bsk::exp(((-bsk::truediv(1000.0f, atom_t2)) * dt)) * wout); + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping(atom_damping, dt, order); + damp_z = bsk::get<0>(t0_); + damp_t = bsk::get<1>(t0_); + } + turn_z = 0.0f; + turn_t = 0.0f; + if (bsk::truth(moving)) { + auto t1_ = _flow(atom_flow, dt, order); + turn_z = bsk::get<0>(t1_); + turn_t = bsk::get<1>(t1_); + } + auto recovery = (1.0f - e1); + e1 = (e1 * damp_z); + e2 = (e2 * damp_t); + off_cos = 1.0f; + off_sin = 0.0f; + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + // Flow winds the transverse states through the same rotation + // off-resonance does, so the two phases add before either is taken. + auto off_phase = (((-6.283185307179586f * atom_b0) * dt) + turn_t); + off_cos = bsk::cos(off_phase); + off_sin = bsk::sin(off_phase); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // Both pools take the same off-resonance and the same per-order + // damping; what separates them is the chemical shift, which the + // exchange operator already carries. + auto t2_ = _two_pool_transverse_step(bsk::truediv(1000.0f, atom_t2), atom_r2_bound, atom_exchange, atom_bound, atom_free, atom_shift, dt, wout); + auto x11r = bsk::get<0>(t2_); + auto x11i = bsk::get<1>(t2_); + auto x12r = bsk::get<2>(t2_); + auto x12i = bsk::get<3>(t2_); + auto x21r = bsk::get<4>(t2_); + auto x21i = bsk::get<5>(t2_); + auto x22r = bsk::get<6>(t2_); + auto x22i = bsk::get<7>(t2_); + auto mixed_pr = ((((x11r * fplus_real) - (x11i * fplus_imag)) + (x12r * bplus_real)) - (x12i * bplus_imag)); + auto mixed_pi = ((((x11r * fplus_imag) + (x11i * fplus_real)) + (x12r * bplus_imag)) + (x12i * bplus_real)); + auto mixed_br = ((((x21r * fplus_real) - (x21i * fplus_imag)) + (x22r * bplus_real)) - (x22i * bplus_imag)); + auto mixed_bi = ((((x21r * fplus_imag) + (x21i * fplus_real)) + (x22r * bplus_imag)) + (x22i * bplus_real)); + // ``F-`` follows the conjugate of the operator entry by entry, not + // its transpose: it is the conjugate state, and the map it takes is + // the conjugate map. + auto mixed_mr = ((((x11r * fminus_real) + (x11i * fminus_imag)) + (x12r * bminus_real)) + (x12i * bminus_imag)); + auto mixed_mi = ((((x11r * fminus_imag) - (x11i * fminus_real)) + (x12r * bminus_imag)) - (x12i * bminus_real)); + auto mixed_nr = ((((x21r * fminus_real) + (x21i * fminus_imag)) + (x22r * bminus_real)) + (x22i * bminus_imag)); + auto mixed_ni = ((((x21r * fminus_imag) - (x21i * fminus_real)) + (x22r * bminus_imag)) - (x22i * bminus_real)); + fplus_real = (damp_t * ((mixed_pr * off_cos) - (mixed_pi * off_sin))); + fplus_imag = (damp_t * ((mixed_pr * off_sin) + (mixed_pi * off_cos))); + bplus_real = (damp_t * ((mixed_br * off_cos) - (mixed_bi * off_sin))); + bplus_imag = (damp_t * ((mixed_br * off_sin) + (mixed_bi * off_cos))); + fminus_real = (damp_t * ((mixed_mr * off_cos) + (mixed_mi * off_sin))); + fminus_imag = (damp_t * (((-mixed_mr) * off_sin) + (mixed_mi * off_cos))); + bminus_real = (damp_t * ((mixed_nr * off_cos) + (mixed_ni * off_sin))); + bminus_imag = (damp_t * (((-mixed_nr) * off_sin) + (mixed_ni * off_cos))); + } else { + old_real = fplus_real; + fplus_real = (e2 * ((old_real * off_cos) - (fplus_imag * off_sin))); + fplus_imag = (e2 * ((old_real * off_sin) + (fplus_imag * off_cos))); + old_real = fminus_real; + fminus_real = (e2 * ((old_real * off_cos) + (fminus_imag * off_sin))); + fminus_imag = (e2 * (((-old_real) * off_sin) + (fminus_imag * off_cos))); + } + // The longitudinal states carry a phase of their own, which nothing + // else in the state machine gives them. + turn_cos = 1.0f; + turn_sin = 0.0f; + if (bsk::truth(moving)) { + turn_cos = bsk::cos(turn_z); + turn_sin = bsk::sin(turn_z); + } + if (bsk::truth((pools == 3))) { + // Three pools mix through a 3x3 formed in double; every pool takes + // the same per-order damping and flow phase, their order-n states + // describing one dephasing configuration. + if (bsk::truth(tabulated)) { + auto t3_ = _three_pool_from_table(pool_table, bsk::ld(((duration_row + event_base) + event), active_atom, 0), atom, atom_count, active_atom, wout, atom_free, atom_bound, atom_semisolid); + t11 = bsk::get<0>(t3_); + t12 = bsk::get<1>(t3_); + t13 = bsk::get<2>(t3_); + t21 = bsk::get<3>(t3_); + t22 = bsk::get<4>(t3_); + t23 = bsk::get<5>(t3_); + t31 = bsk::get<6>(t3_); + t32 = bsk::get<7>(t3_); + t33 = bsk::get<8>(t3_); + grow_free = bsk::get<9>(t3_); + grow_pool_b = bsk::get<10>(t3_); + grow_semisolid = bsk::get<11>(t3_); + } else { + auto t4_ = _three_pool_step(bsk::truediv(1000.0f, atom_t1), atom_r1_bound, atom_r1_semisolid, atom_exchange, atom_semisolid_exchange, atom_bound, atom_semisolid, dt, wout, narrow); + t11 = bsk::get<0>(t4_); + t12 = bsk::get<1>(t4_); + t13 = bsk::get<2>(t4_); + t21 = bsk::get<3>(t4_); + t22 = bsk::get<4>(t4_); + t23 = bsk::get<5>(t4_); + t31 = bsk::get<6>(t4_); + t32 = bsk::get<7>(t4_); + t33 = bsk::get<8>(t4_); + grow_free = bsk::get<9>(t4_); + grow_pool_b = bsk::get<10>(t4_); + grow_semisolid = bsk::get<11>(t4_); + } + free_real = (((t11 * longitudinal_real) + (t12 * bound_real)) + (t13 * semisolid_real)); + free_imag = (((t11 * longitudinal_imag) + (t12 * bound_imag)) + (t13 * semisolid_imag)); + held_real = (((t21 * longitudinal_real) + (t22 * bound_real)) + (t23 * semisolid_real)); + held_imag = (((t21 * longitudinal_imag) + (t22 * bound_imag)) + (t23 * semisolid_imag)); + auto stuck_real = (((t31 * longitudinal_real) + (t32 * bound_real)) + (t33 * semisolid_real)); + auto stuck_imag = (((t31 * longitudinal_imag) + (t32 * bound_imag)) + (t33 * semisolid_imag)); + longitudinal_real = (damp_z * ((free_real * turn_cos) - (free_imag * turn_sin))); + longitudinal_imag = (damp_z * ((free_real * turn_sin) + (free_imag * turn_cos))); + bound_real = (damp_z * ((held_real * turn_cos) - (held_imag * turn_sin))); + bound_imag = (damp_z * ((held_real * turn_sin) + (held_imag * turn_cos))); + semisolid_real = (damp_z * ((stuck_real * turn_cos) - (stuck_imag * turn_sin))); + semisolid_imag = (damp_z * ((stuck_real * turn_sin) + (stuck_imag * turn_cos))); + longitudinal_real = (longitudinal_real + bsk::where((state == 0), grow_free, 0.0f)); + bound_real = (bound_real + bsk::where((state == 0), grow_pool_b, 0.0f)); + semisolid_real = (semisolid_real + bsk::where((state == 0), grow_semisolid, 0.0f)); + } else if (bsk::truth((pools > 0))) { + // The exchange operator is a property of the interval, not of a + // dephasing order, so it is formed once and the per-order damping + // multiplies it. Both pools take that damping and the flow phase: + // their order-n states describe one dephasing configuration, and a + // second pool has no diffusion coefficient of its own to damp by. + auto t5_ = _two_pool_step(bsk::truediv(1000.0f, atom_t1), atom_r1_bound, atom_exchange, atom_bound, dt, wout); + auto e11 = bsk::get<0>(t5_); + auto e12 = bsk::get<1>(t5_); + auto e21 = bsk::get<2>(t5_); + auto e22 = bsk::get<3>(t5_); + grow_free = bsk::get<4>(t5_); + auto grow_bound = bsk::get<5>(t5_); + free_real = ((e11 * longitudinal_real) + (e12 * bound_real)); + free_imag = ((e11 * longitudinal_imag) + (e12 * bound_imag)); + held_real = ((e21 * longitudinal_real) + (e22 * bound_real)); + held_imag = ((e21 * longitudinal_imag) + (e22 * bound_imag)); + longitudinal_real = (damp_z * ((free_real * turn_cos) - (free_imag * turn_sin))); + longitudinal_imag = (damp_z * ((free_real * turn_sin) + (free_imag * turn_cos))); + bound_real = (damp_z * ((held_real * turn_cos) - (held_imag * turn_sin))); + bound_imag = (damp_z * ((held_real * turn_sin) + (held_imag * turn_cos))); + longitudinal_real = (longitudinal_real + bsk::where((state == 0), grow_free, 0.0f)); + bound_real = (bound_real + bsk::where((state == 0), grow_bound, 0.0f)); + } else { + old_real = longitudinal_real; + longitudinal_real = (e1 * ((old_real * turn_cos) - (longitudinal_imag * turn_sin))); + longitudinal_imag = (e1 * ((old_real * turn_sin) + (longitudinal_imag * turn_cos))); + longitudinal_real = (longitudinal_real + bsk::where((state == 0), recovery, 0.0f)); + } + // Every program reads the same event, so a branch on what it does is + // taken by all of them alike: an event pays only for what it does. + auto event_action = bsk::cast(bsk::ld((action + event))); + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t6_ = _shift(fplus_real, fplus_imag, fminus_real, fminus_imag, state, state_mask, state_count); + fplus_real = bsk::get<0>(t6_); + fplus_imag = bsk::get<1>(t6_); + fminus_real = bsk::get<2>(t6_); + fminus_imag = bsk::get<3>(t6_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t7_ = _shift(bplus_real, bplus_imag, bminus_real, bminus_imag, state, state_mask, state_count); + bplus_real = bsk::get<0>(t7_); + bplus_imag = bsk::get<1>(t7_); + bminus_real = bsk::get<2>(t7_); + bminus_imag = bsk::get<3>(t7_); + } + } + auto event_kind = bsk::ld((kind + event)); + auto is_rf = (event_kind == 1); + auto is_inversion = (bsk::band(event_action, 4) != 0); + auto invert = bsk::band(is_rf, is_inversion); + longitudinal_real = bsk::where(invert, ((-atom_inversion) * longitudinal_real), longitudinal_real); + longitudinal_imag = bsk::where(invert, ((-atom_inversion) * longitudinal_imag), longitudinal_imag); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // A chemically exchanging pool is free water and turns over like + // any other; a semisolid one is saturated instead, which its own + // saturation term already carries. + bound_real = bsk::where(invert, ((-atom_inversion) * bound_real), bound_real); + bound_imag = bsk::where(invert, ((-atom_inversion) * bound_imag), bound_imag); + } + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse a row of its own. + if (bsk::truth(shimmed)) { + auto row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + b1_cos = bsk::cos(atom_b1_phase); + b1_sin = bsk::sin(atom_b1_phase); + } + } + if (bsk::truth(bsk::band((event_kind == 1), (bsk::band(event_action, 4) == 0)))) { + auto alpha = (_event_value(flip, event_base, event, active_atom, single_train) * atom_b1); + // The pulse's phase, read off the cosine and sine the launch took + // of it, turned by the transmit field's own. + auto cos_event = _event_value(phase_cos, event_base, event, active_atom, single_train); + auto sin_event = _event_value(phase_sin, event_base, event, active_atom, single_train); + auto cos_phi = bsk::fma(cos_event, b1_cos, (-(sin_event * b1_sin))); + auto sin_phi = bsk::fma(sin_event, b1_cos, (cos_event * b1_sin)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + // Either pair is built at zero RF phase, which turns the rotation + // axis and so reaches ``b`` alone. + if (bsk::truth(dynamic)) { + // Already integrated at this pulse's own flip, so the flip is + // inside the pair rather than read against it. + pair = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, active_atom); + } else { + pair = _profile_pair(profile, _table_row(profile_index, event, location, locations), alpha, profile_bins, profile_step); + } + auto turn_r = cos_phi; + auto turn_i = (-sin_phi); + spun_br = ((bsk::get<2>(pair) * turn_r) - (bsk::get<3>(pair) * turn_i)); + spun_bi = ((bsk::get<2>(pair) * turn_i) + (bsk::get<3>(pair) * turn_r)); + auto t8_ = _rotate_spinor(bsk::get<0>(pair), bsk::get<1>(pair), spun_br, spun_bi, fplus_real, fplus_imag, fminus_real, fminus_imag, longitudinal_real, longitudinal_imag); + shaped_pr = bsk::get<0>(t8_); + shaped_pi = bsk::get<1>(t8_); + shaped_mr = bsk::get<2>(t8_); + shaped_mi = bsk::get<3>(t8_); + shaped_zr = bsk::get<4>(t8_); + shaped_zi = bsk::get<5>(t8_); + } + auto t9_ = _sincos(alpha); + auto sine = bsk::get<0>(t9_); + auto cosine = bsk::get<1>(t9_); + auto cos_2phi = bsk::fma(cos_phi, cos_phi, (-(sin_phi * sin_phi))); + auto sin_2phi = ((2.0f * sin_phi) * cos_phi); + auto t10_ = _rotate_flip_phase(cosine, sine, cos_phi, sin_phi, cos_2phi, sin_2phi, fplus_real, fplus_imag, fminus_real, fminus_imag, longitudinal_real, longitudinal_imag); + rotated_pr = bsk::get<0>(t10_); + rotated_pi = bsk::get<1>(t10_); + rotated_mr = bsk::get<2>(t10_); + rotated_mi = bsk::get<3>(t10_); + rotated_zr = bsk::get<4>(t10_); + rotated_zi = bsk::get<5>(t10_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t11_ = _rotate_flip_phase(cosine, sine, cos_phi, sin_phi, cos_2phi, sin_2phi, bplus_real, bplus_imag, bminus_real, bminus_imag, bound_real, bound_imag); + b_rot_pr = bsk::get<0>(t11_); + b_rot_pi = bsk::get<1>(t11_); + b_rot_mr = bsk::get<2>(t11_); + b_rot_mi = bsk::get<3>(t11_); + b_rot_zr = bsk::get<4>(t11_); + b_rot_zi = bsk::get<5>(t11_); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + // The same pulse, the same rotation: a chemical shift moves + // where a pool precesses, not what a pulse does to it. + auto t12_ = _rotate_spinor(bsk::get<0>(pair), bsk::get<1>(pair), spun_br, spun_bi, bplus_real, bplus_imag, bminus_real, bminus_imag, bound_real, bound_imag); + b_rot_pr = bsk::get<0>(t12_); + b_rot_pi = bsk::get<1>(t12_); + b_rot_mr = bsk::get<2>(t12_); + b_rot_mi = bsk::get<3>(t12_); + b_rot_zr = bsk::get<4>(t12_); + b_rot_zi = bsk::get<5>(t12_); + } + } + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + rotated_pr = shaped_pr; + rotated_pi = shaped_pi; + rotated_mr = shaped_mr; + rotated_mi = shaped_mi; + rotated_zr = shaped_zr; + rotated_zi = shaped_zi; + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + bplus_real = b_rot_pr; + bplus_imag = b_rot_pi; + bminus_real = b_rot_mr; + bminus_imag = b_rot_mi; + bound_real = b_rot_zr; + bound_imag = b_rot_zi; + } + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + // The semisolid pool absorbs the power the pulse deposits, so it + // reads the bare flip the transmit field gives the voxel -- not the + // slice-shaped rotation the free pool takes from the table. + auto offset = (bsk::ld((rf_frequency + event)) - atom_b0); + auto absorbed = bsk::exp((((bsk::ld((saturation + event)) * alpha) * alpha) * _lineshape_at(lineshape, offset, lineshape_bins, lineshape_step))); + if (bsk::truth((pools == 1))) { + bound_real = (absorbed * bound_real); + bound_imag = (absorbed * bound_imag); + } else { + semisolid_real = (absorbed * semisolid_real); + semisolid_imag = (absorbed * semisolid_imag); + } + } + fplus_real = rotated_pr; + fplus_imag = rotated_pi; + fminus_real = rotated_mr; + fminus_imag = rotated_mi; + longitudinal_real = rotated_zr; + longitudinal_imag = rotated_zi; + } + if (bsk::truth(bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)))) { + auto adc_cos = _event_value(phase_cos, event_base, event, active_atom, single_train); + auto adc_sin = _event_value(phase_sin, event_base, event, active_atom, single_train); + // A coil sees the whole voxel, so what it records is the sum over + // pools; each pool's share is already in its own state. + read_real = fplus_real; + read_imag = fplus_imag; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + read_real = (fplus_real + bplus_real); + read_imag = (fplus_imag + bplus_imag); + } + auto signal_real = (atom_m0 * ((read_real * adc_cos) + (read_imag * adc_sin))); + auto signal_imag = (atom_m0 * ((read_imag * adc_cos) - (read_real * adc_sin))); + auto out_ = bsk::ld((output_index + event)); + auto output_offset = ((problem * output_count) + out_); + auto output_mask = bsk::band(bsk::band(active_atom, (state == 0)), (out_ >= 0)); + bsk::st(((output_real + output_offset) + state), signal_real, output_mask); + bsk::st(((output_imag + output_offset) + state), signal_imag, output_mask); + } + if (bsk::truth((bsk::band(event_action, 18) != 0))) { + auto t13_ = _shift(fplus_real, fplus_imag, fminus_real, fminus_imag, state, state_mask, state_count); + fplus_real = bsk::get<0>(t13_); + fplus_imag = bsk::get<1>(t13_); + fminus_real = bsk::get<2>(t13_); + fminus_imag = bsk::get<3>(t13_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t14_ = _shift(bplus_real, bplus_imag, bminus_real, bminus_imag, state, state_mask, state_count); + bplus_real = bsk::get<0>(t14_); + bplus_imag = bsk::get<1>(t14_); + bminus_real = bsk::get<2>(t14_); + bminus_imag = bsk::get<3>(t14_); + } + } + if (bsk::truth((bsk::band(event_action, 8) != 0))) { + fplus_real = empty; + fplus_imag = empty; + fminus_real = empty; + fminus_imag = empty; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + bplus_real = empty; + bplus_imag = empty; + bminus_real = empty; + bminus_imag = empty; + } + } + } +} + +// Fill one row of the three-pool operator table. +// +// The row is ``expm((K - diag(R1)) dt)`` with no washout applied, so the +// event that reads it supplies its own attenuation. Laid out +// ``(rows, 9, voxels)`` -- entry-major over the voxel axis -- so the nine +// loads an event makes are each coalesced. +BSK_HD void _three_pool_table_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* durations, std::int32_t* rows, float* table, std::int64_t voxel_count, std::int64_t BLOCK, std::int64_t narrow) { + auto row = bsk::ld((rows + bsk::program_id(0))); + auto atom = ((bsk::program_id(1) * BLOCK) + bsk::arange_x()); + auto live = (atom < voxel_count); + auto dt = bsk::ld((durations + row)); + auto fraction_b = bsk::ld((pool_b_fraction + atom), live, 0.0f); + auto fraction_c = bsk::ld((bound_fraction + atom), live, 0.0f); + auto t0_ = _three_pool_step(bsk::truediv(1000.0f, bsk::ld((t1 + atom), live, 1.0f)), bsk::truediv(1000.0f, bsk::ld((t1_pool_b + atom), live, 1.0f)), bsk::truediv(1000.0f, bsk::ld((t1_bound + atom), live, 1.0f)), bsk::ld((pool_b_exchange + atom), live, 0.0f), bsk::ld((bound_exchange + atom), live, 0.0f), fraction_b, fraction_c, dt, (1.0f + (0.0f * dt)), narrow); + auto e00 = bsk::get<0>(t0_); + auto e01 = bsk::get<1>(t0_); + auto e02 = bsk::get<2>(t0_); + auto e10 = bsk::get<3>(t0_); + auto e11 = bsk::get<4>(t0_); + auto e12 = bsk::get<5>(t0_); + auto e20 = bsk::get<6>(t0_); + auto e21 = bsk::get<7>(t0_); + auto e22 = bsk::get<8>(t0_); + auto base = ((table + (row * (9 * voxel_count))) + atom); + bsk::st((base + (0 * voxel_count)), e00, live); + bsk::st((base + (1 * voxel_count)), e01, live); + bsk::st((base + (2 * voxel_count)), e02, live); + bsk::st((base + (3 * voxel_count)), e10, live); + bsk::st((base + (4 * voxel_count)), e11, live); + bsk::st((base + (5 * voxel_count)), e12, live); + bsk::st((base + (6 * voxel_count)), e20, live); + bsk::st((base + (7 * voxel_count)), e21, live); + bsk::st((base + (8 * voxel_count)), e22, live); +} + +// Fill one row of the three-pool operator table, value and direction. +// +// The row holds nine undamped entries and the nine a direction through the +// tissue gives them, both at ``d_dt`` of zero -- the interval's own share of +// the direction is ``A1 C d_dt``, which the reading event adds because +// ``d_dt`` is its own and the row's is not. Laid out ``(rows, 18, voxels)``, +// the tangent following the value. +BSK_HD void _three_pool_table_jvp_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* d_t1, float* d_t1_pool_b, float* d_t1_bound, float* d_pool_b_exchange, float* d_bound_exchange, float* d_pool_b_fraction, float* d_bound_fraction, float* durations, std::int32_t* rows, float* table, std::int64_t voxel_count, std::int64_t BLOCK, std::int64_t narrow) { + auto row = bsk::ld((rows + bsk::program_id(0))); + auto atom = ((bsk::program_id(1) * BLOCK) + bsk::arange_x()); + auto live = (atom < voxel_count); + auto dt = bsk::ld((durations + row)); + auto nil = (0.0f * dt); + auto value_t1 = bsk::ld((t1 + atom), live, 1.0f); + auto value_t1b = bsk::ld((t1_pool_b + atom), live, 1.0f); + auto value_t1c = bsk::ld((t1_bound + atom), live, 1.0f); + auto t0_ = _three_pool_step_jvp(bsk::truediv(1000.0f, value_t1), bsk::truediv((-1000.0f * bsk::ld((d_t1 + atom), live, 0.0f)), (value_t1 * value_t1)), bsk::truediv(1000.0f, value_t1b), bsk::truediv((-1000.0f * bsk::ld((d_t1_pool_b + atom), live, 0.0f)), (value_t1b * value_t1b)), bsk::truediv(1000.0f, value_t1c), bsk::truediv((-1000.0f * bsk::ld((d_t1_bound + atom), live, 0.0f)), (value_t1c * value_t1c)), bsk::ld((pool_b_exchange + atom), live, 0.0f), bsk::ld((d_pool_b_exchange + atom), live, 0.0f), bsk::ld((bound_exchange + atom), live, 0.0f), bsk::ld((d_bound_exchange + atom), live, 0.0f), bsk::ld((pool_b_fraction + atom), live, 0.0f), bsk::ld((d_pool_b_fraction + atom), live, 0.0f), bsk::ld((bound_fraction + atom), live, 0.0f), bsk::ld((d_bound_fraction + atom), live, 0.0f), dt, nil, (1.0f + nil), nil, narrow); + auto e00 = bsk::get<0>(t0_); + auto e01 = bsk::get<1>(t0_); + auto e02 = bsk::get<2>(t0_); + auto e10 = bsk::get<3>(t0_); + auto e11 = bsk::get<4>(t0_); + auto e12 = bsk::get<5>(t0_); + auto e20 = bsk::get<6>(t0_); + auto e21 = bsk::get<7>(t0_); + auto e22 = bsk::get<8>(t0_); + auto d00 = bsk::get<12>(t0_); + auto d01 = bsk::get<13>(t0_); + auto d02 = bsk::get<14>(t0_); + auto d10 = bsk::get<15>(t0_); + auto d11 = bsk::get<16>(t0_); + auto d12 = bsk::get<17>(t0_); + auto d20 = bsk::get<18>(t0_); + auto d21 = bsk::get<19>(t0_); + auto d22 = bsk::get<20>(t0_); + auto base = ((table + (row * (18 * voxel_count))) + atom); + bsk::st((base + (0 * voxel_count)), e00, live); + bsk::st((base + (1 * voxel_count)), e01, live); + bsk::st((base + (2 * voxel_count)), e02, live); + bsk::st((base + (3 * voxel_count)), e10, live); + bsk::st((base + (4 * voxel_count)), e11, live); + bsk::st((base + (5 * voxel_count)), e12, live); + bsk::st((base + (6 * voxel_count)), e20, live); + bsk::st((base + (7 * voxel_count)), e21, live); + bsk::st((base + (8 * voxel_count)), e22, live); + bsk::st((base + (9 * voxel_count)), d00, live); + bsk::st((base + (10 * voxel_count)), d01, live); + bsk::st((base + (11 * voxel_count)), d02, live); + bsk::st((base + (12 * voxel_count)), d10, live); + bsk::st((base + (13 * voxel_count)), d11, live); + bsk::st((base + (14 * voxel_count)), d12, live); + bsk::st((base + (15 * voxel_count)), d20, live); + bsk::st((base + (16 * voxel_count)), d21, live); + bsk::st((base + (17 * voxel_count)), d22, live); +} + +template +BSK_HD auto _shift_real(const T0& plus, const T1& minus, const T2& state, const T3& state_mask, const T4& state_count) { + auto shifted_plus = bsk::where(bsk::band((state > 0), state_mask), _up(plus, state), 0.0f); + auto shifted_minus = bsk::where(bsk::band(((state + 1) < state_count), state_mask), _down(minus, state), 0.0f); + return bsk::make_tup(bsk::where((state == 0), (-shifted_minus), shifted_plus), shifted_minus); +} + +// Transpose of ``_shift_real``. +// +// The ``a0 = -b0`` coupling sends the incoming plus adjoint back onto minus, +// at the index the minus shift moves it to. +template +BSK_HD auto _shift_real_adjoint(const T0& plus_bar, const T1& minus_bar, const T2& state, const T3& state_mask, const T4& state_count) { + bsk::tile_t | 0, 3)> shifted_minus{}; + auto carry = (-bsk::where(state_mask, _first(plus_bar, state), 0.0f)); + auto shifted_plus = bsk::where(bsk::band(((state + 1) < state_count), state_mask), _down(plus_bar, state), 0.0f); + shifted_minus = bsk::where(bsk::band((state > 0), state_mask), _up(minus_bar, state), 0.0f); + shifted_minus = bsk::where((state == 1), (shifted_minus + carry), shifted_minus); + return bsk::make_tup(shifted_plus, shifted_minus); +} + +BSK_HD void _epg_real_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_duration, float* trajectory_value, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V adjoint_mv{}; + bsk::V adjoint_pv{}; + bsk::V alpha_bar_terms_value{}; + bsk::V alpha_value{}; + bsk::V atom_b1{}; + bsk::V atom_damping{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V bare1_value{}; + bsk::V bare2_value{}; + bsk::V chs_value{}; + bsk::V cosine_value{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bool do_shift{}; + bsk::V dt_value{}; + bsk::V duration_gain_value{}; + bsk::V e1_value{}; + bsk::V e2_value{}; + std::int64_t event{}; + std::int32_t event_action{}; + bsk::V event_flip{}; + std::int32_t event_kind{}; + bsk::V grad_b1_value{}; + bsk::V grad_damping_value{}; + bsk::V grad_e1_value{}; + bsk::V grad_inversion_value{}; + bsk::V grad_m0_value{}; + bsk::V grad_t1_value{}; + bsk::V grad_t2_value{}; + bsk::V half_sine_value{}; + bool invert{}; + bool is_inversion{}; + bool is_rf{}; + bsk::V long_bar_value{}; + bsk::V long_value{}; + bsk::V minus_bar_value{}; + bsk::V minus_value{}; + bsk::V plus_bar_value{}; + bsk::V plus_value{}; + bool pre_shift{}; + bsk::V problem{}; + bsk::V pulse_b1{}; + bsk::V recovery_value{}; + bool rotate{}; + bsk::V rotated_mbv{}; + bsk::V rotated_mv{}; + bsk::V rotated_pbv{}; + bsk::V rotated_pv{}; + bsk::V rotated_zbv{}; + bsk::V rotated_zv{}; + bsk::V row_m_value{}; + bsk::V row_p_value{}; + bsk::V row_z_value{}; + bsk::V shifted_mv{}; + bsk::V shifted_pv{}; + std::int64_t shim_row{}; + bsk::V shs_value{}; + bsk::V sine_value{}; + bsk::V slot{}; + bool spoil{}; + bsk::V spread_value{}; + bsk::V stage_mv{}; + bsk::V stage_pv{}; + problem = (problem_base + (bsk::program_id(0) * problems)); + problem = (problem + bsk::arange_y()); + auto state = bsk::arange_x(); + // The grid rounds up to whole tiles, so the last program of a wave reaches + // past it. Those problems are real, but their trajectory rows belong to a + // later launch and do not exist yet. + auto active_atom = (problem < problem_end); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + // The trajectory holds the state entering every event: three planes of + // configuration orders. + auto record_stride = (3 * state_count); + auto trajectory = ((((problem - problem_base) * event_count) * record_stride) + state); + auto minus_plane = state_count; + auto long_plane = (2 * state_count); + auto empty = bsk::full(0); + plus_value = empty; + minus_value = empty; + long_value = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + atom_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + } + auto order = bsk::cast(state); + auto longitudinal_weight = (order * order); + auto transverse_weight = ((longitudinal_weight + order) + 0.3333333333333333f); + auto rate1_value = bsk::truediv(1000.0f, atom_t1); + auto rate2_value = bsk::truediv(1000.0f, atom_t2); + auto event_base = (train * event_count); + for (std::int64_t event = 0; event < event_count; event += 1) { + slot = (trajectory + (event * record_stride)); + bsk::st((trajectory_value + slot), plus_value, state_mask); + bsk::st(((trajectory_value + slot) + minus_plane), minus_value, state_mask); + bsk::st(((trajectory_value + slot) + long_plane), long_value, state_mask); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + e1_value = bsk::exp(((-rate1_value) * dt_value)); + e2_value = bsk::exp(((-rate2_value) * dt_value)); + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping(atom_damping, dt_value, order); + damp_z = bsk::get<0>(t0_); + damp_t = bsk::get<1>(t0_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + recovery_value = (1.0f - e1_value); + bare1_value = e1_value; + bare2_value = e2_value; + e1_value = (bare1_value * damp_z); + e2_value = (bare2_value * damp_t); + plus_value = (plus_value * e2_value); + minus_value = (minus_value * e2_value); + long_value = (long_value * e1_value); + long_value = (long_value + bsk::where((state == 0), recovery_value, 0.0f)); + event_action = bsk::cast(bsk::ld((action + event))); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t1_ = _shift_real(plus_value, minus_value, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t1_); + shifted_mv = bsk::get<1>(t1_); + plus_value = bsk::where(pre_shift, shifted_pv, plus_value); + minus_value = bsk::where(pre_shift, shifted_mv, minus_value); + event_kind = bsk::ld((kind + event)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + long_value = bsk::where(invert, ((-atom_inversion) * long_value), long_value); + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse the row of the shim it drives. + if (bsk::truth(shimmed)) { + shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + } + alpha_value = (event_flip * pulse_b1); + cosine_value = bsk::cos(alpha_value); + sine_value = bsk::sin(alpha_value); + chs_value = (0.5f * (1.0f + cosine_value)); + shs_value = (0.5f * (1.0f - cosine_value)); + half_sine_value = (0.5f * sine_value); + rotated_pv = ((chs_value * plus_value) + (shs_value * minus_value)); + rotated_pv = (rotated_pv - (sine_value * long_value)); + rotated_mv = ((shs_value * plus_value) + (chs_value * minus_value)); + rotated_mv = (rotated_mv + (sine_value * long_value)); + rotated_zv = ((half_sine_value * plus_value) - (half_sine_value * minus_value)); + rotated_zv = (rotated_zv + (cosine_value * long_value)); + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + plus_value = bsk::where(rotate, rotated_pv, plus_value); + minus_value = bsk::where(rotate, rotated_mv, minus_value); + long_value = bsk::where(rotate, rotated_zv, long_value); + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + auto t2_ = _shift_real(plus_value, minus_value, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t2_); + shifted_mv = bsk::get<1>(t2_); + plus_value = bsk::where(do_shift, shifted_pv, plus_value); + minus_value = bsk::where(do_shift, shifted_mv, minus_value); + spoil = (bsk::band(event_action, 8) != 0); + plus_value = bsk::where(spoil, 0.0f, plus_value); + minus_value = bsk::where(spoil, 0.0f, minus_value); + } + plus_bar_value = empty; + minus_bar_value = empty; + long_bar_value = empty; + auto zero = bsk::full(0); + grad_t1_value = zero; + grad_t2_value = zero; + grad_m0_value = zero; + grad_b1_value = zero; + grad_inversion_value = zero; + grad_damping_value = zero; + for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + event = ((event_count - 1) - reverse); + slot = (trajectory + (event * record_stride)); + auto entry_pv = bsk::ld((trajectory_value + slot), state_mask, 0.0f); + auto entry_mv = bsk::ld(((trajectory_value + slot) + minus_plane), state_mask, 0.0f); + auto entry_zv = bsk::ld(((trajectory_value + slot) + long_plane), state_mask, 0.0f); + event_action = bsk::cast(bsk::ld((action + event))); + event_kind = bsk::ld((kind + event)); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + e1_value = bsk::exp(((-rate1_value) * dt_value)); + e2_value = bsk::exp(((-rate2_value) * dt_value)); + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t3_ = _damping(atom_damping, dt_value, order); + damp_z = bsk::get<0>(t3_); + damp_t = bsk::get<1>(t3_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + recovery_value = (1.0f - e1_value); + bare1_value = e1_value; + bare2_value = e2_value; + e1_value = (bare1_value * damp_z); + e2_value = (bare2_value * damp_t); + // Replay the intra-event stages from the recorded entry state. + stage_pv = (entry_pv * e2_value); + stage_mv = (entry_mv * e2_value); + auto stage_zv = ((entry_zv * e1_value) + bsk::where((state == 0), recovery_value, 0.0f)); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t4_ = _shift_real(stage_pv, stage_mv, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t4_); + shifted_mv = bsk::get<1>(t4_); + stage_pv = bsk::where(pre_shift, shifted_pv, stage_pv); + stage_mv = bsk::where(pre_shift, shifted_mv, stage_mv); + // Undo the trailing spoil or shift. + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + spoil = (bsk::band(event_action, 8) != 0); + auto t5_ = _shift_real_adjoint(plus_bar_value, minus_bar_value, state, state_mask, state_count); + adjoint_pv = bsk::get<0>(t5_); + adjoint_mv = bsk::get<1>(t5_); + auto trailing = bsk::band(do_shift, bsk::bnot(spoil)); + plus_bar_value = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_pv, plus_bar_value)); + minus_bar_value = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_mv, minus_bar_value)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + grad_inversion_value = (grad_inversion_value + (-bsk::sum_x(bsk::where(invert, (long_bar_value * stage_zv), 0.0f)))); + long_bar_value = bsk::where(invert, ((-atom_inversion) * long_bar_value), long_bar_value); + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse the row of the shim it drives. + if (bsk::truth(shimmed)) { + shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + } + alpha_value = (event_flip * pulse_b1); + cosine_value = bsk::cos(alpha_value); + sine_value = bsk::sin(alpha_value); + chs_value = (0.5f * (1.0f + cosine_value)); + shs_value = (0.5f * (1.0f - cosine_value)); + half_sine_value = (0.5f * sine_value); + // d/dalpha of each output row, contracted with the adjoint. + row_p_value = ((half_sine_value * stage_mv) - (half_sine_value * stage_pv)); + row_p_value = (row_p_value - (cosine_value * stage_zv)); + row_m_value = ((half_sine_value * stage_pv) - (half_sine_value * stage_mv)); + row_m_value = (row_m_value + (cosine_value * stage_zv)); + row_z_value = (((0.5f * cosine_value) * stage_pv) - ((0.5f * cosine_value) * stage_mv)); + row_z_value = (row_z_value - (sine_value * stage_zv)); + alpha_bar_terms_value = (plus_bar_value * row_p_value); + alpha_bar_terms_value = (alpha_bar_terms_value + (minus_bar_value * row_m_value)); + alpha_bar_terms_value = (alpha_bar_terms_value + (long_bar_value * row_z_value)); + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + auto grad_alpha_value = bsk::sum_x(bsk::where(rotate, alpha_bar_terms_value, 0.0f)); + // Transpose of the rotation. + rotated_pbv = ((chs_value * plus_bar_value) + (shs_value * minus_bar_value)); + rotated_pbv = (rotated_pbv + (half_sine_value * long_bar_value)); + rotated_mbv = ((shs_value * plus_bar_value) + (chs_value * minus_bar_value)); + rotated_mbv = (rotated_mbv - (half_sine_value * long_bar_value)); + rotated_zbv = (((-sine_value) * plus_bar_value) + (sine_value * minus_bar_value)); + rotated_zbv = (rotated_zbv + (cosine_value * long_bar_value)); + plus_bar_value = bsk::where(rotate, rotated_pbv, plus_bar_value); + minus_bar_value = bsk::where(rotate, rotated_mbv, minus_bar_value); + long_bar_value = bsk::where(rotate, rotated_zbv, long_bar_value); + auto writes_flip = bsk::band(active_atom, rotate); + bsk::atomic_add(((grad_flip + event_base) + event), (grad_alpha_value * pulse_b1), writes_flip); + if (bsk::truth(shimmed)) { + // A pulse's transmit gradient belongs to the shim it drives, so + // with several it lands in that shim's row rather than in a + // register summed over the whole train. + bsk::atomic_add((((grad_tissue + (3 * atom_count)) + shim_row) + atom), (grad_alpha_value * event_flip), writes_flip); + } else { + grad_b1_value = (grad_b1_value + bsk::where(rotate, (grad_alpha_value * event_flip), 0.0f)); + } + // The sample is i * m0 * plus[0]; only the imaginary seed acts. + auto record = bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)); + auto out_ = bsk::ld((output_index + event)); + auto seed = bsk::ld(((grad_output_imag + (problem * output_count)) + out_), bsk::band(bsk::band(active_atom, record), (out_ >= 0)), 0.0f); + grad_m0_value = (grad_m0_value + bsk::sum_x(bsk::where((state == 0), (seed * stage_pv), 0.0f))); + plus_bar_value = (plus_bar_value + bsk::where((state == 0), (seed * atom_m0), 0.0f)); + auto t6_ = _shift_real_adjoint(plus_bar_value, minus_bar_value, state, state_mask, state_count); + adjoint_pv = bsk::get<0>(t6_); + adjoint_mv = bsk::get<1>(t6_); + plus_bar_value = bsk::where(pre_shift, adjoint_pv, plus_bar_value); + minus_bar_value = bsk::where(pre_shift, adjoint_mv, minus_bar_value); + auto cot2_value = ((plus_bar_value * entry_pv) + (minus_bar_value * entry_mv)); + auto cot1_value = (long_bar_value * entry_zv); + auto grad_e2_value = bsk::sum_x((cot2_value * damp_t)); + grad_e1_value = bsk::sum_x((cot1_value * damp_z)); + grad_e1_value = (grad_e1_value - bsk::sum_x(bsk::where((state == 0), long_bar_value, 0.0f))); + // The rate and the interval multiply every order's b-weight, so both + // take a weighted sum. Order zero has no longitudinal weight, which + // keeps recovery out of this. + spread_value = zero; + if (bsk::truth(diffusing)) { + auto weighted_value = ((((cot1_value * bare1_value) * damp_z) * longitudinal_weight) + (((cot2_value * bare2_value) * damp_t) * transverse_weight)); + spread_value = bsk::sum_x(weighted_value); + grad_damping_value = (grad_damping_value + ((-spread_value) * dt_value)); + } + plus_bar_value = (plus_bar_value * e2_value); + minus_bar_value = (minus_bar_value * e2_value); + long_bar_value = (long_bar_value * e1_value); + auto inverse1_value = bsk::truediv(1000.0f, (atom_t1 * atom_t1)); + auto inverse2_value = bsk::truediv(1000.0f, (atom_t2 * atom_t2)); + grad_t1_value = (grad_t1_value + (grad_e1_value * ((bare1_value * dt_value) * inverse1_value))); + grad_t2_value = (grad_t2_value + (grad_e2_value * ((bare2_value * dt_value) * inverse2_value))); + duration_gain_value = ((-grad_e1_value) * (rate1_value * bare1_value)); + duration_gain_value = (duration_gain_value - (grad_e2_value * (rate2_value * bare2_value))); + duration_gain_value = (duration_gain_value + ((-spread_value) * atom_damping)); + bsk::atomic_add(((grad_duration + event_base) + event), duration_gain_value, active_atom); + } + bsk::atomic_add((grad_tissue + atom), grad_t1_value, active_atom); + bsk::atomic_add(((grad_tissue + atom_count) + atom), grad_t2_value, active_atom); + bsk::atomic_add(((grad_tissue + (2 * atom_count)) + atom), grad_m0_value, active_atom); + if (bsk::truth((!bsk::truth(shimmed)))) { + bsk::atomic_add(((grad_tissue + (3 * atom_count)) + atom), grad_b1_value, active_atom); + } + // The transmit pair takes a row per shim each in the plane the complex + // path allocates, so the rows past it move even though this kernel leaves + // the transmit phase at zero throughout. + auto past_transmit = (2 * (shim_rows - 1)); + bsk::atomic_add(((grad_tissue + ((6 + past_transmit) * atom_count)) + atom), grad_inversion_value, active_atom); + bsk::atomic_add(((grad_tissue + ((7 + past_transmit) * atom_count)) + atom), grad_damping_value, active_atom); +} + +// The pair and its derivative in the flip angle, from the same cubic. +// +// The derivative of a Hermite segment is another polynomial in the same four +// knot values, so reading both costs one extra combination rather than a +// second table. Returned interleaved: each component's value then its slope, +// in the order ``a`` real, ``a`` imaginary, ``b`` real, ``b`` imaginary. +template +BSK_HD auto _profile_pair_slope(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + // d/dtheta is d/du over the knot spacing. + auto g00 = bsk::truediv(((6.0f * u2) - (6.0f * u)), step); + auto g10 = (((3.0f * u2) - (4.0f * u)) + 1.0f); + auto g01 = bsk::truediv(((6.0f * u) - (6.0f * u2)), step); + auto g11 = ((3.0f * u2) - (2.0f * u)); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto near = [&](int c) { return bsk::ld(((profile + base) + c)); }; + auto near_slope = [&](int c) { return bsk::ld((((profile + base) + 4) + c)); }; + auto far = [&](int c) { return bsk::ld((((profile + base) + 8) + c)); }; + auto far_slope = [&](int c) { return bsk::ld((((profile + base) + 12) + c)); }; + auto value = [&](int c) { + return ((((h00 * near(c)) + (h10 * near_slope(c))) + (h01 * far(c))) + (h11 * far_slope(c))); + }; + auto slope = [&](int c) { + return ((((g00 * near(c)) + (g10 * near_slope(c))) + (g01 * far(c))) + (g11 * far_slope(c))); + }; + return bsk::make_tup(value(0), slope(0), value(1), slope(1), value(2), slope(2), value(3), slope(3)); +} + +// One pool through a hard pulse, carried alongside a tangent. +// +// Pulled out of the kernel body so a second pool can take the same +// rotation: a chemical shift moves where a pool precesses, not what a +// pulse does to it. +template +BSK_HD auto _rotate_flip_phase_jvp(const T0& cosine, const T1& dcosine, const T2& sine, const T3& dsine, const T4& cos_phi, const T5& dcos_phi, const T6& sin_phi, const T7& dsin_phi, const T8& cos_2phi, const T9& dcos_2phi, const T10& sin_2phi, const T11& dsin_2phi, const T12& fpr, const T13& fpi, const T14& fmr, const T15& fmi, const T16& zr, const T17& zi, const T18& dfpr, const T19& dfpi, const T20& dfmr, const T21& dfmi, const T22& dzr, const T23& dzi) { + bsk::tile_t | 0, 3)> dmi_a{}; + bsk::tile_t | 0, 3)> dmr_a{}; + bsk::tile_t | 0, 3)> dpi_a{}; + bsk::tile_t | 0, 3)> dpr_a{}; + bsk::tile_t | 0, 3)> rotated_dmi{}; + bsk::tile_t | 0, 3)> rotated_dmr{}; + bsk::tile_t | 0, 3)> rotated_dpi{}; + bsk::tile_t | 0, 3)> rotated_dpr{}; + bsk::tile_t | 0, 3)> rotated_dzi{}; + bsk::tile_t | 0, 3)> rotated_dzr{}; + auto ch = (0.5f * (1.0f + cosine)); + auto sh = (0.5f * (1.0f - cosine)); + auto dch = (0.5f * dcosine); + auto dsh = (-0.5f * dcosine); + auto pr_a = ((cos_2phi * fmr) - (sin_2phi * fmi)); + dpr_a = ((dcos_2phi * fmr) + (cos_2phi * dfmr)); + dpr_a = (dpr_a - ((dsin_2phi * fmi) + (sin_2phi * dfmi))); + auto pr_b = ((sin_phi * zr) + (cos_phi * zi)); + auto dpr_b = ((((dsin_phi * zr) + (sin_phi * dzr)) + (dcos_phi * zi)) + (cos_phi * dzi)); + auto rotated_pr = (((ch * fpr) + (sh * pr_a)) + (sine * pr_b)); + rotated_dpr = ((((dch * fpr) + (ch * dfpr)) + (dsh * pr_a)) + (sh * dpr_a)); + rotated_dpr = (rotated_dpr + ((dsine * pr_b) + (sine * dpr_b))); + auto pi_a = ((sin_2phi * fmr) + (cos_2phi * fmi)); + dpi_a = ((dsin_2phi * fmr) + (sin_2phi * dfmr)); + dpi_a = (dpi_a + ((dcos_2phi * fmi) + (cos_2phi * dfmi))); + auto pi_b = ((sin_phi * zi) - (cos_phi * zr)); + auto dpi_b = ((((dsin_phi * zi) + (sin_phi * dzi)) - (dcos_phi * zr)) - (cos_phi * dzr)); + auto rotated_pi = (((ch * fpi) + (sh * pi_a)) + (sine * pi_b)); + rotated_dpi = ((((dch * fpi) + (ch * dfpi)) + (dsh * pi_a)) + (sh * dpi_a)); + rotated_dpi = (rotated_dpi + ((dsine * pi_b) + (sine * dpi_b))); + auto mr_a = ((cos_2phi * fpr) + (sin_2phi * fpi)); + dmr_a = ((dcos_2phi * fpr) + (cos_2phi * dfpr)); + dmr_a = (dmr_a + ((dsin_2phi * fpi) + (sin_2phi * dfpi))); + auto mr_b = ((sin_phi * zr) - (cos_phi * zi)); + auto dmr_b = ((((dsin_phi * zr) + (sin_phi * dzr)) - (dcos_phi * zi)) - (cos_phi * dzi)); + auto rotated_mr = (((sh * mr_a) + (ch * fmr)) + (sine * mr_b)); + rotated_dmr = ((((dsh * mr_a) + (sh * dmr_a)) + (dch * fmr)) + (ch * dfmr)); + rotated_dmr = (rotated_dmr + ((dsine * mr_b) + (sine * dmr_b))); + auto mi_a = (((-sin_2phi) * fpr) + (cos_2phi * fpi)); + dmi_a = (((-dsin_2phi) * fpr) - (sin_2phi * dfpr)); + dmi_a = (dmi_a + ((dcos_2phi * fpi) + (cos_2phi * dfpi))); + auto mi_b = ((cos_phi * zr) + (sin_phi * zi)); + auto dmi_b = ((((dcos_phi * zr) + (cos_phi * dzr)) + (dsin_phi * zi)) + (sin_phi * dzi)); + auto rotated_mi = (((sh * mi_a) + (ch * fmi)) + (sine * mi_b)); + rotated_dmi = ((((dsh * mi_a) + (sh * dmi_a)) + (dch * fmi)) + (ch * dfmi)); + rotated_dmi = (rotated_dmi + ((dsine * mi_b) + (sine * dmi_b))); + auto zr_a = ((sin_phi * fpr) - (cos_phi * fpi)); + auto dzr_a = ((((dsin_phi * fpr) + (sin_phi * dfpr)) - (dcos_phi * fpi)) - (cos_phi * dfpi)); + auto zr_b = ((sin_phi * fmr) + (cos_phi * fmi)); + auto dzr_b = ((((dsin_phi * fmr) + (sin_phi * dfmr)) + (dcos_phi * fmi)) + (cos_phi * dfmi)); + auto rotated_zr = ((((-0.5f * sine) * zr_a) - ((0.5f * sine) * zr_b)) + (cosine * zr)); + rotated_dzr = (-0.5f * ((dsine * zr_a) + (sine * dzr_a))); + rotated_dzr = (rotated_dzr - (0.5f * ((dsine * zr_b) + (sine * dzr_b)))); + rotated_dzr = (rotated_dzr + ((dcosine * zr) + (cosine * dzr))); + auto zi_a = ((cos_phi * fpr) + (sin_phi * fpi)); + auto dzi_a = ((((dcos_phi * fpr) + (cos_phi * dfpr)) + (dsin_phi * fpi)) + (sin_phi * dfpi)); + auto zi_b = ((cos_phi * fmr) - (sin_phi * fmi)); + auto dzi_b = ((((dcos_phi * fmr) + (cos_phi * dfmr)) - (dsin_phi * fmi)) - (sin_phi * dfmi)); + auto rotated_zi = ((((-0.5f * sine) * zi_a) + ((0.5f * sine) * zi_b)) + (cosine * zi)); + rotated_dzi = (-0.5f * ((dsine * zi_a) + (sine * dzi_a))); + rotated_dzi = (rotated_dzi + (0.5f * ((dsine * zi_b) + (sine * dzi_b)))); + rotated_dzi = (rotated_dzi + ((dcosine * zi) + (cosine * dzi))); + return bsk::make_tup(rotated_pr, rotated_pi, rotated_mr, rotated_mi, rotated_zr, rotated_zi, rotated_dpr, rotated_dpi, rotated_dmr, rotated_dmi, rotated_dzr, rotated_dzi); +} + +BSK_HD void _epg_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_b1_phase, float* tangent_b0, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_velocity, float* tangent_bound_fraction, float* tangent_exchange_rate, float* tangent_t1_bound, float* tangent_pool_b_fraction, float* tangent_pool_b_exchange, float* tangent_t1_pool_b, float* tangent_t2_pool_b, float* tangent_pool_b_shift, float* tangent_duration, float* tangent_flip, float* tangent_phase, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V atom_b0{}; + bsk::V atom_b1{}; + bsk::V atom_b1_phase{}; + bsk::V atom_bound{}; + bsk::V atom_damping{}; + bsk::V atom_exchange{}; + bsk::V atom_flow{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V atom_r1_bound{}; + bsk::V atom_r1_semisolid{}; + bsk::V atom_r2_bound{}; + bsk::V atom_semisolid{}; + bsk::V atom_semisolid_exchange{}; + bsk::V atom_shift{}; + bsk::V atom_washout{}; + bsk::V b_dmi{}; + bsk::V b_dmr{}; + bsk::V b_dpi{}; + bsk::V b_dpr{}; + bsk::V b_dzi{}; + bsk::V b_dzr{}; + bsk::V b_mi{}; + bsk::V b_mr{}; + bsk::V b_pi{}; + bsk::V b_pr{}; + bsk::V b_zi{}; + bsk::V b_zr{}; + bsk::V bi{}; + bsk::V bmi{}; + bsk::V bmr{}; + bsk::V bpi{}; + bsk::V bpr{}; + bsk::V br{}; + bsk::V ci{}; + bsk::V cr{}; + bsk::V d_bound{}; + bsk::V d_damping{}; + bsk::V d_exchange{}; + bsk::V d_flow{}; + bsk::V d_free_i{}; + bsk::V d_free_r{}; + bsk::V d_grow_free{}; + bsk::V d_grow_pool_b{}; + bsk::V d_grow_semisolid{}; + bsk::V d_held_i{}; + bsk::V d_held_r{}; + bsk::V d_r1_bound{}; + bsk::V d_r1_semisolid{}; + bsk::V d_r2_bound{}; + bsk::V d_semisolid{}; + bsk::V d_semisolid_exchange{}; + bsk::V d_shift{}; + bsk::V d_t11{}; + bsk::V d_t12{}; + bsk::V d_t13{}; + bsk::V d_t21{}; + bsk::V d_t22{}; + bsk::V d_t23{}; + bsk::V d_t31{}; + bsk::V d_t32{}; + bsk::V d_t33{}; + bsk::V d_washout{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bsk::V db0{}; + bsk::V db1{}; + bsk::V db1_phase{}; + bsk::V dbi{}; + bsk::V dbmi{}; + bsk::V dbmr{}; + bsk::V dbpi{}; + bsk::V dbpr{}; + bsk::V dbr{}; + bsk::V dci{}; + bsk::V dcr{}; + bsk::V ddamp_t{}; + bsk::V ddamp_z{}; + bsk::V de1{}; + bsk::V de2{}; + bsk::V dfmi{}; + bsk::V dfmr{}; + bsk::V dfpi{}; + bsk::V dfpr{}; + bsk::V dinversion{}; + bsk::V dm0{}; + bsk::V doff_cos{}; + bsk::V doff_sin{}; + bsk::V dot_ai{}; + bsk::V dot_ar{}; + bsk::V dot_bi{}; + bsk::V dot_br{}; + bsk::V dread_i{}; + bsk::V dread_r{}; + bsk::V dspun_hi{}; + bsk::V dspun_hr{}; + bsk::V dturn_cos{}; + bsk::V dturn_sin{}; + bsk::V dturn_t{}; + bsk::V dturn_z{}; + bsk::V dwout{}; + bsk::V dzi{}; + bsk::V dzr{}; + bsk::V e1{}; + bsk::V e2{}; + bsk::V fmi{}; + bsk::V fmr{}; + bsk::V fpi{}; + bsk::V fpr{}; + bsk::V free_i{}; + bsk::V free_r{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; + bsk::V held_dmi{}; + bsk::V held_dmr{}; + bsk::V held_dpi{}; + bsk::V held_dpr{}; + bsk::V held_dzi{}; + bsk::V held_dzr{}; + bsk::V held_i{}; + bsk::V held_mi{}; + bsk::V held_mr{}; + bsk::V held_pi{}; + bsk::V held_pr{}; + bsk::V held_r{}; + bsk::V held_t1{}; + bsk::V held_zi{}; + bsk::V held_zr{}; + bsk::V off_cos{}; + bsk::V off_sin{}; + bsk::V pair_ai{}; + bsk::V pair_ar{}; + bsk::V pair_bi{}; + bsk::V pair_br{}; + bsk::V read_i{}; + bsk::V read_r{}; + bsk::V rotated_dmi{}; + bsk::V rotated_dmr{}; + bsk::V rotated_dpi{}; + bsk::V rotated_dpr{}; + bsk::V rotated_dzi{}; + bsk::V rotated_dzr{}; + bsk::V rotated_mi{}; + bsk::V rotated_mr{}; + bsk::V rotated_pi{}; + bsk::V rotated_pr{}; + bsk::V rotated_zi{}; + bsk::V rotated_zr{}; + bsk::V s_bmi{}; + bsk::V s_bmr{}; + bsk::V s_bpi{}; + bsk::V s_bpr{}; + bsk::V s_dbmi{}; + bsk::V s_dbmr{}; + bsk::V s_dbpi{}; + bsk::V s_dbpr{}; + bsk::V shaped_dmi{}; + bsk::V shaped_dmr{}; + bsk::V shaped_dpi{}; + bsk::V shaped_dpr{}; + bsk::V shaped_dzi{}; + bsk::V shaped_dzr{}; + bsk::V shaped_mi{}; + bsk::V shaped_mr{}; + bsk::V shaped_pi{}; + bsk::V shaped_pr{}; + bsk::V shaped_zi{}; + bsk::V shaped_zr{}; + bsk::V shifted_dmi{}; + bsk::V shifted_dmr{}; + bsk::V shifted_dpi{}; + bsk::V shifted_dpr{}; + bsk::V shifted_mi{}; + bsk::V shifted_mr{}; + bsk::V shifted_pi{}; + bsk::V shifted_pr{}; + bsk::V signal_imag{}; + bsk::V signal_real{}; + bsk::V spun_hi{}; + bsk::V spun_hr{}; + bsk::V t11{}; + bsk::V t12{}; + bsk::V t13{}; + bsk::V t21{}; + bsk::V t22{}; + bsk::V t23{}; + bsk::V t31{}; + bsk::V t32{}; + bsk::V t33{}; + bsk::V turn_cos{}; + bsk::V turn_sin{}; + bsk::V turn_t{}; + bsk::V turn_z{}; + bsk::V wout{}; + bsk::V zi{}; + bsk::V zr{}; + auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < (train_count * atom_count)); + // A partial block carries lanes with no problem behind them, and they must + // take no part in a reduction or a store. + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + // Voxels are spread over the slice voxel-major, so a voxel's place along + // the slice is its index modulo the profile's width. One pulse shape holds + // that many consecutive rows, and the event says which shape it drives. + auto location = bsk::mod(atom, locations); + auto empty = bsk::full(0); + fpr = empty; + fpi = empty; + fmr = empty; + fmi = empty; + // Equilibrium is split between the pools, so a direction along the bound + // fraction moves magnetization from one to the other before a single event + // has run. + atom_bound = 0.0f; + d_bound = 0.0f; + atom_exchange = 0.0f; + d_exchange = 0.0f; + atom_r1_bound = 0.0f; + d_r1_bound = 0.0f; + atom_r2_bound = 0.0f; + d_r2_bound = 0.0f; + atom_shift = 0.0f; + d_shift = 0.0f; + atom_semisolid = 0.0f; + d_semisolid = 0.0f; + atom_semisolid_exchange = 0.0f; + d_semisolid_exchange = 0.0f; + atom_r1_semisolid = 0.0f; + d_r1_semisolid = 0.0f; + if (bsk::truth((pools == 1))) { + atom_bound = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + d_bound = bsk::ld((tangent_bound_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + d_exchange = bsk::ld((tangent_exchange_rate + scalar_atom), active_atom, 0.0f); + held_t1 = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + atom_r1_bound = bsk::truediv(1000.0f, held_t1); + d_r1_bound = bsk::truediv((-1000.0f * bsk::ld((tangent_t1_bound + scalar_atom), active_atom, 0.0f)), (held_t1 * held_t1)); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + atom_bound = bsk::ld((pool_b_fraction + scalar_atom), active_atom, 0.0f); + d_bound = bsk::ld((tangent_pool_b_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((pool_b_exchange + scalar_atom), active_atom, 0.0f); + d_exchange = bsk::ld((tangent_pool_b_exchange + scalar_atom), active_atom, 0.0f); + held_t1 = bsk::ld((t1_pool_b + scalar_atom), active_atom, 1.0f); + atom_r1_bound = bsk::truediv(1000.0f, held_t1); + d_r1_bound = bsk::truediv((-1000.0f * bsk::ld((tangent_t1_pool_b + scalar_atom), active_atom, 0.0f)), (held_t1 * held_t1)); + auto held_t2 = bsk::ld((t2_pool_b + scalar_atom), active_atom, 1.0f); + atom_r2_bound = bsk::truediv(1000.0f, held_t2); + d_r2_bound = bsk::truediv((-1000.0f * bsk::ld((tangent_t2_pool_b + scalar_atom), active_atom, 0.0f)), (held_t2 * held_t2)); + atom_shift = bsk::ld((pool_b_shift + scalar_atom), active_atom, 0.0f); + d_shift = bsk::ld((tangent_pool_b_shift + scalar_atom), active_atom, 0.0f); + } + if (bsk::truth((pools == 3))) { + atom_semisolid = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + d_semisolid = bsk::ld((tangent_bound_fraction + scalar_atom), active_atom, 0.0f); + atom_semisolid_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + d_semisolid_exchange = bsk::ld((tangent_exchange_rate + scalar_atom), active_atom, 0.0f); + auto held_semisolid = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + atom_r1_semisolid = bsk::truediv(1000.0f, held_semisolid); + d_r1_semisolid = bsk::truediv((-1000.0f * bsk::ld((tangent_t1_bound + scalar_atom), active_atom, 0.0f)), (held_semisolid * held_semisolid)); + } + auto atom_free = ((1.0f - atom_bound) - atom_semisolid); + auto d_free = ((-d_bound) - d_semisolid); + zr = (empty + bsk::where((state == 0), atom_free, 0.0f)); + zi = empty; + br = (empty + bsk::where((state == 0), (atom_bound + 0.0f), 0.0f)); + bi = empty; + cr = (empty + bsk::where((state == 0), (atom_semisolid + 0.0f), 0.0f)); + ci = empty; + dcr = (empty + bsk::where((state == 0), (d_semisolid + 0.0f), 0.0f)); + dci = empty; + dfpr = empty; + dfpi = empty; + dfmr = empty; + dfmi = empty; + dzr = (empty + bsk::where((state == 0), ((-d_bound) - d_semisolid), 0.0f)); + dzi = empty; + dbr = (empty + bsk::where((state == 0), (d_bound + 0.0f), 0.0f)); + dbi = empty; + bpr = empty; + bpi = empty; + bmr = empty; + bmi = empty; + dbpr = empty; + dbpi = empty; + dbmr = empty; + dbmi = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_b1_phase = 0.0f; + atom_b0 = 0.0f; + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld((b1_phase + scalar_atom), active_atom, 0.0f); + atom_b0 = bsk::ld((b0 + scalar_atom), active_atom, 0.0f); + } + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + atom_damping = 0.0f; + d_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + d_damping = bsk::ld((tangent_diffusion + scalar_atom), active_atom, 0.0f); + } + atom_flow = 0.0f; + d_flow = 0.0f; + atom_washout = 0.0f; + d_washout = 0.0f; + if (bsk::truth(moving)) { + auto atom_velocity = bsk::ld((velocity + scalar_atom), active_atom, 0.0f); + auto d_velocity = bsk::ld((tangent_velocity + scalar_atom), active_atom, 0.0f); + atom_flow = (atom_velocity * flow_scale); + d_flow = (d_velocity * flow_scale); + // |v| has no derivative at the origin, so a still voxel contributes + // none. + auto direction = (bsk::cast((atom_velocity > 0.0f)) - bsk::cast((atom_velocity < 0.0f))); + atom_washout = (bsk::abs(atom_velocity) * washout_scale); + d_washout = ((direction * d_velocity) * washout_scale); + } + auto order = bsk::cast(state); + auto dt1 = bsk::ld((tangent_t1 + atom), active_atom, 0.0f); + auto dt2 = bsk::ld((tangent_t2 + atom), active_atom, 0.0f); + dm0 = 0.0f; + if (bsk::truth(density)) { + dm0 = bsk::ld((tangent_m0 + scalar_atom), active_atom, 0.0f); + } + db1 = 0.0f; + if (bsk::truth(transmit)) { + db1 = bsk::ld((tangent_b1 + scalar_atom), active_atom, 0.0f); + } + db1_phase = 0.0f; + db0 = 0.0f; + if (bsk::truth(off_axis)) { + db1_phase = bsk::ld((tangent_b1_phase + scalar_atom), active_atom, 0.0f); + db0 = bsk::ld((tangent_b0 + scalar_atom), active_atom, 0.0f); + } + dinversion = 0.0f; + if (bsk::truth(inverting)) { + dinversion = bsk::ld((tangent_inversion_efficiency + scalar_atom), active_atom, 0.0f); + } + auto event_base = (train * event_count); + for (std::int64_t event = 0; event < event_count; event += 1) { + auto event_dt = _event_value(duration, event_base, event, active_atom, single_train); + auto ddt = _event_value(tangent_duration, event_base, event, active_atom, single_train); + auto r1 = bsk::truediv(1000.0f, atom_t1); + auto r2 = bsk::truediv(1000.0f, atom_t2); + wout = 1.0f; + dwout = 0.0f; + if (bsk::truth(moving)) { + auto t0_ = _washout_jvp(atom_washout, d_washout, event_dt, ddt); + wout = bsk::get<0>(t0_); + dwout = bsk::get<1>(t0_); + } + auto dry1 = bsk::exp(((-r1) * event_dt)); + auto dry2 = bsk::exp(((-r2) * event_dt)); + e1 = (dry1 * wout); + e2 = (dry2 * wout); + de1 = ((e1 * (bsk::truediv(((1000.0f * event_dt) * dt1), (atom_t1 * atom_t1)) - (r1 * ddt))) + (dry1 * dwout)); + de2 = ((e2 * (bsk::truediv(((1000.0f * event_dt) * dt2), (atom_t2 * atom_t2)) - (r2 * ddt))) + (dry2 * dwout)); + damp_z = 1.0f; + ddamp_z = 0.0f; + damp_t = 1.0f; + ddamp_t = 0.0f; + if (bsk::truth(diffusing)) { + auto t1_ = _damping_jvp(atom_damping, d_damping, event_dt, ddt, order); + damp_z = bsk::get<0>(t1_); + ddamp_z = bsk::get<1>(t1_); + damp_t = bsk::get<2>(t1_); + ddamp_t = bsk::get<3>(t1_); + } + // Order zero is undamped, so the recovery term keeps the bare factor. + auto t2_ = bsk::make_tup((1.0f - e1), (-de1)); + auto recovery = bsk::get<0>(t2_); + auto drecovery = bsk::get<1>(t2_); + de1 = ((de1 * damp_z) + (e1 * ddamp_z)); + e1 = (e1 * damp_z); + de2 = ((de2 * damp_t) + (e2 * ddamp_t)); + e2 = (e2 * damp_t); + turn_z = 0.0f; + turn_t = 0.0f; + dturn_z = 0.0f; + dturn_t = 0.0f; + if (bsk::truth(moving)) { + auto t3_ = _flow(atom_flow, event_dt, order); + turn_z = bsk::get<0>(t3_); + turn_t = bsk::get<1>(t3_); + auto d_turn = ((d_flow * event_dt) + (atom_flow * ddt)); + dturn_z = ((-order) * d_turn); + dturn_t = ((-(order + 0.5f)) * d_turn); + } + off_cos = 1.0f; + off_sin = 0.0f; + doff_cos = 0.0f; + doff_sin = 0.0f; + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + // Flow winds the transverse states through the same rotation + // off-resonance does, so the two phases add before either is taken. + auto off_phase = (((-6.283185307179586f * atom_b0) * event_dt) + turn_t); + auto doff_phase = ((-6.283185307179586f * ((db0 * event_dt) + (atom_b0 * ddt))) + dturn_t); + off_cos = bsk::cos(off_phase); + off_sin = bsk::sin(off_phase); + doff_cos = ((-off_sin) * doff_phase); + doff_sin = (off_cos * doff_phase); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // Both pools take the same off-resonance and per-order damping; + // what separates them is the chemical shift, which the exchange + // operator already carries. + auto t4_ = _two_pool_transverse_step_jvp(r2, bsk::truediv((-1000.0f * dt2), (atom_t2 * atom_t2)), atom_r2_bound, d_r2_bound, atom_exchange, d_exchange, atom_bound, d_bound, atom_free, d_free, atom_shift, d_shift, event_dt, ddt, wout, dwout); + auto x11r = bsk::get<0>(t4_); + auto x11i = bsk::get<1>(t4_); + auto x12r = bsk::get<2>(t4_); + auto x12i = bsk::get<3>(t4_); + auto x21r = bsk::get<4>(t4_); + auto x21i = bsk::get<5>(t4_); + auto x22r = bsk::get<6>(t4_); + auto x22i = bsk::get<7>(t4_); + auto d11r = bsk::get<8>(t4_); + auto d11i = bsk::get<9>(t4_); + auto d12r = bsk::get<10>(t4_); + auto d12i = bsk::get<11>(t4_); + auto d21r = bsk::get<12>(t4_); + auto d21i = bsk::get<13>(t4_); + auto d22r = bsk::get<14>(t4_); + auto d22i = bsk::get<15>(t4_); + auto mix_pr = ((((x11r * fpr) - (x11i * fpi)) + (x12r * bpr)) - (x12i * bpi)); + auto mix_pi = ((((x11r * fpi) + (x11i * fpr)) + (x12r * bpi)) + (x12i * bpr)); + auto dmix_pr = ((((((((d11r * fpr) + (x11r * dfpr)) - (d11i * fpi)) - (x11i * dfpi)) + (d12r * bpr)) + (x12r * dbpr)) - (d12i * bpi)) - (x12i * dbpi)); + auto dmix_pi = ((((((((d11r * fpi) + (x11r * dfpi)) + (d11i * fpr)) + (x11i * dfpr)) + (d12r * bpi)) + (x12r * dbpi)) + (d12i * bpr)) + (x12i * dbpr)); + auto mix_br = ((((x21r * fpr) - (x21i * fpi)) + (x22r * bpr)) - (x22i * bpi)); + auto mix_bi = ((((x21r * fpi) + (x21i * fpr)) + (x22r * bpi)) + (x22i * bpr)); + auto dmix_br = ((((((((d21r * fpr) + (x21r * dfpr)) - (d21i * fpi)) - (x21i * dfpi)) + (d22r * bpr)) + (x22r * dbpr)) - (d22i * bpi)) - (x22i * dbpi)); + auto dmix_bi = ((((((((d21r * fpi) + (x21r * dfpi)) + (d21i * fpr)) + (x21i * dfpr)) + (d22r * bpi)) + (x22r * dbpi)) + (d22i * bpr)) + (x22i * dbpr)); + // ``F-`` follows the conjugate of the operator entry by entry. + auto mix_mr = ((((x11r * fmr) + (x11i * fmi)) + (x12r * bmr)) + (x12i * bmi)); + auto mix_mi = ((((x11r * fmi) - (x11i * fmr)) + (x12r * bmi)) - (x12i * bmr)); + auto dmix_mr = ((((((((d11r * fmr) + (x11r * dfmr)) + (d11i * fmi)) + (x11i * dfmi)) + (d12r * bmr)) + (x12r * dbmr)) + (d12i * bmi)) + (x12i * dbmi)); + auto dmix_mi = ((((((((d11r * fmi) + (x11r * dfmi)) - (d11i * fmr)) - (x11i * dfmr)) + (d12r * bmi)) + (x12r * dbmi)) - (d12i * bmr)) - (x12i * dbmr)); + auto mix_nr = ((((x21r * fmr) + (x21i * fmi)) + (x22r * bmr)) + (x22i * bmi)); + auto mix_ni = ((((x21r * fmi) - (x21i * fmr)) + (x22r * bmi)) - (x22i * bmr)); + auto dmix_nr = ((((((((d21r * fmr) + (x21r * dfmr)) + (d21i * fmi)) + (x21i * dfmi)) + (d22r * bmr)) + (x22r * dbmr)) + (d22i * bmi)) + (x22i * dbmi)); + auto dmix_ni = ((((((((d21r * fmi) + (x21r * dfmi)) - (d21i * fmr)) - (x21i * dfmr)) + (d22r * bmi)) + (x22r * dbmi)) - (d22i * bmr)) - (x22i * dbmr)); + // The damping and off-resonance both pools share, applied after. + auto carry_r = (damp_t * off_cos); + auto carry_i = (damp_t * off_sin); + auto dcarry_r = ((ddamp_t * off_cos) + (damp_t * doff_cos)); + auto dcarry_i = ((ddamp_t * off_sin) + (damp_t * doff_sin)); + fpr = ((mix_pr * carry_r) - (mix_pi * carry_i)); + fpi = ((mix_pr * carry_i) + (mix_pi * carry_r)); + dfpr = ((((dmix_pr * carry_r) + (mix_pr * dcarry_r)) - (dmix_pi * carry_i)) - (mix_pi * dcarry_i)); + dfpi = ((((dmix_pr * carry_i) + (mix_pr * dcarry_i)) + (dmix_pi * carry_r)) + (mix_pi * dcarry_r)); + bpr = ((mix_br * carry_r) - (mix_bi * carry_i)); + bpi = ((mix_br * carry_i) + (mix_bi * carry_r)); + dbpr = ((((dmix_br * carry_r) + (mix_br * dcarry_r)) - (dmix_bi * carry_i)) - (mix_bi * dcarry_i)); + dbpi = ((((dmix_br * carry_i) + (mix_br * dcarry_i)) + (dmix_bi * carry_r)) + (mix_bi * dcarry_r)); + fmr = ((mix_mr * carry_r) + (mix_mi * carry_i)); + fmi = (((-mix_mr) * carry_i) + (mix_mi * carry_r)); + dfmr = ((((dmix_mr * carry_r) + (mix_mr * dcarry_r)) + (dmix_mi * carry_i)) + (mix_mi * dcarry_i)); + dfmi = (((((-dmix_mr) * carry_i) - (mix_mr * dcarry_i)) + (dmix_mi * carry_r)) + (mix_mi * dcarry_r)); + bmr = ((mix_nr * carry_r) + (mix_ni * carry_i)); + bmi = (((-mix_nr) * carry_i) + (mix_ni * carry_r)); + dbmr = ((((dmix_nr * carry_r) + (mix_nr * dcarry_r)) + (dmix_ni * carry_i)) + (mix_ni * dcarry_i)); + dbmi = (((((-dmix_nr) * carry_i) - (mix_nr * dcarry_i)) + (dmix_ni * carry_r)) + (mix_ni * dcarry_r)); + } else { + auto old_fpr = fpr; + auto old_fpi = fpi; + auto old_dfpr = dfpr; + auto old_dfpi = dfpi; + fpr = (e2 * ((old_fpr * off_cos) - (old_fpi * off_sin))); + fpi = (e2 * ((old_fpr * off_sin) + (old_fpi * off_cos))); + dfpr = (de2 * ((old_fpr * off_cos) - (old_fpi * off_sin))); + dfpr = (dfpr + (e2 * ((((old_dfpr * off_cos) + (old_fpr * doff_cos)) - (old_dfpi * off_sin)) - (old_fpi * doff_sin)))); + dfpi = (de2 * ((old_fpr * off_sin) + (old_fpi * off_cos))); + dfpi = (dfpi + (e2 * ((((old_dfpr * off_sin) + (old_fpr * doff_sin)) + (old_dfpi * off_cos)) + (old_fpi * doff_cos)))); + auto old_fmr = fmr; + auto old_fmi = fmi; + auto old_dfmr = dfmr; + auto old_dfmi = dfmi; + fmr = (e2 * ((old_fmr * off_cos) + (old_fmi * off_sin))); + fmi = (e2 * (((-old_fmr) * off_sin) + (old_fmi * off_cos))); + dfmr = (de2 * ((old_fmr * off_cos) + (old_fmi * off_sin))); + dfmr = (dfmr + (e2 * ((((old_dfmr * off_cos) + (old_fmr * doff_cos)) + (old_dfmi * off_sin)) + (old_fmi * doff_sin)))); + dfmi = (de2 * (((-old_fmr) * off_sin) + (old_fmi * off_cos))); + dfmi = (dfmi + (e2 * (((((-old_dfmr) * off_sin) - (old_fmr * doff_sin)) + (old_dfmi * off_cos)) + (old_fmi * doff_cos)))); + } + // The longitudinal states carry a phase of their own, which nothing + // else in the state machine gives them. + turn_cos = 1.0f; + turn_sin = 0.0f; + dturn_cos = 0.0f; + dturn_sin = 0.0f; + if (bsk::truth(moving)) { + turn_cos = bsk::cos(turn_z); + turn_sin = bsk::sin(turn_z); + dturn_cos = ((-turn_sin) * dturn_z); + dturn_sin = (turn_cos * dturn_z); + } + auto old_zr = zr; + auto old_zi = zi; + auto old_dzr = dzr; + auto old_dzi = dzi; + auto spun_zr = ((old_zr * turn_cos) - (old_zi * turn_sin)); + auto spun_zi = ((old_zr * turn_sin) + (old_zi * turn_cos)); + auto dspun_zr = ((((old_dzr * turn_cos) + (old_zr * dturn_cos)) - (old_dzi * turn_sin)) - (old_zi * dturn_sin)); + auto dspun_zi = ((((old_dzr * turn_sin) + (old_zr * dturn_sin)) + (old_dzi * turn_cos)) + (old_zi * dturn_cos)); + if (bsk::truth((pools == 3))) { + // Three pools mix through a 3x3 formed in double, tangent and all: + // a direction through an operator this ill-conditioned needs the + // width as much as the value does. + if (bsk::truth(tabulated)) { + auto t5_ = _three_pool_from_table_jvp(pool_table, bsk::ld(((duration_row + event_base) + event), active_atom, 0), atom, atom_count, active_atom, r1, atom_r1_bound, atom_r1_semisolid, atom_exchange, atom_semisolid_exchange, atom_bound, d_bound, atom_semisolid, d_semisolid, ddt, wout, dwout); + t11 = bsk::get<0>(t5_); + t12 = bsk::get<1>(t5_); + t13 = bsk::get<2>(t5_); + t21 = bsk::get<3>(t5_); + t22 = bsk::get<4>(t5_); + t23 = bsk::get<5>(t5_); + t31 = bsk::get<6>(t5_); + t32 = bsk::get<7>(t5_); + t33 = bsk::get<8>(t5_); + grow_free = bsk::get<9>(t5_); + grow_pool_b = bsk::get<10>(t5_); + grow_semisolid = bsk::get<11>(t5_); + d_t11 = bsk::get<12>(t5_); + d_t12 = bsk::get<13>(t5_); + d_t13 = bsk::get<14>(t5_); + d_t21 = bsk::get<15>(t5_); + d_t22 = bsk::get<16>(t5_); + d_t23 = bsk::get<17>(t5_); + d_t31 = bsk::get<18>(t5_); + d_t32 = bsk::get<19>(t5_); + d_t33 = bsk::get<20>(t5_); + d_grow_free = bsk::get<21>(t5_); + d_grow_pool_b = bsk::get<22>(t5_); + d_grow_semisolid = bsk::get<23>(t5_); + } else { + auto t6_ = _three_pool_step_jvp(r1, bsk::truediv((-1000.0f * dt1), (atom_t1 * atom_t1)), atom_r1_bound, d_r1_bound, atom_r1_semisolid, d_r1_semisolid, atom_exchange, d_exchange, atom_semisolid_exchange, d_semisolid_exchange, atom_bound, d_bound, atom_semisolid, d_semisolid, event_dt, ddt, wout, dwout, narrow); + t11 = bsk::get<0>(t6_); + t12 = bsk::get<1>(t6_); + t13 = bsk::get<2>(t6_); + t21 = bsk::get<3>(t6_); + t22 = bsk::get<4>(t6_); + t23 = bsk::get<5>(t6_); + t31 = bsk::get<6>(t6_); + t32 = bsk::get<7>(t6_); + t33 = bsk::get<8>(t6_); + grow_free = bsk::get<9>(t6_); + grow_pool_b = bsk::get<10>(t6_); + grow_semisolid = bsk::get<11>(t6_); + d_t11 = bsk::get<12>(t6_); + d_t12 = bsk::get<13>(t6_); + d_t13 = bsk::get<14>(t6_); + d_t21 = bsk::get<15>(t6_); + d_t22 = bsk::get<16>(t6_); + d_t23 = bsk::get<17>(t6_); + d_t31 = bsk::get<18>(t6_); + d_t32 = bsk::get<19>(t6_); + d_t33 = bsk::get<20>(t6_); + d_grow_free = bsk::get<21>(t6_); + d_grow_pool_b = bsk::get<22>(t6_); + d_grow_semisolid = bsk::get<23>(t6_); + } + spun_hr = ((br * turn_cos) - (bi * turn_sin)); + spun_hi = ((br * turn_sin) + (bi * turn_cos)); + dspun_hr = ((((dbr * turn_cos) + (br * dturn_cos)) - (dbi * turn_sin)) - (bi * dturn_sin)); + dspun_hi = ((((dbr * turn_sin) + (br * dturn_sin)) + (dbi * turn_cos)) + (bi * dturn_cos)); + auto spun_cr = ((cr * turn_cos) - (ci * turn_sin)); + auto spun_ci = ((cr * turn_sin) + (ci * turn_cos)); + auto dspun_cr = ((((dcr * turn_cos) + (cr * dturn_cos)) - (dci * turn_sin)) - (ci * dturn_sin)); + auto dspun_ci = ((((dcr * turn_sin) + (cr * dturn_sin)) + (dci * turn_cos)) + (ci * dturn_cos)); + free_r = (((t11 * spun_zr) + (t12 * spun_hr)) + (t13 * spun_cr)); + free_i = (((t11 * spun_zi) + (t12 * spun_hi)) + (t13 * spun_ci)); + held_r = (((t21 * spun_zr) + (t22 * spun_hr)) + (t23 * spun_cr)); + held_i = (((t21 * spun_zi) + (t22 * spun_hi)) + (t23 * spun_ci)); + auto stuck_r = (((t31 * spun_zr) + (t32 * spun_hr)) + (t33 * spun_cr)); + auto stuck_i = (((t31 * spun_zi) + (t32 * spun_hi)) + (t33 * spun_ci)); + d_free_r = ((((((d_t11 * spun_zr) + (t11 * dspun_zr)) + (d_t12 * spun_hr)) + (t12 * dspun_hr)) + (d_t13 * spun_cr)) + (t13 * dspun_cr)); + d_free_i = ((((((d_t11 * spun_zi) + (t11 * dspun_zi)) + (d_t12 * spun_hi)) + (t12 * dspun_hi)) + (d_t13 * spun_ci)) + (t13 * dspun_ci)); + d_held_r = ((((((d_t21 * spun_zr) + (t21 * dspun_zr)) + (d_t22 * spun_hr)) + (t22 * dspun_hr)) + (d_t23 * spun_cr)) + (t23 * dspun_cr)); + d_held_i = ((((((d_t21 * spun_zi) + (t21 * dspun_zi)) + (d_t22 * spun_hi)) + (t22 * dspun_hi)) + (d_t23 * spun_ci)) + (t23 * dspun_ci)); + auto d_stuck_r = ((((((d_t31 * spun_zr) + (t31 * dspun_zr)) + (d_t32 * spun_hr)) + (t32 * dspun_hr)) + (d_t33 * spun_cr)) + (t33 * dspun_cr)); + auto d_stuck_i = ((((((d_t31 * spun_zi) + (t31 * dspun_zi)) + (d_t32 * spun_hi)) + (t32 * dspun_hi)) + (d_t33 * spun_ci)) + (t33 * dspun_ci)); + zr = ((damp_z * free_r) + bsk::where((state == 0), grow_free, 0.0f)); + zi = (damp_z * free_i); + dzr = (((ddamp_z * free_r) + (damp_z * d_free_r)) + bsk::where((state == 0), d_grow_free, 0.0f)); + dzi = ((ddamp_z * free_i) + (damp_z * d_free_i)); + br = ((damp_z * held_r) + bsk::where((state == 0), grow_pool_b, 0.0f)); + bi = (damp_z * held_i); + dbr = (((ddamp_z * held_r) + (damp_z * d_held_r)) + bsk::where((state == 0), d_grow_pool_b, 0.0f)); + dbi = ((ddamp_z * held_i) + (damp_z * d_held_i)); + cr = ((damp_z * stuck_r) + bsk::where((state == 0), grow_semisolid, 0.0f)); + ci = (damp_z * stuck_i); + dcr = (((ddamp_z * stuck_r) + (damp_z * d_stuck_r)) + bsk::where((state == 0), d_grow_semisolid, 0.0f)); + dci = ((ddamp_z * stuck_i) + (damp_z * d_stuck_i)); + } else if (bsk::truth((pools > 0))) { + // The exchange operator belongs to the interval, not to a dephasing + // order, so it is formed once and carries its own tangent; the + // per-order damping multiplies both pools, whose order-n states + // describe one dephasing configuration. + auto t7_ = _two_pool_step_jvp(r1, bsk::truediv((-1000.0f * dt1), (atom_t1 * atom_t1)), atom_r1_bound, d_r1_bound, atom_exchange, d_exchange, atom_bound, d_bound, event_dt, ddt, wout, dwout); + auto e11 = bsk::get<0>(t7_); + auto e12 = bsk::get<1>(t7_); + auto e21 = bsk::get<2>(t7_); + auto e22 = bsk::get<3>(t7_); + grow_free = bsk::get<4>(t7_); + auto grow_bound = bsk::get<5>(t7_); + auto d_e11 = bsk::get<6>(t7_); + auto d_e12 = bsk::get<7>(t7_); + auto d_e21 = bsk::get<8>(t7_); + auto d_e22 = bsk::get<9>(t7_); + d_grow_free = bsk::get<10>(t7_); + auto d_grow_bound = bsk::get<11>(t7_); + auto old_br = br; + auto old_bi = bi; + auto old_dbr = dbr; + auto old_dbi = dbi; + spun_hr = ((old_br * turn_cos) - (old_bi * turn_sin)); + spun_hi = ((old_br * turn_sin) + (old_bi * turn_cos)); + dspun_hr = ((((old_dbr * turn_cos) + (old_br * dturn_cos)) - (old_dbi * turn_sin)) - (old_bi * dturn_sin)); + dspun_hi = ((((old_dbr * turn_sin) + (old_br * dturn_sin)) + (old_dbi * turn_cos)) + (old_bi * dturn_cos)); + free_r = ((e11 * spun_zr) + (e12 * spun_hr)); + free_i = ((e11 * spun_zi) + (e12 * spun_hi)); + held_r = ((e21 * spun_zr) + (e22 * spun_hr)); + held_i = ((e21 * spun_zi) + (e22 * spun_hi)); + d_free_r = ((((d_e11 * spun_zr) + (e11 * dspun_zr)) + (d_e12 * spun_hr)) + (e12 * dspun_hr)); + d_free_i = ((((d_e11 * spun_zi) + (e11 * dspun_zi)) + (d_e12 * spun_hi)) + (e12 * dspun_hi)); + d_held_r = ((((d_e21 * spun_zr) + (e21 * dspun_zr)) + (d_e22 * spun_hr)) + (e22 * dspun_hr)); + d_held_i = ((((d_e21 * spun_zi) + (e21 * dspun_zi)) + (d_e22 * spun_hi)) + (e22 * dspun_hi)); + zr = ((damp_z * free_r) + bsk::where((state == 0), grow_free, 0.0f)); + zi = (damp_z * free_i); + dzr = (((ddamp_z * free_r) + (damp_z * d_free_r)) + bsk::where((state == 0), d_grow_free, 0.0f)); + dzi = ((ddamp_z * free_i) + (damp_z * d_free_i)); + br = ((damp_z * held_r) + bsk::where((state == 0), grow_bound, 0.0f)); + bi = (damp_z * held_i); + dbr = (((ddamp_z * held_r) + (damp_z * d_held_r)) + bsk::where((state == 0), d_grow_bound, 0.0f)); + dbi = ((ddamp_z * held_i) + (damp_z * d_held_i)); + } else { + dzr = (((dspun_zr * e1) + (spun_zr * de1)) + bsk::where((state == 0), drecovery, 0.0f)); + dzi = ((dspun_zi * e1) + (spun_zi * de1)); + zr = ((spun_zr * e1) + bsk::where((state == 0), recovery, 0.0f)); + zi = (spun_zi * e1); + } + auto event_action = bsk::cast(bsk::ld((action + event))); + auto pre_shift = (bsk::band(event_action, 1) != 0); + auto t8_ = _shift(fpr, fpi, fmr, fmi, state, state_mask, state_count); + shifted_pr = bsk::get<0>(t8_); + shifted_pi = bsk::get<1>(t8_); + shifted_mr = bsk::get<2>(t8_); + shifted_mi = bsk::get<3>(t8_); + auto t9_ = _shift(dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count); + shifted_dpr = bsk::get<0>(t9_); + shifted_dpi = bsk::get<1>(t9_); + shifted_dmr = bsk::get<2>(t9_); + shifted_dmi = bsk::get<3>(t9_); + fpr = bsk::where(pre_shift, shifted_pr, fpr); + fpi = bsk::where(pre_shift, shifted_pi, fpi); + fmr = bsk::where(pre_shift, shifted_mr, fmr); + fmi = bsk::where(pre_shift, shifted_mi, fmi); + dfpr = bsk::where(pre_shift, shifted_dpr, dfpr); + dfpi = bsk::where(pre_shift, shifted_dpi, dfpi); + dfmr = bsk::where(pre_shift, shifted_dmr, dfmr); + dfmi = bsk::where(pre_shift, shifted_dmi, dfmi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t10_ = _shift(bpr, bpi, bmr, bmi, state, state_mask, state_count); + s_bpr = bsk::get<0>(t10_); + s_bpi = bsk::get<1>(t10_); + s_bmr = bsk::get<2>(t10_); + s_bmi = bsk::get<3>(t10_); + auto t11_ = _shift(dbpr, dbpi, dbmr, dbmi, state, state_mask, state_count); + s_dbpr = bsk::get<0>(t11_); + s_dbpi = bsk::get<1>(t11_); + s_dbmr = bsk::get<2>(t11_); + s_dbmi = bsk::get<3>(t11_); + bpr = bsk::where(pre_shift, s_bpr, bpr); + bpi = bsk::where(pre_shift, s_bpi, bpi); + bmr = bsk::where(pre_shift, s_bmr, bmr); + bmi = bsk::where(pre_shift, s_bmi, bmi); + dbpr = bsk::where(pre_shift, s_dbpr, dbpr); + dbpi = bsk::where(pre_shift, s_dbpi, dbpi); + dbmr = bsk::where(pre_shift, s_dbmr, dbmr); + dbmi = bsk::where(pre_shift, s_dbmi, dbmi); + } + auto event_kind = bsk::ld((kind + event)); + auto is_rf = (event_kind == 1); + auto is_inversion = (bsk::band(event_action, 4) != 0); + auto invert = bsk::band(is_rf, is_inversion); + dzr = bsk::where(invert, (((-dinversion) * zr) - (atom_inversion * dzr)), dzr); + dzi = bsk::where(invert, (((-dinversion) * zi) - (atom_inversion * dzi)), dzi); + zr = bsk::where(invert, ((-atom_inversion) * zr), zr); + zi = bsk::where(invert, ((-atom_inversion) * zi), zi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // A chemically exchanging pool is free water and turns over like + // any other; a semisolid one is saturated instead. + dbr = bsk::where(invert, (((-dinversion) * br) - (atom_inversion * dbr)), dbr); + dbi = bsk::where(invert, (((-dinversion) * bi) - (atom_inversion * dbi)), dbi); + br = bsk::where(invert, ((-atom_inversion) * br), br); + bi = bsk::where(invert, ((-atom_inversion) * bi), bi); + } + auto event_flip = _event_value(flip, event_base, event, active_atom, single_train); + auto event_phase = _event_value(phase, event_base, event, active_atom, single_train); + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse a row of its own. + if (bsk::truth(shimmed)) { + auto row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + db1 = bsk::ld(((tangent_b1 + row) + atom), active_atom, 0.0f); + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + db1_phase = bsk::ld(((tangent_b1_phase + row) + atom), active_atom, 0.0f); + } + } + auto alpha = (event_flip * atom_b1); + auto dalpha = ((_event_value(tangent_flip, event_base, event, active_atom, single_train) * atom_b1) + (event_flip * db1)); + auto phi = (event_phase + atom_b1_phase); + auto dphi = (_event_value(tangent_phase, event_base, event, active_atom, single_train) + db1_phase); + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + // The semisolid pool absorbs the power the pulse deposits, so it reads + // the bare flip the transmit field gives the voxel. The offset + // reaches it through the voxel's own off-resonance, which is where + // the lineshape's slope enters a forward direction. + auto offset = (bsk::ld((rf_frequency + event)) - atom_b0); + auto t12_ = _lineshape_at_slope(lineshape, offset, lineshape_bins, lineshape_step); + auto shape = bsk::get<0>(t12_); + auto shape_slope = bsk::get<1>(t12_); + auto deposited = bsk::ld((saturation + event)); + auto absorbed = bsk::exp((((deposited * alpha) * alpha) * shape)); + auto d_exponent = (deposited * ((((2.0f * alpha) * dalpha) * shape) - (((alpha * alpha) * shape_slope) * db0))); + auto saturating = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((pools == 1))) { + dbr = bsk::where(saturating, (absorbed * (dbr + (br * d_exponent))), dbr); + dbi = bsk::where(saturating, (absorbed * (dbi + (bi * d_exponent))), dbi); + br = bsk::where(saturating, (absorbed * br), br); + bi = bsk::where(saturating, (absorbed * bi), bi); + } else { + dcr = bsk::where(saturating, (absorbed * (dcr + (cr * d_exponent))), dcr); + dci = bsk::where(saturating, (absorbed * (dci + (ci * d_exponent))), dci); + cr = bsk::where(saturating, (absorbed * cr), cr); + ci = bsk::where(saturating, (absorbed * ci), ci); + } + } + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + if (bsk::truth(dynamic)) { + // The array was resolved outside the kernel, so a direction + // along it arrives already carried through the pulse integral. + auto held = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, active_atom); + auto moved = _dynamic_pair_at(pair_direction, pair_index, event_base, event, atom, atom_count, active_atom); + auto t13_ = held; + pair_ar = bsk::get<0>(t13_); + pair_ai = bsk::get<1>(t13_); + pair_br = bsk::get<2>(t13_); + pair_bi = bsk::get<3>(t13_); + auto t14_ = moved; + dot_ar = bsk::get<0>(t14_); + dot_ai = bsk::get<1>(t14_); + dot_br = bsk::get<2>(t14_); + dot_bi = bsk::get<3>(t14_); + } else { + auto read = _profile_pair_slope(profile, _table_row(profile_index, event, location, locations), alpha, profile_bins, profile_step); + // The flip angle carries the tangent into the table. + auto t15_ = bsk::make_tup(bsk::get<0>(read), bsk::get<2>(read)); + pair_ar = bsk::get<0>(t15_); + pair_ai = bsk::get<1>(t15_); + auto t16_ = bsk::make_tup(bsk::get<4>(read), bsk::get<6>(read)); + pair_br = bsk::get<0>(t16_); + pair_bi = bsk::get<1>(t16_); + auto t17_ = bsk::make_tup((bsk::get<1>(read) * dalpha), (bsk::get<3>(read) * dalpha)); + dot_ar = bsk::get<0>(t17_); + dot_ai = bsk::get<1>(t17_); + auto t18_ = bsk::make_tup((bsk::get<5>(read) * dalpha), (bsk::get<7>(read) * dalpha)); + dot_br = bsk::get<0>(t18_); + dot_bi = bsk::get<1>(t18_); + } + // The RF phase turns the axis after the pair comes out, and so + // reaches ``b`` alone. + auto turn_r = bsk::cos(phi); + auto turn_i = (-bsk::sin(phi)); + auto spun_br = ((pair_br * turn_r) - (pair_bi * turn_i)); + auto spun_bi = ((pair_br * turn_i) + (pair_bi * turn_r)); + auto slope_br = dot_br; + auto slope_bi = dot_bi; + auto t19_ = _rotate_spinor_dual(pair_ar, pair_ai, spun_br, spun_bi, dot_ar, dot_ai, (((slope_br * turn_r) - (slope_bi * turn_i)) + (dphi * spun_bi)), (((slope_br * turn_i) + (slope_bi * turn_r)) - (dphi * spun_br)), fpr, fpi, fmr, fmi, zr, zi, dfpr, dfpi, dfmr, dfmi, dzr, dzi); + shaped_pr = bsk::get<0>(t19_); + shaped_pi = bsk::get<1>(t19_); + shaped_mr = bsk::get<2>(t19_); + shaped_mi = bsk::get<3>(t19_); + shaped_zr = bsk::get<4>(t19_); + shaped_zi = bsk::get<5>(t19_); + shaped_dpr = bsk::get<6>(t19_); + shaped_dpi = bsk::get<7>(t19_); + shaped_dmr = bsk::get<8>(t19_); + shaped_dmi = bsk::get<9>(t19_); + shaped_dzr = bsk::get<10>(t19_); + shaped_dzi = bsk::get<11>(t19_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // The same pulse, the same rotation. + auto t20_ = _rotate_spinor_dual(pair_ar, pair_ai, spun_br, spun_bi, dot_ar, dot_ai, (((slope_br * turn_r) - (slope_bi * turn_i)) + (dphi * spun_bi)), (((slope_br * turn_i) + (slope_bi * turn_r)) - (dphi * spun_br)), bpr, bpi, bmr, bmi, br, bi, dbpr, dbpi, dbmr, dbmi, dbr, dbi); + held_pr = bsk::get<0>(t20_); + held_pi = bsk::get<1>(t20_); + held_mr = bsk::get<2>(t20_); + held_mi = bsk::get<3>(t20_); + held_zr = bsk::get<4>(t20_); + held_zi = bsk::get<5>(t20_); + held_dpr = bsk::get<6>(t20_); + held_dpi = bsk::get<7>(t20_); + held_dmr = bsk::get<8>(t20_); + held_dmi = bsk::get<9>(t20_); + held_dzr = bsk::get<10>(t20_); + held_dzi = bsk::get<11>(t20_); + } + } + auto cosine = bsk::cos(alpha); + auto sine = bsk::sin(alpha); + auto dcosine = ((-sine) * dalpha); + auto dsine = (cosine * dalpha); + auto cos_phi = bsk::cos(phi); + auto sin_phi = bsk::sin(phi); + auto cos_2phi = bsk::cos((2.0f * phi)); + auto sin_2phi = bsk::sin((2.0f * phi)); + auto dcos_phi = ((-sin_phi) * dphi); + auto dsin_phi = (cos_phi * dphi); + auto dcos_2phi = ((-2.0f * sin_2phi) * dphi); + auto dsin_2phi = ((2.0f * cos_2phi) * dphi); + auto t21_ = _rotate_flip_phase_jvp(cosine, dcosine, sine, dsine, cos_phi, dcos_phi, sin_phi, dsin_phi, cos_2phi, dcos_2phi, sin_2phi, dsin_2phi, fpr, fpi, fmr, fmi, zr, zi, dfpr, dfpi, dfmr, dfmi, dzr, dzi); + rotated_pr = bsk::get<0>(t21_); + rotated_pi = bsk::get<1>(t21_); + rotated_mr = bsk::get<2>(t21_); + rotated_mi = bsk::get<3>(t21_); + rotated_zr = bsk::get<4>(t21_); + rotated_zi = bsk::get<5>(t21_); + rotated_dpr = bsk::get<6>(t21_); + rotated_dpi = bsk::get<7>(t21_); + rotated_dmr = bsk::get<8>(t21_); + rotated_dmi = bsk::get<9>(t21_); + rotated_dzr = bsk::get<10>(t21_); + rotated_dzi = bsk::get<11>(t21_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t22_ = _rotate_flip_phase_jvp(cosine, dcosine, sine, dsine, cos_phi, dcos_phi, sin_phi, dsin_phi, cos_2phi, dcos_2phi, sin_2phi, dsin_2phi, bpr, bpi, bmr, bmi, br, bi, dbpr, dbpi, dbmr, dbmi, dbr, dbi); + b_pr = bsk::get<0>(t22_); + b_pi = bsk::get<1>(t22_); + b_mr = bsk::get<2>(t22_); + b_mi = bsk::get<3>(t22_); + b_zr = bsk::get<4>(t22_); + b_zi = bsk::get<5>(t22_); + b_dpr = bsk::get<6>(t22_); + b_dpi = bsk::get<7>(t22_); + b_dmr = bsk::get<8>(t22_); + b_dmi = bsk::get<9>(t22_); + b_dzr = bsk::get<10>(t22_); + b_dzi = bsk::get<11>(t22_); + } + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + rotated_pr = shaped_pr; + rotated_pi = shaped_pi; + rotated_mr = shaped_mr; + rotated_mi = shaped_mi; + rotated_zr = shaped_zr; + rotated_zi = shaped_zi; + rotated_dpr = shaped_dpr; + rotated_dpi = shaped_dpi; + rotated_dmr = shaped_dmr; + rotated_dmi = shaped_dmi; + rotated_dzr = shaped_dzr; + rotated_dzi = shaped_dzi; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + b_pr = held_pr; + b_pi = held_pi; + b_mr = held_mr; + b_mi = held_mi; + b_zr = held_zr; + b_zi = held_zi; + b_dpr = held_dpr; + b_dpi = held_dpi; + b_dmr = held_dmr; + b_dmi = held_dmi; + b_dzr = held_dzr; + b_dzi = held_dzi; + } + } + auto rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + fpr = bsk::where(rotate, rotated_pr, fpr); + fpi = bsk::where(rotate, rotated_pi, fpi); + fmr = bsk::where(rotate, rotated_mr, fmr); + fmi = bsk::where(rotate, rotated_mi, fmi); + zr = bsk::where(rotate, rotated_zr, zr); + zi = bsk::where(rotate, rotated_zi, zi); + dfpr = bsk::where(rotate, rotated_dpr, dfpr); + dfpi = bsk::where(rotate, rotated_dpi, dfpi); + dfmr = bsk::where(rotate, rotated_dmr, dfmr); + dfmi = bsk::where(rotate, rotated_dmi, dfmi); + dzr = bsk::where(rotate, rotated_dzr, dzr); + dzi = bsk::where(rotate, rotated_dzi, dzi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + bpr = bsk::where(rotate, b_pr, bpr); + bpi = bsk::where(rotate, b_pi, bpi); + bmr = bsk::where(rotate, b_mr, bmr); + bmi = bsk::where(rotate, b_mi, bmi); + br = bsk::where(rotate, b_zr, br); + bi = bsk::where(rotate, b_zi, bi); + dbpr = bsk::where(rotate, b_dpr, dbpr); + dbpi = bsk::where(rotate, b_dpi, dbpi); + dbmr = bsk::where(rotate, b_dmr, dbmr); + dbmi = bsk::where(rotate, b_dmi, dbmi); + dbr = bsk::where(rotate, b_dzr, dbr); + dbi = bsk::where(rotate, b_dzi, dbi); + } + auto record = bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)); + auto adc_cos = bsk::cos(event_phase); + auto adc_sin = bsk::sin(event_phase); + auto dadc_phase = _event_value(tangent_phase, event_base, event, active_atom, single_train); + auto dadc_cos = ((-adc_sin) * dadc_phase); + auto dadc_sin = (adc_cos * dadc_phase); + read_r = fpr; + read_i = fpi; + dread_r = dfpr; + dread_i = dfpi; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + read_r = (fpr + bpr); + read_i = (fpi + bpi); + dread_r = (dfpr + dbpr); + dread_i = (dfpi + dbpi); + } + signal_real = (dm0 * ((read_r * adc_cos) + (read_i * adc_sin))); + signal_real = (signal_real + (atom_m0 * ((((dread_r * adc_cos) + (read_r * dadc_cos)) + (dread_i * adc_sin)) + (read_i * dadc_sin)))); + signal_imag = (dm0 * ((read_i * adc_cos) - (read_r * adc_sin))); + signal_imag = (signal_imag + (atom_m0 * ((((dread_i * adc_cos) + (read_i * dadc_cos)) - (dread_r * adc_sin)) - (read_r * dadc_sin)))); + auto out_ = bsk::ld((output_index + event)); + auto output_offset = ((problem * output_count) + out_); + auto output_mask = bsk::band(bsk::band(bsk::band(active_atom, (state == 0)), record), (out_ >= 0)); + bsk::st(((output_real + output_offset) + state), signal_real, output_mask); + bsk::st(((output_imag + output_offset) + state), signal_imag, output_mask); + auto do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + auto t23_ = _shift(fpr, fpi, fmr, fmi, state, state_mask, state_count); + shifted_pr = bsk::get<0>(t23_); + shifted_pi = bsk::get<1>(t23_); + shifted_mr = bsk::get<2>(t23_); + shifted_mi = bsk::get<3>(t23_); + auto t24_ = _shift(dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count); + shifted_dpr = bsk::get<0>(t24_); + shifted_dpi = bsk::get<1>(t24_); + shifted_dmr = bsk::get<2>(t24_); + shifted_dmi = bsk::get<3>(t24_); + fpr = bsk::where(do_shift, shifted_pr, fpr); + fpi = bsk::where(do_shift, shifted_pi, fpi); + fmr = bsk::where(do_shift, shifted_mr, fmr); + fmi = bsk::where(do_shift, shifted_mi, fmi); + dfpr = bsk::where(do_shift, shifted_dpr, dfpr); + dfpi = bsk::where(do_shift, shifted_dpi, dfpi); + dfmr = bsk::where(do_shift, shifted_dmr, dfmr); + dfmi = bsk::where(do_shift, shifted_dmi, dfmi); + auto spoil = (bsk::band(event_action, 8) != 0); + fpr = bsk::where(spoil, 0.0f, fpr); + fpi = bsk::where(spoil, 0.0f, fpi); + fmr = bsk::where(spoil, 0.0f, fmr); + fmi = bsk::where(spoil, 0.0f, fmi); + dfpr = bsk::where(spoil, 0.0f, dfpr); + dfpi = bsk::where(spoil, 0.0f, dfpi); + dfmr = bsk::where(spoil, 0.0f, dfmr); + dfmi = bsk::where(spoil, 0.0f, dfmi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t25_ = _shift(bpr, bpi, bmr, bmi, state, state_mask, state_count); + s_bpr = bsk::get<0>(t25_); + s_bpi = bsk::get<1>(t25_); + s_bmr = bsk::get<2>(t25_); + s_bmi = bsk::get<3>(t25_); + auto t26_ = _shift(dbpr, dbpi, dbmr, dbmi, state, state_mask, state_count); + s_dbpr = bsk::get<0>(t26_); + s_dbpi = bsk::get<1>(t26_); + s_dbmr = bsk::get<2>(t26_); + s_dbmi = bsk::get<3>(t26_); + bpr = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_bpr, bpr)); + bpi = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_bpi, bpi)); + bmr = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_bmr, bmr)); + bmi = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_bmi, bmi)); + dbpr = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_dbpr, dbpr)); + dbpi = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_dbpi, dbpi)); + dbmr = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_dbmr, dbmr)); + dbmi = bsk::where(spoil, 0.0f, bsk::where(do_shift, s_dbmi, dbmi)); + } + } +} + +BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_duration, float* dot_flip, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_value, float* trajectory_tangent, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V adjoint_mt{}; + bsk::V adjoint_mv{}; + bsk::V adjoint_pt{}; + bsk::V adjoint_pv{}; + bsk::V alpha_bar_terms_tangent{}; + bsk::V alpha_bar_terms_value{}; + bsk::V alpha_tangent{}; + bsk::V alpha_value{}; + bsk::V atom_b1{}; + bsk::V atom_damping{}; + bsk::V atom_dot_b1{}; + bsk::V atom_dot_damping{}; + bsk::V atom_dot_inversion{}; + bsk::V atom_dot_m0{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V bare1_tangent{}; + bsk::V bare1_value{}; + bsk::V bare2_tangent{}; + bsk::V bare2_value{}; + bsk::V chs_tangent{}; + bsk::V chs_value{}; + bsk::V cosine_tangent{}; + bsk::V cosine_value{}; + bsk::V damp_t{}; + bsk::V damp_t_tangent{}; + bsk::V damp_z{}; + bsk::V damp_z_tangent{}; + bool do_shift{}; + bsk::V dt_tangent{}; + bsk::V dt_value{}; + bsk::V duration_gain_tangent{}; + bsk::V duration_gain_value{}; + bsk::V e1_tangent{}; + bsk::V e1_value{}; + bsk::V e2_tangent{}; + bsk::V e2_value{}; + std::int64_t event{}; + std::int32_t event_action{}; + bsk::V event_dot_flip{}; + bsk::V event_flip{}; + std::int32_t event_kind{}; + bsk::V flip_gain_tangent{}; + bsk::V grad_b1_tangent{}; + bsk::V grad_b1_value{}; + bsk::V grad_damping_tangent{}; + bsk::V grad_damping_value{}; + bsk::V grad_e1_tangent{}; + bsk::V grad_e1_value{}; + bsk::V grad_inversion_tangent{}; + bsk::V grad_inversion_value{}; + bsk::V grad_m0_tangent{}; + bsk::V grad_m0_value{}; + bsk::V grad_t1_tangent{}; + bsk::V grad_t1_value{}; + bsk::V grad_t2_tangent{}; + bsk::V grad_t2_value{}; + bsk::V half_sine_tangent{}; + bsk::V half_sine_value{}; + bool invert{}; + bsk::V inverted_tangent{}; + bool is_inversion{}; + bool is_rf{}; + bsk::V long_bar_tangent{}; + bsk::V long_bar_value{}; + bsk::V long_tangent{}; + bsk::V long_value{}; + bsk::V minus_bar_tangent{}; + bsk::V minus_bar_value{}; + bsk::V minus_tangent{}; + bsk::V minus_value{}; + bsk::V plus_bar_tangent{}; + bsk::V plus_bar_value{}; + bsk::V plus_tangent{}; + bsk::V plus_value{}; + bool pre_shift{}; + bsk::V problem{}; + bsk::V pulse_b1{}; + bsk::V pulse_dot_b1{}; + bsk::V recovery_tangent{}; + bsk::V recovery_value{}; + bool rotate{}; + bsk::V rotated_mbt{}; + bsk::V rotated_mbv{}; + bsk::V rotated_mt{}; + bsk::V rotated_mv{}; + bsk::V rotated_pbt{}; + bsk::V rotated_pbv{}; + bsk::V rotated_pt{}; + bsk::V rotated_pv{}; + bsk::V rotated_zbt{}; + bsk::V rotated_zbv{}; + bsk::V rotated_zt{}; + bsk::V rotated_zv{}; + bsk::V row_m_tangent{}; + bsk::V row_m_value{}; + bsk::V row_p_tangent{}; + bsk::V row_p_value{}; + bsk::V row_z_tangent{}; + bsk::V row_z_value{}; + bsk::V scale1_tangent{}; + bsk::V scale2_tangent{}; + bsk::V shifted_mt{}; + bsk::V shifted_mv{}; + bsk::V shifted_pt{}; + bsk::V shifted_pv{}; + std::int64_t shim_row{}; + bsk::V shs_tangent{}; + bsk::V shs_value{}; + bsk::V sine_tangent{}; + bsk::V sine_value{}; + bsk::V slot{}; + bool spoil{}; + bsk::V spread_tangent{}; + bsk::V spread_value{}; + bsk::V stage_mt{}; + bsk::V stage_mv{}; + bsk::V stage_pt{}; + bsk::V stage_pv{}; + bsk::V stage_zt{}; + problem = (problem_base + (bsk::program_id(0) * problems)); + problem = (problem + bsk::arange_y()); + auto state = bsk::arange_x(); + // The grid rounds up to whole tiles, so the last program of a wave reaches + // past it. Those problems are real, but their trajectory rows belong to a + // later launch and do not exist yet. + auto active_atom = (problem < problem_end); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + // The trajectory holds the state entering every event: three planes of + // configuration orders, for the value and the tangent alike. + auto record_stride = (3 * state_count); + auto trajectory = ((((problem - problem_base) * event_count) * record_stride) + state); + auto minus_plane = state_count; + auto long_plane = (2 * state_count); + auto empty = bsk::full(0); + plus_value = empty; + plus_tangent = empty; + minus_value = empty; + minus_tangent = empty; + long_value = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + long_tangent = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + auto atom_dot_t1 = bsk::ld((dot_t1 + atom), active_atom, 0.0f); + auto atom_dot_t2 = bsk::ld((dot_t2 + atom), active_atom, 0.0f); + atom_dot_m0 = 0.0f; + if (bsk::truth(density)) { + atom_dot_m0 = bsk::ld((dot_m0 + scalar_atom), active_atom, 0.0f); + } + atom_dot_b1 = 0.0f; + if (bsk::truth(transmit)) { + atom_dot_b1 = bsk::ld((dot_b1 + scalar_atom), active_atom, 0.0f); + } + atom_dot_inversion = 0.0f; + if (bsk::truth(inverting)) { + atom_dot_inversion = bsk::ld((dot_inversion_efficiency + scalar_atom), active_atom, 0.0f); + } + atom_damping = 0.0f; + atom_dot_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + atom_dot_damping = bsk::ld((dot_diffusion + scalar_atom), active_atom, 0.0f); + } + auto order = bsk::cast(state); + auto longitudinal_weight = (order * order); + auto transverse_weight = ((longitudinal_weight + order) + 0.3333333333333333f); + auto rate1_value = bsk::truediv(1000.0f, atom_t1); + auto rate1_tangent = bsk::truediv((-1000.0f * atom_dot_t1), (atom_t1 * atom_t1)); + auto rate2_value = bsk::truediv(1000.0f, atom_t2); + auto rate2_tangent = bsk::truediv((-1000.0f * atom_dot_t2), (atom_t2 * atom_t2)); + auto event_base = (train * event_count); + for (std::int64_t event = 0; event < event_count; event += 1) { + slot = (trajectory + (event * record_stride)); + bsk::st((trajectory_value + slot), plus_value, state_mask); + bsk::st(((trajectory_value + slot) + minus_plane), minus_value, state_mask); + bsk::st(((trajectory_value + slot) + long_plane), long_value, state_mask); + bsk::st((trajectory_tangent + slot), plus_tangent, state_mask); + bsk::st(((trajectory_tangent + slot) + minus_plane), minus_tangent, state_mask); + bsk::st(((trajectory_tangent + slot) + long_plane), long_tangent, state_mask); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + dt_tangent = _event_value(dot_duration, event_base, event, active_atom, single_train); + e1_value = bsk::exp(((-rate1_value) * dt_value)); + e1_tangent = ((-e1_value) * ((rate1_value * dt_tangent) + (rate1_tangent * dt_value))); + e2_value = bsk::exp(((-rate2_value) * dt_value)); + e2_tangent = ((-e2_value) * ((rate2_value * dt_tangent) + (rate2_tangent * dt_value))); + damp_z = 1.0f; + damp_z_tangent = 0.0f; + damp_t = 1.0f; + damp_t_tangent = 0.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping_jvp(atom_damping, atom_dot_damping, dt_value, dt_tangent, order); + damp_z = bsk::get<0>(t0_); + damp_z_tangent = bsk::get<1>(t0_); + damp_t = bsk::get<2>(t0_); + damp_t_tangent = bsk::get<3>(t0_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + auto t1_ = bsk::make_tup((1.0f - e1_value), (-e1_tangent)); + recovery_value = bsk::get<0>(t1_); + recovery_tangent = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup(e1_value, e1_tangent); + bare1_value = bsk::get<0>(t2_); + bare1_tangent = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup(e2_value, e2_tangent); + bare2_value = bsk::get<0>(t3_); + bare2_tangent = bsk::get<1>(t3_); + e1_tangent = ((e1_tangent * damp_z) + (bare1_value * damp_z_tangent)); + e1_value = (bare1_value * damp_z); + e2_tangent = ((e2_tangent * damp_t) + (bare2_value * damp_t_tangent)); + e2_value = (bare2_value * damp_t); + plus_tangent = ((plus_value * e2_tangent) + (plus_tangent * e2_value)); + plus_value = (plus_value * e2_value); + minus_tangent = ((minus_value * e2_tangent) + (minus_tangent * e2_value)); + minus_value = (minus_value * e2_value); + long_tangent = ((long_value * e1_tangent) + (long_tangent * e1_value)); + long_value = (long_value * e1_value); + long_value = (long_value + bsk::where((state == 0), recovery_value, 0.0f)); + long_tangent = (long_tangent + bsk::where((state == 0), recovery_tangent, 0.0f)); + event_action = bsk::cast(bsk::ld((action + event))); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t4_ = _shift_real(plus_value, minus_value, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t4_); + shifted_mv = bsk::get<1>(t4_); + auto t5_ = _shift_real(plus_tangent, minus_tangent, state, state_mask, state_count); + shifted_pt = bsk::get<0>(t5_); + shifted_mt = bsk::get<1>(t5_); + plus_value = bsk::where(pre_shift, shifted_pv, plus_value); + minus_value = bsk::where(pre_shift, shifted_mv, minus_value); + plus_tangent = bsk::where(pre_shift, shifted_pt, plus_tangent); + minus_tangent = bsk::where(pre_shift, shifted_mt, minus_tangent); + event_kind = bsk::ld((kind + event)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + auto inverted_value = ((-atom_inversion) * long_value); + inverted_tangent = ((-atom_inversion) * long_tangent); + inverted_tangent = (inverted_tangent - (atom_dot_inversion * long_value)); + long_value = bsk::where(invert, inverted_value, long_value); + long_tangent = bsk::where(invert, inverted_tangent, long_tangent); + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_dot_flip = _event_value(dot_flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + pulse_dot_b1 = atom_dot_b1; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse the row of the shim it drives. + if (bsk::truth(shimmed)) { + shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + pulse_dot_b1 = bsk::ld(((dot_b1 + shim_row) + atom), active_atom, 0.0f); + } + alpha_value = (event_flip * pulse_b1); + alpha_tangent = ((event_dot_flip * pulse_b1) + (event_flip * pulse_dot_b1)); + cosine_value = bsk::cos(alpha_value); + sine_value = bsk::sin(alpha_value); + cosine_tangent = ((-sine_value) * alpha_tangent); + sine_tangent = (cosine_value * alpha_tangent); + chs_value = (0.5f * (1.0f + cosine_value)); + chs_tangent = (0.5f * cosine_tangent); + shs_value = (0.5f * (1.0f - cosine_value)); + shs_tangent = (-0.5f * cosine_tangent); + half_sine_value = (0.5f * sine_value); + half_sine_tangent = (0.5f * sine_tangent); + rotated_pv = ((chs_value * plus_value) + (shs_value * minus_value)); + rotated_pv = (rotated_pv - (sine_value * long_value)); + rotated_pt = ((chs_value * plus_tangent) + (chs_tangent * plus_value)); + rotated_pt = (rotated_pt + ((shs_value * minus_tangent) + (shs_tangent * minus_value))); + rotated_pt = (rotated_pt - ((sine_value * long_tangent) + (sine_tangent * long_value))); + rotated_mv = ((shs_value * plus_value) + (chs_value * minus_value)); + rotated_mv = (rotated_mv + (sine_value * long_value)); + rotated_mt = ((shs_value * plus_tangent) + (shs_tangent * plus_value)); + rotated_mt = (rotated_mt + ((chs_value * minus_tangent) + (chs_tangent * minus_value))); + rotated_mt = (rotated_mt + ((sine_value * long_tangent) + (sine_tangent * long_value))); + rotated_zv = ((half_sine_value * plus_value) - (half_sine_value * minus_value)); + rotated_zv = (rotated_zv + (cosine_value * long_value)); + rotated_zt = ((half_sine_value * plus_tangent) + (half_sine_tangent * plus_value)); + rotated_zt = (rotated_zt - ((half_sine_value * minus_tangent) + (half_sine_tangent * minus_value))); + rotated_zt = (rotated_zt + ((cosine_value * long_tangent) + (cosine_tangent * long_value))); + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + plus_value = bsk::where(rotate, rotated_pv, plus_value); + plus_tangent = bsk::where(rotate, rotated_pt, plus_tangent); + minus_value = bsk::where(rotate, rotated_mv, minus_value); + minus_tangent = bsk::where(rotate, rotated_mt, minus_tangent); + long_value = bsk::where(rotate, rotated_zv, long_value); + long_tangent = bsk::where(rotate, rotated_zt, long_tangent); + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + auto t6_ = _shift_real(plus_value, minus_value, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t6_); + shifted_mv = bsk::get<1>(t6_); + auto t7_ = _shift_real(plus_tangent, minus_tangent, state, state_mask, state_count); + shifted_pt = bsk::get<0>(t7_); + shifted_mt = bsk::get<1>(t7_); + plus_value = bsk::where(do_shift, shifted_pv, plus_value); + minus_value = bsk::where(do_shift, shifted_mv, minus_value); + plus_tangent = bsk::where(do_shift, shifted_pt, plus_tangent); + minus_tangent = bsk::where(do_shift, shifted_mt, minus_tangent); + spoil = (bsk::band(event_action, 8) != 0); + plus_value = bsk::where(spoil, 0.0f, plus_value); + minus_value = bsk::where(spoil, 0.0f, minus_value); + plus_tangent = bsk::where(spoil, 0.0f, plus_tangent); + minus_tangent = bsk::where(spoil, 0.0f, minus_tangent); + } + plus_bar_value = empty; + plus_bar_tangent = empty; + minus_bar_value = empty; + minus_bar_tangent = empty; + long_bar_value = empty; + long_bar_tangent = empty; + auto zero = bsk::full(0); + grad_t1_value = zero; + grad_t1_tangent = zero; + grad_t2_value = zero; + grad_t2_tangent = zero; + grad_m0_value = zero; + grad_m0_tangent = zero; + grad_b1_value = zero; + grad_b1_tangent = zero; + grad_inversion_value = zero; + grad_inversion_tangent = zero; + grad_damping_value = zero; + grad_damping_tangent = zero; + for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + event = ((event_count - 1) - reverse); + slot = (trajectory + (event * record_stride)); + auto entry_pv = bsk::ld((trajectory_value + slot), state_mask, 0.0f); + auto entry_mv = bsk::ld(((trajectory_value + slot) + minus_plane), state_mask, 0.0f); + auto entry_zv = bsk::ld(((trajectory_value + slot) + long_plane), state_mask, 0.0f); + auto entry_pt = bsk::ld((trajectory_tangent + slot), state_mask, 0.0f); + auto entry_mt = bsk::ld(((trajectory_tangent + slot) + minus_plane), state_mask, 0.0f); + auto entry_zt = bsk::ld(((trajectory_tangent + slot) + long_plane), state_mask, 0.0f); + event_action = bsk::cast(bsk::ld((action + event))); + event_kind = bsk::ld((kind + event)); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + dt_tangent = _event_value(dot_duration, event_base, event, active_atom, single_train); + e1_value = bsk::exp(((-rate1_value) * dt_value)); + e1_tangent = ((-e1_value) * ((rate1_value * dt_tangent) + (rate1_tangent * dt_value))); + e2_value = bsk::exp(((-rate2_value) * dt_value)); + e2_tangent = ((-e2_value) * ((rate2_value * dt_tangent) + (rate2_tangent * dt_value))); + damp_z = 1.0f; + damp_z_tangent = 0.0f; + damp_t = 1.0f; + damp_t_tangent = 0.0f; + if (bsk::truth(diffusing)) { + auto t8_ = _damping_jvp(atom_damping, atom_dot_damping, dt_value, dt_tangent, order); + damp_z = bsk::get<0>(t8_); + damp_z_tangent = bsk::get<1>(t8_); + damp_t = bsk::get<2>(t8_); + damp_t_tangent = bsk::get<3>(t8_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + auto t9_ = bsk::make_tup((1.0f - e1_value), (-e1_tangent)); + recovery_value = bsk::get<0>(t9_); + recovery_tangent = bsk::get<1>(t9_); + auto t10_ = bsk::make_tup(e1_value, e1_tangent); + bare1_value = bsk::get<0>(t10_); + bare1_tangent = bsk::get<1>(t10_); + auto t11_ = bsk::make_tup(e2_value, e2_tangent); + bare2_value = bsk::get<0>(t11_); + bare2_tangent = bsk::get<1>(t11_); + e1_tangent = ((e1_tangent * damp_z) + (bare1_value * damp_z_tangent)); + e1_value = (bare1_value * damp_z); + e2_tangent = ((e2_tangent * damp_t) + (bare2_value * damp_t_tangent)); + e2_value = (bare2_value * damp_t); + // Replay the intra-event stages from the recorded entry state. + stage_pv = (entry_pv * e2_value); + stage_pt = ((entry_pv * e2_tangent) + (entry_pt * e2_value)); + stage_mv = (entry_mv * e2_value); + stage_mt = ((entry_mv * e2_tangent) + (entry_mt * e2_value)); + auto stage_zv = ((entry_zv * e1_value) + bsk::where((state == 0), recovery_value, 0.0f)); + stage_zt = ((entry_zv * e1_tangent) + (entry_zt * e1_value)); + stage_zt = (stage_zt + bsk::where((state == 0), recovery_tangent, 0.0f)); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t12_ = _shift_real(stage_pv, stage_mv, state, state_mask, state_count); + shifted_pv = bsk::get<0>(t12_); + shifted_mv = bsk::get<1>(t12_); + auto t13_ = _shift_real(stage_pt, stage_mt, state, state_mask, state_count); + shifted_pt = bsk::get<0>(t13_); + shifted_mt = bsk::get<1>(t13_); + stage_pv = bsk::where(pre_shift, shifted_pv, stage_pv); + stage_mv = bsk::where(pre_shift, shifted_mv, stage_mv); + stage_pt = bsk::where(pre_shift, shifted_pt, stage_pt); + stage_mt = bsk::where(pre_shift, shifted_mt, stage_mt); + // Undo the trailing spoil or shift. + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + spoil = (bsk::band(event_action, 8) != 0); + auto t14_ = _shift_real_adjoint(plus_bar_value, minus_bar_value, state, state_mask, state_count); + adjoint_pv = bsk::get<0>(t14_); + adjoint_mv = bsk::get<1>(t14_); + auto t15_ = _shift_real_adjoint(plus_bar_tangent, minus_bar_tangent, state, state_mask, state_count); + adjoint_pt = bsk::get<0>(t15_); + adjoint_mt = bsk::get<1>(t15_); + auto trailing = bsk::band(do_shift, bsk::bnot(spoil)); + plus_bar_value = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_pv, plus_bar_value)); + minus_bar_value = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_mv, minus_bar_value)); + plus_bar_tangent = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_pt, plus_bar_tangent)); + minus_bar_tangent = bsk::where(spoil, 0.0f, bsk::where(trailing, adjoint_mt, minus_bar_tangent)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + auto inversion_gain = (-bsk::sum_x(bsk::where(invert, (long_bar_value * stage_zv), 0.0f))); + auto inversion_gain_tangent = (-bsk::sum_x(bsk::where(invert, ((long_bar_value * stage_zt) + (long_bar_tangent * stage_zv)), 0.0f))); + grad_inversion_value = (grad_inversion_value + inversion_gain); + grad_inversion_tangent = (grad_inversion_tangent + inversion_gain_tangent); + auto inverted_bar_value = ((-atom_inversion) * long_bar_value); + auto inverted_bar_tangent = (((-atom_inversion) * long_bar_tangent) - (atom_dot_inversion * long_bar_value)); + long_bar_value = bsk::where(invert, inverted_bar_value, long_bar_value); + long_bar_tangent = bsk::where(invert, inverted_bar_tangent, long_bar_tangent); + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_dot_flip = _event_value(dot_flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + pulse_dot_b1 = atom_dot_b1; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse the row of the shim it drives. + if (bsk::truth(shimmed)) { + shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + pulse_dot_b1 = bsk::ld(((dot_b1 + shim_row) + atom), active_atom, 0.0f); + } + alpha_value = (event_flip * pulse_b1); + alpha_tangent = ((event_dot_flip * pulse_b1) + (event_flip * pulse_dot_b1)); + cosine_value = bsk::cos(alpha_value); + sine_value = bsk::sin(alpha_value); + cosine_tangent = ((-sine_value) * alpha_tangent); + sine_tangent = (cosine_value * alpha_tangent); + chs_value = (0.5f * (1.0f + cosine_value)); + chs_tangent = (0.5f * cosine_tangent); + shs_value = (0.5f * (1.0f - cosine_value)); + shs_tangent = (-0.5f * cosine_tangent); + half_sine_value = (0.5f * sine_value); + half_sine_tangent = (0.5f * sine_tangent); + // d/dalpha of each output row, contracted with the adjoint. + row_p_value = ((half_sine_value * stage_mv) - (half_sine_value * stage_pv)); + row_p_value = (row_p_value - (cosine_value * stage_zv)); + row_p_tangent = ((half_sine_value * stage_mt) + (half_sine_tangent * stage_mv)); + row_p_tangent = (row_p_tangent - ((half_sine_value * stage_pt) + (half_sine_tangent * stage_pv))); + row_p_tangent = (row_p_tangent - ((cosine_value * stage_zt) + (cosine_tangent * stage_zv))); + row_m_value = ((half_sine_value * stage_pv) - (half_sine_value * stage_mv)); + row_m_value = (row_m_value + (cosine_value * stage_zv)); + row_m_tangent = ((half_sine_value * stage_pt) + (half_sine_tangent * stage_pv)); + row_m_tangent = (row_m_tangent - ((half_sine_value * stage_mt) + (half_sine_tangent * stage_mv))); + row_m_tangent = (row_m_tangent + ((cosine_value * stage_zt) + (cosine_tangent * stage_zv))); + row_z_value = (((0.5f * cosine_value) * stage_pv) - ((0.5f * cosine_value) * stage_mv)); + row_z_value = (row_z_value - (sine_value * stage_zv)); + row_z_tangent = (0.5f * ((cosine_value * stage_pt) + (cosine_tangent * stage_pv))); + row_z_tangent = (row_z_tangent - (0.5f * ((cosine_value * stage_mt) + (cosine_tangent * stage_mv)))); + row_z_tangent = (row_z_tangent - ((sine_value * stage_zt) + (sine_tangent * stage_zv))); + alpha_bar_terms_value = (plus_bar_value * row_p_value); + alpha_bar_terms_value = (alpha_bar_terms_value + (minus_bar_value * row_m_value)); + alpha_bar_terms_value = (alpha_bar_terms_value + (long_bar_value * row_z_value)); + alpha_bar_terms_tangent = (plus_bar_value * row_p_tangent); + alpha_bar_terms_tangent = (alpha_bar_terms_tangent + (plus_bar_tangent * row_p_value)); + alpha_bar_terms_tangent = (alpha_bar_terms_tangent + (minus_bar_value * row_m_tangent)); + alpha_bar_terms_tangent = (alpha_bar_terms_tangent + (minus_bar_tangent * row_m_value)); + alpha_bar_terms_tangent = (alpha_bar_terms_tangent + (long_bar_value * row_z_tangent)); + alpha_bar_terms_tangent = (alpha_bar_terms_tangent + (long_bar_tangent * row_z_value)); + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + auto grad_alpha_value = bsk::sum_x(bsk::where(rotate, alpha_bar_terms_value, 0.0f)); + auto grad_alpha_tangent = bsk::sum_x(bsk::where(rotate, alpha_bar_terms_tangent, 0.0f)); + // Transpose of the rotation. + rotated_pbv = ((chs_value * plus_bar_value) + (shs_value * minus_bar_value)); + rotated_pbv = (rotated_pbv + (half_sine_value * long_bar_value)); + rotated_pbt = ((chs_value * plus_bar_tangent) + (chs_tangent * plus_bar_value)); + rotated_pbt = (rotated_pbt + ((shs_value * minus_bar_tangent) + (shs_tangent * minus_bar_value))); + rotated_pbt = (rotated_pbt + (half_sine_value * long_bar_tangent)); + rotated_pbt = (rotated_pbt + (half_sine_tangent * long_bar_value)); + rotated_mbv = ((shs_value * plus_bar_value) + (chs_value * minus_bar_value)); + rotated_mbv = (rotated_mbv - (half_sine_value * long_bar_value)); + rotated_mbt = ((shs_value * plus_bar_tangent) + (shs_tangent * plus_bar_value)); + rotated_mbt = (rotated_mbt + ((chs_value * minus_bar_tangent) + (chs_tangent * minus_bar_value))); + rotated_mbt = (rotated_mbt - (half_sine_value * long_bar_tangent)); + rotated_mbt = (rotated_mbt - (half_sine_tangent * long_bar_value)); + rotated_zbv = (((-sine_value) * plus_bar_value) + (sine_value * minus_bar_value)); + rotated_zbv = (rotated_zbv + (cosine_value * long_bar_value)); + rotated_zbt = (((-sine_value) * plus_bar_tangent) - (sine_tangent * plus_bar_value)); + rotated_zbt = (rotated_zbt + ((sine_value * minus_bar_tangent) + (sine_tangent * minus_bar_value))); + rotated_zbt = (rotated_zbt + ((cosine_value * long_bar_tangent) + (cosine_tangent * long_bar_value))); + plus_bar_value = bsk::where(rotate, rotated_pbv, plus_bar_value); + plus_bar_tangent = bsk::where(rotate, rotated_pbt, plus_bar_tangent); + minus_bar_value = bsk::where(rotate, rotated_mbv, minus_bar_value); + minus_bar_tangent = bsk::where(rotate, rotated_mbt, minus_bar_tangent); + long_bar_value = bsk::where(rotate, rotated_zbv, long_bar_value); + long_bar_tangent = bsk::where(rotate, rotated_zbt, long_bar_tangent); + auto flip_gain_value = (grad_alpha_value * pulse_b1); + flip_gain_tangent = (grad_alpha_tangent * pulse_b1); + flip_gain_tangent = (flip_gain_tangent + (grad_alpha_value * pulse_dot_b1)); + auto writes_flip = bsk::band(active_atom, rotate); + bsk::atomic_add(((grad_flip_value + event_base) + event), flip_gain_value, writes_flip); + bsk::atomic_add(((grad_flip_tangent + event_base) + event), flip_gain_tangent, writes_flip); + if (bsk::truth(shimmed)) { + // A pulse's transmit gradient belongs to the shim it drives, so + // with several it lands in that shim's row rather than in a + // register summed over the whole train. + bsk::atomic_add((((grad_tissue_value + (3 * atom_count)) + shim_row) + atom), (grad_alpha_value * event_flip), writes_flip); + bsk::atomic_add((((grad_tissue_tangent + (3 * atom_count)) + shim_row) + atom), ((grad_alpha_tangent * event_flip) + (grad_alpha_value * event_dot_flip)), writes_flip); + } else { + grad_b1_value = (grad_b1_value + bsk::where(rotate, (grad_alpha_value * event_flip), 0.0f)); + grad_b1_tangent = (grad_b1_tangent + bsk::where(rotate, ((grad_alpha_tangent * event_flip) + (grad_alpha_value * event_dot_flip)), 0.0f)); + } + // The sample is i * m0 * plus[0]; only the imaginary seed acts. + auto record = bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)); + auto out_ = bsk::ld((output_index + event)); + auto seed = bsk::ld(((grad_output_imag + (problem * output_count)) + out_), bsk::band(bsk::band(active_atom, record), (out_ >= 0)), 0.0f); + grad_m0_value = (grad_m0_value + bsk::sum_x(bsk::where((state == 0), (seed * stage_pv), 0.0f))); + grad_m0_tangent = (grad_m0_tangent + bsk::sum_x(bsk::where((state == 0), (seed * stage_pt), 0.0f))); + plus_bar_value = (plus_bar_value + bsk::where((state == 0), (seed * atom_m0), 0.0f)); + plus_bar_tangent = (plus_bar_tangent + bsk::where((state == 0), (seed * atom_dot_m0), 0.0f)); + auto t16_ = _shift_real_adjoint(plus_bar_value, minus_bar_value, state, state_mask, state_count); + adjoint_pv = bsk::get<0>(t16_); + adjoint_mv = bsk::get<1>(t16_); + auto t17_ = _shift_real_adjoint(plus_bar_tangent, minus_bar_tangent, state, state_mask, state_count); + adjoint_pt = bsk::get<0>(t17_); + adjoint_mt = bsk::get<1>(t17_); + plus_bar_value = bsk::where(pre_shift, adjoint_pv, plus_bar_value); + minus_bar_value = bsk::where(pre_shift, adjoint_mv, minus_bar_value); + plus_bar_tangent = bsk::where(pre_shift, adjoint_pt, plus_bar_tangent); + minus_bar_tangent = bsk::where(pre_shift, adjoint_mt, minus_bar_tangent); + auto cot2_value = ((plus_bar_value * entry_pv) + (minus_bar_value * entry_mv)); + auto cot2_tangent = ((((plus_bar_value * entry_pt) + (plus_bar_tangent * entry_pv)) + (minus_bar_value * entry_mt)) + (minus_bar_tangent * entry_mv)); + auto cot1_value = (long_bar_value * entry_zv); + auto cot1_tangent = ((long_bar_value * entry_zt) + (long_bar_tangent * entry_zv)); + auto grad_e2_value = bsk::sum_x((cot2_value * damp_t)); + auto grad_e2_tangent = bsk::sum_x(((cot2_value * damp_t_tangent) + (cot2_tangent * damp_t))); + grad_e1_value = bsk::sum_x((cot1_value * damp_z)); + grad_e1_value = (grad_e1_value - bsk::sum_x(bsk::where((state == 0), long_bar_value, 0.0f))); + grad_e1_tangent = bsk::sum_x(((cot1_value * damp_z_tangent) + (cot1_tangent * damp_z))); + grad_e1_tangent = (grad_e1_tangent - bsk::sum_x(bsk::where((state == 0), long_bar_tangent, 0.0f))); + // The rate and the interval multiply every order's b-weight, so both + // take a weighted sum. Order zero has no longitudinal weight, which + // keeps recovery out of this. + spread_value = zero; + spread_tangent = zero; + if (bsk::truth(diffusing)) { + auto weighted_value = ((((cot1_value * bare1_value) * damp_z) * longitudinal_weight) + (((cot2_value * bare2_value) * damp_t) * transverse_weight)); + auto weighted_tangent = ((((((cot1_tangent * bare1_value) * damp_z) + ((cot1_value * bare1_tangent) * damp_z)) + ((cot1_value * bare1_value) * damp_z_tangent)) * longitudinal_weight) + (((((cot2_tangent * bare2_value) * damp_t) + ((cot2_value * bare2_tangent) * damp_t)) + ((cot2_value * bare2_value) * damp_t_tangent)) * transverse_weight)); + spread_value = bsk::sum_x(weighted_value); + spread_tangent = bsk::sum_x(weighted_tangent); + grad_damping_value = (grad_damping_value + ((-spread_value) * dt_value)); + grad_damping_tangent = (grad_damping_tangent + (-((spread_value * dt_tangent) + (spread_tangent * dt_value)))); + } + plus_bar_tangent = ((plus_bar_value * e2_tangent) + (plus_bar_tangent * e2_value)); + plus_bar_value = (plus_bar_value * e2_value); + minus_bar_tangent = ((minus_bar_value * e2_tangent) + (minus_bar_tangent * e2_value)); + minus_bar_value = (minus_bar_value * e2_value); + long_bar_tangent = ((long_bar_value * e1_tangent) + (long_bar_tangent * e1_value)); + long_bar_value = (long_bar_value * e1_value); + auto inverse1_value = bsk::truediv(1000.0f, (atom_t1 * atom_t1)); + auto inverse1_tangent = bsk::truediv((-2000.0f * atom_dot_t1), ((atom_t1 * atom_t1) * atom_t1)); + auto inverse2_value = bsk::truediv(1000.0f, (atom_t2 * atom_t2)); + auto inverse2_tangent = bsk::truediv((-2000.0f * atom_dot_t2), ((atom_t2 * atom_t2) * atom_t2)); + auto scale1_value = ((bare1_value * dt_value) * inverse1_value); + scale1_tangent = ((bare1_tangent * dt_value) * inverse1_value); + scale1_tangent = (scale1_tangent + ((bare1_value * dt_tangent) * inverse1_value)); + scale1_tangent = (scale1_tangent + ((bare1_value * dt_value) * inverse1_tangent)); + auto scale2_value = ((bare2_value * dt_value) * inverse2_value); + scale2_tangent = ((bare2_tangent * dt_value) * inverse2_value); + scale2_tangent = (scale2_tangent + ((bare2_value * dt_tangent) * inverse2_value)); + scale2_tangent = (scale2_tangent + ((bare2_value * dt_value) * inverse2_tangent)); + grad_t1_value = (grad_t1_value + (grad_e1_value * scale1_value)); + grad_t1_tangent = (grad_t1_tangent + (grad_e1_value * scale1_tangent)); + grad_t1_tangent = (grad_t1_tangent + (grad_e1_tangent * scale1_value)); + grad_t2_value = (grad_t2_value + (grad_e2_value * scale2_value)); + grad_t2_tangent = (grad_t2_tangent + (grad_e2_value * scale2_tangent)); + grad_t2_tangent = (grad_t2_tangent + (grad_e2_tangent * scale2_value)); + auto decay1_value = (rate1_value * bare1_value); + auto decay1_tangent = ((rate1_value * bare1_tangent) + (rate1_tangent * bare1_value)); + auto decay2_value = (rate2_value * bare2_value); + auto decay2_tangent = ((rate2_value * bare2_tangent) + (rate2_tangent * bare2_value)); + duration_gain_value = ((-grad_e1_value) * decay1_value); + duration_gain_value = (duration_gain_value - (grad_e2_value * decay2_value)); + duration_gain_tangent = (-((grad_e1_value * decay1_tangent) + (grad_e1_tangent * decay1_value))); + duration_gain_tangent = (duration_gain_tangent - ((grad_e2_value * decay2_tangent) + (grad_e2_tangent * decay2_value))); + duration_gain_value = (duration_gain_value + ((-spread_value) * atom_damping)); + duration_gain_tangent = (duration_gain_tangent + (-((spread_value * atom_dot_damping) + (spread_tangent * atom_damping)))); + bsk::atomic_add(((grad_duration_value + event_base) + event), duration_gain_value, active_atom); + bsk::atomic_add(((grad_duration_tangent + event_base) + event), duration_gain_tangent, active_atom); + } + bsk::atomic_add((grad_tissue_value + atom), grad_t1_value, active_atom); + bsk::atomic_add((grad_tissue_tangent + atom), grad_t1_tangent, active_atom); + bsk::atomic_add(((grad_tissue_value + atom_count) + atom), grad_t2_value, active_atom); + bsk::atomic_add(((grad_tissue_tangent + atom_count) + atom), grad_t2_tangent, active_atom); + bsk::atomic_add(((grad_tissue_value + (2 * atom_count)) + atom), grad_m0_value, active_atom); + bsk::atomic_add(((grad_tissue_tangent + (2 * atom_count)) + atom), grad_m0_tangent, active_atom); + if (bsk::truth((!bsk::truth(shimmed)))) { + bsk::atomic_add(((grad_tissue_value + (3 * atom_count)) + atom), grad_b1_value, active_atom); + bsk::atomic_add(((grad_tissue_tangent + (3 * atom_count)) + atom), grad_b1_tangent, active_atom); + } + // The transmit pair takes a row per shim each in the plane the complex + // path allocates, so the rows past it move even though this kernel leaves + // the transmit phase at zero throughout. + auto past_transmit = (2 * (shim_rows - 1)); + bsk::atomic_add(((grad_tissue_value + ((6 + past_transmit) * atom_count)) + atom), grad_inversion_value, active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((6 + past_transmit) * atom_count)) + atom), grad_inversion_tangent, active_atom); + bsk::atomic_add(((grad_tissue_value + ((7 + past_transmit) * atom_count)) + atom), grad_damping_value, active_atom); + bsk::atomic_add(((grad_tissue_tangent + ((7 + past_transmit) * atom_count)) + atom), grad_damping_tangent, active_atom); +} + +BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V alpha{}; + bsk::V atom_b1{}; + bsk::V atom_damping{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bsk::V dt{}; + bsk::V e1{}; + bsk::V e2{}; + bsk::V last_dt{}; + bsk::V longitudinal{}; + bsk::V minus{}; + bsk::V plus{}; + bsk::V pulse_b1{}; + bsk::V relaxes{}; + auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < (train_count * atom_count)); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + auto empty = bsk::full(0); + plus = empty; + minus = empty; + longitudinal = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + auto rate1 = bsk::truediv(1000.0f, atom_t1); + auto rate2 = bsk::truediv(1000.0f, atom_t2); + atom_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + } + auto order = bsk::cast(state); + // The relaxation factors depend on the event only through its duration, and + // a train repeats its intervals: an interval as long as the last one reuses + // the factors rather than taking the two exponentials again. Where several + // trains share the program the durations differ across its lanes and there + // is nothing uniform to compare, so only a single-train launch memoizes. + last_dt = -1.0f; + e1 = ((rate1 * 0.0f) + 1.0f); + e2 = ((rate2 * 0.0f) + 1.0f); + auto event_base = (train * event_count); + // Two events to an iteration. A repetition is several events -- a pulse, + // a sample, an interval -- so the loop runs longer than the sequence is + // repetitions, and unrolling lets one back-edge and one set of event + // bookkeeping serve two of them. Two is where it stops paying: four was + // measured slower, and the body is already large enough that widening it + // costs registers. + for (std::int64_t event = 0; event < event_count; event += 1) { + // Read here rather than through the helper: one train gives a duration + // the whole program shares, and the skip and the memo below both want + // to compare it as the single number it is. + if (bsk::truth(single_train)) { + dt = bsk::ld((duration + event)); + } else { + dt = bsk::ld(((duration + event_base) + event), active_atom, 0.0f); + } + // An event of no duration relaxes nothing: both factors are one and the + // recovery term is zero. Half the events of a spoiled repetition are + // instantaneous, and reducing over the trains this program carries makes + // that a branch the whole program agrees on rather than a tile of + // multiplies by one. + if (bsk::truth(single_train)) { + relaxes = (dt != 0.0f); + } else { + relaxes = (bsk::max_all(dt) != 0.0f); + } + if (bsk::truth(relaxes)) { + if (bsk::truth(single_train)) { + if (bsk::truth((dt != last_dt))) { + e1 = bsk::exp(((-rate1) * dt)); + e2 = bsk::exp(((-rate2) * dt)); + last_dt = dt; + } + } else { + e1 = bsk::exp(((-rate1) * dt)); + e2 = bsk::exp(((-rate2) * dt)); + } + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping(atom_damping, dt, order); + damp_z = bsk::get<0>(t0_); + damp_t = bsk::get<1>(t0_); + } + auto recovery = (1.0f - e1); + plus = (plus * (e2 * damp_t)); + minus = (minus * (e2 * damp_t)); + longitudinal = ((longitudinal * (e1 * damp_z)) + bsk::where((state == 0), recovery, 0.0f)); + } + // Every flag below is read from a per-event array with no atom index, so + // it is uniform across the program and can steer real control flow. A + // `tl.where` would make every event pay for every operator: a spoiled + // repetition is four events and needs one rotation and one shift. + auto event_action = bsk::cast(bsk::ld((action + event))); + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t1_ = _shift_real(plus, minus, state, state_mask, state_count); + plus = bsk::get<0>(t1_); + minus = bsk::get<1>(t1_); + } + auto event_kind = bsk::ld((kind + event)); + auto is_rf = (event_kind == 1); + auto is_inversion = (bsk::band(event_action, 4) != 0); + if (bsk::truth((bsk::truth(is_rf) && bsk::truth(is_inversion)))) { + longitudinal = ((-atom_inversion) * longitudinal); + } else if (bsk::truth(is_rf)) { + alpha = _event_value(flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + // One shim is the whole sequence's transmit field, loaded once + // above; several give each pulse the row of the shim it drives. + if (bsk::truth((bsk::truth(shimmed) && bsk::truth(transmit)))) { + auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + alpha = (alpha * pulse_b1); + auto cosine = bsk::cos(alpha); + auto sine = bsk::sin(alpha); + auto cosine_half_sq = (0.5f * (1.0f + cosine)); + auto sine_half_sq = (0.5f * (1.0f - cosine)); + auto half_sine = (0.5f * sine); + auto rotated_p = (((cosine_half_sq * plus) + (sine_half_sq * minus)) - (sine * longitudinal)); + auto rotated_m = (((sine_half_sq * plus) + (cosine_half_sq * minus)) + (sine * longitudinal)); + longitudinal = (((half_sine * plus) - (half_sine * minus)) + (cosine * longitudinal)); + plus = rotated_p; + minus = rotated_m; + } + if (bsk::truth((bsk::truth((bsk::band(event_action, 32) != 0)) && bsk::truth((event_kind == 2))))) { + auto out_ = bsk::ld((output_index + event)); + auto output_offset = ((problem * output_count) + out_); + auto output_mask = bsk::band(bsk::band(active_atom, (state == 0)), (out_ >= 0)); + bsk::st(((output_real + output_offset) + state), empty, output_mask); + bsk::st(((output_imag + output_offset) + state), (atom_m0 * plus), output_mask); + } + if (bsk::truth((bsk::truth((bsk::band(event_action, 2) != 0)) || bsk::truth((bsk::band(event_action, 16) != 0))))) { + auto t2_ = _shift_real(plus, minus, state, state_mask, state_count); + plus = bsk::get<0>(t2_); + minus = bsk::get<1>(t2_); + } + if (bsk::truth((bsk::band(event_action, 8) != 0))) { + plus = empty; + minus = empty; + } + } +} + +// Seven of the nine rotation coefficients; the rest follow by symmetry. +// +// ``t11`` repeats ``t00`` and ``t10`` is the conjugate of ``t01``, so the +// caller derives those. Feeding ``(cos, sin)`` gives the rotation itself and +// ``(sin, cos)`` rearranged gives its derivative in the flip angle, which is +// why this is one routine rather than two. +template +BSK_HD auto _rotation_coefficients(const T0& a, const T1& b, const T2& c, const T3& d, const T4& p1r, const T5& p1i, const T6& p2r, const T7& p2i, const T8& pcr, const T9& pci) { + auto t00 = bsk::make_tup(a, (0.0f * a)); + auto t01 = bsk::make_tup((b * p2r), (b * p2i)); + auto t02 = _complex_mul((0.0f * c), (-c), p1r, p1i); + auto t12 = _complex_mul((0.0f * c), c, pcr, pci); + auto t20 = _complex_mul((0.0f * c), (-0.5f * c), pcr, pci); + auto t21 = _complex_mul((0.0f * c), (0.5f * c), p1r, p1i); + auto t22 = bsk::make_tup(d, (0.0f * d)); + return bsk::make_tup(t00, t01, t02, t12, t20, t21, t22); +} + +// The spinor rotation's adjoint, carrying no forward direction. +// +// Returns the cotangent on the Cayley-Klein pair and the three state +// cotangents sent back through the conjugate transpose. Every entry of the +// matrix is a product of two factors drawn from the pair and its conjugate, +// so the pair's two Wirtinger halves are linear in the outer product of the +// seed with the state the rotation acted on -- a closed form rather than a +// differentiated matrix. +template +BSK_HD auto _spinor_adjoint(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& spr, const T5& spi, const T6& smr, const T7& smi, const T8& rzr, const T9& rzi, const T10& pbr, const T11& pbi, const T12& mbr, const T13& mbi, const T14& zbr, const T15& zbi) { + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n0{}; + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n1{}; + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n2{}; + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto t0_ = bsk::make_tup(aa_r, (-aa_i)); + auto t00_r = bsk::get<0>(t0_); + auto t00_i = bsk::get<1>(t0_); + auto t1_ = bsk::make_tup((-bb_r), bb_i); + auto t01_r = bsk::get<0>(t1_); + auto t01_i = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i)); + auto t02_r = bsk::get<0>(t2_); + auto t02_i = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup((-bb_r), (-bb_i)); + auto t10_r = bsk::get<0>(t3_); + auto t10_i = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(aa_r, aa_i); + auto t11_r = bsk::get<0>(t4_); + auto t11_i = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i)); + auto t12_r = bsk::get<0>(t5_); + auto t12_i = bsk::get<1>(t5_); + auto t6_ = bsk::make_tup(cross_r, cross_i); + auto t20_r = bsk::get<0>(t6_); + auto t20_i = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(cross_r, (-cross_i)); + auto t21_r = bsk::get<0>(t7_); + auto t21_i = bsk::get<1>(t7_); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + // ``m[i][j] = conj(seed_i) * state_j``: the outer product the pair's + // derivative is linear in. + auto m00 = _complex_mul(pbr, (-pbi), spr, spi); + auto m01 = _complex_mul(pbr, (-pbi), smr, smi); + auto m02 = _complex_mul(pbr, (-pbi), rzr, rzi); + auto m10 = _complex_mul(mbr, (-mbi), spr, spi); + auto m11 = _complex_mul(mbr, (-mbi), smr, smi); + auto m12 = _complex_mul(mbr, (-mbi), rzr, rzi); + auto m20 = _complex_mul(zbr, (-zbi), spr, spi); + auto m21 = _complex_mul(zbr, (-zbi), smr, smi); + auto m22 = _complex_mul(zbr, (-zbi), rzr, rzi); + auto hca = _complex_mul(ar, ai, bsk::get<0>(m11), bsk::get<1>(m11)); + auto hcb = _complex_mul(br, bi, bsk::get<0>(m12), bsk::get<1>(m12)); + auto hcc = _complex_mul(br, (-bi), bsk::get<0>(m21), bsk::get<1>(m21)); + auto hcd = _complex_mul(ar, (-ai), bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_conj_a_r = ((((2.0f * bsk::get<0>(hca)) - (2.0f * bsk::get<0>(hcb))) + bsk::get<0>(hcc)) + bsk::get<0>(hcd)); + auto holding_conj_a_i = ((((2.0f * bsk::get<1>(hca)) - (2.0f * bsk::get<1>(hcb))) + bsk::get<1>(hcc)) + bsk::get<1>(hcd)); + auto ha = _complex_mul(ar, (-ai), bsk::get<0>(m00), bsk::get<1>(m00)); + auto hb = _complex_mul(br, (-bi), bsk::get<0>(m02), bsk::get<1>(m02)); + auto hc = _complex_mul(br, bi, bsk::get<0>(m20), bsk::get<1>(m20)); + auto hd = _complex_mul(ar, ai, bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_a_r = ((((2.0f * bsk::get<0>(ha)) - (2.0f * bsk::get<0>(hb))) + bsk::get<0>(hc)) + bsk::get<0>(hd)); + auto holding_a_i = ((((2.0f * bsk::get<1>(ha)) - (2.0f * bsk::get<1>(hb))) + bsk::get<1>(hc)) + bsk::get<1>(hd)); + auto ka = _complex_mul(br, bi, bsk::get<0>(m10), bsk::get<1>(m10)); + auto kb = _complex_mul(ar, ai, bsk::get<0>(m12), bsk::get<1>(m12)); + auto kc = _complex_mul(ar, (-ai), bsk::get<0>(m20), bsk::get<1>(m20)); + auto kd = _complex_mul(br, (-bi), bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_conj_b_r = ((((-2.0f * bsk::get<0>(ka)) - (2.0f * bsk::get<0>(kb))) + bsk::get<0>(kc)) - bsk::get<0>(kd)); + auto holding_conj_b_i = ((((-2.0f * bsk::get<1>(ka)) - (2.0f * bsk::get<1>(kb))) + bsk::get<1>(kc)) - bsk::get<1>(kd)); + auto la = _complex_mul(br, (-bi), bsk::get<0>(m01), bsk::get<1>(m01)); + auto lb = _complex_mul(ar, (-ai), bsk::get<0>(m02), bsk::get<1>(m02)); + auto lc = _complex_mul(ar, ai, bsk::get<0>(m21), bsk::get<1>(m21)); + auto ld_ = _complex_mul(br, bi, bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_b_r = ((((-2.0f * bsk::get<0>(la)) - (2.0f * bsk::get<0>(lb))) + bsk::get<0>(lc)) - bsk::get<0>(ld_)); + auto holding_b_i = ((((-2.0f * bsk::get<1>(la)) - (2.0f * bsk::get<1>(lb))) + bsk::get<1>(lc)) - bsk::get<1>(ld_)); + auto grad_a_r = (holding_conj_a_r + holding_a_r); + auto grad_a_i = ((-holding_conj_a_i) + holding_a_i); + auto grad_b_r = (holding_conj_b_r + holding_b_r); + auto grad_b_i = ((-holding_conj_b_i) + holding_b_i); + n0 = _complex_mul(t00_r, (-t00_i), pbr, pbi); + n1 = _complex_mul(t10_r, (-t10_i), mbr, mbi); + n2 = _complex_mul(t20_r, (-t20_i), zbr, zbi); + auto t8_ = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2))); + auto next_pr = bsk::get<0>(t8_); + auto next_pi = bsk::get<1>(t8_); + n0 = _complex_mul(t01_r, (-t01_i), pbr, pbi); + n1 = _complex_mul(t11_r, (-t11_i), mbr, mbi); + n2 = _complex_mul(t21_r, (-t21_i), zbr, zbi); + auto t9_ = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2))); + auto next_mr = bsk::get<0>(t9_); + auto next_mi = bsk::get<1>(t9_); + n0 = _complex_mul(t02_r, (-t02_i), pbr, pbi); + n1 = _complex_mul(t12_r, (-t12_i), mbr, mbi); + auto next_zr = ((bsk::get<0>(n0) + bsk::get<0>(n1)) + (t22 * zbr)); + auto next_zi = ((bsk::get<1>(n0) + bsk::get<1>(n1)) + (t22 * zbi)); + return bsk::make_tup(grad_a_r, grad_a_i, grad_b_r, grad_b_i, next_pr, next_pi, next_mr, next_mi, next_zr, next_zi); +} + +// Send the cotangent on one pulse's rotation to its row. +// +// Summed over the dephasing orders first: the pair multiplies every one of +// them, so what reaches the row is the sum. The block is padded to a power of +// two and the orders past the last carry whatever the sweep left there, so +// the sum is taken over the orders that exist rather than over the block. +template +BSK_HD auto _store_pair_gradient(const T0& grad_pair, const T1& pair_index, const T2& event_base, const T3& event, const T4& atom, const T5& atom_count, const T6& turning, const T7& mask, const T8& state_mask, const T9& grad_ar, const T10& grad_ai, const T11& grad_br, const T12& grad_bi) { + auto row = bsk::cast(bsk::ld(((pair_index + event_base) + event))); + auto entry = (((row * atom_count) + atom) * 4); + auto keep = bsk::band(turning, state_mask); + bsk::atomic_add(((grad_pair + entry) + 0), bsk::sum_x(bsk::where(keep, grad_ar, 0.0f)), mask); + bsk::atomic_add(((grad_pair + entry) + 1), bsk::sum_x(bsk::where(keep, grad_ai, 0.0f)), mask); + bsk::atomic_add(((grad_pair + entry) + 2), bsk::sum_x(bsk::where(keep, grad_br, 0.0f)), mask); + bsk::atomic_add(((grad_pair + entry) + 3), bsk::sum_x(bsk::where(keep, grad_bi, 0.0f)), mask); +} + +// What one interval's cotangents give its length and its attenuation. +// +// The generator is proportional to the interval, so ``dE/d(dt) == A1 E`` and +// the length's gradient needs the operator and 27 multiplies rather than the +// eigenvalues -- which is what lets every other gradient be pooled over the +// events that share a length while this one stays per event. +// +// Returns the two contractions, in the order +// :func:`_three_pool_step_adjoint_jvp` returns them. +template +BSK_HD auto _three_pool_interval_adjoint(const T0& table, const T1& row, const T2& atom, const T3& voxel_count, const T4& mask, const T5& r1_free, const T6& r1_pool_b, const T7& r1_bound, const T8& exchange_b, const T9& exchange_c, const T10& fraction_b, const T11& fraction_c, const T12& attenuation, const T13& b11, const T14& b12, const T15& b13, const T16& b21, const T17& b22, const T18& b23, const T19& b31, const T20& b32, const T21& b33, const T22& bfree, const T23& bpool_b, const T24& bbound) { + auto free = ((1.0f - fraction_b) - fraction_c); + auto a00 = ((((-exchange_b) * fraction_b) - (exchange_c * fraction_c)) - r1_free); + auto a01 = (exchange_b * free); + auto a02 = (exchange_c * free); + auto a10 = (exchange_b * fraction_b); + auto a11 = (((-exchange_b) * free) - r1_pool_b); + auto a20 = (exchange_c * fraction_c); + auto a22 = (((-exchange_c) * free) - r1_bound); + auto base = ((table + (row * (9 * voxel_count))) + atom); + auto c00 = bsk::ld((base + (0 * voxel_count)), mask, 0.0f); + auto c01 = bsk::ld((base + (1 * voxel_count)), mask, 0.0f); + auto c02 = bsk::ld((base + (2 * voxel_count)), mask, 0.0f); + auto c10 = bsk::ld((base + (3 * voxel_count)), mask, 0.0f); + auto c11 = bsk::ld((base + (4 * voxel_count)), mask, 0.0f); + auto c12 = bsk::ld((base + (5 * voxel_count)), mask, 0.0f); + auto c20 = bsk::ld((base + (6 * voxel_count)), mask, 0.0f); + auto c21 = bsk::ld((base + (7 * voxel_count)), mask, 0.0f); + auto c22 = bsk::ld((base + (8 * voxel_count)), mask, 0.0f); + // A1 C, the second and third rows of A1 having no entry off their own pool. + auto p00 = (((a00 * c00) + (a01 * c10)) + (a02 * c20)); + auto p01 = (((a00 * c01) + (a01 * c11)) + (a02 * c21)); + auto p02 = (((a00 * c02) + (a01 * c12)) + (a02 * c22)); + auto p10 = ((a10 * c00) + (a11 * c10)); + auto p11 = ((a10 * c01) + (a11 * c11)); + auto p12 = ((a10 * c02) + (a11 * c12)); + auto p20 = ((a20 * c00) + (a22 * c20)); + auto p21 = ((a20 * c01) + (a22 * c21)); + auto p22 = ((a20 * c02) + (a22 * c22)); + auto grad_dt = (attenuation * _three_pool_contract(p00, p01, p02, p10, p11, p12, p20, p21, p22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c)); + auto grad_att = _three_pool_contract(c00, c01, c02, c10, c11, c12, c20, c21, c22, b11, b12, b13, b21, b22, b23, b31, b32, b33, bfree, bpool_b, bbound, free, fraction_b, fraction_c); + return bsk::make_tup(grad_dt, grad_att); +} + +BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* lineshape, float* profile, std::int32_t* profile_index, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, std::int64_t row_count, float* grad_pair, float* grad_output_real, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_phase, float* grad_duration, float* trajectory_r, float* trajectory_i, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, std::int64_t shim_rows, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t recording, std::int64_t block_states, std::int64_t problems) { + bsk::V _d11{}; + bsk::V _d12{}; + bsk::V _d21{}; + bsk::V _d22{}; + bsk::V _dgb{}; + bsk::V _dgf{}; + bsk::V _dgs{}; + bsk::V _drb{}; + bsk::V _drf{}; + bsk::V _dw11{}; + bsk::V _dw12{}; + bsk::V _dw13{}; + bsk::V _dw21{}; + bsk::V _dw22{}; + bsk::V _dw23{}; + bsk::V _dw31{}; + bsk::V _dw32{}; + bsk::V _dw33{}; + bsk::V _q1{}; + bsk::V _q2{}; + bsk::V _q3{}; + bsk::V _q4{}; + bsk::V _q5{}; + bsk::V _q6{}; + bsk::V _q7{}; + bsk::V _q8{}; + bsk::V _q9{}; + bsk::V a11i{}; + bsk::V a11r{}; + bsk::V a12i{}; + bsk::V a12r{}; + bsk::V a21i{}; + bsk::V a21r{}; + bsk::V a22i{}; + bsk::V a22r{}; + bsk::V absorbed_value{}; + bsk::tup, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V> across{}; + bsk::tup, bsk::V> add1{}; + bsk::tup, bsk::V> add2{}; + bsk::V alpha_v{}; + bsk::V alpha_value{}; + bsk::V angle_value{}; + bsk::V atom_b0{}; + bsk::V atom_b1{}; + bsk::V atom_b1_phase{}; + bsk::V atom_bound{}; + bsk::V atom_damping{}; + bsk::V atom_exchange{}; + bsk::V atom_flow{}; + bsk::V atom_free{}; + bsk::V atom_inv{}; + bsk::V atom_m0{}; + bsk::V atom_semisolid{}; + bsk::V atom_semisolid_exchange{}; + bsk::V atom_shift{}; + bsk::V atom_t1b{}; + bsk::V atom_t1c{}; + bsk::V atom_t2b{}; + bsk::V atom_washout{}; + bsk::V attenuation_v{}; + bsk::V avi{}; + bsk::V avr{}; + bsk::V back_att{}; + bsk::V back_bound{}; + bsk::V back_dt{}; + bsk::V back_exch{}; + bsk::V back_i{}; + bsk::V back_mi{}; + bsk::V back_mr{}; + bsk::V back_pi{}; + bsk::V back_pr{}; + bsk::V back_r{}; + bsk::V back_r1{}; + bsk::V back_r1b{}; + bsk::V back_r1c{}; + bsk::V back_semi{}; + bsk::V back_sexch{}; + bsk::V back_zi{}; + bsk::V back_zr{}; + bsk::V bare1_value{}; + bsk::V bare2_value{}; + bsk::V bare_cot_v{}; + std::int32_t base_row{}; + bsk::V bmvi{}; + bsk::V bmvr{}; + bsk::V bpvi{}; + bsk::V bpvr{}; + bsk::V bvi{}; + bsk::V bvr{}; + bsk::V cari{}; + bsk::V carr{}; + bsk::V col_bi{}; + bsk::V col_br{}; + bsk::V col_ci{}; + bsk::V col_cr{}; + bsk::V col_fi{}; + bsk::V col_fr{}; + bsk::V cos_value{}; + bsk::V cot2_v{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bsk::V direction{}; + bool do_shift{}; + bsk::V dt_value{}; + bsk::V duration_v{}; + bsk::V e1_v{}; + bsk::V e1_value{}; + bsk::V e2_value{}; + std::int64_t event{}; + std::int32_t event_action{}; + bsk::V event_flip{}; + std::int32_t event_kind{}; + bsk::V event_phase{}; + float event_saturation{}; + bsk::V f11i{}; + bsk::V f11r{}; + bsk::V f12i{}; + bsk::V f12r{}; + bsk::V g21i{}; + bsk::V g21r{}; + bsk::V g22i{}; + bsk::V g22r{}; + bsk::V g_b0v{}; + bsk::V g_b1pv{}; + bsk::V g_b1v{}; + bsk::V g_boundv{}; + bsk::V g_diffv{}; + bsk::V g_exchv{}; + bsk::V g_flowv{}; + bsk::V g_invv{}; + bsk::V g_m0v{}; + bsk::V g_semiv{}; + bsk::V g_sexchv{}; + bsk::V g_shiftv{}; + bsk::V g_t1bv{}; + bsk::V g_t1cv{}; + bsk::V g_t1v{}; + bsk::V g_t2bv{}; + bsk::V g_t2v{}; + bsk::V g_washv{}; + bsk::V grad_alpha_v{}; + bsk::V grad_angle_v{}; + bsk::V grad_e1_v{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; + bsk::V h11i{}; + bsk::V h11r{}; + bsk::V h12i{}; + bsk::V h12r{}; + bsk::V held{}; + bsk::V hold_value{}; + bool invert{}; + bool is_inversion{}; + bool is_rf{}; + bsk::V k21i{}; + bsk::V k21r{}; + bsk::V k22i{}; + bsk::V k22r{}; + bsk::V long_damp_v{}; + bsk::V lvi{}; + bsk::V lvr{}; + bsk::V mbvi{}; + bsk::V mbvr{}; + bsk::V mix_bi{}; + bsk::V mix_br{}; + bsk::V mix_ci{}; + bsk::V mix_cr{}; + bsk::V mix_fi{}; + bsk::V mix_fr{}; + bsk::V mvi{}; + bsk::V mvr{}; + bsk::tup, bsk::V> n0{}; + bsk::tup, bsk::V> n1{}; + bsk::tup, bsk::V> n2{}; + bsk::V nil{}; + bsk::V offset_value{}; + bsk::V ovi{}; + bsk::V ovr{}; + bsk::V p1i{}; + bsk::V p1r{}; + bsk::V p2i{}; + bsk::V p2r{}; + bsk::tup, bsk::V, bsk::V, bsk::V> pair{}; + bsk::V part_i{}; + bsk::V part_r{}; + bsk::V pbvi{}; + bsk::V pbvr{}; + bsk::V pe11{}; + bsk::V pe12{}; + bsk::V pe21{}; + bsk::V pe22{}; + bsk::V per_angle_v{}; + bsk::V per_state{}; + bsk::V phi_v{}; + bsk::V phi_value{}; + bsk::V pool_angle_v{}; + bsk::V pool_back_mi{}; + bsk::V pool_back_mr{}; + bsk::V pool_back_pi{}; + bsk::V pool_back_pr{}; + bsk::V pool_back_zi{}; + bsk::V pool_back_zr{}; + bsk::V pool_row{}; + bsk::V pool_shaped_mbi{}; + bsk::V pool_shaped_mbr{}; + bsk::V pool_shaped_pbi{}; + bsk::V pool_shaped_pbr{}; + bsk::V pool_shaped_zbi{}; + bsk::V pool_shaped_zbr{}; + bsk::V poolbi{}; + bsk::V poolbr{}; + bsk::V poolvi{}; + bsk::V poolvr{}; + bsk::V power_value{}; + bool pre_shift{}; + bsk::V prec_b{}; + bsk::V prec_f{}; + bsk::V problem{}; + bsk::V pulse_b1{}; + bsk::V pulse_b1_phase{}; + bsk::V pvi{}; + bsk::V pvr{}; + bsk::tup, bsk::V> q0{}; + bsk::tup, bsk::V> q1{}; + bsk::tup, bsk::V> q2{}; + bsk::V qi{}; + bsk::V qr{}; + bsk::V r1b_value{}; + bsk::V r1c_value{}; + bsk::V r2b_value{}; + bsk::V rbmvi{}; + bsk::V rbmvr{}; + bsk::V rbpvi{}; + bsk::V rbpvr{}; + bsk::V rbvi{}; + bsk::V rbvr{}; + bsk::V rcvi{}; + bsk::V rcvr{}; + bsk::V reci{}; + bsk::V recovery_value{}; + bsk::V recr{}; + bsk::V rmvi{}; + bsk::V rmvr{}; + bool rotate{}; + std::int64_t row{}; + bsk::tup, bsk::V> row0{}; + bsk::V rpvi{}; + bsk::V rpvr{}; + bsk::V rzvi{}; + bsk::V rzvr{}; + bsk::V sat_alpha_v{}; + bsk::V sat_b0_v{}; + bool saturating{}; + bsk::V sbi{}; + bsk::V sbmvi{}; + bsk::V sbmvr{}; + bsk::V sbpvi{}; + bsk::V sbpvr{}; + bsk::V sbr{}; + bsk::V semibi{}; + bsk::V semibr{}; + bsk::V semivi{}; + bsk::V semivr{}; + bsk::V sfi{}; + bsk::V sfr{}; + bsk::V shape_value{}; + bsk::V shaped_ai{}; + bsk::V shaped_ar{}; + bsk::V shaped_bi{}; + bsk::V shaped_br{}; + bsk::V shaped_mbi{}; + bsk::V shaped_mbr{}; + bsk::V shaped_pbi{}; + bsk::V shaped_pbr{}; + bsk::V shaped_zbi{}; + bsk::V shaped_zbr{}; + bsk::V sin_value{}; + bsk::V slope_ai{}; + bsk::V slope_ar{}; + bsk::V slope_bi{}; + bsk::V slope_br{}; + bsk::V slot{}; + bsk::V spin_i{}; + bsk::V spin_r{}; + bool spoil{}; + bsk::V spread_v{}; + bsk::V spun_bi{}; + bsk::V spun_br{}; + bsk::V spun_fi{}; + bsk::V spun_fr{}; + bsk::V spun_mi{}; + bsk::V spun_mr{}; + bsk::V spun_pi{}; + bsk::V spun_pr{}; + bsk::V spun_zi{}; + bsk::V spun_zr{}; + bsk::V svi{}; + bsk::V svr{}; + bsk::V szi{}; + bsk::V szr{}; + bsk::tup, bsk::V> t00{}; + bsk::tup, bsk::V> t01{}; + bsk::tup, bsk::V> t02{}; + bsk::tup, bsk::V> t12{}; + bsk::tup, bsk::V> t20{}; + bsk::tup, bsk::V> t21{}; + bsk::tup, bsk::V> t22{}; + bsk::V three_a00{}; + bsk::V three_a01{}; + bsk::V three_a02{}; + bsk::V three_a10{}; + bsk::V three_a11{}; + bsk::V three_a20{}; + bsk::V three_a22{}; + bsk::V three_angle{}; + bsk::V three_argument{}; + bsk::V three_centre{}; + bsk::V three_cube{}; + bsk::V three_d_a00{}; + bsk::V three_d_a01{}; + bsk::V three_d_a02{}; + bsk::V three_d_a10{}; + bsk::V three_d_a11{}; + bsk::V three_d_a20{}; + bsk::V three_d_a22{}; + bsk::V three_d_angle{}; + bsk::V three_d_centre{}; + bsk::V three_d_determinant{}; + bsk::V three_d_first{}; + bsk::V three_d_free{}; + bsk::V three_d_guarded{}; + bsk::V three_d_high{}; + bsk::V three_d_leading{}; + bsk::V three_d_lift{}; + bsk::V three_d_low{}; + bsk::V three_d_middle{}; + bsk::V three_d_minors{}; + bsk::V three_d_pool_b{}; + bsk::V three_d_pool_c{}; + bsk::V three_d_q00{}; + bsk::V three_d_q01{}; + bsk::V three_d_q02{}; + bsk::V three_d_q10{}; + bsk::V three_d_q11{}; + bsk::V three_d_q12{}; + bsk::V three_d_q20{}; + bsk::V three_d_q21{}; + bsk::V three_d_q22{}; + bsk::V three_d_radius{}; + bsk::V three_d_raw{}; + bsk::V three_d_s00{}; + bsk::V three_d_s11{}; + bsk::V three_d_s22{}; + bsk::V three_d_second{}; + bsk::V three_d_sum_flat{}; + bsk::V three_d_sum_linear{}; + bsk::V three_d_sum_square{}; + bsk::V three_d_trailing{}; + bsk::V three_def_00{}; + bsk::V three_def_01{}; + bsk::V three_def_02{}; + bsk::V three_def_10{}; + bsk::V three_def_11{}; + bsk::V three_def_12{}; + bsk::V three_def_20{}; + bsk::V three_def_21{}; + bsk::V three_def_22{}; + bsk::V three_determinant{}; + bsk::V three_dif_00{}; + bsk::V three_dif_01{}; + bsk::V three_dif_02{}; + bsk::V three_dif_10{}; + bsk::V three_dif_11{}; + bsk::V three_dif_12{}; + bsk::V three_dif_20{}; + bsk::V three_dif_21{}; + bsk::V three_dif_22{}; + bsk::V three_first{}; + bsk::V three_free{}; + bsk::V three_guarded{}; + bsk::V three_high{}; + bsk::V three_inside_limit{}; + bsk::V three_leading{}; + bsk::V three_lift{}; + bsk::V three_low{}; + bsk::V three_middle{}; + bsk::V three_minors{}; + bsk::V three_pool_b{}; + bsk::V three_pool_c{}; + bsk::V three_q00{}; + bsk::V three_q01{}; + bsk::V three_q02{}; + bsk::V three_q10{}; + bsk::V three_q11{}; + bsk::V three_q12{}; + bsk::V three_q20{}; + bsk::V three_q21{}; + bsk::V three_q22{}; + bsk::V three_radius{}; + bsk::V three_raw{}; + bsk::V three_s00{}; + bsk::V three_s11{}; + bsk::V three_s22{}; + bsk::V three_second{}; + bsk::V three_sum_flat{}; + bsk::V three_sum_linear{}; + bsk::V three_sum_square{}; + bsk::V three_trailing{}; + bsk::V turn_t{}; + bsk::V turn_z{}; + bsk::V turned_mi{}; + bsk::V turned_mr{}; + bsk::V turned_pi{}; + bsk::V turned_pr{}; + bsk::V turned_zi{}; + bsk::V turned_zr{}; + bsk::V two_pool_dt_v{}; + bsk::tup, bsk::V> u1{}; + bsk::tup, bsk::V> u2{}; + bsk::V ubvi{}; + bsk::V ubvr{}; + bsk::V ui{}; + bsk::V ur{}; + bsk::V vi_{}; + bsk::V vr_{}; + bsk::tup, bsk::V> w0{}; + bsk::tup, bsk::V> w1{}; + bsk::V w11{}; + bsk::V w12{}; + bsk::V w13{}; + bsk::tup, bsk::V> w2{}; + bsk::V w21{}; + bsk::V w22{}; + bsk::V w23{}; + bsk::V w31{}; + bsk::V w32{}; + bsk::V w33{}; + bsk::V wash_v{}; + bsk::V wbvi{}; + bsk::V wbvr{}; + bsk::V wound_v{}; + bsk::V wout_value{}; + bsk::V wvi{}; + bsk::V wvr{}; + bsk::V xbmvi{}; + bsk::V xbmvr{}; + bsk::V xbpvi{}; + bsk::V xbpvr{}; + bsk::V xbvi{}; + bsk::V xbvr{}; + bsk::V xcvi{}; + bsk::V xcvr{}; + bsk::V xversal_att{}; + bsk::V xversal_dt{}; + bsk::V yi{}; + bsk::V yr{}; + bsk::V zangle_v{}; + bsk::V zbvi{}; + bsk::V zbvr{}; + bsk::V zvi{}; + bsk::V zvr{}; + problem = (problem_base + (bsk::program_id(0) * problems)); + problem = (problem + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < problem_end); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + auto local = (problem - problem_base); + auto record_stride = (bsk::select(bsk::truth((pools == 3)), 7, bsk::select(bsk::truth((pools == 2)), 6, bsk::select(bsk::truth((pools == 1)), 4, 3))) * state_count); + auto trajectory = (((local * event_count) * record_stride) + state); + auto minus_plane = state_count; + auto long_plane = (2 * state_count); + auto bound_plane = (3 * state_count); + auto bplus_plane = (4 * state_count); + auto bminus_plane = (5 * state_count); + auto semisolid_plane = (6 * state_count); + auto empty = bsk::full(0); + pvr = empty; + pvi = empty; + mvr = empty; + mvi = empty; + zvr = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + zvi = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_b1_phase = 0.0f; + atom_b0 = 0.0f; + if (bsk::truth(off_axis)) { + atom_b1_phase = bsk::ld((b1_phase + scalar_atom), active_atom, 0.0f); + atom_b0 = bsk::ld((b0 + scalar_atom), active_atom, 0.0f); + } + atom_inv = 1.0f; + if (bsk::truth(inverting)) { + atom_inv = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + atom_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + } + atom_flow = 0.0f; + direction = 0.0f; + atom_washout = 0.0f; + if (bsk::truth(moving)) { + auto atom_velocity = bsk::ld((velocity + scalar_atom), active_atom, 0.0f); + atom_flow = (atom_velocity * flow_scale); + // |v| has no derivative at the origin, so a still voxel contributes + // none. + direction = (bsk::cast((atom_velocity > 0.0f)) - bsk::cast((atom_velocity < 0.0f))); + atom_washout = (bsk::abs(atom_velocity) * washout_scale); + } + auto order = bsk::cast(state); + auto longitudinal_weight = (order * order); + auto transverse_weight = ((longitudinal_weight + order) + 0.3333333333333333f); + auto r1_value = bsk::truediv(1000.0f, atom_t1); + auto r2_value = bsk::truediv(1000.0f, atom_t2); + auto location = bsk::mod(atom, locations); + // A semisolid pool rides along as a plane of its own: the pulse deposits + // into it and it exchanges with the free water, so the reverse sweep cannot + // replay it from the free pool's. + atom_bound = 0.0f; + atom_exchange = 0.0f; + atom_t1b = 1.0f; + atom_t2b = 1.0f; + atom_shift = 0.0f; + r1b_value = 0.0f; + r2b_value = 0.0f; + atom_semisolid = 0.0f; + atom_semisolid_exchange = 0.0f; + atom_t1c = 1.0f; + r1c_value = 0.0f; + atom_free = 1.0f; + poolvr = empty; + poolvi = empty; + bpvr = empty; + bpvi = empty; + bmvr = empty; + bmvi = empty; + semivr = empty; + semivi = empty; + if (bsk::truth((pools == 1))) { + atom_bound = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + atom_t1b = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + r1b_value = bsk::truediv(1000.0f, atom_t1b); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + atom_bound = bsk::ld((pool_b_fraction + scalar_atom), active_atom, 0.0f); + atom_exchange = bsk::ld((pool_b_exchange + scalar_atom), active_atom, 0.0f); + atom_t1b = bsk::ld((t1_pool_b + scalar_atom), active_atom, 1.0f); + r1b_value = bsk::truediv(1000.0f, atom_t1b); + atom_t2b = bsk::ld((t2_pool_b + scalar_atom), active_atom, 1.0f); + r2b_value = bsk::truediv(1000.0f, atom_t2b); + atom_shift = bsk::ld((pool_b_shift + scalar_atom), active_atom, 0.0f); + } + if (bsk::truth((pools == 3))) { + // The semisolid pool takes the rows a run with it alone would take, so + // the two second pools never contend for one. + atom_semisolid = bsk::ld((bound_fraction + scalar_atom), active_atom, 0.0f); + atom_semisolid_exchange = bsk::ld((exchange_rate + scalar_atom), active_atom, 0.0f); + atom_t1c = bsk::ld((t1_bound + scalar_atom), active_atom, 1.0f); + r1c_value = bsk::truediv(1000.0f, atom_t1c); + semivr = (empty + bsk::where((state == 0), (atom_semisolid + 0.0f), 0.0f)); + } + if (bsk::truth((pools > 0))) { + // The fractions split the equilibrium at t = 0. + atom_free = ((1.0f - atom_bound) - atom_semisolid); + zvr = (empty + bsk::where((state == 0), atom_free, 0.0f)); + poolvr = (empty + bsk::where((state == 0), (atom_bound + 0.0f), 0.0f)); + } + auto event_base = (train * event_count); + // The forward half records the trajectory the reverse half walks back, + // and the two are launched separately: each compiles the sweep it is + // asked for and no more. + if (bsk::truth(recording)) { + for (std::int64_t event = 0; event < event_count; event += 1) { + slot = (trajectory + (event * record_stride)); + bsk::st((trajectory_r + slot), pvr, state_mask); + bsk::st((trajectory_i + slot), pvi, state_mask); + bsk::st(((trajectory_r + slot) + minus_plane), mvr, state_mask); + bsk::st(((trajectory_i + slot) + minus_plane), mvi, state_mask); + bsk::st(((trajectory_r + slot) + long_plane), zvr, state_mask); + bsk::st(((trajectory_i + slot) + long_plane), zvi, state_mask); + if (bsk::truth((pools > 0))) { + bsk::st(((trajectory_r + slot) + bound_plane), poolvr, state_mask); + bsk::st(((trajectory_i + slot) + bound_plane), poolvi, state_mask); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + bsk::st(((trajectory_r + slot) + bplus_plane), bpvr, state_mask); + bsk::st(((trajectory_i + slot) + bplus_plane), bpvi, state_mask); + bsk::st(((trajectory_r + slot) + bminus_plane), bmvr, state_mask); + bsk::st(((trajectory_i + slot) + bminus_plane), bmvi, state_mask); + } + if (bsk::truth((pools == 3))) { + bsk::st(((trajectory_r + slot) + semisolid_plane), semivr, state_mask); + bsk::st(((trajectory_i + slot) + semisolid_plane), semivi, state_mask); + } + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + wout_value = 1.0f; + if (bsk::truth(moving)) { + wout_value = _washout(atom_washout, dt_value); + } + e1_value = (bsk::exp(((-r1_value) * dt_value)) * wout_value); + e2_value = (bsk::exp(((-r2_value) * dt_value)) * wout_value); + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping(atom_damping, dt_value, order); + damp_z = bsk::get<0>(t0_); + damp_t = bsk::get<1>(t0_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + recovery_value = (1.0f - e1_value); + bare1_value = e1_value; + bare2_value = e2_value; + e1_value = (bare1_value * damp_z); + e2_value = (bare2_value * damp_t); + turn_t = 0.0f; + auto t1_ = bsk::make_tup(1.0f, 0.0f); + szr = bsk::get<0>(t1_); + szi = bsk::get<1>(t1_); + if (bsk::truth(moving)) { + auto t2_ = _flow(atom_flow, dt_value, order); + turn_z = bsk::get<0>(t2_); + turn_t = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup(bsk::cos(turn_z), bsk::sin(turn_z)); + szr = bsk::get<0>(t3_); + szi = bsk::get<1>(t3_); + } + auto t4_ = bsk::make_tup(1.0f, 0.0f); + qr = bsk::get<0>(t4_); + qi = bsk::get<1>(t4_); + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + angle_value = ((-6.283185307179586f * (atom_b0 * dt_value)) + turn_t); + auto t5_ = bsk::make_tup(bsk::cos(angle_value), bsk::sin(angle_value)); + qr = bsk::get<0>(t5_); + qi = bsk::get<1>(t5_); + } + auto t6_ = bsk::make_tup((e2_value * qr), (e2_value * qi)); + ovr = bsk::get<0>(t6_); + ovi = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup((e1_value * szr), (e1_value * szi)); + lvr = bsk::get<0>(t7_); + lvi = bsk::get<1>(t7_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // With an exchanging pool the transverse relaxation sits inside the + // operator instead of in the scalar the free pool alone multiplies. + across = _two_pool_transverse_step_jvp(r2_value, 0.0f, r2b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, atom_free, 0.0f, atom_shift, 0.0f, dt_value, 0.0f, wout_value, 0.0f); + auto t8_ = bsk::make_tup(bsk::get<0>(across), bsk::get<1>(across)); + a11r = bsk::get<0>(t8_); + a11i = bsk::get<1>(t8_); + auto t9_ = bsk::make_tup(bsk::get<2>(across), bsk::get<3>(across)); + a12r = bsk::get<0>(t9_); + a12i = bsk::get<1>(t9_); + auto t10_ = bsk::make_tup(bsk::get<4>(across), bsk::get<5>(across)); + a21r = bsk::get<0>(t10_); + a21i = bsk::get<1>(t10_); + auto t11_ = bsk::make_tup(bsk::get<6>(across), bsk::get<7>(across)); + a22r = bsk::get<0>(t11_); + a22i = bsk::get<1>(t11_); + auto t12_ = bsk::make_tup((damp_t * qr), (damp_t * qi)); + carr = bsk::get<0>(t12_); + cari = bsk::get<1>(t12_); + auto t13_ = _complex_mul(a11r, a11i, pvr, pvi); + f11r = bsk::get<0>(t13_); + f11i = bsk::get<1>(t13_); + auto t14_ = _complex_mul(a12r, a12i, bpvr, bpvi); + f12r = bsk::get<0>(t14_); + f12i = bsk::get<1>(t14_); + auto t15_ = _complex_mul(a21r, a21i, pvr, pvi); + g21r = bsk::get<0>(t15_); + g21i = bsk::get<1>(t15_); + auto t16_ = _complex_mul(a22r, a22i, bpvr, bpvi); + g22r = bsk::get<0>(t16_); + g22i = bsk::get<1>(t16_); + // ``F-`` takes the conjugate of the operator entry by entry, not its + // transpose: it is the conjugate state following the conjugate map. + auto t17_ = _complex_mul(a11r, (-a11i), mvr, mvi); + h11r = bsk::get<0>(t17_); + h11i = bsk::get<1>(t17_); + auto t18_ = _complex_mul(a12r, (-a12i), bmvr, bmvi); + h12r = bsk::get<0>(t18_); + h12i = bsk::get<1>(t18_); + auto t19_ = _complex_mul(a21r, (-a21i), mvr, mvi); + k21r = bsk::get<0>(t19_); + k21i = bsk::get<1>(t19_); + auto t20_ = _complex_mul(a22r, (-a22i), bmvr, bmvi); + k22r = bsk::get<0>(t20_); + k22i = bsk::get<1>(t20_); + auto t21_ = _complex_mul((f11r + f12r), (f11i + f12i), carr, cari); + pvr = bsk::get<0>(t21_); + pvi = bsk::get<1>(t21_); + auto t22_ = _complex_mul((g21r + g22r), (g21i + g22i), carr, cari); + bpvr = bsk::get<0>(t22_); + bpvi = bsk::get<1>(t22_); + auto t23_ = _complex_mul((h11r + h12r), (h11i + h12i), carr, (-cari)); + mvr = bsk::get<0>(t23_); + mvi = bsk::get<1>(t23_); + auto t24_ = _complex_mul((k21r + k22r), (k21i + k22i), carr, (-cari)); + bmvr = bsk::get<0>(t24_); + bmvi = bsk::get<1>(t24_); + } else { + auto t25_ = _complex_mul(ovr, ovi, pvr, pvi); + pvr = bsk::get<0>(t25_); + pvi = bsk::get<1>(t25_); + auto t26_ = _complex_mul(ovr, (-ovi), mvr, mvi); + mvr = bsk::get<0>(t26_); + mvi = bsk::get<1>(t26_); + } + if (bsk::truth((pools == 3))) { + // Three pools mix through a 3x3 formed once for the interval; each + // second pool exchanges with the free water and not with the other. + nil = (0.0f * dt_value); + hold_value = (wout_value + nil); + if (bsk::truth(tabulated)) { + auto t27_ = _three_pool_from_table(pool_table, bsk::ld(((duration_row + event_base) + event), active_atom, 0), atom, atom_count, active_atom, hold_value, atom_free, atom_bound, atom_semisolid); + w11 = bsk::get<0>(t27_); + w12 = bsk::get<1>(t27_); + w13 = bsk::get<2>(t27_); + w21 = bsk::get<3>(t27_); + w22 = bsk::get<4>(t27_); + w23 = bsk::get<5>(t27_); + w31 = bsk::get<6>(t27_); + w32 = bsk::get<7>(t27_); + w33 = bsk::get<8>(t27_); + grow_free = bsk::get<9>(t27_); + grow_pool_b = bsk::get<10>(t27_); + grow_semisolid = bsk::get<11>(t27_); + } else { + auto t28_ = _three_pool_step_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, dt_value, nil, hold_value, nil, narrow); + w11 = bsk::get<0>(t28_); + w12 = bsk::get<1>(t28_); + w13 = bsk::get<2>(t28_); + w21 = bsk::get<3>(t28_); + w22 = bsk::get<4>(t28_); + w23 = bsk::get<5>(t28_); + w31 = bsk::get<6>(t28_); + w32 = bsk::get<7>(t28_); + w33 = bsk::get<8>(t28_); + grow_free = bsk::get<9>(t28_); + grow_pool_b = bsk::get<10>(t28_); + grow_semisolid = bsk::get<11>(t28_); + _dw11 = bsk::get<12>(t28_); + _dw12 = bsk::get<13>(t28_); + _dw13 = bsk::get<14>(t28_); + _dw21 = bsk::get<15>(t28_); + _dw22 = bsk::get<16>(t28_); + _dw23 = bsk::get<17>(t28_); + _dw31 = bsk::get<18>(t28_); + _dw32 = bsk::get<19>(t28_); + _dw33 = bsk::get<20>(t28_); + _dgf = bsk::get<21>(t28_); + _dgb = bsk::get<22>(t28_); + _dgs = bsk::get<23>(t28_); + } + auto t29_ = bsk::make_tup((damp_z * szr), (damp_z * szi)); + spin_r = bsk::get<0>(t29_); + spin_i = bsk::get<1>(t29_); + mix_fr = (((w11 * zvr) + (w12 * poolvr)) + (w13 * semivr)); + mix_fi = (((w11 * zvi) + (w12 * poolvi)) + (w13 * semivi)); + mix_br = (((w21 * zvr) + (w22 * poolvr)) + (w23 * semivr)); + mix_bi = (((w21 * zvi) + (w22 * poolvi)) + (w23 * semivi)); + mix_cr = (((w31 * zvr) + (w32 * poolvr)) + (w33 * semivr)); + mix_ci = (((w31 * zvi) + (w32 * poolvi)) + (w33 * semivi)); + auto t30_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + zvr = bsk::get<0>(t30_); + zvi = bsk::get<1>(t30_); + auto t31_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + poolvr = bsk::get<0>(t31_); + poolvi = bsk::get<1>(t31_); + auto t32_ = _complex_mul(spin_r, spin_i, mix_cr, mix_ci); + semivr = bsk::get<0>(t32_); + semivi = bsk::get<1>(t32_); + zvr = (zvr + bsk::where((state == 0), grow_free, 0.0f)); + poolvr = (poolvr + bsk::where((state == 0), grow_pool_b, 0.0f)); + semivr = (semivr + bsk::where((state == 0), grow_semisolid, 0.0f)); + } else if (bsk::truth((pools > 0))) { + // The pools exchange while they relax, so the longitudinal step is a + // 2x2 the interval forms once and the per-order damping and turn + // multiply. Read from the dual helper with no direction to follow: + // what only its tangents reach, the compiler drops. + auto t33_ = _two_pool_step_jvp(r1_value, 0.0f, r1b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, dt_value, 0.0f, wout_value, 0.0f); + pe11 = bsk::get<0>(t33_); + pe12 = bsk::get<1>(t33_); + pe21 = bsk::get<2>(t33_); + pe22 = bsk::get<3>(t33_); + prec_f = bsk::get<4>(t33_); + prec_b = bsk::get<5>(t33_); + _d11 = bsk::get<6>(t33_); + _d12 = bsk::get<7>(t33_); + _d21 = bsk::get<8>(t33_); + _d22 = bsk::get<9>(t33_); + _drf = bsk::get<10>(t33_); + _drb = bsk::get<11>(t33_); + auto t34_ = bsk::make_tup((damp_z * szr), (damp_z * szi)); + spin_r = bsk::get<0>(t34_); + spin_i = bsk::get<1>(t34_); + mix_fr = ((pe11 * zvr) + (pe12 * poolvr)); + mix_fi = ((pe11 * zvi) + (pe12 * poolvi)); + mix_br = ((pe21 * zvr) + (pe22 * poolvr)); + mix_bi = ((pe21 * zvi) + (pe22 * poolvi)); + auto t35_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + zvr = bsk::get<0>(t35_); + zvi = bsk::get<1>(t35_); + auto t36_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + poolvr = bsk::get<0>(t36_); + poolvi = bsk::get<1>(t36_); + zvr = (zvr + bsk::where((state == 0), prec_f, 0.0f)); + poolvr = (poolvr + bsk::where((state == 0), prec_b, 0.0f)); + } else { + auto t37_ = _complex_mul(lvr, lvi, zvr, zvi); + zvr = bsk::get<0>(t37_); + zvi = bsk::get<1>(t37_); + zvr = (zvr + bsk::where((state == 0), recovery_value, 0.0f)); + } + event_action = bsk::cast(bsk::ld((action + event))); + pre_shift = (bsk::band(event_action, 1) != 0); + auto t38_ = _shift(pvr, pvi, mvr, mvi, state, state_mask, state_count); + svr = bsk::get<0>(t38_); + svi = bsk::get<1>(t38_); + wvr = bsk::get<2>(t38_); + wvi = bsk::get<3>(t38_); + pvr = bsk::where(pre_shift, svr, pvr); + pvi = bsk::where(pre_shift, svi, pvi); + mvr = bsk::where(pre_shift, wvr, mvr); + mvi = bsk::where(pre_shift, wvi, mvi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t39_ = _shift(bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count); + svr = bsk::get<0>(t39_); + svi = bsk::get<1>(t39_); + wvr = bsk::get<2>(t39_); + wvi = bsk::get<3>(t39_); + bpvr = bsk::where(pre_shift, svr, bpvr); + bpvi = bsk::where(pre_shift, svi, bpvi); + bmvr = bsk::where(pre_shift, wvr, bmvr); + bmvi = bsk::where(pre_shift, wvi, bmvi); + } + event_kind = bsk::ld((kind + event)); + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + zvr = bsk::where(invert, ((-atom_inv) * zvr), zvr); + zvi = bsk::where(invert, ((-atom_inv) * zvi), zvi); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // A chemically exchanging pool is free water and inverts like any + // other; a semisolid one is saturated instead, by the pulse's own + // saturation term. + poolvr = bsk::where(invert, ((-atom_inv) * poolvr), poolvr); + poolvi = bsk::where(invert, ((-atom_inv) * poolvi), poolvi); + } + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_phase = _event_value(phase, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + pulse_b1_phase = atom_b1_phase; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse a row of its own. + if (bsk::truth(shimmed)) { + row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + if (bsk::truth(off_axis)) { + pulse_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + } + } + alpha_value = (event_flip * pulse_b1); + phi_value = (event_phase + pulse_b1_phase); + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + // The pool absorbs the power the pulse deposits, read at the offset + // the pulse is played less the voxel's own. + offset_value = (bsk::ld((rf_frequency + event)) - atom_b0); + auto t40_ = _lineshape_at_slope(lineshape, offset_value, lineshape_bins, lineshape_step); + shape_value = bsk::get<0>(t40_); + auto _shape_slope = bsk::get<1>(t40_); + event_saturation = bsk::ld((saturation + event)); + power_value = ((event_saturation * alpha_value) * alpha_value); + absorbed_value = bsk::exp((power_value * shape_value)); + saturating = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((pools == 1))) { + poolvr = bsk::where(saturating, (absorbed_value * poolvr), poolvr); + poolvi = bsk::where(saturating, (absorbed_value * poolvi), poolvi); + } else { + semivr = bsk::where(saturating, (absorbed_value * semivr), semivr); + semivi = bsk::where(saturating, (absorbed_value * semivi), semivi); + } + } + cos_value = bsk::cos(alpha_value); + sin_value = bsk::sin(alpha_value); + auto t41_ = bsk::make_tup(bsk::cos(phi_value), bsk::sin(phi_value)); + p1r = bsk::get<0>(t41_); + p1i = bsk::get<1>(t41_); + auto t42_ = _complex_mul(p1r, p1i, p1r, p1i); + p2r = bsk::get<0>(t42_); + p2i = bsk::get<1>(t42_); + auto t43_ = _rotation_coefficients((0.5f * (1.0f + cos_value)), (0.5f * (1.0f - cos_value)), sin_value, cos_value, p1r, p1i, p2r, p2i, p1r, (-p1i)); + t00 = bsk::get<0>(t43_); + t01 = bsk::get<1>(t43_); + t02 = bsk::get<2>(t43_); + t12 = bsk::get<3>(t43_); + t20 = bsk::get<4>(t43_); + t21 = bsk::get<5>(t43_); + t22 = bsk::get<6>(t43_); + auto a0 = _complex_mul(bsk::get<0>(t00), bsk::get<1>(t00), pvr, pvi); + auto a1 = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), mvr, mvi); + auto a2 = _complex_mul(bsk::get<0>(t02), bsk::get<1>(t02), zvr, zvi); + auto b0_ = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), pvr, pvi); + auto b1_ = _complex_mul(bsk::get<0>(t00), bsk::get<1>(t00), mvr, mvi); + auto b2 = _complex_mul(bsk::get<0>(t12), bsk::get<1>(t12), zvr, zvi); + auto c0 = _complex_mul(bsk::get<0>(t20), bsk::get<1>(t20), pvr, pvi); + auto c1 = _complex_mul(bsk::get<0>(t21), bsk::get<1>(t21), mvr, mvi); + auto c2 = _complex_mul(bsk::get<0>(t22), bsk::get<1>(t22), zvr, zvi); + turned_pr = ((bsk::get<0>(a0) + bsk::get<0>(a1)) + bsk::get<0>(a2)); + turned_pi = ((bsk::get<1>(a0) + bsk::get<1>(a1)) + bsk::get<1>(a2)); + turned_mr = ((bsk::get<0>(b0_) + bsk::get<0>(b1_)) + bsk::get<0>(b2)); + turned_mi = ((bsk::get<1>(b0_) + bsk::get<1>(b1_)) + bsk::get<1>(b2)); + turned_zr = ((bsk::get<0>(c0) + bsk::get<0>(c1)) + bsk::get<0>(c2)); + turned_zi = ((bsk::get<1>(c0) + bsk::get<1>(c1)) + bsk::get<1>(c2)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + if (bsk::truth(dynamic)) { + pair = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, active_atom); + auto t44_ = bsk::make_tup(bsk::get<0>(pair), bsk::get<1>(pair)); + shaped_ar = bsk::get<0>(t44_); + shaped_ai = bsk::get<1>(t44_); + // The pair is integrated at zero RF phase, so the event's own + // phase turns the axis afterwards. + auto t45_ = _complex_mul(bsk::get<2>(pair), bsk::get<3>(pair), p1r, (-p1i)); + shaped_br = bsk::get<0>(t45_); + shaped_bi = bsk::get<1>(t45_); + } else { + auto t46_ = _profile_pair(profile, _table_row(profile_index, event, location, locations), alpha_value, profile_bins, profile_step); + shaped_ar = bsk::get<0>(t46_); + shaped_ai = bsk::get<1>(t46_); + shaped_br = bsk::get<2>(t46_); + shaped_bi = bsk::get<3>(t46_); + auto t47_ = _complex_mul(shaped_br, shaped_bi, p1r, (-p1i)); + shaped_br = bsk::get<0>(t47_); + shaped_bi = bsk::get<1>(t47_); + } + auto t48_ = _rotate_spinor(shaped_ar, shaped_ai, shaped_br, shaped_bi, pvr, pvi, mvr, mvi, zvr, zvi); + turned_pr = bsk::get<0>(t48_); + turned_pi = bsk::get<1>(t48_); + turned_mr = bsk::get<2>(t48_); + turned_mi = bsk::get<3>(t48_); + turned_zr = bsk::get<4>(t48_); + turned_zi = bsk::get<5>(t48_); + } + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // The same pulse, the same rotation. A chemical shift moves where a + // pool precesses, not what a pulse does to it. + auto e0 = _complex_mul(bsk::get<0>(t00), bsk::get<1>(t00), bpvr, bpvi); + auto e1_ = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), bmvr, bmvi); + auto e2_ = _complex_mul(bsk::get<0>(t02), bsk::get<1>(t02), poolvr, poolvi); + auto f0 = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), bpvr, bpvi); + auto f1 = _complex_mul(bsk::get<0>(t00), bsk::get<1>(t00), bmvr, bmvi); + auto f2 = _complex_mul(bsk::get<0>(t12), bsk::get<1>(t12), poolvr, poolvi); + auto h0 = _complex_mul(bsk::get<0>(t20), bsk::get<1>(t20), bpvr, bpvi); + auto h1 = _complex_mul(bsk::get<0>(t21), bsk::get<1>(t21), bmvr, bmvi); + auto h2 = _complex_mul(bsk::get<0>(t22), bsk::get<1>(t22), poolvr, poolvi); + auto t49_ = bsk::make_tup(((bsk::get<0>(e0) + bsk::get<0>(e1_)) + bsk::get<0>(e2_)), ((bsk::get<1>(e0) + bsk::get<1>(e1_)) + bsk::get<1>(e2_))); + spun_pr = bsk::get<0>(t49_); + spun_pi = bsk::get<1>(t49_); + auto t50_ = bsk::make_tup(((bsk::get<0>(f0) + bsk::get<0>(f1)) + bsk::get<0>(f2)), ((bsk::get<1>(f0) + bsk::get<1>(f1)) + bsk::get<1>(f2))); + spun_mr = bsk::get<0>(t50_); + spun_mi = bsk::get<1>(t50_); + auto t51_ = bsk::make_tup(((bsk::get<0>(h0) + bsk::get<0>(h1)) + bsk::get<0>(h2)), ((bsk::get<1>(h0) + bsk::get<1>(h1)) + bsk::get<1>(h2))); + spun_zr = bsk::get<0>(t51_); + spun_zi = bsk::get<1>(t51_); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + auto t52_ = _rotate_spinor(shaped_ar, shaped_ai, shaped_br, shaped_bi, bpvr, bpvi, bmvr, bmvi, poolvr, poolvi); + spun_pr = bsk::get<0>(t52_); + spun_pi = bsk::get<1>(t52_); + spun_mr = bsk::get<2>(t52_); + spun_mi = bsk::get<3>(t52_); + spun_zr = bsk::get<4>(t52_); + spun_zi = bsk::get<5>(t52_); + } + bpvr = bsk::where(rotate, spun_pr, bpvr); + bpvi = bsk::where(rotate, spun_pi, bpvi); + bmvr = bsk::where(rotate, spun_mr, bmvr); + bmvi = bsk::where(rotate, spun_mi, bmvi); + poolvr = bsk::where(rotate, spun_zr, poolvr); + poolvi = bsk::where(rotate, spun_zi, poolvi); + } + pvr = bsk::where(rotate, turned_pr, pvr); + pvi = bsk::where(rotate, turned_pi, pvi); + mvr = bsk::where(rotate, turned_mr, mvr); + mvi = bsk::where(rotate, turned_mi, mvi); + zvr = bsk::where(rotate, turned_zr, zvr); + zvi = bsk::where(rotate, turned_zi, zvi); + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t53_ = _shift(bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count); + svr = bsk::get<0>(t53_); + svi = bsk::get<1>(t53_); + wvr = bsk::get<2>(t53_); + wvi = bsk::get<3>(t53_); + auto spoil_b = (bsk::band(event_action, 8) != 0); + bpvr = bsk::where(spoil_b, 0.0f, bsk::where(do_shift, svr, bpvr)); + bpvi = bsk::where(spoil_b, 0.0f, bsk::where(do_shift, svi, bpvi)); + bmvr = bsk::where(spoil_b, 0.0f, bsk::where(do_shift, wvr, bmvr)); + bmvi = bsk::where(spoil_b, 0.0f, bsk::where(do_shift, wvi, bmvi)); + } + auto t54_ = _shift(pvr, pvi, mvr, mvi, state, state_mask, state_count); + svr = bsk::get<0>(t54_); + svi = bsk::get<1>(t54_); + wvr = bsk::get<2>(t54_); + wvi = bsk::get<3>(t54_); + pvr = bsk::where(do_shift, svr, pvr); + pvi = bsk::where(do_shift, svi, pvi); + mvr = bsk::where(do_shift, wvr, mvr); + mvi = bsk::where(do_shift, wvi, mvi); + spoil = (bsk::band(event_action, 8) != 0); + pvr = bsk::where(spoil, 0.0f, pvr); + pvi = bsk::where(spoil, 0.0f, pvi); + mvr = bsk::where(spoil, 0.0f, mvr); + mvi = bsk::where(spoil, 0.0f, mvi); + } + return; + } + // ---- reverse ---- + pbvr = empty; + pbvi = empty; + mbvr = empty; + mbvi = empty; + zbvr = empty; + zbvi = empty; + auto zero = bsk::full(0); + g_diffv = zero; + g_flowv = zero; + g_washv = zero; + g_t1v = zero; + g_t2v = zero; + g_m0v = zero; + g_b1v = zero; + g_b1pv = zero; + g_b0v = zero; + g_invv = zero; + g_boundv = zero; + g_exchv = zero; + g_t1bv = zero; + g_t2bv = zero; + g_shiftv = zero; + g_semiv = zero; + g_sexchv = zero; + g_t1cv = zero; + poolbr = empty; + poolbi = empty; + semibr = empty; + semibi = empty; + ubvr = empty; + ubvi = empty; + wbvr = empty; + wbvi = empty; + for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + event = ((event_count - 1) - reverse); + slot = (trajectory + (event * record_stride)); + auto xpvr = bsk::ld((trajectory_r + slot), state_mask, 0.0f); + auto xpvi = bsk::ld((trajectory_i + slot), state_mask, 0.0f); + auto xmvr = bsk::ld(((trajectory_r + slot) + minus_plane), state_mask, 0.0f); + auto xmvi = bsk::ld(((trajectory_i + slot) + minus_plane), state_mask, 0.0f); + auto xzvr = bsk::ld(((trajectory_r + slot) + long_plane), state_mask, 0.0f); + auto xzvi = bsk::ld(((trajectory_i + slot) + long_plane), state_mask, 0.0f); + xbvr = empty; + xbvi = empty; + xcvr = empty; + xcvi = empty; + xbpvr = empty; + xbpvi = empty; + xbmvr = empty; + xbmvi = empty; + rbpvr = empty; + rbpvi = empty; + rbmvr = empty; + rbmvi = empty; + if (bsk::truth((pools > 0))) { + xbvr = bsk::ld(((trajectory_r + slot) + bound_plane), state_mask, 0.0f); + xbvi = bsk::ld(((trajectory_i + slot) + bound_plane), state_mask, 0.0f); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + xbpvr = bsk::ld(((trajectory_r + slot) + bplus_plane), state_mask, 0.0f); + xbpvi = bsk::ld(((trajectory_i + slot) + bplus_plane), state_mask, 0.0f); + xbmvr = bsk::ld(((trajectory_r + slot) + bminus_plane), state_mask, 0.0f); + xbmvi = bsk::ld(((trajectory_i + slot) + bminus_plane), state_mask, 0.0f); + } + if (bsk::truth((pools == 3))) { + xcvr = bsk::ld(((trajectory_r + slot) + semisolid_plane), state_mask, 0.0f); + xcvi = bsk::ld(((trajectory_i + slot) + semisolid_plane), state_mask, 0.0f); + } + event_action = bsk::cast(bsk::ld((action + event))); + event_kind = bsk::ld((kind + event)); + dt_value = _event_value(duration, event_base, event, active_atom, single_train); + wout_value = 1.0f; + if (bsk::truth(moving)) { + wout_value = _washout(atom_washout, dt_value); + } + auto dry1_value = bsk::exp(((-r1_value) * dt_value)); + auto dry2_value = bsk::exp(((-r2_value) * dt_value)); + e1_value = (dry1_value * wout_value); + e2_value = (dry2_value * wout_value); + damp_z = 1.0f; + damp_t = 1.0f; + if (bsk::truth(diffusing)) { + auto t55_ = _damping(atom_damping, dt_value, order); + damp_z = bsk::get<0>(t55_); + damp_t = bsk::get<1>(t55_); + } + // Order zero is undamped, so recovery keeps the bare longitudinal factor. + recovery_value = (1.0f - e1_value); + bare1_value = e1_value; + bare2_value = e2_value; + e1_value = (bare1_value * damp_z); + e2_value = (bare2_value * damp_t); + turn_t = 0.0f; + auto t56_ = bsk::make_tup(1.0f, 0.0f); + szr = bsk::get<0>(t56_); + szi = bsk::get<1>(t56_); + if (bsk::truth(moving)) { + auto t57_ = _flow(atom_flow, dt_value, order); + turn_z = bsk::get<0>(t57_); + turn_t = bsk::get<1>(t57_); + auto t58_ = bsk::make_tup(bsk::cos(turn_z), bsk::sin(turn_z)); + szr = bsk::get<0>(t58_); + szi = bsk::get<1>(t58_); + } + auto t59_ = bsk::make_tup(1.0f, 0.0f); + qr = bsk::get<0>(t59_); + qi = bsk::get<1>(t59_); + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + angle_value = ((-6.283185307179586f * (atom_b0 * dt_value)) + turn_t); + auto t60_ = bsk::make_tup(bsk::cos(angle_value), bsk::sin(angle_value)); + qr = bsk::get<0>(t60_); + qi = bsk::get<1>(t60_); + } + auto t61_ = bsk::make_tup((e2_value * qr), (e2_value * qi)); + ovr = bsk::get<0>(t61_); + ovi = bsk::get<1>(t61_); + auto t62_ = bsk::make_tup((e1_value * szr), (e1_value * szi)); + lvr = bsk::get<0>(t62_); + lvi = bsk::get<1>(t62_); + // Replay the intra-event stages from the recorded entry state. + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // With an exchanging pool the transverse relaxation sits inside the + // operator instead of in the scalar the free pool alone multiplies. + across = _two_pool_transverse_step_jvp(r2_value, 0.0f, r2b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, atom_free, 0.0f, atom_shift, 0.0f, dt_value, 0.0f, wout_value, 0.0f); + auto t63_ = bsk::make_tup(bsk::get<0>(across), bsk::get<1>(across)); + a11r = bsk::get<0>(t63_); + a11i = bsk::get<1>(t63_); + auto t64_ = bsk::make_tup(bsk::get<2>(across), bsk::get<3>(across)); + a12r = bsk::get<0>(t64_); + a12i = bsk::get<1>(t64_); + auto t65_ = bsk::make_tup(bsk::get<4>(across), bsk::get<5>(across)); + a21r = bsk::get<0>(t65_); + a21i = bsk::get<1>(t65_); + auto t66_ = bsk::make_tup(bsk::get<6>(across), bsk::get<7>(across)); + a22r = bsk::get<0>(t66_); + a22i = bsk::get<1>(t66_); + auto t67_ = bsk::make_tup((damp_t * qr), (damp_t * qi)); + carr = bsk::get<0>(t67_); + cari = bsk::get<1>(t67_); + auto t68_ = _complex_mul(a11r, a11i, xpvr, xpvi); + f11r = bsk::get<0>(t68_); + f11i = bsk::get<1>(t68_); + auto t69_ = _complex_mul(a12r, a12i, xbpvr, xbpvi); + f12r = bsk::get<0>(t69_); + f12i = bsk::get<1>(t69_); + auto t70_ = _complex_mul(a21r, a21i, xpvr, xpvi); + g21r = bsk::get<0>(t70_); + g21i = bsk::get<1>(t70_); + auto t71_ = _complex_mul(a22r, a22i, xbpvr, xbpvi); + g22r = bsk::get<0>(t71_); + g22i = bsk::get<1>(t71_); + // ``F-`` takes the conjugate of the operator entry by entry, not its + // transpose: it is the conjugate state following the conjugate map. + auto t72_ = _complex_mul(a11r, (-a11i), xmvr, xmvi); + h11r = bsk::get<0>(t72_); + h11i = bsk::get<1>(t72_); + auto t73_ = _complex_mul(a12r, (-a12i), xbmvr, xbmvi); + h12r = bsk::get<0>(t73_); + h12i = bsk::get<1>(t73_); + auto t74_ = _complex_mul(a21r, (-a21i), xmvr, xmvi); + k21r = bsk::get<0>(t74_); + k21i = bsk::get<1>(t74_); + auto t75_ = _complex_mul(a22r, (-a22i), xbmvr, xbmvi); + k22r = bsk::get<0>(t75_); + k22i = bsk::get<1>(t75_); + auto t76_ = _complex_mul((f11r + f12r), (f11i + f12i), carr, cari); + rpvr = bsk::get<0>(t76_); + rpvi = bsk::get<1>(t76_); + auto t77_ = _complex_mul((g21r + g22r), (g21i + g22i), carr, cari); + rbpvr = bsk::get<0>(t77_); + rbpvi = bsk::get<1>(t77_); + auto t78_ = _complex_mul((h11r + h12r), (h11i + h12i), carr, (-cari)); + rmvr = bsk::get<0>(t78_); + rmvi = bsk::get<1>(t78_); + auto t79_ = _complex_mul((k21r + k22r), (k21i + k22i), carr, (-cari)); + rbmvr = bsk::get<0>(t79_); + rbmvi = bsk::get<1>(t79_); + } else { + auto t80_ = _complex_mul(ovr, ovi, xpvr, xpvi); + rpvr = bsk::get<0>(t80_); + rpvi = bsk::get<1>(t80_); + auto t81_ = _complex_mul(ovr, (-ovi), xmvr, xmvi); + rmvr = bsk::get<0>(t81_); + rmvi = bsk::get<1>(t81_); + } + rbvr = empty; + rbvi = empty; + rcvr = empty; + rcvi = empty; + if (bsk::truth((pools == 3))) { + nil = (0.0f * dt_value); + hold_value = (wout_value + nil); + if (bsk::truth(tabulated)) { + // The walk back needs the operator itself, which the row + // already holds -- and pooling the cotangents took what + // the eigenvalues were formed for, so nothing here reads + // them. + pool_row = bsk::ld(((duration_row + event_base) + event), active_atom, 0); + auto t82_ = _three_pool_from_table(pool_table, pool_row, atom, atom_count, active_atom, hold_value, atom_free, atom_bound, atom_semisolid); + w11 = bsk::get<0>(t82_); + w12 = bsk::get<1>(t82_); + w13 = bsk::get<2>(t82_); + w21 = bsk::get<3>(t82_); + w22 = bsk::get<4>(t82_); + w23 = bsk::get<5>(t82_); + w31 = bsk::get<6>(t82_); + w32 = bsk::get<7>(t82_); + w33 = bsk::get<8>(t82_); + grow_free = bsk::get<9>(t82_); + grow_pool_b = bsk::get<10>(t82_); + grow_semisolid = bsk::get<11>(t82_); + } else { + auto t83_ = _three_pool_pieces_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, dt_value, nil, narrow); + three_free = bsk::get<0>(t83_); + three_d_free = bsk::get<1>(t83_); + three_pool_b = bsk::get<2>(t83_); + three_d_pool_b = bsk::get<3>(t83_); + three_pool_c = bsk::get<4>(t83_); + three_d_pool_c = bsk::get<5>(t83_); + three_a00 = bsk::get<6>(t83_); + three_d_a00 = bsk::get<7>(t83_); + three_a01 = bsk::get<8>(t83_); + three_d_a01 = bsk::get<9>(t83_); + three_a02 = bsk::get<10>(t83_); + three_d_a02 = bsk::get<11>(t83_); + three_a10 = bsk::get<12>(t83_); + three_d_a10 = bsk::get<13>(t83_); + three_a11 = bsk::get<14>(t83_); + three_d_a11 = bsk::get<15>(t83_); + three_a20 = bsk::get<16>(t83_); + three_d_a20 = bsk::get<17>(t83_); + three_a22 = bsk::get<18>(t83_); + three_d_a22 = bsk::get<19>(t83_); + three_s00 = bsk::get<20>(t83_); + three_d_s00 = bsk::get<21>(t83_); + three_s11 = bsk::get<22>(t83_); + three_d_s11 = bsk::get<23>(t83_); + three_s22 = bsk::get<24>(t83_); + three_d_s22 = bsk::get<25>(t83_); + three_minors = bsk::get<26>(t83_); + three_d_minors = bsk::get<27>(t83_); + three_sum_flat = bsk::get<28>(t83_); + three_sum_linear = bsk::get<29>(t83_); + three_sum_square = bsk::get<30>(t83_); + three_d_sum_flat = bsk::get<31>(t83_); + three_d_sum_linear = bsk::get<32>(t83_); + three_d_sum_square = bsk::get<33>(t83_); + three_lift = bsk::get<34>(t83_); + three_d_lift = bsk::get<35>(t83_); + three_low = bsk::get<36>(t83_); + three_middle = bsk::get<37>(t83_); + three_d_low = bsk::get<38>(t83_); + three_d_middle = bsk::get<39>(t83_); + three_leading = bsk::get<40>(t83_); + three_d_leading = bsk::get<41>(t83_); + three_first = bsk::get<42>(t83_); + three_d_first = bsk::get<43>(t83_); + three_second = bsk::get<44>(t83_); + three_d_second = bsk::get<45>(t83_); + three_determinant = bsk::get<46>(t83_); + three_d_determinant = bsk::get<47>(t83_); + three_high = bsk::get<48>(t83_); + three_d_high = bsk::get<49>(t83_); + three_radius = bsk::get<50>(t83_); + three_d_radius = bsk::get<51>(t83_); + three_cube = bsk::get<52>(t83_); + three_raw = bsk::get<53>(t83_); + three_d_raw = bsk::get<54>(t83_); + three_argument = bsk::get<55>(t83_); + three_inside_limit = bsk::get<56>(t83_); + three_angle = bsk::get<57>(t83_); + three_d_angle = bsk::get<58>(t83_); + three_centre = bsk::get<59>(t83_); + three_d_centre = bsk::get<60>(t83_); + three_trailing = bsk::get<61>(t83_); + three_d_trailing = bsk::get<62>(t83_); + three_guarded = bsk::get<63>(t83_); + three_d_guarded = bsk::get<64>(t83_); + three_q00 = bsk::get<65>(t83_); + three_d_q00 = bsk::get<66>(t83_); + three_q01 = bsk::get<67>(t83_); + three_d_q01 = bsk::get<68>(t83_); + three_q02 = bsk::get<69>(t83_); + three_d_q02 = bsk::get<70>(t83_); + three_q10 = bsk::get<71>(t83_); + three_d_q10 = bsk::get<72>(t83_); + three_q11 = bsk::get<73>(t83_); + three_d_q11 = bsk::get<74>(t83_); + three_q12 = bsk::get<75>(t83_); + three_d_q12 = bsk::get<76>(t83_); + three_q20 = bsk::get<77>(t83_); + three_d_q20 = bsk::get<78>(t83_); + three_q21 = bsk::get<79>(t83_); + three_d_q21 = bsk::get<80>(t83_); + three_q22 = bsk::get<81>(t83_); + three_d_q22 = bsk::get<82>(t83_); + auto t84_ = _three_pool_assemble_jvp(three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, narrow); + three_def_00 = bsk::get<0>(t84_); + three_dif_00 = bsk::get<1>(t84_); + three_def_01 = bsk::get<2>(t84_); + three_dif_01 = bsk::get<3>(t84_); + three_def_02 = bsk::get<4>(t84_); + three_dif_02 = bsk::get<5>(t84_); + three_def_10 = bsk::get<6>(t84_); + three_dif_10 = bsk::get<7>(t84_); + three_def_11 = bsk::get<8>(t84_); + three_dif_11 = bsk::get<9>(t84_); + three_def_12 = bsk::get<10>(t84_); + three_dif_12 = bsk::get<11>(t84_); + three_def_20 = bsk::get<12>(t84_); + three_dif_20 = bsk::get<13>(t84_); + three_def_21 = bsk::get<14>(t84_); + three_dif_21 = bsk::get<15>(t84_); + three_def_22 = bsk::get<16>(t84_); + three_dif_22 = bsk::get<17>(t84_); + auto t85_ = _three_pool_weigh_jvp(three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, hold_value, nil, narrow); + w11 = bsk::get<0>(t85_); + w12 = bsk::get<1>(t85_); + w13 = bsk::get<2>(t85_); + w21 = bsk::get<3>(t85_); + w22 = bsk::get<4>(t85_); + w23 = bsk::get<5>(t85_); + w31 = bsk::get<6>(t85_); + w32 = bsk::get<7>(t85_); + w33 = bsk::get<8>(t85_); + grow_free = bsk::get<9>(t85_); + grow_pool_b = bsk::get<10>(t85_); + grow_semisolid = bsk::get<11>(t85_); + _dw11 = bsk::get<12>(t85_); + _dw12 = bsk::get<13>(t85_); + _dw13 = bsk::get<14>(t85_); + _dw21 = bsk::get<15>(t85_); + _dw22 = bsk::get<16>(t85_); + _dw23 = bsk::get<17>(t85_); + _dw31 = bsk::get<18>(t85_); + _dw32 = bsk::get<19>(t85_); + _dw33 = bsk::get<20>(t85_); + _dgf = bsk::get<21>(t85_); + _dgb = bsk::get<22>(t85_); + _dgs = bsk::get<23>(t85_); + // The operator is O(1) once formed, so the per-order loop below + // takes it at the width the states are carried in. + w11 = bsk::cast(w11); + w12 = bsk::cast(w12); + w13 = bsk::cast(w13); + w21 = bsk::cast(w21); + w22 = bsk::cast(w22); + w23 = bsk::cast(w23); + w31 = bsk::cast(w31); + w32 = bsk::cast(w32); + w33 = bsk::cast(w33); + grow_free = bsk::cast(grow_free); + grow_pool_b = bsk::cast(grow_pool_b); + grow_semisolid = bsk::cast(grow_semisolid); + } + auto t86_ = bsk::make_tup((damp_z * szr), (damp_z * szi)); + spin_r = bsk::get<0>(t86_); + spin_i = bsk::get<1>(t86_); + mix_fr = (((w11 * xzvr) + (w12 * xbvr)) + (w13 * xcvr)); + mix_fi = (((w11 * xzvi) + (w12 * xbvi)) + (w13 * xcvi)); + mix_br = (((w21 * xzvr) + (w22 * xbvr)) + (w23 * xcvr)); + mix_bi = (((w21 * xzvi) + (w22 * xbvi)) + (w23 * xcvi)); + mix_cr = (((w31 * xzvr) + (w32 * xbvr)) + (w33 * xcvr)); + mix_ci = (((w31 * xzvi) + (w32 * xbvi)) + (w33 * xcvi)); + auto t87_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + rzvr = bsk::get<0>(t87_); + rzvi = bsk::get<1>(t87_); + auto t88_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + rbvr = bsk::get<0>(t88_); + rbvi = bsk::get<1>(t88_); + auto t89_ = _complex_mul(spin_r, spin_i, mix_cr, mix_ci); + rcvr = bsk::get<0>(t89_); + rcvi = bsk::get<1>(t89_); + rzvr = (rzvr + bsk::where((state == 0), grow_free, 0.0f)); + rbvr = (rbvr + bsk::where((state == 0), grow_pool_b, 0.0f)); + rcvr = (rcvr + bsk::where((state == 0), grow_semisolid, 0.0f)); + } else if (bsk::truth((pools > 0))) { + auto t90_ = _two_pool_step_jvp(r1_value, 0.0f, r1b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, dt_value, 0.0f, wout_value, 0.0f); + pe11 = bsk::get<0>(t90_); + pe12 = bsk::get<1>(t90_); + pe21 = bsk::get<2>(t90_); + pe22 = bsk::get<3>(t90_); + prec_f = bsk::get<4>(t90_); + prec_b = bsk::get<5>(t90_); + _d11 = bsk::get<6>(t90_); + _d12 = bsk::get<7>(t90_); + _d21 = bsk::get<8>(t90_); + _d22 = bsk::get<9>(t90_); + _drf = bsk::get<10>(t90_); + _drb = bsk::get<11>(t90_); + auto t91_ = bsk::make_tup((damp_z * szr), (damp_z * szi)); + spin_r = bsk::get<0>(t91_); + spin_i = bsk::get<1>(t91_); + mix_fr = ((pe11 * xzvr) + (pe12 * xbvr)); + mix_fi = ((pe11 * xzvi) + (pe12 * xbvi)); + mix_br = ((pe21 * xzvr) + (pe22 * xbvr)); + mix_bi = ((pe21 * xzvi) + (pe22 * xbvi)); + auto t92_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + rzvr = bsk::get<0>(t92_); + rzvi = bsk::get<1>(t92_); + auto t93_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + rbvr = bsk::get<0>(t93_); + rbvi = bsk::get<1>(t93_); + rzvr = (rzvr + bsk::where((state == 0), prec_f, 0.0f)); + rbvr = (rbvr + bsk::where((state == 0), prec_b, 0.0f)); + } else { + auto t94_ = _complex_mul(lvr, lvi, xzvr, xzvi); + rzvr = bsk::get<0>(t94_); + rzvi = bsk::get<1>(t94_); + rzvr = (rzvr + bsk::where((state == 0), recovery_value, 0.0f)); + } + pre_shift = (bsk::band(event_action, 1) != 0); + auto t95_ = _shift(rpvr, rpvi, rmvr, rmvi, state, state_mask, state_count); + svr = bsk::get<0>(t95_); + svi = bsk::get<1>(t95_); + wvr = bsk::get<2>(t95_); + wvi = bsk::get<3>(t95_); + auto spvr = bsk::where(pre_shift, svr, rpvr); + auto spvi = bsk::where(pre_shift, svi, rpvi); + auto smvr = bsk::where(pre_shift, wvr, rmvr); + auto smvi = bsk::where(pre_shift, wvi, rmvi); + sbpvr = empty; + sbpvi = empty; + sbmvr = empty; + sbmvi = empty; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t96_ = _shift(rbpvr, rbpvi, rbmvr, rbmvi, state, state_mask, state_count); + svr = bsk::get<0>(t96_); + svi = bsk::get<1>(t96_); + wvr = bsk::get<2>(t96_); + wvi = bsk::get<3>(t96_); + sbpvr = bsk::where(pre_shift, svr, rbpvr); + sbpvi = bsk::where(pre_shift, svi, rbpvi); + sbmvr = bsk::where(pre_shift, wvr, rbmvr); + sbmvi = bsk::where(pre_shift, wvi, rbmvi); + } + // Undo the trailing spoil or shift. + do_shift = bsk::bor((bsk::band(event_action, 2) != 0), (bsk::band(event_action, 16) != 0)); + spoil = (bsk::band(event_action, 8) != 0); + auto t97_ = _shift_adjoint(pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count); + avr = bsk::get<0>(t97_); + avi = bsk::get<1>(t97_); + bvr = bsk::get<2>(t97_); + bvi = bsk::get<3>(t97_); + auto trailing = bsk::band(do_shift, bsk::bnot(spoil)); + pbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, avr, pbvr)); + pbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, avi, pbvi)); + mbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, bvr, mbvr)); + mbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, bvi, mbvi)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t98_ = _shift_adjoint(ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count); + avr = bsk::get<0>(t98_); + avi = bsk::get<1>(t98_); + bvr = bsk::get<2>(t98_); + bvi = bsk::get<3>(t98_); + ubvr = bsk::where(spoil, 0.0f, bsk::where(trailing, avr, ubvr)); + ubvi = bsk::where(spoil, 0.0f, bsk::where(trailing, avi, ubvi)); + wbvr = bsk::where(spoil, 0.0f, bsk::where(trailing, bvr, wbvr)); + wbvi = bsk::where(spoil, 0.0f, bsk::where(trailing, bvi, wbvi)); + } + event_flip = _event_value(flip, event_base, event, active_atom, single_train); + event_phase = _event_value(phase, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + pulse_b1_phase = atom_b1_phase; + if (bsk::truth(shimmed)) { + row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + row) + atom), active_atom, 1.0f); + } + if (bsk::truth(off_axis)) { + pulse_b1_phase = bsk::ld(((b1_phase + row) + atom), active_atom, 0.0f); + } + } + // ---- recorded sample ---- + auto record = bsk::band((bsk::band(event_action, 32) != 0), (event_kind == 2)); + auto out_ = bsk::ld((output_index + event)); + auto seed_mask = bsk::band(bsk::band(active_atom, record), (out_ >= 0)); + auto seed_real = bsk::ld(((grad_output_real + (problem * output_count)) + out_), seed_mask, 0.0f); + auto seed_imag = bsk::ld(((grad_output_imag + (problem * output_count)) + out_), seed_mask, 0.0f); + auto t99_ = bsk::make_tup(bsk::cos((-event_phase)), bsk::sin((-event_phase))); + auto dvr = bsk::get<0>(t99_); + auto dvi = bsk::get<1>(t99_); + // grad_m0 = Re(conj(seed) * recorded * demodulation) + auto t100_ = bsk::make_tup(spvr, spvi); + recr = bsk::get<0>(t100_); + reci = bsk::get<1>(t100_); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t101_ = bsk::make_tup((spvr + sbpvr), (spvi + sbpvi)); + recr = bsk::get<0>(t101_); + reci = bsk::get<1>(t101_); + } + auto t102_ = _complex_mul(recr, reci, dvr, dvi); + auto wr = bsk::get<0>(t102_); + auto wi = bsk::get<1>(t102_); + g_m0v = (g_m0v + bsk::sum_x(bsk::where((state == 0), ((seed_real * wr) + (seed_imag * wi)), 0.0f))); + // grad_phase = Re(conj(seed) * m0 * recorded * (-i) * demodulation) + auto t103_ = bsk::make_tup((atom_m0 * recr), (atom_m0 * reci)); + yr = bsk::get<0>(t103_); + yi = bsk::get<1>(t103_); + auto t104_ = bsk::make_tup(yi, (-yr)); + yr = bsk::get<0>(t104_); + yi = bsk::get<1>(t104_); + auto t105_ = _complex_mul(yr, yi, dvr, dvi); + yr = bsk::get<0>(t105_); + yi = bsk::get<1>(t105_); + bsk::atomic_add(((grad_phase + event_base) + event), bsk::sum_x(bsk::where((state == 0), ((seed_real * yr) + (seed_imag * yi)), 0.0f)), seed_mask); + // fplus_bar[0] += conj(m0 * demodulation) * seed + auto t106_ = bsk::make_tup((atom_m0 * dvr), (atom_m0 * dvi)); + auto kr = bsk::get<0>(t106_); + auto ki = bsk::get<1>(t106_); + auto t107_ = _complex_mul(kr, (-ki), seed_real, seed_imag); + auto sr = bsk::get<0>(t107_); + auto si = bsk::get<1>(t107_); + pbvr = (pbvr + bsk::where((state == 0), sr, 0.0f)); + pbvi = (pbvi + bsk::where((state == 0), si, 0.0f)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + ubvr = (ubvr + bsk::where((state == 0), sr, 0.0f)); + ubvi = (ubvi + bsk::where((state == 0), si, 0.0f)); + } + // ---- RF adjoint ---- + is_rf = (event_kind == 1); + is_inversion = (bsk::band(event_action, 4) != 0); + invert = bsk::band(is_rf, is_inversion); + g_invv = (g_invv + bsk::sum_x(bsk::where(invert, ((zbvr * (-rzvr)) + (zbvi * (-rzvi))), 0.0f))); + zbvr = bsk::where(invert, ((-atom_inv) * zbvr), zbvr); + zbvi = bsk::where(invert, ((-atom_inv) * zbvi), zbvi); + alpha_value = (event_flip * pulse_b1); + phi_value = (event_phase + pulse_b1_phase); + cos_value = bsk::cos(alpha_value); + sin_value = bsk::sin(alpha_value); + auto t108_ = bsk::make_tup(bsk::cos(phi_value), bsk::sin(phi_value)); + p1r = bsk::get<0>(t108_); + p1i = bsk::get<1>(t108_); + auto t109_ = _complex_mul(p1r, p1i, p1r, p1i); + p2r = bsk::get<0>(t109_); + p2i = bsk::get<1>(t109_); + auto t110_ = _rotation_coefficients((0.5f * (1.0f + cos_value)), (0.5f * (1.0f - cos_value)), sin_value, cos_value, p1r, p1i, p2r, p2i, p1r, (-p1i)); + t00 = bsk::get<0>(t110_); + t01 = bsk::get<1>(t110_); + t02 = bsk::get<2>(t110_); + t12 = bsk::get<3>(t110_); + t20 = bsk::get<4>(t110_); + t21 = bsk::get<5>(t110_); + t22 = bsk::get<6>(t110_); + auto t111_ = _rotation_coefficients((-0.5f * sin_value), (0.5f * sin_value), cos_value, (-sin_value), p1r, p1i, p2r, p2i, p1r, (-p1i)); + auto d00 = bsk::get<0>(t111_); + auto d01 = bsk::get<1>(t111_); + auto d02 = bsk::get<2>(t111_); + auto d12 = bsk::get<3>(t111_); + auto d20 = bsk::get<4>(t111_); + auto d21 = bsk::get<5>(t111_); + auto d22 = bsk::get<6>(t111_); + sat_alpha_v = zero; + sat_b0_v = zero; + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + // The pulse scales every order of the pool by one real number, so + // its cotangent is a single sum over the states it multiplied. + offset_value = (bsk::ld((rf_frequency + event)) - atom_b0); + auto t112_ = _lineshape_at_slope(lineshape, offset_value, lineshape_bins, lineshape_step); + shape_value = bsk::get<0>(t112_); + auto shape_slope = bsk::get<1>(t112_); + event_saturation = bsk::ld((saturation + event)); + power_value = ((event_saturation * alpha_value) * alpha_value); + absorbed_value = bsk::exp((power_value * shape_value)); + if (bsk::truth((pools == 1))) { + per_state = ((poolbr * rbvr) + (poolbi * rbvi)); + } else { + per_state = ((semibr * rcvr) + (semibi * rcvi)); + } + auto grad_absorbed = bsk::sum_x(per_state); + auto grad_exponent = (grad_absorbed * absorbed_value); + auto twice = (event_saturation * 2.0f); + sat_alpha_v = (grad_exponent * ((twice * alpha_value) * shape_value)); + // The lineshape is read at the pulse's offset from the voxel, so a + // step in the voxel's own off-resonance moves the read the other way. + sat_b0_v = ((-grad_exponent) * (power_value * shape_slope)); + saturating = bsk::band(is_rf, bsk::bnot(is_inversion)); + if (bsk::truth((pools == 1))) { + poolbr = bsk::where(saturating, (absorbed_value * poolbr), poolbr); + poolbi = bsk::where(saturating, (absorbed_value * poolbi), poolbi); + } else { + semibr = bsk::where(saturating, (absorbed_value * semibr), semibr); + semibi = bsk::where(saturating, (absorbed_value * semibi), semibi); + } + } + // d/dalpha, contracted with the adjoint. + row0 = _complex_mul(bsk::get<0>(d00), bsk::get<1>(d00), spvr, spvi); + add1 = _complex_mul(bsk::get<0>(d01), bsk::get<1>(d01), smvr, smvi); + add2 = _complex_mul(bsk::get<0>(d02), bsk::get<1>(d02), rzvr, rzvi); + alpha_v = (pbvr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2))); + alpha_v = (alpha_v + (pbvi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + row0 = _complex_mul(bsk::get<0>(d01), (-bsk::get<1>(d01)), spvr, spvi); + add1 = _complex_mul(bsk::get<0>(d00), bsk::get<1>(d00), smvr, smvi); + add2 = _complex_mul(bsk::get<0>(d12), bsk::get<1>(d12), rzvr, rzvi); + alpha_v = (alpha_v + (mbvr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)))); + alpha_v = (alpha_v + (mbvi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + row0 = _complex_mul(bsk::get<0>(d20), bsk::get<1>(d20), spvr, spvi); + add1 = _complex_mul(bsk::get<0>(d21), bsk::get<1>(d21), smvr, smvi); + add2 = _complex_mul(bsk::get<0>(d22), bsk::get<1>(d22), rzvr, rzvi); + alpha_v = (alpha_v + (zbvr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)))); + alpha_v = (alpha_v + (zbvi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + // d/dphi, where only the phase factors carry the dependence. + u1 = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), smvr, smvi); + u2 = _complex_mul(bsk::get<0>(t02), bsk::get<1>(t02), rzvr, rzvi); + auto t113_ = bsk::make_tup((-((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2))), ((2.0f * bsk::get<0>(u1)) + bsk::get<0>(u2))); + ur = bsk::get<0>(t113_); + ui = bsk::get<1>(t113_); + phi_v = ((pbvr * ur) + (pbvi * ui)); + u1 = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), spvr, spvi); + u2 = _complex_mul(bsk::get<0>(t12), bsk::get<1>(t12), rzvr, rzvi); + auto t114_ = bsk::make_tup(((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2)), ((-2.0f * bsk::get<0>(u1)) - bsk::get<0>(u2))); + ur = bsk::get<0>(t114_); + ui = bsk::get<1>(t114_); + phi_v = (phi_v + ((mbvr * ur) + (mbvi * ui))); + u1 = _complex_mul(bsk::get<0>(t20), bsk::get<1>(t20), spvr, spvi); + u2 = _complex_mul(bsk::get<0>(t21), bsk::get<1>(t21), smvr, smvi); + auto t115_ = bsk::make_tup((-(bsk::get<1>(u2) - bsk::get<1>(u1))), (bsk::get<0>(u2) - bsk::get<0>(u1))); + ur = bsk::get<0>(t115_); + ui = bsk::get<1>(t115_); + phi_v = (phi_v + ((zbvr * ur) + (zbvi * ui))); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + auto t116_ = bsk::make_tup(0.0f, 0.0f, 0.0f, 0.0f); + slope_ar = bsk::get<0>(t116_); + slope_ai = bsk::get<1>(t116_); + slope_br = bsk::get<2>(t116_); + slope_bi = bsk::get<3>(t116_); + if (bsk::truth(dynamic)) { + pair = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, active_atom); + auto t117_ = bsk::make_tup(bsk::get<0>(pair), bsk::get<1>(pair)); + shaped_ar = bsk::get<0>(t117_); + shaped_ai = bsk::get<1>(t117_); + auto t118_ = _complex_mul(bsk::get<2>(pair), bsk::get<3>(pair), p1r, (-p1i)); + shaped_br = bsk::get<0>(t118_); + shaped_bi = bsk::get<1>(t118_); + } else { + auto t119_ = _profile_pair_slope(profile, _table_row(profile_index, event, location, locations), alpha_value, profile_bins, profile_step); + shaped_ar = bsk::get<0>(t119_); + slope_ar = bsk::get<1>(t119_); + shaped_ai = bsk::get<2>(t119_); + slope_ai = bsk::get<3>(t119_); + shaped_br = bsk::get<4>(t119_); + slope_br = bsk::get<5>(t119_); + shaped_bi = bsk::get<6>(t119_); + slope_bi = bsk::get<7>(t119_); + auto t120_ = _complex_mul(shaped_br, shaped_bi, p1r, (-p1i)); + shaped_br = bsk::get<0>(t120_); + shaped_bi = bsk::get<1>(t120_); + auto t121_ = _complex_mul(slope_br, slope_bi, p1r, (-p1i)); + slope_br = bsk::get<0>(t121_); + slope_bi = bsk::get<1>(t121_); + } + auto t122_ = _spinor_adjoint(shaped_ar, shaped_ai, shaped_br, shaped_bi, spvr, spvi, smvr, smvi, rzvr, rzvi, pbvr, pbvi, mbvr, mbvi, zbvr, zbvi); + auto grad_ar = bsk::get<0>(t122_); + auto grad_ai = bsk::get<1>(t122_); + auto grad_br = bsk::get<2>(t122_); + auto grad_bi = bsk::get<3>(t122_); + shaped_pbr = bsk::get<4>(t122_); + shaped_pbi = bsk::get<5>(t122_); + shaped_mbr = bsk::get<6>(t122_); + shaped_mbi = bsk::get<7>(t122_); + shaped_zbr = bsk::get<8>(t122_); + shaped_zbi = bsk::get<9>(t122_); + if (bsk::truth(dynamic)) { + // The flip is inside the pair rather than read against it, so + // it has no gradient here: the cotangent goes out on the + // rotation and whatever integrated it carries the rest. ``b`` + // was turned by the phase after the pair came out, so the + // cotangent turns back the other way. + alpha_v = (alpha_v * 0.0f); + auto t123_ = _complex_mul(grad_br, grad_bi, p1r, p1i); + back_r = bsk::get<0>(t123_); + back_i = bsk::get<1>(t123_); + _store_pair_gradient(grad_pair, pair_index, event_base, event, atom, atom_count, bsk::band(is_rf, bsk::bnot(is_inversion)), active_atom, state_mask, grad_ar, grad_ai, back_r, back_i); + } else { + alpha_v = ((grad_ar * slope_ar) + (grad_ai * slope_ai)); + alpha_v = (alpha_v + ((grad_br * slope_br) + (grad_bi * slope_bi))); + } + // d(b e^{-i phi})/dphi is -i times it, and nothing else moves. + phi_v = ((grad_br * shaped_bi) - (grad_bi * shaped_br)); + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + auto t124_ = _spinor_adjoint(shaped_ar, shaped_ai, shaped_br, shaped_bi, sbpvr, sbpvi, sbmvr, sbmvi, rbvr, rbvi, ubvr, ubvi, wbvr, wbvi, poolbr, poolbi); + auto pool_ar = bsk::get<0>(t124_); + auto pool_ai = bsk::get<1>(t124_); + auto pool_pair_br = bsk::get<2>(t124_); + auto pool_pair_bi = bsk::get<3>(t124_); + pool_shaped_pbr = bsk::get<4>(t124_); + pool_shaped_pbi = bsk::get<5>(t124_); + pool_shaped_mbr = bsk::get<6>(t124_); + pool_shaped_mbi = bsk::get<7>(t124_); + pool_shaped_zbr = bsk::get<8>(t124_); + pool_shaped_zbi = bsk::get<9>(t124_); + if (bsk::truth(dynamic)) { + // The same pulse turned this pool, so its cotangent lands + // on the same row. + auto t125_ = _complex_mul(pool_pair_br, pool_pair_bi, p1r, p1i); + back_r = bsk::get<0>(t125_); + back_i = bsk::get<1>(t125_); + _store_pair_gradient(grad_pair, pair_index, event_base, event, atom, atom_count, bsk::band(is_rf, bsk::bnot(is_inversion)), active_atom, state_mask, pool_ar, pool_ai, back_r, back_i); + } else { + alpha_v = (alpha_v + ((pool_ar * slope_ar) + (pool_ai * slope_ai))); + alpha_v = (alpha_v + ((pool_pair_br * slope_br) + (pool_pair_bi * slope_bi))); + } + phi_v = (phi_v + ((pool_pair_br * shaped_bi) - (pool_pair_bi * shaped_br))); + } + } + if (bsk::truth((bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3)))) && bsk::truth((!bsk::truth(profiled))) && bsk::truth((!bsk::truth(dynamic)))))) { + // The same pulse turns the exchanging pool, so its cotangent adds to + // the flip and phase the free pool already left. + row0 = _complex_mul(bsk::get<0>(d00), bsk::get<1>(d00), sbpvr, sbpvi); + add1 = _complex_mul(bsk::get<0>(d01), bsk::get<1>(d01), sbmvr, sbmvi); + add2 = _complex_mul(bsk::get<0>(d02), bsk::get<1>(d02), rbvr, rbvi); + alpha_v = (alpha_v + (ubvr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)))); + alpha_v = (alpha_v + (ubvi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + row0 = _complex_mul(bsk::get<0>(d01), (-bsk::get<1>(d01)), sbpvr, sbpvi); + add1 = _complex_mul(bsk::get<0>(d00), bsk::get<1>(d00), sbmvr, sbmvi); + add2 = _complex_mul(bsk::get<0>(d12), bsk::get<1>(d12), rbvr, rbvi); + alpha_v = (alpha_v + (wbvr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)))); + alpha_v = (alpha_v + (wbvi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + row0 = _complex_mul(bsk::get<0>(d20), bsk::get<1>(d20), sbpvr, sbpvi); + add1 = _complex_mul(bsk::get<0>(d21), bsk::get<1>(d21), sbmvr, sbmvi); + add2 = _complex_mul(bsk::get<0>(d22), bsk::get<1>(d22), rbvr, rbvi); + alpha_v = (alpha_v + (poolbr * ((bsk::get<0>(row0) + bsk::get<0>(add1)) + bsk::get<0>(add2)))); + alpha_v = (alpha_v + (poolbi * ((bsk::get<1>(row0) + bsk::get<1>(add1)) + bsk::get<1>(add2)))); + u1 = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), sbmvr, sbmvi); + u2 = _complex_mul(bsk::get<0>(t02), bsk::get<1>(t02), rbvr, rbvi); + auto t126_ = bsk::make_tup((-((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2))), ((2.0f * bsk::get<0>(u1)) + bsk::get<0>(u2))); + ur = bsk::get<0>(t126_); + ui = bsk::get<1>(t126_); + phi_v = (phi_v + ((ubvr * ur) + (ubvi * ui))); + u1 = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), sbpvr, sbpvi); + u2 = _complex_mul(bsk::get<0>(t12), bsk::get<1>(t12), rbvr, rbvi); + auto t127_ = bsk::make_tup(((2.0f * bsk::get<1>(u1)) + bsk::get<1>(u2)), ((-2.0f * bsk::get<0>(u1)) - bsk::get<0>(u2))); + ur = bsk::get<0>(t127_); + ui = bsk::get<1>(t127_); + phi_v = (phi_v + ((wbvr * ur) + (wbvi * ui))); + u1 = _complex_mul(bsk::get<0>(t20), bsk::get<1>(t20), sbpvr, sbpvi); + u2 = _complex_mul(bsk::get<0>(t21), bsk::get<1>(t21), sbmvr, sbmvi); + auto t128_ = bsk::make_tup((-(bsk::get<1>(u2) - bsk::get<1>(u1))), (bsk::get<0>(u2) - bsk::get<0>(u1))); + ur = bsk::get<0>(t128_); + ui = bsk::get<1>(t128_); + phi_v = (phi_v + ((poolbr * ur) + (poolbi * ui))); + } + rotate = bsk::band(is_rf, bsk::bnot(is_inversion)); + grad_alpha_v = bsk::sum_x(bsk::where(rotate, alpha_v, 0.0f)); + auto grad_phi_v = bsk::sum_x(bsk::where(rotate, phi_v, 0.0f)); + if (bsk::truth((bsk::truth((pools == 1)) || bsk::truth((pools == 3))))) { + auto turning = bsk::where(rotate, 1.0f, 0.0f); + grad_alpha_v = (grad_alpha_v + (sat_alpha_v * turning)); + g_b0v = (g_b0v + (sat_b0_v * turning)); + } + // Conjugate transpose of the rotation. + n0 = _complex_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), pbvr, pbvi); + n1 = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), mbvr, mbvi); + n2 = _complex_mul(bsk::get<0>(t20), (-bsk::get<1>(t20)), zbvr, zbvi); + q0 = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), pbvr, pbvi); + q1 = _complex_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), mbvr, mbvi); + q2 = _complex_mul(bsk::get<0>(t21), (-bsk::get<1>(t21)), zbvr, zbvi); + w0 = _complex_mul(bsk::get<0>(t02), (-bsk::get<1>(t02)), pbvr, pbvi); + w1 = _complex_mul(bsk::get<0>(t12), (-bsk::get<1>(t12)), mbvr, mbvi); + w2 = _complex_mul(bsk::get<0>(t22), (-bsk::get<1>(t22)), zbvr, zbvi); + back_pr = ((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)); + back_pi = ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2)); + back_mr = ((bsk::get<0>(q0) + bsk::get<0>(q1)) + bsk::get<0>(q2)); + back_mi = ((bsk::get<1>(q0) + bsk::get<1>(q1)) + bsk::get<1>(q2)); + back_zr = ((bsk::get<0>(w0) + bsk::get<0>(w1)) + bsk::get<0>(w2)); + back_zi = ((bsk::get<1>(w0) + bsk::get<1>(w1)) + bsk::get<1>(w2)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + // A shaped pulse turned the states, so its own adjoint is what + // goes back rather than the instant rotation's. + auto t129_ = bsk::make_tup(shaped_pbr, shaped_pbi); + back_pr = bsk::get<0>(t129_); + back_pi = bsk::get<1>(t129_); + auto t130_ = bsk::make_tup(shaped_mbr, shaped_mbi); + back_mr = bsk::get<0>(t130_); + back_mi = bsk::get<1>(t130_); + auto t131_ = bsk::make_tup(shaped_zbr, shaped_zbi); + back_zr = bsk::get<0>(t131_); + back_zi = bsk::get<1>(t131_); + } + pbvr = bsk::where(rotate, back_pr, pbvr); + pbvi = bsk::where(rotate, back_pi, pbvi); + mbvr = bsk::where(rotate, back_mr, mbvr); + mbvi = bsk::where(rotate, back_mi, mbvi); + zbvr = bsk::where(rotate, back_zr, zbvr); + zbvi = bsk::where(rotate, back_zi, zbvi); + auto writes_flip = bsk::band(active_atom, rotate); + bsk::atomic_add(((grad_flip + event_base) + event), (grad_alpha_v * pulse_b1), writes_flip); + bsk::atomic_add(((grad_phase + event_base) + event), grad_phi_v, writes_flip); + if (bsk::truth(shimmed)) { + // A pulse's transmit gradient belongs to the shim it drives, so with + // several it lands in that shim's row rather than in a register + // summed over the whole train. + bsk::atomic_add((((grad_tissue + (3 * atom_count)) + row) + atom), (grad_alpha_v * event_flip), writes_flip); + bsk::atomic_add((((grad_tissue + (((4 + shim_rows) - 1) * atom_count)) + row) + atom), grad_phi_v, writes_flip); + } else { + g_b1v = (g_b1v + (grad_alpha_v * event_flip)); + g_b1pv = (g_b1pv + grad_phi_v); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + n0 = _complex_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), ubvr, ubvi); + n1 = _complex_mul(bsk::get<0>(t01), bsk::get<1>(t01), wbvr, wbvi); + n2 = _complex_mul(bsk::get<0>(t20), (-bsk::get<1>(t20)), poolbr, poolbi); + q0 = _complex_mul(bsk::get<0>(t01), (-bsk::get<1>(t01)), ubvr, ubvi); + q1 = _complex_mul(bsk::get<0>(t00), (-bsk::get<1>(t00)), wbvr, wbvi); + q2 = _complex_mul(bsk::get<0>(t21), (-bsk::get<1>(t21)), poolbr, poolbi); + w0 = _complex_mul(bsk::get<0>(t02), (-bsk::get<1>(t02)), ubvr, ubvi); + w1 = _complex_mul(bsk::get<0>(t12), (-bsk::get<1>(t12)), wbvr, wbvi); + w2 = _complex_mul(bsk::get<0>(t22), (-bsk::get<1>(t22)), poolbr, poolbi); + pool_back_pr = ((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)); + pool_back_pi = ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2)); + pool_back_mr = ((bsk::get<0>(q0) + bsk::get<0>(q1)) + bsk::get<0>(q2)); + pool_back_mi = ((bsk::get<1>(q0) + bsk::get<1>(q1)) + bsk::get<1>(q2)); + pool_back_zr = ((bsk::get<0>(w0) + bsk::get<0>(w1)) + bsk::get<0>(w2)); + pool_back_zi = ((bsk::get<1>(w0) + bsk::get<1>(w1)) + bsk::get<1>(w2)); + if (bsk::truth((bsk::truth(profiled) || bsk::truth(dynamic)))) { + // A shaped pulse turned this pool too, so its own adjoint is + // what goes back rather than the instant rotation's. + auto t132_ = bsk::make_tup(pool_shaped_pbr, pool_shaped_pbi); + pool_back_pr = bsk::get<0>(t132_); + pool_back_pi = bsk::get<1>(t132_); + auto t133_ = bsk::make_tup(pool_shaped_mbr, pool_shaped_mbi); + pool_back_mr = bsk::get<0>(t133_); + pool_back_mi = bsk::get<1>(t133_); + auto t134_ = bsk::make_tup(pool_shaped_zbr, pool_shaped_zbi); + pool_back_zr = bsk::get<0>(t134_); + pool_back_zi = bsk::get<1>(t134_); + } + ubvr = bsk::where(rotate, pool_back_pr, ubvr); + ubvi = bsk::where(rotate, pool_back_pi, ubvi); + wbvr = bsk::where(rotate, pool_back_mr, wbvr); + wbvi = bsk::where(rotate, pool_back_mi, wbvi); + poolbr = bsk::where(rotate, pool_back_zr, poolbr); + poolbi = bsk::where(rotate, pool_back_zi, poolbi); + // An inversion turns the exchanging pool's longitudinal state as + // well, so the efficiency carries what both left behind. + g_invv = (g_invv + bsk::sum_x(bsk::where(invert, ((poolbr * (-rbvr)) + (poolbi * (-rbvi))), 0.0f))); + poolbr = bsk::where(invert, ((-atom_inv) * poolbr), poolbr); + poolbi = bsk::where(invert, ((-atom_inv) * poolbi), poolbi); + auto t135_ = _shift_adjoint(ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count); + avr = bsk::get<0>(t135_); + avi = bsk::get<1>(t135_); + bvr = bsk::get<2>(t135_); + bvi = bsk::get<3>(t135_); + ubvr = bsk::where(pre_shift, avr, ubvr); + ubvi = bsk::where(pre_shift, avi, ubvi); + wbvr = bsk::where(pre_shift, bvr, wbvr); + wbvi = bsk::where(pre_shift, bvi, wbvi); + } + auto t136_ = _shift_adjoint(pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count); + avr = bsk::get<0>(t136_); + avi = bsk::get<1>(t136_); + bvr = bsk::get<2>(t136_); + bvi = bsk::get<3>(t136_); + pbvr = bsk::where(pre_shift, avr, pbvr); + pbvi = bsk::where(pre_shift, avi, pbvi); + mbvr = bsk::where(pre_shift, bvr, mbvr); + mbvi = bsk::where(pre_shift, bvi, mbvi); + // ---- relaxation and off-resonance adjoint ---- + // The damping is homogeneous of degree one in every transverse state it + // acts on, so its gradient times the damping itself is the cotangent + // taken against the states the interval leaves. + auto pq = _complex_mul(qr, qi, xpvr, xpvi); + auto mq = _complex_mul(qr, (-qi), xmvr, xmvi); + bare_cot_v = ((pbvr * bsk::get<0>(pq)) + (pbvi * bsk::get<1>(pq))); + bare_cot_v = (bare_cot_v + ((mbvr * bsk::get<0>(mq)) + (mbvi * bsk::get<1>(mq)))); + auto grad_e2_v = bsk::sum_x((bare_cot_v * damp_t)); + cot2_v = ((bare_cot_v * bare2_value) * damp_t); + pool_angle_v = empty; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // With an exchanging pool the damping sits inside the operator, so + // the cotangent the interval leaves is taken against the states it + // produced rather than against a scalar the free pool multiplies. + auto plus_r = ((((pbvr * rpvr) + (pbvi * rpvi)) + (ubvr * rbpvr)) + (ubvi * rbpvi)); + auto plus_i = ((((pbvr * rpvi) - (pbvi * rpvr)) + (ubvr * rbpvi)) - (ubvi * rbpvr)); + auto minus_r = ((((mbvr * rmvr) + (mbvi * rmvi)) + (wbvr * rbmvr)) + (wbvi * rbmvi)); + auto minus_i = ((((mbvr * rmvi) - (mbvi * rmvr)) + (wbvr * rbmvi)) - (wbvi * rbmvr)); + cot2_v = (plus_r + minus_r); + pool_angle_v = (minus_i - plus_i); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + // ``F-`` follows the conjugate of the operator, so its cotangent + // lands on the entry itself rather than on the conjugate of it. + auto t137_ = bsk::make_tup(carr, cari); + auto def_r = bsk::get<0>(t137_); + auto def_i = bsk::get<1>(t137_); + auto t138_ = _complex_mul(((((pbvr * xpvr) + (pbvi * xpvi)) + (mbvr * xmvr)) + (mbvi * xmvi)), ((((pbvr * xpvi) - (pbvi * xpvr)) - (mbvr * xmvi)) + (mbvi * xmvr)), def_r, def_i); + auto t11r = bsk::get<0>(t138_); + auto t11i = bsk::get<1>(t138_); + auto t139_ = _complex_mul(((((pbvr * xbpvr) + (pbvi * xbpvi)) + (mbvr * xbmvr)) + (mbvi * xbmvi)), ((((pbvr * xbpvi) - (pbvi * xbpvr)) - (mbvr * xbmvi)) + (mbvi * xbmvr)), def_r, def_i); + auto t12r = bsk::get<0>(t139_); + auto t12i = bsk::get<1>(t139_); + auto t140_ = _complex_mul(((((ubvr * xpvr) + (ubvi * xpvi)) + (wbvr * xmvr)) + (wbvi * xmvi)), ((((ubvr * xpvi) - (ubvi * xpvr)) - (wbvr * xmvi)) + (wbvi * xmvr)), def_r, def_i); + auto t21r = bsk::get<0>(t140_); + auto t21i = bsk::get<1>(t140_); + auto t141_ = _complex_mul(((((ubvr * xbpvr) + (ubvi * xbpvi)) + (wbvr * xbmvr)) + (wbvi * xbmvi)), ((((ubvr * xbpvi) - (ubvi * xbpvr)) - (wbvr * xbmvi)) + (wbvi * xbmvr)), def_r, def_i); + auto t22r = bsk::get<0>(t141_); + auto t22i = bsk::get<1>(t141_); + auto qbar11 = bsk::make_tup(bsk::sum_x(t11r), bsk::sum_x(t11i), zero, zero); + auto qbar12 = bsk::make_tup(bsk::sum_x(t12r), bsk::sum_x(t12i), zero, zero); + auto qbar21 = bsk::make_tup(bsk::sum_x(t21r), bsk::sum_x(t21i), zero, zero); + auto qbar22 = bsk::make_tup(bsk::sum_x(t22r), bsk::sum_x(t22i), zero, zero); + auto t142_ = _two_pool_transverse_adjoint_jvp(r2_value, 0.0f, r2b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, atom_free, 0.0f, atom_shift, 0.0f, dt_value, 0.0f, wout_value, 0.0f, qbar11, qbar12, qbar21, qbar22); + auto back_r2 = bsk::get<0>(t142_); + _q1 = bsk::get<1>(t142_); + auto back_r2b = bsk::get<2>(t142_); + _q2 = bsk::get<3>(t142_); + auto back_xexch = bsk::get<4>(t142_); + _q3 = bsk::get<5>(t142_); + auto back_xbound = bsk::get<6>(t142_); + _q4 = bsk::get<7>(t142_); + auto back_xfree = bsk::get<8>(t142_); + _q5 = bsk::get<9>(t142_); + auto back_shift = bsk::get<10>(t142_); + _q6 = bsk::get<11>(t142_); + auto back_xdt = bsk::get<12>(t142_); + _q7 = bsk::get<13>(t142_); + auto back_xatt = bsk::get<14>(t142_); + _q8 = bsk::get<15>(t142_); + g_t2v = (g_t2v + (back_r2 * bsk::truediv(-1000.0f, (atom_t2 * atom_t2)))); + g_t2bv = (g_t2bv + (back_r2b * bsk::truediv(-1000.0f, (atom_t2b * atom_t2b)))); + g_exchv = (g_exchv + back_xexch); + // The free fraction is one less the pool's, so what reaches it + // arrives at the pool's own with the sign turned. + g_boundv = (g_boundv + (back_xbound - back_xfree)); + if (bsk::truth((pools == 3))) { + // The free share is one less both fractions, so what the + // transverse operator leaves on it reaches the semisolid too. + g_semiv = (g_semiv - back_xfree); + } + g_shiftv = (g_shiftv + back_shift); + xversal_dt = back_xdt; + xversal_att = back_xatt; + // The pool's transverse cotangents go back through the same + // operator, transposed. + auto t143_ = _complex_mul(a11r, (-a11i), pbvr, pbvi); + ur = bsk::get<0>(t143_); + ui = bsk::get<1>(t143_); + auto t144_ = _complex_mul(a21r, (-a21i), ubvr, ubvi); + vr_ = bsk::get<0>(t144_); + vi_ = bsk::get<1>(t144_); + auto t145_ = _complex_mul((ur + vr_), (ui + vi_), carr, (-cari)); + auto nub_pr = bsk::get<0>(t145_); + auto nub_pi = bsk::get<1>(t145_); + auto t146_ = _complex_mul(a12r, (-a12i), pbvr, pbvi); + ur = bsk::get<0>(t146_); + ui = bsk::get<1>(t146_); + auto t147_ = _complex_mul(a22r, (-a22i), ubvr, ubvi); + vr_ = bsk::get<0>(t147_); + vi_ = bsk::get<1>(t147_); + auto t148_ = _complex_mul((ur + vr_), (ui + vi_), carr, (-cari)); + auto nub_qr = bsk::get<0>(t148_); + auto nub_qi = bsk::get<1>(t148_); + auto t149_ = _complex_mul(a11r, a11i, mbvr, mbvi); + ur = bsk::get<0>(t149_); + ui = bsk::get<1>(t149_); + auto t150_ = _complex_mul(a21r, a21i, wbvr, wbvi); + vr_ = bsk::get<0>(t150_); + vi_ = bsk::get<1>(t150_); + auto t151_ = _complex_mul((ur + vr_), (ui + vi_), carr, cari); + auto nwb_pr = bsk::get<0>(t151_); + auto nwb_pi = bsk::get<1>(t151_); + auto t152_ = _complex_mul(a12r, a12i, mbvr, mbvi); + ur = bsk::get<0>(t152_); + ui = bsk::get<1>(t152_); + auto t153_ = _complex_mul(a22r, a22i, wbvr, wbvi); + vr_ = bsk::get<0>(t153_); + vi_ = bsk::get<1>(t153_); + auto t154_ = _complex_mul((ur + vr_), (ui + vi_), carr, cari); + auto nwb_qr = bsk::get<0>(t154_); + auto nwb_qi = bsk::get<1>(t154_); + auto t155_ = bsk::make_tup(nub_pr, nub_pi); + pbvr = bsk::get<0>(t155_); + pbvi = bsk::get<1>(t155_); + auto t156_ = bsk::make_tup(nub_qr, nub_qi); + ubvr = bsk::get<0>(t156_); + ubvi = bsk::get<1>(t156_); + auto t157_ = bsk::make_tup(nwb_pr, nwb_pi); + mbvr = bsk::get<0>(t157_); + mbvi = bsk::get<1>(t157_); + auto t158_ = bsk::make_tup(nwb_qr, nwb_qi); + wbvr = bsk::get<0>(t158_); + wbvi = bsk::get<1>(t158_); + } + per_angle_v = pool_angle_v; + if (bsk::truth((bsk::truth((pools != 2)) && bsk::truth((pools != 3)) && bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))))) { + auto po = _complex_mul(ovr, ovi, xpvr, xpvi); + auto mo = _complex_mul(ovr, (-ovi), xmvr, xmvi); + // A turn of the transverse states and the off-resonance angle are + // the same derivative; only the weight each order carries differs. + per_angle_v = ((pbvr * (-bsk::get<1>(po))) + (pbvi * bsk::get<0>(po))); + per_angle_v = (per_angle_v - ((mbvr * (-bsk::get<1>(mo))) + (mbvi * bsk::get<0>(mo)))); + } + grad_angle_v = zero; + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + grad_angle_v = bsk::sum_x(per_angle_v); + } + e1_v = empty; + grad_e1_v = zero; + long_damp_v = empty; + attenuation_v = zero; + two_pool_dt_v = zero; + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + attenuation_v = (attenuation_v + xversal_att); + two_pool_dt_v = (two_pool_dt_v + xversal_dt); + } + zangle_v = empty; + if (bsk::truth((pools == 3))) { + // The nine entries of the mixing operator and the three + // recoveries, summed over the orders that share them, then pushed + // back through the closed form once for the whole interval. + auto t159_ = _complex_mul(spin_r, spin_i, xzvr, xzvi); + spun_fr = bsk::get<0>(t159_); + spun_fi = bsk::get<1>(t159_); + auto t160_ = _complex_mul(spin_r, spin_i, xbvr, xbvi); + spun_br = bsk::get<0>(t160_); + spun_bi = bsk::get<1>(t160_); + auto t161_ = _complex_mul(spin_r, spin_i, xcvr, xcvi); + auto spun_cr = bsk::get<0>(t161_); + auto spun_ci = bsk::get<1>(t161_); + auto e11_v = ((zbvr * spun_fr) + (zbvi * spun_fi)); + auto e12_v = ((zbvr * spun_br) + (zbvi * spun_bi)); + auto e13_v = ((zbvr * spun_cr) + (zbvi * spun_ci)); + auto e21_v = ((poolbr * spun_fr) + (poolbi * spun_fi)); + auto e22_v = ((poolbr * spun_br) + (poolbi * spun_bi)); + auto e23_v = ((poolbr * spun_cr) + (poolbi * spun_ci)); + auto e31_v = ((semibr * spun_fr) + (semibi * spun_fi)); + auto e32_v = ((semibr * spun_br) + (semibi * spun_bi)); + auto e33_v = ((semibr * spun_cr) + (semibi * spun_ci)); + if (bsk::truth(tabulated)) { + // Every gradient but the interval's own is linear in these + // twelve, so the events sharing a length pool them here and + // pay the closed form once each after the walk back. + auto bar11 = bsk::sum_x(e11_v); + auto bar12 = bsk::sum_x(e12_v); + auto bar13 = bsk::sum_x(e13_v); + auto bar21 = bsk::sum_x(e21_v); + auto bar22 = bsk::sum_x(e22_v); + auto bar23 = bsk::sum_x(e23_v); + auto bar31 = bsk::sum_x(e31_v); + auto bar32 = bsk::sum_x(e32_v); + auto bar33 = bsk::sum_x(e33_v); + auto bar_free = bsk::sum_x(bsk::where((state == 0), zbvr, nil)); + auto bar_pool_b = bsk::sum_x(bsk::where((state == 0), poolbr, nil)); + auto bar_bound = bsk::sum_x(bsk::where((state == 0), semibr, nil)); + held = (pool_bars + (((local * row_count) + pool_row) * 12)); + bsk::st((held + 0), (bsk::ld((held + 0), active_atom, 0.0f) + bar11), active_atom); + bsk::st((held + 1), (bsk::ld((held + 1), active_atom, 0.0f) + bar12), active_atom); + bsk::st((held + 2), (bsk::ld((held + 2), active_atom, 0.0f) + bar13), active_atom); + bsk::st((held + 3), (bsk::ld((held + 3), active_atom, 0.0f) + bar21), active_atom); + bsk::st((held + 4), (bsk::ld((held + 4), active_atom, 0.0f) + bar22), active_atom); + bsk::st((held + 5), (bsk::ld((held + 5), active_atom, 0.0f) + bar23), active_atom); + bsk::st((held + 6), (bsk::ld((held + 6), active_atom, 0.0f) + bar31), active_atom); + bsk::st((held + 7), (bsk::ld((held + 7), active_atom, 0.0f) + bar32), active_atom); + bsk::st((held + 8), (bsk::ld((held + 8), active_atom, 0.0f) + bar33), active_atom); + bsk::st((held + 9), (bsk::ld((held + 9), active_atom, 0.0f) + bar_free), active_atom); + bsk::st((held + 10), (bsk::ld((held + 10), active_atom, 0.0f) + bar_pool_b), active_atom); + bsk::st((held + 11), (bsk::ld((held + 11), active_atom, 0.0f) + bar_bound), active_atom); + auto t162_ = _three_pool_interval_adjoint(pool_table, pool_row, atom, atom_count, active_atom, r1_value, r1b_value, r1c_value, atom_exchange, atom_semisolid_exchange, atom_bound, atom_semisolid, hold_value, bar11, bar12, bar13, bar21, bar22, bar23, bar31, bar32, bar33, bar_free, bar_pool_b, bar_bound); + back_dt = bsk::get<0>(t162_); + back_att = bsk::get<1>(t162_); + attenuation_v = (attenuation_v + back_att); + two_pool_dt_v = (two_pool_dt_v + back_dt); + } else { + auto t163_ = _three_pool_step_adjoint_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, dt_value, nil, hold_value, nil, bsk::sum_x(e11_v), nil, bsk::sum_x(e12_v), nil, bsk::sum_x(e13_v), nil, bsk::sum_x(e21_v), nil, bsk::sum_x(e22_v), nil, bsk::sum_x(e23_v), nil, bsk::sum_x(e31_v), nil, bsk::sum_x(e32_v), nil, bsk::sum_x(e33_v), nil, bsk::sum_x(bsk::where((state == 0), zbvr, nil)), nil, bsk::sum_x(bsk::where((state == 0), poolbr, nil)), nil, bsk::sum_x(bsk::where((state == 0), semibr, nil)), nil, three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, narrow); + back_r1 = bsk::get<0>(t163_); + back_r1b = bsk::get<1>(t163_); + back_r1c = bsk::get<2>(t163_); + back_exch = bsk::get<3>(t163_); + back_sexch = bsk::get<4>(t163_); + back_bound = bsk::get<5>(t163_); + back_semi = bsk::get<6>(t163_); + back_dt = bsk::get<7>(t163_); + back_att = bsk::get<8>(t163_); + _q1 = bsk::get<9>(t163_); + _q2 = bsk::get<10>(t163_); + _q3 = bsk::get<11>(t163_); + _q4 = bsk::get<12>(t163_); + _q5 = bsk::get<13>(t163_); + _q6 = bsk::get<14>(t163_); + _q7 = bsk::get<15>(t163_); + _q8 = bsk::get<16>(t163_); + _q9 = bsk::get<17>(t163_); + g_t1v = (g_t1v + (back_r1 * bsk::truediv(-1000.0f, (atom_t1 * atom_t1)))); + g_t1bv = (g_t1bv + (back_r1b * bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)))); + g_t1cv = (g_t1cv + (back_r1c * bsk::truediv(-1000.0f, (atom_t1c * atom_t1c)))); + g_exchv = (g_exchv + back_exch); + g_sexchv = (g_sexchv + back_sexch); + g_boundv = (g_boundv + back_bound); + g_semiv = (g_semiv + back_semi); + // Both halves of the interval reach the same two, so the + // transverse pass has already put its share here. + attenuation_v = (attenuation_v + back_att); + two_pool_dt_v = (two_pool_dt_v + back_dt); + } + // All three pools take the same per-order damping and turn, so each + // collects the cotangent of the mixture that reached it. + auto t164_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + sfr = bsk::get<0>(t164_); + sfi = bsk::get<1>(t164_); + auto t165_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + sbr = bsk::get<0>(t165_); + sbi = bsk::get<1>(t165_); + auto t166_ = _complex_mul(spin_r, spin_i, mix_cr, mix_ci); + auto scr = bsk::get<0>(t166_); + auto sci = bsk::get<1>(t166_); + long_damp_v = ((((zbvr * sfr) + (zbvi * sfi)) + ((poolbr * sbr) + (poolbi * sbi))) + ((semibr * scr) + (semibi * sci))); + if (bsk::truth(moving)) { + zangle_v = ((((zbvr * (-sfi)) + (zbvi * sfr)) + ((poolbr * (-sbi)) + (poolbi * sbr))) + ((semibr * (-sci)) + (semibi * scr))); + } + auto t167_ = _complex_mul((w11 * spin_r), (-(w11 * spin_i)), zbvr, zbvi); + col_fr = bsk::get<0>(t167_); + col_fi = bsk::get<1>(t167_); + auto t168_ = _complex_mul((w21 * spin_r), (-(w21 * spin_i)), poolbr, poolbi); + part_r = bsk::get<0>(t168_); + part_i = bsk::get<1>(t168_); + auto t169_ = bsk::make_tup((col_fr + part_r), (col_fi + part_i)); + col_fr = bsk::get<0>(t169_); + col_fi = bsk::get<1>(t169_); + auto t170_ = _complex_mul((w31 * spin_r), (-(w31 * spin_i)), semibr, semibi); + part_r = bsk::get<0>(t170_); + part_i = bsk::get<1>(t170_); + auto t171_ = bsk::make_tup((col_fr + part_r), (col_fi + part_i)); + col_fr = bsk::get<0>(t171_); + col_fi = bsk::get<1>(t171_); + auto t172_ = _complex_mul((w12 * spin_r), (-(w12 * spin_i)), zbvr, zbvi); + col_br = bsk::get<0>(t172_); + col_bi = bsk::get<1>(t172_); + auto t173_ = _complex_mul((w22 * spin_r), (-(w22 * spin_i)), poolbr, poolbi); + part_r = bsk::get<0>(t173_); + part_i = bsk::get<1>(t173_); + auto t174_ = bsk::make_tup((col_br + part_r), (col_bi + part_i)); + col_br = bsk::get<0>(t174_); + col_bi = bsk::get<1>(t174_); + auto t175_ = _complex_mul((w32 * spin_r), (-(w32 * spin_i)), semibr, semibi); + part_r = bsk::get<0>(t175_); + part_i = bsk::get<1>(t175_); + auto t176_ = bsk::make_tup((col_br + part_r), (col_bi + part_i)); + col_br = bsk::get<0>(t176_); + col_bi = bsk::get<1>(t176_); + auto t177_ = _complex_mul((w13 * spin_r), (-(w13 * spin_i)), zbvr, zbvi); + col_cr = bsk::get<0>(t177_); + col_ci = bsk::get<1>(t177_); + auto t178_ = _complex_mul((w23 * spin_r), (-(w23 * spin_i)), poolbr, poolbi); + part_r = bsk::get<0>(t178_); + part_i = bsk::get<1>(t178_); + auto t179_ = bsk::make_tup((col_cr + part_r), (col_ci + part_i)); + col_cr = bsk::get<0>(t179_); + col_ci = bsk::get<1>(t179_); + auto t180_ = _complex_mul((w33 * spin_r), (-(w33 * spin_i)), semibr, semibi); + part_r = bsk::get<0>(t180_); + part_i = bsk::get<1>(t180_); + auto t181_ = bsk::make_tup((col_cr + part_r), (col_ci + part_i)); + col_cr = bsk::get<0>(t181_); + col_ci = bsk::get<1>(t181_); + auto t182_ = bsk::make_tup(col_fr, col_fi); + zbvr = bsk::get<0>(t182_); + zbvi = bsk::get<1>(t182_); + auto t183_ = bsk::make_tup(col_br, col_bi); + poolbr = bsk::get<0>(t183_); + poolbi = bsk::get<1>(t183_); + auto t184_ = bsk::make_tup(col_cr, col_ci); + semibr = bsk::get<0>(t184_); + semibi = bsk::get<1>(t184_); + } else if (bsk::truth((pools > 0))) { + // The four entries of the exchange operator and the two recoveries, + // summed over the orders that share them, then pushed back through + // the closed form once for the whole interval. + auto t185_ = _complex_mul(spin_r, spin_i, xzvr, xzvi); + spun_fr = bsk::get<0>(t185_); + spun_fi = bsk::get<1>(t185_); + auto t186_ = _complex_mul(spin_r, spin_i, xbvr, xbvi); + spun_br = bsk::get<0>(t186_); + spun_bi = bsk::get<1>(t186_); + auto bar_e11 = bsk::sum_x(((zbvr * spun_fr) + (zbvi * spun_fi))); + auto bar_e12 = bsk::sum_x(((zbvr * spun_br) + (zbvi * spun_bi))); + auto bar_e21 = bsk::sum_x(((poolbr * spun_fr) + (poolbi * spun_fi))); + auto bar_e22 = bsk::sum_x(((poolbr * spun_br) + (poolbi * spun_bi))); + auto rec_f = bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f)); + auto rec_b = bsk::sum_x(bsk::where((state == 0), poolbr, 0.0f)); + auto t187_ = _two_pool_step_adjoint_jvp(r1_value, 0.0f, r1b_value, 0.0f, atom_exchange, 0.0f, atom_bound, 0.0f, dt_value, 0.0f, wout_value, 0.0f, bar_e11, 0.0f, bar_e12, 0.0f, bar_e21, 0.0f, bar_e22, 0.0f, rec_f, 0.0f, rec_b, 0.0f); + back_r1 = bsk::get<0>(t187_); + back_r1b = bsk::get<1>(t187_); + back_exch = bsk::get<2>(t187_); + back_bound = bsk::get<3>(t187_); + back_dt = bsk::get<4>(t187_); + back_att = bsk::get<5>(t187_); + auto _t1 = bsk::get<6>(t187_); + auto _t2 = bsk::get<7>(t187_); + auto _t3 = bsk::get<8>(t187_); + auto _t4 = bsk::get<9>(t187_); + auto _t5 = bsk::get<10>(t187_); + auto _t6 = bsk::get<11>(t187_); + // r1 = 1000/t1, so a rate gradient reaches the time through the + // square of it. + g_t1v = (g_t1v + (back_r1 * bsk::truediv(-1000.0f, (atom_t1 * atom_t1)))); + g_t1bv = (g_t1bv + (back_r1b * bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)))); + g_exchv = (g_exchv + back_exch); + g_boundv = (g_boundv + back_bound); + // Both halves of the interval reach the same two, so the + // transverse pass has already put its share here. + attenuation_v = (attenuation_v + back_att); + two_pool_dt_v = (two_pool_dt_v + back_dt); + // Both pools take the same per-order damping and turn, so each + // collects the cotangent of the mixture that reached it. + auto t188_ = _complex_mul(spin_r, spin_i, mix_fr, mix_fi); + sfr = bsk::get<0>(t188_); + sfi = bsk::get<1>(t188_); + auto t189_ = _complex_mul(spin_r, spin_i, mix_br, mix_bi); + sbr = bsk::get<0>(t189_); + sbi = bsk::get<1>(t189_); + long_damp_v = (((zbvr * sfr) + (zbvi * sfi)) + ((poolbr * sbr) + (poolbi * sbi))); + if (bsk::truth(moving)) { + zangle_v = (((zbvr * (-sfi)) + (zbvi * sfr)) + ((poolbr * (-sbi)) + (poolbi * sbr))); + } + auto t190_ = _complex_mul((pe11 * spin_r), (-(pe11 * spin_i)), zbvr, zbvi); + back_zr = bsk::get<0>(t190_); + back_zi = bsk::get<1>(t190_); + auto t191_ = _complex_mul((pe21 * spin_r), (-(pe21 * spin_i)), poolbr, poolbi); + auto cross_zr = bsk::get<0>(t191_); + auto cross_zi = bsk::get<1>(t191_); + auto t192_ = _complex_mul((pe12 * spin_r), (-(pe12 * spin_i)), zbvr, zbvi); + auto back_br = bsk::get<0>(t192_); + auto back_bi = bsk::get<1>(t192_); + auto t193_ = _complex_mul((pe22 * spin_r), (-(pe22 * spin_i)), poolbr, poolbi); + auto cross_br = bsk::get<0>(t193_); + auto cross_bi = bsk::get<1>(t193_); + poolbr = (back_br + cross_br); + poolbi = (back_bi + cross_bi); + zbvr = (back_zr + cross_zr); + zbvi = (back_zi + cross_zi); + } else { + auto spun = _complex_mul(szr, szi, xzvr, xzvi); + e1_v = ((zbvr * bsk::get<0>(spun)) + (zbvi * bsk::get<1>(spun))); + grad_e1_v = bsk::sum_x((e1_v * damp_z)); + grad_e1_v = (grad_e1_v - bsk::sum_x(bsk::where((state == 0), zbvr, 0.0f))); + // The longitudinal states turn too, and by a whole order rather + // than the transverse half-order more. + if (bsk::truth(moving)) { + auto zo = _complex_mul(lvr, lvi, xzvr, xzvi); + zangle_v = ((zbvr * (-bsk::get<1>(zo))) + (zbvi * bsk::get<0>(zo))); + } + long_damp_v = ((e1_v * bare1_value) * damp_z); + auto t194_ = _complex_mul(lvr, (-lvi), zbvr, zbvi); + zbvr = bsk::get<0>(t194_); + zbvi = bsk::get<1>(t194_); + } + spread_v = zero; + if (bsk::truth(diffusing)) { + // The rate and the interval multiply every order's b-weight, so + // both take a weighted sum rather than one scalar. Order zero + // carries no longitudinal weight, which keeps recovery out of this. + auto weighted_v = ((long_damp_v * longitudinal_weight) + (cot2_v * transverse_weight)); + spread_v = bsk::sum_x(weighted_v); + g_diffv = (g_diffv + ((-spread_v) * dt_value)); + } + wound_v = zero; + wash_v = zero; + if (bsk::truth(moving)) { + wound_v = bsk::sum_x(((per_angle_v * (order + 0.5f)) + (zangle_v * order))); + g_flowv = (g_flowv + ((-wound_v) * dt_value)); + // Washout scales both relaxation factors, so its gradient is the + // one they already carry, taken against the factors before that + // scaling. Past the clamp the interval has replaced the voxel + // outright and nothing further depends on the rate. + auto live = bsk::cast(((atom_washout * dt_value) < 1.0f)); + auto transverse_dry = bsk::select(bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3)))), zero, (grad_e2_v * dry2_value)); + wash_v = ((-live) * (((grad_e1_v * dry1_value) + transverse_dry) + attenuation_v)); + g_washv = (g_washv + (wash_v * dt_value)); + } + if (bsk::truth((bsk::truth((pools != 2)) && bsk::truth((pools != 3))))) { + auto t195_ = _complex_mul(ovr, (-ovi), pbvr, pbvi); + pbvr = bsk::get<0>(t195_); + pbvi = bsk::get<1>(t195_); + auto t196_ = _complex_mul(ovr, ovi, mbvr, mbvi); + mbvr = bsk::get<0>(t196_); + mbvi = bsk::get<1>(t196_); + } + auto inverse1_value = bsk::truediv(1000.0f, (atom_t1 * atom_t1)); + auto inverse2_value = bsk::truediv(1000.0f, (atom_t2 * atom_t2)); + g_t1v = (g_t1v + (grad_e1_v * ((bare1_value * dt_value) * inverse1_value))); + if (bsk::truth((bsk::truth((pools != 2)) && bsk::truth((pools != 3))))) { + g_t2v = (g_t2v + (grad_e2_v * ((bare2_value * dt_value) * inverse2_value))); + } + auto turn = -6.283185307179586f; + g_b0v = (g_b0v + (grad_angle_v * (turn * dt_value))); + duration_v = ((-grad_e1_v) * (r1_value * bare1_value)); + if (bsk::truth((bsk::truth((pools != 2)) && bsk::truth((pools != 3))))) { + duration_v = (duration_v - (grad_e2_v * (r2_value * bare2_value))); + } + duration_v = (duration_v + ((grad_angle_v * (turn * atom_b0)) + two_pool_dt_v)); + duration_v = (duration_v + (((-spread_v) * atom_damping) - (wound_v * atom_flow))); + duration_v = (duration_v + (wash_v * atom_washout)); + bsk::atomic_add(((grad_duration + event_base) + event), duration_v, active_atom); + } + if (bsk::truth((bsk::truth((pools == 3)) && bsk::truth(tabulated)))) { + // One closed form per distinct length rather than one per event. The + // walk back pooled the cotangents the eigenvalues are pushed through, + // and the closed form is linear in them, so the pieces of the sum are + // the sum of the pieces. + for (std::int64_t row = 0; row < row_count; row += 1) { + held = (pool_bars + (((local * row_count) + row) * 12)); + auto row_dt = (bsk::ld((pool_durations + row)) + zero); + auto one_att = bsk::select(bsk::truth(moving), _washout(atom_washout, row_dt), (1.0f + (0.0f * row_dt))); + nil = (0.0f * row_dt); + auto t197_ = _three_pool_pieces_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, row_dt, nil, narrow); + three_free = bsk::get<0>(t197_); + three_d_free = bsk::get<1>(t197_); + three_pool_b = bsk::get<2>(t197_); + three_d_pool_b = bsk::get<3>(t197_); + three_pool_c = bsk::get<4>(t197_); + three_d_pool_c = bsk::get<5>(t197_); + three_a00 = bsk::get<6>(t197_); + three_d_a00 = bsk::get<7>(t197_); + three_a01 = bsk::get<8>(t197_); + three_d_a01 = bsk::get<9>(t197_); + three_a02 = bsk::get<10>(t197_); + three_d_a02 = bsk::get<11>(t197_); + three_a10 = bsk::get<12>(t197_); + three_d_a10 = bsk::get<13>(t197_); + three_a11 = bsk::get<14>(t197_); + three_d_a11 = bsk::get<15>(t197_); + three_a20 = bsk::get<16>(t197_); + three_d_a20 = bsk::get<17>(t197_); + three_a22 = bsk::get<18>(t197_); + three_d_a22 = bsk::get<19>(t197_); + three_s00 = bsk::get<20>(t197_); + three_d_s00 = bsk::get<21>(t197_); + three_s11 = bsk::get<22>(t197_); + three_d_s11 = bsk::get<23>(t197_); + three_s22 = bsk::get<24>(t197_); + three_d_s22 = bsk::get<25>(t197_); + three_minors = bsk::get<26>(t197_); + three_d_minors = bsk::get<27>(t197_); + three_sum_flat = bsk::get<28>(t197_); + three_sum_linear = bsk::get<29>(t197_); + three_sum_square = bsk::get<30>(t197_); + three_d_sum_flat = bsk::get<31>(t197_); + three_d_sum_linear = bsk::get<32>(t197_); + three_d_sum_square = bsk::get<33>(t197_); + three_lift = bsk::get<34>(t197_); + three_d_lift = bsk::get<35>(t197_); + three_low = bsk::get<36>(t197_); + three_middle = bsk::get<37>(t197_); + three_d_low = bsk::get<38>(t197_); + three_d_middle = bsk::get<39>(t197_); + three_leading = bsk::get<40>(t197_); + three_d_leading = bsk::get<41>(t197_); + three_first = bsk::get<42>(t197_); + three_d_first = bsk::get<43>(t197_); + three_second = bsk::get<44>(t197_); + three_d_second = bsk::get<45>(t197_); + three_determinant = bsk::get<46>(t197_); + three_d_determinant = bsk::get<47>(t197_); + three_high = bsk::get<48>(t197_); + three_d_high = bsk::get<49>(t197_); + three_radius = bsk::get<50>(t197_); + three_d_radius = bsk::get<51>(t197_); + three_cube = bsk::get<52>(t197_); + three_raw = bsk::get<53>(t197_); + three_d_raw = bsk::get<54>(t197_); + three_argument = bsk::get<55>(t197_); + three_inside_limit = bsk::get<56>(t197_); + three_angle = bsk::get<57>(t197_); + three_d_angle = bsk::get<58>(t197_); + three_centre = bsk::get<59>(t197_); + three_d_centre = bsk::get<60>(t197_); + three_trailing = bsk::get<61>(t197_); + three_d_trailing = bsk::get<62>(t197_); + three_guarded = bsk::get<63>(t197_); + three_d_guarded = bsk::get<64>(t197_); + three_q00 = bsk::get<65>(t197_); + three_d_q00 = bsk::get<66>(t197_); + three_q01 = bsk::get<67>(t197_); + three_d_q01 = bsk::get<68>(t197_); + three_q02 = bsk::get<69>(t197_); + three_d_q02 = bsk::get<70>(t197_); + three_q10 = bsk::get<71>(t197_); + three_d_q10 = bsk::get<72>(t197_); + three_q11 = bsk::get<73>(t197_); + three_d_q11 = bsk::get<74>(t197_); + three_q12 = bsk::get<75>(t197_); + three_d_q12 = bsk::get<76>(t197_); + three_q20 = bsk::get<77>(t197_); + three_d_q20 = bsk::get<78>(t197_); + three_q21 = bsk::get<79>(t197_); + three_d_q21 = bsk::get<80>(t197_); + three_q22 = bsk::get<81>(t197_); + three_d_q22 = bsk::get<82>(t197_); + auto t198_ = _three_pool_assemble_jvp(three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, narrow); + three_def_00 = bsk::get<0>(t198_); + three_dif_00 = bsk::get<1>(t198_); + three_def_01 = bsk::get<2>(t198_); + three_dif_01 = bsk::get<3>(t198_); + three_def_02 = bsk::get<4>(t198_); + three_dif_02 = bsk::get<5>(t198_); + three_def_10 = bsk::get<6>(t198_); + three_dif_10 = bsk::get<7>(t198_); + three_def_11 = bsk::get<8>(t198_); + three_dif_11 = bsk::get<9>(t198_); + three_def_12 = bsk::get<10>(t198_); + three_dif_12 = bsk::get<11>(t198_); + three_def_20 = bsk::get<12>(t198_); + three_dif_20 = bsk::get<13>(t198_); + three_def_21 = bsk::get<14>(t198_); + three_dif_21 = bsk::get<15>(t198_); + three_def_22 = bsk::get<16>(t198_); + three_dif_22 = bsk::get<17>(t198_); + auto t199_ = _three_pool_step_adjoint_jvp(r1_value, nil, r1b_value, nil, r1c_value, nil, atom_exchange, nil, atom_semisolid_exchange, nil, atom_bound, nil, atom_semisolid, nil, row_dt, nil, one_att, nil, bsk::ld((held + 0), active_atom, 0.0f), nil, bsk::ld((held + 1), active_atom, 0.0f), nil, bsk::ld((held + 2), active_atom, 0.0f), nil, bsk::ld((held + 3), active_atom, 0.0f), nil, bsk::ld((held + 4), active_atom, 0.0f), nil, bsk::ld((held + 5), active_atom, 0.0f), nil, bsk::ld((held + 6), active_atom, 0.0f), nil, bsk::ld((held + 7), active_atom, 0.0f), nil, bsk::ld((held + 8), active_atom, 0.0f), nil, bsk::ld((held + 9), active_atom, 0.0f), nil, bsk::ld((held + 10), active_atom, 0.0f), nil, bsk::ld((held + 11), active_atom, 0.0f), nil, three_free, three_d_free, three_pool_b, three_d_pool_b, three_pool_c, three_d_pool_c, three_a00, three_d_a00, three_a01, three_d_a01, three_a02, three_d_a02, three_a10, three_d_a10, three_a11, three_d_a11, three_a20, three_d_a20, three_a22, three_d_a22, three_s00, three_d_s00, three_s11, three_d_s11, three_s22, three_d_s22, three_minors, three_d_minors, three_sum_flat, three_sum_linear, three_sum_square, three_d_sum_flat, three_d_sum_linear, three_d_sum_square, three_lift, three_d_lift, three_low, three_middle, three_d_low, three_d_middle, three_leading, three_d_leading, three_first, three_d_first, three_second, three_d_second, three_determinant, three_d_determinant, three_high, three_d_high, three_radius, three_d_radius, three_cube, three_raw, three_d_raw, three_argument, three_inside_limit, three_angle, three_d_angle, three_centre, three_d_centre, three_trailing, three_d_trailing, three_guarded, three_d_guarded, three_q00, three_d_q00, three_q01, three_d_q01, three_q02, three_d_q02, three_q10, three_d_q10, three_q11, three_d_q11, three_q12, three_d_q12, three_q20, three_d_q20, three_q21, three_d_q21, three_q22, three_d_q22, three_def_00, three_dif_00, three_def_01, three_dif_01, three_def_02, three_dif_02, three_def_10, three_dif_10, three_def_11, three_dif_11, three_def_12, three_dif_12, three_def_20, three_dif_20, three_def_21, three_dif_21, three_def_22, three_dif_22, narrow); + back_r1 = bsk::get<0>(t199_); + back_r1b = bsk::get<1>(t199_); + back_r1c = bsk::get<2>(t199_); + back_exch = bsk::get<3>(t199_); + back_sexch = bsk::get<4>(t199_); + back_bound = bsk::get<5>(t199_); + back_semi = bsk::get<6>(t199_); + back_dt = bsk::get<7>(t199_); + back_att = bsk::get<8>(t199_); + _q1 = bsk::get<9>(t199_); + _q2 = bsk::get<10>(t199_); + _q3 = bsk::get<11>(t199_); + _q4 = bsk::get<12>(t199_); + _q5 = bsk::get<13>(t199_); + _q6 = bsk::get<14>(t199_); + _q7 = bsk::get<15>(t199_); + _q8 = bsk::get<16>(t199_); + _q9 = bsk::get<17>(t199_); + g_t1v = (g_t1v + (back_r1 * bsk::truediv(-1000.0f, (atom_t1 * atom_t1)))); + g_t1bv = (g_t1bv + (back_r1b * bsk::truediv(-1000.0f, (atom_t1b * atom_t1b)))); + g_t1cv = (g_t1cv + (back_r1c * bsk::truediv(-1000.0f, (atom_t1c * atom_t1c)))); + g_exchv = (g_exchv + back_exch); + g_sexchv = (g_sexchv + back_sexch); + g_boundv = (g_boundv + back_bound); + g_semiv = (g_semiv + back_semi); + } + } + auto velocity_v = ((g_flowv * flow_scale) + ((g_washv * direction) * washout_scale)); + auto values = bsk::make_tup(g_t1v, g_t2v, g_m0v, g_b1v, g_b1pv, g_b0v, g_invv, g_diffv, velocity_v); + if (bsk::truth((pools > 0))) { + // The fraction also sets where each pool starts, which the walk back + // reaches last. + g_boundv = (g_boundv + bsk::sum_x(bsk::where((state == 0), (poolbr - zbvr), 0.0f))); + } + if (bsk::truth((pools == 3))) { + g_semiv = (g_semiv + bsk::sum_x(bsk::where((state == 0), (semibr - zbvr), 0.0f))); + auto semisolid_row = (9 + (2 * (shim_rows - 1))); + bsk::atomic_add(((grad_tissue + (semisolid_row * atom_count)) + atom), g_semiv, active_atom); + bsk::atomic_add(((grad_tissue + ((semisolid_row + 1) * atom_count)) + atom), g_sexchv, active_atom); + bsk::atomic_add(((grad_tissue + ((semisolid_row + 2) * atom_count)) + atom), g_t1cv, active_atom); + } + if (bsk::truth((bsk::truth((pools == 2)) || bsk::truth((pools == 3))))) { + base_row = (12 + (2 * (shim_rows - 1))); + bsk::atomic_add(((grad_tissue + (base_row * atom_count)) + atom), g_boundv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 1) * atom_count)) + atom), g_exchv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 2) * atom_count)) + atom), g_t1bv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 3) * atom_count)) + atom), g_t2bv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 4) * atom_count)) + atom), g_shiftv, active_atom); + } + if (bsk::truth((pools == 1))) { + base_row = (9 + (2 * (shim_rows - 1))); + bsk::atomic_add(((grad_tissue + (base_row * atom_count)) + atom), g_boundv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 1) * atom_count)) + atom), g_exchv, active_atom); + bsk::atomic_add(((grad_tissue + ((base_row + 2) * atom_count)) + atom), g_t1bv, active_atom); + } + bsk::static_for<0, 9, 1>([&](auto parameter_c) { + constexpr std::int64_t parameter = decltype(parameter_c)::value; + // The transmit pair went to its shim's row above when there is more + // than one; the rest sit past whatever rows that pair took. + if (bsk::truth((bsk::truth((!bsk::truth(shimmed))) || bsk::truth((bsk::truth((parameter != 3)) && bsk::truth((parameter != 4))))))) { + auto plane = bsk::select(bsk::truth((parameter < 3)), parameter, (parameter + (2 * (shim_rows - 1)))); + bsk::atomic_add(((grad_tissue + (plane * atom_count)) + atom), bsk::get(values), active_atom); + } + }); +} + +BSK_HD void _epg_real_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_duration, float* tangent_flip, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { + bsk::V atom_b1{}; + bsk::V atom_damping{}; + bsk::V atom_inversion{}; + bsk::V atom_m0{}; + bsk::V damp_t{}; + bsk::V damp_z{}; + bsk::V ddamp_t{}; + bsk::V ddamp_z{}; + bsk::V dot_b1{}; + bsk::V dot_damping{}; + bsk::V dot_e1{}; + bsk::V dot_e2{}; + bsk::V dot_inversion{}; + bsk::V dot_longitudinal{}; + bsk::V dot_m0{}; + bsk::V dot_minus{}; + bsk::V dot_plus{}; + bsk::V e1{}; + bsk::V e2{}; + bsk::V longitudinal{}; + bsk::V minus{}; + bsk::V plus{}; + bsk::V pulse_b1{}; + bsk::V pulse_dot_b1{}; + bsk::V rotated_dm{}; + bsk::V rotated_dp{}; + bsk::V rotated_dz{}; + auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); + auto state = bsk::arange_x(); + auto active_atom = (problem < (train_count * atom_count)); + auto state_mask = bsk::band((state < state_count), active_atom); + auto atom = bsk::mod(problem, atom_count); + // A property given as one value for the whole tissue is read at one + // address by every voxel, which is a stride of zero through it. + auto scalar_atom = (atom * atom_stride); + auto train = bsk::floordiv(problem, atom_count); + auto empty = bsk::full(0); + plus = empty; + minus = empty; + longitudinal = (empty + bsk::where((state == 0), 1.0f, 0.0f)); + dot_plus = empty; + dot_minus = empty; + dot_longitudinal = empty; + auto atom_t1 = bsk::ld((t1 + atom), active_atom, 1.0f); + auto atom_t2 = bsk::ld((t2 + atom), active_atom, 1.0f); + atom_m0 = 1.0f; + if (bsk::truth(density)) { + atom_m0 = bsk::ld((m0 + scalar_atom), active_atom, 0.0f); + } + atom_b1 = 1.0f; + if (bsk::truth(transmit)) { + atom_b1 = bsk::ld((b1 + scalar_atom), active_atom, 1.0f); + } + atom_inversion = 1.0f; + if (bsk::truth(inverting)) { + atom_inversion = bsk::ld((inversion_efficiency + scalar_atom), active_atom, 1.0f); + } + auto dot_t1 = bsk::ld((tangent_t1 + atom), active_atom, 0.0f); + auto dot_t2 = bsk::ld((tangent_t2 + atom), active_atom, 0.0f); + dot_m0 = 0.0f; + if (bsk::truth(density)) { + dot_m0 = bsk::ld((tangent_m0 + scalar_atom), active_atom, 0.0f); + } + dot_b1 = 0.0f; + if (bsk::truth(transmit)) { + dot_b1 = bsk::ld((tangent_b1 + scalar_atom), active_atom, 0.0f); + } + dot_inversion = 0.0f; + if (bsk::truth(inverting)) { + dot_inversion = bsk::ld((tangent_inversion_efficiency + scalar_atom), active_atom, 0.0f); + } + auto rate1 = bsk::truediv(1000.0f, atom_t1); + auto rate2 = bsk::truediv(1000.0f, atom_t2); + atom_damping = 0.0f; + dot_damping = 0.0f; + if (bsk::truth(diffusing)) { + atom_damping = bsk::ld((diffusion + scalar_atom), active_atom, 0.0f); + dot_damping = bsk::ld((tangent_diffusion + scalar_atom), active_atom, 0.0f); + } + auto order = bsk::cast(state); + auto event_base = (train * event_count); + // Two events to an iteration. A repetition is several events -- a pulse, + // a sample, an interval -- so the loop runs longer than the sequence is + // repetitions, and unrolling lets one back-edge and one set of event + // bookkeeping serve two of them. Two is where it stops paying: four was + // measured slower, and the body is already large enough that widening it + // costs registers. + for (std::int64_t event = 0; event < event_count; event += 1) { + auto dt = _event_value(duration, event_base, event, active_atom, single_train); + auto dot_dt = _event_value(tangent_duration, event_base, event, active_atom, single_train); + // An event of no duration relaxes nothing, and carries no tangent along + // the relaxation either: both factors are one and both their derivatives + // are zero. + if (bsk::truth((bsk::truth((bsk::max_all(dt) != 0.0f)) || bsk::truth((bsk::max_all(dot_dt) != 0.0f))))) { + e1 = bsk::exp(((-rate1) * dt)); + e2 = bsk::exp(((-rate2) * dt)); + dot_e1 = (e1 * (bsk::truediv(((1000.0f * dt) * dot_t1), (atom_t1 * atom_t1)) - (rate1 * dot_dt))); + dot_e2 = (e2 * (bsk::truediv(((1000.0f * dt) * dot_t2), (atom_t2 * atom_t2)) - (rate2 * dot_dt))); + damp_z = 1.0f; + ddamp_z = 0.0f; + damp_t = 1.0f; + ddamp_t = 0.0f; + if (bsk::truth(diffusing)) { + auto t0_ = _damping_jvp(atom_damping, dot_damping, dt, dot_dt, order); + damp_z = bsk::get<0>(t0_); + ddamp_z = bsk::get<1>(t0_); + damp_t = bsk::get<2>(t0_); + ddamp_t = bsk::get<3>(t0_); + } + // Order zero is undamped, so the recovery keeps the bare factor. + auto t1_ = bsk::make_tup((1.0f - e1), (-dot_e1)); + auto recovery = bsk::get<0>(t1_); + auto dot_recovery = bsk::get<1>(t1_); + dot_e1 = ((dot_e1 * damp_z) + (e1 * ddamp_z)); + e1 = (e1 * damp_z); + dot_e2 = ((dot_e2 * damp_t) + (e2 * ddamp_t)); + e2 = (e2 * damp_t); + dot_plus = ((dot_plus * e2) + (plus * dot_e2)); + dot_minus = ((dot_minus * e2) + (minus * dot_e2)); + dot_longitudinal = ((dot_longitudinal * e1) + (longitudinal * dot_e1)); + dot_longitudinal = (dot_longitudinal + bsk::where((state == 0), dot_recovery, 0.0f)); + plus = (plus * e2); + minus = (minus * e2); + longitudinal = ((longitudinal * e1) + bsk::where((state == 0), recovery, 0.0f)); + } + // Every flag below is read from a per-event array with no atom index, so + // it is uniform across the program and can steer real control flow. + auto event_action = bsk::cast(bsk::ld((action + event))); + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t2_ = _shift_real(plus, minus, state, state_mask, state_count); + plus = bsk::get<0>(t2_); + minus = bsk::get<1>(t2_); + auto t3_ = _shift_real(dot_plus, dot_minus, state, state_mask, state_count); + dot_plus = bsk::get<0>(t3_); + dot_minus = bsk::get<1>(t3_); + } + auto event_kind = bsk::ld((kind + event)); + auto is_rf = (event_kind == 1); + auto is_inversion = (bsk::band(event_action, 4) != 0); + if (bsk::truth((bsk::truth(is_rf) && bsk::truth(is_inversion)))) { + dot_longitudinal = (((-atom_inversion) * dot_longitudinal) - (dot_inversion * longitudinal)); + longitudinal = ((-atom_inversion) * longitudinal); + } else if (bsk::truth(is_rf)) { + auto event_flip = _event_value(flip, event_base, event, active_atom, single_train); + auto dot_flip = _event_value(tangent_flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + pulse_dot_b1 = dot_b1; + // One shim is the whole sequence's transmit field, loaded once above; + // several give each pulse the row of the shim it drives. + if (bsk::truth(shimmed)) { + auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + if (bsk::truth(transmit)) { + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + pulse_dot_b1 = bsk::ld(((tangent_b1 + shim_row) + atom), active_atom, 0.0f); + } + auto alpha = (event_flip * pulse_b1); + auto dot_alpha = ((dot_flip * pulse_b1) + (event_flip * pulse_dot_b1)); + auto cosine = bsk::cos(alpha); + auto sine = bsk::sin(alpha); + auto cosine_half_sq = (0.5f * (1.0f + cosine)); + auto sine_half_sq = (0.5f * (1.0f - cosine)); + auto half_sine = (0.5f * sine); + auto dot_cosine = ((-sine) * dot_alpha); + auto dot_sine = (cosine * dot_alpha); + auto dot_cosine_half_sq = ((-0.5f * sine) * dot_alpha); + auto dot_sine_half_sq = ((0.5f * sine) * dot_alpha); + auto dot_half_sine = ((0.5f * cosine) * dot_alpha); + rotated_dp = ((cosine_half_sq * dot_plus) + (dot_cosine_half_sq * plus)); + rotated_dp = (rotated_dp + ((sine_half_sq * dot_minus) + (dot_sine_half_sq * minus))); + rotated_dp = (rotated_dp - ((sine * dot_longitudinal) + (dot_sine * longitudinal))); + rotated_dm = ((sine_half_sq * dot_plus) + (dot_sine_half_sq * plus)); + rotated_dm = (rotated_dm + ((cosine_half_sq * dot_minus) + (dot_cosine_half_sq * minus))); + rotated_dm = (rotated_dm + ((sine * dot_longitudinal) + (dot_sine * longitudinal))); + rotated_dz = ((half_sine * dot_plus) + (dot_half_sine * plus)); + rotated_dz = (rotated_dz - ((half_sine * dot_minus) + (dot_half_sine * minus))); + rotated_dz = (rotated_dz + ((cosine * dot_longitudinal) + (dot_cosine * longitudinal))); + auto rotated_p = (((cosine_half_sq * plus) + (sine_half_sq * minus)) - (sine * longitudinal)); + auto rotated_m = (((sine_half_sq * plus) + (cosine_half_sq * minus)) + (sine * longitudinal)); + auto rotated_z = (((half_sine * plus) - (half_sine * minus)) + (cosine * longitudinal)); + plus = rotated_p; + minus = rotated_m; + longitudinal = rotated_z; + dot_plus = rotated_dp; + dot_minus = rotated_dm; + dot_longitudinal = rotated_dz; + } + if (bsk::truth((bsk::truth((bsk::band(event_action, 32) != 0)) && bsk::truth((event_kind == 2))))) { + auto out_ = bsk::ld((output_index + event)); + auto output_offset = ((problem * output_count) + out_); + auto output_mask = bsk::band(bsk::band(active_atom, (state == 0)), (out_ >= 0)); + auto signal_imag = ((dot_m0 * plus) + (atom_m0 * dot_plus)); + bsk::st(((output_real + output_offset) + state), empty, output_mask); + bsk::st(((output_imag + output_offset) + state), signal_imag, output_mask); + } + if (bsk::truth((bsk::truth((bsk::band(event_action, 2) != 0)) || bsk::truth((bsk::band(event_action, 16) != 0))))) { + auto t4_ = _shift_real(plus, minus, state, state_mask, state_count); + plus = bsk::get<0>(t4_); + minus = bsk::get<1>(t4_); + auto t5_ = _shift_real(dot_plus, dot_minus, state, state_mask, state_count); + dot_plus = bsk::get<0>(t5_); + dot_minus = bsk::get<1>(t5_); + } + if (bsk::truth((bsk::band(event_action, 8) != 0))) { + plus = empty; + minus = empty; + dot_plus = empty; + dot_minus = empty; + } + } +} diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu new file mode 100644 index 00000000..476c3747 --- /dev/null +++ b/src/blochsim/_gpu.cu @@ -0,0 +1,99 @@ +// The GPU kernels, compiled ahead of time for the card. +// +// One CUDA block runs one program: its threads hold a tile an element each, +// x along a row and y across the rows. The module links the CUDA runtime +// statically, so a machine needs the driver and nothing else, and it calls no +// PyTorch API: a launch takes the addresses of the tensors' data and the +// stream PyTorch is queueing on. +// +// The kernels themselves are compiled one to a file, from _gpu_kernel.cu.in. +#define BLOCHSIM_TABLE_ONLY 1 + +#include "_launch.hpp" + +#include + +#define BLOCHSIM_DEVICE_ENTRY(name) \ + __global__ void kernel##name##_256(bsk::Arguments arguments, int z); \ + __global__ void kernel##name##_1024(bsk::Arguments arguments, int z); +BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_ENTRY) +#undef BLOCHSIM_DEVICE_ENTRY + +namespace { + +// Each kernel bounded to 256 threads, then to 1024. +#define BLOCHSIM_DEVICE_POINTER(name) \ + {reinterpret_cast(&kernel##name##_256), \ + reinterpret_cast(&kernel##name##_1024)}, +const void* const KERNEL_FUNCTIONS[][2] = {BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_POINTER)}; +#undef BLOCHSIM_DEVICE_POINTER + +PyObject* cuda_error(cudaError_t status, const char* what) { + PyErr_Format(PyExc_RuntimeError, "%s: %s", what, cudaGetErrorString(status)); + return nullptr; +} + +PyObject* launch(PyObject*, PyObject* args) { + PyObject* name = nullptr; + PyObject* grid = nullptr; + PyObject* values = nullptr; + int device = 0; + unsigned long long stream = 0; + if (!PyArg_ParseTuple(args, "UOOiK", &name, &grid, &values, &device, &stream)) { + return nullptr; + } + blochsim_launch::Launch request; + if (!blochsim_launch::read_launch(name, grid, values, request)) { + return nullptr; + } + if (request.grid[0] > 2147483647LL || request.grid[1] > 65535) { + PyErr_SetString(PyExc_ValueError, "the grid is larger than a launch can hold"); + return nullptr; + } + if (request.grid[0] == 0 || request.grid[1] == 0) { + Py_RETURN_NONE; + } + // The caller's current device is put back, whatever this one was. + int previous = 0; + cudaError_t status = cudaGetDevice(&previous); + if (status == cudaSuccess && previous != device) { + status = cudaSetDevice(device); + } + if (status != cudaSuccess) { + return cuda_error(status, "selecting the device"); + } + const dim3 blocks(static_cast(request.grid[0]), static_cast(request.grid[1])); + const dim3 threads(static_cast(request.block[0]), static_cast(request.block[1])); + // A row's reduction and gather go through one word per thread, and a + // product with an operator over pools through as many again. + const std::size_t shared = + sizeof(unsigned long long) * (threads.x * threads.y + bsk::MAX_Z * bsk::MAX_Z); + void* parameters[] = {&request.arguments, &request.z}; + const int bounded = threads.x * threads.y > 256 ? 1 : 0; + status = cudaLaunchKernel(KERNEL_FUNCTIONS[request.kernel][bounded], blocks, threads, parameters, + shared, reinterpret_cast(stream)); + if (previous != device) { + cudaSetDevice(previous); + } + if (status != cudaSuccess) { + return cuda_error(status, bsk::KERNELS[request.kernel].name); + } + Py_RETURN_NONE; +} + +PyMethodDef METHODS[] = { + {"kernels", blochsim_launch::kernel_table, METH_NOARGS, + "Each kernel's parameter names and kinds."}, + {"launch", launch, METH_VARARGS, + "Queue a kernel over a grid of programs on a device's stream."}, + {nullptr, nullptr, 0, nullptr}, +}; + +PyModuleDef MODULE = { + PyModuleDef_HEAD_INIT, "_gpu", "The GPU kernels, compiled for the card.", -1, METHODS, + nullptr, nullptr, nullptr, nullptr, +}; + +} // namespace + +PyMODINIT_FUNC PyInit__gpu(void) { return PyModule_Create(&MODULE); } diff --git a/src/blochsim/_gpu_host.cpp b/src/blochsim/_gpu_host.cpp new file mode 100644 index 00000000..03a1c448 --- /dev/null +++ b/src/blochsim/_gpu_host.cpp @@ -0,0 +1,66 @@ +// The GPU kernels run on the host, one program at a time, over host buffers. +// +// The same source the CUDA build compiles, with a tile held as a whole array +// rather than an element per thread. It is how the kernels are checked on a +// machine with no card; nothing in the package dispatches to it otherwise. +#include "_launch.hpp" + +#include + +namespace { + +using host_call = void (*)(const bsk::Arg*); + +#define BLOCHSIM_HOST_ENTRY(name) &bsk::call##name, +constexpr host_call CALLS[] = {BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_HOST_ENTRY)}; +#undef BLOCHSIM_HOST_ENTRY + +PyObject* launch(PyObject*, PyObject* args) { + PyObject* name = nullptr; + PyObject* grid = nullptr; + PyObject* values = nullptr; + if (!PyArg_ParseTuple(args, "UOO", &name, &grid, &values)) { + return nullptr; + } + blochsim_launch::Launch request; + if (!blochsim_launch::read_launch(name, grid, values, request)) { + return nullptr; + } + bool failed = false; + Py_BEGIN_ALLOW_THREADS + try { + bsk::program.nx = request.block[0]; + bsk::program.ny = request.block[1]; + bsk::program.nz = request.z; + for (std::int64_t y = 0; y < request.grid[1]; ++y) { + for (std::int64_t x = 0; x < request.grid[0]; ++x) { + bsk::program.pid[0] = x; + bsk::program.pid[1] = y; + CALLS[request.kernel](request.arguments.a); + } + } + } catch (const std::bad_alloc&) { + failed = true; + } + Py_END_ALLOW_THREADS + if (failed) { + return PyErr_NoMemory(); + } + Py_RETURN_NONE; +} + +PyMethodDef METHODS[] = { + {"kernels", blochsim_launch::kernel_table, METH_NOARGS, + "Each kernel's parameter names and kinds."}, + {"launch", launch, METH_VARARGS, "Run a kernel over a grid of programs on the host."}, + {nullptr, nullptr, 0, nullptr}, +}; + +PyModuleDef MODULE = { + PyModuleDef_HEAD_INIT, "_gpu_host", "The GPU kernels, run on the host.", -1, METHODS, + nullptr, nullptr, nullptr, nullptr, +}; + +} // namespace + +PyMODINIT_FUNC PyInit__gpu_host(void) { return PyModule_Create(&MODULE); } diff --git a/src/blochsim/_gpu_kernel.cu.in b/src/blochsim/_gpu_kernel.cu.in new file mode 100644 index 00000000..233418f2 --- /dev/null +++ b/src/blochsim/_gpu_kernel.cu.in @@ -0,0 +1,20 @@ +// One kernel, compiled for the card. CMake writes one of these per kernel, +// so the kernels compile in parallel; @NAME@ is the kernel's name. +// +// Each comes twice, bounded to 256 threads and to 1024. The register file is +// divided among a block's threads, so the bound decides how many registers a +// thread may keep before it spills; _gpu.cu launches the narrower variant +// whenever a block fits it. +#define BLOCHSIM_SIMT 1 + +#include "_kernels.hpp" + +__global__ void __launch_bounds__(256) kernel@NAME@_256(bsk::Arguments arguments, int z) { + bsk::enter(z); + bsk::call@NAME@(arguments.a); +} + +__global__ void __launch_bounds__(1024) kernel@NAME@_1024(bsk::Arguments arguments, int z) { + bsk::enter(z); + bsk::call@NAME@(arguments.a); +} diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py new file mode 100644 index 00000000..0235a3fe --- /dev/null +++ b/src/blochsim/_gpu_launch.py @@ -0,0 +1,106 @@ +"""Launching the compiled GPU kernels from the tensors a simulation holds. + +A kernel is named and indexed with its grid, then called with its arguments in +the order its signature lists them, positionally or by name. Tensors on a card +are queued on PyTorch's current stream for that card; tensors on the host run +the same kernel source compiled for the host, one program at a time, which is +how the kernels are checked without a card. +""" + +from __future__ import annotations + +__all__: list[str] = [] + +from functools import cache +from typing import Any + +import torch + + +def next_power_of_2(value: int) -> int: + """The smallest power of two no less than ``value``.""" + return 1 if value <= 1 else 1 << (int(value) - 1).bit_length() + + +def cdiv(numerator: int, denominator: int) -> int: + """``numerator / denominator`` rounded up.""" + return -(-int(numerator) // int(denominator)) + + +@cache +def _module(device_type: str) -> Any: + if device_type == "cuda": + from blochsim import _gpu + + return _gpu + from blochsim import _gpu_host + + return _gpu_host + + +@cache +def _signature(name: str) -> tuple[tuple[str, ...], str]: + params, kinds = _module("cuda" if available() else "cpu").kernels()[name] + return tuple(params.split(",")), kinds + + +@cache +def available() -> bool: + """Whether this installation carries the kernels compiled for a card.""" + try: + _module("cuda") + except ImportError: + return False + return True + + +class Kernel: + """A compiled kernel, launched as ``kernel[grid](*arguments)``.""" + + def __init__(self, name: str) -> None: + self.name = name + + def __getitem__(self, grid: tuple[int, ...]) -> Any: + def run(*args: Any, **kwargs: Any) -> None: + self.launch(tuple(int(count) for count in grid), args, kwargs) + + return run + + def launch( + self, grid: tuple[int, ...], args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> None: + names, kinds = _signature(self.name) + values = list(args) + for name in names[len(args) :]: + if name not in kwargs: + raise TypeError(f"{self.name} is missing argument {name!r}") + values.append(kwargs[name]) + device = None + packed: list[int | float] = [] + for name, kind, value in zip(names, kinds, values, strict=True): + if value is None or (kind == "p" and isinstance(value, int)): + # An argument the launch leaves out, which the kernel never reads. + packed.append(0) + continue + if kind == "p": + if not isinstance(value, torch.Tensor): + raise TypeError(f"{self.name}: {name} must be a tensor") + if device is None: + device = value.device + elif value.device != device: + raise ValueError( + f"{self.name}: {name} is on {value.device}, not {device}" + ) + packed.append(value.data_ptr()) + elif kind == "f": + packed.append(float(value)) + else: + packed.append(int(value)) + if device is None or device.type == "cpu": + _module("cpu").launch(self.name, grid, tuple(packed)) + return + index = ( + device.index if device.index is not None else torch.cuda.current_device() + ) + stream = torch.cuda.current_stream(device).cuda_stream + _module("cuda").launch(self.name, grid, tuple(packed), index, stream) diff --git a/src/blochsim/_kernels.hpp b/src/blochsim/_kernels.hpp new file mode 100644 index 00000000..50ce00f1 --- /dev/null +++ b/src/blochsim/_kernels.hpp @@ -0,0 +1,842 @@ +// What each kernel takes, in order, and which of its arguments size the +// block; and, unless ``BLOCHSIM_TABLE_ONLY``, the kernels themselves with a +// call that unpacks a launch's arguments into each. +#pragma once + +#include + +#include "_tile.hpp" +#ifndef BLOCHSIM_TABLE_ONLY +namespace epg { +#include "_epg_kernels.hpp" +} // namespace epg +namespace pools { +#include "_pools_kernels.hpp" +} // namespace pools +namespace perk { +#include "_perk_kernels.hpp" +} // namespace perk +#endif + +namespace bsk { + +// One argument as the launcher receives it from Python. +union Arg { + void* p; + std::int64_t i; + double f; +}; + +constexpr int MAX_ARGUMENTS = 128; + +struct Arguments { + Arg a[MAX_ARGUMENTS]; +}; + +struct KernelInfo { + const char* name; + // Comma-separated parameter names, in order. + const char* params; + // One letter per parameter: p a pointer, i an integer, f a float. + const char* kinds; + // The parameters whose values are the block's width along x and y, or + // -1 where the program has no such axis. + int x; + int y; + // The parameter whose value is the length of z, or -1. + int z; +}; + +#ifndef BLOCHSIM_TABLE_ONLY +BSK_HD void call_three_pool_table_jvp_kernel(const Arg* a) { + epg::_three_pool_table_jvp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + a[17].i, + a[18].i, + a[19].i); +} + +BSK_HD void call_three_pool_table_kernel(const Arg* a) { + epg::_three_pool_table_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + a[10].i, + a[11].i, + a[12].i); +} + +BSK_HD void call_epg_vjp_kernel(const Arg* a) { + epg::_epg_vjp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + a[35].i, + static_cast(a[36].p), + static_cast(a[37].p), + static_cast(a[38].p), + static_cast(a[39].p), + static_cast(a[40].p), + static_cast(a[41].p), + static_cast(a[42].p), + static_cast(a[43].p), + static_cast(a[44].p), + a[45].i, + a[46].i, + a[47].i, + a[48].i, + a[49].i, + a[50].i, + static_cast(a[51].f), + static_cast(a[52].f), + a[53].i, + static_cast(a[54].f), + static_cast(a[55].f), + a[56].i, + a[57].i, + a[58].i, + a[59].i, + a[60].i, + a[61].i, + a[62].i, + a[63].i, + a[64].i, + a[65].i, + a[66].i, + a[67].i, + a[68].i, + a[69].i, + a[70].i, + a[71].i, + a[72].i, + a[73].i, + a[74].i, + a[75].i, + a[76].i, + a[77].i); +} + +BSK_HD void call_epg_vjp_jvp_kernel(const Arg* a) { + epg::_epg_vjp_jvp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + static_cast(a[35].p), + static_cast(a[36].p), + static_cast(a[37].p), + static_cast(a[38].p), + static_cast(a[39].p), + static_cast(a[40].p), + static_cast(a[41].p), + static_cast(a[42].p), + static_cast(a[43].p), + static_cast(a[44].p), + static_cast(a[45].p), + static_cast(a[46].p), + static_cast(a[47].p), + static_cast(a[48].p), + static_cast(a[49].p), + static_cast(a[50].p), + static_cast(a[51].p), + static_cast(a[52].p), + static_cast(a[53].p), + static_cast(a[54].p), + static_cast(a[55].p), + static_cast(a[56].p), + static_cast(a[57].p), + a[58].i, + static_cast(a[59].p), + static_cast(a[60].p), + static_cast(a[61].p), + static_cast(a[62].p), + static_cast(a[63].p), + static_cast(a[64].p), + static_cast(a[65].p), + static_cast(a[66].p), + static_cast(a[67].p), + static_cast(a[68].p), + static_cast(a[69].p), + static_cast(a[70].p), + static_cast(a[71].p), + static_cast(a[72].p), + a[73].i, + a[74].i, + a[75].i, + a[76].i, + a[77].i, + a[78].i, + static_cast(a[79].f), + static_cast(a[80].f), + static_cast(a[81].f), + static_cast(a[82].f), + a[83].i, + a[84].i, + a[85].i, + a[86].i, + a[87].i, + a[88].i, + a[89].i, + a[90].i, + a[91].i, + a[92].i, + a[93].i, + a[94].i, + a[95].i, + a[96].i, + a[97].i, + a[98].i, + a[99].i, + a[100].i, + a[101].i, + a[102].i, + a[103].i, + a[104].i, + a[105].i, + a[106].i); +} + +BSK_HD void call_epg_real_vjp_jvp_kernel(const Arg* a) { + epg::_epg_real_vjp_jvp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + a[29].i, + a[30].i, + a[31].i, + a[32].i, + a[33].i, + a[34].i, + a[35].i, + a[36].i, + a[37].i, + a[38].i, + a[39].i, + a[40].i, + a[41].i, + a[42].i, + a[43].i, + a[44].i, + a[45].i); +} + +BSK_HD void call_epg_real_vjp_kernel(const Arg* a) { + epg::_epg_real_vjp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + a[17].i, + a[18].i, + a[19].i, + a[20].i, + a[21].i, + a[22].i, + a[23].i, + a[24].i, + a[25].i, + a[26].i, + a[27].i, + a[28].i, + a[29].i, + a[30].i, + a[31].i, + a[32].i, + a[33].i); +} + +BSK_HD void call_epg_real_kernel(const Arg* a) { + epg::_epg_real_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + a[14].i, + a[15].i, + a[16].i, + a[17].i, + a[18].i, + a[19].i, + a[20].i, + a[21].i, + a[22].i, + a[23].i, + a[24].i, + a[25].i, + a[26].i, + a[27].i); +} + +BSK_HD void call_epg_real_jvp_kernel(const Arg* a) { + epg::_epg_real_jvp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + a[22].i, + a[23].i, + a[24].i, + a[25].i, + a[26].i, + a[27].i, + a[28].i, + a[29].i, + a[30].i, + a[31].i, + a[32].i, + a[33].i, + a[34].i, + a[35].i); +} + +BSK_HD void call_epg_kernel(const Arg* a) { + epg::_epg_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + static_cast(a[35].p), + static_cast(a[36].p), + a[37].i, + a[38].i, + a[39].i, + a[40].i, + static_cast(a[41].f), + static_cast(a[42].f), + static_cast(a[43].f), + static_cast(a[44].f), + a[45].i, + a[46].i, + a[47].i, + a[48].i, + a[49].i, + a[50].i, + a[51].i, + a[52].i, + a[53].i, + a[54].i, + a[55].i, + a[56].i, + a[57].i, + a[58].i, + a[59].i, + a[60].i, + a[61].i, + a[62].i, + a[63].i, + a[64].i, + a[65].i, + a[66].i); +} + +BSK_HD void call_epg_jvp_kernel(const Arg* a) { + epg::_epg_jvp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + static_cast(a[35].p), + static_cast(a[36].p), + static_cast(a[37].p), + static_cast(a[38].p), + static_cast(a[39].p), + static_cast(a[40].p), + static_cast(a[41].p), + static_cast(a[42].p), + static_cast(a[43].p), + static_cast(a[44].p), + static_cast(a[45].p), + static_cast(a[46].p), + static_cast(a[47].p), + static_cast(a[48].p), + static_cast(a[49].p), + static_cast(a[50].p), + static_cast(a[51].p), + static_cast(a[52].p), + static_cast(a[53].p), + static_cast(a[54].p), + static_cast(a[55].p), + a[56].i, + a[57].i, + a[58].i, + a[59].i, + static_cast(a[60].f), + static_cast(a[61].f), + static_cast(a[62].f), + static_cast(a[63].f), + a[64].i, + a[65].i, + a[66].i, + a[67].i, + a[68].i, + a[69].i, + a[70].i, + a[71].i, + a[72].i, + a[73].i, + a[74].i, + a[75].i, + a[76].i, + a[77].i, + a[78].i, + a[79].i, + a[80].i, + a[81].i, + a[82].i, + a[83].i, + a[84].i, + a[85].i); +} + +BSK_HD void call_pooled_kernel(const Arg* a) { + pools::_pooled_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + static_cast(a[35].p), + static_cast(a[36].p), + static_cast(a[37].p), + a[38].i, + a[39].i, + a[40].i, + a[41].i, + a[42].i, + a[43].i, + static_cast(a[44].f), + static_cast(a[45].f), + static_cast(a[46].f), + static_cast(a[47].f), + a[48].i, + a[49].i, + a[50].i, + a[51].i, + a[52].i, + a[53].i, + a[54].i, + a[55].i, + a[56].i, + a[57].i, + a[58].i, + a[59].i, + a[60].i, + a[61].i, + a[62].i, + a[63].i, + a[64].i, + a[65].i, + a[66].i, + a[67].i, + a[68].i, + a[69].i, + a[70].i); +} + +BSK_HD void call_pooled_adjoint_kernel(const Arg* a) { + pools::_pooled_adjoint_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + static_cast(a[7].p), + static_cast(a[8].p), + static_cast(a[9].p), + static_cast(a[10].p), + static_cast(a[11].p), + static_cast(a[12].p), + static_cast(a[13].p), + static_cast(a[14].p), + static_cast(a[15].p), + static_cast(a[16].p), + static_cast(a[17].p), + static_cast(a[18].p), + static_cast(a[19].p), + static_cast(a[20].p), + static_cast(a[21].p), + static_cast(a[22].p), + static_cast(a[23].p), + static_cast(a[24].p), + static_cast(a[25].p), + static_cast(a[26].p), + static_cast(a[27].p), + static_cast(a[28].p), + static_cast(a[29].p), + static_cast(a[30].p), + static_cast(a[31].p), + static_cast(a[32].p), + static_cast(a[33].p), + static_cast(a[34].p), + static_cast(a[35].p), + static_cast(a[36].p), + static_cast(a[37].p), + static_cast(a[38].p), + static_cast(a[39].p), + static_cast(a[40].p), + static_cast(a[41].p), + static_cast(a[42].p), + static_cast(a[43].p), + static_cast(a[44].p), + static_cast(a[45].p), + static_cast(a[46].p), + static_cast(a[47].p), + static_cast(a[48].p), + static_cast(a[49].p), + a[50].i, + a[51].i, + a[52].i, + a[53].i, + a[54].i, + a[55].i, + static_cast(a[56].f), + static_cast(a[57].f), + static_cast(a[58].f), + static_cast(a[59].f), + a[60].i, + a[61].i, + a[62].i, + a[63].i, + a[64].i, + a[65].i, + a[66].i, + a[67].i, + a[68].i, + a[69].i, + a[70].i, + a[71].i, + a[72].i, + a[73].i, + a[74].i, + a[75].i, + a[76].i, + a[77].i, + a[78].i, + a[79].i, + a[80].i, + a[81].i, + a[82].i, + a[83].i, + a[84].i, + a[85].i, + a[86].i, + a[87].i, + a[88].i); +} + +BSK_HD void call_regress_kernel(const Arg* a) { + perk::_regress_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + static_cast(a[6].p), + a[7].i, + a[8].i, + a[9].i, + a[10].i, + static_cast(a[11].f), + a[12].i); +} + +BSK_HD void call_regress_vjp_kernel(const Arg* a) { + perk::_regress_vjp_kernel( + static_cast(a[0].p), + static_cast(a[1].p), + static_cast(a[2].p), + static_cast(a[3].p), + static_cast(a[4].p), + static_cast(a[5].p), + a[6].i, + a[7].i, + a[8].i, + a[9].i, + static_cast(a[10].f), + a[11].i); +} + +#endif + +inline constexpr KernelInfo KERNELS[] = { + {"_three_pool_table_jvp_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,d_t1,d_t1_pool_b,d_t1_bound,d_pool_b_exchange,d_bound_exchange,d_pool_b_fraction,d_bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "pppppppppppppppppiii", 18, -1, -1}, + {"_three_pool_table_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "ppppppppppiii", 11, -1, -1}, + {"_epg_vjp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,lineshape,profile,profile_index,pairs,pair_index,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_pair,grad_output_real,grad_output_imag,grad_tissue,grad_flip,grad_phase,grad_duration,trajectory_r,trajectory_i,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,shim_rows,profile_step,lineshape_step,state_count,single_train,atom_stride,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,recording,block_states,problems", "pppppppppppppppppppppppppppppppppppipppppppppiiiiiiffiffiiiiiiiiiiiiiiiiiiiiii", 76, 77, -1}, + {"_epg_vjp_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,grad_pair_value,grad_pair_tangent,dot_t1,dot_t2,dot_m0,dot_b1,dot_b1_phase,dot_b0,dot_inversion_efficiency,dot_diffusion,dot_velocity,dot_bound_fraction,dot_exchange_rate,dot_t1_bound,dot_pool_b_fraction,dot_pool_b_exchange,dot_t1_pool_b,dot_t2_pool_b,dot_pool_b_shift,dot_duration,dot_flip,dot_phase,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_output_real,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_phase_value,grad_phase_tangent,grad_duration_value,grad_duration_tangent,trajectory_vr,trajectory_vi,trajectory_tr,trajectory_ti,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,directed,off_axis,moving,diffusing,transmit,density,inverting,broadened,lineshape_bins,pools,narrow,tabulated,recording,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppppippppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiii", 105, 106, -1}, + {"_epg_real_vjp_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,dot_t1,dot_t2,dot_m0,dot_b1,dot_inversion_efficiency,dot_diffusion,dot_duration,dot_flip,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_duration_value,grad_duration_tangent,trajectory_value,trajectory_tangent,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppiiiiiiiiiiiiiiiii", 44, 45, -1}, + {"_epg_real_vjp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,grad_output_imag,grad_tissue,grad_flip,grad_duration,trajectory_value,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppiiiiiiiiiiiiiiiii", 32, 33, -1}, + {"_epg_real_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppiiiiiiiiiiiiii", 26, 27, -1}, + {"_epg_real_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_inversion_efficiency,tangent_diffusion,tangent_duration,tangent_flip,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppiiiiiiiiiiiiii", 34, 35, -1}, + {"_epg_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,bound_exchange,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,phase_cos,phase_sin,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 65, 66, -1}, + {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1}, + {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69}, + {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87}, + {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "pppppppiiiifi", 12, -1, -1}, + {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "ppppppiiiifi", 11, -1, -1}, +}; + +#define BLOCHSIM_FOR_EACH_KERNEL(X) \ + X(_three_pool_table_jvp_kernel) \ + X(_three_pool_table_kernel) \ + X(_epg_vjp_kernel) \ + X(_epg_vjp_jvp_kernel) \ + X(_epg_real_vjp_jvp_kernel) \ + X(_epg_real_vjp_kernel) \ + X(_epg_real_kernel) \ + X(_epg_real_jvp_kernel) \ + X(_epg_kernel) \ + X(_epg_jvp_kernel) \ + X(_pooled_kernel) \ + X(_pooled_adjoint_kernel) \ + X(_regress_kernel) \ + X(_regress_vjp_kernel) + +} // namespace bsk diff --git a/src/blochsim/_launch.hpp b/src/blochsim/_launch.hpp new file mode 100644 index 00000000..57035c0e --- /dev/null +++ b/src/blochsim/_launch.hpp @@ -0,0 +1,117 @@ +// The Python side of a kernel launch, shared by the CUDA build and the host +// one: the table of kernels, and an argument tuple read into the array a +// kernel takes. Host code only; the kernels themselves are in _kernels.hpp. +#pragma once + +#define PY_SSIZE_T_CLEAN +#include + +#include +#include + +#include "_kernels.hpp" + +namespace blochsim_launch { + +inline int find_kernel(const char* name) { + int index = 0; + for (const auto& info : bsk::KERNELS) { + if (std::strcmp(info.name, name) == 0) { + return index; + } + ++index; + } + return -1; +} + +inline PyObject* kernel_table(PyObject*, PyObject*) { + PyObject* table = PyDict_New(); + if (table == nullptr) { + return nullptr; + } + for (const auto& info : bsk::KERNELS) { + PyObject* entry = Py_BuildValue("(ss)", info.params, info.kinds); + if (entry == nullptr || PyDict_SetItemString(table, info.name, entry) < 0) { + Py_XDECREF(entry); + Py_DECREF(table); + return nullptr; + } + Py_DECREF(entry); + } + return table; +} + +struct Launch { + int kernel = -1; + std::int64_t grid[2] = {1, 1}; + int block[2] = {1, 1}; + // The length of z, which a program holds in each thread. + int z = 1; + bsk::Arguments arguments{}; +}; + +// ``name``, ``grid`` and ``args`` of a launch: the grid a tuple of one or two +// program counts, the arguments one Python number per parameter, pointers as +// the integer address of their first element. +inline bool read_launch(PyObject* name, PyObject* grid, PyObject* args, Launch& launch) { + const char* text = PyUnicode_AsUTF8AndSize(name, nullptr); + if (text == nullptr) { + return false; + } + launch.kernel = find_kernel(text); + if (launch.kernel < 0) { + PyErr_Format(PyExc_KeyError, "no kernel named %s", text); + return false; + } + const bsk::KernelInfo& info = bsk::KERNELS[launch.kernel]; + if (!PyTuple_Check(grid) || PyTuple_Size(grid) < 1 || PyTuple_Size(grid) > 2) { + PyErr_SetString(PyExc_ValueError, "the grid is a tuple of one or two counts"); + return false; + } + for (Py_ssize_t axis = 0; axis < PyTuple_Size(grid); ++axis) { + launch.grid[axis] = PyLong_AsLongLong(PyTuple_GetItem(grid, axis)); + if (PyErr_Occurred()) { + return false; + } + } + const Py_ssize_t count = static_cast(std::strlen(info.kinds)); + if (!PyTuple_Check(args) || PyTuple_Size(args) != count) { + PyErr_Format(PyExc_TypeError, "%s takes %zd arguments", info.name, count); + return false; + } + for (Py_ssize_t i = 0; i < count; ++i) { + PyObject* item = PyTuple_GetItem(args, i); + bsk::Arg& arg = launch.arguments.a[i]; + switch (info.kinds[i]) { + case 'p': + arg.p = PyLong_AsVoidPtr(item); + break; + case 'f': + arg.f = PyFloat_AsDouble(item); + break; + default: + arg.i = PyLong_AsLongLong(item); + break; + } + if (PyErr_Occurred()) { + PyErr_Format(PyExc_TypeError, "argument %zd of %s is not a number", i, info.name); + return false; + } + } + launch.block[0] = info.x < 0 ? 1 : static_cast(launch.arguments.a[info.x].i); + launch.block[1] = info.y < 0 ? 1 : static_cast(launch.arguments.a[info.y].i); + launch.z = info.z < 0 ? 1 : static_cast(launch.arguments.a[info.z].i); + if (launch.block[0] < 1 || launch.block[1] < 1 || launch.block[0] * launch.block[1] > 1024) { + PyErr_Format(PyExc_ValueError, "%s: a block of %d by %d threads is more than a card runs", + info.name, launch.block[0], launch.block[1]); + return false; + } + if (launch.z < 1 || launch.z > bsk::MAX_Z) { + PyErr_Format(PyExc_ValueError, "%s: %d pools is more than a program holds", info.name, + launch.z); + return false; + } + return true; +} + +} // namespace blochsim_launch diff --git a/src/blochsim/_perk_kernels.hpp b/src/blochsim/_perk_kernels.hpp new file mode 100644 index 00000000..b77ffb4b --- /dev/null +++ b/src/blochsim/_perk_kernels.hpp @@ -0,0 +1,113 @@ +// The PERK feature map and its regression, fused, one voxel per thread. +// +// y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T +// +// A block of features is formed and consumed into the output accumulator in +// registers, so the ``(voxels, features)`` matrix never exists. The adjoint +// does the same and forms the angles again rather than keeping them. Every +// array is contiguous and row-major. + +// Features formed at once, which is how often a voxel's signal is read. +constexpr int FEATURE_BLOCK = 32; +// Parameters accumulated at once by the forward pass. +constexpr int PARAMETER_BLOCK = 16; +// Contrasts accumulated at once by the adjoint. +constexpr int CONTRAST_BLOCK = 32; + +// The angles ``W @ x + b`` of features ``first`` onward, as many as there are. +template +BSK_HD void _angles(const float* signal, const float* frequency, const float* phase, + const Voxel& voxel, const Live& live, std::int64_t contrasts, + std::int64_t features, std::int64_t first, bsk::V* angle) { + for (int j = 0; j < FEATURE_BLOCK; ++j) { + angle[j] = bsk::V(first + j < features ? phase[first + j] : 0.0f); + } + for (std::int64_t contrast = 0; contrast < contrasts; ++contrast) { + const auto value = bsk::ld(signal + voxel * contrasts + contrast, live, 0.0f); + for (int j = 0; j < FEATURE_BLOCK; ++j) { + if (first + j < features) { + angle[j] = angle[j] + value * frequency[(first + j) * contrasts + contrast]; + } + } + } +} + +// One block of voxels, from signal to parameters. +BSK_HD void _regress_kernel(const float* signal, const float* frequency, const float* phase, + const float* feature_mean, const float* weight, + const float* parameter_mean, float* output, std::int64_t voxels, + std::int64_t contrasts, std::int64_t features, + std::int64_t parameters, float scale, std::int64_t BLOCK_VOXELS) { + const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); + const auto live = voxel < voxels; + bsk::V angle[FEATURE_BLOCK]; + for (std::int64_t base = 0; base < parameters; base += PARAMETER_BLOCK) { + bsk::V total[PARAMETER_BLOCK]; + for (int k = 0; k < PARAMETER_BLOCK; ++k) { + total[k] = bsk::V(0.0f); + } + for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { + _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); + for (int j = 0; j < FEATURE_BLOCK; ++j) { + if (first + j >= features) { + break; + } + const auto mapped = scale * bsk::cos(angle[j]) - feature_mean[first + j]; + for (int k = 0; k < PARAMETER_BLOCK; ++k) { + if (base + k < parameters) { + total[k] = total[k] + mapped * weight[(base + k) * features + first + j]; + } + } + } + } + for (int k = 0; k < PARAMETER_BLOCK; ++k) { + if (base + k < parameters) { + bsk::st(output + voxel * parameters + base + k, + total[k] + parameter_mean[base + k], live); + } + } + } +} + +// The derivative of one block of voxels with respect to their signals. +BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, const float* phase, + const float* weight, const float* cotangent, float* output, + std::int64_t voxels, std::int64_t contrasts, + std::int64_t features, std::int64_t parameters, float scale, + std::int64_t BLOCK_VOXELS) { + const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); + const auto live = voxel < voxels; + bsk::V angle[FEATURE_BLOCK]; + for (std::int64_t base = 0; base < contrasts; base += CONTRAST_BLOCK) { + bsk::V gradient[CONTRAST_BLOCK]; + for (int c = 0; c < CONTRAST_BLOCK; ++c) { + gradient[c] = bsk::V(0.0f); + } + for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { + _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); + for (int j = 0; j < FEATURE_BLOCK; ++j) { + if (first + j >= features) { + break; + } + bsk::V through(0.0f); + for (std::int64_t p = 0; p < parameters; ++p) { + through = through + + bsk::ld(cotangent + voxel * parameters + p, live, 0.0f) + * weight[p * features + first + j]; + } + through = through * (-scale) * bsk::sin(angle[j]); + for (int c = 0; c < CONTRAST_BLOCK; ++c) { + if (base + c < contrasts) { + gradient[c] = gradient[c] + + through * frequency[(first + j) * contrasts + base + c]; + } + } + } + } + for (int c = 0; c < CONTRAST_BLOCK; ++c) { + if (base + c < contrasts) { + bsk::st(output + voxel * contrasts + base + c, gradient[c], live); + } + } + } +} diff --git a/src/blochsim/_pools_kernels.hpp b/src/blochsim/_pools_kernels.hpp new file mode 100644 index 00000000..d84a47b7 --- /dev/null +++ b/src/blochsim/_pools_kernels.hpp @@ -0,0 +1,1851 @@ +// The kernels for two or more exchanging pools, forward and adjoint, which +// carry the pools along the tile's y and z axes. Written over the tiles of +// _tile.hpp and included by _kernels.hpp inside ``pools``. + +// The lineshape, its slope and its curvature, from the same cubic. +// +// The table covers the magnitude, so the slope changes sign with the offset +// and the curvature does not: an even function's second derivative is even. +template +BSK_HD auto _lineshape_at_curve(const T0& lineshape, const T1& offset_hz, const T2& bins, const T3& step) { + auto last = (bins - 1); + auto magnitude = bsk::truediv(bsk::abs(offset_hz), step); + auto scaled = bsk::minimum(magnitude, (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto base = (bsk::cast(lower) * 2); + auto near = bsk::ld((lineshape + base)); + auto near_slope = bsk::ld(((lineshape + base) + 1)); + auto far = bsk::ld(((lineshape + base) + 2)); + auto far_slope = bsk::ld(((lineshape + base) + 3)); + auto value = (((((((2.0f * u3) - (3.0f * u2)) + 1.0f) * near) + ((((u3 - (2.0f * u2)) + u) * step) * near_slope)) + (((-2.0f * u3) + (3.0f * u2)) * far)) + (((u3 - u2) * step) * far_slope)); + auto direction = bsk::where((offset_hz < 0.0f), -1.0f, 1.0f); + auto slope = (direction * (((bsk::truediv((((6.0f * u2) - (6.0f * u)) * near), step) + ((((3.0f * u2) - (4.0f * u)) + 1.0f) * near_slope)) + bsk::truediv((((-6.0f * u2) + (6.0f * u)) * far), step)) + (((3.0f * u2) - (2.0f * u)) * far_slope))); + auto curve = (((bsk::truediv((((12.0f * u) - 6.0f) * near), (step * step)) + bsk::truediv((((6.0f * u) - 4.0f) * near_slope), step)) + bsk::truediv((((-12.0f * u) + 6.0f) * far), (step * step))) + bsk::truediv((((6.0f * u) - 2.0f) * far_slope), step)); + auto beyond = (magnitude > last); + return bsk::make_tup(value, bsk::where(beyond, 0.0f, slope), bsk::where(beyond, 0.0f, curve)); +} + +// The lineshape and its derivative in the *signed* offset. +// +// The table covers the magnitude, so the slope changes sign with the offset; +// past the last knot the read is constant and the slope is zero. +template +BSK_HD auto _lineshape_at_slope(const T0& lineshape, const T1& offset_hz, const T2& bins, const T3& step) { + auto last = (bins - 1); + auto magnitude = bsk::truediv(bsk::abs(offset_hz), step); + auto scaled = bsk::minimum(magnitude, (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto base = (bsk::cast(lower) * 2); + auto near = bsk::ld((lineshape + base)); + auto near_slope = bsk::ld(((lineshape + base) + 1)); + auto far = bsk::ld(((lineshape + base) + 2)); + auto far_slope = bsk::ld(((lineshape + base) + 3)); + auto value = (((((((2.0f * u3) - (3.0f * u2)) + 1.0f) * near) + ((((u3 - (2.0f * u2)) + u) * step) * near_slope)) + (((-2.0f * u3) + (3.0f * u2)) * far)) + (((u3 - u2) * step) * far_slope)); + auto direction = bsk::where((offset_hz < 0.0f), -1.0f, 1.0f); + auto slope = (direction * (((bsk::truediv((((6.0f * u2) - (6.0f * u)) * near), step) + ((((3.0f * u2) - (4.0f * u)) + 1.0f) * near_slope)) + bsk::truediv((((-6.0f * u2) + (6.0f * u)) * far), step)) + (((3.0f * u2) - (2.0f * u)) * far_slope))); + return bsk::make_tup(value, bsk::where((magnitude > last), 0.0f, slope)); +} + +// Two real duals multiplied. +template +BSK_HD auto _rmul(const T0& x, const T1& y, const T2& following) { + using Ret = bsk::tup; + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup((bsk::get<0>(x) * bsk::get<0>(y)), ((bsk::get<1>(x) * bsk::get<0>(y)) + (bsk::get<0>(x) * bsk::get<1>(y))))); + } else { + return bsk::convert(bsk::make_tup((bsk::get<0>(x) * bsk::get<0>(y)), 0.0f)); + } +} + +// The semisolid pool's saturation by a pulse, and the lineshape it read. +// +// Returns ``exp(saturation * alpha^2 * G(offset))``, ``G`` and its slope in +// the offset, each a real dual. +template +BSK_HD auto _absorption(const T0& lineshape, const T1& rf_frequency, const T2& saturation, const T3& event, const T4& alpha, const T5& b0, const T6& lineshape_bins, const T7& lineshape_step, const T8& following) { + float shape{}; + bsk::tup shape_dual{}; + float slope{}; + bsk::tup slope_dual{}; + auto offset = (bsk::ld((rf_frequency + event)) - bsk::get<0>(b0)); + auto deposited = bsk::ld((saturation + event)); + if (bsk::truth(following)) { + auto t0_ = _lineshape_at_curve(lineshape, offset, lineshape_bins, lineshape_step); + shape = bsk::get<0>(t0_); + slope = bsk::get<1>(t0_); + auto curve = bsk::get<2>(t0_); + // The lineshape is read at the pulse's offset from the voxel, so a + // step in the voxel's own off-resonance moves the read the other way. + shape_dual = bsk::make_tup(shape, (slope * (-bsk::get<1>(b0)))); + slope_dual = bsk::make_tup(slope, (curve * (-bsk::get<1>(b0)))); + } else { + auto t1_ = _lineshape_at_slope(lineshape, offset, lineshape_bins, lineshape_step); + shape = bsk::get<0>(t1_); + slope = bsk::get<1>(t1_); + shape_dual = bsk::make_tup(shape, 0.0f); + slope_dual = bsk::make_tup(slope, 0.0f); + } + auto exponent = _rmul(bsk::make_tup(deposited, 0.0f), _rmul(_rmul(alpha, alpha, following), shape_dual, following), following); + auto absorbed = bsk::exp(bsk::get<0>(exponent)); + return bsk::make_tup(bsk::make_tup(absorbed, (absorbed * bsk::get<1>(exponent))), shape_dual, slope_dual, deposited); +} + +// Two complex duals multiplied. +template +BSK_HD auto _cmul(const T0& x, const T1& y, const T2& following) { + using Ret = bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>>; + auto real = ((bsk::get<0>(x) * bsk::get<0>(y)) - (bsk::get<1>(x) * bsk::get<1>(y))); + auto imag = ((bsk::get<0>(x) * bsk::get<1>(y)) + (bsk::get<1>(x) * bsk::get<0>(y))); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(real, imag, ((((bsk::get<2>(x) * bsk::get<0>(y)) - (bsk::get<3>(x) * bsk::get<1>(y))) + (bsk::get<0>(x) * bsk::get<2>(y))) - (bsk::get<1>(x) * bsk::get<3>(y))), ((((bsk::get<2>(x) * bsk::get<1>(y)) + (bsk::get<3>(x) * bsk::get<0>(y))) + (bsk::get<0>(x) * bsk::get<3>(y))) + (bsk::get<1>(x) * bsk::get<2>(y))))); + } else { + return bsk::convert(bsk::make_tup(real, imag, 0.0f, 0.0f)); + } +} + +// A real dual times a complex one. +template +BSK_HD auto _cscale(const T0& r, const T1& x, const T2& following) { + using Ret = bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>>; + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup((bsk::get<0>(r) * bsk::get<0>(x)), (bsk::get<0>(r) * bsk::get<1>(x)), ((bsk::get<1>(r) * bsk::get<0>(x)) + (bsk::get<0>(r) * bsk::get<2>(x))), ((bsk::get<1>(r) * bsk::get<1>(x)) + (bsk::get<0>(r) * bsk::get<3>(x))))); + } else { + return bsk::convert(bsk::make_tup((bsk::get<0>(r) * bsk::get<0>(x)), (bsk::get<0>(r) * bsk::get<1>(x)), 0.0f, 0.0f)); + } +} + +// The rotation a pulse performs at this voxel, read rather than read off. +// +// A tabulated pair covers a shape's every pulse because a static array +// reaches the rotation through one complex scalar; this one is integrated per +// pulse per voxel, so there is nothing to interpolate and the read is four +// floats. The row runs per train and per event, as the flip does. +template +BSK_HD auto _dynamic_pair_at(const T0& pairs, const T1& pair_index, const T2& event_base, const T3& event, const T4& atom, const T5& atom_count, const T6& mask) { + auto row = bsk::cast(bsk::ld(((pair_index + event_base) + event))); + auto entry = (pairs + (((row * atom_count) + atom) * 4)); + return bsk::make_tup(bsk::ld((entry + 0), mask, 1.0f), bsk::ld((entry + 1), mask, 0.0f), bsk::ld((entry + 2), mask, 0.0f), bsk::ld((entry + 3), mask, 0.0f)); +} + +// ``exp(i * angle)`` for a real dual angle. +template +BSK_HD auto _dual_polar(const T0& angle_value, const T1& angle_tangent) { + auto cosine = bsk::cos(angle_value); + auto sine = bsk::sin(angle_value); + return bsk::make_tup(cosine, sine, ((-sine) * angle_tangent), (cosine * angle_tangent)); +} + +template +BSK_HD auto _complex_mul(const T0& a_real, const T1& a_imag, const T2& b_real, const T3& b_imag) { + return bsk::make_tup(((a_real * b_real) - (a_imag * b_imag)), ((a_real * b_imag) + (a_imag * b_real))); +} + +// Product of two dual complex numbers. +template +BSK_HD auto _dual_mul(const T0& a_vr, const T1& a_vi, const T2& a_tr, const T3& a_ti, const T4& b_vr, const T5& b_vi, const T6& b_tr, const T7& b_ti) { + auto t0_ = _complex_mul(a_vr, a_vi, b_vr, b_vi); + auto value_real = bsk::get<0>(t0_); + auto value_imag = bsk::get<1>(t0_); + auto t1_ = _complex_mul(a_tr, a_ti, b_vr, b_vi); + auto left_real = bsk::get<0>(t1_); + auto left_imag = bsk::get<1>(t1_); + auto t2_ = _complex_mul(a_vr, a_vi, b_tr, b_ti); + auto right_real = bsk::get<0>(t2_); + auto right_imag = bsk::get<1>(t2_); + return bsk::make_tup(value_real, value_imag, (left_real + right_real), (left_imag + right_imag)); +} + +// Two dual complex numbers multiplied. +template +BSK_HD auto _dual_product(const T0& x, const T1& y) { + return _dual_mul(bsk::get<0>(x), bsk::get<1>(x), bsk::get<2>(x), bsk::get<3>(x), bsk::get<0>(y), bsk::get<1>(y), bsk::get<2>(y), bsk::get<3>(y)); +} + +// The rotation and the direction along it, with the phase applied. +// +// Shaped exactly as :func:`_profiled_pair_dual` returns, so the spinor +// operator and its adjoint read one from the other without knowing which +// they were handed. A pass that follows no direction holds the rotation +// still, and ``directed`` keeps the read for one out of the kernel. +template +BSK_HD auto _dynamic_pair_dual_at(const T0& pairs, const T1& pair_direction, const T2& pair_index, const T3& event_base, const T4& event, const T5& atom, const T6& atom_count, const T7& mask, const T8& phi_value, const T9& phi_tangent, const T10& directed) { + bsk::tup moved{}; + auto held = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, mask); + auto still = (bsk::get<0>(held) * 0.0f); + moved = bsk::make_tup(still, still, still, still); + if (bsk::truth(directed)) { + moved = _dynamic_pair_at(pair_direction, pair_index, event_base, event, atom, atom_count, mask); + } + auto a = bsk::make_tup(bsk::get<0>(held), bsk::get<1>(held), bsk::get<0>(moved), bsk::get<1>(moved)); + auto b = bsk::make_tup(bsk::get<2>(held), bsk::get<3>(held), bsk::get<2>(moved), bsk::get<3>(moved)); + auto turn = _dual_polar((-phi_value), (-phi_tangent)); + return bsk::make_tup(a, _dual_product(b, turn)); +} + +// Table entries at ``offset``, moving with the table's direction and slope. +// +// An event reads the row of its own interval length; a pass following a +// direction in that length moves every entry along the row's slope, which +// sits ``sloped_at`` further into the table. +template +BSK_HD auto _entries(const T0& slot, const T1& directions, const T2& offset, const T3& mask, const T4& along, const T5& sloped_at, const T6& directed, const T7& sloped, const T8& following) { + using Ret = bsk::tup | 0, 6)>, bsk::tile_t | 0, 6)>>; + bsk::tile_t | 0, 6)> tangent{}; + auto value = bsk::ld((slot + offset), mask, 0.0f); + if (bsk::truth(following)) { + tangent = (value * 0.0f); + if (bsk::truth(directed)) { + tangent = (tangent + bsk::ld((directions + offset), mask, 0.0f)); + } + if (bsk::truth(sloped)) { + tangent = (tangent + (bsk::ld(((slot + sloped_at) + offset), mask, 0.0f) * along)); + } + return bsk::convert(bsk::make_tup(value, tangent)); + } else { + return bsk::convert(bsk::make_tup(value, 0.0f)); + } +} + +// ``exp(i angle)`` for a real dual angle. +template +BSK_HD auto _polar(const T0& angle, const T1& following) { + using Ret = bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>>; + auto cosine = bsk::cos(bsk::get<0>(angle)); + auto sine = bsk::sin(bsk::get<0>(angle)); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(cosine, sine, ((-sine) * bsk::get<1>(angle)), (cosine * bsk::get<1>(angle)))); + } else { + return bsk::convert(bsk::make_tup(cosine, sine, 0.0f, 0.0f)); + } +} + +// What an interval does to every pool alike, per dephasing order. +// +// Returns the fraction washout leaves, the transverse and longitudinal +// factors before washout, and the per-order damping weights. +template +BSK_HD auto _factors(const T0& dt, const T1& damping_rate, const T2& b0, const T3& flow_rate, const T4& washout_rate, const T5& order, const T6& off_axis, const T7& moving, const T8& diffusing, const T9& following) { + bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>> damp_t{}; + bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>> damp_z{}; + bsk::tup turn{}; + bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>> unit_t{}; + bsk::tup | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>, bsk::tile_t | 0, 1)>> unit_z{}; + bsk::tup wout{}; + auto squared = (order * order); + auto transverse_weight = ((squared + order) + 0.3333333333333333f); + damp_z = bsk::make_tup(((order * 0.0f) + 1.0f), 0.0f); + damp_t = bsk::make_tup(((order * 0.0f) + 1.0f), 0.0f); + if (bsk::truth(diffusing)) { + auto b_factor = _rmul(damping_rate, dt, following); + auto z = bsk::exp(((-squared) * bsk::get<0>(b_factor))); + auto t = bsk::exp(((-transverse_weight) * bsk::get<0>(b_factor))); + if (bsk::truth(following)) { + damp_z = bsk::make_tup(z, (z * ((-squared) * bsk::get<1>(b_factor)))); + damp_t = bsk::make_tup(t, (t * ((-transverse_weight) * bsk::get<1>(b_factor)))); + } else { + damp_z = bsk::make_tup(z, 0.0f); + damp_t = bsk::make_tup(t, 0.0f); + } + } + wout = bsk::make_tup(1.0f, 0.0f); + if (bsk::truth(moving)) { + auto fraction = (bsk::get<0>(washout_rate) * bsk::get<0>(dt)); + auto left = (1.0f - bsk::minimum(fraction, 1.0f)); + if (bsk::truth(following)) { + wout = bsk::make_tup(left, bsk::where((fraction < 1.0f), (-((bsk::get<1>(washout_rate) * bsk::get<0>(dt)) + (bsk::get<0>(washout_rate) * bsk::get<1>(dt)))), 0.0f)); + } else { + wout = bsk::make_tup(left, 0.0f); + } + } + unit_t = bsk::make_tup(bsk::get<0>(damp_t), (bsk::get<0>(damp_t) * 0.0f), bsk::get<1>(damp_t), 0.0f); + unit_z = bsk::make_tup(bsk::get<0>(damp_z), (bsk::get<0>(damp_z) * 0.0f), bsk::get<1>(damp_z), 0.0f); + if (bsk::truth((bsk::truth(off_axis) || bsk::truth(moving)))) { + auto angle = _rmul(bsk::make_tup(-6.283185307179586f, 0.0f), _rmul(b0, dt, following), following); + turn = bsk::make_tup(0.0f, 0.0f); + if (bsk::truth(moving)) { + turn = _rmul(flow_rate, dt, following); + } + auto half = (-(order + 0.5f)); + auto theta = bsk::make_tup((bsk::get<0>(angle) + (half * bsk::get<0>(turn))), (bsk::get<1>(angle) + (half * bsk::get<1>(turn)))); + unit_t = _cscale(damp_t, _polar(theta, following), following); + if (bsk::truth(moving)) { + unit_z = _cscale(damp_z, _polar(bsk::make_tup(((-order) * bsk::get<0>(turn)), ((-order) * bsk::get<1>(turn))), following), following); + } + } + if (bsk::truth(following)) { + // Every tangent a tile, so the operator products can take them. + unit_t = bsk::make_tup(bsk::get<0>(unit_t), bsk::get<1>(unit_t), (bsk::get<2>(unit_t) + (order * 0.0f)), (bsk::get<3>(unit_t) + (order * 0.0f))); + unit_z = bsk::make_tup(bsk::get<0>(unit_z), bsk::get<1>(unit_z), (bsk::get<2>(unit_z) + (order * 0.0f)), (bsk::get<3>(unit_z) + (order * 0.0f))); + } + return bsk::make_tup(wout, unit_t, unit_z, squared, transverse_weight); +} + +// A hard pulse as its Cayley-Klein pair, with the pair's slope in the flip. +// +// ``a = cos(alpha / 2)`` and ``b = -i sin(alpha / 2) exp(-i phi)``, the +// rotation ``_rotate_flip_phase`` performs. +template +BSK_HD auto _hard_pair(const T0& alpha, const T1& phi, const T2& following) { + bsk::tup a{}; + bsk::tup b{}; + bsk::tup slope_a{}; + bsk::tup slope_b{}; + auto half = (0.5f * bsk::get<0>(alpha)); + auto cosine = bsk::cos(half); + auto sine = bsk::sin(half); + auto nothing = (cosine * 0.0f); + auto turn = _polar(bsk::make_tup((-bsk::get<0>(phi)), (-bsk::get<1>(phi))), following); + if (bsk::truth(following)) { + a = bsk::make_tup(cosine, nothing, ((-0.5f * sine) * bsk::get<1>(alpha)), nothing); + b = bsk::make_tup(nothing, (-sine), nothing, ((-0.5f * cosine) * bsk::get<1>(alpha))); + slope_a = bsk::make_tup((-0.5f * sine), nothing, ((-0.25f * cosine) * bsk::get<1>(alpha)), nothing); + slope_b = bsk::make_tup(nothing, (-0.5f * cosine), nothing, ((0.25f * sine) * bsk::get<1>(alpha))); + } else { + a = bsk::make_tup(cosine, nothing, 0.0f, 0.0f); + b = bsk::make_tup(nothing, (-sine), 0.0f, 0.0f); + slope_a = bsk::make_tup((-0.5f * sine), nothing, 0.0f, 0.0f); + slope_b = bsk::make_tup(nothing, (-0.5f * cosine), 0.0f, 0.0f); + } + return bsk::make_tup(a, _cmul(b, turn, following), slope_a, _cmul(slope_b, turn, following), turn); +} + +// One interval's longitudinal operator, its recovery and transverse operator. +template +BSK_HD auto _operators(const T0& slot, const T1& directions, const T2& row_offset, const T3& along, const T4& slope_offset, const T5& pool, const T6& column, const T7& n, const T8& m, const T9& directed, const T10& sloped, const T11& following) { + auto longitudinal = _entries(slot, directions, ((row_offset + (pool * n)) + column), bsk::band((pool < n), (column < n)), along, slope_offset, directed, sloped, following); + auto restored = _entries(slot, directions, ((row_offset + (n * n)) + pool), (pool < n), along, slope_offset, directed, sloped, following); + auto across = (((row_offset + (n * n)) + n) + (2 * ((pool * m) + column))); + auto carried = bsk::band((pool < m), (column < m)); + auto real = _entries(slot, directions, across, carried, along, slope_offset, directed, sloped, following); + auto imag = _entries(slot, directions, (across + 1), carried, along, slope_offset, directed, sloped, following); + return bsk::make_tup(longitudinal, restored, bsk::make_tup(bsk::get<0>(real), bsk::get<0>(imag), bsk::get<1>(real), bsk::get<1>(imag))); +} + +// The Cayley-Klein pair the transition table holds at this flip angle. +// +// Cubic Hermite between the two knots bracketing ``theta``, clamped at both +// ends: a cubic run off its grid leaves the unit circle. Each knot is eight +// floats -- the pair then its slope, real before imaginary -- so the two a +// read needs are sixteen contiguous ones. +template +BSK_HD auto _profile_pair(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto component = [&](int c) { + auto near = bsk::ld(((profile + base) + c)); + auto near_slope = bsk::ld((((profile + base) + 4) + c)); + auto far = bsk::ld((((profile + base) + 8) + c)); + auto far_slope = bsk::ld((((profile + base) + 12) + c)); + return ((((h00 * near) + (h10 * near_slope)) + (h01 * far)) + (h11 * far_slope)); + }; + return bsk::make_tup(component(0), component(1), component(2), component(3)); +} + +// The pair and its derivative in the flip angle, from the same cubic. +// +// The derivative of a Hermite segment is another polynomial in the same four +// knot values, so reading both costs one extra combination rather than a +// second table. Returned interleaved: each component's value then its slope, +// in the order ``a`` real, ``a`` imaginary, ``b`` real, ``b`` imaginary. +template +BSK_HD auto _profile_pair_slope(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + // d/dtheta is d/du over the knot spacing. + auto g00 = bsk::truediv(((6.0f * u2) - (6.0f * u)), step); + auto g10 = (((3.0f * u2) - (4.0f * u)) + 1.0f); + auto g01 = bsk::truediv(((6.0f * u) - (6.0f * u2)), step); + auto g11 = ((3.0f * u2) - (2.0f * u)); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto near = [&](int c) { return bsk::ld(((profile + base) + c)); }; + auto near_slope = [&](int c) { return bsk::ld((((profile + base) + 4) + c)); }; + auto far = [&](int c) { return bsk::ld((((profile + base) + 8) + c)); }; + auto far_slope = [&](int c) { return bsk::ld((((profile + base) + 12) + c)); }; + auto value = [&](int c) { + return ((((h00 * near(c)) + (h10 * near_slope(c))) + (h01 * far(c))) + (h11 * far_slope(c))); + }; + auto slope = [&](int c) { + return ((((g00 * near(c)) + (g10 * near_slope(c))) + (g01 * far(c))) + (g11 * far_slope(c))); + }; + return bsk::make_tup(value(0), slope(0), value(1), slope(1), value(2), slope(2), value(3), slope(3)); +} + +// A buffer entry and the direction along it, or its identity when absent. +template +BSK_HD auto _read(const T0& values, const T1& directions, const T2& at, const T3& live, const T4& identity, const T5& following) { + using Ret = bsk::tup; + if (bsk::truth(live)) { + auto value = bsk::ld((values + at)); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(value, bsk::ld((directions + at)))); + } else { + return bsk::convert(bsk::make_tup(value, 0.0f)); + } + } else { + return bsk::convert(bsk::make_tup(identity, 0.0f)); + } +} + +// ``operator @ planes`` over the pools, or its transpose's. +template +BSK_HD auto _times(const T0& operator_, const T1& planes, const T2& transposed) { + return bsk::times(operator_, planes, bsk::truth(transposed)); +} + +// A complex dual operator applied to complex dual pool tiles. +template +BSK_HD auto _apply(const T0& operator_, const T1& planes, const T2& conjugate, const T3& transposed, const T4& following) { + using Ret = bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>>; + bsk::tile_t | 0, 6)> ei{}; + bsk::tile_t | 0, 6)> eti{}; + auto er = bsk::get<0>(operator_); + ei = bsk::get<1>(operator_); + if (bsk::truth(conjugate)) { + ei = (-ei); + } + auto real = (_times(er, bsk::get<0>(planes), transposed) - _times(ei, bsk::get<1>(planes), transposed)); + auto imag = (_times(er, bsk::get<1>(planes), transposed) + _times(ei, bsk::get<0>(planes), transposed)); + if (bsk::truth(following)) { + auto etr = bsk::get<2>(operator_); + eti = bsk::get<3>(operator_); + if (bsk::truth(conjugate)) { + eti = (-eti); + } + auto tangent_real = (((_times(etr, bsk::get<0>(planes), transposed) - _times(eti, bsk::get<1>(planes), transposed)) + _times(er, bsk::get<2>(planes), transposed)) - _times(ei, bsk::get<3>(planes), transposed)); + auto tangent_imag = (((_times(etr, bsk::get<1>(planes), transposed) + _times(eti, bsk::get<0>(planes), transposed)) + _times(er, bsk::get<3>(planes), transposed)) + _times(ei, bsk::get<2>(planes), transposed)); + return bsk::convert(bsk::make_tup(real, imag, tangent_real, tangent_imag)); + } else { + return bsk::convert(bsk::make_tup(real, imag, 0.0f, 0.0f)); + } +} + +// A real dual operator applied to complex dual pool tiles. +template +BSK_HD auto _apply_real(const T0& operator_, const T1& planes, const T2& transposed, const T3& following) { + using Ret = bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>>; + auto real = _times(bsk::get<0>(operator_), bsk::get<0>(planes), transposed); + auto imag = _times(bsk::get<0>(operator_), bsk::get<1>(planes), transposed); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(real, imag, (_times(bsk::get<1>(operator_), bsk::get<0>(planes), transposed) + _times(bsk::get<0>(operator_), bsk::get<2>(planes), transposed)), (_times(bsk::get<1>(operator_), bsk::get<1>(planes), transposed) + _times(bsk::get<0>(operator_), bsk::get<3>(planes), transposed)))); + } else { + return bsk::convert(bsk::make_tup(real, imag, 0.0f, 0.0f)); + } +} + +template +BSK_HD auto _conj(const T0& x) { + return bsk::make_tup(bsk::get<0>(x), (-bsk::get<1>(x)), bsk::get<2>(x), (-bsk::get<3>(x))); +} + +// One interval's relaxation and exchange over every order. +// +// Returns the three states it leaves and the operator products before the +// per-order factors, which the adjoint reuses. +template +BSK_HD auto _relax(const T0& plus, const T1& minus, const T2& longitudinal, const T3& transverse_op, const T4& longitudinal_op, const T5& restored, const T6& equilibrium, const T7& wout, const T8& carried, const T9& spin, const T10& state, const T11& following) { + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>, bsk::tile_t | 0, 3)>> out_z{}; + auto mixed_plus = _apply(transverse_op, plus, false, false, following); + auto mixed_minus = _apply(transverse_op, minus, true, false, following); + auto mixed_z = _apply_real(longitudinal_op, longitudinal, false, following); + auto out_plus = _cmul(carried, mixed_plus, following); + auto out_minus = _cmul(_conj(carried), mixed_minus, following); + out_z = _cmul(spin, mixed_z, following); + // Inflowing spins arrive at equilibrium, so washout scales what the pools + // held and not what they recover towards. + auto origin = (state == 0); + auto grown = (bsk::get<0>(equilibrium) - (bsk::get<0>(wout) * bsk::get<0>(restored))); + if (bsk::truth(following)) { + auto grown_tangent = (bsk::get<1>(equilibrium) - ((bsk::get<1>(wout) * bsk::get<0>(restored)) + (bsk::get<0>(wout) * bsk::get<1>(restored)))); + out_z = bsk::make_tup((bsk::get<0>(out_z) + bsk::where(origin, grown, 0.0f)), bsk::get<1>(out_z), (bsk::get<2>(out_z) + bsk::where(origin, grown_tangent, 0.0f)), bsk::get<3>(out_z)); + } else { + out_z = bsk::make_tup((bsk::get<0>(out_z) + bsk::where(origin, grown, 0.0f)), bsk::get<1>(out_z), 0.0f, 0.0f); + } + return bsk::make_tup(out_plus, out_minus, out_z, mixed_plus, mixed_minus, mixed_z); +} + +// The rotation named by its Cayley-Klein pair, applied to the states. +// +// T = [ conj(a)^2 -conj(b)^2 -2 conj(a b) ] +// [ -b^2 a^2 -2 a b ] +// [ conj(a) b a conj(b) |a|^2-|b|^2 ] +template +BSK_HD auto _rotate_spinor(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& fp_r, const T5& fp_i, const T6& fm_r, const T7& fm_i, const T8& z_r, const T9& z_i) { + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto t0_ = bsk::make_tup(aa_r, (-aa_i)); + auto t00_r = bsk::get<0>(t0_); + auto t00_i = bsk::get<1>(t0_); + auto t1_ = bsk::make_tup((-bb_r), bb_i); + auto t01_r = bsk::get<0>(t1_); + auto t01_i = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i)); + auto t02_r = bsk::get<0>(t2_); + auto t02_i = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup((-bb_r), (-bb_i)); + auto t10_r = bsk::get<0>(t3_); + auto t10_i = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(aa_r, aa_i); + auto t11_r = bsk::get<0>(t4_); + auto t11_i = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i)); + auto t12_r = bsk::get<0>(t5_); + auto t12_i = bsk::get<1>(t5_); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto t6_ = bsk::make_tup(cross_r, cross_i); + auto t20_r = bsk::get<0>(t6_); + auto t20_i = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(cross_r, (-cross_i)); + auto t21_r = bsk::get<0>(t7_); + auto t21_i = bsk::get<1>(t7_); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + auto out_pr = ((((((t00_r * fp_r) - (t00_i * fp_i)) + (t01_r * fm_r)) - (t01_i * fm_i)) + (t02_r * z_r)) - (t02_i * z_i)); + auto out_pi = ((((((t00_r * fp_i) + (t00_i * fp_r)) + (t01_r * fm_i)) + (t01_i * fm_r)) + (t02_r * z_i)) + (t02_i * z_r)); + auto out_mr = ((((((t10_r * fp_r) - (t10_i * fp_i)) + (t11_r * fm_r)) - (t11_i * fm_i)) + (t12_r * z_r)) - (t12_i * z_i)); + auto out_mi = ((((((t10_r * fp_i) + (t10_i * fp_r)) + (t11_r * fm_i)) + (t11_i * fm_r)) + (t12_r * z_i)) + (t12_i * z_r)); + auto out_zr = (((((t20_r * fp_r) - (t20_i * fp_i)) + (t21_r * fm_r)) - (t21_i * fm_i)) + (t22 * z_r)); + auto out_zi = (((((t20_r * fp_i) + (t20_i * fp_r)) + (t21_r * fm_i)) + (t21_i * fm_r)) + (t22 * z_i)); + return bsk::make_tup(out_pr, out_pi, out_mr, out_mi, out_zr, out_zi); +} + +// One row of the rotation applied to the states, values only. +template +BSK_HD auto _dual_row(const T0& first, const T1& second, const T2& third, const T3& fp_r, const T4& fp_i, const T5& fm_r, const T6& fm_i, const T7& z_r, const T8& z_i) { + auto real = ((((((bsk::get<0>(first) * fp_r) - (bsk::get<1>(first) * fp_i)) + (bsk::get<0>(second) * fm_r)) - (bsk::get<1>(second) * fm_i)) + (bsk::get<0>(third) * z_r)) - (bsk::get<1>(third) * z_i)); + auto imag = ((((((bsk::get<0>(first) * fp_i) + (bsk::get<1>(first) * fp_r)) + (bsk::get<0>(second) * fm_i)) + (bsk::get<1>(second) * fm_r)) + (bsk::get<0>(third) * z_i)) + (bsk::get<1>(third) * z_r)); + return bsk::make_tup(real, imag); +} + +// The rotation's nine coefficients and their tangents. +// +// Every entry is a product of two factors drawn from the pair and its +// conjugate, so five products carry all nine: ``a^2``, ``b^2``, ``a b``, +// ``a conj(b)`` and the norm difference. +template +BSK_HD auto _spinor_coefficients(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& dar, const T5& dai, const T6& dbr, const T7& dbi) { + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto daa_r = (2.0f * ((ar * dar) - (ai * dai))); + auto daa_i = (2.0f * ((dar * ai) + (ar * dai))); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto dbb_r = (2.0f * ((br * dbr) - (bi * dbi))); + auto dbb_i = (2.0f * ((dbr * bi) + (br * dbi))); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto dab_r = ((((dar * br) + (ar * dbr)) - (dai * bi)) - (ai * dbi)); + auto dab_i = ((((dar * bi) + (ar * dbi)) + (dai * br)) + (ai * dbr)); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto dcross_r = ((((dar * br) + (ar * dbr)) + (dai * bi)) + (ai * dbi)); + auto dcross_i = ((((dar * bi) + (ar * dbi)) - (dai * br)) - (ai * dbr)); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + auto dt22 = (2.0f * ((((ar * dar) + (ai * dai)) - (br * dbr)) - (bi * dbi))); + return bsk::make_tup(bsk::make_tup(aa_r, (-aa_i), daa_r, (-daa_i)), bsk::make_tup((-bb_r), bb_i, (-dbb_r), dbb_i), bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i), (-2.0f * dab_r), (2.0f * dab_i)), bsk::make_tup((-bb_r), (-bb_i), (-dbb_r), (-dbb_i)), bsk::make_tup(aa_r, aa_i, daa_r, daa_i), bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i), (-2.0f * dab_r), (-2.0f * dab_i)), bsk::make_tup(cross_r, cross_i, dcross_r, dcross_i), bsk::make_tup(cross_r, (-cross_i), dcross_r, (-dcross_i)), bsk::make_tup(t22, (0.0f * t22), dt22, (0.0f * dt22))); +} + +// The same row built from the coefficients' tangents instead. +template +BSK_HD auto _tangent_row(const T0& first, const T1& second, const T2& third, const T3& fp_r, const T4& fp_i, const T5& fm_r, const T6& fm_i, const T7& z_r, const T8& z_i) { + auto real = ((((((bsk::get<2>(first) * fp_r) - (bsk::get<3>(first) * fp_i)) + (bsk::get<2>(second) * fm_r)) - (bsk::get<3>(second) * fm_i)) + (bsk::get<2>(third) * z_r)) - (bsk::get<3>(third) * z_i)); + auto imag = ((((((bsk::get<2>(first) * fp_i) + (bsk::get<3>(first) * fp_r)) + (bsk::get<2>(second) * fm_i)) + (bsk::get<3>(second) * fm_r)) + (bsk::get<2>(third) * z_i)) + (bsk::get<3>(third) * z_r)); + return bsk::make_tup(real, imag); +} + +// The spinor rotation carrying a forward-mode tangent. +// +// Both the states and the pair naming the rotation move, so the tangent is +// ``T dx + dT x``. +template +BSK_HD auto _rotate_spinor_dual(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& dar, const T5& dai, const T6& dbr, const T7& dbi, const T8& fp_r, const T9& fp_i, const T10& fm_r, const T11& fm_i, const T12& z_r, const T13& z_i, const T14& dfp_r, const T15& dfp_i, const T16& dfm_r, const T17& dfm_i, const T18& dz_r, const T19& dz_i) { + auto t0_ = _spinor_coefficients(ar, ai, br, bi, dar, dai, dbr, dbi); + auto t00 = bsk::get<0>(t0_); + auto t01 = bsk::get<1>(t0_); + auto t02 = bsk::get<2>(t0_); + auto t10 = bsk::get<3>(t0_); + auto t11 = bsk::get<4>(t0_); + auto t12 = bsk::get<5>(t0_); + auto t20 = bsk::get<6>(t0_); + auto t21 = bsk::get<7>(t0_); + auto t22 = bsk::get<8>(t0_); + auto t1_ = _dual_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_pr = bsk::get<0>(t1_); + auto out_pi = bsk::get<1>(t1_); + auto t2_ = _dual_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_mr = bsk::get<0>(t2_); + auto out_mi = bsk::get<1>(t2_); + auto t3_ = _dual_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto out_zr = bsk::get<0>(t3_); + auto out_zi = bsk::get<1>(t3_); + auto t4_ = _dual_row(t00, t01, t02, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dpr = bsk::get<0>(t4_); + auto dpi = bsk::get<1>(t4_); + auto t5_ = _dual_row(t10, t11, t12, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dmr = bsk::get<0>(t5_); + auto dmi = bsk::get<1>(t5_); + auto t6_ = _dual_row(t20, t21, t22, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i); + auto dzr = bsk::get<0>(t6_); + auto dzi = bsk::get<1>(t6_); + auto t7_ = _tangent_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tpr = bsk::get<0>(t7_); + auto tpi = bsk::get<1>(t7_); + auto t8_ = _tangent_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tmr = bsk::get<0>(t8_); + auto tmi = bsk::get<1>(t8_); + auto t9_ = _tangent_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i); + auto tzr = bsk::get<0>(t9_); + auto tzi = bsk::get<1>(t9_); + return bsk::make_tup(out_pr, out_pi, out_mr, out_mi, out_zr, out_zi, (dpr + tpr), (dpi + tpi), (dmr + tmr), (dmi + tmi), (dzr + tzr), (dzi + tzi)); +} + +// ``values`` moved one order down: ``result[k] = values[k + 1]``. +// +// The top order has no neighbour to read, so it reads itself and the caller +// masks it away. +template +BSK_HD auto _down(const T0& values, const T1& state) { + return bsk::gather_x(values, bsk::minimum(state + 1, bsk::width_x() - 1)); +} + +// ``values`` moved one configuration order up: ``result[k] = values[k - 1]``. +// +// Order zero is left to the caller, which fills it from the sequence's own +// boundary condition rather than from a neighbour. +template +BSK_HD auto _up(const T0& values, const T1& state) { + return bsk::gather_x(values, bsk::maximum(state - 1, 0)); +} + +template +BSK_HD auto _shift(const T0& fplus_real, const T1& fplus_imag, const T2& fminus_real, const T3& fminus_imag, const T4& state, const T5& state_mask, const T6& state_count) { + bsk::tile_t | 0, 3)> plus_imag{}; + bsk::tile_t | 0, 3)> plus_real{}; + auto keep_up = bsk::band((state > 0), state_mask); + auto keep_down = bsk::band(((state + 1) < state_count), state_mask); + plus_real = bsk::where(keep_up, _up(fplus_real, state), 0.0f); + plus_imag = bsk::where(keep_up, _up(fplus_imag, state), 0.0f); + auto minus_real = bsk::where(keep_down, _down(fminus_real, state), 0.0f); + auto minus_imag = bsk::where(keep_down, _down(fminus_imag, state), 0.0f); + plus_real = bsk::where((state == 0), minus_real, plus_real); + plus_imag = bsk::where((state == 0), (-minus_imag), plus_imag); + return bsk::make_tup(plus_real, plus_imag, minus_real, minus_imag); +} + +// Which row of the stacked tables this pulse reads. +// +// Its own shape's block of ``locations`` rows, then the voxel's place along +// the slice. +template +BSK_HD auto _table_row(const T0& profile_index, const T1& event, const T2& location, const T3& locations) { + return ((bsk::cast(bsk::ld((profile_index + event))) * locations) + location); +} + +// The sum of a tile over the entries ``mask`` keeps. +template +BSK_HD auto _total(const T0& value, const T1& mask) { + return bsk::sum_y(bsk::sum_x(bsk::where(mask, value, 0.0f))); +} + +BSK_HD void _pooled_kernel(float* m0, float* b1, float* b1_phase, float* b0, float* efficiency, float* diffusion, float* velocity, float* dm0, float* db1, float* db1_phase, float* db0, float* defficiency, float* ddiffusion, float* dvelocity, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* dduration, float* dflip, float* dphase, float* table, float* dtable, std::int32_t* pool_index, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* dpairs, float* output_real, float* output_imag, float* trajectory, std::int64_t base, std::int64_t atom_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t rows, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t locations, std::int64_t profile_bins, std::int64_t lineshape_bins, std::int64_t n, std::int64_t m, std::int64_t blocks, std::int64_t planes, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t profiled, std::int64_t dynamic, std::int64_t directed_pairs, std::int64_t directed_table, std::int64_t following, std::int64_t keep, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t P, std::int64_t S) { + bsk::tup a{}; + bsk::tup b{}; + bsk::V dfmi{}; + bsk::V dfmr{}; + bsk::V dfpi{}; + bsk::V dfpr{}; + bsk::V dzi{}; + bsk::V dzr{}; + bsk::V fmi{}; + bsk::V fmr{}; + bsk::V fpi{}; + bsk::V fpr{}; + bsk::tup pulse_b1{}; + bsk::tup pulse_b1_phase{}; + bsk::tup read{}; + bsk::tup spun{}; + bsk::tup, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V, bsk::V> turned{}; + bsk::tup washout_rate{}; + bsk::V zi{}; + bsk::V zr{}; + auto problem = (bsk::cast(bsk::program_id(0)) + base); + auto atom = bsk::mod(problem, atom_count); + auto train = bsk::floordiv(problem, atom_count); + auto event_base = (train * event_count); + auto voxel_at = (atom * atom_stride); + auto location = bsk::mod(atom, locations); + auto live = (problem >= 0); + auto pool = bsk::arange_y(); + auto column = bsk::arange_z(); + auto state = bsk::arange_x(); + auto state_mask = (state < state_count); + auto order = bsk::cast(state); + auto exchanging = (pool < m); + auto semisolid = (pool == (n - 1)); + auto row_width = (((n * n) + n) + ((2 * m) * m)); + auto width = (n + ((rows * row_width) * blocks)); + auto slot = (table + (atom * width)); + auto directions = (dtable + (atom * width)); + auto density_of = _read(m0, dm0, voxel_at, density, 1.0f, following); + auto voxel_b1 = _read(b1, db1, voxel_at, transmit, 1.0f, following); + auto voxel_b1_phase = _read(b1_phase, db1_phase, voxel_at, off_axis, 0.0f, following); + auto voxel_b0 = _read(b0, db0, voxel_at, off_axis, 0.0f, following); + auto inversion = _read(efficiency, defficiency, voxel_at, inverting, 1.0f, following); + auto damping_rate = _read(diffusion, ddiffusion, voxel_at, diffusing, 0.0f, following); + auto moved = _read(velocity, dvelocity, voxel_at, moving, 0.0f, following); + auto flow_rate = bsk::make_tup((flow_scale * bsk::get<0>(moved)), (flow_scale * bsk::get<1>(moved))); + washout_rate = bsk::make_tup(0.0f, 0.0f); + if (bsk::truth(moving)) { + auto heading = (bsk::where((bsk::get<0>(moved) > 0.0f), 1.0f, 0.0f) - bsk::where((bsk::get<0>(moved) < 0.0f), 1.0f, 0.0f)); + washout_rate = bsk::make_tup((washout_scale * bsk::abs(bsk::get<0>(moved))), ((washout_scale * heading) * bsk::get<1>(moved))); + } + auto equilibrium = _entries(slot, directions, pool, (pool < n), 0.0f, 0, directed_table, false, following); + auto zero = bsk::full(0); + fpr = zero; + fpi = zero; + fmr = zero; + fmi = zero; + zr = bsk::where((state == 0), bsk::get<0>(equilibrium), 0.0f); + zi = zero; + dfpr = zero; + dfpi = zero; + dfmr = zero; + dfmi = zero; + dzr = zero; + dzi = zero; + if (bsk::truth(following)) { + dzr = bsk::where((state == 0), bsk::get<1>(equilibrium), 0.0f); + } + auto tile = ((pool * S) + state); + for (std::int64_t event = 0; event < event_count; event += 1) { + if (bsk::truth(keep)) { + auto at = (trajectory + ((((problem - base) * event_count) + event) * ((planes * P) * S))); + bsk::st(((at + ((0 * P) * S)) + tile), fpr); + bsk::st(((at + ((1 * P) * S)) + tile), fpi); + bsk::st(((at + ((2 * P) * S)) + tile), fmr); + bsk::st(((at + ((3 * P) * S)) + tile), fmi); + bsk::st(((at + ((4 * P) * S)) + tile), zr); + bsk::st(((at + ((5 * P) * S)) + tile), zi); + if (bsk::truth(following)) { + bsk::st(((at + ((6 * P) * S)) + tile), dfpr); + bsk::st(((at + ((7 * P) * S)) + tile), dfpi); + bsk::st(((at + ((8 * P) * S)) + tile), dfmr); + bsk::st(((at + ((9 * P) * S)) + tile), dfmi); + bsk::st(((at + ((10 * P) * S)) + tile), dzr); + bsk::st(((at + ((11 * P) * S)) + tile), dzi); + } + } + auto dt = _read((duration + event_base), (dduration + event_base), event, true, 0.0f, following); + auto t0_ = _factors(dt, damping_rate, voxel_b0, flow_rate, washout_rate, order, off_axis, moving, diffusing, following); + auto wout = bsk::get<0>(t0_); + auto unit_t = bsk::get<1>(t0_); + auto unit_z = bsk::get<2>(t0_); + auto _squared = bsk::get<3>(t0_); + auto _weight = bsk::get<4>(t0_); + auto carried = _cscale(wout, unit_t, following); + auto spin = _cscale(wout, unit_z, following); + auto row = bsk::cast(bsk::ld(((pool_index + event_base) + event))); + auto t1_ = _operators(slot, directions, (n + (row * row_width)), bsk::get<1>(dt), (rows * row_width), pool, column, n, m, directed_table, (blocks > 1), following); + auto longitudinal_op = bsk::get<0>(t1_); + auto restored = bsk::get<1>(t1_); + auto transverse_op = bsk::get<2>(t1_); + auto t2_ = _relax(bsk::make_tup(fpr, fpi, dfpr, dfpi), bsk::make_tup(fmr, fmi, dfmr, dfmi), bsk::make_tup(zr, zi, dzr, dzi), transverse_op, longitudinal_op, restored, equilibrium, wout, carried, spin, state, following); + auto plus = bsk::get<0>(t2_); + auto minus = bsk::get<1>(t2_); + auto longitudinal = bsk::get<2>(t2_); + auto _mp = bsk::get<3>(t2_); + auto _mm = bsk::get<4>(t2_); + auto _mz = bsk::get<5>(t2_); + fpr = bsk::get<0>(plus); + fpi = bsk::get<1>(plus); + fmr = bsk::get<0>(minus); + fmi = bsk::get<1>(minus); + zr = bsk::get<0>(longitudinal); + zi = bsk::get<1>(longitudinal); + if (bsk::truth(following)) { + dfpr = bsk::get<2>(plus); + dfpi = bsk::get<3>(plus); + dfmr = bsk::get<2>(minus); + dfmi = bsk::get<3>(minus); + dzr = bsk::get<2>(longitudinal); + dzi = bsk::get<3>(longitudinal); + } + auto event_action = bsk::cast(bsk::ld((action + event))); + auto event_kind = bsk::cast(bsk::ld((kind + event))); + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t3_ = _shift(fpr, fpi, fmr, fmi, state, state_mask, state_count); + fpr = bsk::get<0>(t3_); + fpi = bsk::get<1>(t3_); + fmr = bsk::get<2>(t3_); + fmi = bsk::get<3>(t3_); + if (bsk::truth(following)) { + auto t4_ = _shift(dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count); + dfpr = bsk::get<0>(t4_); + dfpi = bsk::get<1>(t4_); + dfmr = bsk::get<2>(t4_); + dfmi = bsk::get<3>(t4_); + } + } + if (bsk::truth((event_kind == 1))) { + if (bsk::truth((bsk::band(event_action, 4) != 0))) { + // Every exchanging pool is free water and inverts like it; a + // semisolid one is saturated by the pulse's own term. + if (bsk::truth(following)) { + dzr = bsk::where(exchanging, (-((bsk::get<1>(inversion) * zr) + (bsk::get<0>(inversion) * dzr))), dzr); + dzi = bsk::where(exchanging, (-((bsk::get<1>(inversion) * zi) + (bsk::get<0>(inversion) * dzi))), dzi); + } + zr = bsk::where(exchanging, ((-bsk::get<0>(inversion)) * zr), zr); + zi = bsk::where(exchanging, ((-bsk::get<0>(inversion)) * zi), zi); + } else { + pulse_b1 = voxel_b1; + pulse_b1_phase = voxel_b1_phase; + if (bsk::truth(shimmed)) { + auto transmit_at = ((bsk::cast(bsk::ld((shim_index + event))) * atom_count) + atom); + pulse_b1 = _read(b1, db1, transmit_at, transmit, 1.0f, following); + pulse_b1_phase = _read(b1_phase, db1_phase, transmit_at, true, 0.0f, following); + } + auto nominal = _read((flip + event_base), (dflip + event_base), event, true, 0.0f, following); + auto played = _read((phase + event_base), (dphase + event_base), event, true, 0.0f, following); + auto alpha = _rmul(nominal, pulse_b1, following); + auto phi = bsk::make_tup((bsk::get<0>(played) + bsk::get<0>(pulse_b1_phase)), (bsk::get<1>(played) + bsk::get<1>(pulse_b1_phase))); + if (bsk::truth((n > m))) { + auto t5_ = _absorption(lineshape, rf_frequency, saturation, event, alpha, voxel_b0, lineshape_bins, lineshape_step, following); + auto absorbed = bsk::get<0>(t5_); + auto _shape = bsk::get<1>(t5_); + auto _slope = bsk::get<2>(t5_); + auto _deposited = bsk::get<3>(t5_); + if (bsk::truth(following)) { + dzr = bsk::where(semisolid, ((bsk::get<1>(absorbed) * zr) + (bsk::get<0>(absorbed) * dzr)), dzr); + dzi = bsk::where(semisolid, ((bsk::get<1>(absorbed) * zi) + (bsk::get<0>(absorbed) * dzi)), dzi); + } + zr = bsk::where(semisolid, (bsk::get<0>(absorbed) * zr), zr); + zi = bsk::where(semisolid, (bsk::get<0>(absorbed) * zi), zi); + } + if (bsk::truth(dynamic)) { + if (bsk::truth(following)) { + auto t6_ = _dynamic_pair_dual_at(pairs, dpairs, pair_index, event_base, event, atom, atom_count, live, bsk::get<0>(phi), bsk::get<1>(phi), directed_pairs); + a = bsk::get<0>(t6_); + spun = bsk::get<1>(t6_); + } else { + auto held = _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, live); + a = bsk::make_tup(bsk::get<0>(held), bsk::get<1>(held), 0.0f, 0.0f); + spun = _cmul(bsk::make_tup(bsk::get<2>(held), bsk::get<3>(held), 0.0f, 0.0f), _polar(bsk::make_tup((-bsk::get<0>(phi)), 0.0f), following), following); + } + } else if (bsk::truth(profiled)) { + auto at_row = _table_row(profile_index, event, location, locations); + auto turn = _polar(bsk::make_tup((-bsk::get<0>(phi)), (-bsk::get<1>(phi))), following); + if (bsk::truth(following)) { + read = _profile_pair_slope(profile, at_row, bsk::get<0>(alpha), profile_bins, profile_step); + a = bsk::make_tup(bsk::get<0>(read), bsk::get<2>(read), (bsk::get<1>(read) * bsk::get<1>(alpha)), (bsk::get<3>(read) * bsk::get<1>(alpha))); + b = bsk::make_tup(bsk::get<4>(read), bsk::get<6>(read), (bsk::get<5>(read) * bsk::get<1>(alpha)), (bsk::get<7>(read) * bsk::get<1>(alpha))); + } else { + read = _profile_pair(profile, at_row, bsk::get<0>(alpha), profile_bins, profile_step); + a = bsk::make_tup(bsk::get<0>(read), bsk::get<1>(read), 0.0f, 0.0f); + b = bsk::make_tup(bsk::get<2>(read), bsk::get<3>(read), 0.0f, 0.0f); + } + spun = _cmul(b, turn, following); + } else { + auto t7_ = _hard_pair(alpha, phi, following); + a = bsk::get<0>(t7_); + spun = bsk::get<1>(t7_); + auto _sa = bsk::get<2>(t7_); + auto _sb = bsk::get<3>(t7_); + auto _turn = bsk::get<4>(t7_); + } + if (bsk::truth(following)) { + turned = _rotate_spinor_dual(bsk::get<0>(a), bsk::get<1>(a), bsk::get<0>(spun), bsk::get<1>(spun), bsk::get<2>(a), bsk::get<3>(a), bsk::get<2>(spun), bsk::get<3>(spun), fpr, fpi, fmr, fmi, zr, zi, dfpr, dfpi, dfmr, dfmi, dzr, dzi); + dfpr = bsk::where(exchanging, bsk::get<6>(turned), dfpr); + dfpi = bsk::where(exchanging, bsk::get<7>(turned), dfpi); + dfmr = bsk::where(exchanging, bsk::get<8>(turned), dfmr); + dfmi = bsk::where(exchanging, bsk::get<9>(turned), dfmi); + dzr = bsk::where(exchanging, bsk::get<10>(turned), dzr); + dzi = bsk::where(exchanging, bsk::get<11>(turned), dzi); + } else { + turned = _rotate_spinor(bsk::get<0>(a), bsk::get<1>(a), bsk::get<0>(spun), bsk::get<1>(spun), fpr, fpi, fmr, fmi, zr, zi); + } + fpr = bsk::where(exchanging, bsk::get<0>(turned), fpr); + fpi = bsk::where(exchanging, bsk::get<1>(turned), fpi); + fmr = bsk::where(exchanging, bsk::get<2>(turned), fmr); + fmi = bsk::where(exchanging, bsk::get<3>(turned), fmi); + zr = bsk::where(exchanging, bsk::get<4>(turned), zr); + zi = bsk::where(exchanging, bsk::get<5>(turned), zi); + } + } + if (bsk::truth((!bsk::truth(keep)))) { + if (bsk::truth(bsk::band((event_kind == 2), (bsk::band(event_action, 32) != 0)))) { + auto origin = (state == 0); + auto recorded = bsk::make_tup(_total(fpr, origin), _total(fpi, origin), _total(dfpr, origin), _total(dfpi, origin)); + auto read_phase = _read((phase + event_base), (dphase + event_base), event, true, 0.0f, following); + auto signal_ = _cscale(density_of, _cmul(recorded, _polar(bsk::make_tup((-bsk::get<0>(read_phase)), (-bsk::get<1>(read_phase))), following), following), following); + auto out_ = bsk::ld((output_index + event)); + auto written = ((problem * output_count) + out_); + if (bsk::truth(following)) { + bsk::st((output_real + written), bsk::get<2>(signal_), (out_ >= 0)); + bsk::st((output_imag + written), bsk::get<3>(signal_), (out_ >= 0)); + } else { + bsk::st((output_real + written), bsk::get<0>(signal_), (out_ >= 0)); + bsk::st((output_imag + written), bsk::get<1>(signal_), (out_ >= 0)); + } + } + } + if (bsk::truth((bsk::band(event_action, 2) != 0))) { + auto t8_ = _shift(fpr, fpi, fmr, fmi, state, state_mask, state_count); + fpr = bsk::get<0>(t8_); + fpi = bsk::get<1>(t8_); + fmr = bsk::get<2>(t8_); + fmi = bsk::get<3>(t8_); + if (bsk::truth(following)) { + auto t9_ = _shift(dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count); + dfpr = bsk::get<0>(t9_); + dfpi = bsk::get<1>(t9_); + dfmr = bsk::get<2>(t9_); + dfmi = bsk::get<3>(t9_); + } + } + if (bsk::truth((bsk::band(event_action, 8) != 0))) { + fpr = zero; + fpi = zero; + fmr = zero; + fmi = zero; + if (bsk::truth(following)) { + dfpr = zero; + dfpi = zero; + dfmr = zero; + dfmi = zero; + } + } else if (bsk::truth((bsk::band(event_action, 16) != 0))) { + auto t10_ = _shift(fpr, fpi, fmr, fmi, state, state_mask, state_count); + fpr = bsk::get<0>(t10_); + fpi = bsk::get<1>(t10_); + fmr = bsk::get<2>(t10_); + fmi = bsk::get<3>(t10_); + if (bsk::truth(following)) { + auto t11_ = _shift(dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count); + dfpr = bsk::get<0>(t11_); + dfpi = bsk::get<1>(t11_); + dfmr = bsk::get<2>(t11_); + dfmi = bsk::get<3>(t11_); + } + } + } +} + +template +BSK_HD auto _cadd(const T0& x, const T1& y) { + return bsk::make_tup((bsk::get<0>(x) + bsk::get<0>(y)), (bsk::get<1>(x) + bsk::get<1>(y)), (bsk::get<2>(x) + bsk::get<2>(y)), (bsk::get<3>(x) + bsk::get<3>(y))); +} + +// ``sum_k left[i, k] right[j, k]``: the operator a pair of tiles makes. +template +BSK_HD auto _outer(const T0& left, const T1& right) { + return bsk::outer(left, right); +} + +// ``sum_k left[i, k] right[j, k]`` of two complex dual tiles. +template +BSK_HD auto _couter(const T0& left, const T1& right, const T2& following) { + using Ret = bsk::tup | 4, 6)>, bsk::tile_t | 4, 6)>, bsk::tile_t | 4, 6)>, bsk::tile_t | 4, 6)>>; + auto real = (_outer(bsk::get<0>(left), bsk::get<0>(right)) - _outer(bsk::get<1>(left), bsk::get<1>(right))); + auto imag = (_outer(bsk::get<0>(left), bsk::get<1>(right)) + _outer(bsk::get<1>(left), bsk::get<0>(right))); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(real, imag, (((_outer(bsk::get<2>(left), bsk::get<0>(right)) - _outer(bsk::get<3>(left), bsk::get<1>(right))) + _outer(bsk::get<0>(left), bsk::get<2>(right))) - _outer(bsk::get<1>(left), bsk::get<3>(right))), (((_outer(bsk::get<2>(left), bsk::get<1>(right)) + _outer(bsk::get<3>(left), bsk::get<0>(right))) + _outer(bsk::get<0>(left), bsk::get<3>(right))) + _outer(bsk::get<1>(left), bsk::get<2>(right))))); + } else { + return bsk::convert(bsk::make_tup(real, imag, 0.0f, 0.0f)); + } +} + +// The pair, its slope and its curvature in the flip angle. +// +// The second-order pass differentiates the read twice, and a Hermite segment +// is a cubic, so all three come from the same four knot values. Returned in +// threes per component: value, slope, curvature. +template +BSK_HD auto _profile_pair_curve(const T0& profile, const T1& row, const T2& theta, const T3& bins, const T4& step) { + auto last = (bins - 1); + auto scaled = bsk::minimum(bsk::maximum(bsk::truediv(theta, step), 0.0f), (last + 0.0f)); + auto lower = bsk::minimum(bsk::floor(scaled), (last - 1.0f)); + auto u = (scaled - lower); + auto u2 = (u * u); + auto u3 = (u2 * u); + auto h00 = (((2.0f * u3) - (3.0f * u2)) + 1.0f); + auto h10 = (((u3 - (2.0f * u2)) + u) * step); + auto h01 = (((-2.0f) * u3) + (3.0f * u2)); + auto h11 = ((u3 - u2) * step); + auto g00 = bsk::truediv(((6.0f * u2) - (6.0f * u)), step); + auto g10 = (((3.0f * u2) - (4.0f * u)) + 1.0f); + auto g01 = bsk::truediv(((6.0f * u) - (6.0f * u2)), step); + auto g11 = ((3.0f * u2) - (2.0f * u)); + auto c00 = bsk::truediv(((12.0f * u) - 6.0f), (step * step)); + auto c10 = bsk::truediv(((6.0f * u) - 4.0f), step); + auto c01 = bsk::truediv((6.0f - (12.0f * u)), (step * step)); + auto c11 = bsk::truediv(((6.0f * u) - 2.0f), step); + auto base = (((row * bins) + bsk::cast(lower)) * 8); + auto near = [&](int c) { return bsk::ld(((profile + base) + c)); }; + auto near_slope = [&](int c) { return bsk::ld((((profile + base) + 4) + c)); }; + auto far = [&](int c) { return bsk::ld((((profile + base) + 8) + c)); }; + auto far_slope = [&](int c) { return bsk::ld((((profile + base) + 12) + c)); }; + auto value = [&](int c) { + return ((((h00 * near(c)) + (h10 * near_slope(c))) + (h01 * far(c))) + (h11 * far_slope(c))); + }; + auto slope = [&](int c) { + return ((((g00 * near(c)) + (g10 * near_slope(c))) + (g01 * far(c))) + (g11 * far_slope(c))); + }; + auto curve = [&](int c) { + return ((((c00 * near(c)) + (c10 * near_slope(c))) + (c01 * far(c))) + (c11 * far_slope(c))); + }; + return bsk::make_tup(value(0), slope(0), curve(0), value(1), slope(1), curve(1), + value(2), slope(2), curve(2), value(3), slope(3), curve(3)); +} + +// The pair a shaped pulse turns through, and its slope, as duals. +// +// The flip angle carries the tangent into the table, so the pair's tangent is +// the stored slope and the slope's own tangent is the segment's curvature. +// The RF phase turns the axis once the pair is out, and so reaches ``b``. +template +BSK_HD auto _profiled_pair_dual(const T0& profile, const T1& row, const T2& alpha_value, const T3& alpha_tangent, const T4& phi_value, const T5& phi_tangent, const T6& bins, const T7& step) { + auto read = _profile_pair_curve(profile, row, alpha_value, bins, step); + auto a = bsk::make_tup(bsk::get<0>(read), bsk::get<3>(read), (bsk::get<1>(read) * alpha_tangent), (bsk::get<4>(read) * alpha_tangent)); + auto slope_a = bsk::make_tup(bsk::get<1>(read), bsk::get<4>(read), (bsk::get<2>(read) * alpha_tangent), (bsk::get<5>(read) * alpha_tangent)); + auto b = bsk::make_tup(bsk::get<6>(read), bsk::get<9>(read), (bsk::get<7>(read) * alpha_tangent), (bsk::get<10>(read) * alpha_tangent)); + auto slope_b = bsk::make_tup(bsk::get<7>(read), bsk::get<10>(read), (bsk::get<8>(read) * alpha_tangent), (bsk::get<11>(read) * alpha_tangent)); + auto turn = _dual_polar((-phi_value), (-phi_tangent)); + return bsk::make_tup(a, _dual_product(b, turn), slope_a, _dual_product(slope_b, turn)); +} + +// ``Re(conj(x) y)`` as a real dual. +template +BSK_HD auto _re_dot(const T0& x, const T1& y, const T2& following) { + using Ret = bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>>; + auto value = ((bsk::get<0>(x) * bsk::get<0>(y)) + (bsk::get<1>(x) * bsk::get<1>(y))); + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(value, ((((bsk::get<2>(x) * bsk::get<0>(y)) + (bsk::get<3>(x) * bsk::get<1>(y))) + (bsk::get<0>(x) * bsk::get<2>(y))) + (bsk::get<1>(x) * bsk::get<3>(y))))); + } else { + return bsk::convert(bsk::make_tup(value, 0.0f)); + } +} + +// A real dual tile summed to a real dual scalar. +template +BSK_HD auto _rtotal(const T0& x, const T1& mask, const T2& following) { + using Ret = bsk::tup; + if (bsk::truth(following)) { + return bsk::convert(bsk::make_tup(_total(bsk::get<0>(x), mask), _total(bsk::get<1>(x), mask))); + } else { + return bsk::convert(bsk::make_tup(_total(bsk::get<0>(x), mask), 0.0f)); + } +} + +// Order zero of ``values``, spread across every order. +template +BSK_HD auto _first(const T0& values, const T1& state) { + return bsk::gather_x(values, state * 0); +} + +// Transpose of ``_shift``. +// +// The conjugate refill at order zero sends the incoming plus adjoint back +// onto minus, conjugated, at the index the minus shift moves it to. +template +BSK_HD auto _shift_adjoint(const T0& plus_bar_real, const T1& plus_bar_imag, const T2& minus_bar_real, const T3& minus_bar_imag, const T4& state, const T5& state_mask, const T6& state_count) { + bsk::tile_t | 0, 3)> shifted_mi{}; + bsk::tile_t | 0, 3)> shifted_mr{}; + auto carry_real = bsk::where(state_mask, _first(plus_bar_real, state), 0.0f); + auto carry_imag = (-bsk::where(state_mask, _first(plus_bar_imag, state), 0.0f)); + auto forward = bsk::band(((state + 1) < state_count), state_mask); + auto backward = bsk::band((state > 0), state_mask); + auto shifted_pr = bsk::where(forward, _down(plus_bar_real, state), 0.0f); + auto shifted_pi = bsk::where(forward, _down(plus_bar_imag, state), 0.0f); + shifted_mr = bsk::where(backward, _up(minus_bar_real, state), 0.0f); + shifted_mi = bsk::where(backward, _up(minus_bar_imag, state), 0.0f); + shifted_mr = bsk::where((state == 1), (shifted_mr + carry_real), shifted_mr); + shifted_mi = bsk::where((state == 1), (shifted_mi + carry_imag), shifted_mi); + return bsk::make_tup(shifted_pr, shifted_pi, shifted_mr, shifted_mi); +} + +// The spinor rotation's adjoint, carrying no forward direction. +// +// Returns the cotangent on the Cayley-Klein pair and the three state +// cotangents sent back through the conjugate transpose. Every entry of the +// matrix is a product of two factors drawn from the pair and its conjugate, +// so the pair's two Wirtinger halves are linear in the outer product of the +// seed with the state the rotation acted on -- a closed form rather than a +// differentiated matrix. +template +BSK_HD auto _spinor_adjoint(const T0& ar, const T1& ai, const T2& br, const T3& bi, const T4& spr, const T5& spi, const T6& smr, const T7& smi, const T8& rzr, const T9& rzi, const T10& pbr, const T11& pbi, const T12& mbr, const T13& mbi, const T14& zbr, const T15& zbi) { + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n0{}; + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n1{}; + bsk::tup | 0, 3)>, bsk::tile_t | 0, 3)>> n2{}; + auto aa_r = ((ar * ar) - (ai * ai)); + auto aa_i = ((2.0f * ar) * ai); + auto bb_r = ((br * br) - (bi * bi)); + auto bb_i = ((2.0f * br) * bi); + auto ab_r = ((ar * br) - (ai * bi)); + auto ab_i = ((ar * bi) + (ai * br)); + auto cross_r = ((ar * br) + (ai * bi)); + auto cross_i = ((ar * bi) - (ai * br)); + auto t0_ = bsk::make_tup(aa_r, (-aa_i)); + auto t00_r = bsk::get<0>(t0_); + auto t00_i = bsk::get<1>(t0_); + auto t1_ = bsk::make_tup((-bb_r), bb_i); + auto t01_r = bsk::get<0>(t1_); + auto t01_i = bsk::get<1>(t1_); + auto t2_ = bsk::make_tup((-2.0f * ab_r), (2.0f * ab_i)); + auto t02_r = bsk::get<0>(t2_); + auto t02_i = bsk::get<1>(t2_); + auto t3_ = bsk::make_tup((-bb_r), (-bb_i)); + auto t10_r = bsk::get<0>(t3_); + auto t10_i = bsk::get<1>(t3_); + auto t4_ = bsk::make_tup(aa_r, aa_i); + auto t11_r = bsk::get<0>(t4_); + auto t11_i = bsk::get<1>(t4_); + auto t5_ = bsk::make_tup((-2.0f * ab_r), (-2.0f * ab_i)); + auto t12_r = bsk::get<0>(t5_); + auto t12_i = bsk::get<1>(t5_); + auto t6_ = bsk::make_tup(cross_r, cross_i); + auto t20_r = bsk::get<0>(t6_); + auto t20_i = bsk::get<1>(t6_); + auto t7_ = bsk::make_tup(cross_r, (-cross_i)); + auto t21_r = bsk::get<0>(t7_); + auto t21_i = bsk::get<1>(t7_); + auto t22 = ((((ar * ar) + (ai * ai)) - (br * br)) - (bi * bi)); + // ``m[i][j] = conj(seed_i) * state_j``: the outer product the pair's + // derivative is linear in. + auto m00 = _complex_mul(pbr, (-pbi), spr, spi); + auto m01 = _complex_mul(pbr, (-pbi), smr, smi); + auto m02 = _complex_mul(pbr, (-pbi), rzr, rzi); + auto m10 = _complex_mul(mbr, (-mbi), spr, spi); + auto m11 = _complex_mul(mbr, (-mbi), smr, smi); + auto m12 = _complex_mul(mbr, (-mbi), rzr, rzi); + auto m20 = _complex_mul(zbr, (-zbi), spr, spi); + auto m21 = _complex_mul(zbr, (-zbi), smr, smi); + auto m22 = _complex_mul(zbr, (-zbi), rzr, rzi); + auto hca = _complex_mul(ar, ai, bsk::get<0>(m11), bsk::get<1>(m11)); + auto hcb = _complex_mul(br, bi, bsk::get<0>(m12), bsk::get<1>(m12)); + auto hcc = _complex_mul(br, (-bi), bsk::get<0>(m21), bsk::get<1>(m21)); + auto hcd = _complex_mul(ar, (-ai), bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_conj_a_r = ((((2.0f * bsk::get<0>(hca)) - (2.0f * bsk::get<0>(hcb))) + bsk::get<0>(hcc)) + bsk::get<0>(hcd)); + auto holding_conj_a_i = ((((2.0f * bsk::get<1>(hca)) - (2.0f * bsk::get<1>(hcb))) + bsk::get<1>(hcc)) + bsk::get<1>(hcd)); + auto ha = _complex_mul(ar, (-ai), bsk::get<0>(m00), bsk::get<1>(m00)); + auto hb = _complex_mul(br, (-bi), bsk::get<0>(m02), bsk::get<1>(m02)); + auto hc = _complex_mul(br, bi, bsk::get<0>(m20), bsk::get<1>(m20)); + auto hd = _complex_mul(ar, ai, bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_a_r = ((((2.0f * bsk::get<0>(ha)) - (2.0f * bsk::get<0>(hb))) + bsk::get<0>(hc)) + bsk::get<0>(hd)); + auto holding_a_i = ((((2.0f * bsk::get<1>(ha)) - (2.0f * bsk::get<1>(hb))) + bsk::get<1>(hc)) + bsk::get<1>(hd)); + auto ka = _complex_mul(br, bi, bsk::get<0>(m10), bsk::get<1>(m10)); + auto kb = _complex_mul(ar, ai, bsk::get<0>(m12), bsk::get<1>(m12)); + auto kc = _complex_mul(ar, (-ai), bsk::get<0>(m20), bsk::get<1>(m20)); + auto kd = _complex_mul(br, (-bi), bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_conj_b_r = ((((-2.0f * bsk::get<0>(ka)) - (2.0f * bsk::get<0>(kb))) + bsk::get<0>(kc)) - bsk::get<0>(kd)); + auto holding_conj_b_i = ((((-2.0f * bsk::get<1>(ka)) - (2.0f * bsk::get<1>(kb))) + bsk::get<1>(kc)) - bsk::get<1>(kd)); + auto la = _complex_mul(br, (-bi), bsk::get<0>(m01), bsk::get<1>(m01)); + auto lb = _complex_mul(ar, (-ai), bsk::get<0>(m02), bsk::get<1>(m02)); + auto lc = _complex_mul(ar, ai, bsk::get<0>(m21), bsk::get<1>(m21)); + auto ld_ = _complex_mul(br, bi, bsk::get<0>(m22), bsk::get<1>(m22)); + auto holding_b_r = ((((-2.0f * bsk::get<0>(la)) - (2.0f * bsk::get<0>(lb))) + bsk::get<0>(lc)) - bsk::get<0>(ld_)); + auto holding_b_i = ((((-2.0f * bsk::get<1>(la)) - (2.0f * bsk::get<1>(lb))) + bsk::get<1>(lc)) - bsk::get<1>(ld_)); + auto grad_a_r = (holding_conj_a_r + holding_a_r); + auto grad_a_i = ((-holding_conj_a_i) + holding_a_i); + auto grad_b_r = (holding_conj_b_r + holding_b_r); + auto grad_b_i = ((-holding_conj_b_i) + holding_b_i); + n0 = _complex_mul(t00_r, (-t00_i), pbr, pbi); + n1 = _complex_mul(t10_r, (-t10_i), mbr, mbi); + n2 = _complex_mul(t20_r, (-t20_i), zbr, zbi); + auto t8_ = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2))); + auto next_pr = bsk::get<0>(t8_); + auto next_pi = bsk::get<1>(t8_); + n0 = _complex_mul(t01_r, (-t01_i), pbr, pbi); + n1 = _complex_mul(t11_r, (-t11_i), mbr, mbi); + n2 = _complex_mul(t21_r, (-t21_i), zbr, zbi); + auto t9_ = bsk::make_tup(((bsk::get<0>(n0) + bsk::get<0>(n1)) + bsk::get<0>(n2)), ((bsk::get<1>(n0) + bsk::get<1>(n1)) + bsk::get<1>(n2))); + auto next_mr = bsk::get<0>(t9_); + auto next_mi = bsk::get<1>(t9_); + n0 = _complex_mul(t02_r, (-t02_i), pbr, pbi); + n1 = _complex_mul(t12_r, (-t12_i), mbr, mbi); + auto next_zr = ((bsk::get<0>(n0) + bsk::get<0>(n1)) + (t22 * zbr)); + auto next_zi = ((bsk::get<1>(n0) + bsk::get<1>(n1)) + (t22 * zbi)); + return bsk::make_tup(grad_a_r, grad_a_i, grad_b_r, grad_b_i, next_pr, next_pi, next_mr, next_mi, next_zr, next_zi); +} + +// A dual complex number's conjugate, both halves. +template +BSK_HD auto _dual_conj(const T0& z) { + return bsk::make_tup(bsk::get<0>(z), (-bsk::get<1>(z)), bsk::get<2>(z), (-bsk::get<3>(z))); +} + +// Four dual complex numbers added. +template +BSK_HD auto _dual_sum(const T0& first, const T1& second, const T2& third, const T3& fourth) { + return bsk::make_tup((((bsk::get<0>(first) + bsk::get<0>(second)) + bsk::get<0>(third)) + bsk::get<0>(fourth)), (((bsk::get<1>(first) + bsk::get<1>(second)) + bsk::get<1>(third)) + bsk::get<1>(fourth)), (((bsk::get<2>(first) + bsk::get<2>(second)) + bsk::get<2>(third)) + bsk::get<2>(fourth)), (((bsk::get<3>(first) + bsk::get<3>(second)) + bsk::get<3>(third)) + bsk::get<3>(fourth))); +} + +// A dual complex number scaled by a real constant. +template +BSK_HD auto _dual_weigh(const T0& z, const T1& factor) { + return bsk::make_tup((factor * bsk::get<0>(z)), (factor * bsk::get<1>(z)), (factor * bsk::get<2>(z)), (factor * bsk::get<3>(z))); +} + +// The spinor rotation's adjoint, on dual numbers. +// +// Returns the cotangent on the Cayley-Klein pair and the three state +// cotangents sent back through the conjugate transpose. Every entry of the +// matrix is a product of two factors drawn from the pair and its conjugate, +// so the pair's two Wirtinger halves are linear in the outer product of the +// seed with the state the rotation acted on -- which is why this is a closed +// form rather than a differentiated matrix. +template +BSK_HD auto _spinor_adjoint_dual(const T0& a, const T1& b, const T2& sp, const T3& sm, const T4& rz, const T5& pb, const T6& mb, const T7& zb) { + auto t0_ = _spinor_coefficients(bsk::get<0>(a), bsk::get<1>(a), bsk::get<0>(b), bsk::get<1>(b), bsk::get<2>(a), bsk::get<3>(a), bsk::get<2>(b), bsk::get<3>(b)); + auto t00 = bsk::get<0>(t0_); + auto t01 = bsk::get<1>(t0_); + auto t02 = bsk::get<2>(t0_); + auto t10 = bsk::get<3>(t0_); + auto t11 = bsk::get<4>(t0_); + auto t12 = bsk::get<5>(t0_); + auto t20 = bsk::get<6>(t0_); + auto t21 = bsk::get<7>(t0_); + auto t22 = bsk::get<8>(t0_); + auto conj_pb = _dual_conj(pb); + auto conj_mb = _dual_conj(mb); + auto conj_zb = _dual_conj(zb); + auto m00 = _dual_product(conj_pb, sp); + auto m01 = _dual_product(conj_pb, sm); + auto m02 = _dual_product(conj_pb, rz); + auto m10 = _dual_product(conj_mb, sp); + auto m11 = _dual_product(conj_mb, sm); + auto m12 = _dual_product(conj_mb, rz); + auto m20 = _dual_product(conj_zb, sp); + auto m21 = _dual_product(conj_zb, sm); + auto m22 = _dual_product(conj_zb, rz); + auto conj_a = _dual_conj(a); + auto conj_b = _dual_conj(b); + auto holding_conj_a = _dual_sum(_dual_weigh(_dual_product(a, m11), 2.0f), _dual_weigh(_dual_product(b, m12), -2.0f), _dual_product(conj_b, m21), _dual_product(conj_a, m22)); + auto holding_a = _dual_sum(_dual_weigh(_dual_product(conj_a, m00), 2.0f), _dual_weigh(_dual_product(conj_b, m02), -2.0f), _dual_product(b, m20), _dual_product(a, m22)); + auto holding_conj_b = _dual_sum(_dual_weigh(_dual_product(b, m10), -2.0f), _dual_weigh(_dual_product(a, m12), -2.0f), _dual_product(conj_a, m20), _dual_weigh(_dual_product(conj_b, m22), -1.0f)); + auto holding_b = _dual_sum(_dual_weigh(_dual_product(conj_b, m01), -2.0f), _dual_weigh(_dual_product(conj_a, m02), -2.0f), _dual_product(a, m21), _dual_weigh(_dual_product(b, m22), -1.0f)); + auto zero = _dual_weigh(m00, 0.0f); + auto grad_a = _dual_sum(_dual_conj(holding_conj_a), holding_a, zero, zero); + auto grad_b = _dual_sum(_dual_conj(holding_conj_b), holding_b, zero, zero); + auto next_pb = _dual_sum(_dual_product(_dual_conj(t00), pb), _dual_product(_dual_conj(t10), mb), _dual_product(_dual_conj(t20), zb), zero); + auto next_mb = _dual_sum(_dual_product(_dual_conj(t01), pb), _dual_product(_dual_conj(t11), mb), _dual_product(_dual_conj(t21), zb), zero); + auto next_zb = _dual_sum(_dual_product(_dual_conj(t02), pb), _dual_product(_dual_conj(t12), mb), _dual_product(_dual_conj(t22), zb), zero); + return bsk::make_tup(grad_a, grad_b, next_pb, next_mb, next_zb); +} + +BSK_HD void _pooled_adjoint_kernel(float* m0, float* b1, float* b1_phase, float* b0, float* efficiency, float* diffusion, float* velocity, float* dm0, float* db1, float* db1_phase, float* db0, float* defficiency, float* ddiffusion, float* dvelocity, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* dduration, float* dflip, float* dphase, float* table, float* dtable, std::int32_t* pool_index, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* dpairs, float* grad_real, float* grad_imag, float* grad_tissue, float* dgrad_tissue, float* grad_duration, float* dgrad_duration, float* grad_flip, float* dgrad_flip, float* grad_phase, float* dgrad_phase, float* grad_table, float* dgrad_table, float* grad_pairs, float* dgrad_pairs, float* trajectory, std::int64_t base, std::int64_t atom_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t rows, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t locations, std::int64_t profile_bins, std::int64_t lineshape_bins, std::int64_t m0_row, std::int64_t b1_row, std::int64_t b1_phase_row, std::int64_t b0_row, std::int64_t efficiency_row, std::int64_t diffusion_row, std::int64_t velocity_row, std::int64_t n, std::int64_t m, std::int64_t blocks, std::int64_t planes, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t profiled, std::int64_t dynamic, std::int64_t directed_pairs, std::int64_t directed_table, std::int64_t following, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t P, std::int64_t S) { + bsk::tup a{}; + bsk::tup, bsk::V, bsk::V, bsk::V> back_z{}; + float curve_b0{}; + float curve_damping{}; + bsk::V curve_eq{}; + float curve_flow{}; + float curve_washout{}; + bsk::V dmbi{}; + bsk::V dmbr{}; + bsk::V dpbi{}; + bsk::V dpbr{}; + bsk::V dsmi{}; + bsk::V dsmr{}; + bsk::V dspi{}; + bsk::V dspr{}; + bsk::V dzbi{}; + bsk::V dzbr{}; + bsk::tup fraction_grad{}; + bsk::tup grad_a{}; + bsk::tup grad_alpha{}; + bsk::tup grad_angle{}; + bsk::tup grad_b{}; + float grad_b0{}; + float grad_b1{}; + float grad_b1_phase{}; + bsk::tup grad_b_factor{}; + float grad_damping{}; + float grad_efficiency{}; + bsk::V grad_eq{}; + float grad_flow{}; + bsk::V grad_m0{}; + bsk::tup grad_turn{}; + float grad_washout{}; + bsk::tup grad_wout{}; + float heading{}; + std::int32_t held{}; + bsk::tup, bsk::V> longitudinal_grad{}; + bsk::V mbi{}; + bsk::V mbr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> minus_bar{}; + bsk::tup, bsk::V, bsk::V, bsk::V> minus_in{}; + bsk::V pbi{}; + bsk::V pbr{}; + bsk::tup, bsk::V, bsk::V, bsk::V> plus_bar{}; + bsk::tup, bsk::V, bsk::V, bsk::V> plus_in{}; + bsk::tup pulse_b1{}; + bsk::tup pulse_b1_phase{}; + bsk::tup, bsk::V, float, float> seed{}; + bsk::tup slope_a{}; + bsk::tup slope_b{}; + bsk::V smi{}; + bsk::V smr{}; + bsk::V spi{}; + bsk::V spr{}; + bsk::tup spun{}; + bsk::tup table_duration{}; + bsk::tup taken{}; + std::int32_t transmit_row{}; + bsk::tup washout_rate{}; + bsk::tup, bsk::V, bsk::V, bsk::V> z_bar{}; + bsk::tup, bsk::V, bsk::V, bsk::V> z_in{}; + bsk::V zbi{}; + bsk::V zbr{}; + auto problem = (bsk::cast(bsk::program_id(0)) + base); + auto atom = bsk::mod(problem, atom_count); + auto train = bsk::floordiv(problem, atom_count); + auto event_base = (train * event_count); + auto voxel_at = (atom * atom_stride); + auto location = bsk::mod(atom, locations); + auto live = (problem >= 0); + auto pool = bsk::arange_y(); + auto column = bsk::arange_z(); + auto state = bsk::arange_x(); + auto state_mask = (state < state_count); + auto order = bsk::cast(state); + auto origin = (state == 0); + auto exchanging = (pool < m); + auto semisolid = (pool == (n - 1)); + auto held_rows = (pool < n); + auto square = bsk::band((pool < n), (column < n)); + auto across_mask = bsk::band((pool < m), (column < m)); + auto live_exchanging = bsk::band(exchanging, state_mask); + auto row_width = (((n * n) + n) + ((2 * m) * m)); + auto width = (n + ((rows * row_width) * blocks)); + auto slot = (table + (atom * width)); + auto directions = (dtable + (atom * width)); + auto slot_grad = (grad_table + (problem * width)); + auto slot_curve = (dgrad_table + (problem * width)); + auto density_of = _read(m0, dm0, voxel_at, density, 1.0f, following); + auto voxel_b1 = _read(b1, db1, voxel_at, transmit, 1.0f, following); + auto voxel_b1_phase = _read(b1_phase, db1_phase, voxel_at, off_axis, 0.0f, following); + auto voxel_b0 = _read(b0, db0, voxel_at, off_axis, 0.0f, following); + auto inversion = _read(efficiency, defficiency, voxel_at, inverting, 1.0f, following); + auto damping_rate = _read(diffusion, ddiffusion, voxel_at, diffusing, 0.0f, following); + auto moved = _read(velocity, dvelocity, voxel_at, moving, 0.0f, following); + auto flow_rate = bsk::make_tup((flow_scale * bsk::get<0>(moved)), (flow_scale * bsk::get<1>(moved))); + washout_rate = bsk::make_tup(0.0f, 0.0f); + heading = 0.0f; + if (bsk::truth(moving)) { + heading = (bsk::where((bsk::get<0>(moved) > 0.0f), 1.0f, 0.0f) - bsk::where((bsk::get<0>(moved) < 0.0f), 1.0f, 0.0f)); + washout_rate = bsk::make_tup((washout_scale * bsk::abs(bsk::get<0>(moved))), ((washout_scale * heading) * bsk::get<1>(moved))); + } + auto equilibrium = _entries(slot, directions, pool, held_rows, 0.0f, 0, directed_table, false, following); + auto zero = bsk::full(0); + pbr = zero; + pbi = zero; + mbr = zero; + mbi = zero; + zbr = zero; + zbi = zero; + dpbr = zero; + dpbi = zero; + dmbr = zero; + dmbi = zero; + dzbr = zero; + dzbi = zero; + grad_eq = (bsk::get<0>(equilibrium) * 0.0f); + curve_eq = (bsk::get<0>(equilibrium) * 0.0f); + grad_m0 = 0.0f; + grad_b1 = 0.0f; + grad_b1_phase = 0.0f; + grad_b0 = 0.0f; + grad_efficiency = 0.0f; + grad_damping = 0.0f; + grad_flow = 0.0f; + grad_washout = 0.0f; + // Only what every interval adds to is carried in the derivative plane; + // what a pulse or a readout adds is stored as the branch reaches it. + curve_b0 = 0.0f; + curve_damping = 0.0f; + curve_flow = 0.0f; + curve_washout = 0.0f; + // Transmit gradients are summed per shim: the running pair is flushed to + // its row whenever the walk back reaches a pulse on a different one. + held = 0; + auto tile = ((pool * S) + state); + for (std::int64_t step = 0; step < event_count; step += 1) { + auto event = ((event_count - 1) - step); + auto at = (trajectory + ((((problem - base) * event_count) + event) * ((planes * P) * S))); + if (bsk::truth(following)) { + plus_in = bsk::make_tup(bsk::ld(((at + ((0 * P) * S)) + tile)), bsk::ld(((at + ((1 * P) * S)) + tile)), bsk::ld(((at + ((6 * P) * S)) + tile)), bsk::ld(((at + ((7 * P) * S)) + tile))); + minus_in = bsk::make_tup(bsk::ld(((at + ((2 * P) * S)) + tile)), bsk::ld(((at + ((3 * P) * S)) + tile)), bsk::ld(((at + ((8 * P) * S)) + tile)), bsk::ld(((at + ((9 * P) * S)) + tile))); + z_in = bsk::make_tup(bsk::ld(((at + ((4 * P) * S)) + tile)), bsk::ld(((at + ((5 * P) * S)) + tile)), bsk::ld(((at + ((10 * P) * S)) + tile)), bsk::ld(((at + ((11 * P) * S)) + tile))); + } else { + plus_in = bsk::make_tup(bsk::ld((at + tile)), bsk::ld(((at + (P * S)) + tile)), 0.0f, 0.0f); + minus_in = bsk::make_tup(bsk::ld(((at + ((2 * P) * S)) + tile)), bsk::ld(((at + ((3 * P) * S)) + tile)), 0.0f, 0.0f); + z_in = bsk::make_tup(bsk::ld(((at + ((4 * P) * S)) + tile)), bsk::ld(((at + ((5 * P) * S)) + tile)), 0.0f, 0.0f); + } + auto dt = _read((duration + event_base), (dduration + event_base), event, true, 0.0f, following); + auto t0_ = _factors(dt, damping_rate, voxel_b0, flow_rate, washout_rate, order, off_axis, moving, diffusing, following); + auto wout = bsk::get<0>(t0_); + auto unit_t = bsk::get<1>(t0_); + auto unit_z = bsk::get<2>(t0_); + auto squared = bsk::get<3>(t0_); + auto weight = bsk::get<4>(t0_); + auto carried = _cscale(wout, unit_t, following); + auto spin = _cscale(wout, unit_z, following); + auto row = bsk::cast(bsk::ld(((pool_index + event_base) + event))); + auto row_offset = (n + (row * row_width)); + auto t1_ = _operators(slot, directions, row_offset, bsk::get<1>(dt), (rows * row_width), pool, column, n, m, directed_table, (blocks > 1), following); + auto longitudinal_op = bsk::get<0>(t1_); + auto restored = bsk::get<1>(t1_); + auto transverse_op = bsk::get<2>(t1_); + // Replay the interval to recover the states the event acted on. + auto t2_ = _relax(plus_in, minus_in, z_in, transverse_op, longitudinal_op, restored, equilibrium, wout, carried, spin, state, following); + auto relaxed_plus = bsk::get<0>(t2_); + auto relaxed_minus = bsk::get<1>(t2_); + auto relaxed_z = bsk::get<2>(t2_); + auto mixed_plus = bsk::get<3>(t2_); + auto mixed_minus = bsk::get<4>(t2_); + auto mixed_z = bsk::get<5>(t2_); + spr = bsk::get<0>(relaxed_plus); + spi = bsk::get<1>(relaxed_plus); + smr = bsk::get<0>(relaxed_minus); + smi = bsk::get<1>(relaxed_minus); + dspr = zero; + dspi = zero; + dsmr = zero; + dsmi = zero; + if (bsk::truth(following)) { + dspr = bsk::get<2>(relaxed_plus); + dspi = bsk::get<3>(relaxed_plus); + dsmr = bsk::get<2>(relaxed_minus); + dsmi = bsk::get<3>(relaxed_minus); + } + auto event_action = bsk::cast(bsk::ld((action + event))); + auto event_kind = bsk::cast(bsk::ld((kind + event))); + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t3_ = _shift(spr, spi, smr, smi, state, state_mask, state_count); + spr = bsk::get<0>(t3_); + spi = bsk::get<1>(t3_); + smr = bsk::get<2>(t3_); + smi = bsk::get<3>(t3_); + if (bsk::truth(following)) { + auto t4_ = _shift(dspr, dspi, dsmr, dsmi, state, state_mask, state_count); + dspr = bsk::get<0>(t4_); + dspi = bsk::get<1>(t4_); + dsmr = bsk::get<2>(t4_); + dsmi = bsk::get<3>(t4_); + } + } + // The trailing shift or spoil. + if (bsk::truth((bsk::band(event_action, 8) != 0))) { + pbr = zero; + pbi = zero; + mbr = zero; + mbi = zero; + if (bsk::truth(following)) { + dpbr = zero; + dpbi = zero; + dmbr = zero; + dmbi = zero; + } + } else if (bsk::truth((bsk::band(event_action, 16) != 0))) { + auto t5_ = _shift_adjoint(pbr, pbi, mbr, mbi, state, state_mask, state_count); + pbr = bsk::get<0>(t5_); + pbi = bsk::get<1>(t5_); + mbr = bsk::get<2>(t5_); + mbi = bsk::get<3>(t5_); + if (bsk::truth(following)) { + auto t6_ = _shift_adjoint(dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count); + dpbr = bsk::get<0>(t6_); + dpbi = bsk::get<1>(t6_); + dmbr = bsk::get<2>(t6_); + dmbi = bsk::get<3>(t6_); + } + } + if (bsk::truth((bsk::band(event_action, 2) != 0))) { + auto t7_ = _shift_adjoint(pbr, pbi, mbr, mbi, state, state_mask, state_count); + pbr = bsk::get<0>(t7_); + pbi = bsk::get<1>(t7_); + mbr = bsk::get<2>(t7_); + mbi = bsk::get<3>(t7_); + if (bsk::truth(following)) { + auto t8_ = _shift_adjoint(dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count); + dpbr = bsk::get<0>(t8_); + dpbi = bsk::get<1>(t8_); + dmbr = bsk::get<2>(t8_); + dmbi = bsk::get<3>(t8_); + } + } + // The readout. It carries no pulse, so what it recorded is the state + // the pre-shift left. + auto out_ = bsk::ld((output_index + event)); + if (bsk::truth(bsk::band(bsk::band((event_kind == 2), (bsk::band(event_action, 32) != 0)), (out_ >= 0)))) { + auto index_ = ((problem * output_count) + out_); + seed = bsk::make_tup(bsk::ld((grad_real + index_)), bsk::ld((grad_imag + index_)), 0.0f, 0.0f); + auto read_phase = _read((phase + event_base), (dphase + event_base), event, true, 0.0f, following); + auto demodulation = _polar(bsk::make_tup((-bsk::get<0>(read_phase)), (-bsk::get<1>(read_phase))), following); + auto recorded = bsk::make_tup(_total(spr, origin), _total(spi, origin), _total(dspr, origin), _total(dspi, origin)); + auto density_grad = _re_dot(seed, _cmul(recorded, demodulation, following), following); + auto turned = bsk::make_tup(bsk::get<1>(demodulation), (-bsk::get<0>(demodulation)), bsk::get<3>(demodulation), (-bsk::get<2>(demodulation))); + auto phase_grad = _re_dot(seed, _cscale(density_of, _cmul(recorded, turned, following), following), following); + grad_m0 = (grad_m0 + bsk::get<0>(density_grad)); + bsk::atomic_add(((grad_phase + event_base) + event), bsk::get<0>(phase_grad)); + auto weighted = _cmul(_conj(_cscale(density_of, demodulation, following)), seed, following); + auto put = bsk::band(origin, exchanging); + pbr = (pbr + bsk::where(put, bsk::get<0>(weighted), 0.0f)); + pbi = (pbi + bsk::where(put, bsk::get<1>(weighted), 0.0f)); + if (bsk::truth(following)) { + bsk::atomic_add(((dgrad_tissue + (m0_row * atom_count)) + atom), bsk::get<1>(density_grad)); + bsk::atomic_add(((dgrad_phase + event_base) + event), bsk::get<1>(phase_grad)); + dpbr = (dpbr + bsk::where(put, bsk::get<2>(weighted), 0.0f)); + dpbi = (dpbi + bsk::where(put, bsk::get<3>(weighted), 0.0f)); + } + } + // The pulse. + if (bsk::truth((event_kind == 1))) { + if (bsk::truth((bsk::band(event_action, 4) != 0))) { + auto bar = bsk::make_tup(zbr, zbi, dzbr, dzbi); + taken = _rtotal(_re_dot(bar, relaxed_z, following), live_exchanging, following); + grad_efficiency = (grad_efficiency - bsk::get<0>(taken)); + if (bsk::truth(following)) { + bsk::atomic_add(((dgrad_tissue + (efficiency_row * atom_count)) + atom), (-bsk::get<1>(taken))); + dzbr = bsk::where(exchanging, (-((bsk::get<1>(inversion) * zbr) + (bsk::get<0>(inversion) * dzbr))), dzbr); + dzbi = bsk::where(exchanging, (-((bsk::get<1>(inversion) * zbi) + (bsk::get<0>(inversion) * dzbi))), dzbi); + } + zbr = bsk::where(exchanging, ((-bsk::get<0>(inversion)) * zbr), zbr); + zbi = bsk::where(exchanging, ((-bsk::get<0>(inversion)) * zbi), zbi); + } else { + pulse_b1 = voxel_b1; + pulse_b1_phase = voxel_b1_phase; + transmit_row = 0; + if (bsk::truth(shimmed)) { + auto shim = bsk::cast(bsk::ld((shim_index + event))); + auto changed = (shim != held); + auto b1_at = ((bsk::cast((b1_row + held)) * atom_count) + atom); + auto b1_phase_at = ((bsk::cast((b1_phase_row + held)) * atom_count) + atom); + bsk::atomic_add((grad_tissue + b1_at), grad_b1, changed); + bsk::atomic_add((grad_tissue + b1_phase_at), grad_b1_phase, changed); + grad_b1 = bsk::where(changed, 0.0f, grad_b1); + grad_b1_phase = bsk::where(changed, 0.0f, grad_b1_phase); + held = shim; + transmit_row = shim; + auto transmit_at = ((bsk::cast(shim) * atom_count) + atom); + pulse_b1 = _read(b1, db1, transmit_at, transmit, 1.0f, following); + pulse_b1_phase = _read(b1_phase, db1_phase, transmit_at, true, 0.0f, following); + } + auto nominal = _read((flip + event_base), (dflip + event_base), event, true, 0.0f, following); + auto played = _read((phase + event_base), (dphase + event_base), event, true, 0.0f, following); + auto alpha = _rmul(nominal, pulse_b1, following); + auto phi = bsk::make_tup((bsk::get<0>(played) + bsk::get<0>(pulse_b1_phase)), (bsk::get<1>(played) + bsk::get<1>(pulse_b1_phase))); + auto turn = _polar(bsk::make_tup((-bsk::get<0>(phi)), (-bsk::get<1>(phi))), following); + if (bsk::truth(dynamic)) { + auto t9_ = _dynamic_pair_dual_at(pairs, dpairs, pair_index, event_base, event, atom, atom_count, live, bsk::get<0>(phi), bsk::get<1>(phi), directed_pairs); + a = bsk::get<0>(t9_); + spun = bsk::get<1>(t9_); + slope_a = a; + slope_b = spun; + } else if (bsk::truth(profiled)) { + auto t10_ = _profiled_pair_dual(profile, _table_row(profile_index, event, location, locations), bsk::get<0>(alpha), bsk::get<1>(alpha), bsk::get<0>(phi), bsk::get<1>(phi), profile_bins, profile_step); + a = bsk::get<0>(t10_); + spun = bsk::get<1>(t10_); + slope_a = bsk::get<2>(t10_); + slope_b = bsk::get<3>(t10_); + } else { + auto t11_ = _hard_pair(alpha, phi, following); + a = bsk::get<0>(t11_); + spun = bsk::get<1>(t11_); + slope_a = bsk::get<2>(t11_); + slope_b = bsk::get<3>(t11_); + auto _turn = bsk::get<4>(t11_); + } + if (bsk::truth(following)) { + auto t12_ = _spinor_adjoint_dual(a, spun, bsk::make_tup(spr, spi, dspr, dspi), bsk::make_tup(smr, smi, dsmr, dsmi), relaxed_z, bsk::make_tup(pbr, pbi, dpbr, dpbi), bsk::make_tup(mbr, mbi, dmbr, dmbi), bsk::make_tup(zbr, zbi, dzbr, dzbi)); + auto pair_a = bsk::get<0>(t12_); + auto pair_b = bsk::get<1>(t12_); + auto back_p = bsk::get<2>(t12_); + auto back_m = bsk::get<3>(t12_); + back_z = bsk::get<4>(t12_); + grad_a = bsk::make_tup(_total(bsk::get<0>(pair_a), live_exchanging), _total(bsk::get<1>(pair_a), live_exchanging), _total(bsk::get<2>(pair_a), live_exchanging), _total(bsk::get<3>(pair_a), live_exchanging)); + grad_b = bsk::make_tup(_total(bsk::get<0>(pair_b), live_exchanging), _total(bsk::get<1>(pair_b), live_exchanging), _total(bsk::get<2>(pair_b), live_exchanging), _total(bsk::get<3>(pair_b), live_exchanging)); + dpbr = bsk::where(exchanging, bsk::get<2>(back_p), dpbr); + dpbi = bsk::where(exchanging, bsk::get<3>(back_p), dpbi); + dmbr = bsk::where(exchanging, bsk::get<2>(back_m), dmbr); + dmbi = bsk::where(exchanging, bsk::get<3>(back_m), dmbi); + dzbr = bsk::where(exchanging, bsk::get<2>(back_z), dzbr); + dzbi = bsk::where(exchanging, bsk::get<3>(back_z), dzbi); + pbr = bsk::where(exchanging, bsk::get<0>(back_p), pbr); + pbi = bsk::where(exchanging, bsk::get<1>(back_p), pbi); + mbr = bsk::where(exchanging, bsk::get<0>(back_m), mbr); + mbi = bsk::where(exchanging, bsk::get<1>(back_m), mbi); + zbr = bsk::where(exchanging, bsk::get<0>(back_z), zbr); + zbi = bsk::where(exchanging, bsk::get<1>(back_z), zbi); + } else { + auto back = _spinor_adjoint(bsk::get<0>(a), bsk::get<1>(a), bsk::get<0>(spun), bsk::get<1>(spun), spr, spi, smr, smi, bsk::get<0>(relaxed_z), bsk::get<1>(relaxed_z), pbr, pbi, mbr, mbi, zbr, zbi); + grad_a = bsk::make_tup(_total(bsk::get<0>(back), live_exchanging), _total(bsk::get<1>(back), live_exchanging), 0.0f, 0.0f); + grad_b = bsk::make_tup(_total(bsk::get<2>(back), live_exchanging), _total(bsk::get<3>(back), live_exchanging), 0.0f, 0.0f); + pbr = bsk::where(exchanging, bsk::get<4>(back), pbr); + pbi = bsk::where(exchanging, bsk::get<5>(back), pbi); + mbr = bsk::where(exchanging, bsk::get<6>(back), mbr); + mbi = bsk::where(exchanging, bsk::get<7>(back), mbi); + zbr = bsk::where(exchanging, bsk::get<8>(back), zbr); + zbi = bsk::where(exchanging, bsk::get<9>(back), zbi); + } + // The RF phase turns the axis once the pair is out, so it + // reaches ``b`` alone -- under every mode. + auto grad_phi = _re_dot(grad_b, bsk::make_tup(bsk::get<1>(spun), (-bsk::get<0>(spun)), bsk::get<3>(spun), (-bsk::get<2>(spun))), following); + grad_alpha = bsk::make_tup(0.0f, 0.0f); + if (bsk::truth(dynamic)) { + // The flip is inside the pair rather than read against it, + // so the cotangent goes out on the pair; ``b`` was turned + // by the phase after the pair came out, so it turns back. + auto unturned = _cmul(grad_b, _conj(turn), following); + auto entry = (((bsk::cast(bsk::ld(((pair_index + event_base) + event))) * atom_count) + atom) * 4); + bsk::atomic_add(((grad_pairs + entry) + 0), bsk::get<0>(grad_a)); + bsk::atomic_add(((grad_pairs + entry) + 1), bsk::get<1>(grad_a)); + bsk::atomic_add(((grad_pairs + entry) + 2), bsk::get<0>(unturned)); + bsk::atomic_add(((grad_pairs + entry) + 3), bsk::get<1>(unturned)); + if (bsk::truth(following)) { + bsk::atomic_add(((dgrad_pairs + entry) + 0), bsk::get<2>(grad_a)); + bsk::atomic_add(((dgrad_pairs + entry) + 1), bsk::get<3>(grad_a)); + bsk::atomic_add(((dgrad_pairs + entry) + 2), bsk::get<2>(unturned)); + bsk::atomic_add(((dgrad_pairs + entry) + 3), bsk::get<3>(unturned)); + } + } else { + auto along_a = _re_dot(grad_a, slope_a, following); + auto along_b = _re_dot(grad_b, slope_b, following); + grad_alpha = bsk::make_tup((bsk::get<0>(along_a) + bsk::get<0>(along_b)), (bsk::get<1>(along_a) + bsk::get<1>(along_b))); + } + if (bsk::truth((n > m))) { + // The pulse scales every order of the semisolid pool by one + // real number, so its cotangent is one sum over the states. + auto t13_ = _absorption(lineshape, rf_frequency, saturation, event, alpha, voxel_b0, lineshape_bins, lineshape_step, following); + auto absorbed = bsk::get<0>(t13_); + auto shape = bsk::get<1>(t13_); + auto slope = bsk::get<2>(t13_); + auto deposited = bsk::get<3>(t13_); + taken = _rtotal(_re_dot(bsk::make_tup(zbr, zbi, dzbr, dzbi), relaxed_z, following), bsk::band(semisolid, state_mask), following); + if (bsk::truth(following)) { + dzbr = bsk::where(semisolid, ((bsk::get<1>(absorbed) * zbr) + (bsk::get<0>(absorbed) * dzbr)), dzbr); + dzbi = bsk::where(semisolid, ((bsk::get<1>(absorbed) * zbi) + (bsk::get<0>(absorbed) * dzbi)), dzbi); + } + zbr = bsk::where(semisolid, (bsk::get<0>(absorbed) * zbr), zbr); + zbi = bsk::where(semisolid, (bsk::get<0>(absorbed) * zbi), zbi); + auto exponent = _rmul(taken, absorbed, following); + auto swing = _rmul(bsk::make_tup((2.0f * deposited), 0.0f), _rmul(_rmul(exponent, alpha, following), shape, following), following); + grad_alpha = bsk::make_tup((bsk::get<0>(grad_alpha) + bsk::get<0>(swing)), (bsk::get<1>(grad_alpha) + bsk::get<1>(swing))); + auto shifted = _rmul(bsk::make_tup(deposited, 0.0f), _rmul(_rmul(_rmul(exponent, alpha, following), alpha, following), slope, following), following); + grad_b0 = (grad_b0 - bsk::get<0>(shifted)); + if (bsk::truth(following)) { + bsk::atomic_add(((dgrad_tissue + (b0_row * atom_count)) + atom), (-bsk::get<1>(shifted))); + } + } + auto flip_grad = _rmul(grad_alpha, pulse_b1, following); + auto b1_grad = _rmul(grad_alpha, nominal, following); + bsk::atomic_add(((grad_flip + event_base) + event), bsk::get<0>(flip_grad)); + bsk::atomic_add(((grad_phase + event_base) + event), bsk::get<0>(grad_phi)); + grad_b1 = (grad_b1 + bsk::get<0>(b1_grad)); + grad_b1_phase = (grad_b1_phase + bsk::get<0>(grad_phi)); + if (bsk::truth(following)) { + bsk::atomic_add(((dgrad_flip + event_base) + event), bsk::get<1>(flip_grad)); + bsk::atomic_add(((dgrad_phase + event_base) + event), bsk::get<1>(grad_phi)); + bsk::atomic_add(((dgrad_tissue + (bsk::cast((b1_row + transmit_row)) * atom_count)) + atom), bsk::get<1>(b1_grad)); + bsk::atomic_add(((dgrad_tissue + (bsk::cast((b1_phase_row + transmit_row)) * atom_count)) + atom), bsk::get<1>(grad_phi)); + } + } + } + if (bsk::truth((bsk::band(event_action, 1) != 0))) { + auto t14_ = _shift_adjoint(pbr, pbi, mbr, mbi, state, state_mask, state_count); + pbr = bsk::get<0>(t14_); + pbi = bsk::get<1>(t14_); + mbr = bsk::get<2>(t14_); + mbi = bsk::get<3>(t14_); + if (bsk::truth(following)) { + auto t15_ = _shift_adjoint(dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count); + dpbr = bsk::get<0>(t15_); + dpbi = bsk::get<1>(t15_); + dmbr = bsk::get<2>(t15_); + dmbi = bsk::get<3>(t15_); + } + } + // The interval. Order zero also carries the recovery, which is the + // equilibrium less what washout leaves of the operator applied to it. + plus_bar = bsk::make_tup(pbr, pbi, dpbr, dpbi); + minus_bar = bsk::make_tup(mbr, mbi, dmbr, dmbi); + z_bar = bsk::make_tup(zbr, zbi, dzbr, dzbi); + if (bsk::truth((!bsk::truth(following)))) { + plus_bar = bsk::make_tup(pbr, pbi, 0.0f, 0.0f); + minus_bar = bsk::make_tup(mbr, mbi, 0.0f, 0.0f); + z_bar = bsk::make_tup(zbr, zbi, 0.0f, 0.0f); + } + seed = bsk::make_tup(bsk::sum_x(bsk::where(origin, zbr, 0.0f)), bsk::sum_x(bsk::where(origin, dzbr, 0.0f))); + grad_eq = (grad_eq + bsk::get<0>(seed)); + auto restored_grad = bsk::make_tup((-(bsk::get<0>(wout) * bsk::get<0>(seed))), (-((bsk::get<1>(wout) * bsk::get<0>(seed)) + (bsk::get<0>(wout) * bsk::get<1>(seed))))); + grad_wout = bsk::make_tup((-_total((bsk::get<0>(seed) * bsk::get<0>(restored)), held_rows)), (-_total(((bsk::get<1>(seed) * bsk::get<0>(restored)) + (bsk::get<0>(seed) * bsk::get<1>(restored))), held_rows))); + if (bsk::truth(following)) { + curve_eq = (curve_eq + bsk::get<1>(seed)); + } + auto out_plus = _cmul(_conj(plus_bar), _cmul(carried, mixed_plus, following), following); + auto out_minus = _cmul(_conj(minus_bar), _cmul(_conj(carried), mixed_minus, following), following); + auto out_z = _cmul(_conj(z_bar), _cmul(spin, mixed_z, following), following); + // The damping is homogeneous of degree one in every state it acts on, + // so its gradient times the damping itself is the cotangent taken + // against the states the interval leaves; the turns are the same + // derivatives with an imaginary weight. + auto transverse_scaled = bsk::sum_y((bsk::get<0>(out_plus) + bsk::get<0>(out_minus))); + auto transverse_angle = bsk::sum_y((bsk::get<1>(out_minus) - bsk::get<1>(out_plus))); + auto longitudinal_scaled = bsk::sum_y(bsk::get<0>(out_z)); + auto longitudinal_angle = bsk::sum_y((-bsk::get<1>(out_z))); + auto wout_plus = _re_dot(plus_bar, _cmul(unit_t, mixed_plus, following), following); + auto wout_minus = _re_dot(minus_bar, _cmul(_conj(unit_t), mixed_minus, following), following); + auto wout_z = _re_dot(z_bar, _cmul(unit_z, mixed_z, following), following); + grad_wout = bsk::make_tup((bsk::get<0>(grad_wout) + _total(((bsk::get<0>(wout_plus) + bsk::get<0>(wout_minus)) + bsk::get<0>(wout_z)), state_mask)), bsk::get<1>(grad_wout)); + auto half = (order + 0.5f); + grad_angle = bsk::make_tup(_total(transverse_angle, state_mask), 0.0f); + grad_b_factor = bsk::make_tup((-_total(((weight * transverse_scaled) + (squared * longitudinal_scaled)), state_mask)), 0.0f); + grad_turn = bsk::make_tup((-_total(((half * transverse_angle) + (order * longitudinal_angle)), state_mask)), 0.0f); + if (bsk::truth(following)) { + grad_wout = bsk::make_tup(bsk::get<0>(grad_wout), (bsk::get<1>(grad_wout) + _total(((bsk::get<1>(wout_plus) + bsk::get<1>(wout_minus)) + bsk::get<1>(wout_z)), state_mask))); + auto transverse_scaled_t = bsk::sum_y((bsk::get<2>(out_plus) + bsk::get<2>(out_minus))); + auto transverse_angle_t = bsk::sum_y((bsk::get<3>(out_minus) - bsk::get<3>(out_plus))); + auto longitudinal_scaled_t = bsk::sum_y(bsk::get<2>(out_z)); + auto longitudinal_angle_t = bsk::sum_y((-bsk::get<3>(out_z))); + grad_angle = bsk::make_tup(bsk::get<0>(grad_angle), _total(transverse_angle_t, state_mask)); + grad_b_factor = bsk::make_tup(bsk::get<0>(grad_b_factor), (-_total(((weight * transverse_scaled_t) + (squared * longitudinal_scaled_t)), state_mask))); + grad_turn = bsk::make_tup(bsk::get<0>(grad_turn), (-_total(((half * transverse_angle_t) + (order * longitudinal_angle_t)), state_mask))); + } + // The operators' own entries. ``F-`` follows the conjugate of the + // transverse operator, so its cotangent lands on the entry itself. + auto released = _conj(carried); + auto transverse_grad = _cadd(_couter(_cmul(plus_bar, released, following), _conj(plus_in), following), _couter(_cmul(_conj(minus_bar), released, following), minus_in, following)); + auto weighed = _cmul(_conj(z_bar), spin, following); + longitudinal_grad = bsk::make_tup((_outer(bsk::get<0>(weighed), bsk::get<0>(z_in)) - _outer(bsk::get<1>(weighed), bsk::get<1>(z_in))), 0.0f); + if (bsk::truth(following)) { + longitudinal_grad = bsk::make_tup(bsk::get<0>(longitudinal_grad), (((_outer(bsk::get<2>(weighed), bsk::get<0>(z_in)) - _outer(bsk::get<3>(weighed), bsk::get<1>(z_in))) + _outer(bsk::get<0>(weighed), bsk::get<2>(z_in))) - _outer(bsk::get<1>(weighed), bsk::get<3>(z_in)))); + } + // The cotangents back through the interval. + auto back_plus = _cmul(released, _apply(transverse_op, plus_bar, true, true, following), following); + auto back_minus = _cmul(carried, _apply(transverse_op, minus_bar, false, true, following), following); + back_z = _apply_real(longitudinal_op, _cmul(_conj(spin), z_bar, following), true, following); + pbr = bsk::get<0>(back_plus); + pbi = bsk::get<1>(back_plus); + mbr = bsk::get<0>(back_minus); + mbi = bsk::get<1>(back_minus); + zbr = bsk::get<0>(back_z); + zbi = bsk::get<1>(back_z); + if (bsk::truth(following)) { + dpbr = bsk::get<2>(back_plus); + dpbi = bsk::get<3>(back_plus); + dmbr = bsk::get<2>(back_minus); + dmbi = bsk::get<3>(back_minus); + dzbr = bsk::get<2>(back_z); + dzbi = bsk::get<3>(back_z); + } + // The row this event read is shared with every event of its length; + // its cotangent is summed into it, and reaches the event's own length + // through the row's slope. + auto longitudinal_at = ((row_offset + (pool * n)) + column); + auto restored_at = ((row_offset + (n * n)) + pool); + auto across = (((row_offset + (n * n)) + n) + (2 * ((pool * m) + column))); + bsk::atomic_add((slot_grad + longitudinal_at), bsk::get<0>(longitudinal_grad), square); + bsk::atomic_add((slot_grad + restored_at), bsk::get<0>(restored_grad), held_rows); + bsk::atomic_add((slot_grad + across), bsk::get<0>(transverse_grad), across_mask); + bsk::atomic_add(((slot_grad + across) + 1), bsk::get<1>(transverse_grad), across_mask); + if (bsk::truth(following)) { + bsk::atomic_add((slot_curve + longitudinal_at), bsk::get<1>(longitudinal_grad), square); + bsk::atomic_add((slot_curve + restored_at), bsk::get<1>(restored_grad), held_rows); + bsk::atomic_add((slot_curve + across), bsk::get<2>(transverse_grad), across_mask); + bsk::atomic_add(((slot_curve + across) + 1), bsk::get<3>(transverse_grad), across_mask); + } + table_duration = bsk::make_tup(0.0f, 0.0f); + if (bsk::truth((blocks > 1))) { + auto further = (rows * row_width); + auto slope_z = _entries(slot, directions, (further + longitudinal_at), square, bsk::get<1>(dt), further, directed_table, true, following); + auto slope_restored = _entries(slot, directions, (further + restored_at), held_rows, bsk::get<1>(dt), further, directed_table, true, following); + auto slope_real = _entries(slot, directions, (further + across), across_mask, bsk::get<1>(dt), further, directed_table, true, following); + auto slope_imag = _entries(slot, directions, ((further + across) + 1), across_mask, bsk::get<1>(dt), further, directed_table, true, following); + table_duration = bsk::make_tup(((_total((bsk::get<0>(longitudinal_grad) * bsk::get<0>(slope_z)), square) + _total((bsk::get<0>(restored_grad) * bsk::get<0>(slope_restored)), held_rows)) + _total(((bsk::get<0>(transverse_grad) * bsk::get<0>(slope_real)) + (bsk::get<1>(transverse_grad) * bsk::get<0>(slope_imag))), across_mask)), 0.0f); + if (bsk::truth(following)) { + table_duration = bsk::make_tup(bsk::get<0>(table_duration), ((_total(((bsk::get<1>(longitudinal_grad) * bsk::get<0>(slope_z)) + (bsk::get<0>(longitudinal_grad) * bsk::get<1>(slope_z))), square) + _total(((bsk::get<1>(restored_grad) * bsk::get<0>(slope_restored)) + (bsk::get<0>(restored_grad) * bsk::get<1>(slope_restored))), held_rows)) + _total(((((bsk::get<2>(transverse_grad) * bsk::get<0>(slope_real)) + (bsk::get<0>(transverse_grad) * bsk::get<1>(slope_real))) + (bsk::get<3>(transverse_grad) * bsk::get<0>(slope_imag))) + (bsk::get<1>(transverse_grad) * bsk::get<1>(slope_imag))), across_mask))); + // At the row's own length the slope reaches the output only + // through a direction in that length, so only the tangent + // plane takes this. + bsk::atomic_add(((slot_curve + further) + longitudinal_at), (bsk::get<0>(longitudinal_grad) * bsk::get<1>(dt)), square); + bsk::atomic_add(((slot_curve + further) + restored_at), (bsk::get<0>(restored_grad) * bsk::get<1>(dt)), held_rows); + bsk::atomic_add(((slot_curve + further) + across), (bsk::get<0>(transverse_grad) * bsk::get<1>(dt)), across_mask); + bsk::atomic_add((((slot_curve + further) + across) + 1), (bsk::get<1>(transverse_grad) * bsk::get<1>(dt)), across_mask); + } + } + // Washout scales every factor the interval applies and the recovery it + // leaves; past the clamp nothing depends on the rate. + fraction_grad = bsk::make_tup(0.0f, 0.0f); + if (bsk::truth(moving)) { + auto inside = ((bsk::get<0>(washout_rate) * bsk::get<0>(dt)) < 1.0f); + fraction_grad = bsk::make_tup(bsk::where(inside, (-bsk::get<0>(grad_wout)), 0.0f), bsk::where(inside, (-bsk::get<1>(grad_wout)), 0.0f)); + } + auto angle_rate = _rmul(bsk::make_tup(-6.283185307179586f, 0.0f), dt, following); + auto b0_grad = _rmul(grad_angle, angle_rate, following); + auto damping_grad = _rmul(grad_b_factor, dt, following); + auto flow_grad = _rmul(grad_turn, dt, following); + auto washout_grad = _rmul(fraction_grad, dt, following); + auto duration_grad = _rmul(grad_angle, _rmul(bsk::make_tup(-6.283185307179586f, 0.0f), voxel_b0, following), following); + auto through_damping = _rmul(grad_b_factor, damping_rate, following); + auto through_flow = _rmul(grad_turn, flow_rate, following); + auto through_washout = _rmul(fraction_grad, washout_rate, following); + grad_b0 = (grad_b0 + bsk::get<0>(b0_grad)); + grad_damping = (grad_damping + bsk::get<0>(damping_grad)); + grad_flow = (grad_flow + bsk::get<0>(flow_grad)); + grad_washout = (grad_washout + bsk::get<0>(washout_grad)); + bsk::atomic_add(((grad_duration + event_base) + event), ((((bsk::get<0>(duration_grad) + bsk::get<0>(through_damping)) + bsk::get<0>(through_flow)) + bsk::get<0>(through_washout)) + bsk::get<0>(table_duration))); + if (bsk::truth(following)) { + curve_b0 = (curve_b0 + bsk::get<1>(b0_grad)); + curve_damping = (curve_damping + bsk::get<1>(damping_grad)); + curve_flow = (curve_flow + bsk::get<1>(flow_grad)); + curve_washout = (curve_washout + bsk::get<1>(washout_grad)); + bsk::atomic_add(((dgrad_duration + event_base) + event), ((((bsk::get<1>(duration_grad) + bsk::get<1>(through_damping)) + bsk::get<1>(through_flow)) + bsk::get<1>(through_washout)) + bsk::get<1>(table_duration))); + } + } + // The equilibrium is also where every pool starts, which the walk back + // reaches last. + grad_eq = (grad_eq + bsk::sum_x(bsk::where(origin, zbr, 0.0f))); + bsk::atomic_add((slot_grad + pool), grad_eq, held_rows); + bsk::atomic_add(((grad_tissue + (m0_row * atom_count)) + atom), grad_m0); + bsk::atomic_add(((grad_tissue + (bsk::cast((b1_row + held)) * atom_count)) + atom), grad_b1); + bsk::atomic_add(((grad_tissue + (bsk::cast((b1_phase_row + held)) * atom_count)) + atom), grad_b1_phase); + bsk::atomic_add(((grad_tissue + (b0_row * atom_count)) + atom), grad_b0); + bsk::atomic_add(((grad_tissue + (efficiency_row * atom_count)) + atom), grad_efficiency); + bsk::atomic_add(((grad_tissue + (diffusion_row * atom_count)) + atom), grad_damping); + // One buffer drives two rates, so the velocity gradient is the sum of what + // each geometry carries back. + bsk::atomic_add(((grad_tissue + (velocity_row * atom_count)) + atom), ((flow_scale * grad_flow) + ((heading * washout_scale) * grad_washout))); + if (bsk::truth(following)) { + curve_eq = (curve_eq + bsk::sum_x(bsk::where(origin, dzbr, 0.0f))); + bsk::atomic_add((slot_curve + pool), curve_eq, held_rows); + bsk::atomic_add(((dgrad_tissue + (b0_row * atom_count)) + atom), curve_b0); + bsk::atomic_add(((dgrad_tissue + (diffusion_row * atom_count)) + atom), curve_damping); + bsk::atomic_add(((dgrad_tissue + (velocity_row * atom_count)) + atom), ((flow_scale * curve_flow) + ((heading * washout_scale) * curve_washout))); + } +} diff --git a/src/blochsim/_tile.hpp b/src/blochsim/_tile.hpp new file mode 100644 index 00000000..272aec61 --- /dev/null +++ b/src/blochsim/_tile.hpp @@ -0,0 +1,1120 @@ +// The tile model the GPU kernels are written in, for a CUDA block and for the +// host. +// +// A kernel program holds a tile of at most two axes: ``x`` runs along the +// threads of one row of the block and ``y`` across its rows. A third, ``z``, +// stands in for ``x`` where a tile is square in a short axis -- an operator +// over pools -- and is held as an array in each thread. A value carries, in +// its type, the axes it varies along (``AX``: bit 0 for x, bit 1 for y, bit 2 +// for z), so +// a value of one row's problem is not mistaken for a value of every state of +// it: a sum along x of a value that does not vary along x is the value itself, +// and a store or an atomic of it is made once rather than once per thread. +// Values that vary along neither are plain C++ scalars, uniform over the block. +// +// Under ``BLOCHSIM_SIMT`` (the CUDA build) a tile is one element per thread, +// and the operations across a row are warp shuffles or shared memory. Without +// it, a tile is the whole array, so the same kernel source runs on the host one +// program at a time: that is the build the tests run without a card. +#pragma once + +#include +#include +#include +#include +#include + +#if defined(BLOCHSIM_SIMT) +#define BSK_HD __device__ __forceinline__ +#define BSK_CONSTEXPR __host__ __device__ constexpr +#else +#define BSK_CONSTEXPR constexpr +#include +#include +#define BSK_HD inline +#endif + +namespace bsk { + +// --------------------------------------------------------------------------- +// The launch a program belongs to. +// --------------------------------------------------------------------------- + +#if defined(BLOCHSIM_SIMT) + +extern __shared__ unsigned long long shared_words[]; + +// The length of z, set once by a kernel's entry. +__shared__ int z_width; + +BSK_HD int width_x() { return static_cast(blockDim.x); } +BSK_HD int width_y() { return static_cast(blockDim.y); } +BSK_HD int width_z() { return z_width; } + +__device__ __forceinline__ void enter(int nz) { + if (threadIdx.x == 0 && threadIdx.y == 0) { + z_width = nz; + } + __syncthreads(); +} +BSK_HD std::int64_t program_id(int axis) { + return axis == 0 ? static_cast(blockIdx.x) + : static_cast(blockIdx.y); +} + +#else + +struct HostProgram { + int nx = 1; + int ny = 1; + int nz = 1; + std::int64_t pid[2] = {0, 0}; +}; + +inline thread_local HostProgram program; + +inline int width_x() { return program.nx; } +inline int width_y() { return program.ny; } +inline int width_z() { return program.nz; } +inline std::int64_t program_id(int axis) { return program.pid[axis]; } + +#endif + +// --------------------------------------------------------------------------- +// Tiles. +// --------------------------------------------------------------------------- + +// The longest z a tile holds. +constexpr int MAX_Z = 8; + +template +struct V; + +template +struct tile_traits { + static constexpr int axes = 0; + using element = A; +}; + +template +struct tile_traits> { + static constexpr int axes = AX; + using element = T; +}; + +template +using element_t = typename tile_traits>::element; + +template +constexpr int axes_of = tile_traits>::axes; + +template +constexpr bool any_tile = ((axes_of != 0) || ...); + +#if defined(BLOCHSIM_SIMT) + +template +struct V { + static_assert(AX > 0 && AX < 8 && AX != 5 && AX != 7, "a tile varies along x or z, y, or both"); + static constexpr int lanes = (AX & 4) ? MAX_Z : 1; + T v[lanes]; + + V() = default; + + template || std::is_pointer_v, int> = 0> + BSK_HD V(U scalar) { + for (int z = 0; z < lanes; ++z) v[z] = static_cast(scalar); + } + + template = 0> + BSK_HD V(const V& other) { + for (int z = 0; z < lanes; ++z) v[z] = static_cast(other.v[(BX & 4) ? z : 0]); + } +}; + +// The element a thread holds at ``z``. +template +BSK_HD decltype(auto) element(const A& a, int z = 0) { + if constexpr (axes_of == 0) { + return a; + } else if constexpr ((axes_of & 4) != 0) { + return (a.v[z]); + } else { + return (a.v[0]); + } +} + +// Lanes of z past its length are left unset, so nothing reads memory for them. +template +BSK_HD auto zip(F f, const A&... a) { + constexpr int AX = (0 | ... | axes_of); + if constexpr (AX == 0) { + return f(a...); + } else { + using R = decltype(f(element(a)...)); + V out; + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + out.v[z] = f(element(a, z)...); + } + } + } else { + out.v[0] = f(element(a)...); + } + return out; + } +} + +#else + +template +inline std::size_t count() { + return static_cast((AX & 1) ? program.nx : 1) + * static_cast((AX & 2) ? program.ny : 1) + * static_cast((AX & 4) ? program.nz : 1); +} + +// The extent of a loop over the axes in AX. +template +inline int extent_x() { return (AX & 1) ? program.nx : 1; } +template +inline int extent_y() { return (AX & 2) ? program.ny : 1; } +template +inline int extent_z() { return (AX & 4) ? program.nz : 1; } + +template +struct V { + static_assert(AX > 0 && AX < 8 && AX != 5 && AX != 7, "a tile varies along x or z, y, or both"); + // std::vector holds no addressable elements. + using Stored = std::conditional_t, unsigned char, T>; + std::vector v; + + V() : v(count()) {} + + template || std::is_pointer_v, int> = 0> + V(U scalar) : v(count(), static_cast(static_cast(scalar))) {} + + template = 0> + V(const V& other) : v(count()) { + for (int y = 0; y < extent_y(); ++y) { + for (int x = 0; x < extent_x(); ++x) { + for (int z = 0; z < extent_z(); ++z) { + v[index(y, x, z)] = static_cast(static_cast(other.at(y, x, z))); + } + } + } + } + + std::size_t index(int y, int x, int z) const { + const std::size_t row = static_cast(extent_x()) * extent_z(); + return static_cast((AX & 2) ? y : 0) * row + + static_cast((AX & 1) ? x : 0) * extent_z() + + static_cast((AX & 4) ? z : 0); + } + const Stored& at(int y, int x, int z = 0) const { return v[index(y, x, z)]; } + Stored& at(int y, int x, int z = 0) { return v[index(y, x, z)]; } +}; + +template +inline decltype(auto) element(const A& a, int y, int x, int z = 0) { + if constexpr (axes_of == 0) { + return a; + } else if constexpr (std::is_same_v, bool>) { + return static_cast(a.at(y, x, z)); + } else { + return a.at(y, x, z); + } +} + +template +inline auto zip(F f, const A&... a) { + constexpr int AX = (0 | ... | axes_of); + if constexpr (AX == 0) { + return f(a...); + } else { + using R = decltype(f(element(a, 0, 0)...)); + V out; + for (int y = 0; y < extent_y(); ++y) { + for (int x = 0; x < extent_x(); ++x) { + for (int z = 0; z < extent_z(); ++z) { + out.at(y, x, z) = f(element(a, y, x, z)...); + } + } + } + return out; + } +} + +#endif + +// --------------------------------------------------------------------------- +// Element-wise arithmetic. Between two scalars the language's own operators +// apply; these take over as soon as one side is a tile. +// --------------------------------------------------------------------------- + +#define BSK_BINARY(op) \ + template , int> = 0> \ + BSK_HD auto operator op(const A& a, const B& b) { \ + return zip([](auto x, auto y) { return x op y; }, a, b); \ + } + +BSK_BINARY(+) +BSK_BINARY(-) +BSK_BINARY(*) +BSK_BINARY(<) +BSK_BINARY(<=) +BSK_BINARY(>) +BSK_BINARY(>=) +BSK_BINARY(==) +BSK_BINARY(!=) +BSK_BINARY(<<) +BSK_BINARY(>>) +#undef BSK_BINARY + +template +BSK_HD auto operator-(const V& a) { + return zip([](auto x) { return -x; }, a); +} + +// Python's ``/`` divides integers into a float; ``//`` and ``%`` keep them. +template +BSK_HD auto truediv(const A& a, const B& b) { + return zip( + [](auto x, auto y) { + if constexpr (std::is_integral_v && std::is_integral_v) { + return static_cast(x) / static_cast(y); + } else { + return x / y; + } + }, + a, b); +} + +template +BSK_HD auto floordiv(const A& a, const B& b) { + return zip( + [](auto x, auto y) { + if constexpr (std::is_integral_v && std::is_integral_v) { + return x / y; + } else { + return floor(x / y); + } + }, + a, b); +} + +template +BSK_HD auto mod(const A& a, const B& b) { + return zip( + [](auto x, auto y) { + if constexpr (std::is_integral_v && std::is_integral_v) { + return x % y; + } else { + return fmod(x, y); + } + }, + a, b); +} + +// ``&``, ``|``, ``^`` and ``~`` are logical on masks and bitwise on integers, +// and a mask stays a mask. +template +BSK_HD auto bit_and(X x, Y y) { + if constexpr (std::is_same_v && std::is_same_v) { + return static_cast(x && y); + } else { + return x & y; + } +} +template +BSK_HD auto bit_or(X x, Y y) { + if constexpr (std::is_same_v && std::is_same_v) { + return static_cast(x || y); + } else { + return x | y; + } +} +template +BSK_HD auto bit_xor(X x, Y y) { + if constexpr (std::is_same_v && std::is_same_v) { + return static_cast(x != y); + } else { + return x ^ y; + } +} +template +BSK_HD auto bit_not(X x) { + if constexpr (std::is_same_v) { + return !x; + } else { + return ~x; + } +} + +template +BSK_HD auto band(const A& a, const B& b) { + return zip([](auto x, auto y) { return bit_and(x, y); }, a, b); +} +template +BSK_HD auto bor(const A& a, const B& b) { + return zip([](auto x, auto y) { return bit_or(x, y); }, a, b); +} +template +BSK_HD auto bxor(const A& a, const B& b) { + return zip([](auto x, auto y) { return bit_xor(x, y); }, a, b); +} +template +BSK_HD auto bnot(const A& a) { + return zip([](auto x) { return bit_not(x); }, a); +} + +template +BSK_HD auto cast(const A& a) { + return zip([](auto x) { return static_cast(x); }, a); +} + +template +BSK_HD auto where(const C& c, const A& a, const B& b) { + return zip( + [](auto test, auto yes, auto no) { + using R = decltype(yes + no); + return test ? static_cast(yes) : static_cast(no); + }, + c, a, b); +} + +// A branch condition. Every thread of a block takes the same branch, so a +// tile reaching one holds the same value everywhere it is read. +template +BSK_HD bool truth(const A& a) { + if constexpr (axes_of == 0) { + return static_cast(a); + } else { + return static_cast(a.v[0]); + } +} + +template +BSK_HD auto select(bool test, const A& a, const B& b) { + return where(test, a, b); +} + +// --------------------------------------------------------------------------- +// Element-wise functions, in the precision of their argument. +// --------------------------------------------------------------------------- + +BSK_HD float s_exp(float x) { return expf(x); } +BSK_HD double s_exp(double x) { return ::exp(x); } +BSK_HD float s_cos(float x) { return cosf(x); } +BSK_HD double s_cos(double x) { return ::cos(x); } +BSK_HD float s_sin(float x) { return sinf(x); } +BSK_HD double s_sin(double x) { return ::sin(x); } +BSK_HD float s_sqrt(float x) { return sqrtf(x); } +BSK_HD double s_sqrt(double x) { return ::sqrt(x); } +BSK_HD float s_floor(float x) { return floorf(x); } +BSK_HD double s_floor(double x) { return ::floor(x); } +BSK_HD float s_rint(float x) { return rintf(x); } +BSK_HD double s_rint(double x) { return ::rint(x); } +BSK_HD float s_acos(float x) { return acosf(x); } +BSK_HD double s_acos(double x) { return ::acos(x); } +BSK_HD float s_fma(float x, float y, float z) { return fmaf(x, y, z); } +BSK_HD double s_fma(double x, double y, double z) { return ::fma(x, y, z); } +template +BSK_HD X s_abs(X x) { + if constexpr (std::is_floating_point_v) { + return x < X(0) ? -x : (x == X(0) ? X(0) : x); + } else { + return x < X(0) ? -x : x; + } +} +// ``tl.minimum`` and ``tl.maximum`` return the number where one side is NaN. +template +BSK_HD auto s_min(X x, Y y) { + using R = decltype(x + y); + if constexpr (std::is_floating_point_v) { + return static_cast(fmin(static_cast(x), static_cast(y))); + } else { + return static_cast(x) < static_cast(y) ? static_cast(x) : static_cast(y); + } +} +template +BSK_HD auto s_max(X x, Y y) { + using R = decltype(x + y); + if constexpr (std::is_floating_point_v) { + return static_cast(fmax(static_cast(x), static_cast(y))); + } else { + return static_cast(x) > static_cast(y) ? static_cast(x) : static_cast(y); + } +} + +#define BSK_UNARY(name) \ + template \ + BSK_HD auto name(const A& a) { \ + return zip([](auto x) { return s_##name(x); }, a); \ + } +BSK_UNARY(exp) +BSK_UNARY(cos) +BSK_UNARY(sin) +BSK_UNARY(sqrt) +BSK_UNARY(floor) +BSK_UNARY(rint) +BSK_UNARY(acos) +BSK_UNARY(abs) +#undef BSK_UNARY + +template +BSK_HD auto minimum(const A& a, const B& b) { + return zip([](auto x, auto y) { return s_min(x, y); }, a, b); +} +template +BSK_HD auto maximum(const A& a, const B& b) { + return zip([](auto x, auto y) { return s_max(x, y); }, a, b); +} +template +BSK_HD auto fma(const A& a, const B& b, const C& c) { + return zip( + [](auto x, auto y, auto z) { + using R = decltype(x * y + z); + return s_fma(static_cast(x), static_cast(y), static_cast(z)); + }, + a, b, c); +} + +// --------------------------------------------------------------------------- +// Indices. +// --------------------------------------------------------------------------- + +BSK_HD V arange_x() { +#if defined(BLOCHSIM_SIMT) + V out; + out.v[0] = static_cast(threadIdx.x); + return out; +#else + V out; + for (int x = 0; x < program.nx; ++x) out.v[static_cast(x)] = x; + return out; +#endif +} + +BSK_HD V arange_y() { +#if defined(BLOCHSIM_SIMT) + V out; + out.v[0] = static_cast(threadIdx.y); + return out; +#else + V out; + for (int y = 0; y < program.ny; ++y) out.v[static_cast(y)] = y; + return out; +#endif +} + +BSK_HD V arange_z() { + V out; +#if defined(BLOCHSIM_SIMT) + for (int z = 0; z < MAX_Z; ++z) out.v[z] = z; +#else + for (int z = 0; z < program.nz; ++z) out.v[static_cast(z)] = z; +#endif + return out; +} + +template +BSK_HD auto full(T value) { + if constexpr (AX == 0) { + return value; + } else { + return V(value); + } +} + +template +BSK_HD auto zeros_like(const A& a) { + using T = element_t; + if constexpr (axes_of == 0) { + return T(0); + } else { + return V>(T(0)); + } +} + +// --------------------------------------------------------------------------- +// Memory. +// --------------------------------------------------------------------------- + +template +BSK_HD auto ld(const P& pointer) { + return zip([](auto p) { return *p; }, pointer); +} + +template +BSK_HD auto ld(const P& pointer, const M& mask, const O& other) { + return zip( + [](auto p, auto m, auto o) { + using T = std::remove_cv_t>; + return m ? *p : static_cast(o); + }, + pointer, mask, other); +} + +#if defined(BLOCHSIM_SIMT) + +// An address that does not vary along an axis is held by every thread along +// it; only the first of them writes, with the value that thread holds. +template +BSK_HD bool writes() { + return ((AX & 1) || threadIdx.x == 0) && ((AX & 2) || threadIdx.y == 0); +} + +template +BSK_HD void each_written(F f) { + if (!writes()) { + return; + } + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + f(z); + } + } + } else { + f(0); + } +} + +template +BSK_HD void st(const P& pointer, const T& value, const M& mask) { + constexpr int AX = axes_of

| axes_of; + each_written([&](int z) { + if (element(mask, z)) { + auto p = element(pointer, z); + *p = static_cast>(element(value, z)); + } + }); +} + +template +BSK_HD void atomic_add(const P& pointer, const T& value, const M& mask) { + constexpr int AX = axes_of

| axes_of; + each_written([&](int z) { + if (element(mask, z)) { + auto p = element(pointer, z); + atomicAdd(p, static_cast>(element(value, z))); + } + }); +} + +#else + +template +inline void st(const P& pointer, const T& value, const M& mask) { + constexpr int AX = axes_of

| axes_of; + for (int y = 0; y < extent_y(); ++y) { + for (int x = 0; x < extent_x(); ++x) { + for (int z = 0; z < extent_z(); ++z) { + if (element(mask, y, x, z)) { + auto p = element(pointer, y, x, z); + *p = static_cast>(element(value, y, x, z)); + } + } + } + } +} + +template +inline void atomic_add(const P& pointer, const T& value, const M& mask) { + constexpr int AX = axes_of

| axes_of; + for (int y = 0; y < extent_y(); ++y) { + for (int x = 0; x < extent_x(); ++x) { + for (int z = 0; z < extent_z(); ++z) { + if (element(mask, y, x, z)) { + auto p = element(pointer, y, x, z); + *p += static_cast>(element(value, y, x, z)); + } + } + } + } +} + +#endif + +template +BSK_HD void st(const P& pointer, const T& value) { + st(pointer, value, true); +} + +template +BSK_HD void atomic_add(const P& pointer, const T& value) { + atomic_add(pointer, value, true); +} + +// --------------------------------------------------------------------------- +// Across a row or a column of the block. +// --------------------------------------------------------------------------- + +#if defined(BLOCHSIM_SIMT) + +template +BSK_HD T shuffle_xor(T value, int offset, int width) { + if constexpr (std::is_same_v) { + return __shfl_xor_sync(0xffffffffu, static_cast(value), offset, width) != 0; + } else { + return __shfl_xor_sync(0xffffffffu, value, offset, width); + } +} + +template +BSK_HD T shuffle(T value, int lane, int width) { + if constexpr (std::is_same_v) { + return __shfl_sync(0xffffffffu, static_cast(value), lane, width) != 0; + } else { + return __shfl_sync(0xffffffffu, value, lane, width); + } +} + +// Every thread of the block runs these together: a kernel's control flow +// depends on nothing a single thread holds. +template +BSK_HD T reduce_row(T value, Op op) { + const int nx = width_x(); + const int width = nx < 32 ? nx : 32; + for (int offset = width >> 1; offset > 0; offset >>= 1) { + value = op(value, shuffle_xor(value, offset, width)); + } + if (nx <= 32) { + return value; + } + T* words = reinterpret_cast(shared_words); + const int warps = nx >> 5; + __syncthreads(); + if ((threadIdx.x & 31) == 0) { + words[threadIdx.y * warps + (threadIdx.x >> 5)] = value; + } + __syncthreads(); + T total = words[threadIdx.y * warps]; + for (int w = 1; w < warps; ++w) { + total = op(total, words[threadIdx.y * warps + w]); + } + return total; +} + +template +BSK_HD T reduce_column(T value, Op op) { + T* words = reinterpret_cast(shared_words); + __syncthreads(); + words[threadIdx.y * blockDim.x + threadIdx.x] = value; + __syncthreads(); + T total = words[threadIdx.x]; + for (unsigned y = 1; y < blockDim.y; ++y) { + total = op(total, words[y * blockDim.x + threadIdx.x]); + } + return total; +} + +template +BSK_HD T gather_row(T value, int lane) { + const int nx = width_x(); + if (nx <= 32) { + return shuffle(value, lane, nx); + } + T* words = reinterpret_cast(shared_words); + __syncthreads(); + words[threadIdx.y * nx + threadIdx.x] = value; + __syncthreads(); + return words[threadIdx.y * nx + lane]; +} + +#endif + +struct Add { + template + BSK_HD X operator()(X a, X b) const { return a + b; } +}; +struct Max { + template + BSK_HD X operator()(X a, X b) const { return s_max(a, b); } +}; + +// A value with an axis taken out of AX, held where the reduction left it. +template +BSK_HD auto reduced(const T* lanes) { + if constexpr (AX == 0) { + return lanes[0]; + } else { + V out; + for (int z = 0; z < V::lanes; ++z) out.v[z] = lanes[z]; + return out; + } +} + +template +BSK_HD auto reduce_x(const V& a, Op op) { + if constexpr ((AX & 1) == 0) { + return a; + } else { + constexpr int RX = AX & ~1; +#if defined(BLOCHSIM_SIMT) + T total = reduce_row(a.v[0], op); + return reduced(&total); +#else + if constexpr (RX == 0) { + T total = a.at(0, 0); + for (int x = 1; x < program.nx; ++x) total = op(total, a.at(0, x)); + return total; + } else { + V out; + for (int y = 0; y < program.ny; ++y) { + T total = a.at(y, 0); + for (int x = 1; x < program.nx; ++x) total = op(total, a.at(y, x)); + out.at(y, 0) = total; + } + return out; + } +#endif + } +} + +template +BSK_HD auto reduce_y(const V& a, Op op) { + if constexpr ((AX & 2) == 0) { + return a; + } else { + constexpr int RX = AX & ~2; +#if defined(BLOCHSIM_SIMT) + T totals[V::lanes]; + for (int z = 0; z < V::lanes; ++z) totals[z] = reduce_column(a.v[z], op); + return reduced(totals); +#else + if constexpr (RX == 0) { + T total = a.at(0, 0); + for (int y = 1; y < program.ny; ++y) total = op(total, a.at(y, 0)); + return total; + } else { + V out; + for (int x = 0; x < extent_x(); ++x) { + for (int z = 0; z < extent_z(); ++z) { + T total = a.at(0, x, z); + for (int y = 1; y < program.ny; ++y) total = op(total, a.at(y, x, z)); + out.at(0, x, z) = total; + } + } + return out; + } +#endif + } +} + +template +BSK_HD auto reduce_z(const V& a, Op op) { + if constexpr ((AX & 4) == 0) { + return a; + } else { + constexpr int RX = AX & ~4; +#if defined(BLOCHSIM_SIMT) + T total = a.v[0]; + const int nz = width_z(); + for (int z = 1; z < MAX_Z; ++z) { + if (z < nz) { + total = op(total, a.v[z]); + } + } + return reduced(&total); +#else + if constexpr (RX == 0) { + T total = a.at(0, 0, 0); + for (int z = 1; z < program.nz; ++z) total = op(total, a.at(0, 0, z)); + return total; + } else { + V out; + for (int y = 0; y < program.ny; ++y) { + T total = a.at(y, 0, 0); + for (int z = 1; z < program.nz; ++z) total = op(total, a.at(y, 0, z)); + out.at(y, 0) = total; + } + return out; + } +#endif + } +} + +// ``tl.sum(a, axis=1)`` of a two-axis tile: along x, or along z where the +// tile's second axis is z. +template +BSK_HD auto sum_x(const A& a) { + if constexpr (axes_of == 0) { + return a; + } else if constexpr ((axes_of & 4) != 0) { + return reduce_z(a, Add{}); + } else { + return reduce_x(a, Add{}); + } +} +template +BSK_HD auto sum_y(const A& a) { + if constexpr (axes_of == 0) { + return a; + } else { + return reduce_y(a, Add{}); + } +} +template +BSK_HD auto sum_all(const A& a) { + return sum_y(sum_x(a)); +} +template +BSK_HD auto max_all(const A& a) { + static_assert((axes_of & 4) == 0, "a maximum along z is not taken"); + if constexpr (axes_of == 0) { + return a; + } else { + auto row = reduce_x(a, Max{}); + if constexpr (axes_of == 0) { + return row; + } else { + return reduce_y(row, Max{}); + } + } +} + +// ``values`` read at ``index`` along x, row by row: ``tl.gather(values, index, 1)``. +template +BSK_HD auto gather_x(const V& values, const I& index) { + static_assert(AX & 1, "a gather along x reads a value that varies along x"); +#if defined(BLOCHSIM_SIMT) + V> out; + out.v[0] = gather_row(values.v[0], static_cast(element(index))); + return out; +#else + constexpr int RX = AX | axes_of; + V out; + for (int y = 0; y < ((RX & 2) ? program.ny : 1); ++y) { + for (int x = 0; x < program.nx; ++x) { + out.at(y, x) = values.at(y, static_cast(element(index, y, x))); + } + } + return out; +#endif +} + +// --------------------------------------------------------------------------- +// Square operators along y and z against tiles along y and x. +// --------------------------------------------------------------------------- + +// ``operator @ planes`` over y, or ``operator.T @ planes``: the operator's rows +// along y and its columns along z, the planes' rows along y. +template +BSK_HD auto times(const O& op_in, const P& planes_in, bool transposed) { + using T = decltype(element_t() * element_t

()); + const V op(op_in); + const V planes(planes_in); + V out(T(0)); +#if defined(BLOCHSIM_SIMT) + // The planes, then the operator, through shared memory: a thread reads the + // column of planes beneath its state and the operator's entries it needs. + T* words = reinterpret_cast(shared_words); + const int nx = width_x(); + const int ny = width_y(); + const int nz = width_z(); + T* matrix = words + nx * ny; + __syncthreads(); + words[threadIdx.y * nx + threadIdx.x] = planes.v[0]; + if (threadIdx.x == 0) { + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + matrix[threadIdx.y * nz + z] = op.v[z]; + } + } + } + __syncthreads(); + T total = T(0); + for (int k = 0; k < ny; ++k) { + const T entry = transposed ? matrix[k * nz + threadIdx.y] : matrix[threadIdx.y * nz + k]; + total += entry * words[k * nx + threadIdx.x]; + } + out.v[0] = total; +#else + for (int i = 0; i < program.ny; ++i) { + for (int x = 0; x < program.nx; ++x) { + T total = T(0); + for (int k = 0; k < program.ny; ++k) { + const T entry = transposed ? op.at(k, 0, i) : op.at(i, 0, k); + total += entry * planes.at(k, x); + } + out.at(i, x) = total; + } + } +#endif + return out; +} + +// ``sum_x left[i, x] right[j, x]``: rows ``i`` along y, columns ``j`` along z. +template +BSK_HD auto outer(const L& left_in, const R& right_in) { + using T = decltype(element_t() * element_t()); + const V left(left_in); + const V right(right_in); + V out(T(0)); +#if defined(BLOCHSIM_SIMT) + T* words = reinterpret_cast(shared_words); + const int nx = width_x(); + const int ny = width_y(); + T* rows = words + nx * ny; + __syncthreads(); + rows[threadIdx.y * nx + threadIdx.x] = right.v[0]; + __syncthreads(); + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + out.v[z] = reduce_row(left.v[0] * rows[z * nx + threadIdx.x], Add{}); + } + } +#else + for (int i = 0; i < program.ny; ++i) { + for (int j = 0; j < program.nz; ++j) { + T total = T(0); + for (int x = 0; x < program.nx; ++x) { + total += left.at(i, x) * right.at(j, x); + } + out.at(i, 0, j) = total; + } + } +#endif + return out; +} + +// --------------------------------------------------------------------------- +// Tuples, which is how a helper returns more than one value. +// --------------------------------------------------------------------------- + +template +struct tup; + +template <> +struct tup<> {}; + +template +struct tup { + H head; + tup tail; + + tup() = default; + + // From a tuple of narrower elements, element by element; a shorter one + // fills the front. + template + BSK_HD tup(const tup& other) : head(other.head), tail(other.tail) {} + + BSK_HD tup(const tup<>&) : head(), tail() {} +}; + +template +BSK_HD tup make_tup(H head, T... tail) { + tup out; + out.head = head; + if constexpr (sizeof...(T) > 0) { + out.tail = make_tup(tail...); + } + return out; +} + +template +BSK_HD auto& get(tup& t) { + if constexpr (I == 0) { + return t.head; + } else { + return get(t.tail); + } +} + +template +BSK_HD const auto& get(const tup& t) { + if constexpr (I == 0) { + return t.head; + } else { + return get(t.tail); + } +} + +template +struct tile_traits> { + static constexpr int axes = (0 | ... | tile_traits::axes); + using element = void; +}; + +// ``f`` called with each of ``Start``, ``Start + Step``, ... below ``Stop`` as +// a compile-time constant, for a loop whose counter indexes a tuple. +template +BSK_HD void static_for(F&& f) { + if constexpr (Start < Stop) { + f(std::integral_constant{}); + static_for(f); + } +} + +// The axes any of the arguments varies along. +template +constexpr int joint_axes = (0 | ... | axes_of); + +// The axes a helper's local takes: those of the arguments it was called with +// and those it makes itself, keeping x or z -- never both -- as traced. +BSK_CONSTEXPR int fit_axes(int axes, int traced) { + return (axes & 5) == 5 ? (axes & ~5) | ((traced & 4) ? 4 : 1) : axes; +} + +// A value of element type T varying along AX: a scalar where AX is zero. +template +struct tile_type { + using type = V; +}; + +template +struct tile_type { + using type = T; +}; + +template +using tile_t = typename tile_type::type; + +} // namespace bsk + +namespace bsk { + +// A value carried into a wider type of the same shape: a tile into a wider +// element type or more axes, a tuple element by element. +template +BSK_HD R convert(const A& a); + +template +struct converter { + template + BSK_HD static R run(const A& a) { + return R(a); + } +}; + +template +struct converter> { + template + BSK_HD static tup run(const tup& a) { + tup out; + fill<0>(out, a); + return out; + } + + template + BSK_HD static void fill(O& out, const A& a) { + if constexpr (I < sizeof...(T)) { + get(out) = convert(out))>>(get(a)); + fill(out, a); + } + } +}; + +template +BSK_HD R convert(const A& a) { + return converter::run(a); +} + +} // namespace bsk diff --git a/src/blochsim/estimators/_dictionary.py b/src/blochsim/estimators/_dictionary.py index 42a114f7..c368d542 100644 --- a/src/blochsim/estimators/_dictionary.py +++ b/src/blochsim/estimators/_dictionary.py @@ -99,7 +99,7 @@ class DictionaryMatcher(Estimator): The expensive operation is a matrix product. Torch therefore dispatches directly to the installed CPU BLAS or cuBLAS implementation; a separate - C++ or Triton matrix-multiplication kernel would duplicate a faster + C++ or CUDA matrix-multiplication kernel would duplicate a faster vendor implementation. Chunking bounds the temporary score matrix. References diff --git a/src/blochsim/estimators/_perk.py b/src/blochsim/estimators/_perk.py index 135abf3b..cb5dc782 100644 --- a/src/blochsim/estimators/_perk.py +++ b/src/blochsim/estimators/_perk.py @@ -10,6 +10,7 @@ import torch +from .. import _gpu_launch from .._execution import PER_VOXEL_CROSSOVER, one_device, per_voxel from ._mapping import Estimator @@ -726,13 +727,13 @@ def _loaded(name: str) -> Any: return None -_TRITON = _loaded("_perk_triton") +_GPU = _loaded("_perk_gpu") if _gpu_launch.available() else None _NATIVE = _loaded("_perk_native") def _kernels(device: torch.device) -> Any: """The fused backend for this device, or ``None``.""" - return _TRITON if device.type == "cuda" else _NATIVE + return _GPU if device.type == "cuda" else _NATIVE class _FusedRegression(torch.autograd.Function): @@ -751,7 +752,7 @@ def forward( ) -> torch.Tensor: ctx.save_for_backward(signals, frequency, transposed, phase, weight) if signals.device.type == "cuda": - return _TRITON.regress( + return _GPU.regress( signals, frequency, phase, feature_mean, weight, parameter_mean ) return _NATIVE.regress( @@ -770,7 +771,7 @@ def backward(ctx: Any, cotangent: torch.Tensor) -> tuple[torch.Tensor | None, .. if not ctx.needs_input_grad[0]: return (None,) * 7 gradient = ( - _TRITON.regress_vjp(cotangent, signals, frequency, phase, weight) + _GPU.regress_vjp(cotangent, signals, frequency, phase, weight) if signals.device.type == "cuda" else _NATIVE.regress_vjp( cotangent, signals, frequency, transposed, phase, weight diff --git a/src/blochsim/estimators/_perk_gpu.py b/src/blochsim/estimators/_perk_gpu.py new file mode 100644 index 00000000..b68088b7 --- /dev/null +++ b/src/blochsim/estimators/_perk_gpu.py @@ -0,0 +1,113 @@ +"""The fused PERK kernels on a card. + +Estimating a parameter from a signal is, once PERK is fitted, one line of +arithmetic per voxel:: + + y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T + +Written as Torch operations that line builds the whole ``(voxels, features)`` +matrix, writes it to memory and reads it back to contract it away again. On a +million voxels at a thousand features that is nearly four gigabytes of traffic +carrying nothing the answer needs. The kernels form a block of features and +consume it into the output while it is still in registers, so the matrix never +exists; the adjoint forms the angles again rather than keeping them. +""" + +from __future__ import annotations + +__all__ = ["regress", "regress_vjp"] + +import math + +import torch + +from .._gpu_launch import Kernel, cdiv + +#: Voxels per program, one per thread. +_BLOCK_VOXELS = 128 + +_regress_kernel = Kernel("_regress_kernel") +_regress_vjp_kernel = Kernel("_regress_vjp_kernel") + + +def _ready(tensor: torch.Tensor) -> torch.Tensor: + """A contiguous float32 tensor the kernels can read row by row.""" + return tensor.detach().to(torch.float32).contiguous() + + +def regress( + signals: torch.Tensor, + frequency: torch.Tensor, + phase: torch.Tensor, + feature_mean: torch.Tensor, + weight: torch.Tensor, + parameter_mean: torch.Tensor, +) -> torch.Tensor: + """Estimate parameters from ``(voxels, contrasts)`` signals. + + Returns + ------- + torch.Tensor + ``(voxels, parameters)``. + """ + signals = _ready(signals) + voxels, contrasts = signals.shape + features = frequency.shape[0] + parameters = weight.shape[0] + output = torch.empty( + (voxels, parameters), dtype=torch.float32, device=signals.device + ) + if voxels: + _regress_kernel[(cdiv(voxels, _BLOCK_VOXELS),)]( + signals, + _ready(frequency), + _ready(phase), + _ready(feature_mean), + _ready(weight), + _ready(parameter_mean), + output, + voxels, + contrasts, + features, + parameters, + math.sqrt(2.0 / features), + _BLOCK_VOXELS, + ) + return output + + +def regress_vjp( + cotangent: torch.Tensor, + signals: torch.Tensor, + frequency: torch.Tensor, + phase: torch.Tensor, + weight: torch.Tensor, +) -> torch.Tensor: + """The derivative of :func:`regress` with respect to ``signals``. + + Returns + ------- + torch.Tensor + ``(voxels, contrasts)``. + """ + signals = _ready(signals) + voxels, contrasts = signals.shape + features = frequency.shape[0] + parameters = weight.shape[0] + output = torch.empty_like(signals) + if voxels: + _regress_vjp_kernel[(cdiv(voxels, _BLOCK_VOXELS),)]( + signals, + _ready(frequency), + _ready(phase), + _ready(weight), + _ready(cotangent), + output, + voxels, + contrasts, + features, + parameters, + math.sqrt(2.0 / features), + _BLOCK_VOXELS, + ) + return output diff --git a/src/blochsim/estimators/_perk_triton.py b/src/blochsim/estimators/_perk_triton.py deleted file mode 100644 index 7056444b..00000000 --- a/src/blochsim/estimators/_perk_triton.py +++ /dev/null @@ -1,332 +0,0 @@ -"""Fused Triton kernels for the PERK feature map and its regression. - -Estimating a parameter from a signal is, once PERK is fitted, one line of -arithmetic per voxel:: - - y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T - -Written as Torch operations that line builds the whole ``(voxels, features)`` -matrix, writes it to memory and reads it back to contract it away again. On a -million voxels at a thousand features that is nearly four gigabytes of traffic -carrying nothing the answer needs. - -Here a tile of features is formed and consumed into the output accumulator -while it is still in registers, so the matrix never exists. The adjoint does -the same and recomputes the tile rather than keeping it, for the same reason. -""" - -from __future__ import annotations - -__all__ = ["regress", "regress_vjp"] - -import math - -import torch -import triton -import triton.language as tl - -#: Tile shapes, swept on one card over voxel, feature and contrast blocks of -#: 32 to 128. The feature and contrast blocks are what ``tl.dot`` sees, so they -#: cannot go below the 16 it requires however narrow the problem is. -_BLOCK_VOXELS = 128 -_BLOCK_FEATURES = 64 -_BLOCK_CONTRASTS = 32 -_WARPS = 8 -_STAGES = 2 -_MIN_DOT = 16 - -#: Both contractions run at full float32. -#: -#: The output is a sum of a thousand terms that cancel, so a relative error on -#: each shows up magnified in the sum: measured against a float64 reference, -#: TF32 on either contraction gives 5e-4 to 7e-4 where full float32 gives -#: 6.6e-7 -- which is what Torch's own float32 gives on the same expression. -#: TF32 would be three times faster and three orders of magnitude further from -#: the answer, so the kernel computes the same function the composed path does -#: and the speed comes from not writing the features to memory. -_PRECISION = "ieee" - - -@triton.jit -def _regress_kernel( - signal_ptr, - frequency_ptr, - phase_ptr, - feature_mean_ptr, - weight_ptr, - parameter_mean_ptr, - output_ptr, - voxels, - contrasts, - features, - parameters, - scale, - stride_sv, - stride_sc, - stride_ff, - stride_fc, - stride_wp, - stride_wf, - stride_ov, - stride_op, - BLOCK_VOXELS: tl.constexpr, - BLOCK_FEATURES: tl.constexpr, - BLOCK_CONTRASTS: tl.constexpr, - BLOCK_PARAMETERS: tl.constexpr, - PRECISION: tl.constexpr, -): - """One block of voxels, all the way from signal to parameters.""" - voxel = tl.program_id(0) * BLOCK_VOXELS + tl.arange(0, BLOCK_VOXELS) - live = voxel < voxels - parameter = tl.arange(0, BLOCK_PARAMETERS) - wanted = parameter < parameters - - total = tl.zeros((BLOCK_VOXELS, BLOCK_PARAMETERS), dtype=tl.float32) - for start in range(0, features, BLOCK_FEATURES): - feature = start + tl.arange(0, BLOCK_FEATURES) - present = feature < features - angle = tl.zeros((BLOCK_VOXELS, BLOCK_FEATURES), dtype=tl.float32) - for first in range(0, contrasts, BLOCK_CONTRASTS): - contrast = first + tl.arange(0, BLOCK_CONTRASTS) - here = contrast < contrasts - block = tl.load( - signal_ptr + voxel[:, None] * stride_sv + contrast[None, :] * stride_sc, - mask=live[:, None] & here[None, :], - other=0.0, - ) - rows = tl.load( - frequency_ptr - + feature[:, None] * stride_ff - + contrast[None, :] * stride_fc, - mask=present[:, None] & here[None, :], - other=0.0, - ) - angle += tl.dot(block, tl.trans(rows), input_precision=PRECISION) - shift = tl.load(phase_ptr + feature, mask=present, other=0.0) - centre = tl.load(feature_mean_ptr + feature, mask=present, other=0.0) - mapped = scale * tl.cos(angle + shift[None, :]) - centre[None, :] - mapped = tl.where(present[None, :], mapped, 0.0) - columns = tl.load( - weight_ptr + parameter[:, None] * stride_wp + feature[None, :] * stride_wf, - mask=wanted[:, None] & present[None, :], - other=0.0, - ) - total += tl.dot(mapped, tl.trans(columns), input_precision=PRECISION) - - offset = tl.load(parameter_mean_ptr + parameter, mask=wanted, other=0.0) - tl.store( - output_ptr + voxel[:, None] * stride_ov + parameter[None, :] * stride_op, - total + offset[None, :], - mask=live[:, None] & wanted[None, :], - ) - - -@triton.jit -def _regress_vjp_kernel( - signal_ptr, - frequency_ptr, - phase_ptr, - weight_ptr, - cotangent_ptr, - output_ptr, - voxels, - contrasts, - features, - parameters, - scale, - stride_sv, - stride_sc, - stride_ff, - stride_fc, - stride_wp, - stride_wf, - stride_cv, - stride_cp, - stride_ov, - stride_oc, - BLOCK_VOXELS: tl.constexpr, - BLOCK_FEATURES: tl.constexpr, - BLOCK_CONTRASTS: tl.constexpr, - BLOCK_PARAMETERS: tl.constexpr, - PRECISION: tl.constexpr, -): - """The derivative of one block of voxels with respect to their signals. - - The angle the forward pass formed is rebuilt rather than stored: a tile is - cheaper to compute twice than to carry through memory once, which is the - same trade that makes the forward pass worth fusing. - """ - voxel = tl.program_id(0) * BLOCK_VOXELS + tl.arange(0, BLOCK_VOXELS) - live = voxel < voxels - parameter = tl.arange(0, BLOCK_PARAMETERS) - wanted = parameter < parameters - - seed = tl.load( - cotangent_ptr + voxel[:, None] * stride_cv + parameter[None, :] * stride_cp, - mask=live[:, None] & wanted[None, :], - other=0.0, - ) - - for outer in range(0, contrasts, BLOCK_CONTRASTS): - column = outer + tl.arange(0, BLOCK_CONTRASTS) - writing = column < contrasts - gradient = tl.zeros((BLOCK_VOXELS, BLOCK_CONTRASTS), dtype=tl.float32) - for start in range(0, features, BLOCK_FEATURES): - feature = start + tl.arange(0, BLOCK_FEATURES) - present = feature < features - angle = tl.zeros((BLOCK_VOXELS, BLOCK_FEATURES), dtype=tl.float32) - for first in range(0, contrasts, BLOCK_CONTRASTS): - contrast = first + tl.arange(0, BLOCK_CONTRASTS) - here = contrast < contrasts - block = tl.load( - signal_ptr - + voxel[:, None] * stride_sv - + contrast[None, :] * stride_sc, - mask=live[:, None] & here[None, :], - other=0.0, - ) - rows = tl.load( - frequency_ptr - + feature[:, None] * stride_ff - + contrast[None, :] * stride_fc, - mask=present[:, None] & here[None, :], - other=0.0, - ) - angle += tl.dot(block, tl.trans(rows), input_precision=PRECISION) - shift = tl.load(phase_ptr + feature, mask=present, other=0.0) - columns = tl.load( - weight_ptr - + parameter[:, None] * stride_wp - + feature[None, :] * stride_wf, - mask=wanted[:, None] & present[None, :], - other=0.0, - ) - through = tl.dot(seed, columns, input_precision=PRECISION) - through *= -scale * tl.sin(angle + shift[None, :]) - through = tl.where(present[None, :], through, 0.0) - rows = tl.load( - frequency_ptr - + feature[:, None] * stride_ff - + column[None, :] * stride_fc, - mask=present[:, None] & writing[None, :], - other=0.0, - ) - gradient += tl.dot(through, rows, input_precision=PRECISION) - tl.store( - output_ptr + voxel[:, None] * stride_ov + column[None, :] * stride_oc, - gradient, - mask=live[:, None] & writing[None, :], - ) - - -def _blocks(signals: torch.Tensor, parameters: int) -> dict[str, int]: - """Tile shapes for one problem, never below what ``tl.dot`` accepts.""" - contrasts = signals.shape[-1] - return { - "BLOCK_VOXELS": _BLOCK_VOXELS, - "BLOCK_FEATURES": _BLOCK_FEATURES, - "BLOCK_CONTRASTS": max( - _MIN_DOT, min(_BLOCK_CONTRASTS, triton.next_power_of_2(contrasts)) - ), - "BLOCK_PARAMETERS": max(_MIN_DOT, triton.next_power_of_2(parameters)), - "PRECISION": _PRECISION, - "num_warps": _WARPS, - "num_stages": _STAGES, - } - - -def regress( - signals: torch.Tensor, - frequency: torch.Tensor, - phase: torch.Tensor, - feature_mean: torch.Tensor, - weight: torch.Tensor, - parameter_mean: torch.Tensor, -) -> torch.Tensor: - """Estimate parameters from ``(voxels, contrasts)`` signals. - - Returns - ------- - torch.Tensor - ``(voxels, parameters)``. - """ - signals = signals.contiguous() - voxels, contrasts = signals.shape - features = frequency.shape[0] - parameters = weight.shape[0] - output = torch.empty( - (voxels, parameters), dtype=torch.float32, device=signals.device - ) - grid = (triton.cdiv(voxels, _BLOCK_VOXELS),) - _regress_kernel[grid]( - signals, - frequency, - phase, - feature_mean, - weight, - parameter_mean, - output, - voxels, - contrasts, - features, - parameters, - math.sqrt(2.0 / features), - signals.stride(0), - signals.stride(1), - frequency.stride(0), - frequency.stride(1), - weight.stride(0), - weight.stride(1), - output.stride(0), - output.stride(1), - **_blocks(signals, parameters), - ) - return output - - -def regress_vjp( - cotangent: torch.Tensor, - signals: torch.Tensor, - frequency: torch.Tensor, - phase: torch.Tensor, - weight: torch.Tensor, -) -> torch.Tensor: - """The derivative of :func:`regress` with respect to ``signals``. - - Returns - ------- - torch.Tensor - ``(voxels, contrasts)``. - """ - signals = signals.contiguous() - cotangent = cotangent.contiguous() - voxels, contrasts = signals.shape - features = frequency.shape[0] - parameters = weight.shape[0] - output = torch.empty_like(signals) - grid = (triton.cdiv(voxels, _BLOCK_VOXELS),) - _regress_vjp_kernel[grid]( - signals, - frequency, - phase, - weight, - cotangent, - output, - voxels, - contrasts, - features, - parameters, - math.sqrt(2.0 / features), - signals.stride(0), - signals.stride(1), - frequency.stride(0), - frequency.stride(1), - weight.stride(0), - weight.stride(1), - cotangent.stride(0), - cotangent.stride(1), - output.stride(0), - output.stride(1), - **_blocks(signals, parameters), - ) - return output diff --git a/src/blochsim/sequence/_accelerators.py b/src/blochsim/sequence/_accelerators.py index e9362678..fc78ea90 100644 --- a/src/blochsim/sequence/_accelerators.py +++ b/src/blochsim/sequence/_accelerators.py @@ -14,6 +14,7 @@ import torch +from .. import _gpu_launch from .._execution import ( Lane, _Choice, @@ -1863,7 +1864,7 @@ def _backend_available(device: torch.device) -> bool: if device.type == "cpu": from blochsim import _epg_cpu # noqa: F401 else: - from . import _epg_triton # noqa: F401 + return _gpu_launch.available() except ImportError: return False return True @@ -2874,7 +2875,7 @@ def _run_offloaded( features: frozenset[str] | None = None, ) -> torch.Tensor: """Stream a host-resident volume through the devices, chunk by chunk.""" - from ._epg_triton import simulate_into + from ._epg_gpu import simulate_into train_count = _train_count(events) voxels = tissue[0].numel() @@ -2937,7 +2938,7 @@ def _run_offloaded_jvp( Tissue seeds are per voxel and are chunked with it. Event seeds are shared by every voxel, so they are replicated once per device instead. """ - from ._epg_triton import simulate_jvp_into + from ._epg_gpu import simulate_jvp_into train_count = _train_count(events) voxels = tissue[0].numel() @@ -3007,7 +3008,7 @@ def _run_offloaded_vjp( pass holds -- so for one budget the chunks are wider, which is the second saving on top of the kernel being faster. """ - from ._epg_triton import ( + from ._epg_gpu import ( GradientBuffers, simulate_real_vjp_into, simulate_vjp_into, @@ -3112,7 +3113,7 @@ def _run_offloaded_vjp_jvp( Nothing in the loop reads a device result, which is what lets the host run ahead and keep the lanes fed. """ - from ._epg_triton import AdjointBuffers, simulate_vjp_jvp_into + from ._epg_gpu import AdjointBuffers, simulate_vjp_jvp_into train_count = _train_count(events) voxels = tissue[0].numel() @@ -3362,7 +3363,7 @@ def _run_packed( ) return moved.to(tissue[0].device) if tissue[0].device.type == "cuda": - from ._epg_triton import simulate + from ._epg_gpu import simulate shards = _shard_bounds(_train_count(events)) if shards: @@ -3543,7 +3544,7 @@ def _run_packed_vjp( "adjoint", tissue, events, output_count, state_count, real_axis ) ): - from ._epg_triton import simulate_real_vjp, simulate_vjp + from ._epg_gpu import simulate_real_vjp, simulate_vjp if real_axis == 1: return simulate_real_vjp( @@ -3808,7 +3809,7 @@ def _run_packed_vjp_jvp( origin = tissue[0].device return tuple(tuple(value.to(origin) for value in side) for side in moved) if tissue[0].device.type == "cuda": - from ._epg_triton import simulate_vjp_jvp + from ._epg_gpu import simulate_vjp_jvp shards = _shard_bounds(_train_count(events)) if shards: @@ -4030,7 +4031,7 @@ def _run_packed_jvp( ) return moved.to(tissue[0].device) if tissue[0].device.type == "cuda": - from ._epg_triton import simulate_jvp + from ._epg_gpu import simulate_jvp shards = _shard_bounds(_train_count(events)) if shards: diff --git a/src/blochsim/sequence/_epg_gpu.py b/src/blochsim/sequence/_epg_gpu.py new file mode 100644 index 00000000..0299fad3 --- /dev/null +++ b/src/blochsim/sequence/_epg_gpu.py @@ -0,0 +1,1893 @@ +"""The EPG state machines on a card, through the compiled kernels.""" + +from __future__ import annotations + +__all__: list[str] = [] + +from typing import Any + +import torch + +from .._gpu_launch import Kernel, cdiv, next_power_of_2 +from ._accelerators import _shim_count, _train_count +from ._parameters import FLOAT_NAMES as _FLOAT_NAMES +from ._parameters import ( + NARROW_SPREAD, + NO_GEOMETRY, + Geometry, + narrow_three_pool, + three_pool_spread_rate, + tissue_gradient_bases, + tissue_gradient_height, + tissue_gradient_rows, +) +from ._parameters import ( + feature_flags as _feature_flags, +) + +# Where the two event directions the real-subspace adjoint follows sit among +# the differentiable inputs. Named rather than counted, so a tissue parameter +# added ahead of them moves them instead of silently renaming a neighbour. +_DURATION_SEED = _FLOAT_NAMES.index("duration") +_FLIP_SEED = _FLOAT_NAMES.index("flip") + +_epg_kernel = Kernel("_epg_kernel") +_epg_jvp_kernel = Kernel("_epg_jvp_kernel") +_epg_vjp_kernel = Kernel("_epg_vjp_kernel") +_epg_vjp_jvp_kernel = Kernel("_epg_vjp_jvp_kernel") +_epg_real_kernel = Kernel("_epg_real_kernel") +_epg_real_jvp_kernel = Kernel("_epg_real_jvp_kernel") +_epg_real_vjp_kernel = Kernel("_epg_real_vjp_kernel") +_epg_real_vjp_jvp_kernel = Kernel("_epg_real_vjp_jvp_kernel") +_three_pool_table_kernel = Kernel("_three_pool_table_kernel") +_three_pool_table_jvp_kernel = Kernel("_three_pool_table_jvp_kernel") + + +def _pool_flag(lineshape: Any, exchanging: bool) -> int: + """Which pools a launch is to carry, as the kernels' own constexpr reads it. + + Kept in one place so a launcher cannot describe the tissue one way and the + kernel read it another. + """ + if lineshape is not None and exchanging: + return 3 + if exchanging: + return 2 + return 1 if lineshape is not None else 0 + + +# The operator table holds nine entries per voxel per distinct interval and +# the adjoint's cotangent table twelve, both float32. +_TABLE_FLOATS_PER_ROW = 9 +_BAR_FLOATS_PER_ROW = 12 + +# What the two tables may take of what the card can spare. The trajectory is +# the larger claim on the same memory and is allocated after them. +_TABLE_SHARE = 0.25 +_TABLE_FLOOR_BYTES = 64 << 20 + + +def _three_pool_table_bytes( + tissue: tuple[torch.Tensor, ...], + rows: int, + *, + problems: int | None, + dual: bool, +) -> int: + """What the tables would take -- the operator's, and the adjoint's bars. + + The operator table holds a row of voxels; the cotangent table holds a row + of problems, which is voxels times trains cut to what one chunk carries. + ``problems`` of ``None`` is a caller that builds no cotangent table. + + A ``dual`` launch stores the operator twice over, value and direction, and + pools its cotangents three times over: the value bars, the tangent bars, + and the value bars weighted by each event's own interval direction. + """ + entries = _TABLE_FLOATS_PER_ROW * (2 if dual else 1) + total = int(tissue[0].numel()) * int(rows) * entries + if problems is not None: + pooled = _BAR_FLOATS_PER_ROW * (3 if dual else 1) + total += int(problems) * int(rows) * pooled + return total * 4 + + +def _table_budget(device: torch.device) -> int: + """How many bytes the three-pool tables may claim on this device.""" + if device.type != "cuda": + return _TABLE_FLOOR_BYTES + free, _total = torch.cuda.mem_get_info(device) + return max(_TABLE_FLOOR_BYTES, int(free * _TABLE_SHARE)) + + +def _three_pool_table_jvp( + tissue: tuple[torch.Tensor, ...], + tangents: tuple[torch.Tensor, ...], + durations: torch.Tensor, +) -> torch.Tensor: + """The three-pool operator and a direction through it, per distinct length. + + Parameters + ---------- + tissue: + The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. + tangents: + The directions along them, in the same order. + durations: + The distinct interval lengths, in seconds. + + Returns + ------- + torch.Tensor + ``(rows, 18, voxels)`` float32, undamped and at ``d_dt`` of zero. + """ + voxels = int(tissue[0].numel()) + rows = int(durations.numel()) + table = torch.empty( + (rows, 18, voxels), dtype=torch.float32, device=tissue[0].device + ) + order = (0, 14, 11, 13, 10, 12, 9) + block = min(1024, next_power_of_2(max(voxels, 1))) + spread = durations.abs().to(torch.float64) * three_pool_spread_rate(tissue) + for picked, narrow in ( + (torch.nonzero(spread <= NARROW_SPREAD).flatten(), True), + (torch.nonzero(spread > NARROW_SPREAD).flatten(), False), + ): + if picked.numel() == 0: + continue + _three_pool_table_jvp_kernel[(picked.numel(), cdiv(voxels, block))]( + *(tissue[index] for index in order), + *(tangents[index] for index in order), + durations.to(torch.float32), + picked.to(torch.int32), + table, + voxels, + BLOCK=block, + narrow=narrow, + ) + return table + + +def _tabulate_three_pool( + tissue: tuple[torch.Tensor, ...], + duration: torch.Tensor, + *, + pools: int, + narrow: bool, + problems: int | None = None, + tangents: tuple[torch.Tensor, ...] | None = None, +) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: + """The operator table an event loop should read, or ``None`` to form it. + + Only a wide launch has anything to gain: under ``narrow`` the operator is + already 504 float32 instructions and forming it per event costs less than + a round trip through memory. + + A wide launch is wide because of its longest interval, and one preparation + delay is enough -- so the events that pay the roots in double are mostly + events whose own length would have taken the series. Splitting the table by + row is what lets each interval take the branch its own spread asks for, + which is worth more than sharing a row between events and does not need a + row to be shared at all. + + Parameters + ---------- + tissue: + The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. + duration: + The packed event durations, in seconds. + pools: + Which pool model the launch carries. + narrow: + Whether every interval keeps the eigenvalues close together. + problems: + How many problems a chunk of the adjoint carries, which is the height + of the cotangent table it allocates. ``None`` for a caller that builds + no such table. + tangents: + The directions along the tissue, for a caller that follows one. The + table then carries the direction through the operator beside its + value, at twice the width. + + Returns + ------- + tuple + The per-event row index, the table and the distinct lengths, or + ``(None, None, None)``. + """ + if pools != 3 or narrow: + return None, None, None + distinct, inverse = torch.unique(duration.detach(), return_inverse=True) + # A row costs a formation, an event costs one too, so a train whose + # lengths are all different has nothing to gain and a table to write. + if distinct.numel() >= duration.numel(): + return None, None, None + if _three_pool_table_bytes( + tissue, distinct.numel(), problems=problems, dual=tangents is not None + ) > _table_budget(tissue[0].device): + # A pathological train has as many lengths as events, and the tables + # grow with their product. Forming the operator per event is slower + # and always fits, so that is what an unbounded one falls back to. + return None, None, None + lengths = distinct.to(torch.float32).contiguous() + if tangents is not None: + built = _three_pool_table_jvp(tissue, tangents, lengths) + else: + built = _three_pool_table(tissue, lengths) + return inverse.reshape(duration.shape).to(torch.int32), built, lengths + + +def _three_pool_table( + tissue: tuple[torch.Tensor, ...], durations: torch.Tensor +) -> torch.Tensor: + """The three-pool operator for each distinct interval, over every voxel. + + Parameters + ---------- + tissue: + The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. + durations: + The distinct interval lengths, in seconds. + + Returns + ------- + torch.Tensor + ``(rows, 9, voxels)`` float32, undamped -- the reading event applies + its own washout. + """ + ( + t1, + _t2, + _m0, + _b1, + _b1_phase, + _b0, + _inversion, + _diffusion, + _velocity, + bound_fraction, + bound_exchange, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + _t2_pool_b, + _pool_b_shift, + ) = tissue + voxels = t1.numel() + rows = durations.numel() + table = torch.empty((rows, 9, voxels), dtype=torch.float32, device=t1.device) + # The spread a row reaches is its own length times the rate, so the split + # is exact per row rather than one verdict for the whole table. + spread = durations.abs().to(torch.float64) * three_pool_spread_rate(tissue) + block = min(1024, next_power_of_2(max(voxels, 1))) + narrow_rows = torch.nonzero(spread <= NARROW_SPREAD, as_tuple=False).flatten() + wide_rows = torch.nonzero(spread > NARROW_SPREAD, as_tuple=False).flatten() + for picked, narrow in ((narrow_rows, True), (wide_rows, False)): + if picked.numel() == 0: + continue + _three_pool_table_kernel[(picked.numel(), cdiv(voxels, block))]( + t1, + t1_pool_b, + t1_bound, + pool_b_exchange, + bound_exchange, + pool_b_fraction, + bound_fraction, + durations.to(torch.float32), + picked.to(torch.int32), + table, + voxels, + BLOCK=block, + narrow=narrow, + ) + return table + + +# Elements of the state tile one program carries, one to a thread. +_TILE_ELEMENTS = 64 + + +def _atom_stride(*tuples: tuple[torch.Tensor, ...]) -> int: + """How far to step through a property to reach one voxel's value. + + Zero where every optional property was given as one value for the whole + tissue: each is then read at one address by every voxel and needs no room + per voxel. The relaxation times lead each tuple and are stepped by one + whatever this says, since a tissue is its two relaxation times before it is + anything else. + + One stride serves the values and the directions followed beside them, so a + pass carrying tangents is asked about both: a direction laid out per voxel + has to be stepped through even where the value it follows is one number. + """ + return ( + 0 if all(value.numel() <= 1 for values in tuples for value in values[2:]) else 1 + ) + + +def _problems_per_program(block_states: int) -> int: + """How many independent problems to carry on one program's lane axis. + + A warp's lanes cost about the same whether they are used or not, so packing + several problems into one program is close to free. + + It depends on the state count alone, and deliberately not on how many + problems the launch has. A run cut into chunks would otherwise compile a + different tile from the same run whole, and the two tiles reassociate their + arithmetic differently -- so a streamed volume would answer a little + differently from an unstreamed one, which is a difference a caller has no + way to account for. ``tests/sequence/test_both_pools.py`` pins that. + + The result sizes a block of threads, so it must be a power of two. + """ + widest = max(1, _TILE_ELEMENTS // block_states) + return 1 << (widest.bit_length() - 1) + + +def _output_shape( + train_count: int, atom_count: int, output_count: int +) -> tuple[int, ...]: + """Signal shape, matching what the CPU kernels return.""" + if train_count == 1: + return (atom_count, output_count) + return (train_count, atom_count, output_count) + + +def _only_scalars(flags: dict) -> dict: + """The switches a real-subspace kernel takes. + + Off-resonance and flow are not in its representation to begin with -- it + carries three real planes where the complex kernels carry four -- so it is + given the terms that survive that reduction and no others. + """ + return { + name: flags[name] for name in ("diffusing", "transmit", "density", "inverting") + } + + +def simulate( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + features: frozenset[str] | None = None, + pools: Any = None, +) -> torch.Tensor: + """Run a packed state machine on CUDA and return complex signals. + + ``real_axis`` of 1 selects the real-subspace kernel; see + ``real_subspace_axis`` for when that is legitimate. + """ + if pools is not None: + from . import _pools_gpu + + return _pools_gpu.simulate( + tissue, + events, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + features=features, + pools=pools, + ) + train_count = _train_count(events) + atom_count = tissue[0].numel() + output_real = torch.empty( + _output_shape(train_count, atom_count, output_count), + dtype=torch.float32, + device=tissue[0].device, + ) + output_imag = torch.empty_like(output_real) + simulate_into( + tissue, + events, + output_real, + output_imag, + state_count=state_count, + output_count=output_count, + real_axis=real_axis, + atom_count=atom_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + exchanging=exchanging, + dynamic=dynamic, + features=features, + ) + return torch.complex(output_real, output_imag) + + +def simulate_into( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + output_real: torch.Tensor, + output_imag: torch.Tensor, + *, + state_count: int, + output_count: int, + real_axis: int | None, + atom_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + features: frozenset[str] | None = None, +) -> None: + """Run the forward machine into buffers the caller owns. + + Streaming reuses one set of buffers per chunk, so allocating here would put + an allocation in the loop -- and an allocation that reaches ``cudaMalloc`` + synchronizes the device, which is exactly what the streams exist to avoid. + + ``atom_count`` is given rather than taken from ``tissue`` because a chunk's + buffers are sized for the largest chunk and the last one is shorter. + """ + ( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + bound_exchange, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + ) = tissue + ( + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + ) = events + train_count = _train_count(events) + shims = _shim_count(tissue) + pools = _pool_flag(lineshape, exchanging) + block_states = next_power_of_2(state_count) + total = train_count * atom_count + problems = _problems_per_program(block_states) + grid = (cdiv(total, problems),) + # A kernel argument has to be a tensor even where the branch reading it is + # compiled out, so an unprofiled launch passes one it already has. + table = None if profile is None else profile.packed(t1.device) + pairs = None if dynamic is None else dynamic.packed(t1.device) + pair_rows = ( + None + if dynamic is None + else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) + ) + table_rows = None if profile is None else profile.rows(kind.device) + absorption = None if lineshape is None else lineshape.packed(t1.device) + narrow = narrow_three_pool(tissue, duration, pools=pools) + duration_row, pool_table, _lengths = _tabulate_three_pool( + tissue, duration, pools=pools, narrow=narrow + ) + + # Phases grow without bound under RF spoiling, so their cosines and sines + # are taken once here, in double precision, rather than in every program. + phase_cos = torch.cos(phase.double()).to(torch.float32) + phase_sin = torch.sin(phase.double()).to(torch.float32) + + if real_axis == 1: + _epg_real_kernel[grid]( + t1, + t2, + m0, + b1, + inversion_efficiency, + diffusion, + duration, + kind, + flip, + action, + output_index, + shim_index, + output_real, + output_imag, + atom_count, + train_count, + kind.numel(), + output_count, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + shimmed=_shim_count(tissue) > 1, + **_only_scalars(_feature_flags(features, geometry)), + block_states=block_states, + problems=problems, + ) + return + + _epg_kernel[grid]( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + bound_exchange, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + duration, + kind, + flip, + phase, + phase_cos, + phase_sin, + action, + output_index, + shim_index, + saturation, + rf_frequency, + t1 if table is None else table, + kind if table_rows is None else table_rows, + t1 if absorption is None else absorption, + t1 if pairs is None else pairs, + kind if pair_rows is None else pair_rows, + kind if duration_row is None else duration_row, + t1 if pool_table is None else pool_table, + output_real, + output_imag, + atom_count, + train_count, + kind.numel(), + output_count, + geometry.flow_scale, + geometry.washout_scale, + 1.0 if profile is None else profile.step, + 1.0 if lineshape is None else lineshape.step, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + shim_rows=shims, + shimmed=shims > 1, + locations=1 if profile is None else profile.points, + profiled=profile is not None and profile.bins > 0, + profile_bins=0 if profile is None else profile.bins, + dynamic=dynamic is not None, + broadened=lineshape is not None and lineshape.bins > 0, + lineshape_bins=0 if lineshape is None else lineshape.bins, + pools=pools, + narrow=narrow, + tabulated=pool_table is not None, + **_feature_flags(features, geometry), + block_states=block_states, + problems=problems, + ) + + +def simulate_jvp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tissue_tangents: tuple[torch.Tensor, ...], + event_tangents: tuple[torch.Tensor, torch.Tensor, torch.Tensor], + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + dynamic_direction: Any = None, + features: frozenset[str] | None = None, + pools: Any = None, +) -> torch.Tensor: + """Run one fused state-machine Jacobian-vector product on CUDA. + + ``real_axis`` of 1 selects the real-subspace kernel, which produces no + derivative along ``b1_phase``, ``b0`` or the RF phase -- seeds along those + directions leave the subspace, so the caller must rule them out. + """ + if pools is not None: + from . import _pools_gpu + + return _pools_gpu.simulate_jvp( + tissue, + events, + tissue_tangents, + event_tangents, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + features=features, + pools=pools, + ) + train_count = _train_count(events) + atom_count = tissue[0].numel() + output_real = torch.empty( + _output_shape(train_count, atom_count, output_count), + dtype=torch.float32, + device=tissue[0].device, + ) + output_imag = torch.empty_like(output_real) + simulate_jvp_into( + tissue, + events, + tissue_tangents, + event_tangents, + output_real, + output_imag, + state_count=state_count, + output_count=output_count, + real_axis=real_axis, + atom_count=atom_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + exchanging=exchanging, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + features=features, + ) + return torch.complex(output_real, output_imag) + + +def simulate_jvp_into( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tissue_tangents: tuple[torch.Tensor, ...], + event_tangents: tuple[torch.Tensor, torch.Tensor, torch.Tensor], + output_real: torch.Tensor, + output_imag: torch.Tensor, + *, + state_count: int, + output_count: int, + real_axis: int | None, + atom_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + dynamic_direction: Any = None, + features: frozenset[str] | None = None, +) -> None: + """Run one Jacobian-vector product into buffers the caller owns. + + See ``simulate_into`` for why the streaming path needs this. + """ + ( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + bound_exchange, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + ) = tissue + ( + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + ) = events + tangent_duration, tangent_flip, tangent_phase = event_tangents + train_count = _train_count(events) + pools = _pool_flag(lineshape, exchanging) + shims = _shim_count(tissue) + block_states = next_power_of_2(state_count) + total = train_count * atom_count + problems = _problems_per_program(block_states) + grid = (cdiv(total, problems),) + + if real_axis == 1: + _epg_real_jvp_kernel[grid]( + t1, + t2, + m0, + b1, + inversion_efficiency, + diffusion, + duration, + kind, + flip, + action, + output_index, + shim_index, + tissue_tangents[0], + tissue_tangents[1], + tissue_tangents[2], + tissue_tangents[3], + tissue_tangents[6], + tissue_tangents[7], + tangent_duration, + tangent_flip, + output_real, + output_imag, + atom_count, + train_count, + kind.numel(), + output_count, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue, tissue_tangents), + shimmed=shims > 1, + **_only_scalars(_feature_flags(features, geometry)), + block_states=block_states, + problems=problems, + ) + return + + table = None if profile is None else profile.packed(t1.device) + pairs = None if dynamic is None else dynamic.packed(t1.device) + pair_rows = ( + None + if dynamic is None + else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) + ) + pair_direction = ( + None if dynamic_direction is None else dynamic_direction.to(t1.device) + ) + table_rows = None if profile is None else profile.rows(kind.device) + absorption = None if lineshape is None else lineshape.packed(t1.device) + narrow = narrow_three_pool(tissue, duration, pools=pools) + duration_row, pool_table, _lengths = _tabulate_three_pool( + tissue, duration, pools=pools, narrow=narrow, tangents=tissue_tangents + ) + _epg_jvp_kernel[grid]( + *tissue, + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + *tissue_tangents, + tangent_duration, + tangent_flip, + tangent_phase, + saturation, + rf_frequency, + t1 if table is None else table, + kind if table_rows is None else table_rows, + t1 if absorption is None else absorption, + t1 if pairs is None else pairs, + kind if pair_rows is None else pair_rows, + t1 if pair_direction is None else pair_direction, + kind if duration_row is None else duration_row, + t1 if pool_table is None else pool_table, + output_real, + output_imag, + atom_count, + train_count, + kind.numel(), + output_count, + geometry.flow_scale, + geometry.washout_scale, + 1.0 if profile is None else profile.step, + 1.0 if lineshape is None else lineshape.step, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue, tissue_tangents), + shim_rows=shims, + shimmed=shims > 1, + locations=1 if profile is None else profile.points, + profiled=profile is not None and profile.bins > 0, + profile_bins=0 if profile is None else profile.bins, + dynamic=dynamic is not None, + broadened=lineshape is not None and lineshape.bins > 0, + lineshape_bins=0 if lineshape is None else lineshape.bins, + pools=pools, + narrow=narrow, + tabulated=pool_table is not None, + **_feature_flags(features, geometry), + block_states=block_states, + problems=problems, + ) + + +# How much device memory the recorded trajectory may hold at once. Beyond this +# the problems are run in waves, which the gradient buffers absorb because they +# accumulate rather than being written. +_TRAJECTORY_BUDGET_BYTES = 256 << 20 + + +def _trajectory_wave( + event_count: int, state_count: int, total: int, planes: int, blocks: int = 3 +) -> int: + """How many problems can record their trajectory in one launch.""" + per_problem = event_count * blocks * state_count * planes * 4 + return max(1, min(total, _TRAJECTORY_BUDGET_BYTES // max(1, per_problem))) + + +def simulate_vjp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + dynamic: Any = None, + lineshape: Any = None, + exchanging: bool = False, + features: frozenset[str] | None = None, + pools: Any = None, +) -> tuple[torch.Tensor, ...]: + """The first-order adjoint on CUDA, for a whole volume on one device. + + Returns the gradients in the differentiable-input order -- every tissue + property, then event duration, flip and phase, and the pair's cotangent + where one is given. A shard takes this same kernel a level down and a + streamed volume has chunked launchers of its own; + :func:`blochsim.sequence._accelerators` decides which route a run takes. + + Carrying no forward direction, this records two trajectory planes per + recorded state where that pass records four, and holds one state where it + holds a dual. + """ + if pools is not None: + from . import _pools_gpu + + return _pools_gpu.simulate_vjp( + tissue, + events, + grad_output, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + dynamic=dynamic, + lineshape=lineshape, + features=features, + pools=pools, + ) + ( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + exchange_rate, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + ) = tissue + ( + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + ) = events[:9] + atom_count = t1.numel() + train_count = _train_count(events) + event_count = kind.numel() + total = train_count * atom_count + block_states = next_power_of_2(state_count) + device = t1.device + shims = max(1, b1.numel() // atom_count) if atom_count else 1 + table = None if profile is None else profile.packed(device) + table_rows = None if profile is None else profile.rows().to(device) + pairs = None if dynamic is None else dynamic.packed(device) + pair_rows = ( + None + if dynamic is None + else dynamic.rows_per_event(train_count, event_count).to(device) + ) + grad_pair = None if dynamic is None else torch.zeros_like(pairs) + locations = 1 if profile is None else profile.points + absorption = None if lineshape is None else lineshape.packed(device) + + grad_tissue = torch.zeros( + tissue_gradient_height(shims) * atom_count, + dtype=torch.float32, + device=device, + ) + grad_flip = torch.zeros_like(flip) + grad_phase = torch.zeros_like(phase) + grad_duration = torch.zeros_like(duration) + grad_output = grad_output.resolve_conj() + grad_real = grad_output.real.contiguous() + grad_imag = grad_output.imag.contiguous() + + # A semisolid pool records a plane of its own beside the three the free + # water keeps; a chemically exchanging one three, and the two together + # four. + pools = _pool_flag(lineshape, exchanging) + narrow = narrow_three_pool(tissue, duration, pools=pools) + blocks = 7 if pools == 3 else (6 if pools == 2 else (4 if pools == 1 else 3)) + wave = _trajectory_wave(event_count, state_count, total, 2, blocks) + duration_row, pool_table, pool_durations = _tabulate_three_pool( + tissue, duration, pools=pools, narrow=narrow, problems=wave + ) + row_count = 0 if pool_durations is None else pool_durations.numel() + pool_bars = None + if pool_table is not None: + # A slot per problem the chunk carries, so the walk back accumulates + # into memory it owns and no two programs contend for a row. The + # chunks run one after another, so one chunk's worth is enough. + pool_bars = torch.zeros( + wave * row_count * 12, dtype=torch.float32, device=device + ) + trajectory = [ + torch.empty( + (wave, event_count * blocks * state_count), + dtype=torch.float32, + device=device, + ) + for _ in range(2) + ] + + problems = _problems_per_program(block_states) + for base in range(0, total, wave): + span = min(wave, total - base) + if pool_bars is not None: + # The slots are per chunk, so each chunk starts from nothing. + pool_bars.zero_() + # The trajectory is written by one launch and walked back by the + # next, so each compiles one sweep instead of both. + for recording in (True, False): + _epg_vjp_kernel[(cdiv(span, problems),)]( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + exchange_rate, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + absorption, + table, + table_rows, + pairs, + pair_rows, + duration_row, + pool_table, + pool_bars, + pool_durations, + row_count, + grad_pair, + grad_real, + grad_imag, + grad_tissue, + grad_flip, + grad_phase, + grad_duration, + *trajectory, + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + geometry.flow_scale, + geometry.washout_scale, + shims, + 1.0 if profile is None else profile.step, + 1.0 if lineshape is None else lineshape.step, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + shimmed=shims > 1, + locations=locations, + profiled=profile is not None and profile.bins > 0, + profile_bins=0 if profile is None else profile.bins, + dynamic=dynamic is not None, + broadened=lineshape is not None and lineshape.bins > 0, + lineshape_bins=0 if lineshape is None else lineshape.bins, + pools=pools, + narrow=narrow, + tabulated=pool_table is not None, + recording=recording, + block_states=block_states, + problems=problems, + **_feature_flags(features, geometry), + ) + voxel = tuple( + grad_tissue[base * atom_count : (base + rows) * atom_count] + for base, rows in zip( + tissue_gradient_bases(shims), tissue_gradient_rows(shims), strict=True + ) + ) + if dynamic is not None: + return (*voxel, grad_duration, grad_flip, grad_phase, grad_pair) + return (*voxel, grad_duration, grad_flip, grad_phase) + + +def simulate_real_vjp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + features: frozenset[str] | None = None, +) -> tuple[torch.Tensor, ...]: + """The first-order adjoint through the real subspace, on CUDA. + + Returns the gradients in the differentiable-input order -- every tissue + property, then event duration, flip and phase. The representation divides + the RF phase out, so transmit phase, off-resonance, velocity and RF phase + come back at zero and callers must not ask for those. + + Carrying no forward direction, this records one trajectory plane where the + forward-over-reverse pass records two, and holds one state where it holds a + dual. + """ + ( + t1, + t2, + m0, + b1, + _b1_phase, + _b0, + inversion_efficiency, + diffusion, + *_rest, + ) = tissue + duration, kind, flip, phase, action, output_index, shim_index = events[:7] + atom_count = t1.numel() + train_count = _train_count(events) + event_count = kind.numel() + total = train_count * atom_count + block_states = next_power_of_2(state_count) + device = t1.device + shims = _shim_count(tissue) + + grad_tissue = torch.zeros( + tissue_gradient_height(shims) * atom_count, + dtype=torch.float32, + device=device, + ) + grad_flip = torch.zeros_like(flip) + grad_duration = torch.zeros_like(duration) + grad_phase = torch.zeros_like(phase) + grad_imag = grad_output.resolve_conj().imag.contiguous() + + wave = _trajectory_wave(event_count, state_count, total, 1) + trajectory = torch.empty( + (wave, event_count * 3 * state_count), dtype=torch.float32, device=device + ) + + problems = _problems_per_program(block_states) + for base in range(0, total, wave): + span = min(wave, total - base) + _epg_real_vjp_kernel[(cdiv(span, problems),)]( + t1, + t2, + m0, + b1, + inversion_efficiency, + diffusion, + duration, + kind, + flip, + action, + output_index, + shim_index, + grad_imag, + grad_tissue, + grad_flip, + grad_duration, + trajectory, + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + shim_rows=shims, + shimmed=shims > 1, + **_only_scalars(_feature_flags(features, NO_GEOMETRY)), + block_states=block_states, + problems=problems, + ) + voxel = tuple( + grad_tissue[base * atom_count : (base + rows) * atom_count] + for base, rows in zip( + tissue_gradient_bases(shims), tissue_gradient_rows(shims), strict=True + ) + ) + return (*voxel, grad_duration, grad_flip, grad_phase) + + +class AdjointBuffers: + """Device memory a forward-over-reverse pass writes into. + + Sized for ``chunk`` voxels and reusable for any narrower one. Per-voxel + gradients are cleared before each pass; per-event gradients accumulate over + every pass the buffers serve and are read out with ``event_gradients``. + + ``real_axis`` of 1 halves the state planes, so buffers built for one + representation cannot be handed to the other. + """ + + def __init__( + self, + events: tuple[torch.Tensor, ...], + chunk: int, + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + shims: int = 1, + pools: int = 0, + ) -> None: + ( + duration, + kind, + flip, + phase, + _action, + _output_index, + _shim, + _saturation, + _rf_frequency, + ) = events + device = kind.device + train_count = _train_count(events) + event_count = kind.numel() + self.planes = 2 if real_axis == 1 else 4 + # A bound pool records a fourth block of states per event: the RF + # operator scales it, so the reverse sweep cannot replay it from the + # free pool's. + self.blocks = 3 + (4 if pools == 3 else (3 if pools == 2 else pools)) + self.chunk = chunk + self.shims = shims + self.rows = tissue_gradient_height(shims) + self.state_count = state_count + self.output_count = output_count + self.train_count = train_count + # One dual accumulator per plane: value is the gradient w.r.t. the + # tangent inputs, tangent the gradient w.r.t. the primal ones. + self.tissue = [ + torch.zeros(self.rows * chunk, dtype=torch.float32, device=device) + for _ in range(2) + ] + self.flip = [torch.zeros_like(flip) for _ in range(2)] + self.duration = [torch.zeros_like(duration) for _ in range(2)] + self.phase = [torch.zeros_like(phase) for _ in range(2)] + self.cotangent = [ + torch.empty( + train_count * chunk * output_count, + dtype=torch.float32, + device=device, + ) + for _ in range(2) + ] + self.wave = _trajectory_wave( + event_count, + state_count, + train_count * chunk, + self.planes, + self.blocks, + ) + self.trajectory = [ + torch.empty( + (self.wave, event_count * self.blocks * state_count), + dtype=torch.float32, + device=device, + ) + for _ in range(self.planes) + ] + + def tissue_gradients(self, atom_count: int) -> tuple[tuple[torch.Tensor, ...], ...]: + """The per-voxel gradients of the last pass, one entry per parameter. + + Each is flat and as wide as the buffer it belongs to, so the transmit + pair spans every shim. Ordered to match ``event_gradients``: tangent + plane first. + """ + return tuple( + tuple( + self.tissue[plane][base * atom_count : (base + rows) * atom_count] + for base, rows in zip( + tissue_gradient_bases(self.shims), + tissue_gradient_rows(self.shims), + strict=True, + ) + ) + for plane in (1, 0) + ) + + def event_gradients(self) -> tuple[tuple[torch.Tensor, ...], ...]: + """The per-event gradients summed over every pass so far. + + Ordered ``(duration, flip, phase)`` to match the tail of the + differentiable-input order, tangent plane first. + """ + return tuple( + (self.duration[plane], self.flip[plane], self.phase[plane]) + for plane in (1, 0) + ) + + +class GradientBuffers: + """Device memory a first-order adjoint writes into. + + Half of what the forward-over-reverse pass needs: one accumulator per + gradient rather than a dual, and one trajectory plane per real state rather + than a plane per component of one. Sized for ``chunk`` voxels and reusable + for any narrower one; per-event gradients accumulate over every pass the + buffers serve. + + ``real_axis`` of 1 halves the planes again, so buffers built for one + representation cannot be handed to the other. + """ + + def __init__( + self, + events: tuple[torch.Tensor, ...], + chunk: int, + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + ) -> None: + duration, kind, flip, phase = events[:4] + device = kind.device + train_count = _train_count(events) + event_count = kind.numel() + self.real_axis = real_axis + self.planes = 1 if real_axis == 1 else 2 + self.chunk = chunk + self.rows = tissue_gradient_height(1) + self.state_count = state_count + self.output_count = output_count + self.train_count = train_count + self.tissue = torch.zeros(self.rows * chunk, dtype=torch.float32, device=device) + self.flip = torch.zeros_like(flip) + self.duration = torch.zeros_like(duration) + self.phase = torch.zeros_like(phase) + self.cotangent = [ + torch.empty( + train_count * chunk * output_count, + dtype=torch.float32, + device=device, + ) + for _ in range(2) + ] + self.wave = _trajectory_wave( + event_count, state_count, train_count * chunk, self.planes + ) + self.trajectory = [ + torch.empty( + (self.wave, event_count * 3 * state_count), + dtype=torch.float32, + device=device, + ) + for _ in range(self.planes) + ] + + def tissue_gradients(self, atom_count: int) -> tuple[torch.Tensor, ...]: + """The per-voxel gradients of the last pass, one entry per parameter.""" + return tuple( + self.tissue[base * atom_count : (base + rows) * atom_count] + for base, rows in zip( + tissue_gradient_bases(1), tissue_gradient_rows(1), strict=True + ) + ) + + def event_gradients(self) -> tuple[torch.Tensor, ...]: + """``(duration, flip, phase)``, summed over every pass so far.""" + return (self.duration, self.flip, self.phase) + + +def simulate_vjp_into( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + buffers: GradientBuffers, + *, + state_count: int, + output_count: int, + atom_count: int, + geometry: Geometry = NO_GEOMETRY, + features: frozenset[str] | None = None, +) -> tuple[torch.Tensor, ...]: + """One chunk of a first-order adjoint, into buffers the caller owns. + + ``grad_output`` is already on the device. Returns the per-voxel gradients + of this chunk; the per-event ones accumulate in ``buffers``. + """ + ( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + exchange_rate, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + ) = tissue + ( + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + ) = events[:9] + train_count = _train_count(events) + event_count = kind.numel() + total = train_count * atom_count + block_states = next_power_of_2(state_count) + + buffers.tissue.zero_() + grad_output = grad_output.resolve_conj() + size = total * output_count + grad_real = buffers.cotangent[0][:size] + grad_imag = buffers.cotangent[1][:size] + grad_real.copy_(grad_output.real.reshape(-1)) + grad_imag.copy_(grad_output.imag.reshape(-1)) + + problems = _problems_per_program(block_states) + for base in range(0, total, buffers.wave): + span = min(buffers.wave, total - base) + # The trajectory is written by one launch and walked back by the + # next, so each compiles one sweep instead of both. + for recording in (True, False): + _epg_vjp_kernel[(cdiv(span, problems),)]( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + bound_fraction, + exchange_rate, + t1_bound, + pool_b_fraction, + pool_b_exchange, + t1_pool_b, + t2_pool_b, + pool_b_shift, + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + saturation, + rf_frequency, + None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + 0, + grad_real, + grad_imag, + buffers.tissue, + buffers.flip, + buffers.phase, + buffers.duration, + *buffers.trajectory, + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + geometry.flow_scale, + geometry.washout_scale, + 1, + 1.0, + 1.0, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + shimmed=False, + locations=1, + profiled=False, + profile_bins=0, + dynamic=False, + broadened=False, + lineshape_bins=0, + pools=0, + narrow=False, + tabulated=False, + recording=recording, + block_states=block_states, + problems=problems, + **_feature_flags(features, geometry), + ) + return buffers.tissue_gradients(atom_count) + + +def simulate_real_vjp_into( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + buffers: GradientBuffers, + *, + state_count: int, + output_count: int, + atom_count: int, + features: frozenset[str] | None = None, +) -> tuple[torch.Tensor, ...]: + """The same, for a train the real subspace covers.""" + ( + t1, + t2, + m0, + b1, + _b1_phase, + _b0, + inversion_efficiency, + diffusion, + *_rest, + ) = tissue + duration, kind, flip, _phase, action, output_index, shim_index = events[:7] + train_count = _train_count(events) + event_count = kind.numel() + total = train_count * atom_count + block_states = next_power_of_2(state_count) + + buffers.tissue.zero_() + size = total * output_count + grad_imag = buffers.cotangent[1][:size] + grad_imag.copy_(grad_output.resolve_conj().imag.reshape(-1)) + + problems = _problems_per_program(block_states) + for base in range(0, total, buffers.wave): + span = min(buffers.wave, total - base) + _epg_real_vjp_kernel[(cdiv(span, problems),)]( + t1, + t2, + m0, + b1, + inversion_efficiency, + diffusion, + duration, + kind, + flip, + action, + output_index, + shim_index, + grad_imag, + buffers.tissue, + buffers.flip, + buffers.duration, + buffers.trajectory[0], + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue), + # The streamed route carries one shim, as the complex one does: + # ``GradientBuffers`` sizes its gradient plane for a single row. + shim_rows=1, + shimmed=False, + block_states=block_states, + problems=problems, + **_only_scalars(_feature_flags(features, NO_GEOMETRY)), + ) + return buffers.tissue_gradients(atom_count) + + +def simulate_vjp_jvp_into( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tangents: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + buffers: AdjointBuffers, + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + atom_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + dynamic_direction: Any = None, + dynamic_gradients: tuple[torch.Tensor, torch.Tensor] | None = None, + features: frozenset[str] | None = None, +) -> tuple[tuple[torch.Tensor, ...], ...]: + """Forward-over-reverse for one chunk of voxels, into caller-owned buffers. + + Returns the per-voxel gradients of this chunk -- tangent plane first, one + entry per tissue parameter, views into ``buffers`` that the next call + overwrites. The per-event gradients accumulate inside ``buffers`` instead, + because every chunk contributes to all of them. + + ``atom_count`` is this chunk's width, which may be narrower than the one + the buffers were built for. + """ + ( + t1, + t2, + m0, + b1, + b1_phase, + b0, + inversion_efficiency, + diffusion, + velocity, + _bound_fraction, + _bound_exchange, + _t1_bound, + _pool_b_fraction, + _pool_b_exchange, + _t1_pool_b, + _t2_pool_b, + _pool_b_shift, + ) = tissue + ( + duration, + kind, + flip, + phase, + action, + output_index, + shim_index, + _saturation, + _rf_frequency, + ) = events + train_count = _train_count(events) + pools = _pool_flag(lineshape, exchanging) + event_count = kind.numel() + total = train_count * atom_count + block_states = next_power_of_2(state_count) + real = real_axis == 1 + + grad_output = grad_output.resolve_conj() + size = total * output_count + grad_real, grad_imag = ( + plane[:size].view(grad_output.shape) for plane in buffers.cotangent + ) + grad_real.copy_(grad_output.real) + grad_imag.copy_(grad_output.imag) + grad_tissue = [plane[: buffers.rows * atom_count] for plane in buffers.tissue] + for plane in grad_tissue: + plane.zero_() + grad_flip, grad_duration, grad_phase = buffers.flip, buffers.duration, buffers.phase + trajectory = buffers.trajectory + table = None if profile is None else profile.packed(t1.device) + pairs = None if dynamic is None else dynamic.packed(t1.device) + pair_rows = ( + None + if dynamic is None + else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) + ) + pair_direction = ( + None if dynamic_direction is None else dynamic_direction.to(t1.device) + ) + grad_pair_value = None if dynamic_gradients is None else dynamic_gradients[0] + grad_pair_tangent = None if dynamic_gradients is None else dynamic_gradients[1] + table_rows = None if profile is None else profile.rows(kind.device) + absorption = None if lineshape is None else lineshape.packed(t1.device) + narrow = narrow_three_pool(tissue, duration, pools=pools) + wave = buffers.wave + duration_row, pool_table, pool_durations = _tabulate_three_pool( + tissue, + duration, + pools=pools, + narrow=narrow, + tangents=tangents, + problems=wave, + ) + row_count = 0 if pool_durations is None else pool_durations.numel() + pool_bars = None + if pool_table is not None: + # A slot per problem the chunk carries, so the walk back accumulates + # into memory it owns and no two programs contend for a row. Three + # sets of twelve: the value cotangents, their directions, and the + # value cotangents weighted by each event's own interval direction. + pool_bars = torch.zeros( + wave * row_count * 36, dtype=torch.float32, device=t1.device + ) + + problems = _problems_per_program(block_states) + for base in range(0, total, wave): + span = min(wave, total - base) + if pool_bars is not None: + # The slots are per chunk, so each chunk starts from nothing. + pool_bars.zero_() + grid = (cdiv(span, problems),) + shape = dict( + state_count=state_count, + single_train=train_count == 1, + atom_stride=_atom_stride(tissue, tangents), + block_states=block_states, + problems=problems, + ) + if real: + _epg_real_vjp_jvp_kernel[grid]( + t1, + t2, + m0, + b1, + inversion_efficiency, + diffusion, + duration, + kind, + flip, + action, + output_index, + shim_index, + tangents[0], + tangents[1], + tangents[2], + tangents[3], + tangents[6], + tangents[7], + tangents[_DURATION_SEED], + tangents[_FLIP_SEED], + grad_imag, + *grad_tissue, + *grad_flip, + *grad_duration, + *trajectory, + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + shim_rows=_shim_count(tissue), + shimmed=_shim_count(tissue) > 1, + **_only_scalars(_feature_flags(features, geometry)), + **shape, + ) + else: + # The trajectory is written by one launch and walked back by the + # next, so each compiles one sweep instead of both. + for recording in (True, False): + _epg_vjp_jvp_kernel[grid]( + *tissue, + *events, + t1 if table is None else table, + kind if table_rows is None else table_rows, + t1 if absorption is None else absorption, + t1 if pairs is None else pairs, + kind if pair_rows is None else pair_rows, + t1 if pair_direction is None else pair_direction, + t1 if grad_pair_value is None else grad_pair_value, + t1 if grad_pair_tangent is None else grad_pair_tangent, + *tangents, + kind if duration_row is None else duration_row, + t1 if pool_table is None else pool_table, + t1 if pool_bars is None else pool_bars, + t1 if pool_durations is None else pool_durations, + row_count, + grad_real, + grad_imag, + *grad_tissue, + *grad_flip, + *grad_phase, + *grad_duration, + *trajectory, + base, + base + span, + atom_count, + train_count, + event_count, + output_count, + geometry.flow_scale, + geometry.washout_scale, + 1.0 if profile is None else profile.step, + 1.0 if lineshape is None else lineshape.step, + shim_rows=_shim_count(tissue), + shimmed=_shim_count(tissue) > 1, + locations=1 if profile is None else profile.points, + profiled=profile is not None and profile.bins > 0, + profile_bins=0 if profile is None else profile.bins, + dynamic=dynamic is not None, + directed=dynamic_direction is not None, + broadened=lineshape is not None and lineshape.bins > 0, + lineshape_bins=0 if lineshape is None else lineshape.bins, + pools=pools, + narrow=narrow, + tabulated=pool_table is not None, + recording=recording, + **_feature_flags(features, geometry), + **shape, + ) + + # Plane 1 is the tangent part -> d/d(primal inputs); plane 0 the value part. + return buffers.tissue_gradients(atom_count) + + +def simulate_vjp_jvp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tangents: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + real_axis: int | None = None, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + exchanging: bool = False, + dynamic: Any = None, + dynamic_direction: Any = None, + features: frozenset[str] | None = None, + pools: Any = None, +) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: + """Forward-over-reverse through the state machine on CUDA. + + ``tangents`` follows the differentiable-input order -- every tissue + property, then event duration, flip and phase -- and the two returned + tuples, gradients with respect to the primal inputs then to the tangent + inputs, follow it too. + + ``real_axis`` of 1 selects the real-subspace adjoint. That representation + divides the RF phase out, so it leaves ``b1_phase``, ``b0`` and ``phase`` at + zero and callers must not ask for those; the complex adjoint produces every + one of them. + + Gradients land through atomic accumulation, so repeated runs agree to + floating-point tolerance rather than bit for bit. + """ + if pools is not None: + from . import _pools_gpu + + return _pools_gpu.simulate_vjp_jvp( + tissue, + events, + tangents, + grad_output, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + features=features, + pools=pools, + ) + atom_count = tissue[0].numel() + gradients = None + if dynamic is not None: + held = dynamic.packed(tissue[0].device) + gradients = (torch.zeros_like(held), torch.zeros_like(held)) + buffers = AdjointBuffers( + events, + atom_count, + state_count=state_count, + output_count=output_count, + real_axis=real_axis, + shims=_shim_count(tissue), + pools=_pool_flag(lineshape, exchanging), + ) + voxel_grads = simulate_vjp_jvp_into( + tissue, + events, + tangents, + grad_output, + buffers, + state_count=state_count, + output_count=output_count, + real_axis=real_axis, + atom_count=atom_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + exchanging=exchanging, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + dynamic_gradients=gradients, + features=features, + ) + sides = tuple( + (*voxels, *per_event) + for voxels, per_event in zip( + voxel_grads, buffers.event_gradients(), strict=True + ) + ) + if gradients is None: + return sides + # The value plane is the adjoint and the tangent plane its own derivative, + # which is the split the tissue gradients take; the sides come back in the + # order the caller reads them, curvature first. + return (*sides[0], gradients[1]), (*sides[1], gradients[0]) diff --git a/src/blochsim/sequence/_epg_triton.py b/src/blochsim/sequence/_epg_triton.py deleted file mode 100644 index fe227f58..00000000 --- a/src/blochsim/sequence/_epg_triton.py +++ /dev/null @@ -1,18718 +0,0 @@ -"""Fused Triton kernel for inference-only EPG state machines.""" - -from __future__ import annotations - -__all__: list[str] = [] - -from typing import Any - -import torch -import triton -import triton.language as tl -from triton.language.extra import libdevice - -from ._accelerators import _shim_count, _train_count -from ._parameters import BOUND_POOL_INPUTS as _BOUND_POOL_INPUTS -from ._parameters import EXCHANGE_POOL_INPUTS as _EXCHANGE_POOL_INPUTS -from ._parameters import FLOAT_NAMES as _FLOAT_NAMES -from ._parameters import ( - NARROW_SPREAD, - NO_GEOMETRY, - Geometry, - narrow_three_pool, - three_pool_spread_rate, - tissue_gradient_bases, - tissue_gradient_height, - tissue_gradient_rows, -) -from ._parameters import TISSUE_COUNT as _TISSUE_PARAMETERS -from ._parameters import TRANSMIT_INPUTS as _TRANSMIT_INPUTS -from ._parameters import ( - feature_flags as _feature_flags, -) - -# Triton reads globals only through its own constexpr wrapper. -_TISSUE_COUNT = tl.constexpr(_TISSUE_PARAMETERS) - -# How many tissue parameters the free pool alone accounts for. Both second -# pools' sit past them, and are written by the kernels that carry them; a -# single-pool run leaves their planes at the zero they were cleared to. -_FREE_POOL_COUNT = tl.constexpr( - _TISSUE_PARAMETERS - len(_BOUND_POOL_INPUTS) - len(_EXCHANGE_POOL_INPUTS) -) -_BOUND_ROW = tl.constexpr(_BOUND_POOL_INPUTS[0]) -_POOL_B_ROW = tl.constexpr(_EXCHANGE_POOL_INPUTS[0]) - -# The gradient plane holds a row of voxels per tissue parameter, except that -# the transmit pair holds one per shim. Both sit ahead of everything that -# widens, so a plane's row is its parameter index shifted by the rows the pair -# added ahead of it. -_B1_ROW = tl.constexpr(_TRANSMIT_INPUTS[0]) -_B1_PHASE_ROW = tl.constexpr(_TRANSMIT_INPUTS[1]) - -# Where the two event directions the real-subspace adjoint follows sit among -# the differentiable inputs. Named rather than counted, so a tissue parameter -# added ahead of them moves them instead of silently renaming a neighbour. -_DURATION_SEED = _FLOAT_NAMES.index("duration") -_FLIP_SEED = _FLOAT_NAMES.index("flip") - - -@triton.jit -def _up(values, state): - """``values`` moved one configuration order up: ``result[k] = values[k - 1]``. - - A shift along the state axis of a tile the program already holds. Order - zero is left to the caller, which fills it from the sequence's own boundary - condition rather than from a neighbour. - """ - index = tl.broadcast_to(tl.maximum(state - 1, 0), values.shape) - return tl.gather(values, index, 1) - - -@triton.jit -def _down(values, state): - """``values`` moved one order down: ``result[k] = values[k + 1]``. - - The top order has no neighbour to read, so it reads itself and the caller - masks it away. - """ - index = tl.broadcast_to(tl.minimum(state + 1, values.shape[1] - 1), values.shape) - return tl.gather(values, index, 1) - - -@triton.jit -def _first(values, state): - """Order zero of ``values``, spread across every order.""" - index = tl.broadcast_to(state * 0, values.shape) - return tl.gather(values, index, 1) - - -@triton.jit -def _sincos(x): - """The sine and cosine of ``x`` from one reduction by a quarter turn. - - Cody and Waite's three-part quarter turn and the single-precision Cephes - polynomials on the eighth turn either side of zero, to about an ulp where - ``|x|`` is a flip angle; one reduction serves both where two library calls - would each make their own. - """ - quarter = tl.extra.cuda.libdevice.rint(x * 0.6366197723675814) - r = tl.fma(-quarter, 1.5703125, x) - r = tl.fma(-quarter, 4.837512969970703125e-4, r) - r = tl.fma(-quarter, 7.54978995489188216e-8, r) - r2 = r * r - sine = tl.fma(-1.9515295891e-4, r2, 8.3321608736e-3) - sine = tl.fma(sine, r2, -1.6666654611e-1) - sine = tl.fma(r * r2, sine, r) - cosine = tl.fma(2.443315711809948e-5, r2, -1.388731625493765e-3) - cosine = tl.fma(cosine, r2, 4.166664568298827e-2) - cosine = tl.fma(r2 * r2, cosine, tl.fma(-0.5, r2, 1.0)) - q = quarter.to(tl.int32) & 3 - s = tl.where( - q == 0, sine, tl.where(q == 1, cosine, tl.where(q == 2, -sine, -cosine)) - ) - c = tl.where( - q == 0, cosine, tl.where(q == 1, -sine, tl.where(q == 2, -cosine, sine)) - ) - return s, c - - -@triton.jit -def _event_value(values, event_base, event, active_atom, single_train: tl.constexpr): - """One event's entry of a buffer carrying a row per train. - - ``duration``, ``flip`` and ``phase`` are indexed by the train and the event - and never by the atom, so where there is one train the address is the same - for every lane of the program and the value can be read once. Triton cannot - see that through ``event_base``, and a tile-shaped load emits one - instruction per element the lane holds -- four reads of one number. - """ - if single_train: - # Spread over the program's problems, which a jitted helper has to do - # for itself: both arms of the branch have to hand back the one shape. - return tl.load(values + event) + tl.zeros_like(event_base.to(tl.float32)) - return tl.load(values + event_base + event, mask=active_atom, other=0.0) - - -@triton.jit -def _shift( - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - state, - state_mask, - state_count, -): - keep_up = (state > 0) & state_mask - keep_down = (state + 1 < state_count) & state_mask - plus_real = tl.where(keep_up, _up(fplus_real, state), 0.0) - plus_imag = tl.where(keep_up, _up(fplus_imag, state), 0.0) - minus_real = tl.where(keep_down, _down(fminus_real, state), 0.0) - minus_imag = tl.where(keep_down, _down(fminus_imag, state), 0.0) - plus_real = tl.where(state == 0, minus_real, plus_real) - plus_imag = tl.where(state == 0, -minus_imag, plus_imag) - return plus_real, plus_imag, minus_real, minus_imag - - -@triton.jit -def _shift_real( - plus, - minus, - state, - state_mask, - state_count, -): - shifted_plus = tl.where((state > 0) & state_mask, _up(plus, state), 0.0) - shifted_minus = tl.where( - (state + 1 < state_count) & state_mask, _down(minus, state), 0.0 - ) - return tl.where(state == 0, -shifted_minus, shifted_plus), shifted_minus - - -# --------------------------------------------------------------------------- -# Dual complex arithmetic. -# -# A dual complex number is four planes: the real and imaginary parts of the -# value, then of the tangent. Triton has no structs, so every quantity travels -# as four separate registers and these helpers keep the bookkeeping in one -# place rather than spread through the kernel. -# --------------------------------------------------------------------------- - - -@triton.jit -def _damping(rate, dt, order): - """Longitudinal and transverse diffusion damping for one interval. - - ``rate`` already carries the sequence's gradient geometry, so an interval's - b-factor is that rate times its duration. Order zero has no longitudinal - weight, which is what keeps the recovery term undamped. - """ - b_factor = rate * dt - squared = order * order - return ( - tl.exp(-b_factor * squared), - tl.exp(-b_factor * (squared + order + 0.3333333333333333)), - ) - - -@triton.jit -def _flow(rate, dt, order): - """Phase each dephasing order turns through over one interval. - - ``rate`` already carries the sequence's gradient geometry, so it is the - winding per unit order per second: a longitudinal state at order l turns - through ``l * rate * dt``. The transverse states sit half an order further - along the gradient, which is where the extra half turn comes from. Order - zero is left alone while longitudinal, so the recovery term is unaffected. - """ - turn = rate * dt - return -order * turn, -(order + 0.5) * turn - - -@triton.jit -def _washout(rate, dt): - """The fraction of a voxel's spins that stay put over one interval. - - Inflowing spins are taken to be fully relaxed and unexcited, which makes - washout an affine map of the shape longitudinal recovery already has: - - wout * (Z * e1 + (1 - e1)) + win == Z * (e1 * wout) + (1 - e1 * wout) - - so scaling both relaxation factors by it carries the whole term. Clamped at - one, past which the interval has replaced the voxel outright. - """ - return 1.0 - tl.minimum(rate * dt, 1.0) - - -@triton.jit -def _two_pool_step(r1_free, r1_bound, exchange, bound, dt, attenuation): - """The two-pool longitudinal operator over one interval, and its recovery. - - ``expm((K - diag(R1)) t)`` in the exact 2x2 closed form. Its discriminant - is a square plus a product of two non-negative rates, so the root is real - and the branch a general exponential would need does not exist here. - ``sinh(d)/d`` is taken by series near the origin, where the root has no - derivative of its own. - - The equilibrium each pool relaxes toward is its own fraction, so the - recovery is ``(I - E1) (1 - f, f)`` and needs no solve. Returned as - ``(e11, e12, e21, e22, recovery_free, recovery_bound)``. - """ - free = 1.0 - bound - kab = exchange * bound - kba = exchange * free - l11 = (-kab - r1_free) * dt - l12 = kba * dt - l21 = kab * dt - l22 = (-kba - r1_bound) * dt - - half_trace = 0.5 * (l11 + l22) - half_gap = 0.5 * (l11 - l22) - square = half_gap * half_gap + l12 * l21 - # tau +/- d are the eigenvalues, both non-positive for a decaying system, - # so their exponentials are bounded by one. Formed that way rather than as - # e^tau cosh(d), which over a long interval is an underflow times an - # overflow. - root = tl.sqrt(tl.maximum(square, 0.0)) - upper = tl.exp(half_trace + root) - lower = tl.exp(half_trace - root) - cosine = 0.5 * (upper + lower) - turning = square > 1e-12 - guarded = tl.where(turning, root, 1.0) - scale = tl.where( - turning, - 0.5 * (upper - lower) / guarded, - tl.exp(half_trace) * (1.0 + square / 6.0 + square * square / 120.0), - ) - e11 = attenuation * (cosine + scale * half_gap) - e12 = attenuation * scale * l12 - e21 = attenuation * scale * l21 - e22 = attenuation * (cosine - scale * half_gap) - return ( - e11, - e12, - e21, - e22, - free - (e11 * free + e12 * bound), - bound - (e21 * free + e22 * bound), - ) - - -# The three-pool longitudinal step is the one operator the state machine forms -# in double. A 2x2's closed form loses accuracy like the interval; a 3x3's -# loses it like the square of it, because the answer's entries are order one -# while the terms that build them are order |L dt|^2. That is intrinsic to -# writing the answer as a polynomial in the generator, so it is met with -# precision rather than with rearrangement. It is formed once per interval, not -# once per dephasing order, so the cost stays out of the state loop. -_SPREAD_CUT = tl.constexpr(1.0) -_SINCH_CUT = tl.constexpr(1e-4) -# Where two roots meet the sorted roots have a vertical tangent the operator -# itself does not, so the arc cosine is held a hair off its endpoints. -_ARG_LIMIT = tl.constexpr(1.0 - 1e-16) -_TURN_THIRD = tl.constexpr(2.09439510239319549231) - - -@triton.jit -def _exp_difference(lower, upper, exp_lower, exp_upper): - """``[a, b] exp``, from exponentials the caller has already taken. - - Near the coalescence ``sinh(d)/d`` is even in the gap, so the series is a - polynomial in its square; the exponential of the midpoint is reached from - the lower one by a series too, because over a gap this small it is one. - """ - half = 0.5 * (upper - lower) - near = tl.abs(half) < _SINCH_CUT - square = half * half - # exp(mid) * sinh(half)/half, with both factors expanded about zero. - series = exp_lower * (1.0 + half + 0.5 * square) * (1.0 + square / 6.0) - gap = tl.where(near, 1.0, upper - lower) - return tl.where(near, series, (exp_upper - exp_lower) / gap) - - -@triton.jit -def _three_pool_recovery( - e00, e01, e02, e10, e11, e12, e20, e21, e22, free, pool_b, pool_c -): - """What each pool recovers over the interval, beside the operator itself. - - Returns the nine entries and the three recoveries, narrowed to float32 - once they are an operator. - """ - grow_free = free - (e00 * free + e01 * pool_b + e02 * pool_c) - grow_pool_b = pool_b - (e10 * free + e11 * pool_b + e12 * pool_c) - grow_bound = pool_c - (e20 * free + e21 * pool_b + e22 * pool_c) - return ( - e00.to(tl.float32), - e01.to(tl.float32), - e02.to(tl.float32), - e10.to(tl.float32), - e11.to(tl.float32), - e12.to(tl.float32), - e20.to(tl.float32), - e21.to(tl.float32), - e22.to(tl.float32), - grow_free.to(tl.float32), - grow_pool_b.to(tl.float32), - grow_bound.to(tl.float32), - ) - - -@triton.jit -def _three_pool_step( - r1_free, - r1_pool_b, - r1_bound, - exchange_b, - exchange_c, - fraction_b, - fraction_c, - dt, - attenuation, - narrow: tl.constexpr = False, -): - """``expm((K - diag(R1)) t)`` for free water beside both second pools. - - Free water is pool a, the chemically exchanging pool b and the semisolid - pool c; each second pool exchanges with the free water and not with the - other. Returns the nine entries and the three recoveries, narrowed to - float32 once they are an operator. - - Two branches, by how far apart the eigenvalues are. Where they are close - the exponential's own series is reduced modulo the characteristic - polynomial, which forms no root at all. Where they are far apart the - interpolating polynomial is taken in Newton form at the three roots, each - of which is non-positive, so a long interval cannot overflow. - - ``narrow`` says the caller has bounded the spread below - :data:`blochsim.sequence._parameters.NARROW_SPREAD` for every voxel and - every interval it will pass, so - only the series can be reached. The roots then cost nothing, and the series - holds the answer to float32 without being carried in double -- - :func:`blochsim.sequence._parameters.narrow_three_pool` is what decides it. - """ - work: tl.constexpr = tl.float32 if narrow else tl.float64 - terms: tl.constexpr = 24 if narrow else 16 - step = dt.to(work) - free = (1.0 - fraction_b - fraction_c).to(work) - pool_b = fraction_b.to(work) - pool_c = fraction_c.to(work) - kab = exchange_b.to(work) * pool_b - kba = exchange_b.to(work) * free - kac = exchange_c.to(work) * pool_c - kca = exchange_c.to(work) * free - a00 = (-kab - kac - r1_free.to(work)) * step - a01 = kba * step - a02 = kca * step - a10 = kab * step - a11 = (-kba - r1_pool_b.to(work)) * step - a20 = kac * step - a22 = (-kca - r1_bound.to(work)) * step - - third = (a00 + a11 + a22) / 3.0 - s00 = a00 - third - s11 = a11 - third - s22 = a22 - third - # The two second pools do not exchange, so the generator keeps a pair of - # structural zeros the products below are written around. - minors = s00 * s11 - a01 * a10 + s00 * s22 - a02 * a20 + s11 * s22 - determinant = s00 * s11 * s22 - a01 * (a10 * s22) + a02 * (-s11 * a20) - - # --- close together: the series reduced modulo x^3 + minors x - det --- - flat = 1.0 + 0.0 * third - linear = 0.0 * third - square = 0.0 * third - sum_flat = flat - sum_linear = linear - sum_square = square - factorial = 1.0 - for order in tl.static_range(1, terms): - next_flat = square * determinant - next_linear = flat - square * minors - next_square = linear - flat = next_flat - linear = next_linear - square = next_square - factorial = factorial * order - weight = 1.0 / factorial - sum_flat = sum_flat + weight * flat - sum_linear = sum_linear + weight * linear - sum_square = sum_square + weight * square - q00 = s00 * s00 + a01 * a10 + a02 * a20 - q01 = s00 * a01 + a01 * s11 - q02 = s00 * a02 + a02 * s22 - q10 = a10 * s00 + s11 * a10 - q11 = a10 * a01 + s11 * s11 - q12 = a10 * a02 - q20 = a20 * s00 + s22 * a20 - q21 = a20 * a01 - q22 = a20 * a02 + s22 * s22 - lift = tl.exp(third) - c00 = lift * (sum_flat + sum_linear * s00 + sum_square * q00) - c01 = lift * (sum_linear * a01 + sum_square * q01) - c02 = lift * (sum_linear * a02 + sum_square * q02) - c10 = lift * (sum_linear * a10 + sum_square * q10) - c11 = lift * (sum_flat + sum_linear * s11 + sum_square * q11) - c12 = lift * (sum_square * q12) - c20 = lift * (sum_linear * a20 + sum_square * q20) - c21 = lift * (sum_square * q21) - c22 = lift * (sum_flat + sum_linear * s22 + sum_square * q22) - - # --- far apart: the Newton form at the three roots --- - damp = attenuation.to(work) - if narrow: - e00 = damp * c00 - e01 = damp * c01 - e02 = damp * c02 - e10 = damp * c10 - e11 = damp * c11 - e12 = damp * c12 - e20 = damp * c20 - e21 = damp * c21 - e22 = damp * c22 - return _three_pool_recovery( - e00, e01, e02, e10, e11, e12, e20, e21, e22, free, pool_b, pool_c - ) - radius = tl.sqrt(tl.maximum(-minors * (1.0 / 3.0), 1e-300)) - argument = tl.minimum( - tl.maximum(0.5 * determinant / (radius * radius * radius), -_ARG_LIMIT), - _ARG_LIMIT, - ) - angle = libdevice.acos(argument) / 3.0 - root_a = 2.0 * radius * tl.cos(angle) + third - root_b = 2.0 * radius * tl.cos(angle - _TURN_THIRD) + third - root_c = 2.0 * radius * tl.cos(angle - 2.0 * _TURN_THIRD) + third - low = tl.minimum(tl.minimum(root_a, root_b), root_c) - high = tl.maximum(tl.maximum(root_a, root_b), root_c) - middle = tl.maximum( - tl.minimum(root_a, root_b), tl.minimum(tl.maximum(root_a, root_b), root_c) - ) - # Three exponentials serve every divided difference between them. - leading = tl.exp(low) - centre = tl.exp(middle) - trailing = tl.exp(high) - first = _exp_difference(low, middle, leading, centre) - span = high - low - second = (_exp_difference(middle, high, centre, trailing) - first) / tl.where( - span > 0.0, span, 1.0 - ) - m00 = a00 - low - m11 = a11 - low - m22 = a22 - low - n00 = a00 - middle - n11 = a11 - middle - n22 = a22 - middle - p00 = m00 * n00 + a01 * a10 + a02 * a20 - p01 = m00 * a01 + a01 * n11 - p02 = m00 * a02 + a02 * n22 - p10 = a10 * n00 + m11 * a10 - p11 = a10 * a01 + m11 * n11 - p12 = a10 * a02 - p20 = a20 * n00 + m22 * a20 - p21 = a20 * a01 - p22 = a20 * a02 + m22 * n22 - d00 = leading + first * m00 + second * p00 - d01 = first * a01 + second * p01 - d02 = first * a02 + second * p02 - d10 = first * a10 + second * p10 - d11 = leading + first * m11 + second * p11 - d12 = second * p12 - d20 = first * a20 + second * p20 - d21 = second * p21 - d22 = leading + first * m22 + second * p22 - - # The shifted roots sum to zero, so the sum of their squares is -2 * minors - # and none is larger than the root of that. - close = -2.0 * minors < _SPREAD_CUT * _SPREAD_CUT - e00 = damp * tl.where(close, c00, d00) - e01 = damp * tl.where(close, c01, d01) - e02 = damp * tl.where(close, c02, d02) - e10 = damp * tl.where(close, c10, d10) - e11 = damp * tl.where(close, c11, d11) - e12 = damp * tl.where(close, c12, d12) - e20 = damp * tl.where(close, c20, d20) - e21 = damp * tl.where(close, c21, d21) - e22 = damp * tl.where(close, c22, d22) - return _three_pool_recovery( - e00, e01, e02, e10, e11, e12, e20, e21, e22, free, pool_b, pool_c - ) - - -@triton.jit -def _three_pool_from_table( - table, row, atom, voxel_count, mask, attenuation, free, pool_b, pool_c -): - """Read one interval's three-pool operator, and what each pool recovers. - - The stored row is undamped, so the washout the event carries is applied - here and the three recoveries follow from the damped entries -- which is - what makes one row serve every event of the same length whatever its - washout. - """ - base = table + row * (9 * voxel_count) + atom - return _three_pool_recovery( - attenuation * tl.load(base + 0 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 1 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 2 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 3 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 4 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 5 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 6 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 7 * voxel_count, mask=mask, other=0.0), - attenuation * tl.load(base + 8 * voxel_count, mask=mask, other=0.0), - free, - pool_b, - pool_c, - ) - - -@triton.jit -def _three_pool_from_table_jvp( - table, - row, - atom, - voxel_count, - mask, - r1_free, - r1_pool_b, - r1_bound, - exchange_b, - exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - d_dt, - attenuation, - d_attenuation, -): - """Read one interval's three-pool operator and a direction through it. - - The row carries the tissue's share of the direction; the interval's own is - ``A1 C d_dt``, formed here because ``d_dt`` belongs to the event and the - row is shared. Returns the nine entries and three recoveries with their - tangents, in the order :func:`_three_pool_step_jvp` returns them. - """ - free = 1.0 - fraction_b - fraction_c - d_free = -d_fraction_b - d_fraction_c - a00 = -exchange_b * fraction_b - exchange_c * fraction_c - r1_free - a01 = exchange_b * free - a02 = exchange_c * free - a10 = exchange_b * fraction_b - a11 = -exchange_b * free - r1_pool_b - a20 = exchange_c * fraction_c - a22 = -exchange_c * free - r1_bound - base = table + row * (18 * voxel_count) + atom - c00 = tl.load(base + 0 * voxel_count, mask=mask, other=0.0) - c01 = tl.load(base + 1 * voxel_count, mask=mask, other=0.0) - c02 = tl.load(base + 2 * voxel_count, mask=mask, other=0.0) - c10 = tl.load(base + 3 * voxel_count, mask=mask, other=0.0) - c11 = tl.load(base + 4 * voxel_count, mask=mask, other=0.0) - c12 = tl.load(base + 5 * voxel_count, mask=mask, other=0.0) - c20 = tl.load(base + 6 * voxel_count, mask=mask, other=0.0) - c21 = tl.load(base + 7 * voxel_count, mask=mask, other=0.0) - c22 = tl.load(base + 8 * voxel_count, mask=mask, other=0.0) - # The row's tangent, plus what the event's own interval direction adds. - t00 = tl.load(base + 9 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a00 * c00 + a01 * c10 + a02 * c20 - ) - t01 = tl.load(base + 10 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a00 * c01 + a01 * c11 + a02 * c21 - ) - t02 = tl.load(base + 11 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a00 * c02 + a01 * c12 + a02 * c22 - ) - t10 = tl.load(base + 12 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a10 * c00 + a11 * c10 - ) - t11 = tl.load(base + 13 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a10 * c01 + a11 * c11 - ) - t12 = tl.load(base + 14 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a10 * c02 + a11 * c12 - ) - t20 = tl.load(base + 15 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a20 * c00 + a22 * c20 - ) - t21 = tl.load(base + 16 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a20 * c01 + a22 * c21 - ) - t22 = tl.load(base + 17 * voxel_count, mask=mask, other=0.0) + d_dt * ( - a20 * c02 + a22 * c22 - ) - e00 = attenuation * c00 - e01 = attenuation * c01 - e02 = attenuation * c02 - e10 = attenuation * c10 - e11 = attenuation * c11 - e12 = attenuation * c12 - e20 = attenuation * c20 - e21 = attenuation * c21 - e22 = attenuation * c22 - f00 = d_attenuation * c00 + attenuation * t00 - f01 = d_attenuation * c01 + attenuation * t01 - f02 = d_attenuation * c02 + attenuation * t02 - f10 = d_attenuation * c10 + attenuation * t10 - f11 = d_attenuation * c11 + attenuation * t11 - f12 = d_attenuation * c12 + attenuation * t12 - f20 = d_attenuation * c20 + attenuation * t20 - f21 = d_attenuation * c21 + attenuation * t21 - f22 = d_attenuation * c22 + attenuation * t22 - # The equilibrium the recoveries are taken against moves with the - # fractions, so it carries a direction of its own. - grow_free = free - (e00 * free + e01 * fraction_b + e02 * fraction_c) - grow_pool_b = fraction_b - (e10 * free + e11 * fraction_b + e12 * fraction_c) - grow_bound = fraction_c - (e20 * free + e21 * fraction_b + e22 * fraction_c) - d_grow_free = d_free - ( - f00 * free - + f01 * fraction_b - + f02 * fraction_c - + e00 * d_free - + e01 * d_fraction_b - + e02 * d_fraction_c - ) - d_grow_pool_b = d_fraction_b - ( - f10 * free - + f11 * fraction_b - + f12 * fraction_c - + e10 * d_free - + e11 * d_fraction_b - + e12 * d_fraction_c - ) - d_grow_bound = d_fraction_c - ( - f20 * free - + f21 * fraction_b - + f22 * fraction_c - + e20 * d_free - + e21 * d_fraction_b - + e22 * d_fraction_c - ) - return ( - e00, - e01, - e02, - e10, - e11, - e12, - e20, - e21, - e22, - grow_free, - grow_pool_b, - grow_bound, - f00, - f01, - f02, - f10, - f11, - f12, - f20, - f21, - f22, - d_grow_free, - d_grow_pool_b, - d_grow_bound, - ) - - -@triton.jit -def _three_pool_table_jvp_kernel( - t1, - t1_pool_b, - t1_bound, - pool_b_exchange, - bound_exchange, - pool_b_fraction, - bound_fraction, - d_t1, - d_t1_pool_b, - d_t1_bound, - d_pool_b_exchange, - d_bound_exchange, - d_pool_b_fraction, - d_bound_fraction, - durations, - rows, - table, - voxel_count, - BLOCK: tl.constexpr, - narrow: tl.constexpr, -): - """Fill one row of the three-pool operator table, value and direction. - - The row holds nine undamped entries and the nine a direction through the - tissue gives them, both at ``d_dt`` of zero -- the interval's own share of - the direction is ``A1 C d_dt``, which the reading event adds because - ``d_dt`` is its own and the row's is not. Laid out ``(rows, 18, voxels)``, - the tangent following the value. - """ - row = tl.load(rows + tl.program_id(0)) - atom = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) - live = atom < voxel_count - dt = tl.load(durations + row) - nil = 0.0 * dt - value_t1 = tl.load(t1 + atom, mask=live, other=1.0) - value_t1b = tl.load(t1_pool_b + atom, mask=live, other=1.0) - value_t1c = tl.load(t1_bound + atom, mask=live, other=1.0) - ( - e00, - e01, - e02, - e10, - e11, - e12, - e20, - e21, - e22, - _, - _, - _, - d00, - d01, - d02, - d10, - d11, - d12, - d20, - d21, - d22, - _, - _, - _, - ) = _three_pool_step_jvp( - 1000.0 / value_t1, - -1000.0 * tl.load(d_t1 + atom, mask=live, other=0.0) / (value_t1 * value_t1), - 1000.0 / value_t1b, - -1000.0 - * tl.load(d_t1_pool_b + atom, mask=live, other=0.0) - / (value_t1b * value_t1b), - 1000.0 / value_t1c, - -1000.0 - * tl.load(d_t1_bound + atom, mask=live, other=0.0) - / (value_t1c * value_t1c), - tl.load(pool_b_exchange + atom, mask=live, other=0.0), - tl.load(d_pool_b_exchange + atom, mask=live, other=0.0), - tl.load(bound_exchange + atom, mask=live, other=0.0), - tl.load(d_bound_exchange + atom, mask=live, other=0.0), - tl.load(pool_b_fraction + atom, mask=live, other=0.0), - tl.load(d_pool_b_fraction + atom, mask=live, other=0.0), - tl.load(bound_fraction + atom, mask=live, other=0.0), - tl.load(d_bound_fraction + atom, mask=live, other=0.0), - dt, - nil, - 1.0 + nil, - nil, - narrow, - ) - base = table + row * (18 * voxel_count) + atom - tl.store(base + 0 * voxel_count, e00, mask=live) - tl.store(base + 1 * voxel_count, e01, mask=live) - tl.store(base + 2 * voxel_count, e02, mask=live) - tl.store(base + 3 * voxel_count, e10, mask=live) - tl.store(base + 4 * voxel_count, e11, mask=live) - tl.store(base + 5 * voxel_count, e12, mask=live) - tl.store(base + 6 * voxel_count, e20, mask=live) - tl.store(base + 7 * voxel_count, e21, mask=live) - tl.store(base + 8 * voxel_count, e22, mask=live) - tl.store(base + 9 * voxel_count, d00, mask=live) - tl.store(base + 10 * voxel_count, d01, mask=live) - tl.store(base + 11 * voxel_count, d02, mask=live) - tl.store(base + 12 * voxel_count, d10, mask=live) - tl.store(base + 13 * voxel_count, d11, mask=live) - tl.store(base + 14 * voxel_count, d12, mask=live) - tl.store(base + 15 * voxel_count, d20, mask=live) - tl.store(base + 16 * voxel_count, d21, mask=live) - tl.store(base + 17 * voxel_count, d22, mask=live) - - -@triton.jit -def _three_pool_contract( - x00, - x01, - x02, - x10, - x11, - x12, - x20, - x21, - x22, - e11, - e12, - e13, - e21, - e22, - e23, - e31, - e32, - e33, - rec_free, - rec_pool_b, - rec_bound, - free, - fraction_b, - fraction_c, -): - """Nine cotangents against nine entries, less what the recoveries take. - - A recovery is ``m - E m`` with ``m`` the equilibrium - ``(free, fraction_b, fraction_c)``, so it differentiates through the same - nine entries with the equilibrium contracted out of them. - """ - return ( - e11 * x00 - + e12 * x01 - + e13 * x02 - + e21 * x10 - + e22 * x11 - + e23 * x12 - + e31 * x20 - + e32 * x21 - + e33 * x22 - - rec_free * (x00 * free + x01 * fraction_b + x02 * fraction_c) - - rec_pool_b * (x10 * free + x11 * fraction_b + x12 * fraction_c) - - rec_bound * (x20 * free + x21 * fraction_b + x22 * fraction_c) - ) - - -@triton.jit -def _three_pool_interval_adjoint( - table, - row, - atom, - voxel_count, - mask, - r1_free, - r1_pool_b, - r1_bound, - exchange_b, - exchange_c, - fraction_b, - fraction_c, - attenuation, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, -): - """What one interval's cotangents give its length and its attenuation. - - The generator is proportional to the interval, so ``dE/d(dt) == A1 E`` and - the length's gradient needs the operator and 27 multiplies rather than the - eigenvalues -- which is what lets every other gradient be pooled over the - events that share a length while this one stays per event. - - Returns the two contractions, in the order - :func:`_three_pool_step_adjoint_jvp` returns them. - """ - free = 1.0 - fraction_b - fraction_c - a00 = -exchange_b * fraction_b - exchange_c * fraction_c - r1_free - a01 = exchange_b * free - a02 = exchange_c * free - a10 = exchange_b * fraction_b - a11 = -exchange_b * free - r1_pool_b - a20 = exchange_c * fraction_c - a22 = -exchange_c * free - r1_bound - base = table + row * (9 * voxel_count) + atom - c00 = tl.load(base + 0 * voxel_count, mask=mask, other=0.0) - c01 = tl.load(base + 1 * voxel_count, mask=mask, other=0.0) - c02 = tl.load(base + 2 * voxel_count, mask=mask, other=0.0) - c10 = tl.load(base + 3 * voxel_count, mask=mask, other=0.0) - c11 = tl.load(base + 4 * voxel_count, mask=mask, other=0.0) - c12 = tl.load(base + 5 * voxel_count, mask=mask, other=0.0) - c20 = tl.load(base + 6 * voxel_count, mask=mask, other=0.0) - c21 = tl.load(base + 7 * voxel_count, mask=mask, other=0.0) - c22 = tl.load(base + 8 * voxel_count, mask=mask, other=0.0) - # A1 C, the second and third rows of A1 having no entry off their own pool. - p00 = a00 * c00 + a01 * c10 + a02 * c20 - p01 = a00 * c01 + a01 * c11 + a02 * c21 - p02 = a00 * c02 + a01 * c12 + a02 * c22 - p10 = a10 * c00 + a11 * c10 - p11 = a10 * c01 + a11 * c11 - p12 = a10 * c02 + a11 * c12 - p20 = a20 * c00 + a22 * c20 - p21 = a20 * c01 + a22 * c21 - p22 = a20 * c02 + a22 * c22 - grad_dt = attenuation * _three_pool_contract( - p00, - p01, - p02, - p10, - p11, - p12, - p20, - p21, - p22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - grad_att = _three_pool_contract( - c00, - c01, - c02, - c10, - c11, - c12, - c20, - c21, - c22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - return grad_dt, grad_att - - -@triton.jit -def _three_pool_interval_adjoint_jvp( - table, - row, - atom, - voxel_count, - mask, - r1_free, - d_r1_free, - r1_pool_b, - d_r1_pool_b, - r1_bound, - d_r1_bound, - exchange_b, - d_exchange_b, - exchange_c, - d_exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - d_dt, - attenuation, - d_attenuation, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - tfree, - tpool_b, - tbound, -): - """The interval and the attenuation, from a dual pair's cotangents. - - ``dE/d(dt)`` is ``A1 E``, and the direction that quantity carries follows - from the same generator: with ``C_dot == C_row + A1 C d_dt``, the - derivative in the interval is ``A1_dot C + A1 C_row + A1 A1 C d_dt``. So - three products of the generator against the tabulated row serve what the - eigenvalues would otherwise be re-formed for, and these two quantities are - the only ones that stay per event. - - Returns the interval and attenuation gradients, value then tangent, in the - order :func:`_three_pool_step_adjoint_jvp` returns them. - """ - free = 1.0 - fraction_b - fraction_c - d_free = -d_fraction_b - d_fraction_c - a00 = -exchange_b * fraction_b - exchange_c * fraction_c - r1_free - a01 = exchange_b * free - a02 = exchange_c * free - a10 = exchange_b * fraction_b - a11 = -exchange_b * free - r1_pool_b - a20 = exchange_c * fraction_c - a22 = -exchange_c * free - r1_bound - da00 = ( - -d_exchange_b * fraction_b - - exchange_b * d_fraction_b - - d_exchange_c * fraction_c - - exchange_c * d_fraction_c - - d_r1_free - ) - da01 = d_exchange_b * free + exchange_b * d_free - da02 = d_exchange_c * free + exchange_c * d_free - da10 = d_exchange_b * fraction_b + exchange_b * d_fraction_b - da11 = -d_exchange_b * free - exchange_b * d_free - d_r1_pool_b - da20 = d_exchange_c * fraction_c + exchange_c * d_fraction_c - da22 = -d_exchange_c * free - exchange_c * d_free - d_r1_bound - base = table + row * (18 * voxel_count) + atom - c00 = tl.load(base + 0 * voxel_count, mask=mask, other=0.0) - c01 = tl.load(base + 1 * voxel_count, mask=mask, other=0.0) - c02 = tl.load(base + 2 * voxel_count, mask=mask, other=0.0) - c10 = tl.load(base + 3 * voxel_count, mask=mask, other=0.0) - c11 = tl.load(base + 4 * voxel_count, mask=mask, other=0.0) - c12 = tl.load(base + 5 * voxel_count, mask=mask, other=0.0) - c20 = tl.load(base + 6 * voxel_count, mask=mask, other=0.0) - c21 = tl.load(base + 7 * voxel_count, mask=mask, other=0.0) - c22 = tl.load(base + 8 * voxel_count, mask=mask, other=0.0) - r00 = tl.load(base + 9 * voxel_count, mask=mask, other=0.0) - r01 = tl.load(base + 10 * voxel_count, mask=mask, other=0.0) - r02 = tl.load(base + 11 * voxel_count, mask=mask, other=0.0) - r10 = tl.load(base + 12 * voxel_count, mask=mask, other=0.0) - r11 = tl.load(base + 13 * voxel_count, mask=mask, other=0.0) - r12 = tl.load(base + 14 * voxel_count, mask=mask, other=0.0) - r20 = tl.load(base + 15 * voxel_count, mask=mask, other=0.0) - r21 = tl.load(base + 16 * voxel_count, mask=mask, other=0.0) - r22 = tl.load(base + 17 * voxel_count, mask=mask, other=0.0) - # P = A1 C, Q = A1_dot C + A1 C_row, S = A1 P. - p00 = a00 * c00 + a01 * c10 + a02 * c20 - p01 = a00 * c01 + a01 * c11 + a02 * c21 - p02 = a00 * c02 + a01 * c12 + a02 * c22 - p10 = a10 * c00 + a11 * c10 - p11 = a10 * c01 + a11 * c11 - p12 = a10 * c02 + a11 * c12 - p20 = a20 * c00 + a22 * c20 - p21 = a20 * c01 + a22 * c21 - p22 = a20 * c02 + a22 * c22 - q00 = da00 * c00 + da01 * c10 + da02 * c20 + a00 * r00 + a01 * r10 + a02 * r20 - q01 = da00 * c01 + da01 * c11 + da02 * c21 + a00 * r01 + a01 * r11 + a02 * r21 - q02 = da00 * c02 + da01 * c12 + da02 * c22 + a00 * r02 + a01 * r12 + a02 * r22 - q10 = da10 * c00 + da11 * c10 + a10 * r00 + a11 * r10 - q11 = da10 * c01 + da11 * c11 + a10 * r01 + a11 * r11 - q12 = da10 * c02 + da11 * c12 + a10 * r02 + a11 * r12 - q20 = da20 * c00 + da22 * c20 + a20 * r00 + a22 * r20 - q21 = da20 * c01 + da22 * c21 + a20 * r01 + a22 * r21 - q22 = da20 * c02 + da22 * c22 + a20 * r02 + a22 * r22 - s00 = a00 * p00 + a01 * p10 + a02 * p20 - s01 = a00 * p01 + a01 * p11 + a02 * p21 - s02 = a00 * p02 + a01 * p12 + a02 * p22 - s10 = a10 * p00 + a11 * p10 - s11 = a10 * p01 + a11 * p11 - s12 = a10 * p02 + a11 * p12 - s20 = a20 * p00 + a22 * p20 - s21 = a20 * p01 + a22 * p21 - s22 = a20 * p02 + a22 * p22 - # The direction the tabulated operator carries, and the interval's own - # share of it. - d00 = r00 + p00 * d_dt - d01 = r01 + p01 * d_dt - d02 = r02 + p02 * d_dt - d10 = r10 + p10 * d_dt - d11 = r11 + p11 * d_dt - d12 = r12 + p12 * d_dt - d20 = r20 + p20 * d_dt - d21 = r21 + p21 * d_dt - d22 = r22 + p22 * d_dt - # dE/d(dt) with the attenuation held, and the direction that carries. - g00 = attenuation * p00 - g01 = attenuation * p01 - g02 = attenuation * p02 - g10 = attenuation * p10 - g11 = attenuation * p11 - g12 = attenuation * p12 - g20 = attenuation * p20 - g21 = attenuation * p21 - g22 = attenuation * p22 - w00 = d_attenuation * p00 + attenuation * (q00 + s00 * d_dt) - w01 = d_attenuation * p01 + attenuation * (q01 + s01 * d_dt) - w02 = d_attenuation * p02 + attenuation * (q02 + s02 * d_dt) - w10 = d_attenuation * p10 + attenuation * (q10 + s10 * d_dt) - w11 = d_attenuation * p11 + attenuation * (q11 + s11 * d_dt) - w12 = d_attenuation * p12 + attenuation * (q12 + s12 * d_dt) - w20 = d_attenuation * p20 + attenuation * (q20 + s20 * d_dt) - w21 = d_attenuation * p21 + attenuation * (q21 + s21 * d_dt) - w22 = d_attenuation * p22 + attenuation * (q22 + s22 * d_dt) - # This kernel carries every quantity as a dual pair, so the tangent - # returned beside a gradient is that gradient's own directional - # derivative -- not the gradient with respect to the direction. - grad_dt_v = _three_pool_contract( - g00, - g01, - g02, - g10, - g11, - g12, - g20, - g21, - g22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - grad_dt_t = ( - _three_pool_contract( - g00, - g01, - g02, - g10, - g11, - g12, - g20, - g21, - g22, - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - tfree, - tpool_b, - tbound, - free, - fraction_b, - fraction_c, - ) - + _three_pool_contract( - w00, - w01, - w02, - w10, - w11, - w12, - w20, - w21, - w22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - - ( - bfree * (g00 * d_free + g01 * d_fraction_b + g02 * d_fraction_c) - + bpool_b * (g10 * d_free + g11 * d_fraction_b + g12 * d_fraction_c) - + bbound * (g20 * d_free + g21 * d_fraction_b + g22 * d_fraction_c) - ) - ) - grad_att_v = _three_pool_contract( - c00, - c01, - c02, - c10, - c11, - c12, - c20, - c21, - c22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - grad_att_t = ( - _three_pool_contract( - c00, - c01, - c02, - c10, - c11, - c12, - c20, - c21, - c22, - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - tfree, - tpool_b, - tbound, - free, - fraction_b, - fraction_c, - ) - + _three_pool_contract( - d00, - d01, - d02, - d10, - d11, - d12, - d20, - d21, - d22, - b11, - b12, - b13, - b21, - b22, - b23, - b31, - b32, - b33, - bfree, - bpool_b, - bbound, - free, - fraction_b, - fraction_c, - ) - - ( - bfree * (c00 * d_free + c01 * d_fraction_b + c02 * d_fraction_c) - + bpool_b * (c10 * d_free + c11 * d_fraction_b + c12 * d_fraction_c) - + bbound * (c20 * d_free + c21 * d_fraction_b + c22 * d_fraction_c) - ) - ) - return grad_dt_v, grad_att_v, grad_dt_t, grad_att_t - - -@triton.jit -def _three_pool_table_kernel( - t1, - t1_pool_b, - t1_bound, - pool_b_exchange, - bound_exchange, - pool_b_fraction, - bound_fraction, - durations, - rows, - table, - voxel_count, - BLOCK: tl.constexpr, - narrow: tl.constexpr, -): - """Fill one row of the three-pool operator table. - - The row is ``expm((K - diag(R1)) dt)`` with no washout applied, so the - event that reads it supplies its own attenuation. Laid out - ``(rows, 9, voxels)`` -- entry-major over the voxel axis -- so the nine - loads an event makes are each coalesced. - """ - row = tl.load(rows + tl.program_id(0)) - atom = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) - live = atom < voxel_count - dt = tl.load(durations + row) - fraction_b = tl.load(pool_b_fraction + atom, mask=live, other=0.0) - fraction_c = tl.load(bound_fraction + atom, mask=live, other=0.0) - ( - e00, - e01, - e02, - e10, - e11, - e12, - e20, - e21, - e22, - _, - _, - _, - ) = _three_pool_step( - 1000.0 / tl.load(t1 + atom, mask=live, other=1.0), - 1000.0 / tl.load(t1_pool_b + atom, mask=live, other=1.0), - 1000.0 / tl.load(t1_bound + atom, mask=live, other=1.0), - tl.load(pool_b_exchange + atom, mask=live, other=0.0), - tl.load(bound_exchange + atom, mask=live, other=0.0), - fraction_b, - fraction_c, - dt, - # The helpers narrow their arguments, so the attenuation an undamped - # row wants has to arrive as a tensor rather than a literal one. - 1.0 + 0.0 * dt, - narrow, - ) - base = table + row * (9 * voxel_count) + atom - tl.store(base + 0 * voxel_count, e00, mask=live) - tl.store(base + 1 * voxel_count, e01, mask=live) - tl.store(base + 2 * voxel_count, e02, mask=live) - tl.store(base + 3 * voxel_count, e10, mask=live) - tl.store(base + 4 * voxel_count, e11, mask=live) - tl.store(base + 5 * voxel_count, e12, mask=live) - tl.store(base + 6 * voxel_count, e20, mask=live) - tl.store(base + 7 * voxel_count, e21, mask=live) - tl.store(base + 8 * voxel_count, e22, mask=live) - - -@triton.jit -def _exp_difference_jvp(lower, d_lower, upper, d_upper, low_exp, high_exp): - """``[a, b] exp`` and its directional derivative. - - The derivative of a divided difference is the next one along, - ``d/da [a,b] = [a,a,b]``, which near the coalescence is again a series in - the gap's square rather than a quotient that vanishes over a vanishing - denominator. - """ - half = 0.5 * (upper - lower) - d_half = 0.5 * (d_upper - d_lower) - near = tl.abs(half) < _SINCH_CUT - square = half * half - d_square = 2.0 * half * d_half - # exp(mid) * sinh(half)/half, both factors expanded about zero. - lift = low_exp * (1.0 + half + 0.5 * square) - d_lift = low_exp * (d_lower * (1.0 + half + 0.5 * square) + d_half + half * d_half) - sinch = 1.0 + square / 6.0 + square * square / 120.0 - d_sinch = d_square / 6.0 + 2.0 * square * d_square / 120.0 - series = lift * sinch - d_series = d_lift * sinch + lift * d_sinch - gap = tl.where(near, 1.0, upper - lower) - d_gap = tl.where(near, 0.0, d_upper - d_lower) - quotient = (high_exp - low_exp) / gap - d_quotient = (high_exp * d_upper - low_exp * d_lower - quotient * d_gap) / gap - return tl.where(near, series, quotient), tl.where(near, d_series, d_quotient) - - -@triton.jit -def _three_pool_pieces_jvp( - r1_free, - d_r1_free, - r1_pool_b, - d_r1_pool_b, - r1_bound, - d_r1_bound, - exchange_b, - d_exchange_b, - exchange_c, - d_exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - dt, - d_dt, - narrow: tl.constexpr = False, -): - """The three-pool operator's shared front half, as duals. - - The generator, its two invariants, the series coefficients and the three - roots with the divided differences between them -- everything both the - operator and its reverse sweep are assembled from, computed once in double - so the two cannot drift apart. - """ - work: tl.constexpr = tl.float32 if narrow else tl.float64 - step = dt.to(work) - d_step = d_dt.to(work) - free = (1.0 - fraction_b - fraction_c).to(work) - d_free = (-d_fraction_b - d_fraction_c).to(work) - pool_b = fraction_b.to(work) - d_pool_b = d_fraction_b.to(work) - pool_c = fraction_c.to(work) - d_pool_c = d_fraction_c.to(work) - rate_b = exchange_b.to(work) - d_rate_b = d_exchange_b.to(work) - rate_c = exchange_c.to(work) - d_rate_c = d_exchange_c.to(work) - kab = rate_b * pool_b - d_kab = d_rate_b * pool_b + rate_b * d_pool_b - kba = rate_b * free - d_kba = d_rate_b * free + rate_b * d_free - kac = rate_c * pool_c - d_kac = d_rate_c * pool_c + rate_c * d_pool_c - kca = rate_c * free - d_kca = d_rate_c * free + rate_c * d_free - row_a = -kab - kac - r1_free.to(work) - d_row_a = -d_kab - d_kac - d_r1_free.to(work) - row_b = -kba - r1_pool_b.to(work) - d_row_b = -d_kba - d_r1_pool_b.to(work) - row_c = -kca - r1_bound.to(work) - d_row_c = -d_kca - d_r1_bound.to(work) - a00 = row_a * step - d_a00 = d_row_a * step + row_a * d_step - a01 = kba * step - d_a01 = d_kba * step + kba * d_step - a02 = kca * step - d_a02 = d_kca * step + kca * d_step - a10 = kab * step - d_a10 = d_kab * step + kab * d_step - a11 = row_b * step - d_a11 = d_row_b * step + row_b * d_step - a20 = kac * step - d_a20 = d_kac * step + kac * d_step - a22 = row_c * step - d_a22 = d_row_c * step + row_c * d_step - - third = (a00 + a11 + a22) / 3.0 - d_third = (d_a00 + d_a11 + d_a22) / 3.0 - s00 = a00 - third - d_s00 = d_a00 - d_third - s11 = a11 - third - d_s11 = d_a11 - d_third - s22 = a22 - third - d_s22 = d_a22 - d_third - minors = s00 * s11 - a01 * a10 + s00 * s22 - a02 * a20 + s11 * s22 - d_minors = ( - d_s00 * s11 - + s00 * d_s11 - - d_a01 * a10 - - a01 * d_a10 - + d_s00 * s22 - + s00 * d_s22 - - d_a02 * a20 - - a02 * d_a20 - + d_s11 * s22 - + s11 * d_s22 - ) - determinant = s00 * s11 * s22 - a01 * (a10 * s22) + a02 * (-s11 * a20) - d_determinant = ( - d_s00 * s11 * s22 - + s00 * d_s11 * s22 - + s00 * s11 * d_s22 - - d_a01 * a10 * s22 - - a01 * d_a10 * s22 - - a01 * a10 * d_s22 - - d_a02 * s11 * a20 - - a02 * d_s11 * a20 - - a02 * s11 * d_a20 - ) - - # --- close together: the series reduced modulo x^3 + minors x - det --- - flat = 1.0 + 0.0 * third - linear = 0.0 * third - square = 0.0 * third - d_flat = 0.0 * third - d_linear = 0.0 * third - d_square = 0.0 * third - sum_flat = flat - sum_linear = linear - sum_square = square - d_sum_flat = d_flat - d_sum_linear = d_linear - d_sum_square = d_square - factorial = 1.0 - for order in tl.static_range(1, 16): - next_flat = square * determinant - d_next_flat = d_square * determinant + square * d_determinant - next_linear = flat - square * minors - d_next_linear = d_flat - d_square * minors - square * d_minors - next_square = linear - d_next_square = d_linear - flat = next_flat - linear = next_linear - square = next_square - d_flat = d_next_flat - d_linear = d_next_linear - d_square = d_next_square - factorial = factorial * order - weight = 1.0 / factorial - sum_flat = sum_flat + weight * flat - sum_linear = sum_linear + weight * linear - sum_square = sum_square + weight * square - d_sum_flat = d_sum_flat + weight * d_flat - d_sum_linear = d_sum_linear + weight * d_linear - d_sum_square = d_sum_square + weight * d_square - lift = tl.exp(third) - d_lift = lift * d_third - - # --- far apart: the Newton form at the three roots --- - inside = -minors * (1.0 / 3.0) - d_inside = -d_minors * (1.0 / 3.0) - radius = tl.sqrt(tl.maximum(inside, 1e-300)) - d_radius = tl.where(inside > 0.0, 0.5 * d_inside / radius, 0.0) - cube = radius * radius * radius - raw = 0.5 * determinant / cube - d_raw = (0.5 * d_determinant - raw * 3.0 * radius * radius * d_radius) / cube - inside_limit = (raw > -_ARG_LIMIT) & (raw < _ARG_LIMIT) - argument = tl.minimum(tl.maximum(raw, -_ARG_LIMIT), _ARG_LIMIT) - d_argument = tl.where(inside_limit, d_raw, 0.0) - angle = libdevice.acos(argument) / 3.0 - d_angle = -d_argument / ( - 3.0 * tl.sqrt(tl.maximum(1.0 - argument * argument, 1e-300)) - ) - root_a = 2.0 * radius * tl.cos(angle) + third - d_root_a = ( - 2.0 * d_radius * tl.cos(angle) - - 2.0 * radius * tl.sin(angle) * d_angle - + d_third - ) - root_b = 2.0 * radius * tl.cos(angle - _TURN_THIRD) + third - d_root_b = ( - 2.0 * d_radius * tl.cos(angle - _TURN_THIRD) - - 2.0 * radius * tl.sin(angle - _TURN_THIRD) * d_angle - + d_third - ) - root_c = 2.0 * radius * tl.cos(angle - 2.0 * _TURN_THIRD) + third - d_root_c = ( - 2.0 * d_radius * tl.cos(angle - 2.0 * _TURN_THIRD) - - 2.0 * radius * tl.sin(angle - 2.0 * _TURN_THIRD) * d_angle - + d_third - ) - # Sorting is a permutation, so the tangents follow their own values. - low = tl.minimum(tl.minimum(root_a, root_b), root_c) - high = tl.maximum(tl.maximum(root_a, root_b), root_c) - middle = tl.maximum( - tl.minimum(root_a, root_b), tl.minimum(tl.maximum(root_a, root_b), root_c) - ) - d_low = tl.where( - root_a == low, d_root_a, tl.where(root_b == low, d_root_b, d_root_c) - ) - d_high = tl.where( - root_a == high, d_root_a, tl.where(root_b == high, d_root_b, d_root_c) - ) - d_middle = tl.where( - root_a == middle, d_root_a, tl.where(root_b == middle, d_root_b, d_root_c) - ) - leading = tl.exp(low) - d_leading = leading * d_low - centre = tl.exp(middle) - d_centre = centre * d_middle - trailing = tl.exp(high) - d_trailing = trailing * d_high - first, d_first = _exp_difference_jvp(low, d_low, middle, d_middle, leading, centre) - upper, d_upper = _exp_difference_jvp( - middle, d_middle, high, d_high, centre, trailing - ) - span = high - low - d_span = d_high - d_low - guarded = tl.where(span > 0.0, span, 1.0) - d_guarded = tl.where(span > 0.0, d_span, 0.0) - second = (upper - first) / guarded - d_second = (d_upper - d_first - second * d_guarded) / guarded - - # --- the shifted generator squared, for the series branch --- - q00 = s00 * s00 + a01 * a10 + a02 * a20 - d_q00 = 2.0 * s00 * d_s00 + d_a01 * a10 + a01 * d_a10 + d_a02 * a20 + a02 * d_a20 - q01 = a01 * (s00 + s11) - d_q01 = d_a01 * (s00 + s11) + a01 * (d_s00 + d_s11) - q02 = a02 * (s00 + s22) - d_q02 = d_a02 * (s00 + s22) + a02 * (d_s00 + d_s22) - q10 = a10 * (s00 + s11) - d_q10 = d_a10 * (s00 + s11) + a10 * (d_s00 + d_s11) - q11 = a10 * a01 + s11 * s11 - d_q11 = d_a10 * a01 + a10 * d_a01 + 2.0 * s11 * d_s11 - q12 = a10 * a02 - d_q12 = d_a10 * a02 + a10 * d_a02 - q20 = a20 * (s00 + s22) - d_q20 = d_a20 * (s00 + s22) + a20 * (d_s00 + d_s22) - q21 = a20 * a01 - d_q21 = d_a20 * a01 + a20 * d_a01 - q22 = a20 * a02 + s22 * s22 - d_q22 = d_a20 * a02 + a20 * d_a02 + 2.0 * s22 * d_s22 - - return ( - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - a00, - d_a00, - a01, - d_a01, - a02, - d_a02, - a10, - d_a10, - a11, - d_a11, - a20, - d_a20, - a22, - d_a22, - s00, - d_s00, - s11, - d_s11, - s22, - d_s22, - minors, - d_minors, - sum_flat, - sum_linear, - sum_square, - d_sum_flat, - d_sum_linear, - d_sum_square, - lift, - d_lift, - low, - middle, - d_low, - d_middle, - leading, - d_leading, - first, - d_first, - second, - d_second, - determinant, - d_determinant, - high, - d_high, - radius, - d_radius, - cube, - raw, - d_raw, - argument, - inside_limit, - angle, - d_angle, - centre, - d_centre, - trailing, - d_trailing, - guarded, - d_guarded, - q00, - d_q00, - q01, - d_q01, - q02, - d_q02, - q10, - d_q10, - q11, - d_q11, - q12, - d_q12, - q20, - d_q20, - q21, - d_q21, - q22, - d_q22, - ) - - -@triton.jit -def _three_pool_assemble_jvp( - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - a00, - d_a00, - a01, - d_a01, - a02, - d_a02, - a10, - d_a10, - a11, - d_a11, - a20, - d_a20, - a22, - d_a22, - s00, - d_s00, - s11, - d_s11, - s22, - d_s22, - minors, - d_minors, - sum_flat, - sum_linear, - sum_square, - d_sum_flat, - d_sum_linear, - d_sum_square, - lift, - d_lift, - low, - middle, - d_low, - d_middle, - leading, - d_leading, - first, - d_first, - second, - d_second, - determinant, - d_determinant, - high, - d_high, - radius, - d_radius, - cube, - raw, - d_raw, - argument, - inside_limit, - angle, - d_angle, - centre, - d_centre, - trailing, - d_trailing, - guarded, - d_guarded, - q00, - d_q00, - q01, - d_q01, - q02, - d_q02, - q10, - d_q10, - q11, - d_q11, - q12, - d_q12, - q20, - d_q20, - q21, - d_q21, - q22, - d_q22, - narrow: tl.constexpr = False, -): - """The bare three-pool operator, assembled from its shared pieces. - - Both branches are formed and one is chosen: a ``where`` evaluates each - side, so the divisor each of them carries is guarded whether or not it - is the side taken. - - In double, and before any attenuation -- what a reverse sweep reads, - and what :func:`_three_pool_weigh_jvp` turns into an interval's step. - """ - c00 = lift * (sum_flat + sum_linear * s00 + sum_square * q00) - d_c00 = d_lift * (sum_flat + sum_linear * s00 + sum_square * q00) + lift * ( - d_sum_flat - + d_sum_linear * s00 - + sum_linear * d_s00 - + d_sum_square * q00 - + sum_square * d_q00 - ) - c01 = lift * (sum_linear * a01 + sum_square * q01) - d_c01 = d_lift * (sum_linear * a01 + sum_square * q01) + lift * ( - d_sum_linear * a01 - + sum_linear * d_a01 - + d_sum_square * q01 - + sum_square * d_q01 - ) - c02 = lift * (sum_linear * a02 + sum_square * q02) - d_c02 = d_lift * (sum_linear * a02 + sum_square * q02) + lift * ( - d_sum_linear * a02 - + sum_linear * d_a02 - + d_sum_square * q02 - + sum_square * d_q02 - ) - c10 = lift * (sum_linear * a10 + sum_square * q10) - d_c10 = d_lift * (sum_linear * a10 + sum_square * q10) + lift * ( - d_sum_linear * a10 - + sum_linear * d_a10 - + d_sum_square * q10 - + sum_square * d_q10 - ) - c11 = lift * (sum_flat + sum_linear * s11 + sum_square * q11) - d_c11 = d_lift * (sum_flat + sum_linear * s11 + sum_square * q11) + lift * ( - d_sum_flat - + d_sum_linear * s11 - + sum_linear * d_s11 - + d_sum_square * q11 - + sum_square * d_q11 - ) - c12 = lift * (sum_square * q12) - d_c12 = d_lift * (sum_square * q12) + lift * ( - d_sum_square * q12 + sum_square * d_q12 - ) - c20 = lift * (sum_linear * a20 + sum_square * q20) - d_c20 = d_lift * (sum_linear * a20 + sum_square * q20) + lift * ( - d_sum_linear * a20 - + sum_linear * d_a20 - + d_sum_square * q20 - + sum_square * d_q20 - ) - c21 = lift * (sum_square * q21) - d_c21 = d_lift * (sum_square * q21) + lift * ( - d_sum_square * q21 + sum_square * d_q21 - ) - c22 = lift * (sum_flat + sum_linear * s22 + sum_square * q22) - d_c22 = d_lift * (sum_flat + sum_linear * s22 + sum_square * q22) + lift * ( - d_sum_flat - + d_sum_linear * s22 - + sum_linear * d_s22 - + d_sum_square * q22 - + sum_square * d_q22 - ) - - # --- the Newton form's two factors, for the eigenvalue branch --- - m00 = a00 - low - d_m00 = d_a00 - d_low - m11 = a11 - low - d_m11 = d_a11 - d_low - m22 = a22 - low - d_m22 = d_a22 - d_low - n00 = a00 - middle - d_n00 = d_a00 - d_middle - n11 = a11 - middle - d_n11 = d_a11 - d_middle - n22 = a22 - middle - d_n22 = d_a22 - d_middle - p00 = m00 * n00 + a01 * a10 + a02 * a20 - d_p00 = ( - d_m00 * n00 - + m00 * d_n00 - + d_a01 * a10 - + a01 * d_a10 - + d_a02 * a20 - + a02 * d_a20 - ) - p01 = a01 * (m00 + n11) - d_p01 = d_a01 * (m00 + n11) + a01 * (d_m00 + d_n11) - p02 = a02 * (m00 + n22) - d_p02 = d_a02 * (m00 + n22) + a02 * (d_m00 + d_n22) - p10 = a10 * (n00 + m11) - d_p10 = d_a10 * (n00 + m11) + a10 * (d_n00 + d_m11) - p11 = a10 * a01 + m11 * n11 - d_p11 = d_a10 * a01 + a10 * d_a01 + d_m11 * n11 + m11 * d_n11 - p12 = a10 * a02 - d_p12 = d_a10 * a02 + a10 * d_a02 - p20 = a20 * (n00 + m22) - d_p20 = d_a20 * (n00 + m22) + a20 * (d_n00 + d_m22) - p21 = a20 * a01 - d_p21 = d_a20 * a01 + a20 * d_a01 - p22 = a20 * a02 + m22 * n22 - d_p22 = d_a20 * a02 + a20 * d_a02 + d_m22 * n22 + m22 * d_n22 - - e00 = leading + first * m00 + second * p00 - d_e00 = d_leading + d_first * m00 + first * d_m00 + d_second * p00 + second * d_p00 - e01 = first * a01 + second * p01 - d_e01 = d_first * a01 + first * d_a01 + d_second * p01 + second * d_p01 - e02 = first * a02 + second * p02 - d_e02 = d_first * a02 + first * d_a02 + d_second * p02 + second * d_p02 - e10 = first * a10 + second * p10 - d_e10 = d_first * a10 + first * d_a10 + d_second * p10 + second * d_p10 - e11 = leading + first * m11 + second * p11 - d_e11 = d_leading + d_first * m11 + first * d_m11 + d_second * p11 + second * d_p11 - e12 = second * p12 - d_e12 = d_second * p12 + second * d_p12 - e20 = first * a20 + second * p20 - d_e20 = d_first * a20 + first * d_a20 + d_second * p20 + second * d_p20 - e21 = second * p21 - d_e21 = d_second * p21 + second * d_p21 - e22 = leading + first * m22 + second * p22 - d_e22 = d_leading + d_first * m22 + first * d_m22 + d_second * p22 + second * d_p22 - - if narrow: - # The caller has bounded the spread, so the roots are unreachable and - # everything that leads to them goes with this select. - def_00, dif_00 = c00, d_c00 - def_01, dif_01 = c01, d_c01 - def_02, dif_02 = c02, d_c02 - def_10, dif_10 = c10, d_c10 - def_11, dif_11 = c11, d_c11 - def_12, dif_12 = c12, d_c12 - def_20, dif_20 = c20, d_c20 - def_21, dif_21 = c21, d_c21 - def_22, dif_22 = c22, d_c22 - else: - close = -2.0 * minors < _SPREAD_CUT * _SPREAD_CUT - def_00 = tl.where(close, c00, e00) - dif_00 = tl.where(close, d_c00, d_e00) - def_01 = tl.where(close, c01, e01) - dif_01 = tl.where(close, d_c01, d_e01) - def_02 = tl.where(close, c02, e02) - dif_02 = tl.where(close, d_c02, d_e02) - def_10 = tl.where(close, c10, e10) - dif_10 = tl.where(close, d_c10, d_e10) - def_11 = tl.where(close, c11, e11) - dif_11 = tl.where(close, d_c11, d_e11) - def_12 = tl.where(close, c12, e12) - dif_12 = tl.where(close, d_c12, d_e12) - def_20 = tl.where(close, c20, e20) - dif_20 = tl.where(close, d_c20, d_e20) - def_21 = tl.where(close, c21, e21) - dif_21 = tl.where(close, d_c21, d_e21) - def_22 = tl.where(close, c22, e22) - dif_22 = tl.where(close, d_c22, d_e22) - - return ( - def_00, - dif_00, - def_01, - dif_01, - def_02, - dif_02, - def_10, - dif_10, - def_11, - dif_11, - def_12, - dif_12, - def_20, - dif_20, - def_21, - dif_21, - def_22, - dif_22, - ) - - -@triton.jit -def _three_pool_weigh_jvp( - def_00, - dif_00, - def_01, - dif_01, - def_02, - dif_02, - def_10, - dif_10, - def_11, - dif_11, - def_12, - dif_12, - def_20, - dif_20, - def_21, - dif_21, - def_22, - dif_22, - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - attenuation, - d_attenuation, - narrow: tl.constexpr = False, -): - """An interval's step, from the bare operator and what survives it. - - The recovery is ``(I - E) m0`` rather than a solve, which is what the - equilibrium being a fixed point of the generator buys. - """ - work: tl.constexpr = tl.float32 if narrow else tl.float64 - damp = attenuation.to(work) - d_damp = d_attenuation.to(work) - w00 = damp * def_00 - dw00 = d_damp * def_00 + damp * dif_00 - w01 = damp * def_01 - dw01 = d_damp * def_01 + damp * dif_01 - w02 = damp * def_02 - dw02 = d_damp * def_02 + damp * dif_02 - w10 = damp * def_10 - dw10 = d_damp * def_10 + damp * dif_10 - w11 = damp * def_11 - dw11 = d_damp * def_11 + damp * dif_11 - w12 = damp * def_12 - dw12 = d_damp * def_12 + damp * dif_12 - w20 = damp * def_20 - dw20 = d_damp * def_20 + damp * dif_20 - w21 = damp * def_21 - dw21 = d_damp * def_21 + damp * dif_21 - w22 = damp * def_22 - dw22 = d_damp * def_22 + damp * dif_22 - - grow_free = free - (w00 * free + w01 * pool_b + w02 * pool_c) - d_grow_free = d_free - ( - dw00 * free - + w00 * d_free - + dw01 * pool_b - + w01 * d_pool_b - + dw02 * pool_c - + w02 * d_pool_c - ) - grow_pool_b = pool_b - (w10 * free + w11 * pool_b + w12 * pool_c) - d_grow_pool_b = d_pool_b - ( - dw10 * free - + w10 * d_free - + dw11 * pool_b - + w11 * d_pool_b - + dw12 * pool_c - + w12 * d_pool_c - ) - grow_bound = pool_c - (w20 * free + w21 * pool_b + w22 * pool_c) - d_grow_bound = d_pool_c - ( - dw20 * free - + w20 * d_free - + dw21 * pool_b - + w21 * d_pool_b - + dw22 * pool_c - + w22 * d_pool_c - ) - return ( - w00, - w01, - w02, - w10, - w11, - w12, - w20, - w21, - w22, - grow_free, - grow_pool_b, - grow_bound, - dw00, - dw01, - dw02, - dw10, - dw11, - dw12, - dw20, - dw21, - dw22, - d_grow_free, - d_grow_pool_b, - d_grow_bound, - ) - - -@triton.jit -def _three_pool_step_jvp( - r1_free, - d_r1_free, - r1_pool_b, - d_r1_pool_b, - r1_bound, - d_r1_bound, - exchange_b, - d_exchange_b, - exchange_c, - d_exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - dt, - d_dt, - attenuation, - d_attenuation, - narrow: tl.constexpr = False, -): - """The three-pool longitudinal step and its directional derivative. - - The same closed form :func:`_three_pool_step` evaluates, carried - alongside a tangent and in the same double precision. Returns the - nine entries and three recoveries, then their twelve tangents. - """ - ( - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - a00, - d_a00, - a01, - d_a01, - a02, - d_a02, - a10, - d_a10, - a11, - d_a11, - a20, - d_a20, - a22, - d_a22, - s00, - d_s00, - s11, - d_s11, - s22, - d_s22, - minors, - d_minors, - sum_flat, - sum_linear, - sum_square, - d_sum_flat, - d_sum_linear, - d_sum_square, - lift, - d_lift, - low, - middle, - d_low, - d_middle, - leading, - d_leading, - first, - d_first, - second, - d_second, - determinant, - d_determinant, - high, - d_high, - radius, - d_radius, - cube, - raw, - d_raw, - argument, - inside_limit, - angle, - d_angle, - centre, - d_centre, - trailing, - d_trailing, - guarded, - d_guarded, - q00, - d_q00, - q01, - d_q01, - q02, - d_q02, - q10, - d_q10, - q11, - d_q11, - q12, - d_q12, - q20, - d_q20, - q21, - d_q21, - q22, - d_q22, - ) = _three_pool_pieces_jvp( - r1_free, - d_r1_free, - r1_pool_b, - d_r1_pool_b, - r1_bound, - d_r1_bound, - exchange_b, - d_exchange_b, - exchange_c, - d_exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - dt, - d_dt, - narrow, - ) - ( - def_00, - dif_00, - def_01, - dif_01, - def_02, - dif_02, - def_10, - dif_10, - def_11, - dif_11, - def_12, - dif_12, - def_20, - dif_20, - def_21, - dif_21, - def_22, - dif_22, - ) = _three_pool_assemble_jvp( - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - a00, - d_a00, - a01, - d_a01, - a02, - d_a02, - a10, - d_a10, - a11, - d_a11, - a20, - d_a20, - a22, - d_a22, - s00, - d_s00, - s11, - d_s11, - s22, - d_s22, - minors, - d_minors, - sum_flat, - sum_linear, - sum_square, - d_sum_flat, - d_sum_linear, - d_sum_square, - lift, - d_lift, - low, - middle, - d_low, - d_middle, - leading, - d_leading, - first, - d_first, - second, - d_second, - determinant, - d_determinant, - high, - d_high, - radius, - d_radius, - cube, - raw, - d_raw, - argument, - inside_limit, - angle, - d_angle, - centre, - d_centre, - trailing, - d_trailing, - guarded, - d_guarded, - q00, - d_q00, - q01, - d_q01, - q02, - d_q02, - q10, - d_q10, - q11, - d_q11, - q12, - d_q12, - q20, - d_q20, - q21, - d_q21, - q22, - d_q22, - narrow, - ) - ( - w00, - w01, - w02, - w10, - w11, - w12, - w20, - w21, - w22, - grow_free, - grow_pool_b, - grow_bound, - dw00, - dw01, - dw02, - dw10, - dw11, - dw12, - dw20, - dw21, - dw22, - d_grow_free, - d_grow_pool_b, - d_grow_bound, - ) = _three_pool_weigh_jvp( - def_00, - dif_00, - def_01, - dif_01, - def_02, - dif_02, - def_10, - dif_10, - def_11, - dif_11, - def_12, - dif_12, - def_20, - dif_20, - def_21, - dif_21, - def_22, - dif_22, - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - attenuation, - d_attenuation, - narrow, - ) - return ( - w00.to(tl.float32), - w01.to(tl.float32), - w02.to(tl.float32), - w10.to(tl.float32), - w11.to(tl.float32), - w12.to(tl.float32), - w20.to(tl.float32), - w21.to(tl.float32), - w22.to(tl.float32), - grow_free.to(tl.float32), - grow_pool_b.to(tl.float32), - grow_bound.to(tl.float32), - dw00.to(tl.float32), - dw01.to(tl.float32), - dw02.to(tl.float32), - dw10.to(tl.float32), - dw11.to(tl.float32), - dw12.to(tl.float32), - dw20.to(tl.float32), - dw21.to(tl.float32), - dw22.to(tl.float32), - d_grow_free.to(tl.float32), - d_grow_pool_b.to(tl.float32), - d_grow_bound.to(tl.float32), - ) - - -@triton.jit -def _exp_difference_adjoint_jvp( - lower, - d_lower, - upper, - d_upper, - exp_lower, - d_exp_lower, - exp_upper, - d_exp_upper, - seed, - d_seed, -): - """The reverse of :func:`_exp_difference`, onto both points, on a direction. - - Near the coalescence the slope comes from the same series the value does, - because the difference quotient's own derivative is a cancellation divided - by a small number twice over. - """ - half = 0.5 * (upper - lower) - d_half = 0.5 * (d_upper - d_lower) - near = tl.abs(half) < _SINCH_CUT - poly = 1.0 + half + 0.5 * half * half - d_poly = d_half + half * d_half - even = 1.0 + half * half / 6.0 - d_even = half * d_half / 3.0 - slope = (1.0 + half) * even + poly * half * (1.0 / 3.0) - d_slope = ( - d_half * even - + (1.0 + half) * d_even - + (d_poly * half + poly * d_half) * (1.0 / 3.0) - ) - series = exp_lower * poly * even - d_series = d_exp_lower * poly * even + exp_lower * (d_poly * even + poly * d_even) - swing = 0.5 * exp_lower * slope - d_swing = 0.5 * (d_exp_lower * slope + exp_lower * d_slope) - - gap = tl.where(near, 1.0, upper - lower) - d_gap = tl.where(near, 0.0, d_upper - d_lower) - value = (exp_upper - exp_lower) / gap - d_value = (d_exp_upper - d_exp_lower - value * d_gap) / gap - far_lower = (value - exp_lower) / gap - d_far_lower = (d_value - d_exp_lower - far_lower * d_gap) / gap - far_upper = (exp_upper - value) / gap - d_far_upper = (d_exp_upper - d_value - far_upper * d_gap) / gap - - to_lower = tl.where(near, series - swing, far_lower) - d_to_lower = tl.where(near, d_series - d_swing, d_far_lower) - to_upper = tl.where(near, swing, far_upper) - d_to_upper = tl.where(near, d_swing, d_far_upper) - return ( - seed * to_lower, - d_seed * to_lower + seed * d_to_lower, - seed * to_upper, - d_seed * to_upper + seed * d_to_upper, - ) - - -@triton.jit -def _three_pool_step_adjoint_jvp( - r1_free, - d_r1_free, - r1_pool_b, - d_r1_pool_b, - r1_bound, - d_r1_bound, - exchange_b, - d_exchange_b, - exchange_c, - d_exchange_c, - fraction_b, - d_fraction_b, - fraction_c, - d_fraction_c, - dt, - d_dt, - attenuation, - d_attenuation, - bar_e00, - d_bar_e00, - bar_e01, - d_bar_e01, - bar_e02, - d_bar_e02, - bar_e10, - d_bar_e10, - bar_e11, - d_bar_e11, - bar_e12, - d_bar_e12, - bar_e20, - d_bar_e20, - bar_e21, - d_bar_e21, - bar_e22, - d_bar_e22, - bar_grow_free, - d_bar_grow_free, - bar_grow_pool_b, - d_bar_grow_pool_b, - bar_grow_bound, - d_bar_grow_bound, - free, - d_free, - pool_b, - d_pool_b, - pool_c, - d_pool_c, - a00, - d_a00, - a01, - d_a01, - a02, - d_a02, - a10, - d_a10, - a11, - d_a11, - a20, - d_a20, - a22, - d_a22, - s00, - d_s00, - s11, - d_s11, - s22, - d_s22, - minors, - d_minors, - sum_flat, - sum_linear, - sum_square, - d_sum_flat, - d_sum_linear, - d_sum_square, - lift, - d_lift, - low, - middle, - d_low, - d_middle, - leading, - d_leading, - first, - d_first, - second, - d_second, - determinant, - d_determinant, - high, - d_high, - radius, - d_radius, - cube, - raw, - d_raw, - argument, - inside_limit, - angle, - d_angle, - centre, - d_centre, - trailing, - d_trailing, - guarded, - d_guarded, - q00, - d_q00, - q01, - d_q01, - q02, - d_q02, - q10, - d_q10, - q11, - d_q11, - q12, - d_q12, - q20, - d_q20, - q21, - d_q21, - q22, - d_q22, - def_00, - dif_00, - def_01, - dif_01, - def_02, - dif_02, - def_10, - dif_10, - def_11, - dif_11, - def_12, - dif_12, - def_20, - dif_20, - def_21, - dif_21, - def_22, - dif_22, - narrow: tl.constexpr = False, -): - """The reverse sweep of :func:`_three_pool_step`, carried on a direction. - - Reads the pieces and the bare operator the replay already formed, so an - interval's transcendentals are taken once for the pass rather than once - for each direction through it, and in the same double. - - Both branches are swept, each by the algebra its own forward used, and the - choice between them is made on the cotangents rather than on the way in -- - a ``where`` evaluates both sides, so each side's divisors are guarded. - - The series branch is a polynomial in the two invariants alone, so its - reverse is reached by carrying the recurrence's sensitivity to those two - forward beside it, which needs no history of the sixteen terms. - - Returned as the gradients w.r.t. ``(r1_free, r1_pool_b, r1_bound, - exchange_b, exchange_c, fraction_b, fraction_c, dt, attenuation)`` and - then their nine tangents. - """ - work: tl.constexpr = tl.float32 if narrow else tl.float64 - # --- the recovery and the attenuation, which both branches share --- - damp = attenuation.to(work) - d_damp = d_attenuation.to(work) - r0 = bar_grow_free.to(work) - d_r0 = d_bar_grow_free.to(work) - r1 = bar_grow_pool_b.to(work) - d_r1 = d_bar_grow_pool_b.to(work) - r2 = bar_grow_bound.to(work) - d_r2 = d_bar_grow_bound.to(work) - - y00 = bar_e00.to(work) - r0 * free - d_y00 = d_bar_e00.to(work) - d_r0 * free - r0 * d_free - y01 = bar_e01.to(work) - r0 * pool_b - d_y01 = d_bar_e01.to(work) - d_r0 * pool_b - r0 * d_pool_b - y02 = bar_e02.to(work) - r0 * pool_c - d_y02 = d_bar_e02.to(work) - d_r0 * pool_c - r0 * d_pool_c - y10 = bar_e10.to(work) - r1 * free - d_y10 = d_bar_e10.to(work) - d_r1 * free - r1 * d_free - y11 = bar_e11.to(work) - r1 * pool_b - d_y11 = d_bar_e11.to(work) - d_r1 * pool_b - r1 * d_pool_b - y12 = bar_e12.to(work) - r1 * pool_c - d_y12 = d_bar_e12.to(work) - d_r1 * pool_c - r1 * d_pool_c - y20 = bar_e20.to(work) - r2 * free - d_y20 = d_bar_e20.to(work) - d_r2 * free - r2 * d_free - y21 = bar_e21.to(work) - r2 * pool_b - d_y21 = d_bar_e21.to(work) - d_r2 * pool_b - r2 * d_pool_b - y22 = bar_e22.to(work) - r2 * pool_c - d_y22 = d_bar_e22.to(work) - d_r2 * pool_c - r2 * d_pool_c - - # The bare operator is read rather than the attenuation divided back out - # of the weighed one -- a washed-out interval leaves nothing to divide by. - bar_damp = ( - y00 * def_00 - + y01 * def_01 - + y02 * def_02 - + y10 * def_10 - + y11 * def_11 - + y12 * def_12 - + y20 * def_20 - + y21 * def_21 - + y22 * def_22 - ) - d_bar_damp = ( - d_y00 * def_00 - + y00 * dif_00 - + d_y01 * def_01 - + y01 * dif_01 - + d_y02 * def_02 - + y02 * dif_02 - + d_y10 * def_10 - + y10 * dif_10 - + d_y11 * def_11 - + y11 * dif_11 - + d_y12 * def_12 - + y12 * dif_12 - + d_y20 * def_20 - + y20 * dif_20 - + d_y21 * def_21 - + y21 * dif_21 - + d_y22 * def_22 - + y22 * dif_22 - ) - column_free = r0 * def_00 + r1 * def_10 + r2 * def_20 - d_column_free = ( - d_r0 * def_00 - + r0 * dif_00 - + d_r1 * def_10 - + r1 * dif_10 - + d_r2 * def_20 - + r2 * dif_20 - ) - column_pool_b = r0 * def_01 + r1 * def_11 + r2 * def_21 - d_column_pool_b = ( - d_r0 * def_01 - + r0 * dif_01 - + d_r1 * def_11 - + r1 * dif_11 - + d_r2 * def_21 - + r2 * dif_21 - ) - column_bound = r0 * def_02 + r1 * def_12 + r2 * def_22 - d_column_bound = ( - d_r0 * def_02 - + r0 * dif_02 - + d_r1 * def_12 - + r1 * dif_12 - + d_r2 * def_22 - + r2 * dif_22 - ) - bar_free = r0 - damp * column_free - d_bar_free = d_r0 - d_damp * column_free - damp * d_column_free - bar_pool_b = r1 - damp * column_pool_b - d_bar_pool_b = d_r1 - d_damp * column_pool_b - damp * d_column_pool_b - bar_pool_c = r2 - damp * column_bound - d_bar_pool_c = d_r2 - d_damp * column_bound - damp * d_column_bound - - o00 = damp * y00 - d_o00 = d_damp * y00 + damp * d_y00 - o01 = damp * y01 - d_o01 = d_damp * y01 + damp * d_y01 - o02 = damp * y02 - d_o02 = d_damp * y02 + damp * d_y02 - o10 = damp * y10 - d_o10 = d_damp * y10 + damp * d_y10 - o11 = damp * y11 - d_o11 = d_damp * y11 + damp * d_y11 - o12 = damp * y12 - d_o12 = d_damp * y12 + damp * d_y12 - o20 = damp * y20 - d_o20 = d_damp * y20 + damp * d_y20 - o21 = damp * y21 - d_o21 = d_damp * y21 + damp * d_y21 - o22 = damp * y22 - d_o22 = d_damp * y22 + damp * d_y22 - - # --- close together: the series in the two invariants, run backwards --- - scale00 = o00 * lift - d_scale00 = d_o00 * lift + o00 * d_lift - scale01 = o01 * lift - d_scale01 = d_o01 * lift + o01 * d_lift - scale02 = o02 * lift - d_scale02 = d_o02 * lift + o02 * d_lift - scale10 = o10 * lift - d_scale10 = d_o10 * lift + o10 * d_lift - scale11 = o11 * lift - d_scale11 = d_o11 * lift + o11 * d_lift - scale12 = o12 * lift - d_scale12 = d_o12 * lift + o12 * d_lift - scale20 = o20 * lift - d_scale20 = d_o20 * lift + o20 * d_lift - scale21 = o21 * lift - d_scale21 = d_o21 * lift + o21 * d_lift - scale22 = o22 * lift - d_scale22 = d_o22 * lift + o22 * d_lift - - bar_flat = scale00 + scale11 + scale22 - d_bar_flat = d_scale00 + d_scale11 + d_scale22 - bar_linear = ( - scale00 * s00 - + scale01 * a01 - + scale02 * a02 - + scale10 * a10 - + scale11 * s11 - + scale20 * a20 - + scale22 * s22 - ) - d_bar_linear = ( - d_scale00 * s00 - + scale00 * d_s00 - + d_scale01 * a01 - + scale01 * d_a01 - + d_scale02 * a02 - + scale02 * d_a02 - + d_scale10 * a10 - + scale10 * d_a10 - + d_scale11 * s11 - + scale11 * d_s11 - + d_scale20 * a20 - + scale20 * d_a20 - + d_scale22 * s22 - + scale22 * d_s22 - ) - bar_square = ( - scale00 * q00 - + scale01 * q01 - + scale02 * q02 - + scale10 * q10 - + scale11 * q11 - + scale12 * q12 - + scale20 * q20 - + scale21 * q21 - + scale22 * q22 - ) - d_bar_square = ( - d_scale00 * q00 - + scale00 * d_q00 - + d_scale01 * q01 - + scale01 * d_q01 - + d_scale02 * q02 - + scale02 * d_q02 - + d_scale10 * q10 - + scale10 * d_q10 - + d_scale11 * q11 - + scale11 * d_q11 - + d_scale12 * q12 - + scale12 * d_q12 - + d_scale20 * q20 - + scale20 * d_q20 - + d_scale21 * q21 - + scale21 * d_q21 - + d_scale22 * q22 - + scale22 * d_q22 - ) - # ``lift`` multiplies the whole bracket, so the shift it carries picks up - # the bracket back again -- which is what the three sums contract to. - turn_series = ( - sum_flat * bar_flat + sum_linear * bar_linear + sum_square * bar_square - ) - d_turn_series = ( - d_sum_flat * bar_flat - + sum_flat * d_bar_flat - + d_sum_linear * bar_linear - + sum_linear * d_bar_linear - + d_sum_square * bar_square - + sum_square * d_bar_square - ) - - g00 = sum_square * scale00 - d_g00 = d_sum_square * scale00 + sum_square * d_scale00 - g01 = sum_square * scale01 - d_g01 = d_sum_square * scale01 + sum_square * d_scale01 - g02 = sum_square * scale02 - d_g02 = d_sum_square * scale02 + sum_square * d_scale02 - g10 = sum_square * scale10 - d_g10 = d_sum_square * scale10 + sum_square * d_scale10 - g11 = sum_square * scale11 - d_g11 = d_sum_square * scale11 + sum_square * d_scale11 - g12 = sum_square * scale12 - d_g12 = d_sum_square * scale12 + sum_square * d_scale12 - g20 = sum_square * scale20 - d_g20 = d_sum_square * scale20 + sum_square * d_scale20 - g21 = sum_square * scale21 - d_g21 = d_sum_square * scale21 + sum_square * d_scale21 - g22 = sum_square * scale22 - d_g22 = d_sum_square * scale22 + sum_square * d_scale22 - - # The square's reverse, ``g @ shifted^T + shifted^T @ g``. - v00 = ( - g00 * s00 - + g01 * a01 - + g02 * a02 - + s00 * g00 - + a10 * g10 - + a20 * g20 - + sum_linear * scale00 - ) - d_v00 = ( - d_g00 * s00 - + g00 * d_s00 - + d_g01 * a01 - + g01 * d_a01 - + d_g02 * a02 - + g02 * d_a02 - + d_s00 * g00 - + s00 * d_g00 - + d_a10 * g10 - + a10 * d_g10 - + d_a20 * g20 - + a20 * d_g20 - + d_sum_linear * scale00 - + sum_linear * d_scale00 - ) - v01 = ( - g00 * a10 + g01 * s11 + s00 * g01 + a10 * g11 + a20 * g21 + sum_linear * scale01 - ) - d_v01 = ( - d_g00 * a10 - + g00 * d_a10 - + d_g01 * s11 - + g01 * d_s11 - + d_s00 * g01 - + s00 * d_g01 - + d_a10 * g11 - + a10 * d_g11 - + d_a20 * g21 - + a20 * d_g21 - + d_sum_linear * scale01 - + sum_linear * d_scale01 - ) - v02 = ( - g00 * a20 + g02 * s22 + s00 * g02 + a10 * g12 + a20 * g22 + sum_linear * scale02 - ) - d_v02 = ( - d_g00 * a20 - + g00 * d_a20 - + d_g02 * s22 - + g02 * d_s22 - + d_s00 * g02 - + s00 * d_g02 - + d_a10 * g12 - + a10 * d_g12 - + d_a20 * g22 - + a20 * d_g22 - + d_sum_linear * scale02 - + sum_linear * d_scale02 - ) - v10 = ( - g10 * s00 + g11 * a01 + g12 * a02 + a01 * g00 + s11 * g10 + sum_linear * scale10 - ) - d_v10 = ( - d_g10 * s00 - + g10 * d_s00 - + d_g11 * a01 - + g11 * d_a01 - + d_g12 * a02 - + g12 * d_a02 - + d_a01 * g00 - + a01 * d_g00 - + d_s11 * g10 - + s11 * d_g10 - + d_sum_linear * scale10 - + sum_linear * d_scale10 - ) - v11 = g10 * a10 + g11 * s11 + a01 * g01 + s11 * g11 + sum_linear * scale11 - d_v11 = ( - d_g10 * a10 - + g10 * d_a10 - + d_g11 * s11 - + g11 * d_s11 - + d_a01 * g01 - + a01 * d_g01 - + d_s11 * g11 - + s11 * d_g11 - + d_sum_linear * scale11 - + sum_linear * d_scale11 - ) - v20 = ( - g20 * s00 + g21 * a01 + g22 * a02 + a02 * g00 + s22 * g20 + sum_linear * scale20 - ) - d_v20 = ( - d_g20 * s00 - + g20 * d_s00 - + d_g21 * a01 - + g21 * d_a01 - + d_g22 * a02 - + g22 * d_a02 - + d_a02 * g00 - + a02 * d_g00 - + d_s22 * g20 - + s22 * d_g20 - + d_sum_linear * scale20 - + sum_linear * d_scale20 - ) - v22 = g20 * a20 + g22 * s22 + a02 * g02 + s22 * g22 + sum_linear * scale22 - d_v22 = ( - d_g20 * a20 - + g20 * d_a20 - + d_g22 * s22 - + g22 * d_s22 - + d_a02 * g02 - + a02 * d_g02 - + d_s22 * g22 - + s22 * d_g22 - + d_sum_linear * scale22 - + sum_linear * d_scale22 - ) - turn_series = turn_series - (v00 + v11 + v22) - d_turn_series = d_turn_series - (d_v00 + d_v11 + d_v22) - - # The recurrence's own sensitivity to the two invariants, carried forward - # beside it: two numbers reach the whole series, so their derivatives are - # cheaper to push forward than the sixteen terms are to keep. - flat = 1.0 + 0.0 * a00 - linear = 0.0 * a00 - square = 0.0 * a00 - d_flat = 0.0 * a00 - d_linear = 0.0 * a00 - d_square = 0.0 * a00 - fu = 0.0 * a00 - lu = 0.0 * a00 - su = 0.0 * a00 - d_fu = 0.0 * a00 - d_lu = 0.0 * a00 - d_su = 0.0 * a00 - fv = 0.0 * a00 - lv = 0.0 * a00 - sv = 0.0 * a00 - d_fv = 0.0 * a00 - d_lv = 0.0 * a00 - d_sv = 0.0 * a00 - slope_u_flat = 0.0 * a00 - slope_u_linear = 0.0 * a00 - slope_u_square = 0.0 * a00 - d_slope_u_flat = 0.0 * a00 - d_slope_u_linear = 0.0 * a00 - d_slope_u_square = 0.0 * a00 - slope_v_flat = 0.0 * a00 - slope_v_linear = 0.0 * a00 - slope_v_square = 0.0 * a00 - d_slope_v_flat = 0.0 * a00 - d_slope_v_linear = 0.0 * a00 - d_slope_v_square = 0.0 * a00 - factorial = 1.0 - for order in tl.static_range(1, 16): - next_flat = square * determinant - d_next_flat = d_square * determinant + square * d_determinant - next_linear = flat - square * minors - d_next_linear = d_flat - d_square * minors - square * d_minors - next_square = linear - d_next_square = d_linear - next_fu = su * determinant - d_next_fu = d_su * determinant + su * d_determinant - next_lu = fu - su * minors - square - d_next_lu = d_fu - d_su * minors - su * d_minors - d_square - next_su = lu - d_next_su = d_lu - next_fv = sv * determinant + square - d_next_fv = d_sv * determinant + sv * d_determinant + d_square - next_lv = fv - sv * minors - d_next_lv = d_fv - d_sv * minors - sv * d_minors - next_sv = lv - d_next_sv = d_lv - flat = next_flat - linear = next_linear - square = next_square - d_flat = d_next_flat - d_linear = d_next_linear - d_square = d_next_square - fu = next_fu - lu = next_lu - su = next_su - d_fu = d_next_fu - d_lu = d_next_lu - d_su = d_next_su - fv = next_fv - lv = next_lv - sv = next_sv - d_fv = d_next_fv - d_lv = d_next_lv - d_sv = d_next_sv - factorial = factorial * order - weight = 1.0 / factorial - slope_u_flat = slope_u_flat + weight * fu - slope_u_linear = slope_u_linear + weight * lu - slope_u_square = slope_u_square + weight * su - d_slope_u_flat = d_slope_u_flat + weight * d_fu - d_slope_u_linear = d_slope_u_linear + weight * d_lu - d_slope_u_square = d_slope_u_square + weight * d_su - slope_v_flat = slope_v_flat + weight * fv - slope_v_linear = slope_v_linear + weight * lv - slope_v_square = slope_v_square + weight * sv - d_slope_v_flat = d_slope_v_flat + weight * d_fv - d_slope_v_linear = d_slope_v_linear + weight * d_lv - d_slope_v_square = d_slope_v_square + weight * d_sv - - minors_series = ( - bar_flat * slope_u_flat - + bar_linear * slope_u_linear - + bar_square * slope_u_square - ) - d_minors_series = ( - d_bar_flat * slope_u_flat - + bar_flat * d_slope_u_flat - + d_bar_linear * slope_u_linear - + bar_linear * d_slope_u_linear - + d_bar_square * slope_u_square - + bar_square * d_slope_u_square - ) - determinant_series = ( - bar_flat * slope_v_flat - + bar_linear * slope_v_linear - + bar_square * slope_v_square - ) - d_determinant_series = ( - d_bar_flat * slope_v_flat - + bar_flat * d_slope_v_flat - + d_bar_linear * slope_v_linear - + bar_linear * d_slope_v_linear - + d_bar_square * slope_v_square - + bar_square * d_slope_v_square - ) - - # --- far apart: back through the Newton form and the three roots --- - m00 = a00 - low - d_m00 = d_a00 - d_low - m11 = a11 - low - d_m11 = d_a11 - d_low - m22 = a22 - low - d_m22 = d_a22 - d_low - n00 = a00 - middle - d_n00 = d_a00 - d_middle - n11 = a11 - middle - d_n11 = d_a11 - d_middle - n22 = a22 - middle - d_n22 = d_a22 - d_middle - p00 = m00 * n00 + a01 * a10 + a02 * a20 - d_p00 = ( - d_m00 * n00 - + m00 * d_n00 - + d_a01 * a10 - + a01 * d_a10 - + d_a02 * a20 - + a02 * d_a20 - ) - p01 = a01 * (m00 + n11) - d_p01 = d_a01 * (m00 + n11) + a01 * (d_m00 + d_n11) - p02 = a02 * (m00 + n22) - d_p02 = d_a02 * (m00 + n22) + a02 * (d_m00 + d_n22) - p10 = a10 * (m11 + n00) - d_p10 = d_a10 * (m11 + n00) + a10 * (d_m11 + d_n00) - p11 = m11 * n11 + a01 * a10 - d_p11 = d_m11 * n11 + m11 * d_n11 + d_a01 * a10 + a01 * d_a10 - p12 = a10 * a02 - d_p12 = d_a10 * a02 + a10 * d_a02 - p20 = a20 * (m22 + n00) - d_p20 = d_a20 * (m22 + n00) + a20 * (d_m22 + d_n00) - p21 = a20 * a01 - d_p21 = d_a20 * a01 + a20 * d_a01 - p22 = m22 * n22 + a02 * a20 - d_p22 = d_m22 * n22 + m22 * d_n22 + d_a02 * a20 + a02 * d_a20 - - bar_leading = o00 + o11 + o22 - d_bar_leading = d_o00 + d_o11 + d_o22 - bar_first = ( - o00 * m00 - + o01 * a01 - + o02 * a02 - + o10 * a10 - + o11 * m11 - + o20 * a20 - + o22 * m22 - ) - d_bar_first = ( - d_o00 * m00 - + o00 * d_m00 - + d_o01 * a01 - + o01 * d_a01 - + d_o02 * a02 - + o02 * d_a02 - + d_o10 * a10 - + o10 * d_a10 - + d_o11 * m11 - + o11 * d_m11 - + d_o20 * a20 - + o20 * d_a20 - + d_o22 * m22 - + o22 * d_m22 - ) - bar_second = ( - o00 * p00 - + o01 * p01 - + o02 * p02 - + o10 * p10 - + o11 * p11 - + o12 * p12 - + o20 * p20 - + o21 * p21 - + o22 * p22 - ) - d_bar_second = ( - d_o00 * p00 - + o00 * d_p00 - + d_o01 * p01 - + o01 * d_p01 - + d_o02 * p02 - + o02 * d_p02 - + d_o10 * p10 - + o10 * d_p10 - + d_o11 * p11 - + o11 * d_p11 - + d_o12 * p12 - + o12 * d_p12 - + d_o20 * p20 - + o20 * d_p20 - + d_o21 * p21 - + o21 * d_p21 - + d_o22 * p22 - + o22 * d_p22 - ) - - z00 = second * o00 - d_z00 = d_second * o00 + second * d_o00 - z01 = second * o01 - d_z01 = d_second * o01 + second * d_o01 - z02 = second * o02 - d_z02 = d_second * o02 + second * d_o02 - z10 = second * o10 - d_z10 = d_second * o10 + second * d_o10 - z11 = second * o11 - d_z11 = d_second * o11 + second * d_o11 - z12 = second * o12 - d_z12 = d_second * o12 + second * d_o12 - z20 = second * o20 - d_z20 = d_second * o20 + second * d_o20 - z21 = second * o21 - d_z21 = d_second * o21 + second * d_o21 - z22 = second * o22 - d_z22 = d_second * o22 + second * d_o22 - - # ``z @ n^T``, the product's reverse onto the first factor. - u00 = z00 * n00 + z01 * a01 + z02 * a02 - d_u00 = ( - d_z00 * n00 - + z00 * d_n00 - + d_z01 * a01 - + z01 * d_a01 - + d_z02 * a02 - + z02 * d_a02 - ) - u01 = z00 * a10 + z01 * n11 - d_u01 = d_z00 * a10 + z00 * d_a10 + d_z01 * n11 + z01 * d_n11 - u02 = z00 * a20 + z02 * n22 - d_u02 = d_z00 * a20 + z00 * d_a20 + d_z02 * n22 + z02 * d_n22 - u10 = z10 * n00 + z11 * a01 + z12 * a02 - d_u10 = ( - d_z10 * n00 - + z10 * d_n00 - + d_z11 * a01 - + z11 * d_a01 - + d_z12 * a02 - + z12 * d_a02 - ) - u11 = z10 * a10 + z11 * n11 - d_u11 = d_z10 * a10 + z10 * d_a10 + d_z11 * n11 + z11 * d_n11 - u20 = z20 * n00 + z21 * a01 + z22 * a02 - d_u20 = ( - d_z20 * n00 - + z20 * d_n00 - + d_z21 * a01 - + z21 * d_a01 - + d_z22 * a02 - + z22 * d_a02 - ) - u22 = z20 * a20 + z22 * n22 - d_u22 = d_z20 * a20 + z20 * d_a20 + d_z22 * n22 + z22 * d_n22 - - # ``m^T @ z``, onto the second. - w00 = m00 * z00 + a10 * z10 + a20 * z20 - d_w00 = ( - d_m00 * z00 - + m00 * d_z00 - + d_a10 * z10 - + a10 * d_z10 - + d_a20 * z20 - + a20 * d_z20 - ) - w01 = m00 * z01 + a10 * z11 + a20 * z21 - d_w01 = ( - d_m00 * z01 - + m00 * d_z01 - + d_a10 * z11 - + a10 * d_z11 - + d_a20 * z21 - + a20 * d_z21 - ) - w02 = m00 * z02 + a10 * z12 + a20 * z22 - d_w02 = ( - d_m00 * z02 - + m00 * d_z02 - + d_a10 * z12 - + a10 * d_z12 - + d_a20 * z22 - + a20 * d_z22 - ) - w10 = a01 * z00 + m11 * z10 - d_w10 = d_a01 * z00 + a01 * d_z00 + d_m11 * z10 + m11 * d_z10 - w11 = a01 * z01 + m11 * z11 - d_w11 = d_a01 * z01 + a01 * d_z01 + d_m11 * z11 + m11 * d_z11 - w20 = a02 * z00 + m22 * z20 - d_w20 = d_a02 * z00 + a02 * d_z00 + d_m22 * z20 + m22 * d_z20 - w22 = a02 * z02 + m22 * z22 - d_w22 = d_a02 * z02 + a02 * d_z02 + d_m22 * z22 + m22 * d_z22 - - bar_low = bar_leading * leading - first * (o00 + o11 + o22) - (u00 + u11 + u22) - d_bar_low = ( - d_bar_leading * leading - + bar_leading * d_leading - - d_first * (o00 + o11 + o22) - - first * (d_o00 + d_o11 + d_o22) - - (d_u00 + d_u11 + d_u22) - ) - bar_middle = -(w00 + w11 + w22) - d_bar_middle = -(d_w00 + d_w11 + d_w22) - bar_high = 0.0 * a00 - d_bar_high = 0.0 * a00 - - span = high - low - positive = span > 0.0 - bar_upper = bar_second / guarded - d_bar_upper = (d_bar_second - bar_upper * d_guarded) / guarded - bar_first = bar_first - bar_upper - d_bar_first = d_bar_first - d_bar_upper - bar_span = tl.where(positive, -bar_upper * second, 0.0) - d_bar_span = tl.where(positive, -d_bar_upper * second - bar_upper * d_second, 0.0) - bar_high = bar_high + bar_span - d_bar_high = d_bar_high + d_bar_span - bar_low = bar_low - bar_span - d_bar_low = d_bar_low - d_bar_span - - ( - from_first_low, - d_from_first_low, - from_first_middle, - d_from_first_middle, - ) = _exp_difference_adjoint_jvp( - low, - d_low, - middle, - d_middle, - leading, - d_leading, - centre, - d_centre, - bar_first, - d_bar_first, - ) - ( - from_upper_middle, - d_from_upper_middle, - from_upper_high, - d_from_upper_high, - ) = _exp_difference_adjoint_jvp( - middle, - d_middle, - high, - d_high, - centre, - d_centre, - trailing, - d_trailing, - bar_upper, - d_bar_upper, - ) - bar_low = bar_low + from_first_low - d_bar_low = d_bar_low + d_from_first_low - bar_middle = bar_middle + from_first_middle + from_upper_middle - d_bar_middle = d_bar_middle + d_from_first_middle + d_from_upper_middle - bar_high = bar_high + from_upper_high - d_bar_high = d_bar_high + d_from_upper_high - - # The three roots come off one angle a third of a turn apart, and the - # cosine puts them in a fixed order: the last turn is the lowest, the - # first the highest, whatever the angle is. - swing_low = angle - 2.0 * _TURN_THIRD - swing_middle = angle - _TURN_THIRD - cos_low = tl.cos(swing_low) - cos_middle = tl.cos(swing_middle) - cos_high = tl.cos(angle) - sin_low = tl.sin(swing_low) - sin_middle = tl.sin(swing_middle) - sin_high = tl.sin(angle) - bar_radius = 2.0 * ( - cos_low * bar_low + cos_middle * bar_middle + cos_high * bar_high - ) - d_bar_radius = 2.0 * ( - cos_low * d_bar_low - + cos_middle * d_bar_middle - + cos_high * d_bar_high - - d_angle * (sin_low * bar_low + sin_middle * bar_middle + sin_high * bar_high) - ) - swept = sin_low * bar_low + sin_middle * bar_middle + sin_high * bar_high - d_swept = ( - sin_low * d_bar_low - + sin_middle * d_bar_middle - + sin_high * d_bar_high - + d_angle * (cos_low * bar_low + cos_middle * bar_middle + cos_high * bar_high) - ) - bar_angle = -2.0 * radius * swept - d_bar_angle = -2.0 * (d_radius * swept + radius * d_swept) - turn_roots = bar_low + bar_middle + bar_high - d_turn_roots = d_bar_low + d_bar_middle + d_bar_high - - # ``acos`` is clamped, and where it is the angle no longer moves with the - # cubic's argument -- which is what keeps a double root differentiable. - d_argument = tl.where(inside_limit, d_raw, 0.0) - inner = 1.0 - argument * argument - d_inner = -2.0 * argument * d_argument - stem = tl.sqrt(tl.maximum(inner, 1e-300)) - tilt = -1.0 / (3.0 * stem) - d_tilt = d_inner / (6.0 * stem * stem * stem) - bar_raw = tl.where(inside_limit, bar_angle * tilt, 0.0) - d_bar_raw = tl.where(inside_limit, d_bar_angle * tilt + bar_angle * d_tilt, 0.0) - safe_radius = tl.where(radius > 1e-30, radius, 1.0) - d_safe_radius = tl.where(radius > 1e-30, d_radius, 0.0) - safe_cube = safe_radius * safe_radius * safe_radius - d_safe_cube = 3.0 * safe_radius * safe_radius * d_safe_radius - determinant_roots = 0.5 * bar_raw / safe_cube - d_determinant_roots = ( - 0.5 * d_bar_raw - determinant_roots * d_safe_cube - ) / safe_cube - pull = tl.where(inside_limit, -3.0 * raw * bar_raw / safe_radius, 0.0) - d_pull = tl.where( - inside_limit, - (-3.0 * (d_raw * bar_raw + raw * d_bar_raw) - pull * d_safe_radius) - / safe_radius, - 0.0, - ) - bar_radius = bar_radius + pull - d_bar_radius = d_bar_radius + d_pull - minors_roots = -bar_radius / (6.0 * safe_radius) - d_minors_roots = (-d_bar_radius - minors_roots * 6.0 * d_safe_radius) / ( - 6.0 * safe_radius - ) - - # --- the branch chosen on the cotangents, not on the way in --- - if narrow: - bar_a00, d_bar_a00 = v00, d_v00 - bar_a01, d_bar_a01 = v01, d_v01 - bar_a02, d_bar_a02 = v02, d_v02 - bar_a10, d_bar_a10 = v10, d_v10 - bar_a11, d_bar_a11 = v11, d_v11 - bar_a20, d_bar_a20 = v20, d_v20 - bar_a22, d_bar_a22 = v22, d_v22 - bar_third, d_bar_third = turn_series, d_turn_series - bar_minors, d_bar_minors = minors_series, d_minors_series - bar_determinant, d_bar_determinant = (determinant_series, d_determinant_series) - else: - close = -2.0 * minors < _SPREAD_CUT * _SPREAD_CUT - bar_a00 = tl.where(close, v00, first * o00 + u00 + w00) - d_bar_a00 = tl.where( - close, d_v00, d_first * o00 + first * d_o00 + d_u00 + d_w00 - ) - bar_a01 = tl.where(close, v01, first * o01 + u01 + w01) - d_bar_a01 = tl.where( - close, d_v01, d_first * o01 + first * d_o01 + d_u01 + d_w01 - ) - bar_a02 = tl.where(close, v02, first * o02 + u02 + w02) - d_bar_a02 = tl.where( - close, d_v02, d_first * o02 + first * d_o02 + d_u02 + d_w02 - ) - bar_a10 = tl.where(close, v10, first * o10 + u10 + w10) - d_bar_a10 = tl.where( - close, d_v10, d_first * o10 + first * d_o10 + d_u10 + d_w10 - ) - bar_a11 = tl.where(close, v11, first * o11 + u11 + w11) - d_bar_a11 = tl.where( - close, d_v11, d_first * o11 + first * d_o11 + d_u11 + d_w11 - ) - bar_a20 = tl.where(close, v20, first * o20 + u20 + w20) - d_bar_a20 = tl.where( - close, d_v20, d_first * o20 + first * d_o20 + d_u20 + d_w20 - ) - bar_a22 = tl.where(close, v22, first * o22 + u22 + w22) - d_bar_a22 = tl.where( - close, d_v22, d_first * o22 + first * d_o22 + d_u22 + d_w22 - ) - bar_third = tl.where(close, turn_series, turn_roots) - d_bar_third = tl.where(close, d_turn_series, d_turn_roots) - bar_minors = tl.where(close, minors_series, minors_roots) - d_bar_minors = tl.where(close, d_minors_series, d_minors_roots) - bar_determinant = tl.where(close, determinant_series, determinant_roots) - d_bar_determinant = tl.where(close, d_determinant_series, d_determinant_roots) - - # --- the two invariants back onto the shifted generator --- - cofactor00 = s11 * s22 - d_cofactor00 = d_s11 * s22 + s11 * d_s22 - cofactor11 = s00 * s22 - a02 * a20 - d_cofactor11 = d_s00 * s22 + s00 * d_s22 - d_a02 * a20 - a02 * d_a20 - cofactor22 = s00 * s11 - a01 * a10 - d_cofactor22 = d_s00 * s11 + s00 * d_s11 - d_a01 * a10 - a01 * d_a10 - shift00 = bar_minors * (s11 + s22) + bar_determinant * cofactor00 - d_shift00 = ( - d_bar_minors * (s11 + s22) - + bar_minors * (d_s11 + d_s22) - + d_bar_determinant * cofactor00 - + bar_determinant * d_cofactor00 - ) - shift11 = bar_minors * (s00 + s22) + bar_determinant * cofactor11 - d_shift11 = ( - d_bar_minors * (s00 + s22) - + bar_minors * (d_s00 + d_s22) - + d_bar_determinant * cofactor11 - + bar_determinant * d_cofactor11 - ) - shift22 = bar_minors * (s00 + s11) + bar_determinant * cofactor22 - d_shift22 = ( - d_bar_minors * (s00 + s11) - + bar_minors * (d_s00 + d_s11) - + d_bar_determinant * cofactor22 - + bar_determinant * d_cofactor22 - ) - shift01 = -a10 * (bar_minors + bar_determinant * s22) - d_shift01 = -d_a10 * (bar_minors + bar_determinant * s22) - a10 * ( - d_bar_minors + d_bar_determinant * s22 + bar_determinant * d_s22 - ) - shift10 = -a01 * (bar_minors + bar_determinant * s22) - d_shift10 = -d_a01 * (bar_minors + bar_determinant * s22) - a01 * ( - d_bar_minors + d_bar_determinant * s22 + bar_determinant * d_s22 - ) - shift02 = -a20 * (bar_minors + bar_determinant * s11) - d_shift02 = -d_a20 * (bar_minors + bar_determinant * s11) - a20 * ( - d_bar_minors + d_bar_determinant * s11 + bar_determinant * d_s11 - ) - shift20 = -a02 * (bar_minors + bar_determinant * s11) - d_shift20 = -d_a02 * (bar_minors + bar_determinant * s11) - a02 * ( - d_bar_minors + d_bar_determinant * s11 + bar_determinant * d_s11 - ) - bar_a00 = bar_a00 + shift00 - d_bar_a00 = d_bar_a00 + d_shift00 - bar_a01 = bar_a01 + shift01 - d_bar_a01 = d_bar_a01 + d_shift01 - bar_a02 = bar_a02 + shift02 - d_bar_a02 = d_bar_a02 + d_shift02 - bar_a10 = bar_a10 + shift10 - d_bar_a10 = d_bar_a10 + d_shift10 - bar_a11 = bar_a11 + shift11 - d_bar_a11 = d_bar_a11 + d_shift11 - bar_a20 = bar_a20 + shift20 - d_bar_a20 = d_bar_a20 + d_shift20 - bar_a22 = bar_a22 + shift22 - d_bar_a22 = d_bar_a22 + d_shift22 - bar_third = bar_third - (shift00 + shift11 + shift22) - d_bar_third = d_bar_third - (d_shift00 + d_shift11 + d_shift22) - bar_a00 = bar_a00 + bar_third / 3.0 - d_bar_a00 = d_bar_a00 + d_bar_third / 3.0 - bar_a11 = bar_a11 + bar_third / 3.0 - d_bar_a11 = d_bar_a11 + d_bar_third / 3.0 - bar_a22 = bar_a22 + bar_third / 3.0 - d_bar_a22 = d_bar_a22 + d_bar_third / 3.0 - - # --- the generator back onto the rates, the fractions and the interval --- - step = dt.to(work) - d_step = d_dt.to(work) - rate_b = exchange_b.to(work) - d_rate_b = d_exchange_b.to(work) - rate_c = exchange_c.to(work) - d_rate_c = d_exchange_c.to(work) - kab = rate_b * pool_b - d_kab = d_rate_b * pool_b + rate_b * d_pool_b - kba = rate_b * free - d_kba = d_rate_b * free + rate_b * d_free - kac = rate_c * pool_c - d_kac = d_rate_c * pool_c + rate_c * d_pool_c - kca = rate_c * free - d_kca = d_rate_c * free + rate_c * d_free - row_a = -kab - kac - r1_free.to(work) - d_row_a = -d_kab - d_kac - d_r1_free.to(work) - row_b = -kba - r1_pool_b.to(work) - d_row_b = -d_kba - d_r1_pool_b.to(work) - row_c = -kca - r1_bound.to(work) - d_row_c = -d_kca - d_r1_bound.to(work) - - bar_step = ( - row_a * bar_a00 - + kba * bar_a01 - + kca * bar_a02 - + kab * bar_a10 - + row_b * bar_a11 - + kac * bar_a20 - + row_c * bar_a22 - ) - d_bar_step = ( - d_row_a * bar_a00 - + row_a * d_bar_a00 - + d_kba * bar_a01 - + kba * d_bar_a01 - + d_kca * bar_a02 - + kca * d_bar_a02 - + d_kab * bar_a10 - + kab * d_bar_a10 - + d_row_b * bar_a11 - + row_b * d_bar_a11 - + d_kac * bar_a20 - + kac * d_bar_a20 - + d_row_c * bar_a22 - + row_c * d_bar_a22 - ) - bar_kab = step * (bar_a10 - bar_a00) - d_bar_kab = d_step * (bar_a10 - bar_a00) + step * (d_bar_a10 - d_bar_a00) - bar_kba = step * (bar_a01 - bar_a11) - d_bar_kba = d_step * (bar_a01 - bar_a11) + step * (d_bar_a01 - d_bar_a11) - bar_kac = step * (bar_a20 - bar_a00) - d_bar_kac = d_step * (bar_a20 - bar_a00) + step * (d_bar_a20 - d_bar_a00) - bar_kca = step * (bar_a02 - bar_a22) - d_bar_kca = d_step * (bar_a02 - bar_a22) + step * (d_bar_a02 - d_bar_a22) - - whole_free = bar_free + rate_b * bar_kba + rate_c * bar_kca - d_whole_free = ( - d_bar_free - + d_rate_b * bar_kba - + rate_b * d_bar_kba - + d_rate_c * bar_kca - + rate_c * d_bar_kca - ) - whole_pool_b = bar_pool_b + rate_b * bar_kab - d_whole_pool_b = d_bar_pool_b + d_rate_b * bar_kab + rate_b * d_bar_kab - whole_pool_c = bar_pool_c + rate_c * bar_kac - d_whole_pool_c = d_bar_pool_c + d_rate_c * bar_kac + rate_c * d_bar_kac - - return ( - (-step * bar_a00).to(tl.float32), - (-step * bar_a11).to(tl.float32), - (-step * bar_a22).to(tl.float32), - (pool_b * bar_kab + free * bar_kba).to(tl.float32), - (pool_c * bar_kac + free * bar_kca).to(tl.float32), - (whole_pool_b - whole_free).to(tl.float32), - (whole_pool_c - whole_free).to(tl.float32), - bar_step.to(tl.float32), - bar_damp.to(tl.float32), - (-d_step * bar_a00 - step * d_bar_a00).to(tl.float32), - (-d_step * bar_a11 - step * d_bar_a11).to(tl.float32), - (-d_step * bar_a22 - step * d_bar_a22).to(tl.float32), - ( - d_pool_b * bar_kab - + pool_b * d_bar_kab - + d_free * bar_kba - + free * d_bar_kba - ).to(tl.float32), - ( - d_pool_c * bar_kac - + pool_c * d_bar_kac - + d_free * bar_kca - + free * d_bar_kca - ).to(tl.float32), - (d_whole_pool_b - d_whole_free).to(tl.float32), - (d_whole_pool_c - d_whole_free).to(tl.float32), - d_bar_step.to(tl.float32), - d_bar_damp.to(tl.float32), - ) - - -@triton.jit -def _complex_sqrt(real, imag): - """A square root of a complex number carried as a pair of floats. - - Which of the two it is does not matter here: the only thing that reads it - is even in it, so the branch cut the principal root carries is unreachable. - """ - magnitude = tl.sqrt(real * real + imag * imag) - root_real = tl.sqrt(tl.maximum(0.5 * (magnitude + real), 0.0)) - root_imag = tl.sqrt(tl.maximum(0.5 * (magnitude - real), 0.0)) - return root_real, tl.where(imag < 0.0, -root_imag, root_imag) - - -@triton.jit -def _complex_exp(real, imag): - """``e^z`` for ``z`` carried as a pair of floats.""" - scale = tl.exp(real) - return scale * tl.cos(imag), scale * tl.sin(imag) - - -@triton.jit -def _two_pool_transverse_step( - r2_free, r2_bound, exchange, bound, free, shift_hz, dt, attenuation -): - """The transverse operator of two chemically exchanging pools. - - ``expm((K - diag(R2) - 2 pi i diag(df)) t)``, in the closed form the - longitudinal pair uses -- the numbers have become complex, the algebra has - not. Returned as the four entries, each a pair of floats. - - A semisolid pool holds a share of the voxel without carrying any transverse - magnetization, so it is absent from this 2x2 and present in ``free`` -- how - much free water the exchange sees. - - There is no recovery term: transverse magnetization relaxes toward zero. - """ - kab = exchange * bound - kba = exchange * free - l11 = (-kab - r2_free) * dt - l12 = kba * dt - l21 = kab * dt - l22 = (-kba - r2_bound) * dt - # Only pool b's offset appears: pool a sits at whatever off-resonance the - # free precession already carries the whole voxel through. - l22_imag = -2.0 * 3.141592653589793 * shift_hz * dt - - trace_real = 0.5 * (l11 + l22) - trace_imag = 0.5 * l22_imag - gap_real = 0.5 * (l11 - l22) - gap_imag = -0.5 * l22_imag - square_real = gap_real * gap_real - gap_imag * gap_imag + l12 * l21 - square_imag = 2.0 * gap_real * gap_imag - - root_real, root_imag = _complex_sqrt(square_real, square_imag) - upper_real, upper_imag = _complex_exp( - trace_real + root_real, trace_imag + root_imag - ) - lower_real, lower_imag = _complex_exp( - trace_real - root_real, trace_imag - root_imag - ) - cos_real = 0.5 * (upper_real + lower_real) - cos_imag = 0.5 * (upper_imag + lower_imag) - - # ``sinh(d)/d`` by series near the origin, where the root has no - # derivative of its own. - turning = square_real * square_real + square_imag * square_imag > 1e-24 - half_real = 0.5 * (upper_real - lower_real) - half_imag = 0.5 * (upper_imag - lower_imag) - guard = tl.where(turning, root_real * root_real + root_imag * root_imag, 1.0) - divided_real = (half_real * root_real + half_imag * root_imag) / guard - divided_imag = (half_imag * root_real - half_real * root_imag) / guard - plain_real, plain_imag = _complex_exp(trace_real, trace_imag) - square2_real = square_real * square_real - square_imag * square_imag - square2_imag = 2.0 * square_real * square_imag - poly_real = 1.0 + square_real / 6.0 + square2_real / 120.0 - poly_imag = square_imag / 6.0 + square2_imag / 120.0 - series_real = plain_real * poly_real - plain_imag * poly_imag - series_imag = plain_real * poly_imag + plain_imag * poly_real - scale_real = tl.where(turning, divided_real, series_real) - scale_imag = tl.where(turning, divided_imag, series_imag) - - off_real = scale_real * gap_real - scale_imag * gap_imag - off_imag = scale_real * gap_imag + scale_imag * gap_real - return ( - attenuation * (cos_real + off_real), - attenuation * (cos_imag + off_imag), - attenuation * scale_real * l12, - attenuation * scale_imag * l12, - attenuation * scale_real * l21, - attenuation * scale_imag * l21, - attenuation * (cos_real - off_real), - attenuation * (cos_imag - off_imag), - ) - - -@triton.jit -def _complex_sqrt_jvp(real, imag, d_real, d_imag): - """A complex square root and its directional derivative. - - The derivative divides by twice the root, so a caller keeps the origin -- - where the root has none -- on its series branch. - """ - root_real, root_imag = _complex_sqrt(real, imag) - guard = 2.0 * (root_real * root_real + root_imag * root_imag) - live = guard > 0.0 - guarded = tl.where(live, guard, 1.0) - # dz / (2 w) == dz * conj(2 w) / |2 w|^2 - tangent_real = tl.where( - live, (d_real * root_real + d_imag * root_imag) / guarded, 0.0 - ) - tangent_imag = tl.where( - live, (d_imag * root_real - d_real * root_imag) / guarded, 0.0 - ) - return root_real, root_imag, tangent_real, tangent_imag - - -@triton.jit -def _complex_exp_jvp(real, imag, d_real, d_imag): - """``e^z`` and its directional derivative, which is ``e^z`` times it.""" - value_real, value_imag = _complex_exp(real, imag) - return ( - value_real, - value_imag, - value_real * d_real - value_imag * d_imag, - value_real * d_imag + value_imag * d_real, - ) - - -@triton.jit -def _two_pool_transverse_step_jvp( - r2_free, - d_r2_free, - r2_bound, - d_r2_bound, - exchange, - d_exchange, - bound, - d_bound, - free, - d_free, - shift_hz, - d_shift_hz, - dt, - d_dt, - attenuation, - d_attenuation, -): - """The transverse operator and its directional derivative. - - The same closed form :func:`_two_pool_transverse_step` evaluates, carried - alongside a tangent. Returned as the four entries then their four tangents, - each a pair of floats. - """ - kab = exchange * bound - d_kab = d_exchange * bound + exchange * d_bound - kba = exchange * free - d_kba = d_exchange * free + exchange * d_free - l11 = (-kab - r2_free) * dt - d_l11 = (-d_kab - d_r2_free) * dt + (-kab - r2_free) * d_dt - l12 = kba * dt - d_l12 = d_kba * dt + kba * d_dt - l21 = kab * dt - d_l21 = d_kab * dt + kab * d_dt - l22 = (-kba - r2_bound) * dt - d_l22 = (-d_kba - d_r2_bound) * dt + (-kba - r2_bound) * d_dt - turn = -2.0 * 3.141592653589793 - l22_imag = turn * shift_hz * dt - d_l22_imag = turn * (d_shift_hz * dt + shift_hz * d_dt) - - trace_real = 0.5 * (l11 + l22) - d_trace_real = 0.5 * (d_l11 + d_l22) - trace_imag = 0.5 * l22_imag - d_trace_imag = 0.5 * d_l22_imag - gap_real = 0.5 * (l11 - l22) - d_gap_real = 0.5 * (d_l11 - d_l22) - gap_imag = -0.5 * l22_imag - d_gap_imag = -0.5 * d_l22_imag - - square_real = gap_real * gap_real - gap_imag * gap_imag + l12 * l21 - d_square_real = ( - 2.0 * gap_real * d_gap_real - - 2.0 * gap_imag * d_gap_imag - + d_l12 * l21 - + l12 * d_l21 - ) - square_imag = 2.0 * gap_real * gap_imag - d_square_imag = 2.0 * (d_gap_real * gap_imag + gap_real * d_gap_imag) - - root_real, root_imag, d_root_real, d_root_imag = _complex_sqrt_jvp( - square_real, square_imag, d_square_real, d_square_imag - ) - upper_real, upper_imag, d_upper_real, d_upper_imag = _complex_exp_jvp( - trace_real + root_real, - trace_imag + root_imag, - d_trace_real + d_root_real, - d_trace_imag + d_root_imag, - ) - lower_real, lower_imag, d_lower_real, d_lower_imag = _complex_exp_jvp( - trace_real - root_real, - trace_imag - root_imag, - d_trace_real - d_root_real, - d_trace_imag - d_root_imag, - ) - cos_real = 0.5 * (upper_real + lower_real) - cos_imag = 0.5 * (upper_imag + lower_imag) - d_cos_real = 0.5 * (d_upper_real + d_lower_real) - d_cos_imag = 0.5 * (d_upper_imag + d_lower_imag) - - turning = square_real * square_real + square_imag * square_imag > 1e-24 - half_real = 0.5 * (upper_real - lower_real) - half_imag = 0.5 * (upper_imag - lower_imag) - d_half_real = 0.5 * (d_upper_real - d_lower_real) - d_half_imag = 0.5 * (d_upper_imag - d_lower_imag) - norm = root_real * root_real + root_imag * root_imag - guard = tl.where(turning, norm, 1.0) - d_norm = tl.where( - turning, 2.0 * (root_real * d_root_real + root_imag * d_root_imag), 0.0 - ) - # (a / w) with w complex: a * conj(w) / |w|^2, differentiated as a quotient. - top_real = half_real * root_real + half_imag * root_imag - top_imag = half_imag * root_real - half_real * root_imag - d_top_real = ( - d_half_real * root_real - + half_real * d_root_real - + d_half_imag * root_imag - + half_imag * d_root_imag - ) - d_top_imag = ( - d_half_imag * root_real - + half_imag * d_root_real - - d_half_real * root_imag - - half_real * d_root_imag - ) - divided_real = top_real / guard - divided_imag = top_imag / guard - d_divided_real = (d_top_real - divided_real * d_norm) / guard - d_divided_imag = (d_top_imag - divided_imag * d_norm) / guard - - plain_real, plain_imag, d_plain_real, d_plain_imag = _complex_exp_jvp( - trace_real, trace_imag, d_trace_real, d_trace_imag - ) - square2_real = square_real * square_real - square_imag * square_imag - square2_imag = 2.0 * square_real * square_imag - d_square2_real = ( - 2.0 * square_real * d_square_real - 2.0 * square_imag * d_square_imag - ) - d_square2_imag = 2.0 * (d_square_real * square_imag + square_real * d_square_imag) - poly_real = 1.0 + square_real / 6.0 + square2_real / 120.0 - poly_imag = square_imag / 6.0 + square2_imag / 120.0 - d_poly_real = d_square_real / 6.0 + d_square2_real / 120.0 - d_poly_imag = d_square_imag / 6.0 + d_square2_imag / 120.0 - series_real = plain_real * poly_real - plain_imag * poly_imag - series_imag = plain_real * poly_imag + plain_imag * poly_real - d_series_real = ( - d_plain_real * poly_real - + plain_real * d_poly_real - - d_plain_imag * poly_imag - - plain_imag * d_poly_imag - ) - d_series_imag = ( - d_plain_real * poly_imag - + plain_real * d_poly_imag - + d_plain_imag * poly_real - + plain_imag * d_poly_real - ) - - scale_real = tl.where(turning, divided_real, series_real) - scale_imag = tl.where(turning, divided_imag, series_imag) - d_scale_real = tl.where(turning, d_divided_real, d_series_real) - d_scale_imag = tl.where(turning, d_divided_imag, d_series_imag) - - off_real = scale_real * gap_real - scale_imag * gap_imag - off_imag = scale_real * gap_imag + scale_imag * gap_real - d_off_real = ( - d_scale_real * gap_real - + scale_real * d_gap_real - - d_scale_imag * gap_imag - - scale_imag * d_gap_imag - ) - d_off_imag = ( - d_scale_real * gap_imag - + scale_real * d_gap_imag - + d_scale_imag * gap_real - + scale_imag * d_gap_real - ) - - e11_real = cos_real + off_real - e11_imag = cos_imag + off_imag - d_e11_real = d_cos_real + d_off_real - d_e11_imag = d_cos_imag + d_off_imag - e22_real = cos_real - off_real - e22_imag = cos_imag - off_imag - d_e22_real = d_cos_real - d_off_real - d_e22_imag = d_cos_imag - d_off_imag - return ( - attenuation * e11_real, - attenuation * e11_imag, - attenuation * scale_real * l12, - attenuation * scale_imag * l12, - attenuation * scale_real * l21, - attenuation * scale_imag * l21, - attenuation * e22_real, - attenuation * e22_imag, - d_attenuation * e11_real + attenuation * d_e11_real, - d_attenuation * e11_imag + attenuation * d_e11_imag, - d_attenuation * scale_real * l12 - + attenuation * (d_scale_real * l12 + scale_real * d_l12), - d_attenuation * scale_imag * l12 - + attenuation * (d_scale_imag * l12 + scale_imag * d_l12), - d_attenuation * scale_real * l21 - + attenuation * (d_scale_real * l21 + scale_real * d_l21), - d_attenuation * scale_imag * l21 - + attenuation * (d_scale_imag * l21 + scale_imag * d_l21), - d_attenuation * e22_real + attenuation * d_e22_real, - d_attenuation * e22_imag + attenuation * d_e22_imag, - ) - - -@triton.jit -def _lineshape_at(lineshape, offset_hz, bins, step): - """How well the bound pool absorbs a pulse this far off its resonance. - - Cubic Hermite between the two knots bracketing the offset, taken in - magnitude because the lineshape is even, and clamped at the far end. Each - knot is two floats -- the value then its slope -- so the two a read needs - are four contiguous ones. - """ - last = bins - 1 - scaled = tl.minimum(tl.abs(offset_hz) / step, last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - base = lower.to(tl.int64) * 2 - near = tl.load(lineshape + base) - near_slope = tl.load(lineshape + base + 1) - far = tl.load(lineshape + base + 2) - far_slope = tl.load(lineshape + base + 3) - return ( - (2.0 * u3 - 3.0 * u2 + 1.0) * near - + (u3 - 2.0 * u2 + u) * step * near_slope - + (-2.0 * u3 + 3.0 * u2) * far - + (u3 - u2) * step * far_slope - ) - - -@triton.jit -def _two_pool_step_jvp( - r1_free, - d_r1_free, - r1_bound, - d_r1_bound, - exchange, - d_exchange, - bound, - d_bound, - dt, - d_dt, - attenuation, - d_attenuation, -): - """The two-pool operator and its directional derivative. - - The same closed form :func:`_two_pool_step` evaluates, carried alongside a - tangent. Returned as the six outputs then their six tangents. - """ - free = 1.0 - bound - d_free = -d_bound - kab = exchange * bound - d_kab = d_exchange * bound + exchange * d_bound - kba = exchange * free - d_kba = d_exchange * free + exchange * d_free - l11 = (-kab - r1_free) * dt - d_l11 = (-d_kab - d_r1_free) * dt + (-kab - r1_free) * d_dt - l12 = kba * dt - d_l12 = d_kba * dt + kba * d_dt - l21 = kab * dt - d_l21 = d_kab * dt + kab * d_dt - l22 = (-kba - r1_bound) * dt - d_l22 = (-d_kba - d_r1_bound) * dt + (-kba - r1_bound) * d_dt - - half_trace = 0.5 * (l11 + l22) - d_half_trace = 0.5 * (d_l11 + d_l22) - half_gap = 0.5 * (l11 - l22) - d_half_gap = 0.5 * (d_l11 - d_l22) - square = half_gap * half_gap + l12 * l21 - d_square = 2.0 * half_gap * d_half_gap + d_l12 * l21 + l12 * d_l21 - - turning = square > 1e-12 - root = tl.sqrt(tl.maximum(square, 0.0)) - guarded = tl.where(turning, root, 1.0) - d_root = tl.where(turning, 0.5 * d_square / guarded, 0.0) - - upper = tl.exp(half_trace + root) - d_upper = upper * (d_half_trace + d_root) - lower = tl.exp(half_trace - root) - d_lower = lower * (d_half_trace - d_root) - cosine = 0.5 * (upper + lower) - d_cosine = 0.5 * (d_upper + d_lower) - - # sinh(d)/d by series where the root has no derivative of its own. - plain = tl.exp(half_trace) - d_plain = plain * d_half_trace - poly = 1.0 + square / 6.0 + square * square / 120.0 - d_poly = d_square / 6.0 + square * d_square / 60.0 - scale = tl.where(turning, 0.5 * (upper - lower) / guarded, plain * poly) - d_scale = tl.where( - turning, - 0.5 * (d_upper - d_lower) / guarded - - 0.5 * (upper - lower) * d_root / (guarded * guarded), - d_plain * poly + plain * d_poly, - ) - - e11 = attenuation * (cosine + scale * half_gap) - d_e11 = d_attenuation * (cosine + scale * half_gap) + attenuation * ( - d_cosine + d_scale * half_gap + scale * d_half_gap - ) - e12 = attenuation * scale * l12 - d_e12 = ( - d_attenuation * scale * l12 - + attenuation * d_scale * l12 - + attenuation * scale * d_l12 - ) - e21 = attenuation * scale * l21 - d_e21 = ( - d_attenuation * scale * l21 - + attenuation * d_scale * l21 - + attenuation * scale * d_l21 - ) - e22 = attenuation * (cosine - scale * half_gap) - d_e22 = d_attenuation * (cosine - scale * half_gap) + attenuation * ( - d_cosine - d_scale * half_gap - scale * d_half_gap - ) - - grow_free = free - (e11 * free + e12 * bound) - d_grow_free = d_free - (d_e11 * free + e11 * d_free + d_e12 * bound + e12 * d_bound) - grow_bound = bound - (e21 * free + e22 * bound) - d_grow_bound = d_bound - ( - d_e21 * free + e21 * d_free + d_e22 * bound + e22 * d_bound - ) - return ( - e11, - e12, - e21, - e22, - grow_free, - grow_bound, - d_e11, - d_e12, - d_e21, - d_e22, - d_grow_free, - d_grow_bound, - ) - - -@triton.jit -def _lineshape_at_slope(lineshape, offset_hz, bins, step): - """The lineshape and its derivative in the *signed* offset. - - The table covers the magnitude, so the slope changes sign with the offset; - past the last knot the read is constant and the slope is zero. - """ - last = bins - 1 - magnitude = tl.abs(offset_hz) / step - scaled = tl.minimum(magnitude, last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - base = lower.to(tl.int64) * 2 - near = tl.load(lineshape + base) - near_slope = tl.load(lineshape + base + 1) - far = tl.load(lineshape + base + 2) - far_slope = tl.load(lineshape + base + 3) - value = ( - (2.0 * u3 - 3.0 * u2 + 1.0) * near - + (u3 - 2.0 * u2 + u) * step * near_slope - + (-2.0 * u3 + 3.0 * u2) * far - + (u3 - u2) * step * far_slope - ) - direction = tl.where(offset_hz < 0.0, -1.0, 1.0) - slope = direction * ( - (6.0 * u2 - 6.0 * u) * near / step - + (3.0 * u2 - 4.0 * u + 1.0) * near_slope - + (-6.0 * u2 + 6.0 * u) * far / step - + (3.0 * u2 - 2.0 * u) * far_slope - ) - return value, tl.where(magnitude > last, 0.0, slope) - - -@triton.jit -def _lineshape_at_curve(lineshape, offset_hz, bins, step): - """The lineshape, its slope and its curvature, from the same cubic. - - The table covers the magnitude, so the slope changes sign with the offset - and the curvature does not: an even function's second derivative is even. - """ - last = bins - 1 - magnitude = tl.abs(offset_hz) / step - scaled = tl.minimum(magnitude, last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - base = lower.to(tl.int64) * 2 - near = tl.load(lineshape + base) - near_slope = tl.load(lineshape + base + 1) - far = tl.load(lineshape + base + 2) - far_slope = tl.load(lineshape + base + 3) - value = ( - (2.0 * u3 - 3.0 * u2 + 1.0) * near - + (u3 - 2.0 * u2 + u) * step * near_slope - + (-2.0 * u3 + 3.0 * u2) * far - + (u3 - u2) * step * far_slope - ) - direction = tl.where(offset_hz < 0.0, -1.0, 1.0) - slope = direction * ( - (6.0 * u2 - 6.0 * u) * near / step - + (3.0 * u2 - 4.0 * u + 1.0) * near_slope - + (-6.0 * u2 + 6.0 * u) * far / step - + (3.0 * u2 - 2.0 * u) * far_slope - ) - curve = ( - (12.0 * u - 6.0) * near / (step * step) - + (6.0 * u - 4.0) * near_slope / step - + (-12.0 * u + 6.0) * far / (step * step) - + (6.0 * u - 2.0) * far_slope / step - ) - beyond = magnitude > last - return ( - value, - tl.where(beyond, 0.0, slope), - tl.where(beyond, 0.0, curve), - ) - - -@triton.jit -def _two_pool_step_adjoint_jvp( - r1_free, - d_r1_free, - r1_bound, - d_r1_bound, - exchange, - d_exchange, - bound, - d_bound, - dt, - d_dt, - attenuation, - d_attenuation, - bar_e11, - d_bar_e11, - bar_e12, - d_bar_e12, - bar_e21, - d_bar_e21, - bar_e22, - d_bar_e22, - bar_free, - d_bar_free, - bar_bound, - d_bar_bound, -): - """The reverse sweep of :func:`_two_pool_step`, carried on a direction. - - Recomputes the forward rather than carrying it across the event: the whole - thing is a handful of transcendentals once per interval, against a state - loop that runs per dephasing order. - - Where the discriminant is small the value is still formed from the two - eigenvalues -- a sum, which loses nothing -- but the derivative is taken - from the series, because ``d cosh(d)/d(d^2)`` reached through - ``(e^{t+d} - e^{t-d})/2d`` is a cancellation divided by a small number. - - Returned as the gradients w.r.t. ``(r1_free, r1_bound, exchange, bound, - dt, attenuation)`` then their six tangents. - """ - free = 1.0 - bound - d_free = -d_bound - kab = exchange * bound - d_kab = d_exchange * bound + exchange * d_bound - kba = exchange * free - d_kba = d_exchange * free + exchange * d_free - l11 = (-kab - r1_free) * dt - d_l11 = (-d_kab - d_r1_free) * dt + (-kab - r1_free) * d_dt - l12 = kba * dt - d_l12 = d_kba * dt + kba * d_dt - l21 = kab * dt - d_l21 = d_kab * dt + kab * d_dt - l22 = (-kba - r1_bound) * dt - d_l22 = (-d_kba - d_r1_bound) * dt + (-kba - r1_bound) * d_dt - - half_trace = 0.5 * (l11 + l22) - d_half_trace = 0.5 * (d_l11 + d_l22) - half_gap = 0.5 * (l11 - l22) - d_half_gap = 0.5 * (d_l11 - d_l22) - square = half_gap * half_gap + l12 * l21 - d_square = 2.0 * half_gap * d_half_gap + d_l12 * l21 + l12 * d_l21 - - turning = square > 1e-12 - root = tl.sqrt(tl.maximum(square, 0.0)) - guarded = tl.where(turning, root, 1.0) - d_root = tl.where(turning, 0.5 * d_square / guarded, 0.0) - upper = tl.exp(half_trace + root) - d_upper = upper * (d_half_trace + d_root) - lower = tl.exp(half_trace - root) - d_lower = lower * (d_half_trace - d_root) - cosine = 0.5 * (upper + lower) - d_cosine = 0.5 * (d_upper + d_lower) - plain = tl.exp(half_trace) - d_plain = plain * d_half_trace - poly = 1.0 + square / 6.0 + square * square / 120.0 - d_poly = d_square / 6.0 + square * d_square / 60.0 - scale = tl.where(turning, 0.5 * (upper - lower) / guarded, plain * poly) - d_scale = tl.where( - turning, - 0.5 * (d_upper - d_lower) / guarded - - 0.5 * (upper - lower) * d_root / (guarded * guarded), - d_plain * poly + plain * d_poly, - ) - - bare11 = cosine + scale * half_gap - d_bare11 = d_cosine + d_scale * half_gap + scale * d_half_gap - bare12 = scale * l12 - d_bare12 = d_scale * l12 + scale * d_l12 - bare21 = scale * l21 - d_bare21 = d_scale * l21 + scale * d_l21 - bare22 = cosine - scale * half_gap - d_bare22 = d_cosine - d_scale * half_gap - scale * d_half_gap - - # The recovery reaches the operator's four entries and the two fractions. - carried11 = bar_e11 - bar_free * free - d_carried11 = d_bar_e11 - (d_bar_free * free + bar_free * d_free) - carried12 = bar_e12 - bar_free * bound - d_carried12 = d_bar_e12 - (d_bar_free * bound + bar_free * d_bound) - carried21 = bar_e21 - bar_bound * free - d_carried21 = d_bar_e21 - (d_bar_bound * free + bar_bound * d_free) - carried22 = bar_e22 - bar_bound * bound - d_carried22 = d_bar_e22 - (d_bar_bound * bound + bar_bound * d_bound) - - e11 = attenuation * bare11 - d_e11 = d_attenuation * bare11 + attenuation * d_bare11 - e12 = attenuation * bare12 - d_e12 = d_attenuation * bare12 + attenuation * d_bare12 - e21 = attenuation * bare21 - d_e21 = d_attenuation * bare21 + attenuation * d_bare21 - e22 = attenuation * bare22 - d_e22 = d_attenuation * bare22 + attenuation * d_bare22 - - back_free = bar_free * (1.0 - e11) - bar_bound * e21 - d_back_free = ( - d_bar_free * (1.0 - e11) - - bar_free * d_e11 - - (d_bar_bound * e21 + bar_bound * d_e21) - ) - back_bound = bar_bound * (1.0 - e22) - bar_free * e12 - d_back_bound = ( - d_bar_bound * (1.0 - e22) - - bar_bound * d_e22 - - (d_bar_free * e12 + bar_free * d_e12) - ) - - back_attenuation = ( - carried11 * bare11 - + carried12 * bare12 - + carried21 * bare21 - + carried22 * bare22 - ) - d_back_attenuation = ( - d_carried11 * bare11 - + carried11 * d_bare11 - + d_carried12 * bare12 - + carried12 * d_bare12 - + d_carried21 * bare21 - + carried21 * d_bare21 - + d_carried22 * bare22 - + carried22 * d_bare22 - ) - - scaled11 = attenuation * carried11 - d_scaled11 = d_attenuation * carried11 + attenuation * d_carried11 - scaled12 = attenuation * carried12 - d_scaled12 = d_attenuation * carried12 + attenuation * d_carried12 - scaled21 = attenuation * carried21 - d_scaled21 = d_attenuation * carried21 + attenuation * d_carried21 - scaled22 = attenuation * carried22 - d_scaled22 = d_attenuation * carried22 + attenuation * d_carried22 - - bar_cosine = scaled11 + scaled22 - d_bar_cosine = d_scaled11 + d_scaled22 - gap = scaled11 - scaled22 - d_gap = d_scaled11 - d_scaled22 - bar_scale = gap * half_gap + scaled12 * l12 + scaled21 * l21 - d_bar_scale = ( - d_gap * half_gap - + gap * d_half_gap - + d_scaled12 * l12 - + scaled12 * d_l12 - + d_scaled21 * l21 - + scaled21 * d_l21 - ) - bar_half_gap = scale * gap - d_bar_half_gap = d_scale * gap + scale * d_gap - bar_l12 = scale * scaled12 - d_bar_l12 = d_scale * scaled12 + scale * d_scaled12 - bar_l21 = scale * scaled21 - d_bar_l21 = d_scale * scaled21 + scale * d_scaled21 - - series_trace = bar_cosine * cosine + bar_scale * scale - d_series_trace = ( - d_bar_cosine * cosine - + bar_cosine * d_cosine - + d_bar_scale * scale - + bar_scale * d_scale - ) - cosine_poly = 0.5 + square / 12.0 - d_cosine_poly = d_square / 12.0 - scale_poly = 1.0 / 6.0 + square / 60.0 - d_scale_poly = d_square / 60.0 - series_square = plain * (bar_cosine * cosine_poly + bar_scale * scale_poly) - d_series_square = d_plain * ( - bar_cosine * cosine_poly + bar_scale * scale_poly - ) + plain * ( - d_bar_cosine * cosine_poly - + bar_cosine * d_cosine_poly - + d_bar_scale * scale_poly - + bar_scale * d_scale_poly - ) - - inverse = tl.where(turning, 1.0 / guarded, 0.0) - d_inverse = tl.where(turning, -d_root / (guarded * guarded), 0.0) - bar_upper = 0.5 * (bar_cosine + bar_scale * inverse) - d_bar_upper = 0.5 * (d_bar_cosine + d_bar_scale * inverse + bar_scale * d_inverse) - bar_lower = 0.5 * (bar_cosine - bar_scale * inverse) - d_bar_lower = 0.5 * (d_bar_cosine - d_bar_scale * inverse - bar_scale * d_inverse) - root_trace = bar_upper * upper + bar_lower * lower - d_root_trace = ( - d_bar_upper * upper - + bar_upper * d_upper - + d_bar_lower * lower - + bar_lower * d_lower - ) - bar_root = bar_upper * upper - bar_lower * lower - bar_scale * scale * inverse - d_bar_root = ( - d_bar_upper * upper - + bar_upper * d_upper - - d_bar_lower * lower - - bar_lower * d_lower - - ( - d_bar_scale * scale * inverse - + bar_scale * d_scale * inverse - + bar_scale * scale * d_inverse - ) - ) - root_square = 0.5 * bar_root * inverse - d_root_square = 0.5 * (d_bar_root * inverse + bar_root * d_inverse) - - bar_half_trace = tl.where(turning, root_trace, series_trace) - d_bar_half_trace = tl.where(turning, d_root_trace, d_series_trace) - bar_square = tl.where(turning, root_square, series_square) - d_bar_square = tl.where(turning, d_root_square, d_series_square) - - bar_half_gap += 2.0 * bar_square * half_gap - d_bar_half_gap += 2.0 * (d_bar_square * half_gap + bar_square * d_half_gap) - bar_l12 += bar_square * l21 - d_bar_l12 += d_bar_square * l21 + bar_square * d_l21 - bar_l21 += bar_square * l12 - d_bar_l21 += d_bar_square * l12 + bar_square * d_l12 - - bar_l11 = 0.5 * (bar_half_trace + bar_half_gap) - d_bar_l11 = 0.5 * (d_bar_half_trace + d_bar_half_gap) - bar_l22 = 0.5 * (bar_half_trace - bar_half_gap) - d_bar_l22 = 0.5 * (d_bar_half_trace - d_bar_half_gap) - - bar_kab = (bar_l21 - bar_l11) * dt - d_bar_kab = (d_bar_l21 - d_bar_l11) * dt + (bar_l21 - bar_l11) * d_dt - bar_kba = (bar_l12 - bar_l22) * dt - d_bar_kba = (d_bar_l12 - d_bar_l22) * dt + (bar_l12 - bar_l22) * d_dt - back_dt = ( - bar_l11 * (-kab - r1_free) - + bar_l12 * kba - + bar_l21 * kab - + bar_l22 * (-kba - r1_bound) - ) - d_back_dt = ( - d_bar_l11 * (-kab - r1_free) - + bar_l11 * (-d_kab - d_r1_free) - + d_bar_l12 * kba - + bar_l12 * d_kba - + d_bar_l21 * kab - + bar_l21 * d_kab - + d_bar_l22 * (-kba - r1_bound) - + bar_l22 * (-d_kba - d_r1_bound) - ) - - back_bound += bar_kab * exchange - d_back_bound += d_bar_kab * exchange + bar_kab * d_exchange - back_free += bar_kba * exchange - d_back_free += d_bar_kba * exchange + bar_kba * d_exchange - - return ( - -bar_l11 * dt, - -bar_l22 * dt, - bar_kab * bound + bar_kba * free, - back_bound - back_free, - back_dt, - back_attenuation, - -(d_bar_l11 * dt + bar_l11 * d_dt), - -(d_bar_l22 * dt + bar_l22 * d_dt), - d_bar_kab * bound + bar_kab * d_bound + d_bar_kba * free + bar_kba * d_free, - d_back_bound - d_back_free, - d_back_dt, - d_back_attenuation, - ) - - -@triton.jit -def _washout_jvp(rate, rate_tangent, dt, dt_tangent): - """The same fraction and its directional derivative.""" - fraction = rate * dt - live = fraction < 1.0 - return ( - tl.where(live, 1.0 - fraction, 0.0), - tl.where(live, -(rate_tangent * dt + rate * dt_tangent), 0.0), - ) - - -@triton.jit -def _damping_jvp(rate, rate_tangent, dt, dt_tangent, order): - """Diffusion damping and its directional derivative, per state order.""" - b_factor = rate * dt - b_tangent = rate_tangent * dt + rate * dt_tangent - squared = order * order - transverse_weight = squared + order + 0.3333333333333333 - damp_z = tl.exp(-b_factor * squared) - damp_t = tl.exp(-b_factor * transverse_weight) - return ( - damp_z, - damp_z * (-b_tangent * squared), - damp_t, - damp_t * (-b_tangent * transverse_weight), - ) - - -@triton.jit -def _complex_mul(a_real, a_imag, b_real, b_imag): - return a_real * b_real - a_imag * b_imag, a_real * b_imag + a_imag * b_real - - -@triton.jit -def _dual_mul(a_vr, a_vi, a_tr, a_ti, b_vr, b_vi, b_tr, b_ti): - """Product of two dual complex numbers.""" - value_real, value_imag = _complex_mul(a_vr, a_vi, b_vr, b_vi) - left_real, left_imag = _complex_mul(a_tr, a_ti, b_vr, b_vi) - right_real, right_imag = _complex_mul(a_vr, a_vi, b_tr, b_ti) - return value_real, value_imag, left_real + right_real, left_imag + right_imag - - -@triton.jit -def _dual_scale(scale_value, scale_tangent, vr, vi, tr, ti): - """A real dual number times a complex one.""" - return ( - scale_value * vr, - scale_value * vi, - scale_tangent * vr + scale_value * tr, - scale_tangent * vi + scale_value * ti, - ) - - -@triton.jit -def _dual_times_i(vr, vi, tr, ti): - return -vi, vr, -ti, tr - - -@triton.jit -def _dual_real_conj_mul(a_vr, a_vi, a_tr, a_ti, b_vr, b_vi, b_tr, b_ti): - """``real_part(conj(a) * b)``, the contraction an adjoint asks for.""" - value = a_vr * b_vr + a_vi * b_vi - tangent = a_tr * b_vr + a_ti * b_vi + a_vr * b_tr + a_vi * b_ti - return value, tangent - - -@triton.jit -def _dual_back(entry, d_entry, spin_vr, spin_vi, spin_tr, spin_ti, br, bi, tr, ti): - """``conj(entry * spin)`` against a cotangent, entry and spin both dual. - - One row of a real mixing operator carried through the per-order turn, which - is what a longitudinal cotangent walks back through. - """ - return _dual_mul( - entry * spin_vr, - -(entry * spin_vi), - d_entry * spin_vr + entry * spin_tr, - -(d_entry * spin_vi + entry * spin_ti), - br, - bi, - tr, - ti, - ) - - -@triton.jit -def _dual_polar(angle_value, angle_tangent): - """``exp(i * angle)`` for a real dual angle.""" - cosine = tl.cos(angle_value) - sine = tl.sin(angle_value) - return cosine, sine, -sine * angle_tangent, cosine * angle_tangent - - -@triton.jit -def _shift_adjoint( - plus_bar_real, - plus_bar_imag, - minus_bar_real, - minus_bar_imag, - state, - state_mask, - state_count, -): - """Transpose of ``_shift``. - - The conjugate refill at order zero sends the incoming plus adjoint back - onto minus, conjugated, at the index the minus shift moves it to. - """ - carry_real = tl.where(state_mask, _first(plus_bar_real, state), 0.0) - carry_imag = -tl.where(state_mask, _first(plus_bar_imag, state), 0.0) - forward = (state + 1 < state_count) & state_mask - backward = (state > 0) & state_mask - shifted_pr = tl.where(forward, _down(plus_bar_real, state), 0.0) - shifted_pi = tl.where(forward, _down(plus_bar_imag, state), 0.0) - shifted_mr = tl.where(backward, _up(minus_bar_real, state), 0.0) - shifted_mi = tl.where(backward, _up(minus_bar_imag, state), 0.0) - shifted_mr = tl.where(state == 1, shifted_mr + carry_real, shifted_mr) - shifted_mi = tl.where(state == 1, shifted_mi + carry_imag, shifted_mi) - return shifted_pr, shifted_pi, shifted_mr, shifted_mi - - -@triton.jit -def _shift_real_adjoint( - plus_bar, - minus_bar, - state, - state_mask, - state_count, -): - """Transpose of ``_shift_real``. - - The ``a0 = -b0`` coupling sends the incoming plus adjoint back onto minus, - at the index the minus shift moves it to. - """ - carry = -tl.where(state_mask, _first(plus_bar, state), 0.0) - shifted_plus = tl.where( - (state + 1 < state_count) & state_mask, _down(plus_bar, state), 0.0 - ) - shifted_minus = tl.where((state > 0) & state_mask, _up(minus_bar, state), 0.0) - shifted_minus = tl.where(state == 1, shifted_minus + carry, shifted_minus) - return shifted_plus, shifted_minus - - -@triton.jit -def _table_row(profile_index, event, location, locations): - """Which row of the stacked tables this pulse reads. - - Its own shape's block of ``locations`` rows, then the voxel's place along - the slice. - """ - return tl.load(profile_index + event).to(tl.int64) * locations + location - - -@triton.jit -def _dynamic_pair_at(pairs, pair_index, event_base, event, atom, atom_count, mask): - """The rotation a pulse performs at this voxel, read rather than read off. - - A tabulated pair covers a shape's every pulse because a static array - reaches the rotation through one complex scalar; this one is integrated per - pulse per voxel, so there is nothing to interpolate and the read is four - floats. The row runs per train and per event, as the flip does. - """ - row = tl.load(pair_index + event_base + event).to(tl.int64) - entry = pairs + (row * atom_count + atom) * 4 - return ( - tl.load(entry + 0, mask=mask, other=1.0), - tl.load(entry + 1, mask=mask, other=0.0), - tl.load(entry + 2, mask=mask, other=0.0), - tl.load(entry + 3, mask=mask, other=0.0), - ) - - -@triton.jit -def _profile_pair(profile, row, theta, bins, step): - """The Cayley-Klein pair the transition table holds at this flip angle. - - Cubic Hermite between the two knots bracketing ``theta``, clamped at both - ends: a cubic run off its grid leaves the unit circle. Each knot is eight - floats -- the pair then its slope, real before imaginary -- so the two a - read needs are sixteen contiguous ones. - """ - last = bins - 1 - scaled = tl.minimum(tl.maximum(theta / step, 0.0), last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - h00 = 2.0 * u3 - 3.0 * u2 + 1.0 - h10 = (u3 - 2.0 * u2 + u) * step - h01 = -2.0 * u3 + 3.0 * u2 - h11 = (u3 - u2) * step - - base = (row * bins + lower.to(tl.int64)) * 8 - pair = () - for component in tl.static_range(4): - near = tl.load(profile + base + component) - near_slope = tl.load(profile + base + 4 + component) - far = tl.load(profile + base + 8 + component) - far_slope = tl.load(profile + base + 12 + component) - pair = pair + (h00 * near + h10 * near_slope + h01 * far + h11 * far_slope,) - return pair - - -@triton.jit -def _profile_pair_slope(profile, row, theta, bins, step): - """The pair and its derivative in the flip angle, from the same cubic. - - The derivative of a Hermite segment is another polynomial in the same four - knot values, so reading both costs one extra combination rather than a - second table. Returned interleaved: each component's value then its slope, - in the order ``a`` real, ``a`` imaginary, ``b`` real, ``b`` imaginary. - """ - last = bins - 1 - scaled = tl.minimum(tl.maximum(theta / step, 0.0), last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - h00 = 2.0 * u3 - 3.0 * u2 + 1.0 - h10 = (u3 - 2.0 * u2 + u) * step - h01 = -2.0 * u3 + 3.0 * u2 - h11 = (u3 - u2) * step - # d/dtheta is d/du over the knot spacing. - g00 = (6.0 * u2 - 6.0 * u) / step - g10 = 3.0 * u2 - 4.0 * u + 1.0 - g01 = (6.0 * u - 6.0 * u2) / step - g11 = 3.0 * u2 - 2.0 * u - - base = (row * bins + lower.to(tl.int64)) * 8 - read = () - for component in tl.static_range(4): - near = tl.load(profile + base + component) - near_slope = tl.load(profile + base + 4 + component) - far = tl.load(profile + base + 8 + component) - far_slope = tl.load(profile + base + 12 + component) - read = read + ( - h00 * near + h10 * near_slope + h01 * far + h11 * far_slope, - g00 * near + g10 * near_slope + g01 * far + g11 * far_slope, - ) - return read - - -@triton.jit -def _profile_pair_curve(profile, row, theta, bins, step): - """The pair, its slope and its curvature in the flip angle. - - The second-order pass differentiates the read twice, and a Hermite segment - is a cubic, so all three come from the same four knot values. Returned in - threes per component: value, slope, curvature. - """ - last = bins - 1 - scaled = tl.minimum(tl.maximum(theta / step, 0.0), last + 0.0) - lower = tl.minimum(tl.floor(scaled), last - 1.0) - u = scaled - lower - u2 = u * u - u3 = u2 * u - h00 = 2.0 * u3 - 3.0 * u2 + 1.0 - h10 = (u3 - 2.0 * u2 + u) * step - h01 = -2.0 * u3 + 3.0 * u2 - h11 = (u3 - u2) * step - g00 = (6.0 * u2 - 6.0 * u) / step - g10 = 3.0 * u2 - 4.0 * u + 1.0 - g01 = (6.0 * u - 6.0 * u2) / step - g11 = 3.0 * u2 - 2.0 * u - c00 = (12.0 * u - 6.0) / (step * step) - c10 = (6.0 * u - 4.0) / step - c01 = (6.0 - 12.0 * u) / (step * step) - c11 = (6.0 * u - 2.0) / step - - base = (row * bins + lower.to(tl.int64)) * 8 - read = () - for component in tl.static_range(4): - near = tl.load(profile + base + component) - near_slope = tl.load(profile + base + 4 + component) - far = tl.load(profile + base + 8 + component) - far_slope = tl.load(profile + base + 12 + component) - read = read + ( - h00 * near + h10 * near_slope + h01 * far + h11 * far_slope, - g00 * near + g10 * near_slope + g01 * far + g11 * far_slope, - c00 * near + c10 * near_slope + c01 * far + c11 * far_slope, - ) - return read - - -@triton.jit -def _dual_conj(z): - """A dual complex number's conjugate, both halves.""" - return (z[0], -z[1], z[2], -z[3]) - - -@triton.jit -def _dual_weigh(z, factor): - """A dual complex number scaled by a real constant.""" - return (factor * z[0], factor * z[1], factor * z[2], factor * z[3]) - - -@triton.jit -def _dual_sum(first, second, third, fourth): - """Four dual complex numbers added.""" - return ( - first[0] + second[0] + third[0] + fourth[0], - first[1] + second[1] + third[1] + fourth[1], - first[2] + second[2] + third[2] + fourth[2], - first[3] + second[3] + third[3] + fourth[3], - ) - - -@triton.jit -def _dual_product(x, y): - """Two dual complex numbers multiplied.""" - return _dual_mul(x[0], x[1], x[2], x[3], y[0], y[1], y[2], y[3]) - - -@triton.jit -def _dual_add(x, y): - """Two dual complex numbers added.""" - return (x[0] + y[0], x[1] + y[1], x[2] + y[2], x[3] + y[3]) - - -@triton.jit -def _dual_subtract(x, y): - """One dual complex number less another.""" - return (x[0] - y[0], x[1] - y[1], x[2] - y[2], x[3] - y[3]) - - -@triton.jit -def _dual_reciprocal(z): - """``1/z`` for a dual complex number, and the tangent that goes with it.""" - norm = z[0] * z[0] + z[1] * z[1] - guard = tl.where(norm > 0.0, norm, 1.0) - value_real = z[0] / guard - value_imag = -z[1] / guard - square_real, square_imag = _complex_mul( - value_real, value_imag, value_real, value_imag - ) - tangent_real, tangent_imag = _complex_mul(square_real, square_imag, z[2], z[3]) - return value_real, value_imag, -tangent_real, -tangent_imag - - -@triton.jit -def _two_pool_transverse_adjoint_jvp( - r2_free, - d_r2_free, - r2_bound, - d_r2_bound, - exchange, - d_exchange, - bound, - d_bound, - free, - d_free, - shift_hz, - d_shift_hz, - dt, - d_dt, - attenuation, - d_attenuation, - bar_e11, - bar_e12, - bar_e21, - bar_e22, -): - """The reverse sweep of :func:`_two_pool_transverse_step_jvp`. - - Every step from the four generator entries to the four operator entries is - holomorphic, so the sweep is the longitudinal one with complex numbers in - place of real ones and no conjugates along the way. That holds because the - cotangents arrive as row covectors -- ``bar_e`` is the number with ``dL = - Re(bar_e de)`` -- and only where a complex intermediate meets one of the - real inputs is a real part taken. - - Takes the four cotangents as dual complex quadruples and returns the seven - real gradients, each as a value and a tangent. - """ - zero = 0.0 * dt - kab = exchange * bound - d_kab = d_exchange * bound + exchange * d_bound - kba = exchange * free - d_kba = d_exchange * free + exchange * d_free - turn = -2.0 * 3.141592653589793 - l11 = ( - (-kab - r2_free) * dt, - zero, - (-d_kab - d_r2_free) * dt + (-kab - r2_free) * d_dt, - zero, - ) - l12 = (kba * dt, zero, d_kba * dt + kba * d_dt, zero) - l21 = (kab * dt, zero, d_kab * dt + kab * d_dt, zero) - l22 = ( - (-kba - r2_bound) * dt, - turn * (shift_hz * dt), - (-d_kba - d_r2_bound) * dt + (-kba - r2_bound) * d_dt, - turn * (d_shift_hz * dt + shift_hz * d_dt), - ) - - half_trace = _dual_weigh(_dual_add(l11, l22), 0.5) - half_gap = _dual_weigh(_dual_subtract(l11, l22), 0.5) - square = _dual_add(_dual_product(half_gap, half_gap), _dual_product(l12, l21)) - delta = _complex_sqrt_jvp(square[0], square[1], square[2], square[3]) - upper = _complex_exp_jvp(*_dual_add(half_trace, delta)) - lower = _complex_exp_jvp(*_dual_subtract(half_trace, delta)) - plain = _complex_exp_jvp(half_trace[0], half_trace[1], half_trace[2], half_trace[3]) - cosine = _dual_weigh(_dual_add(upper, lower), 0.5) - - turning = square[0] * square[0] + square[1] * square[1] > 1e-24 - # Off the branch the reciprocal is taken at one instead, so a discriminant - # at the origin never divides anything the series answer then discards. - guarded = ( - tl.where(turning, delta[0], 1.0), - tl.where(turning, delta[1], 0.0), - tl.where(turning, delta[2], 0.0), - tl.where(turning, delta[3], 0.0), - ) - inverse = _dual_reciprocal(guarded) - divided = _dual_product(_dual_weigh(_dual_subtract(upper, lower), 0.5), inverse) - square2 = _dual_product(square, square) - poly = ( - 1.0 + square[0] / 6.0 + square2[0] / 120.0, - square[1] / 6.0 + square2[1] / 120.0, - square[2] / 6.0 + square2[2] / 120.0, - square[3] / 6.0 + square2[3] / 120.0, - ) - series = _dual_product(plain, poly) - scale = ( - tl.where(turning, divided[0], series[0]), - tl.where(turning, divided[1], series[1]), - tl.where(turning, divided[2], series[2]), - tl.where(turning, divided[3], series[3]), - ) - - off = _dual_product(scale, half_gap) - bare_11 = _dual_add(cosine, off) - bare_12 = _dual_product(scale, l12) - bare_21 = _dual_product(scale, l21) - bare_22 = _dual_subtract(cosine, off) - - bar_attenuation = _dual_sum( - _dual_product(bar_e11, bare_11), - _dual_product(bar_e12, bare_12), - _dual_product(bar_e21, bare_21), - _dual_product(bar_e22, bare_22), - ) - scaled_11 = _dual_scale(attenuation, d_attenuation, *bar_e11) - scaled_12 = _dual_scale(attenuation, d_attenuation, *bar_e12) - scaled_21 = _dual_scale(attenuation, d_attenuation, *bar_e21) - scaled_22 = _dual_scale(attenuation, d_attenuation, *bar_e22) - - diagonal = _dual_subtract(scaled_11, scaled_22) - bar_cosine = _dual_add(scaled_11, scaled_22) - bar_scale = _dual_add( - _dual_product(diagonal, half_gap), - _dual_add(_dual_product(scaled_12, l12), _dual_product(scaled_21, l21)), - ) - bar_half_gap = _dual_product(scale, diagonal) - bar_l12 = _dual_product(scale, scaled_12) - bar_l21 = _dual_product(scale, scaled_21) - - series_trace = _dual_add( - _dual_product(bar_cosine, cosine), _dual_product(bar_scale, scale) - ) - series_square = _dual_product( - plain, - _dual_add( - _dual_product( - bar_cosine, - ( - 0.5 + square[0] / 12.0, - square[1] / 12.0, - square[2] / 12.0, - square[3] / 12.0, - ), - ), - _dual_product( - bar_scale, - ( - 0.16666666666666666 + square[0] / 60.0, - square[1] / 60.0, - square[2] / 60.0, - square[3] / 60.0, - ), - ), - ), - ) - bar_upper = _dual_weigh( - _dual_add(bar_cosine, _dual_product(bar_scale, inverse)), 0.5 - ) - bar_lower = _dual_weigh( - _dual_subtract(bar_cosine, _dual_product(bar_scale, inverse)), 0.5 - ) - split_trace = _dual_add( - _dual_product(bar_upper, upper), _dual_product(bar_lower, lower) - ) - bar_delta = _dual_subtract( - _dual_subtract( - _dual_product(bar_upper, upper), _dual_product(bar_lower, lower) - ), - _dual_product(_dual_product(bar_scale, scale), inverse), - ) - split_square = _dual_weigh(_dual_product(bar_delta, inverse), 0.5) - bar_half_trace = ( - tl.where(turning, split_trace[0], series_trace[0]), - tl.where(turning, split_trace[1], series_trace[1]), - tl.where(turning, split_trace[2], series_trace[2]), - tl.where(turning, split_trace[3], series_trace[3]), - ) - bar_square = ( - tl.where(turning, split_square[0], series_square[0]), - tl.where(turning, split_square[1], series_square[1]), - tl.where(turning, split_square[2], series_square[2]), - tl.where(turning, split_square[3], series_square[3]), - ) - - bar_half_gap = _dual_add( - bar_half_gap, _dual_weigh(_dual_product(bar_square, half_gap), 2.0) - ) - bar_l12 = _dual_add(bar_l12, _dual_product(bar_square, l21)) - bar_l21 = _dual_add(bar_l21, _dual_product(bar_square, l12)) - bar_l11 = _dual_weigh(_dual_add(bar_half_trace, bar_half_gap), 0.5) - bar_l22 = _dual_weigh(_dual_subtract(bar_half_trace, bar_half_gap), 0.5) - - bar_kab = _dual_scale(dt, d_dt, *_dual_subtract(bar_l21, bar_l11)) - bar_kba = _dual_scale(dt, d_dt, *_dual_subtract(bar_l12, bar_l22)) - slope_22 = ( - -kba - r2_bound, - turn * shift_hz, - -d_kba - d_r2_bound, - turn * d_shift_hz, - ) - bar_dt = _dual_sum( - _dual_scale(-kab - r2_free, -d_kab - d_r2_free, *bar_l11), - _dual_scale(kba, d_kba, *bar_l12), - _dual_scale(kab, d_kab, *bar_l21), - _dual_product(slope_22, bar_l22), - ) - - r2_free_bar = _dual_scale(-dt, -d_dt, *bar_l11) - r2_bound_bar = _dual_scale(-dt, -d_dt, *bar_l22) - exchange_bar = _dual_add( - _dual_scale(bound, d_bound, *bar_kab), - _dual_scale(free, d_free, *bar_kba), - ) - bound_bar = _dual_scale(exchange, d_exchange, *bar_kab) - free_bar = _dual_scale(exchange, d_exchange, *bar_kba) - shift_bar = _dual_scale(turn * dt, turn * d_dt, *_dual_times_i(*bar_l22)) - return ( - r2_free_bar[0], - r2_free_bar[2], - r2_bound_bar[0], - r2_bound_bar[2], - exchange_bar[0], - exchange_bar[2], - bound_bar[0], - bound_bar[2], - free_bar[0], - free_bar[2], - shift_bar[0], - shift_bar[2], - bar_dt[0], - bar_dt[2], - bar_attenuation[0], - bar_attenuation[2], - ) - - -@triton.jit -def _store_pair_cotangent( - grad_value, - grad_tangent, - pair_index, - event_base, - event, - atom, - atom_count, - turning, - mask, - state_mask, - grad_a, - grad_b, -): - """Send the cotangent on one pulse's rotation to its row. - - Summed over the dephasing orders first: the pair multiplies every one of - them, so what reaches the row is the sum. The value plane is the adjoint - and the tangent plane its own derivative, which is the split every other - gradient here takes. - """ - row = tl.load(pair_index + event_base + event).to(tl.int64) - entry = (row * atom_count + atom) * 4 - # The block is padded to a power of two, and the orders past the last one - # carry whatever the sweep left there -- so the sum is taken over the - # orders that exist rather than over the block. - keep = turning & state_mask - tl.atomic_add( - grad_value + entry + 0, - tl.sum(tl.where(keep, grad_a[0], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_value + entry + 1, - tl.sum(tl.where(keep, grad_a[1], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_value + entry + 2, - tl.sum(tl.where(keep, grad_b[0], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_value + entry + 3, - tl.sum(tl.where(keep, grad_b[1], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_tangent + entry + 0, - tl.sum(tl.where(keep, grad_a[2], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_tangent + entry + 1, - tl.sum(tl.where(keep, grad_a[3], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_tangent + entry + 2, - tl.sum(tl.where(keep, grad_b[2], 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_tangent + entry + 3, - tl.sum(tl.where(keep, grad_b[3], 0.0), axis=1)[:, None], - mask=mask, - ) - - -@triton.jit -def _store_pair_gradient( - grad_pair, - pair_index, - event_base, - event, - atom, - atom_count, - turning, - mask, - state_mask, - grad_ar, - grad_ai, - grad_br, - grad_bi, -): - """Send the cotangent on one pulse's rotation to its row. - - Summed over the dephasing orders first: the pair multiplies every one of - them, so what reaches the row is the sum. The block is padded to a power of - two and the orders past the last carry whatever the sweep left there, so - the sum is taken over the orders that exist rather than over the block. - """ - row = tl.load(pair_index + event_base + event).to(tl.int64) - entry = (row * atom_count + atom) * 4 - keep = turning & state_mask - tl.atomic_add( - grad_pair + entry + 0, - tl.sum(tl.where(keep, grad_ar, 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_pair + entry + 1, - tl.sum(tl.where(keep, grad_ai, 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_pair + entry + 2, - tl.sum(tl.where(keep, grad_br, 0.0), axis=1)[:, None], - mask=mask, - ) - tl.atomic_add( - grad_pair + entry + 3, - tl.sum(tl.where(keep, grad_bi, 0.0), axis=1)[:, None], - mask=mask, - ) - - -@triton.jit -def _dynamic_pair_dual_at( - pairs, - pair_direction, - pair_index, - event_base, - event, - atom, - atom_count, - mask, - phi_value, - phi_tangent, - directed: tl.constexpr, -): - """The rotation and the direction along it, with the phase applied. - - Shaped exactly as :func:`_profiled_pair_dual` returns, so the spinor - operator and its adjoint read one from the other without knowing which - they were handed. A pass that follows no direction holds the rotation - still, and ``directed`` keeps the read for one out of the kernel. - """ - held = _dynamic_pair_at( - pairs, pair_index, event_base, event, atom, atom_count, mask - ) - still = held[0] * 0.0 - moved = (still, still, still, still) - if directed: - moved = _dynamic_pair_at( - pair_direction, - pair_index, - event_base, - event, - atom, - atom_count, - mask, - ) - a = (held[0], held[1], moved[0], moved[1]) - b = (held[2], held[3], moved[2], moved[3]) - turn = _dual_polar(-phi_value, -phi_tangent) - return a, _dual_product(b, turn) - - -@triton.jit -def _profiled_pair_dual( - profile, - row, - alpha_value, - alpha_tangent, - phi_value, - phi_tangent, - bins, - step, -): - """The pair a shaped pulse turns through, and its slope, as duals. - - The flip angle carries the tangent into the table, so the pair's tangent is - the stored slope and the slope's own tangent is the segment's curvature. - The RF phase turns the axis once the pair is out, and so reaches ``b``. - """ - read = _profile_pair_curve(profile, row, alpha_value, bins, step) - a = (read[0], read[3], read[1] * alpha_tangent, read[4] * alpha_tangent) - slope_a = (read[1], read[4], read[2] * alpha_tangent, read[5] * alpha_tangent) - b = (read[6], read[9], read[7] * alpha_tangent, read[10] * alpha_tangent) - slope_b = (read[7], read[10], read[8] * alpha_tangent, read[11] * alpha_tangent) - turn = _dual_polar(-phi_value, -phi_tangent) - return a, _dual_product(b, turn), slope_a, _dual_product(slope_b, turn) - - -@triton.jit -def _spinor_adjoint_dual(a, b, sp, sm, rz, pb, mb, zb): - """The spinor rotation's adjoint, on dual numbers. - - Returns the cotangent on the Cayley-Klein pair and the three state - cotangents sent back through the conjugate transpose. Every entry of the - matrix is a product of two factors drawn from the pair and its conjugate, - so the pair's two Wirtinger halves are linear in the outer product of the - seed with the state the rotation acted on -- which is why this is a closed - form rather than a differentiated matrix. - """ - t00, t01, t02, t10, t11, t12, t20, t21, t22 = _spinor_coefficients( - a[0], a[1], b[0], b[1], a[2], a[3], b[2], b[3] - ) - - conj_pb = _dual_conj(pb) - conj_mb = _dual_conj(mb) - conj_zb = _dual_conj(zb) - m00 = _dual_product(conj_pb, sp) - m01 = _dual_product(conj_pb, sm) - m02 = _dual_product(conj_pb, rz) - m10 = _dual_product(conj_mb, sp) - m11 = _dual_product(conj_mb, sm) - m12 = _dual_product(conj_mb, rz) - m20 = _dual_product(conj_zb, sp) - m21 = _dual_product(conj_zb, sm) - m22 = _dual_product(conj_zb, rz) - - conj_a = _dual_conj(a) - conj_b = _dual_conj(b) - holding_conj_a = _dual_sum( - _dual_weigh(_dual_product(a, m11), 2.0), - _dual_weigh(_dual_product(b, m12), -2.0), - _dual_product(conj_b, m21), - _dual_product(conj_a, m22), - ) - holding_a = _dual_sum( - _dual_weigh(_dual_product(conj_a, m00), 2.0), - _dual_weigh(_dual_product(conj_b, m02), -2.0), - _dual_product(b, m20), - _dual_product(a, m22), - ) - holding_conj_b = _dual_sum( - _dual_weigh(_dual_product(b, m10), -2.0), - _dual_weigh(_dual_product(a, m12), -2.0), - _dual_product(conj_a, m20), - _dual_weigh(_dual_product(conj_b, m22), -1.0), - ) - holding_b = _dual_sum( - _dual_weigh(_dual_product(conj_b, m01), -2.0), - _dual_weigh(_dual_product(conj_a, m02), -2.0), - _dual_product(a, m21), - _dual_weigh(_dual_product(b, m22), -1.0), - ) - zero = _dual_weigh(m00, 0.0) - grad_a = _dual_sum(_dual_conj(holding_conj_a), holding_a, zero, zero) - grad_b = _dual_sum(_dual_conj(holding_conj_b), holding_b, zero, zero) - - next_pb = _dual_sum( - _dual_product(_dual_conj(t00), pb), - _dual_product(_dual_conj(t10), mb), - _dual_product(_dual_conj(t20), zb), - zero, - ) - next_mb = _dual_sum( - _dual_product(_dual_conj(t01), pb), - _dual_product(_dual_conj(t11), mb), - _dual_product(_dual_conj(t21), zb), - zero, - ) - next_zb = _dual_sum( - _dual_product(_dual_conj(t02), pb), - _dual_product(_dual_conj(t12), mb), - _dual_product(_dual_conj(t22), zb), - zero, - ) - return grad_a, grad_b, next_pb, next_mb, next_zb - - -@triton.jit -def _spinor_coefficients(ar, ai, br, bi, dar, dai, dbr, dbi): - """The rotation's nine coefficients and their tangents. - - Every entry is a product of two factors drawn from the pair and its - conjugate, so five products carry all nine: ``a^2``, ``b^2``, ``a b``, - ``a conj(b)`` and the norm difference. - """ - aa_r = ar * ar - ai * ai - aa_i = 2.0 * ar * ai - daa_r = 2.0 * (ar * dar - ai * dai) - daa_i = 2.0 * (dar * ai + ar * dai) - - bb_r = br * br - bi * bi - bb_i = 2.0 * br * bi - dbb_r = 2.0 * (br * dbr - bi * dbi) - dbb_i = 2.0 * (dbr * bi + br * dbi) - - ab_r = ar * br - ai * bi - ab_i = ar * bi + ai * br - dab_r = dar * br + ar * dbr - dai * bi - ai * dbi - dab_i = dar * bi + ar * dbi + dai * br + ai * dbr - - cross_r = ar * br + ai * bi - cross_i = ar * bi - ai * br - dcross_r = dar * br + ar * dbr + dai * bi + ai * dbi - dcross_i = dar * bi + ar * dbi - dai * br - ai * dbr - - t22 = ar * ar + ai * ai - br * br - bi * bi - dt22 = 2.0 * (ar * dar + ai * dai - br * dbr - bi * dbi) - - return ( - (aa_r, -aa_i, daa_r, -daa_i), - (-bb_r, bb_i, -dbb_r, dbb_i), - (-2.0 * ab_r, 2.0 * ab_i, -2.0 * dab_r, 2.0 * dab_i), - (-bb_r, -bb_i, -dbb_r, -dbb_i), - (aa_r, aa_i, daa_r, daa_i), - (-2.0 * ab_r, -2.0 * ab_i, -2.0 * dab_r, -2.0 * dab_i), - (cross_r, cross_i, dcross_r, dcross_i), - (cross_r, -cross_i, dcross_r, -dcross_i), - (t22, 0.0 * t22, dt22, 0.0 * dt22), - ) - - -@triton.jit -def _rotate_spinor_dual( - ar, - ai, - br, - bi, - dar, - dai, - dbr, - dbi, - fp_r, - fp_i, - fm_r, - fm_i, - z_r, - z_i, - dfp_r, - dfp_i, - dfm_r, - dfm_i, - dz_r, - dz_i, -): - """The spinor rotation carrying a forward-mode tangent. - - Both the states and the pair naming the rotation move, so the tangent is - ``T dx + dT x``. - """ - t00, t01, t02, t10, t11, t12, t20, t21, t22 = _spinor_coefficients( - ar, ai, br, bi, dar, dai, dbr, dbi - ) - - out_pr, out_pi = _dual_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - out_mr, out_mi = _dual_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - out_zr, out_zi = _dual_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - - dpr, dpi = _dual_row(t00, t01, t02, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i) - dmr, dmi = _dual_row(t10, t11, t12, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i) - dzr, dzi = _dual_row(t20, t21, t22, dfp_r, dfp_i, dfm_r, dfm_i, dz_r, dz_i) - tpr, tpi = _tangent_row(t00, t01, t02, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - tmr, tmi = _tangent_row(t10, t11, t12, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - tzr, tzi = _tangent_row(t20, t21, t22, fp_r, fp_i, fm_r, fm_i, z_r, z_i) - - return ( - out_pr, - out_pi, - out_mr, - out_mi, - out_zr, - out_zi, - dpr + tpr, - dpi + tpi, - dmr + tmr, - dmi + tmi, - dzr + tzr, - dzi + tzi, - ) - - -@triton.jit -def _dual_row(first, second, third, fp_r, fp_i, fm_r, fm_i, z_r, z_i): - """One row of the rotation applied to the states, values only.""" - real = ( - first[0] * fp_r - - first[1] * fp_i - + second[0] * fm_r - - second[1] * fm_i - + third[0] * z_r - - third[1] * z_i - ) - imag = ( - first[0] * fp_i - + first[1] * fp_r - + second[0] * fm_i - + second[1] * fm_r - + third[0] * z_i - + third[1] * z_r - ) - return real, imag - - -@triton.jit -def _tangent_row(first, second, third, fp_r, fp_i, fm_r, fm_i, z_r, z_i): - """The same row built from the coefficients' tangents instead.""" - real = ( - first[2] * fp_r - - first[3] * fp_i - + second[2] * fm_r - - second[3] * fm_i - + third[2] * z_r - - third[3] * z_i - ) - imag = ( - first[2] * fp_i - + first[3] * fp_r - + second[2] * fm_i - + second[3] * fm_r - + third[2] * z_i - + third[3] * z_r - ) - return real, imag - - -@triton.jit -def _spinor_adjoint( - ar, ai, br, bi, spr, spi, smr, smi, rzr, rzi, pbr, pbi, mbr, mbi, zbr, zbi -): - """The spinor rotation's adjoint, carrying no forward direction. - - Returns the cotangent on the Cayley-Klein pair and the three state - cotangents sent back through the conjugate transpose. Every entry of the - matrix is a product of two factors drawn from the pair and its conjugate, - so the pair's two Wirtinger halves are linear in the outer product of the - seed with the state the rotation acted on -- a closed form rather than a - differentiated matrix. - """ - aa_r = ar * ar - ai * ai - aa_i = 2.0 * ar * ai - bb_r = br * br - bi * bi - bb_i = 2.0 * br * bi - ab_r = ar * br - ai * bi - ab_i = ar * bi + ai * br - cross_r = ar * br + ai * bi - cross_i = ar * bi - ai * br - t00_r, t00_i = aa_r, -aa_i - t01_r, t01_i = -bb_r, bb_i - t02_r, t02_i = -2.0 * ab_r, 2.0 * ab_i - t10_r, t10_i = -bb_r, -bb_i - t11_r, t11_i = aa_r, aa_i - t12_r, t12_i = -2.0 * ab_r, -2.0 * ab_i - t20_r, t20_i = cross_r, cross_i - t21_r, t21_i = cross_r, -cross_i - t22 = ar * ar + ai * ai - br * br - bi * bi - - # ``m[i][j] = conj(seed_i) * state_j``: the outer product the pair's - # derivative is linear in. - m00 = _complex_mul(pbr, -pbi, spr, spi) - m01 = _complex_mul(pbr, -pbi, smr, smi) - m02 = _complex_mul(pbr, -pbi, rzr, rzi) - m10 = _complex_mul(mbr, -mbi, spr, spi) - m11 = _complex_mul(mbr, -mbi, smr, smi) - m12 = _complex_mul(mbr, -mbi, rzr, rzi) - m20 = _complex_mul(zbr, -zbi, spr, spi) - m21 = _complex_mul(zbr, -zbi, smr, smi) - m22 = _complex_mul(zbr, -zbi, rzr, rzi) - - hca = _complex_mul(ar, ai, m11[0], m11[1]) - hcb = _complex_mul(br, bi, m12[0], m12[1]) - hcc = _complex_mul(br, -bi, m21[0], m21[1]) - hcd = _complex_mul(ar, -ai, m22[0], m22[1]) - holding_conj_a_r = 2.0 * hca[0] - 2.0 * hcb[0] + hcc[0] + hcd[0] - holding_conj_a_i = 2.0 * hca[1] - 2.0 * hcb[1] + hcc[1] + hcd[1] - - ha = _complex_mul(ar, -ai, m00[0], m00[1]) - hb = _complex_mul(br, -bi, m02[0], m02[1]) - hc = _complex_mul(br, bi, m20[0], m20[1]) - hd = _complex_mul(ar, ai, m22[0], m22[1]) - holding_a_r = 2.0 * ha[0] - 2.0 * hb[0] + hc[0] + hd[0] - holding_a_i = 2.0 * ha[1] - 2.0 * hb[1] + hc[1] + hd[1] - - ka = _complex_mul(br, bi, m10[0], m10[1]) - kb = _complex_mul(ar, ai, m12[0], m12[1]) - kc = _complex_mul(ar, -ai, m20[0], m20[1]) - kd = _complex_mul(br, -bi, m22[0], m22[1]) - holding_conj_b_r = -2.0 * ka[0] - 2.0 * kb[0] + kc[0] - kd[0] - holding_conj_b_i = -2.0 * ka[1] - 2.0 * kb[1] + kc[1] - kd[1] - - la = _complex_mul(br, -bi, m01[0], m01[1]) - lb = _complex_mul(ar, -ai, m02[0], m02[1]) - lc = _complex_mul(ar, ai, m21[0], m21[1]) - ld = _complex_mul(br, bi, m22[0], m22[1]) - holding_b_r = -2.0 * la[0] - 2.0 * lb[0] + lc[0] - ld[0] - holding_b_i = -2.0 * la[1] - 2.0 * lb[1] + lc[1] - ld[1] - - grad_a_r = holding_conj_a_r + holding_a_r - grad_a_i = -holding_conj_a_i + holding_a_i - grad_b_r = holding_conj_b_r + holding_b_r - grad_b_i = -holding_conj_b_i + holding_b_i - - n0 = _complex_mul(t00_r, -t00_i, pbr, pbi) - n1 = _complex_mul(t10_r, -t10_i, mbr, mbi) - n2 = _complex_mul(t20_r, -t20_i, zbr, zbi) - next_pr, next_pi = n0[0] + n1[0] + n2[0], n0[1] + n1[1] + n2[1] - n0 = _complex_mul(t01_r, -t01_i, pbr, pbi) - n1 = _complex_mul(t11_r, -t11_i, mbr, mbi) - n2 = _complex_mul(t21_r, -t21_i, zbr, zbi) - next_mr, next_mi = n0[0] + n1[0] + n2[0], n0[1] + n1[1] + n2[1] - n0 = _complex_mul(t02_r, -t02_i, pbr, pbi) - n1 = _complex_mul(t12_r, -t12_i, mbr, mbi) - next_zr = n0[0] + n1[0] + t22 * zbr - next_zi = n0[1] + n1[1] + t22 * zbi - return ( - grad_a_r, - grad_a_i, - grad_b_r, - grad_b_i, - next_pr, - next_pi, - next_mr, - next_mi, - next_zr, - next_zi, - ) - - -@triton.jit -def _rotate_spinor(ar, ai, br, bi, fp_r, fp_i, fm_r, fm_i, z_r, z_i): - """The rotation named by its Cayley-Klein pair, applied to the states. - - T = [ conj(a)^2 -conj(b)^2 -2 conj(a b) ] - [ -b^2 a^2 -2 a b ] - [ conj(a) b a conj(b) |a|^2-|b|^2 ] - """ - aa_r = ar * ar - ai * ai - aa_i = 2.0 * ar * ai - bb_r = br * br - bi * bi - bb_i = 2.0 * br * bi - ab_r = ar * br - ai * bi - ab_i = ar * bi + ai * br - - t00_r, t00_i = aa_r, -aa_i - t01_r, t01_i = -bb_r, bb_i - t02_r, t02_i = -2.0 * ab_r, 2.0 * ab_i - t10_r, t10_i = -bb_r, -bb_i - t11_r, t11_i = aa_r, aa_i - t12_r, t12_i = -2.0 * ab_r, -2.0 * ab_i - cross_r = ar * br + ai * bi - cross_i = ar * bi - ai * br - t20_r, t20_i = cross_r, cross_i - t21_r, t21_i = cross_r, -cross_i - t22 = ar * ar + ai * ai - br * br - bi * bi - - out_pr = ( - t00_r * fp_r - - t00_i * fp_i - + t01_r * fm_r - - t01_i * fm_i - + t02_r * z_r - - t02_i * z_i - ) - out_pi = ( - t00_r * fp_i - + t00_i * fp_r - + t01_r * fm_i - + t01_i * fm_r - + t02_r * z_i - + t02_i * z_r - ) - out_mr = ( - t10_r * fp_r - - t10_i * fp_i - + t11_r * fm_r - - t11_i * fm_i - + t12_r * z_r - - t12_i * z_i - ) - out_mi = ( - t10_r * fp_i - + t10_i * fp_r - + t11_r * fm_i - + t11_i * fm_r - + t12_r * z_i - + t12_i * z_r - ) - out_zr = t20_r * fp_r - t20_i * fp_i + t21_r * fm_r - t21_i * fm_i + t22 * z_r - out_zi = t20_r * fp_i + t20_i * fp_r + t21_r * fm_i + t21_i * fm_r + t22 * z_i - return out_pr, out_pi, out_mr, out_mi, out_zr, out_zi - - -@triton.jit -def _rotation_block( - a_value, - a_tangent, - b_value, - b_tangent, - c_value, - c_tangent, - d_value, - d_tangent, - p1r, - p1i, - p1tr, - p1ti, - p2r, - p2i, - p2tr, - p2ti, - pcr, - pci, - pctr, - pcti, -): - """Seven of the nine rotation coefficients; the rest follow by symmetry. - - ``t11`` repeats ``t00`` and ``t10`` is the conjugate of ``t01``, so the - caller derives those. Feeding ``(cos, sin)`` gives the rotation itself and - ``(sin, cos)`` rearranged gives its derivative in the flip angle, which is - why this is one routine rather than two. - """ - t00 = (a_value, 0.0 * a_value, a_tangent, 0.0 * a_tangent) - t01 = _dual_scale(b_value, b_tangent, p2r, p2i, p2tr, p2ti) - t02 = _dual_mul( - 0.0 * c_value, -c_value, 0.0 * c_tangent, -c_tangent, p1r, p1i, p1tr, p1ti - ) - t12 = _dual_mul( - 0.0 * c_value, c_value, 0.0 * c_tangent, c_tangent, pcr, pci, pctr, pcti - ) - t20 = _dual_mul( - 0.0 * c_value, - -0.5 * c_value, - 0.0 * c_tangent, - -0.5 * c_tangent, - pcr, - pci, - pctr, - pcti, - ) - t21 = _dual_mul( - 0.0 * c_value, - 0.5 * c_value, - 0.0 * c_tangent, - 0.5 * c_tangent, - p1r, - p1i, - p1tr, - p1ti, - ) - t22 = (d_value, 0.0 * d_value, d_tangent, 0.0 * d_tangent) - return t00, t01, t02, t12, t20, t21, t22 - - -@triton.jit -def _rotation_coefficients(a, b, c, d, p1r, p1i, p2r, p2i, pcr, pci): - """Seven of the nine rotation coefficients; the rest follow by symmetry. - - ``t11`` repeats ``t00`` and ``t10`` is the conjugate of ``t01``, so the - caller derives those. Feeding ``(cos, sin)`` gives the rotation itself and - ``(sin, cos)`` rearranged gives its derivative in the flip angle, which is - why this is one routine rather than two. - """ - t00 = (a, 0.0 * a) - t01 = (b * p2r, b * p2i) - t02 = _complex_mul(0.0 * c, -c, p1r, p1i) - t12 = _complex_mul(0.0 * c, c, pcr, pci) - t20 = _complex_mul(0.0 * c, -0.5 * c, pcr, pci) - t21 = _complex_mul(0.0 * c, 0.5 * c, p1r, p1i) - t22 = (d, 0.0 * d) - return t00, t01, t02, t12, t20, t21, t22 - - -@triton.jit( - do_not_specialize=["state_count", "locations", "profile_bins", "lineshape_bins"] -) -def _epg_vjp_kernel( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - lineshape, - profile, - profile_index, - pairs, - pair_index, - duration_row, - pool_table, - pool_bars, - pool_durations, - row_count, - grad_pair, - grad_output_real, - grad_output_imag, - grad_tissue, - grad_flip, - grad_phase, - grad_duration, - trajectory_r, - trajectory_i, - problem_base, - problem_end, - atom_count, - train_count, - event_count, - output_count, - flow_scale, - washout_scale, - shim_rows, - profile_step, - lineshape_step, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shimmed: tl.constexpr, - locations, - profiled: tl.constexpr, - profile_bins, - dynamic: tl.constexpr, - broadened: tl.constexpr, - lineshape_bins, - pools: tl.constexpr, - narrow: tl.constexpr, - tabulated: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - recording: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = problem_base + tl.program_id(0) * problems - problem = problem + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < problem_end - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - local = problem - problem_base - record_stride = ( - 7 if pools == 3 else (6 if pools == 2 else (4 if pools == 1 else 3)) - ) * state_count - trajectory = local * event_count * record_stride + state - minus_plane = state_count - long_plane = 2 * state_count - bound_plane = 3 * state_count - bplus_plane = 4 * state_count - bminus_plane = 5 * state_count - semisolid_plane = 6 * state_count - - empty = tl.zeros((problems, block_states), tl.float32) - pvr = empty - pvi = empty - mvr = empty - mvi = empty - zvr = empty + tl.where(state == 0, 1.0, 0.0) - zvi = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_b1_phase = 0.0 - atom_b0 = 0.0 - if off_axis: - atom_b1_phase = tl.load(b1_phase + scalar_atom, mask=active_atom, other=0.0) - atom_b0 = tl.load(b0 + scalar_atom, mask=active_atom, other=0.0) - atom_inv = 1.0 - if inverting: - atom_inv = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - atom_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - atom_flow = 0.0 - direction = 0.0 - atom_washout = 0.0 - if moving: - atom_velocity = tl.load(velocity + scalar_atom, mask=active_atom, other=0.0) - atom_flow = atom_velocity * flow_scale - # |v| has no derivative at the origin, so a still voxel contributes - # none. - direction = (atom_velocity > 0.0).to(tl.float32) - (atom_velocity < 0.0).to( - tl.float32 - ) - atom_washout = tl.abs(atom_velocity) * washout_scale - order = state.to(tl.float32) - longitudinal_weight = order * order - transverse_weight = longitudinal_weight + order + 0.3333333333333333 - r1_value = 1000.0 / atom_t1 - r2_value = 1000.0 / atom_t2 - - location = atom % locations - # A semisolid pool rides along as a plane of its own: the pulse deposits - # into it and it exchanges with the free water, so the reverse sweep cannot - # replay it from the free pool's. - atom_bound = 0.0 - atom_exchange = 0.0 - atom_t1b = 1.0 - atom_t2b = 1.0 - atom_shift = 0.0 - r1b_value = 0.0 - r2b_value = 0.0 - atom_semisolid = 0.0 - atom_semisolid_exchange = 0.0 - atom_t1c = 1.0 - r1c_value = 0.0 - atom_free = 1.0 - poolvr = empty - poolvi = empty - bpvr = empty - bpvi = empty - bmvr = empty - bmvi = empty - semivr = empty - semivi = empty - if pools == 1: - atom_bound = tl.load(bound_fraction + scalar_atom, mask=active_atom, other=0.0) - atom_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - atom_t1b = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - r1b_value = 1000.0 / atom_t1b - if pools == 2 or pools == 3: - atom_bound = tl.load(pool_b_fraction + scalar_atom, mask=active_atom, other=0.0) - atom_exchange = tl.load( - pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - atom_t1b = tl.load(t1_pool_b + scalar_atom, mask=active_atom, other=1.0) - r1b_value = 1000.0 / atom_t1b - atom_t2b = tl.load(t2_pool_b + scalar_atom, mask=active_atom, other=1.0) - r2b_value = 1000.0 / atom_t2b - atom_shift = tl.load(pool_b_shift + scalar_atom, mask=active_atom, other=0.0) - if pools == 3: - # The semisolid pool takes the rows a run with it alone would take, so - # the two second pools never contend for one. - atom_semisolid = tl.load( - bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_semisolid_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - atom_t1c = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - r1c_value = 1000.0 / atom_t1c - semivr = empty + tl.where(state == 0, atom_semisolid + 0.0, 0.0) - if pools > 0: - # The fractions split the equilibrium at t = 0. - atom_free = 1.0 - atom_bound - atom_semisolid - zvr = empty + tl.where(state == 0, atom_free, 0.0) - poolvr = empty + tl.where(state == 0, atom_bound + 0.0, 0.0) - event_base = train * event_count - # The forward half records the trajectory the reverse half walks back, - # and the two are launched separately: each compiles the sweep it is - # asked for and no more. - if recording: - for event in range(0, event_count): - slot = trajectory + event * record_stride - tl.store(trajectory_r + slot, pvr, mask=state_mask) - tl.store(trajectory_i + slot, pvi, mask=state_mask) - tl.store(trajectory_r + slot + minus_plane, mvr, mask=state_mask) - tl.store(trajectory_i + slot + minus_plane, mvi, mask=state_mask) - tl.store(trajectory_r + slot + long_plane, zvr, mask=state_mask) - tl.store(trajectory_i + slot + long_plane, zvi, mask=state_mask) - if pools > 0: - tl.store(trajectory_r + slot + bound_plane, poolvr, mask=state_mask) - tl.store(trajectory_i + slot + bound_plane, poolvi, mask=state_mask) - if pools == 2 or pools == 3: - tl.store(trajectory_r + slot + bplus_plane, bpvr, mask=state_mask) - tl.store(trajectory_i + slot + bplus_plane, bpvi, mask=state_mask) - tl.store(trajectory_r + slot + bminus_plane, bmvr, mask=state_mask) - tl.store(trajectory_i + slot + bminus_plane, bmvi, mask=state_mask) - if pools == 3: - tl.store(trajectory_r + slot + semisolid_plane, semivr, mask=state_mask) - tl.store(trajectory_i + slot + semisolid_plane, semivi, mask=state_mask) - - dt_value = _event_value( - duration, event_base, event, active_atom, single_train - ) - wout_value = 1.0 - if moving: - wout_value = _washout(atom_washout, dt_value) - e1_value = tl.exp(-r1_value * dt_value) * wout_value - e2_value = tl.exp(-r2_value * dt_value) * wout_value - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt_value, order) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value = 1.0 - e1_value - bare1_value = e1_value - bare2_value = e2_value - e1_value = bare1_value * damp_z - e2_value = bare2_value * damp_t - turn_t = 0.0 - szr, szi = 1.0, 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, dt_value, order) - szr, szi = tl.cos(turn_z), tl.sin(turn_z) - qr, qi = 1.0, 0.0 - if off_axis or moving: - angle_value = -2.0 * 3.141592653589793 * (atom_b0 * dt_value) + turn_t - qr, qi = tl.cos(angle_value), tl.sin(angle_value) - ovr, ovi = e2_value * qr, e2_value * qi - lvr, lvi = e1_value * szr, e1_value * szi - - if pools == 2 or pools == 3: - # With an exchanging pool the transverse relaxation sits inside the - # operator instead of in the scalar the free pool alone multiplies. - across = _two_pool_transverse_step_jvp( - r2_value, - 0.0, - r2b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - atom_free, - 0.0, - atom_shift, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - ) - a11r, a11i = across[0], across[1] - a12r, a12i = across[2], across[3] - a21r, a21i = across[4], across[5] - a22r, a22i = across[6], across[7] - carr, cari = damp_t * qr, damp_t * qi - f11r, f11i = _complex_mul(a11r, a11i, pvr, pvi) - f12r, f12i = _complex_mul(a12r, a12i, bpvr, bpvi) - g21r, g21i = _complex_mul(a21r, a21i, pvr, pvi) - g22r, g22i = _complex_mul(a22r, a22i, bpvr, bpvi) - # ``F-`` takes the conjugate of the operator entry by entry, not its - # transpose: it is the conjugate state following the conjugate map. - h11r, h11i = _complex_mul(a11r, -a11i, mvr, mvi) - h12r, h12i = _complex_mul(a12r, -a12i, bmvr, bmvi) - k21r, k21i = _complex_mul(a21r, -a21i, mvr, mvi) - k22r, k22i = _complex_mul(a22r, -a22i, bmvr, bmvi) - pvr, pvi = _complex_mul(f11r + f12r, f11i + f12i, carr, cari) - bpvr, bpvi = _complex_mul(g21r + g22r, g21i + g22i, carr, cari) - mvr, mvi = _complex_mul(h11r + h12r, h11i + h12i, carr, -cari) - bmvr, bmvi = _complex_mul(k21r + k22r, k21i + k22i, carr, -cari) - else: - pvr, pvi = _complex_mul(ovr, ovi, pvr, pvi) - mvr, mvi = _complex_mul(ovr, -ovi, mvr, mvi) - if pools == 3: - # Three pools mix through a 3x3 formed once for the interval; each - # second pool exchanges with the free water and not with the other. - nil = 0.0 * dt_value - hold_value = wout_value + nil - if tabulated: - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - ) = _three_pool_from_table( - pool_table, - tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ), - atom, - atom_count, - active_atom, - hold_value, - atom_free, - atom_bound, - atom_semisolid, - ) - else: - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - _dw11, - _dw12, - _dw13, - _dw21, - _dw22, - _dw23, - _dw31, - _dw32, - _dw33, - _dgf, - _dgb, - _dgs, - ) = _three_pool_step_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - dt_value, - nil, - hold_value, - nil, - narrow, - ) - spin_r, spin_i = damp_z * szr, damp_z * szi - mix_fr = w11 * zvr + w12 * poolvr + w13 * semivr - mix_fi = w11 * zvi + w12 * poolvi + w13 * semivi - mix_br = w21 * zvr + w22 * poolvr + w23 * semivr - mix_bi = w21 * zvi + w22 * poolvi + w23 * semivi - mix_cr = w31 * zvr + w32 * poolvr + w33 * semivr - mix_ci = w31 * zvi + w32 * poolvi + w33 * semivi - zvr, zvi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - poolvr, poolvi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - semivr, semivi = _complex_mul(spin_r, spin_i, mix_cr, mix_ci) - zvr += tl.where(state == 0, grow_free, 0.0) - poolvr += tl.where(state == 0, grow_pool_b, 0.0) - semivr += tl.where(state == 0, grow_semisolid, 0.0) - elif pools > 0: - # The pools exchange while they relax, so the longitudinal step is a - # 2x2 the interval forms once and the per-order damping and turn - # multiply. Read from the dual helper with no direction to follow: - # what only its tangents reach, the compiler drops. - ( - pe11, - pe12, - pe21, - pe22, - prec_f, - prec_b, - _d11, - _d12, - _d21, - _d22, - _drf, - _drb, - ) = _two_pool_step_jvp( - r1_value, - 0.0, - r1b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - ) - spin_r, spin_i = damp_z * szr, damp_z * szi - mix_fr = pe11 * zvr + pe12 * poolvr - mix_fi = pe11 * zvi + pe12 * poolvi - mix_br = pe21 * zvr + pe22 * poolvr - mix_bi = pe21 * zvi + pe22 * poolvi - zvr, zvi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - poolvr, poolvi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - zvr += tl.where(state == 0, prec_f, 0.0) - poolvr += tl.where(state == 0, prec_b, 0.0) - else: - zvr, zvi = _complex_mul(lvr, lvi, zvr, zvi) - zvr += tl.where(state == 0, recovery_value, 0.0) - - event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - svr, svi, wvr, wvi = _shift( - pvr, pvi, mvr, mvi, state, state_mask, state_count - ) - pvr = tl.where(pre_shift, svr, pvr) - pvi = tl.where(pre_shift, svi, pvi) - mvr = tl.where(pre_shift, wvr, mvr) - mvi = tl.where(pre_shift, wvi, mvi) - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count - ) - bpvr = tl.where(pre_shift, svr, bpvr) - bpvi = tl.where(pre_shift, svi, bpvi) - bmvr = tl.where(pre_shift, wvr, bmvr) - bmvi = tl.where(pre_shift, wvi, bmvi) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - zvr = tl.where(invert, -atom_inv * zvr, zvr) - zvi = tl.where(invert, -atom_inv * zvi, zvi) - if pools == 2 or pools == 3: - # A chemically exchanging pool is free water and inverts like any - # other; a semisolid one is saturated instead, by the pulse's own - # saturation term. - poolvr = tl.where(invert, -atom_inv * poolvr, poolvr) - poolvi = tl.where(invert, -atom_inv * poolvi, poolvi) - - event_flip = _event_value( - flip, event_base, event, active_atom, single_train - ) - event_phase = _event_value( - phase, event_base, event, active_atom, single_train - ) - pulse_b1 = atom_b1 - pulse_b1_phase = atom_b1_phase - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse a row of its own. - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - if off_axis: - pulse_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - alpha_value = event_flip * pulse_b1 - phi_value = event_phase + pulse_b1_phase - if pools == 1 or pools == 3: - # The pool absorbs the power the pulse deposits, read at the offset - # the pulse is played less the voxel's own. - offset_value = tl.load(rf_frequency + event) - atom_b0 - shape_value, _shape_slope = _lineshape_at_slope( - lineshape, offset_value, lineshape_bins, lineshape_step - ) - event_saturation = tl.load(saturation + event) - power_value = event_saturation * alpha_value * alpha_value - absorbed_value = tl.exp(power_value * shape_value) - saturating = is_rf & ~is_inversion - if pools == 1: - poolvr = tl.where(saturating, absorbed_value * poolvr, poolvr) - poolvi = tl.where(saturating, absorbed_value * poolvi, poolvi) - else: - semivr = tl.where(saturating, absorbed_value * semivr, semivr) - semivi = tl.where(saturating, absorbed_value * semivi, semivi) - cos_value = tl.cos(alpha_value) - sin_value = tl.sin(alpha_value) - p1r, p1i = tl.cos(phi_value), tl.sin(phi_value) - p2r, p2i = _complex_mul(p1r, p1i, p1r, p1i) - t00, t01, t02, t12, t20, t21, t22 = _rotation_coefficients( - 0.5 * (1.0 + cos_value), - 0.5 * (1.0 - cos_value), - sin_value, - cos_value, - p1r, - p1i, - p2r, - p2i, - p1r, - -p1i, - ) - a0 = _complex_mul(t00[0], t00[1], pvr, pvi) - a1 = _complex_mul(t01[0], t01[1], mvr, mvi) - a2 = _complex_mul(t02[0], t02[1], zvr, zvi) - b0_ = _complex_mul(t01[0], -t01[1], pvr, pvi) - b1_ = _complex_mul(t00[0], t00[1], mvr, mvi) - b2 = _complex_mul(t12[0], t12[1], zvr, zvi) - c0 = _complex_mul(t20[0], t20[1], pvr, pvi) - c1 = _complex_mul(t21[0], t21[1], mvr, mvi) - c2 = _complex_mul(t22[0], t22[1], zvr, zvi) - - turned_pr = a0[0] + a1[0] + a2[0] - turned_pi = a0[1] + a1[1] + a2[1] - turned_mr = b0_[0] + b1_[0] + b2[0] - turned_mi = b0_[1] + b1_[1] + b2[1] - turned_zr = c0[0] + c1[0] + c2[0] - turned_zi = c0[1] + c1[1] + c2[1] - if profiled or dynamic: - if dynamic: - pair = _dynamic_pair_at( - pairs, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - shaped_ar, shaped_ai = pair[0], pair[1] - # The pair is integrated at zero RF phase, so the event's own - # phase turns the axis afterwards. - shaped_br, shaped_bi = _complex_mul(pair[2], pair[3], p1r, -p1i) - else: - shaped_ar, shaped_ai, shaped_br, shaped_bi = _profile_pair( - profile, - _table_row(profile_index, event, location, locations), - alpha_value, - profile_bins, - profile_step, - ) - shaped_br, shaped_bi = _complex_mul(shaped_br, shaped_bi, p1r, -p1i) - ( - turned_pr, - turned_pi, - turned_mr, - turned_mi, - turned_zr, - turned_zi, - ) = _rotate_spinor( - shaped_ar, - shaped_ai, - shaped_br, - shaped_bi, - pvr, - pvi, - mvr, - mvi, - zvr, - zvi, - ) - - rotate = is_rf & ~is_inversion - if pools == 2 or pools == 3: - # The same pulse, the same rotation. A chemical shift moves where a - # pool precesses, not what a pulse does to it. - e0 = _complex_mul(t00[0], t00[1], bpvr, bpvi) - e1_ = _complex_mul(t01[0], t01[1], bmvr, bmvi) - e2_ = _complex_mul(t02[0], t02[1], poolvr, poolvi) - f0 = _complex_mul(t01[0], -t01[1], bpvr, bpvi) - f1 = _complex_mul(t00[0], t00[1], bmvr, bmvi) - f2 = _complex_mul(t12[0], t12[1], poolvr, poolvi) - h0 = _complex_mul(t20[0], t20[1], bpvr, bpvi) - h1 = _complex_mul(t21[0], t21[1], bmvr, bmvi) - h2 = _complex_mul(t22[0], t22[1], poolvr, poolvi) - spun_pr, spun_pi = e0[0] + e1_[0] + e2_[0], e0[1] + e1_[1] + e2_[1] - spun_mr, spun_mi = f0[0] + f1[0] + f2[0], f0[1] + f1[1] + f2[1] - spun_zr, spun_zi = h0[0] + h1[0] + h2[0], h0[1] + h1[1] + h2[1] - if profiled or dynamic: - ( - spun_pr, - spun_pi, - spun_mr, - spun_mi, - spun_zr, - spun_zi, - ) = _rotate_spinor( - shaped_ar, - shaped_ai, - shaped_br, - shaped_bi, - bpvr, - bpvi, - bmvr, - bmvi, - poolvr, - poolvi, - ) - bpvr = tl.where(rotate, spun_pr, bpvr) - bpvi = tl.where(rotate, spun_pi, bpvi) - bmvr = tl.where(rotate, spun_mr, bmvr) - bmvi = tl.where(rotate, spun_mi, bmvi) - poolvr = tl.where(rotate, spun_zr, poolvr) - poolvi = tl.where(rotate, spun_zi, poolvi) - pvr = tl.where(rotate, turned_pr, pvr) - pvi = tl.where(rotate, turned_pi, pvi) - mvr = tl.where(rotate, turned_mr, mvr) - mvi = tl.where(rotate, turned_mi, mvi) - zvr = tl.where(rotate, turned_zr, zvr) - zvi = tl.where(rotate, turned_zi, zvi) - - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count - ) - spoil_b = (event_action & 8) != 0 - bpvr = tl.where(spoil_b, 0.0, tl.where(do_shift, svr, bpvr)) - bpvi = tl.where(spoil_b, 0.0, tl.where(do_shift, svi, bpvi)) - bmvr = tl.where(spoil_b, 0.0, tl.where(do_shift, wvr, bmvr)) - bmvi = tl.where(spoil_b, 0.0, tl.where(do_shift, wvi, bmvi)) - svr, svi, wvr, wvi = _shift( - pvr, pvi, mvr, mvi, state, state_mask, state_count - ) - pvr = tl.where(do_shift, svr, pvr) - pvi = tl.where(do_shift, svi, pvi) - mvr = tl.where(do_shift, wvr, mvr) - mvi = tl.where(do_shift, wvi, mvi) - spoil = (event_action & 8) != 0 - pvr = tl.where(spoil, 0.0, pvr) - pvi = tl.where(spoil, 0.0, pvi) - mvr = tl.where(spoil, 0.0, mvr) - mvi = tl.where(spoil, 0.0, mvi) - return - - # ---- reverse ---- - pbvr = empty - pbvi = empty - mbvr = empty - mbvi = empty - zbvr = empty - zbvi = empty - zero = tl.zeros((problems, 1), tl.float32) - g_diffv = zero - g_flowv = zero - g_washv = zero - g_t1v = zero - g_t2v = zero - g_m0v = zero - g_b1v = zero - g_b1pv = zero - g_b0v = zero - g_invv = zero - g_boundv = zero - g_exchv = zero - g_t1bv = zero - g_t2bv = zero - g_shiftv = zero - g_semiv = zero - g_sexchv = zero - g_t1cv = zero - poolbr = empty - poolbi = empty - semibr = empty - semibi = empty - ubvr = empty - ubvi = empty - wbvr = empty - wbvi = empty - - for reverse in range(0, event_count): - event = event_count - 1 - reverse - slot = trajectory + event * record_stride - xpvr = tl.load(trajectory_r + slot, mask=state_mask, other=0.0) - xpvi = tl.load(trajectory_i + slot, mask=state_mask, other=0.0) - xmvr = tl.load(trajectory_r + slot + minus_plane, mask=state_mask, other=0.0) - xmvi = tl.load(trajectory_i + slot + minus_plane, mask=state_mask, other=0.0) - xzvr = tl.load(trajectory_r + slot + long_plane, mask=state_mask, other=0.0) - xzvi = tl.load(trajectory_i + slot + long_plane, mask=state_mask, other=0.0) - xbvr = empty - xbvi = empty - xcvr = empty - xcvi = empty - xbpvr = empty - xbpvi = empty - xbmvr = empty - xbmvi = empty - rbpvr = empty - rbpvi = empty - rbmvr = empty - rbmvi = empty - if pools > 0: - xbvr = tl.load( - trajectory_r + slot + bound_plane, mask=state_mask, other=0.0 - ) - xbvi = tl.load( - trajectory_i + slot + bound_plane, mask=state_mask, other=0.0 - ) - if pools == 2 or pools == 3: - xbpvr = tl.load( - trajectory_r + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbpvi = tl.load( - trajectory_i + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbmvr = tl.load( - trajectory_r + slot + bminus_plane, mask=state_mask, other=0.0 - ) - xbmvi = tl.load( - trajectory_i + slot + bminus_plane, mask=state_mask, other=0.0 - ) - if pools == 3: - xcvr = tl.load( - trajectory_r + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - xcvi = tl.load( - trajectory_i + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event) - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - wout_value = 1.0 - if moving: - wout_value = _washout(atom_washout, dt_value) - dry1_value = tl.exp(-r1_value * dt_value) - dry2_value = tl.exp(-r2_value * dt_value) - e1_value = dry1_value * wout_value - e2_value = dry2_value * wout_value - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt_value, order) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value = 1.0 - e1_value - bare1_value = e1_value - bare2_value = e2_value - e1_value = bare1_value * damp_z - e2_value = bare2_value * damp_t - turn_t = 0.0 - szr, szi = 1.0, 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, dt_value, order) - szr, szi = tl.cos(turn_z), tl.sin(turn_z) - qr, qi = 1.0, 0.0 - if off_axis or moving: - angle_value = -2.0 * 3.141592653589793 * (atom_b0 * dt_value) + turn_t - qr, qi = tl.cos(angle_value), tl.sin(angle_value) - ovr, ovi = e2_value * qr, e2_value * qi - lvr, lvi = e1_value * szr, e1_value * szi - - # Replay the intra-event stages from the recorded entry state. - if pools == 2 or pools == 3: - # With an exchanging pool the transverse relaxation sits inside the - # operator instead of in the scalar the free pool alone multiplies. - across = _two_pool_transverse_step_jvp( - r2_value, - 0.0, - r2b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - atom_free, - 0.0, - atom_shift, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - ) - a11r, a11i = across[0], across[1] - a12r, a12i = across[2], across[3] - a21r, a21i = across[4], across[5] - a22r, a22i = across[6], across[7] - carr, cari = damp_t * qr, damp_t * qi - f11r, f11i = _complex_mul(a11r, a11i, xpvr, xpvi) - f12r, f12i = _complex_mul(a12r, a12i, xbpvr, xbpvi) - g21r, g21i = _complex_mul(a21r, a21i, xpvr, xpvi) - g22r, g22i = _complex_mul(a22r, a22i, xbpvr, xbpvi) - # ``F-`` takes the conjugate of the operator entry by entry, not its - # transpose: it is the conjugate state following the conjugate map. - h11r, h11i = _complex_mul(a11r, -a11i, xmvr, xmvi) - h12r, h12i = _complex_mul(a12r, -a12i, xbmvr, xbmvi) - k21r, k21i = _complex_mul(a21r, -a21i, xmvr, xmvi) - k22r, k22i = _complex_mul(a22r, -a22i, xbmvr, xbmvi) - rpvr, rpvi = _complex_mul(f11r + f12r, f11i + f12i, carr, cari) - rbpvr, rbpvi = _complex_mul(g21r + g22r, g21i + g22i, carr, cari) - rmvr, rmvi = _complex_mul(h11r + h12r, h11i + h12i, carr, -cari) - rbmvr, rbmvi = _complex_mul(k21r + k22r, k21i + k22i, carr, -cari) - else: - rpvr, rpvi = _complex_mul(ovr, ovi, xpvr, xpvi) - rmvr, rmvi = _complex_mul(ovr, -ovi, xmvr, xmvi) - - rbvr = empty - rbvi = empty - rcvr = empty - rcvi = empty - if pools == 3: - nil = 0.0 * dt_value - hold_value = wout_value + nil - if tabulated: - # The walk back needs the operator itself, which the row - # already holds -- and pooling the cotangents took what - # the eigenvalues were formed for, so nothing here reads - # them. - pool_row = tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ) - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - ) = _three_pool_from_table( - pool_table, - pool_row, - atom, - atom_count, - active_atom, - hold_value, - atom_free, - atom_bound, - atom_semisolid, - ) - else: - # The pieces and the bare operator are kept rather than - # the step alone: the walk back pushes the cotangents - # through them, and forming them once for the interval - # is what keeps this kernel a size a compiler will take. - ( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - ) = _three_pool_pieces_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - dt_value, - nil, - narrow, - ) - ( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - ) = _three_pool_assemble_jvp( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - narrow, - ) - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - _dw11, - _dw12, - _dw13, - _dw21, - _dw22, - _dw23, - _dw31, - _dw32, - _dw33, - _dgf, - _dgb, - _dgs, - ) = _three_pool_weigh_jvp( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - hold_value, - nil, - narrow, - ) - # The operator is O(1) once formed, so the per-order loop below - # takes it at the width the states are carried in. - w11 = w11.to(tl.float32) - w12 = w12.to(tl.float32) - w13 = w13.to(tl.float32) - w21 = w21.to(tl.float32) - w22 = w22.to(tl.float32) - w23 = w23.to(tl.float32) - w31 = w31.to(tl.float32) - w32 = w32.to(tl.float32) - w33 = w33.to(tl.float32) - grow_free = grow_free.to(tl.float32) - grow_pool_b = grow_pool_b.to(tl.float32) - grow_semisolid = grow_semisolid.to(tl.float32) - spin_r, spin_i = damp_z * szr, damp_z * szi - mix_fr = w11 * xzvr + w12 * xbvr + w13 * xcvr - mix_fi = w11 * xzvi + w12 * xbvi + w13 * xcvi - mix_br = w21 * xzvr + w22 * xbvr + w23 * xcvr - mix_bi = w21 * xzvi + w22 * xbvi + w23 * xcvi - mix_cr = w31 * xzvr + w32 * xbvr + w33 * xcvr - mix_ci = w31 * xzvi + w32 * xbvi + w33 * xcvi - rzvr, rzvi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - rbvr, rbvi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - rcvr, rcvi = _complex_mul(spin_r, spin_i, mix_cr, mix_ci) - rzvr += tl.where(state == 0, grow_free, 0.0) - rbvr += tl.where(state == 0, grow_pool_b, 0.0) - rcvr += tl.where(state == 0, grow_semisolid, 0.0) - elif pools > 0: - ( - pe11, - pe12, - pe21, - pe22, - prec_f, - prec_b, - _d11, - _d12, - _d21, - _d22, - _drf, - _drb, - ) = _two_pool_step_jvp( - r1_value, - 0.0, - r1b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - ) - spin_r, spin_i = damp_z * szr, damp_z * szi - mix_fr = pe11 * xzvr + pe12 * xbvr - mix_fi = pe11 * xzvi + pe12 * xbvi - mix_br = pe21 * xzvr + pe22 * xbvr - mix_bi = pe21 * xzvi + pe22 * xbvi - rzvr, rzvi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - rbvr, rbvi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - rzvr += tl.where(state == 0, prec_f, 0.0) - rbvr += tl.where(state == 0, prec_b, 0.0) - else: - rzvr, rzvi = _complex_mul(lvr, lvi, xzvr, xzvi) - rzvr += tl.where(state == 0, recovery_value, 0.0) - - pre_shift = (event_action & 1) != 0 - svr, svi, wvr, wvi = _shift( - rpvr, rpvi, rmvr, rmvi, state, state_mask, state_count - ) - spvr = tl.where(pre_shift, svr, rpvr) - spvi = tl.where(pre_shift, svi, rpvi) - smvr = tl.where(pre_shift, wvr, rmvr) - smvi = tl.where(pre_shift, wvi, rmvi) - sbpvr = empty - sbpvi = empty - sbmvr = empty - sbmvi = empty - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - rbpvr, rbpvi, rbmvr, rbmvi, state, state_mask, state_count - ) - sbpvr = tl.where(pre_shift, svr, rbpvr) - sbpvi = tl.where(pre_shift, svi, rbpvi) - sbmvr = tl.where(pre_shift, wvr, rbmvr) - sbmvi = tl.where(pre_shift, wvi, rbmvi) - - # Undo the trailing spoil or shift. - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - spoil = (event_action & 8) != 0 - avr, avi, bvr, bvi = _shift_adjoint( - pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count - ) - trailing = do_shift & ~spoil - pbvr = tl.where(spoil, 0.0, tl.where(trailing, avr, pbvr)) - pbvi = tl.where(spoil, 0.0, tl.where(trailing, avi, pbvi)) - mbvr = tl.where(spoil, 0.0, tl.where(trailing, bvr, mbvr)) - mbvi = tl.where(spoil, 0.0, tl.where(trailing, bvi, mbvi)) - if pools == 2 or pools == 3: - avr, avi, bvr, bvi = _shift_adjoint( - ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count - ) - ubvr = tl.where(spoil, 0.0, tl.where(trailing, avr, ubvr)) - ubvi = tl.where(spoil, 0.0, tl.where(trailing, avi, ubvi)) - wbvr = tl.where(spoil, 0.0, tl.where(trailing, bvr, wbvr)) - wbvi = tl.where(spoil, 0.0, tl.where(trailing, bvi, wbvi)) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - event_phase = _event_value(phase, event_base, event, active_atom, single_train) - pulse_b1 = atom_b1 - pulse_b1_phase = atom_b1_phase - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - if off_axis: - pulse_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - - # ---- recorded sample ---- - record = ((event_action & 32) != 0) & (event_kind == 2) - out = tl.load(output_index + event) - seed_mask = active_atom & record & (out >= 0) - seed_real = tl.load( - grad_output_real + problem * output_count + out, mask=seed_mask, other=0.0 - ) - seed_imag = tl.load( - grad_output_imag + problem * output_count + out, mask=seed_mask, other=0.0 - ) - dvr, dvi = tl.cos(-event_phase), tl.sin(-event_phase) - # grad_m0 = Re(conj(seed) * recorded * demodulation) - recr, reci = spvr, spvi - if pools == 2 or pools == 3: - recr, reci = spvr + sbpvr, spvi + sbpvi - wr, wi = _complex_mul(recr, reci, dvr, dvi) - g_m0v += tl.sum( - tl.where(state == 0, seed_real * wr + seed_imag * wi, 0.0), axis=1 - )[:, None] - # grad_phase = Re(conj(seed) * m0 * recorded * (-i) * demodulation) - yr, yi = atom_m0 * recr, atom_m0 * reci - yr, yi = yi, -yr - yr, yi = _complex_mul(yr, yi, dvr, dvi) - tl.atomic_add( - grad_phase + event_base + event, - tl.sum(tl.where(state == 0, seed_real * yr + seed_imag * yi, 0.0), axis=1)[ - :, None - ], - mask=seed_mask, - ) - # fplus_bar[0] += conj(m0 * demodulation) * seed - kr, ki = atom_m0 * dvr, atom_m0 * dvi - sr, si = _complex_mul(kr, -ki, seed_real, seed_imag) - pbvr += tl.where(state == 0, sr, 0.0) - pbvi += tl.where(state == 0, si, 0.0) - if pools == 2 or pools == 3: - ubvr += tl.where(state == 0, sr, 0.0) - ubvi += tl.where(state == 0, si, 0.0) - - # ---- RF adjoint ---- - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - g_invv += tl.sum(tl.where(invert, zbvr * -rzvr + zbvi * -rzvi, 0.0), axis=1)[ - :, None - ] - zbvr = tl.where(invert, -atom_inv * zbvr, zbvr) - zbvi = tl.where(invert, -atom_inv * zbvi, zbvi) - - alpha_value = event_flip * pulse_b1 - phi_value = event_phase + pulse_b1_phase - cos_value = tl.cos(alpha_value) - sin_value = tl.sin(alpha_value) - p1r, p1i = tl.cos(phi_value), tl.sin(phi_value) - p2r, p2i = _complex_mul(p1r, p1i, p1r, p1i) - t00, t01, t02, t12, t20, t21, t22 = _rotation_coefficients( - 0.5 * (1.0 + cos_value), - 0.5 * (1.0 - cos_value), - sin_value, - cos_value, - p1r, - p1i, - p2r, - p2i, - p1r, - -p1i, - ) - d00, d01, d02, d12, d20, d21, d22 = _rotation_coefficients( - -0.5 * sin_value, - 0.5 * sin_value, - cos_value, - -sin_value, - p1r, - p1i, - p2r, - p2i, - p1r, - -p1i, - ) - - sat_alpha_v = zero - sat_b0_v = zero - if pools == 1 or pools == 3: - # The pulse scales every order of the pool by one real number, so - # its cotangent is a single sum over the states it multiplied. - offset_value = tl.load(rf_frequency + event) - atom_b0 - shape_value, shape_slope = _lineshape_at_slope( - lineshape, offset_value, lineshape_bins, lineshape_step - ) - event_saturation = tl.load(saturation + event) - power_value = event_saturation * alpha_value * alpha_value - absorbed_value = tl.exp(power_value * shape_value) - if pools == 1: - per_state = poolbr * rbvr + poolbi * rbvi - else: - per_state = semibr * rcvr + semibi * rcvi - grad_absorbed = tl.sum(per_state, axis=1)[:, None] - grad_exponent = grad_absorbed * absorbed_value - twice = event_saturation * 2.0 - sat_alpha_v = grad_exponent * (twice * alpha_value * shape_value) - # The lineshape is read at the pulse's offset from the voxel, so a - # step in the voxel's own off-resonance moves the read the other way. - sat_b0_v = -grad_exponent * (power_value * shape_slope) - saturating = is_rf & ~is_inversion - if pools == 1: - poolbr = tl.where(saturating, absorbed_value * poolbr, poolbr) - poolbi = tl.where(saturating, absorbed_value * poolbi, poolbi) - else: - semibr = tl.where(saturating, absorbed_value * semibr, semibr) - semibi = tl.where(saturating, absorbed_value * semibi, semibi) - - # d/dalpha, contracted with the adjoint. - row0 = _complex_mul(d00[0], d00[1], spvr, spvi) - add1 = _complex_mul(d01[0], d01[1], smvr, smvi) - add2 = _complex_mul(d02[0], d02[1], rzvr, rzvi) - alpha_v = pbvr * (row0[0] + add1[0] + add2[0]) - alpha_v += pbvi * (row0[1] + add1[1] + add2[1]) - row0 = _complex_mul(d01[0], -d01[1], spvr, spvi) - add1 = _complex_mul(d00[0], d00[1], smvr, smvi) - add2 = _complex_mul(d12[0], d12[1], rzvr, rzvi) - alpha_v += mbvr * (row0[0] + add1[0] + add2[0]) - alpha_v += mbvi * (row0[1] + add1[1] + add2[1]) - row0 = _complex_mul(d20[0], d20[1], spvr, spvi) - add1 = _complex_mul(d21[0], d21[1], smvr, smvi) - add2 = _complex_mul(d22[0], d22[1], rzvr, rzvi) - alpha_v += zbvr * (row0[0] + add1[0] + add2[0]) - alpha_v += zbvi * (row0[1] + add1[1] + add2[1]) - - # d/dphi, where only the phase factors carry the dependence. - u1 = _complex_mul(t01[0], t01[1], smvr, smvi) - u2 = _complex_mul(t02[0], t02[1], rzvr, rzvi) - ur, ui = -(2.0 * u1[1] + u2[1]), 2.0 * u1[0] + u2[0] - phi_v = pbvr * ur + pbvi * ui - u1 = _complex_mul(t01[0], -t01[1], spvr, spvi) - u2 = _complex_mul(t12[0], t12[1], rzvr, rzvi) - ur, ui = 2.0 * u1[1] + u2[1], -2.0 * u1[0] - u2[0] - phi_v += mbvr * ur + mbvi * ui - u1 = _complex_mul(t20[0], t20[1], spvr, spvi) - u2 = _complex_mul(t21[0], t21[1], smvr, smvi) - ur, ui = -(u2[1] - u1[1]), u2[0] - u1[0] - phi_v += zbvr * ur + zbvi * ui - - if profiled or dynamic: - slope_ar, slope_ai, slope_br, slope_bi = 0.0, 0.0, 0.0, 0.0 - if dynamic: - pair = _dynamic_pair_at( - pairs, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - shaped_ar, shaped_ai = pair[0], pair[1] - shaped_br, shaped_bi = _complex_mul(pair[2], pair[3], p1r, -p1i) - else: - ( - shaped_ar, - slope_ar, - shaped_ai, - slope_ai, - shaped_br, - slope_br, - shaped_bi, - slope_bi, - ) = _profile_pair_slope( - profile, - _table_row(profile_index, event, location, locations), - alpha_value, - profile_bins, - profile_step, - ) - shaped_br, shaped_bi = _complex_mul(shaped_br, shaped_bi, p1r, -p1i) - slope_br, slope_bi = _complex_mul(slope_br, slope_bi, p1r, -p1i) - ( - grad_ar, - grad_ai, - grad_br, - grad_bi, - shaped_pbr, - shaped_pbi, - shaped_mbr, - shaped_mbi, - shaped_zbr, - shaped_zbi, - ) = _spinor_adjoint( - shaped_ar, - shaped_ai, - shaped_br, - shaped_bi, - spvr, - spvi, - smvr, - smvi, - rzvr, - rzvi, - pbvr, - pbvi, - mbvr, - mbvi, - zbvr, - zbvi, - ) - if dynamic: - # The flip is inside the pair rather than read against it, so - # it has no gradient here: the cotangent goes out on the - # rotation and whatever integrated it carries the rest. ``b`` - # was turned by the phase after the pair came out, so the - # cotangent turns back the other way. - alpha_v = alpha_v * 0.0 - back_r, back_i = _complex_mul(grad_br, grad_bi, p1r, p1i) - _store_pair_gradient( - grad_pair, - pair_index, - event_base, - event, - atom, - atom_count, - is_rf & ~is_inversion, - active_atom, - state_mask, - grad_ar, - grad_ai, - back_r, - back_i, - ) - else: - alpha_v = grad_ar * slope_ar + grad_ai * slope_ai - alpha_v += grad_br * slope_br + grad_bi * slope_bi - # d(b e^{-i phi})/dphi is -i times it, and nothing else moves. - phi_v = grad_br * shaped_bi - grad_bi * shaped_br - if pools == 2 or pools == 3: - ( - pool_ar, - pool_ai, - pool_pair_br, - pool_pair_bi, - pool_shaped_pbr, - pool_shaped_pbi, - pool_shaped_mbr, - pool_shaped_mbi, - pool_shaped_zbr, - pool_shaped_zbi, - ) = _spinor_adjoint( - shaped_ar, - shaped_ai, - shaped_br, - shaped_bi, - sbpvr, - sbpvi, - sbmvr, - sbmvi, - rbvr, - rbvi, - ubvr, - ubvi, - wbvr, - wbvi, - poolbr, - poolbi, - ) - if dynamic: - # The same pulse turned this pool, so its cotangent lands - # on the same row. - back_r, back_i = _complex_mul(pool_pair_br, pool_pair_bi, p1r, p1i) - _store_pair_gradient( - grad_pair, - pair_index, - event_base, - event, - atom, - atom_count, - is_rf & ~is_inversion, - active_atom, - state_mask, - pool_ar, - pool_ai, - back_r, - back_i, - ) - else: - alpha_v += pool_ar * slope_ar + pool_ai * slope_ai - alpha_v += pool_pair_br * slope_br + pool_pair_bi * slope_bi - phi_v += pool_pair_br * shaped_bi - pool_pair_bi * shaped_br - - if (pools == 2 or pools == 3) and not profiled and not dynamic: - # The same pulse turns the exchanging pool, so its cotangent adds to - # the flip and phase the free pool already left. - row0 = _complex_mul(d00[0], d00[1], sbpvr, sbpvi) - add1 = _complex_mul(d01[0], d01[1], sbmvr, sbmvi) - add2 = _complex_mul(d02[0], d02[1], rbvr, rbvi) - alpha_v += ubvr * (row0[0] + add1[0] + add2[0]) - alpha_v += ubvi * (row0[1] + add1[1] + add2[1]) - row0 = _complex_mul(d01[0], -d01[1], sbpvr, sbpvi) - add1 = _complex_mul(d00[0], d00[1], sbmvr, sbmvi) - add2 = _complex_mul(d12[0], d12[1], rbvr, rbvi) - alpha_v += wbvr * (row0[0] + add1[0] + add2[0]) - alpha_v += wbvi * (row0[1] + add1[1] + add2[1]) - row0 = _complex_mul(d20[0], d20[1], sbpvr, sbpvi) - add1 = _complex_mul(d21[0], d21[1], sbmvr, sbmvi) - add2 = _complex_mul(d22[0], d22[1], rbvr, rbvi) - alpha_v += poolbr * (row0[0] + add1[0] + add2[0]) - alpha_v += poolbi * (row0[1] + add1[1] + add2[1]) - u1 = _complex_mul(t01[0], t01[1], sbmvr, sbmvi) - u2 = _complex_mul(t02[0], t02[1], rbvr, rbvi) - ur, ui = -(2.0 * u1[1] + u2[1]), 2.0 * u1[0] + u2[0] - phi_v += ubvr * ur + ubvi * ui - u1 = _complex_mul(t01[0], -t01[1], sbpvr, sbpvi) - u2 = _complex_mul(t12[0], t12[1], rbvr, rbvi) - ur, ui = 2.0 * u1[1] + u2[1], -2.0 * u1[0] - u2[0] - phi_v += wbvr * ur + wbvi * ui - u1 = _complex_mul(t20[0], t20[1], sbpvr, sbpvi) - u2 = _complex_mul(t21[0], t21[1], sbmvr, sbmvi) - ur, ui = -(u2[1] - u1[1]), u2[0] - u1[0] - phi_v += poolbr * ur + poolbi * ui - - rotate = is_rf & ~is_inversion - grad_alpha_v = tl.sum(tl.where(rotate, alpha_v, 0.0), axis=1)[:, None] - grad_phi_v = tl.sum(tl.where(rotate, phi_v, 0.0), axis=1)[:, None] - if pools == 1 or pools == 3: - turning = tl.where(rotate, 1.0, 0.0) - grad_alpha_v += sat_alpha_v * turning - g_b0v += sat_b0_v * turning - - # Conjugate transpose of the rotation. - n0 = _complex_mul(t00[0], -t00[1], pbvr, pbvi) - n1 = _complex_mul(t01[0], t01[1], mbvr, mbvi) - n2 = _complex_mul(t20[0], -t20[1], zbvr, zbvi) - q0 = _complex_mul(t01[0], -t01[1], pbvr, pbvi) - q1 = _complex_mul(t00[0], -t00[1], mbvr, mbvi) - q2 = _complex_mul(t21[0], -t21[1], zbvr, zbvi) - w0 = _complex_mul(t02[0], -t02[1], pbvr, pbvi) - w1 = _complex_mul(t12[0], -t12[1], mbvr, mbvi) - w2 = _complex_mul(t22[0], -t22[1], zbvr, zbvi) - back_pr = n0[0] + n1[0] + n2[0] - back_pi = n0[1] + n1[1] + n2[1] - back_mr = q0[0] + q1[0] + q2[0] - back_mi = q0[1] + q1[1] + q2[1] - back_zr = w0[0] + w1[0] + w2[0] - back_zi = w0[1] + w1[1] + w2[1] - if profiled or dynamic: - # A shaped pulse turned the states, so its own adjoint is what - # goes back rather than the instant rotation's. - back_pr, back_pi = shaped_pbr, shaped_pbi - back_mr, back_mi = shaped_mbr, shaped_mbi - back_zr, back_zi = shaped_zbr, shaped_zbi - pbvr = tl.where(rotate, back_pr, pbvr) - pbvi = tl.where(rotate, back_pi, pbvi) - mbvr = tl.where(rotate, back_mr, mbvr) - mbvi = tl.where(rotate, back_mi, mbvi) - zbvr = tl.where(rotate, back_zr, zbvr) - zbvi = tl.where(rotate, back_zi, zbvi) - - writes_flip = active_atom & rotate - tl.atomic_add( - grad_flip + event_base + event, - grad_alpha_v * pulse_b1, - mask=writes_flip, - ) - tl.atomic_add(grad_phase + event_base + event, grad_phi_v, mask=writes_flip) - if shimmed: - # A pulse's transmit gradient belongs to the shim it drives, so with - # several it lands in that shim's row rather than in a register - # summed over the whole train. - tl.atomic_add( - grad_tissue + _B1_ROW * atom_count + row + atom, - grad_alpha_v * event_flip, - mask=writes_flip, - ) - tl.atomic_add( - grad_tissue + (_B1_PHASE_ROW + shim_rows - 1) * atom_count + row + atom, - grad_phi_v, - mask=writes_flip, - ) - else: - g_b1v += grad_alpha_v * event_flip - g_b1pv += grad_phi_v - - if pools == 2 or pools == 3: - n0 = _complex_mul(t00[0], -t00[1], ubvr, ubvi) - n1 = _complex_mul(t01[0], t01[1], wbvr, wbvi) - n2 = _complex_mul(t20[0], -t20[1], poolbr, poolbi) - q0 = _complex_mul(t01[0], -t01[1], ubvr, ubvi) - q1 = _complex_mul(t00[0], -t00[1], wbvr, wbvi) - q2 = _complex_mul(t21[0], -t21[1], poolbr, poolbi) - w0 = _complex_mul(t02[0], -t02[1], ubvr, ubvi) - w1 = _complex_mul(t12[0], -t12[1], wbvr, wbvi) - w2 = _complex_mul(t22[0], -t22[1], poolbr, poolbi) - pool_back_pr = n0[0] + n1[0] + n2[0] - pool_back_pi = n0[1] + n1[1] + n2[1] - pool_back_mr = q0[0] + q1[0] + q2[0] - pool_back_mi = q0[1] + q1[1] + q2[1] - pool_back_zr = w0[0] + w1[0] + w2[0] - pool_back_zi = w0[1] + w1[1] + w2[1] - if profiled or dynamic: - # A shaped pulse turned this pool too, so its own adjoint is - # what goes back rather than the instant rotation's. - pool_back_pr, pool_back_pi = pool_shaped_pbr, pool_shaped_pbi - pool_back_mr, pool_back_mi = pool_shaped_mbr, pool_shaped_mbi - pool_back_zr, pool_back_zi = pool_shaped_zbr, pool_shaped_zbi - ubvr = tl.where(rotate, pool_back_pr, ubvr) - ubvi = tl.where(rotate, pool_back_pi, ubvi) - wbvr = tl.where(rotate, pool_back_mr, wbvr) - wbvi = tl.where(rotate, pool_back_mi, wbvi) - poolbr = tl.where(rotate, pool_back_zr, poolbr) - poolbi = tl.where(rotate, pool_back_zi, poolbi) - # An inversion turns the exchanging pool's longitudinal state as - # well, so the efficiency carries what both left behind. - g_invv += tl.sum( - tl.where(invert, poolbr * -rbvr + poolbi * -rbvi, 0.0), axis=1 - )[:, None] - poolbr = tl.where(invert, -atom_inv * poolbr, poolbr) - poolbi = tl.where(invert, -atom_inv * poolbi, poolbi) - avr, avi, bvr, bvi = _shift_adjoint( - ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count - ) - ubvr = tl.where(pre_shift, avr, ubvr) - ubvi = tl.where(pre_shift, avi, ubvi) - wbvr = tl.where(pre_shift, bvr, wbvr) - wbvi = tl.where(pre_shift, bvi, wbvi) - avr, avi, bvr, bvi = _shift_adjoint( - pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count - ) - pbvr = tl.where(pre_shift, avr, pbvr) - pbvi = tl.where(pre_shift, avi, pbvi) - mbvr = tl.where(pre_shift, bvr, mbvr) - mbvi = tl.where(pre_shift, bvi, mbvi) - - # ---- relaxation and off-resonance adjoint ---- - # The damping is homogeneous of degree one in every transverse state it - # acts on, so its gradient times the damping itself is the cotangent - # taken against the states the interval leaves. - pq = _complex_mul(qr, qi, xpvr, xpvi) - mq = _complex_mul(qr, -qi, xmvr, xmvi) - bare_cot_v = pbvr * pq[0] + pbvi * pq[1] - bare_cot_v += mbvr * mq[0] + mbvi * mq[1] - grad_e2_v = tl.sum(bare_cot_v * damp_t, axis=1)[:, None] - cot2_v = bare_cot_v * bare2_value * damp_t - pool_angle_v = empty - if pools == 2 or pools == 3: - # With an exchanging pool the damping sits inside the operator, so - # the cotangent the interval leaves is taken against the states it - # produced rather than against a scalar the free pool multiplies. - plus_r = pbvr * rpvr + pbvi * rpvi + ubvr * rbpvr + ubvi * rbpvi - plus_i = pbvr * rpvi - pbvi * rpvr + ubvr * rbpvi - ubvi * rbpvr - minus_r = mbvr * rmvr + mbvi * rmvi + wbvr * rbmvr + wbvi * rbmvi - minus_i = mbvr * rmvi - mbvi * rmvr + wbvr * rbmvi - wbvi * rbmvr - cot2_v = plus_r + minus_r - pool_angle_v = minus_i - plus_i - if pools == 2 or pools == 3: - # ``F-`` follows the conjugate of the operator, so its cotangent - # lands on the entry itself rather than on the conjugate of it. - def_r, def_i = carr, cari - t11r, t11i = _complex_mul( - pbvr * xpvr + pbvi * xpvi + mbvr * xmvr + mbvi * xmvi, - pbvr * xpvi - pbvi * xpvr - mbvr * xmvi + mbvi * xmvr, - def_r, - def_i, - ) - t12r, t12i = _complex_mul( - pbvr * xbpvr + pbvi * xbpvi + mbvr * xbmvr + mbvi * xbmvi, - pbvr * xbpvi - pbvi * xbpvr - mbvr * xbmvi + mbvi * xbmvr, - def_r, - def_i, - ) - t21r, t21i = _complex_mul( - ubvr * xpvr + ubvi * xpvi + wbvr * xmvr + wbvi * xmvi, - ubvr * xpvi - ubvi * xpvr - wbvr * xmvi + wbvi * xmvr, - def_r, - def_i, - ) - t22r, t22i = _complex_mul( - ubvr * xbpvr + ubvi * xbpvi + wbvr * xbmvr + wbvi * xbmvi, - ubvr * xbpvi - ubvi * xbpvr - wbvr * xbmvi + wbvi * xbmvr, - def_r, - def_i, - ) - bar11 = ( - tl.sum(t11r, axis=1)[:, None], - tl.sum(t11i, axis=1)[:, None], - zero, - zero, - ) - bar12 = ( - tl.sum(t12r, axis=1)[:, None], - tl.sum(t12i, axis=1)[:, None], - zero, - zero, - ) - bar21 = ( - tl.sum(t21r, axis=1)[:, None], - tl.sum(t21i, axis=1)[:, None], - zero, - zero, - ) - bar22 = ( - tl.sum(t22r, axis=1)[:, None], - tl.sum(t22i, axis=1)[:, None], - zero, - zero, - ) - ( - back_r2, - _q1, - back_r2b, - _q2, - back_xexch, - _q3, - back_xbound, - _q4, - back_xfree, - _q5, - back_shift, - _q6, - back_xdt, - _q7, - back_xatt, - _q8, - ) = _two_pool_transverse_adjoint_jvp( - r2_value, - 0.0, - r2b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - atom_free, - 0.0, - atom_shift, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - bar11, - bar12, - bar21, - bar22, - ) - g_t2v += back_r2 * (-1000.0 / (atom_t2 * atom_t2)) - g_t2bv += back_r2b * (-1000.0 / (atom_t2b * atom_t2b)) - g_exchv += back_xexch - # The free fraction is one less the pool's, so what reaches it - # arrives at the pool's own with the sign turned. - g_boundv += back_xbound - back_xfree - if pools == 3: - # The free share is one less both fractions, so what the - # transverse operator leaves on it reaches the semisolid too. - g_semiv -= back_xfree - g_shiftv += back_shift - xversal_dt = back_xdt - xversal_att = back_xatt - # The pool's transverse cotangents go back through the same - # operator, transposed. - ur, ui = _complex_mul(a11r, -a11i, pbvr, pbvi) - vr_, vi_ = _complex_mul(a21r, -a21i, ubvr, ubvi) - nub_pr, nub_pi = _complex_mul(ur + vr_, ui + vi_, carr, -cari) - ur, ui = _complex_mul(a12r, -a12i, pbvr, pbvi) - vr_, vi_ = _complex_mul(a22r, -a22i, ubvr, ubvi) - nub_qr, nub_qi = _complex_mul(ur + vr_, ui + vi_, carr, -cari) - ur, ui = _complex_mul(a11r, a11i, mbvr, mbvi) - vr_, vi_ = _complex_mul(a21r, a21i, wbvr, wbvi) - nwb_pr, nwb_pi = _complex_mul(ur + vr_, ui + vi_, carr, cari) - ur, ui = _complex_mul(a12r, a12i, mbvr, mbvi) - vr_, vi_ = _complex_mul(a22r, a22i, wbvr, wbvi) - nwb_qr, nwb_qi = _complex_mul(ur + vr_, ui + vi_, carr, cari) - pbvr, pbvi = nub_pr, nub_pi - ubvr, ubvi = nub_qr, nub_qi - mbvr, mbvi = nwb_pr, nwb_pi - wbvr, wbvi = nwb_qr, nwb_qi - - per_angle_v = pool_angle_v - if pools != 2 and pools != 3 and (off_axis or moving): - po = _complex_mul(ovr, ovi, xpvr, xpvi) - mo = _complex_mul(ovr, -ovi, xmvr, xmvi) - # A turn of the transverse states and the off-resonance angle are - # the same derivative; only the weight each order carries differs. - per_angle_v = pbvr * -po[1] + pbvi * po[0] - per_angle_v -= mbvr * -mo[1] + mbvi * mo[0] - grad_angle_v = zero - if off_axis or moving: - grad_angle_v = tl.sum(per_angle_v, axis=1)[:, None] - - e1_v = empty - grad_e1_v = zero - long_damp_v = empty - attenuation_v = zero - two_pool_dt_v = zero - if pools == 2 or pools == 3: - attenuation_v += xversal_att - two_pool_dt_v += xversal_dt - zangle_v = empty - if pools == 3: - # The nine entries of the mixing operator and the three - # recoveries, summed over the orders that share them, then pushed - # back through the closed form once for the whole interval. - spun_fr, spun_fi = _complex_mul(spin_r, spin_i, xzvr, xzvi) - spun_br, spun_bi = _complex_mul(spin_r, spin_i, xbvr, xbvi) - spun_cr, spun_ci = _complex_mul(spin_r, spin_i, xcvr, xcvi) - e11_v = zbvr * spun_fr + zbvi * spun_fi - e12_v = zbvr * spun_br + zbvi * spun_bi - e13_v = zbvr * spun_cr + zbvi * spun_ci - e21_v = poolbr * spun_fr + poolbi * spun_fi - e22_v = poolbr * spun_br + poolbi * spun_bi - e23_v = poolbr * spun_cr + poolbi * spun_ci - e31_v = semibr * spun_fr + semibi * spun_fi - e32_v = semibr * spun_br + semibi * spun_bi - e33_v = semibr * spun_cr + semibi * spun_ci - if tabulated: - # Every gradient but the interval's own is linear in these - # twelve, so the events sharing a length pool them here and - # pay the closed form once each after the walk back. - bar11 = tl.sum(e11_v, axis=1)[:, None] - bar12 = tl.sum(e12_v, axis=1)[:, None] - bar13 = tl.sum(e13_v, axis=1)[:, None] - bar21 = tl.sum(e21_v, axis=1)[:, None] - bar22 = tl.sum(e22_v, axis=1)[:, None] - bar23 = tl.sum(e23_v, axis=1)[:, None] - bar31 = tl.sum(e31_v, axis=1)[:, None] - bar32 = tl.sum(e32_v, axis=1)[:, None] - bar33 = tl.sum(e33_v, axis=1)[:, None] - bar_free = tl.sum(tl.where(state == 0, zbvr, nil), axis=1)[:, None] - bar_pool_b = tl.sum(tl.where(state == 0, poolbr, nil), axis=1)[:, None] - bar_bound = tl.sum(tl.where(state == 0, semibr, nil), axis=1)[:, None] - held = pool_bars + (local * row_count + pool_row) * 12 - tl.store( - held + 0, - tl.load(held + 0, mask=active_atom, other=0.0) + bar11, - mask=active_atom, - ) - tl.store( - held + 1, - tl.load(held + 1, mask=active_atom, other=0.0) + bar12, - mask=active_atom, - ) - tl.store( - held + 2, - tl.load(held + 2, mask=active_atom, other=0.0) + bar13, - mask=active_atom, - ) - tl.store( - held + 3, - tl.load(held + 3, mask=active_atom, other=0.0) + bar21, - mask=active_atom, - ) - tl.store( - held + 4, - tl.load(held + 4, mask=active_atom, other=0.0) + bar22, - mask=active_atom, - ) - tl.store( - held + 5, - tl.load(held + 5, mask=active_atom, other=0.0) + bar23, - mask=active_atom, - ) - tl.store( - held + 6, - tl.load(held + 6, mask=active_atom, other=0.0) + bar31, - mask=active_atom, - ) - tl.store( - held + 7, - tl.load(held + 7, mask=active_atom, other=0.0) + bar32, - mask=active_atom, - ) - tl.store( - held + 8, - tl.load(held + 8, mask=active_atom, other=0.0) + bar33, - mask=active_atom, - ) - tl.store( - held + 9, - tl.load(held + 9, mask=active_atom, other=0.0) + bar_free, - mask=active_atom, - ) - tl.store( - held + 10, - tl.load(held + 10, mask=active_atom, other=0.0) + bar_pool_b, - mask=active_atom, - ) - tl.store( - held + 11, - tl.load(held + 11, mask=active_atom, other=0.0) + bar_bound, - mask=active_atom, - ) - back_dt, back_att = _three_pool_interval_adjoint( - pool_table, - pool_row, - atom, - atom_count, - active_atom, - r1_value, - r1b_value, - r1c_value, - atom_exchange, - atom_semisolid_exchange, - atom_bound, - atom_semisolid, - hold_value, - bar11, - bar12, - bar13, - bar21, - bar22, - bar23, - bar31, - bar32, - bar33, - bar_free, - bar_pool_b, - bar_bound, - ) - attenuation_v += back_att - two_pool_dt_v += back_dt - else: - ( - back_r1, - back_r1b, - back_r1c, - back_exch, - back_sexch, - back_bound, - back_semi, - back_dt, - back_att, - _q1, - _q2, - _q3, - _q4, - _q5, - _q6, - _q7, - _q8, - _q9, - ) = _three_pool_step_adjoint_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - dt_value, - nil, - hold_value, - nil, - tl.sum(e11_v, axis=1)[:, None], - nil, - tl.sum(e12_v, axis=1)[:, None], - nil, - tl.sum(e13_v, axis=1)[:, None], - nil, - tl.sum(e21_v, axis=1)[:, None], - nil, - tl.sum(e22_v, axis=1)[:, None], - nil, - tl.sum(e23_v, axis=1)[:, None], - nil, - tl.sum(e31_v, axis=1)[:, None], - nil, - tl.sum(e32_v, axis=1)[:, None], - nil, - tl.sum(e33_v, axis=1)[:, None], - nil, - tl.sum(tl.where(state == 0, zbvr, nil), axis=1)[:, None], - nil, - tl.sum(tl.where(state == 0, poolbr, nil), axis=1)[:, None], - nil, - tl.sum(tl.where(state == 0, semibr, nil), axis=1)[:, None], - nil, - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - narrow, - ) - g_t1v += back_r1 * (-1000.0 / (atom_t1 * atom_t1)) - g_t1bv += back_r1b * (-1000.0 / (atom_t1b * atom_t1b)) - g_t1cv += back_r1c * (-1000.0 / (atom_t1c * atom_t1c)) - g_exchv += back_exch - g_sexchv += back_sexch - g_boundv += back_bound - g_semiv += back_semi - # Both halves of the interval reach the same two, so the - # transverse pass has already put its share here. - attenuation_v += back_att - two_pool_dt_v += back_dt - # All three pools take the same per-order damping and turn, so each - # collects the cotangent of the mixture that reached it. - sfr, sfi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - sbr, sbi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - scr, sci = _complex_mul(spin_r, spin_i, mix_cr, mix_ci) - long_damp_v = ( - (zbvr * sfr + zbvi * sfi) - + (poolbr * sbr + poolbi * sbi) - + (semibr * scr + semibi * sci) - ) - if moving: - zangle_v = ( - (zbvr * -sfi + zbvi * sfr) - + (poolbr * -sbi + poolbi * sbr) - + (semibr * -sci + semibi * scr) - ) - col_fr, col_fi = _complex_mul(w11 * spin_r, -(w11 * spin_i), zbvr, zbvi) - part_r, part_i = _complex_mul(w21 * spin_r, -(w21 * spin_i), poolbr, poolbi) - col_fr, col_fi = col_fr + part_r, col_fi + part_i - part_r, part_i = _complex_mul(w31 * spin_r, -(w31 * spin_i), semibr, semibi) - col_fr, col_fi = col_fr + part_r, col_fi + part_i - col_br, col_bi = _complex_mul(w12 * spin_r, -(w12 * spin_i), zbvr, zbvi) - part_r, part_i = _complex_mul(w22 * spin_r, -(w22 * spin_i), poolbr, poolbi) - col_br, col_bi = col_br + part_r, col_bi + part_i - part_r, part_i = _complex_mul(w32 * spin_r, -(w32 * spin_i), semibr, semibi) - col_br, col_bi = col_br + part_r, col_bi + part_i - col_cr, col_ci = _complex_mul(w13 * spin_r, -(w13 * spin_i), zbvr, zbvi) - part_r, part_i = _complex_mul(w23 * spin_r, -(w23 * spin_i), poolbr, poolbi) - col_cr, col_ci = col_cr + part_r, col_ci + part_i - part_r, part_i = _complex_mul(w33 * spin_r, -(w33 * spin_i), semibr, semibi) - col_cr, col_ci = col_cr + part_r, col_ci + part_i - zbvr, zbvi = col_fr, col_fi - poolbr, poolbi = col_br, col_bi - semibr, semibi = col_cr, col_ci - elif pools > 0: - # The four entries of the exchange operator and the two recoveries, - # summed over the orders that share them, then pushed back through - # the closed form once for the whole interval. - spun_fr, spun_fi = _complex_mul(spin_r, spin_i, xzvr, xzvi) - spun_br, spun_bi = _complex_mul(spin_r, spin_i, xbvr, xbvi) - bar_e11 = tl.sum(zbvr * spun_fr + zbvi * spun_fi, axis=1)[:, None] - bar_e12 = tl.sum(zbvr * spun_br + zbvi * spun_bi, axis=1)[:, None] - bar_e21 = tl.sum(poolbr * spun_fr + poolbi * spun_fi, axis=1)[:, None] - bar_e22 = tl.sum(poolbr * spun_br + poolbi * spun_bi, axis=1)[:, None] - rec_f = tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None] - rec_b = tl.sum(tl.where(state == 0, poolbr, 0.0), axis=1)[:, None] - ( - back_r1, - back_r1b, - back_exch, - back_bound, - back_dt, - back_att, - _t1, - _t2, - _t3, - _t4, - _t5, - _t6, - ) = _two_pool_step_adjoint_jvp( - r1_value, - 0.0, - r1b_value, - 0.0, - atom_exchange, - 0.0, - atom_bound, - 0.0, - dt_value, - 0.0, - wout_value, - 0.0, - bar_e11, - 0.0, - bar_e12, - 0.0, - bar_e21, - 0.0, - bar_e22, - 0.0, - rec_f, - 0.0, - rec_b, - 0.0, - ) - # r1 = 1000/t1, so a rate gradient reaches the time through the - # square of it. - g_t1v += back_r1 * (-1000.0 / (atom_t1 * atom_t1)) - g_t1bv += back_r1b * (-1000.0 / (atom_t1b * atom_t1b)) - g_exchv += back_exch - g_boundv += back_bound - # Both halves of the interval reach the same two, so the - # transverse pass has already put its share here. - attenuation_v += back_att - two_pool_dt_v += back_dt - # Both pools take the same per-order damping and turn, so each - # collects the cotangent of the mixture that reached it. - sfr, sfi = _complex_mul(spin_r, spin_i, mix_fr, mix_fi) - sbr, sbi = _complex_mul(spin_r, spin_i, mix_br, mix_bi) - long_damp_v = (zbvr * sfr + zbvi * sfi) + (poolbr * sbr + poolbi * sbi) - if moving: - zangle_v = (zbvr * -sfi + zbvi * sfr) + (poolbr * -sbi + poolbi * sbr) - back_zr, back_zi = _complex_mul(pe11 * spin_r, -(pe11 * spin_i), zbvr, zbvi) - cross_zr, cross_zi = _complex_mul( - pe21 * spin_r, -(pe21 * spin_i), poolbr, poolbi - ) - back_br, back_bi = _complex_mul(pe12 * spin_r, -(pe12 * spin_i), zbvr, zbvi) - cross_br, cross_bi = _complex_mul( - pe22 * spin_r, -(pe22 * spin_i), poolbr, poolbi - ) - poolbr = back_br + cross_br - poolbi = back_bi + cross_bi - zbvr = back_zr + cross_zr - zbvi = back_zi + cross_zi - else: - spun = _complex_mul(szr, szi, xzvr, xzvi) - e1_v = zbvr * spun[0] + zbvi * spun[1] - grad_e1_v = tl.sum(e1_v * damp_z, axis=1)[:, None] - grad_e1_v -= tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None] - # The longitudinal states turn too, and by a whole order rather - # than the transverse half-order more. - if moving: - zo = _complex_mul(lvr, lvi, xzvr, xzvi) - zangle_v = zbvr * -zo[1] + zbvi * zo[0] - long_damp_v = e1_v * bare1_value * damp_z - zbvr, zbvi = _complex_mul(lvr, -lvi, zbvr, zbvi) - - spread_v = zero - if diffusing: - # The rate and the interval multiply every order's b-weight, so - # both take a weighted sum rather than one scalar. Order zero - # carries no longitudinal weight, which keeps recovery out of this. - weighted_v = long_damp_v * longitudinal_weight + cot2_v * transverse_weight - spread_v = tl.sum(weighted_v, axis=1)[:, None] - g_diffv += -spread_v * dt_value - - wound_v = zero - wash_v = zero - if moving: - wound_v = tl.sum(per_angle_v * (order + 0.5) + zangle_v * order, axis=1)[ - :, None - ] - g_flowv += -wound_v * dt_value - # Washout scales both relaxation factors, so its gradient is the - # one they already carry, taken against the factors before that - # scaling. Past the clamp the interval has replaced the voxel - # outright and nothing further depends on the rate. - live = (atom_washout * dt_value < 1.0).to(tl.float32) - transverse_dry = ( - zero if pools == 2 or pools == 3 else grad_e2_v * dry2_value - ) - wash_v = -live * (grad_e1_v * dry1_value + transverse_dry + attenuation_v) - g_washv += wash_v * dt_value - - if pools != 2 and pools != 3: - pbvr, pbvi = _complex_mul(ovr, -ovi, pbvr, pbvi) - mbvr, mbvi = _complex_mul(ovr, ovi, mbvr, mbvi) - - inverse1_value = 1000.0 / (atom_t1 * atom_t1) - inverse2_value = 1000.0 / (atom_t2 * atom_t2) - g_t1v += grad_e1_v * (bare1_value * dt_value * inverse1_value) - if pools != 2 and pools != 3: - g_t2v += grad_e2_v * (bare2_value * dt_value * inverse2_value) - - turn = -2.0 * 3.141592653589793 - g_b0v += grad_angle_v * (turn * dt_value) - - duration_v = -grad_e1_v * (r1_value * bare1_value) - if pools != 2 and pools != 3: - duration_v -= grad_e2_v * (r2_value * bare2_value) - duration_v += grad_angle_v * (turn * atom_b0) + two_pool_dt_v - duration_v += -spread_v * atom_damping - wound_v * atom_flow - duration_v += wash_v * atom_washout - tl.atomic_add(grad_duration + event_base + event, duration_v, mask=active_atom) - - if pools == 3 and tabulated: - # One closed form per distinct length rather than one per event. The - # walk back pooled the cotangents the eigenvalues are pushed through, - # and the closed form is linear in them, so the pieces of the sum are - # the sum of the pieces. - for row in range(0, row_count): - held = pool_bars + (local * row_count + row) * 12 - row_dt = tl.load(pool_durations + row) + zero - one_att = _washout(atom_washout, row_dt) if moving else 1.0 + 0.0 * row_dt - nil = 0.0 * row_dt - ( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - ) = _three_pool_pieces_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - row_dt, - nil, - narrow, - ) - ( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - ) = _three_pool_assemble_jvp( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - narrow, - ) - ( - back_r1, - back_r1b, - back_r1c, - back_exch, - back_sexch, - back_bound, - back_semi, - back_dt, - back_att, - _q1, - _q2, - _q3, - _q4, - _q5, - _q6, - _q7, - _q8, - _q9, - ) = _three_pool_step_adjoint_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - row_dt, - nil, - one_att, - nil, - tl.load(held + 0, mask=active_atom, other=0.0), - nil, - tl.load(held + 1, mask=active_atom, other=0.0), - nil, - tl.load(held + 2, mask=active_atom, other=0.0), - nil, - tl.load(held + 3, mask=active_atom, other=0.0), - nil, - tl.load(held + 4, mask=active_atom, other=0.0), - nil, - tl.load(held + 5, mask=active_atom, other=0.0), - nil, - tl.load(held + 6, mask=active_atom, other=0.0), - nil, - tl.load(held + 7, mask=active_atom, other=0.0), - nil, - tl.load(held + 8, mask=active_atom, other=0.0), - nil, - tl.load(held + 9, mask=active_atom, other=0.0), - nil, - tl.load(held + 10, mask=active_atom, other=0.0), - nil, - tl.load(held + 11, mask=active_atom, other=0.0), - nil, - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - narrow, - ) - g_t1v += back_r1 * (-1000.0 / (atom_t1 * atom_t1)) - g_t1bv += back_r1b * (-1000.0 / (atom_t1b * atom_t1b)) - g_t1cv += back_r1c * (-1000.0 / (atom_t1c * atom_t1c)) - g_exchv += back_exch - g_sexchv += back_sexch - g_boundv += back_bound - g_semiv += back_semi - - velocity_v = g_flowv * flow_scale + g_washv * direction * washout_scale - values = ( - g_t1v, - g_t2v, - g_m0v, - g_b1v, - g_b1pv, - g_b0v, - g_invv, - g_diffv, - velocity_v, - ) - if pools > 0: - # The fraction also sets where each pool starts, which the walk back - # reaches last. - g_boundv += tl.sum(tl.where(state == 0, poolbr - zbvr, 0.0), axis=1)[:, None] - if pools == 3: - g_semiv += tl.sum(tl.where(state == 0, semibr - zbvr, 0.0), axis=1)[:, None] - semisolid_row = _BOUND_ROW + 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue + semisolid_row * atom_count + atom, - g_semiv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (semisolid_row + 1) * atom_count + atom, - g_sexchv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (semisolid_row + 2) * atom_count + atom, - g_t1cv, - mask=active_atom, - ) - if pools == 2 or pools == 3: - base_row = _POOL_B_ROW + 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue + base_row * atom_count + atom, - g_boundv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 1) * atom_count + atom, - g_exchv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 2) * atom_count + atom, - g_t1bv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 3) * atom_count + atom, - g_t2bv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 4) * atom_count + atom, - g_shiftv, - mask=active_atom, - ) - if pools == 1: - base_row = _BOUND_ROW + 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue + base_row * atom_count + atom, - g_boundv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 1) * atom_count + atom, - g_exchv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (base_row + 2) * atom_count + atom, - g_t1bv, - mask=active_atom, - ) - for parameter in tl.static_range(_FREE_POOL_COUNT): - # The transmit pair went to its shim's row above when there is more - # than one; the rest sit past whatever rows that pair took. - if not shimmed or (parameter != _B1_ROW and parameter != _B1_PHASE_ROW): - plane = ( - parameter if parameter < _B1_ROW else parameter + 2 * (shim_rows - 1) - ) - tl.atomic_add( - grad_tissue + plane * atom_count + atom, - values[parameter], - mask=active_atom, - ) - - -@triton.jit( - do_not_specialize=["state_count", "locations", "profile_bins", "lineshape_bins"] -) -def _epg_vjp_jvp_kernel( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - profile, - profile_index, - lineshape, - pairs, - pair_index, - pair_direction, - grad_pair_value, - grad_pair_tangent, - dot_t1, - dot_t2, - dot_m0, - dot_b1, - dot_b1_phase, - dot_b0, - dot_inversion_efficiency, - dot_diffusion, - dot_velocity, - dot_bound_fraction, - dot_exchange_rate, - dot_t1_bound, - dot_pool_b_fraction, - dot_pool_b_exchange, - dot_t1_pool_b, - dot_t2_pool_b, - dot_pool_b_shift, - dot_duration, - dot_flip, - dot_phase, - duration_row, - pool_table, - pool_bars, - pool_durations, - row_count, - grad_output_real, - grad_output_imag, - grad_tissue_value, - grad_tissue_tangent, - grad_flip_value, - grad_flip_tangent, - grad_phase_value, - grad_phase_tangent, - grad_duration_value, - grad_duration_tangent, - trajectory_vr, - trajectory_vi, - trajectory_tr, - trajectory_ti, - problem_base, - problem_end, - atom_count, - train_count, - event_count, - output_count, - flow_scale, - washout_scale, - profile_step, - lineshape_step, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shim_rows, - shimmed: tl.constexpr, - locations, - profiled: tl.constexpr, - profile_bins, - dynamic: tl.constexpr, - directed: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - broadened: tl.constexpr, - lineshape_bins, - pools: tl.constexpr, - narrow: tl.constexpr, - tabulated: tl.constexpr, - recording: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = problem_base + tl.program_id(0) * problems - problem = problem + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < problem_end - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - # Voxels are spread over the slice voxel-major, so a voxel's place along - # the slice is its index modulo the profile's width. One pulse shape holds - # that many consecutive rows, and the event says which shape it drives. - location = atom % locations - local = problem - problem_base - # A second pool rides along as planes of its own: it enters an event as its - # own vector and the RF operator acts on it, so the reverse sweep cannot - # replay it from the free pool's. A semisolid pool adds one plane, a - # chemically exchanging one three, and the two together add four. - record_stride = ( - 7 if pools == 3 else (6 if pools == 2 else (4 if pools == 1 else 3)) - ) * state_count - trajectory = local * event_count * record_stride + state - minus_plane = state_count - long_plane = 2 * state_count - bound_plane = 3 * state_count - bplus_plane = 4 * state_count - bminus_plane = 5 * state_count - semisolid_plane = 6 * state_count - - empty = tl.zeros((problems, block_states), tl.float32) - pvr = empty - pvi = empty - ptr = empty - pti = empty - mvr = empty - mvi = empty - mtr = empty - mti = empty - bvr = empty - bvi = empty - btr = empty - bti = empty - bpvr = empty - bpvi = empty - bptr = empty - bpti = empty - bmvr = empty - bmvi = empty - bmtr = empty - bmti = empty - cvr = empty - cvi = empty - ctr = empty - cti = empty - atom_bound = 0.0 - d_boundf = 0.0 - atom_exchange = 0.0 - d_exchange = 0.0 - atom_t1b = 1.0 - d_t1b = 0.0 - r1b_value = 0.0 - r1b_tangent = 0.0 - atom_t2b = 1.0 - d_t2b = 0.0 - r2b_value = 0.0 - r2b_tangent = 0.0 - atom_shift = 0.0 - d_shift = 0.0 - atom_semisolid = 0.0 - d_semisolidf = 0.0 - atom_semisolid_exchange = 0.0 - d_semisolid_exchange = 0.0 - r1c_value = 0.0 - r1c_tangent = 0.0 - if pools == 1: - atom_bound = tl.load(bound_fraction + scalar_atom, mask=active_atom, other=0.0) - d_boundf = tl.load( - dot_bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - d_exchange = tl.load( - dot_exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - atom_t1b = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - d_t1b = tl.load(dot_t1_bound + scalar_atom, mask=active_atom, other=0.0) - r1b_value = 1000.0 / atom_t1b - r1b_tangent = -1000.0 * d_t1b / (atom_t1b * atom_t1b) - if pools == 2 or pools == 3: - atom_bound = tl.load(pool_b_fraction + scalar_atom, mask=active_atom, other=0.0) - d_boundf = tl.load( - dot_pool_b_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_exchange = tl.load( - pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - d_exchange = tl.load( - dot_pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - atom_t1b = tl.load(t1_pool_b + scalar_atom, mask=active_atom, other=1.0) - d_t1b = tl.load(dot_t1_pool_b + scalar_atom, mask=active_atom, other=0.0) - r1b_value = 1000.0 / atom_t1b - r1b_tangent = -1000.0 * d_t1b / (atom_t1b * atom_t1b) - atom_t2b = tl.load(t2_pool_b + scalar_atom, mask=active_atom, other=1.0) - d_t2b = tl.load(dot_t2_pool_b + scalar_atom, mask=active_atom, other=0.0) - r2b_value = 1000.0 / atom_t2b - r2b_tangent = -1000.0 * d_t2b / (atom_t2b * atom_t2b) - atom_shift = tl.load(pool_b_shift + scalar_atom, mask=active_atom, other=0.0) - d_shift = tl.load(dot_pool_b_shift + scalar_atom, mask=active_atom, other=0.0) - if pools == 3: - atom_semisolid = tl.load( - bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - d_semisolidf = tl.load( - dot_bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_semisolid_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - d_semisolid_exchange = tl.load( - dot_exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - held_semisolid = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - d_semisolid_t1 = tl.load( - dot_t1_bound + scalar_atom, mask=active_atom, other=0.0 - ) - r1c_value = 1000.0 / held_semisolid - r1c_tangent = -1000.0 * d_semisolid_t1 / (held_semisolid * held_semisolid) - cvr = empty + tl.where(state == 0, atom_semisolid + 0.0, 0.0) - ctr = empty + tl.where(state == 0, d_semisolidf + 0.0, 0.0) - if pools > 0: - atom_free = 1.0 - atom_bound - atom_semisolid - d_free = -d_boundf - d_semisolidf - zvr = empty + tl.where(state == 0, atom_free, 0.0) - ztr = empty + tl.where(state == 0, d_free, 0.0) - bvr = empty + tl.where(state == 0, atom_bound + 0.0, 0.0) - btr = empty + tl.where(state == 0, d_boundf + 0.0, 0.0) - else: - atom_free = 1.0 + 0.0 * atom_bound - d_free = 0.0 * atom_bound - zvr = empty + tl.where(state == 0, 1.0, 0.0) - ztr = empty - zvi = empty - zti = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_b1_phase = 0.0 - atom_b0 = 0.0 - if off_axis: - atom_b1_phase = tl.load(b1_phase + scalar_atom, mask=active_atom, other=0.0) - atom_b0 = tl.load(b0 + scalar_atom, mask=active_atom, other=0.0) - atom_inv = 1.0 - if inverting: - atom_inv = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - d_t1 = tl.load(dot_t1 + atom, mask=active_atom, other=0.0) - d_t2 = tl.load(dot_t2 + atom, mask=active_atom, other=0.0) - d_m0 = 0.0 - if density: - d_m0 = tl.load(dot_m0 + scalar_atom, mask=active_atom, other=0.0) - d_b1 = 0.0 - if transmit: - d_b1 = tl.load(dot_b1 + scalar_atom, mask=active_atom, other=0.0) - d_b1_phase = 0.0 - d_b0 = 0.0 - if off_axis: - d_b1_phase = tl.load(dot_b1_phase + scalar_atom, mask=active_atom, other=0.0) - d_b0 = tl.load(dot_b0 + scalar_atom, mask=active_atom, other=0.0) - d_inv = 0.0 - if inverting: - d_inv = tl.load( - dot_inversion_efficiency + scalar_atom, mask=active_atom, other=0.0 - ) - atom_damping = 0.0 - d_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - d_damping = tl.load(dot_diffusion + scalar_atom, mask=active_atom, other=0.0) - atom_flow = 0.0 - d_flow = 0.0 - direction = 0.0 - atom_washout = 0.0 - d_washout = 0.0 - if moving: - atom_velocity = tl.load(velocity + scalar_atom, mask=active_atom, other=0.0) - d_velocity = tl.load(dot_velocity + scalar_atom, mask=active_atom, other=0.0) - atom_flow = atom_velocity * flow_scale - d_flow = d_velocity * flow_scale - # |v| has no derivative at the origin, so a still voxel contributes - # none. - direction = (atom_velocity > 0.0).to(tl.float32) - (atom_velocity < 0.0).to( - tl.float32 - ) - atom_washout = tl.abs(atom_velocity) * washout_scale - d_washout = direction * d_velocity * washout_scale - order = state.to(tl.float32) - longitudinal_weight = order * order - transverse_weight = longitudinal_weight + order + 0.3333333333333333 - r1_value = 1000.0 / atom_t1 - r1_tangent = -1000.0 * d_t1 / (atom_t1 * atom_t1) - r2_value = 1000.0 / atom_t2 - r2_tangent = -1000.0 * d_t2 / (atom_t2 * atom_t2) - - event_base = train * event_count - # The forward half records the trajectory the reverse half walks back, - # and the two are launched separately: each compiles the sweep it is - # asked for and no more. - if recording: - for event in range(0, event_count): - slot = trajectory + event * record_stride - tl.store(trajectory_vr + slot, pvr, mask=state_mask) - tl.store(trajectory_vi + slot, pvi, mask=state_mask) - tl.store(trajectory_tr + slot, ptr, mask=state_mask) - tl.store(trajectory_ti + slot, pti, mask=state_mask) - tl.store(trajectory_vr + slot + minus_plane, mvr, mask=state_mask) - tl.store(trajectory_vi + slot + minus_plane, mvi, mask=state_mask) - tl.store(trajectory_tr + slot + minus_plane, mtr, mask=state_mask) - tl.store(trajectory_ti + slot + minus_plane, mti, mask=state_mask) - tl.store(trajectory_vr + slot + long_plane, zvr, mask=state_mask) - tl.store(trajectory_vi + slot + long_plane, zvi, mask=state_mask) - tl.store(trajectory_tr + slot + long_plane, ztr, mask=state_mask) - tl.store(trajectory_ti + slot + long_plane, zti, mask=state_mask) - if pools > 0: - tl.store(trajectory_vr + slot + bound_plane, bvr, mask=state_mask) - tl.store(trajectory_vi + slot + bound_plane, bvi, mask=state_mask) - tl.store(trajectory_tr + slot + bound_plane, btr, mask=state_mask) - tl.store(trajectory_ti + slot + bound_plane, bti, mask=state_mask) - if pools == 2 or pools == 3: - tl.store(trajectory_vr + slot + bplus_plane, bpvr, mask=state_mask) - tl.store(trajectory_vi + slot + bplus_plane, bpvi, mask=state_mask) - tl.store(trajectory_tr + slot + bplus_plane, bptr, mask=state_mask) - tl.store(trajectory_ti + slot + bplus_plane, bpti, mask=state_mask) - tl.store(trajectory_vr + slot + bminus_plane, bmvr, mask=state_mask) - tl.store(trajectory_vi + slot + bminus_plane, bmvi, mask=state_mask) - tl.store(trajectory_tr + slot + bminus_plane, bmtr, mask=state_mask) - tl.store(trajectory_ti + slot + bminus_plane, bmti, mask=state_mask) - if pools == 3: - tl.store(trajectory_vr + slot + semisolid_plane, cvr, mask=state_mask) - tl.store(trajectory_vi + slot + semisolid_plane, cvi, mask=state_mask) - tl.store(trajectory_tr + slot + semisolid_plane, ctr, mask=state_mask) - tl.store(trajectory_ti + slot + semisolid_plane, cti, mask=state_mask) - - dt_value = _event_value( - duration, event_base, event, active_atom, single_train - ) - dt_tangent = _event_value( - dot_duration, event_base, event, active_atom, single_train - ) - wout_value = 1.0 - wout_tangent = 0.0 - if moving: - wout_value, wout_tangent = _washout_jvp( - atom_washout, d_washout, dt_value, dt_tangent - ) - dry1_value = tl.exp(-r1_value * dt_value) - dry1_tangent = -dry1_value * (r1_value * dt_tangent + r1_tangent * dt_value) - dry2_value = tl.exp(-r2_value * dt_value) - dry2_tangent = -dry2_value * (r2_value * dt_tangent + r2_tangent * dt_value) - e1_value = dry1_value * wout_value - e1_tangent = dry1_tangent * wout_value + dry1_value * wout_tangent - e2_value = dry2_value * wout_value - e2_tangent = dry2_tangent * wout_value + dry2_value * wout_tangent - damp_z = 1.0 - damp_z_tangent = 0.0 - damp_t = 1.0 - damp_t_tangent = 0.0 - if diffusing: - damp_z, damp_z_tangent, damp_t, damp_t_tangent = _damping_jvp( - atom_damping, d_damping, dt_value, dt_tangent, order - ) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value, recovery_tangent = 1.0 - e1_value, -e1_tangent - bare1_value, bare1_tangent = e1_value, e1_tangent - bare2_value, bare2_tangent = e2_value, e2_tangent - e1_tangent = e1_tangent * damp_z + bare1_value * damp_z_tangent - e1_value = bare1_value * damp_z - e2_tangent = e2_tangent * damp_t + bare2_value * damp_t_tangent - e2_value = bare2_value * damp_t - turn_t = 0.0 - dturn_t = 0.0 - szr, szi, sztr, szti = 1.0, 0.0, 0.0, 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, dt_value, order) - d_turn = d_flow * dt_value + atom_flow * dt_tangent - dturn_z = -order * d_turn - dturn_t = -(order + 0.5) * d_turn - szr, szi, sztr, szti = _dual_polar(turn_z, dturn_z) - qr, qi, qtr, qti = 1.0, 0.0, 0.0, 0.0 - if off_axis or moving: - angle_value = -2.0 * 3.141592653589793 * (atom_b0 * dt_value) + turn_t - angle_tangent = ( - -2.0 * 3.141592653589793 * (d_b0 * dt_value + atom_b0 * dt_tangent) - + dturn_t - ) - qr, qi, qtr, qti = _dual_polar(angle_value, angle_tangent) - ovr, ovi, otr, oti = _dual_scale(e2_value, e2_tangent, qr, qi, qtr, qti) - lvr, lvi, ltr, lti = _dual_scale(e1_value, e1_tangent, szr, szi, sztr, szti) - - # The damping and the off-resonance turn both pools take; with an - # exchanging one the relaxation itself sits inside the operator instead - # of in the scalar the free pool alone multiplies by. - carried = _dual_scale(damp_t, damp_t_tangent, qr, qi, qtr, qti) - if pools == 2 or pools == 3: - across = _two_pool_transverse_step_jvp( - r2_value, - r2_tangent, - r2b_value, - r2b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - atom_free, - d_free, - atom_shift, - d_shift, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - ) - a11 = (across[0], across[1], across[8], across[9]) - a12 = (across[2], across[3], across[10], across[11]) - a21 = (across[4], across[5], across[12], across[13]) - a22 = (across[6], across[7], across[14], across[15]) - free_plus = (pvr, pvi, ptr, pti) - pool_plus = (bpvr, bpvi, bptr, bpti) - free_minus = (mvr, mvi, mtr, mti) - pool_minus = (bmvr, bmvi, bmtr, bmti) - conjugated = _dual_conj(carried) - # ``F-`` takes the conjugate of the operator entry by entry, not - # its transpose: it is the conjugate state following the conjugate - # map. - pvr, pvi, ptr, pti = _dual_product( - _dual_add( - _dual_product(a11, free_plus), _dual_product(a12, pool_plus) - ), - carried, - ) - bpvr, bpvi, bptr, bpti = _dual_product( - _dual_add( - _dual_product(a21, free_plus), _dual_product(a22, pool_plus) - ), - carried, - ) - mvr, mvi, mtr, mti = _dual_product( - _dual_add( - _dual_product(_dual_conj(a11), free_minus), - _dual_product(_dual_conj(a12), pool_minus), - ), - conjugated, - ) - bmvr, bmvi, bmtr, bmti = _dual_product( - _dual_add( - _dual_product(_dual_conj(a21), free_minus), - _dual_product(_dual_conj(a22), pool_minus), - ), - conjugated, - ) - else: - pvr, pvi, ptr, pti = _dual_mul(ovr, ovi, otr, oti, pvr, pvi, ptr, pti) - mvr, mvi, mtr, mti = _dual_mul(ovr, -ovi, otr, -oti, mvr, mvi, mtr, mti) - if pools == 3: - # Three pools mix through a 3x3 formed in double, tangent and all: - # a direction through an operator this ill-conditioned needs the - # width as much as the value does. - if tabulated: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - d_t11, - d_t12, - d_t13, - d_t21, - d_t22, - d_t23, - d_t31, - d_t32, - d_t33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_from_table_jvp( - pool_table, - tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ), - atom, - atom_count, - active_atom, - r1_value, - r1b_value, - r1c_value, - atom_exchange, - atom_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_tangent, - wout_value, - wout_tangent, - ) - else: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - d_t11, - d_t12, - d_t13, - d_t21, - d_t22, - d_t23, - d_t31, - d_t32, - d_t33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_step_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - narrow, - ) - spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti) - was_free = (zvr, zvi, ztr, zti) - was_pool_b = (bvr, bvi, btr, bti) - was_semisolid = (cvr, cvi, ctr, cti) - mixed_free = _dual_add( - _dual_add( - _dual_scale( - t11, - d_t11, - was_free[0], - was_free[1], - was_free[2], - was_free[3], - ), - _dual_scale( - t12, - d_t12, - was_pool_b[0], - was_pool_b[1], - was_pool_b[2], - was_pool_b[3], - ), - ), - _dual_scale( - t13, - d_t13, - was_semisolid[0], - was_semisolid[1], - was_semisolid[2], - was_semisolid[3], - ), - ) - mixed_pool_b = _dual_add( - _dual_add( - _dual_scale( - t21, - d_t21, - was_free[0], - was_free[1], - was_free[2], - was_free[3], - ), - _dual_scale( - t22, - d_t22, - was_pool_b[0], - was_pool_b[1], - was_pool_b[2], - was_pool_b[3], - ), - ), - _dual_scale( - t23, - d_t23, - was_semisolid[0], - was_semisolid[1], - was_semisolid[2], - was_semisolid[3], - ), - ) - mixed_semisolid = _dual_add( - _dual_add( - _dual_scale( - t31, - d_t31, - was_free[0], - was_free[1], - was_free[2], - was_free[3], - ), - _dual_scale( - t32, - d_t32, - was_pool_b[0], - was_pool_b[1], - was_pool_b[2], - was_pool_b[3], - ), - ), - _dual_scale( - t33, - d_t33, - was_semisolid[0], - was_semisolid[1], - was_semisolid[2], - was_semisolid[3], - ), - ) - zvr, zvi, ztr, zti = _dual_mul( - spin[0], - spin[1], - spin[2], - spin[3], - mixed_free[0], - mixed_free[1], - mixed_free[2], - mixed_free[3], - ) - bvr, bvi, btr, bti = _dual_mul( - spin[0], - spin[1], - spin[2], - spin[3], - mixed_pool_b[0], - mixed_pool_b[1], - mixed_pool_b[2], - mixed_pool_b[3], - ) - cvr, cvi, ctr, cti = _dual_mul( - spin[0], - spin[1], - spin[2], - spin[3], - mixed_semisolid[0], - mixed_semisolid[1], - mixed_semisolid[2], - mixed_semisolid[3], - ) - zvr += tl.where(state == 0, grow_free, 0.0) - ztr += tl.where(state == 0, d_grow_free, 0.0) - bvr += tl.where(state == 0, grow_pool_b, 0.0) - btr += tl.where(state == 0, d_grow_pool_b, 0.0) - cvr += tl.where(state == 0, grow_semisolid, 0.0) - ctr += tl.where(state == 0, d_grow_semisolid, 0.0) - elif pools > 0: - # The exchange operator is a property of the interval, not of a - # dephasing order, so it is formed once and the per-order damping - # and turn multiply it. - ( - pe11, - pe12, - pe21, - pe22, - prec_f, - prec_b, - de11, - de12, - de21, - de22, - drec_f, - drec_b, - ) = _two_pool_step_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - ) - spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti) - free_part = _dual_scale(pe11, de11, zvr, zvi, ztr, zti) - cross_in = _dual_scale(pe12, de12, bvr, bvi, btr, bti) - cross_out = _dual_scale(pe21, de21, zvr, zvi, ztr, zti) - bound_part = _dual_scale(pe22, de22, bvr, bvi, btr, bti) - zvr, zvi, ztr, zti = _dual_mul( - spin[0], - spin[1], - spin[2], - spin[3], - free_part[0] + cross_in[0], - free_part[1] + cross_in[1], - free_part[2] + cross_in[2], - free_part[3] + cross_in[3], - ) - bvr, bvi, btr, bti = _dual_mul( - spin[0], - spin[1], - spin[2], - spin[3], - cross_out[0] + bound_part[0], - cross_out[1] + bound_part[1], - cross_out[2] + bound_part[2], - cross_out[3] + bound_part[3], - ) - zvr += tl.where(state == 0, prec_f, 0.0) - ztr += tl.where(state == 0, drec_f, 0.0) - bvr += tl.where(state == 0, prec_b, 0.0) - btr += tl.where(state == 0, drec_b, 0.0) - else: - zvr, zvi, ztr, zti = _dual_mul(lvr, lvi, ltr, lti, zvr, zvi, ztr, zti) - zvr += tl.where(state == 0, recovery_value, 0.0) - ztr += tl.where(state == 0, recovery_tangent, 0.0) - - event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - svr, svi, wvr, wvi = _shift( - pvr, pvi, mvr, mvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - ptr, pti, mtr, mti, state, state_mask, state_count - ) - pvr = tl.where(pre_shift, svr, pvr) - pvi = tl.where(pre_shift, svi, pvi) - ptr = tl.where(pre_shift, str_, ptr) - pti = tl.where(pre_shift, sti, pti) - mvr = tl.where(pre_shift, wvr, mvr) - mvi = tl.where(pre_shift, wvi, mvi) - mtr = tl.where(pre_shift, wtr, mtr) - mti = tl.where(pre_shift, wti, mti) - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - bptr, bpti, bmtr, bmti, state, state_mask, state_count - ) - bpvr = tl.where(pre_shift, svr, bpvr) - bpvi = tl.where(pre_shift, svi, bpvi) - bptr = tl.where(pre_shift, str_, bptr) - bpti = tl.where(pre_shift, sti, bpti) - bmvr = tl.where(pre_shift, wvr, bmvr) - bmvi = tl.where(pre_shift, wvi, bmvi) - bmtr = tl.where(pre_shift, wtr, bmtr) - bmti = tl.where(pre_shift, wti, bmti) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - ivr, ivi, itr, iti = _dual_scale(-atom_inv, -d_inv, zvr, zvi, ztr, zti) - zvr = tl.where(invert, ivr, zvr) - zvi = tl.where(invert, ivi, zvi) - ztr = tl.where(invert, itr, ztr) - zti = tl.where(invert, iti, zti) - if pools == 2 or pools == 3: - # A semisolid pool is saturated by an adiabatic sweep rather than - # turned over; a chemically exchanging one is free water and - # inverts like any other. - ivr, ivi, itr, iti = _dual_scale(-atom_inv, -d_inv, bvr, bvi, btr, bti) - bvr = tl.where(invert, ivr, bvr) - bvi = tl.where(invert, ivi, bvi) - btr = tl.where(invert, itr, btr) - bti = tl.where(invert, iti, bti) - - event_flip = _event_value( - flip, event_base, event, active_atom, single_train - ) - event_dot_flip = _event_value( - dot_flip, event_base, event, active_atom, single_train - ) - event_phase = _event_value( - phase, event_base, event, active_atom, single_train - ) - event_dot_phase = _event_value( - dot_phase, event_base, event, active_atom, single_train - ) - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse a row of its own. - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - if off_axis: - atom_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - d_b1 = tl.load(dot_b1 + row + atom, mask=active_atom, other=0.0) - if off_axis: - d_b1_phase = tl.load( - dot_b1_phase + row + atom, mask=active_atom, other=0.0 - ) - alpha_value = event_flip * atom_b1 - alpha_tangent = event_dot_flip * atom_b1 + event_flip * d_b1 - phi_value = event_phase + atom_b1_phase - phi_tangent = event_dot_phase + d_b1_phase - if pools == 1 or pools == 3: - # The semisolid pool absorbs the power the pulse deposits, so it - # reads the bare flip the transmit field gives the voxel -- not the - # slice-shaped rotation the free pool takes from the table. - offset_value = tl.load(rf_frequency + event) - atom_b0 - shape_value, shape_slope = _lineshape_at_slope( - lineshape, offset_value, lineshape_bins, lineshape_step - ) - shape_tangent = shape_slope * -d_b0 - event_saturation = tl.load(saturation + event) - power_value = event_saturation * alpha_value * alpha_value - power_tangent = event_saturation * 2.0 * alpha_value * alpha_tangent - absorbed_value = tl.exp(power_value * shape_value) - absorbed_tangent = absorbed_value * ( - power_tangent * shape_value + power_value * shape_tangent - ) - saturating = is_rf & ~is_inversion - if pools == 1: - sat_b = _dual_scale( - absorbed_value, absorbed_tangent, bvr, bvi, btr, bti - ) - bvr = tl.where(saturating, sat_b[0], bvr) - bvi = tl.where(saturating, sat_b[1], bvi) - btr = tl.where(saturating, sat_b[2], btr) - bti = tl.where(saturating, sat_b[3], bti) - else: - sat_c = _dual_scale( - absorbed_value, absorbed_tangent, cvr, cvi, ctr, cti - ) - cvr = tl.where(saturating, sat_c[0], cvr) - cvi = tl.where(saturating, sat_c[1], cvi) - ctr = tl.where(saturating, sat_c[2], ctr) - cti = tl.where(saturating, sat_c[3], cti) - cos_value = tl.cos(alpha_value) - sin_value = tl.sin(alpha_value) - cos_tangent = -sin_value * alpha_tangent - sin_tangent = cos_value * alpha_tangent - p1r, p1i, p1tr, p1ti = _dual_polar(phi_value, phi_tangent) - p2r, p2i, p2tr, p2ti = _dual_mul(p1r, p1i, p1tr, p1ti, p1r, p1i, p1tr, p1ti) - t00, t01, t02, t12, t20, t21, t22 = _rotation_block( - 0.5 * (1.0 + cos_value), - 0.5 * cos_tangent, - 0.5 * (1.0 - cos_value), - -0.5 * cos_tangent, - sin_value, - sin_tangent, - cos_value, - cos_tangent, - p1r, - p1i, - p1tr, - p1ti, - p2r, - p2i, - p2tr, - p2ti, - p1r, - -p1i, - p1tr, - -p1ti, - ) - a0 = _dual_mul(t00[0], t00[1], t00[2], t00[3], pvr, pvi, ptr, pti) - a1 = _dual_mul(t01[0], t01[1], t01[2], t01[3], mvr, mvi, mtr, mti) - a2 = _dual_mul(t02[0], t02[1], t02[2], t02[3], zvr, zvi, ztr, zti) - b0_ = _dual_mul(t01[0], -t01[1], t01[2], -t01[3], pvr, pvi, ptr, pti) - b1_ = _dual_mul(t00[0], t00[1], t00[2], t00[3], mvr, mvi, mtr, mti) - b2 = _dual_mul(t12[0], t12[1], t12[2], t12[3], zvr, zvi, ztr, zti) - c0 = _dual_mul(t20[0], t20[1], t20[2], t20[3], pvr, pvi, ptr, pti) - c1 = _dual_mul(t21[0], t21[1], t21[2], t21[3], mvr, mvi, mtr, mti) - c2 = _dual_mul(t22[0], t22[1], t22[2], t22[3], zvr, zvi, ztr, zti) - - turned_pvr = a0[0] + a1[0] + a2[0] - turned_pvi = a0[1] + a1[1] + a2[1] - turned_ptr = a0[2] + a1[2] + a2[2] - turned_pti = a0[3] + a1[3] + a2[3] - turned_mvr = b0_[0] + b1_[0] + b2[0] - turned_mvi = b0_[1] + b1_[1] + b2[1] - turned_mtr = b0_[2] + b1_[2] + b2[2] - turned_mti = b0_[3] + b1_[3] + b2[3] - turned_zvr = c0[0] + c1[0] + c2[0] - turned_zvi = c0[1] + c1[1] + c2[1] - turned_ztr = c0[2] + c1[2] + c2[2] - turned_zti = c0[3] + c1[3] + c2[3] - if profiled or dynamic: - if dynamic: - shaped_a, shaped_b = _dynamic_pair_dual_at( - pairs, - pair_direction, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - phi_value, - phi_tangent, - directed, - ) - else: - shaped_a, shaped_b, _, _ = _profiled_pair_dual( - profile, - _table_row(profile_index, event, location, locations), - alpha_value, - alpha_tangent, - phi_value, - phi_tangent, - profile_bins, - profile_step, - ) - ( - turned_pvr, - turned_pvi, - turned_mvr, - turned_mvi, - turned_zvr, - turned_zvi, - turned_ptr, - turned_pti, - turned_mtr, - turned_mti, - turned_ztr, - turned_zti, - ) = _rotate_spinor_dual( - shaped_a[0], - shaped_a[1], - shaped_b[0], - shaped_b[1], - shaped_a[2], - shaped_a[3], - shaped_b[2], - shaped_b[3], - pvr, - pvi, - mvr, - mvi, - zvr, - zvi, - ptr, - pti, - mtr, - mti, - ztr, - zti, - ) - - rotate = is_rf & ~is_inversion - if pools == 2 or pools == 3: - # The same pulse, the same rotation. A chemical shift moves where a - # pool precesses, not what a pulse does to it. - e0 = _dual_mul(t00[0], t00[1], t00[2], t00[3], bpvr, bpvi, bptr, bpti) - e1_ = _dual_mul(t01[0], t01[1], t01[2], t01[3], bmvr, bmvi, bmtr, bmti) - e2_ = _dual_mul(t02[0], t02[1], t02[2], t02[3], bvr, bvi, btr, bti) - f0 = _dual_mul(t01[0], -t01[1], t01[2], -t01[3], bpvr, bpvi, bptr, bpti) - f1 = _dual_mul(t00[0], t00[1], t00[2], t00[3], bmvr, bmvi, bmtr, bmti) - f2 = _dual_mul(t12[0], t12[1], t12[2], t12[3], bvr, bvi, btr, bti) - h0 = _dual_mul(t20[0], t20[1], t20[2], t20[3], bpvr, bpvi, bptr, bpti) - h1 = _dual_mul(t21[0], t21[1], t21[2], t21[3], bmvr, bmvi, bmtr, bmti) - h2 = _dual_mul(t22[0], t22[1], t22[2], t22[3], bvr, bvi, btr, bti) - spun_pvr = e0[0] + e1_[0] + e2_[0] - spun_pvi = e0[1] + e1_[1] + e2_[1] - spun_ptr = e0[2] + e1_[2] + e2_[2] - spun_pti = e0[3] + e1_[3] + e2_[3] - spun_mvr = f0[0] + f1[0] + f2[0] - spun_mvi = f0[1] + f1[1] + f2[1] - spun_mtr = f0[2] + f1[2] + f2[2] - spun_mti = f0[3] + f1[3] + f2[3] - spun_zvr = h0[0] + h1[0] + h2[0] - spun_zvi = h0[1] + h1[1] + h2[1] - spun_ztr = h0[2] + h1[2] + h2[2] - spun_zti = h0[3] + h1[3] + h2[3] - if profiled or dynamic: - ( - spun_pvr, - spun_pvi, - spun_mvr, - spun_mvi, - spun_zvr, - spun_zvi, - spun_ptr, - spun_pti, - spun_mtr, - spun_mti, - spun_ztr, - spun_zti, - ) = _rotate_spinor_dual( - shaped_a[0], - shaped_a[1], - shaped_b[0], - shaped_b[1], - shaped_a[2], - shaped_a[3], - shaped_b[2], - shaped_b[3], - bpvr, - bpvi, - bmvr, - bmvi, - bvr, - bvi, - bptr, - bpti, - bmtr, - bmti, - btr, - bti, - ) - bpvr = tl.where(rotate, spun_pvr, bpvr) - bpvi = tl.where(rotate, spun_pvi, bpvi) - bptr = tl.where(rotate, spun_ptr, bptr) - bpti = tl.where(rotate, spun_pti, bpti) - bmvr = tl.where(rotate, spun_mvr, bmvr) - bmvi = tl.where(rotate, spun_mvi, bmvi) - bmtr = tl.where(rotate, spun_mtr, bmtr) - bmti = tl.where(rotate, spun_mti, bmti) - bvr = tl.where(rotate, spun_zvr, bvr) - bvi = tl.where(rotate, spun_zvi, bvi) - btr = tl.where(rotate, spun_ztr, btr) - bti = tl.where(rotate, spun_zti, bti) - pvr = tl.where(rotate, turned_pvr, pvr) - pvi = tl.where(rotate, turned_pvi, pvi) - ptr = tl.where(rotate, turned_ptr, ptr) - pti = tl.where(rotate, turned_pti, pti) - mvr = tl.where(rotate, turned_mvr, mvr) - mvi = tl.where(rotate, turned_mvi, mvi) - mtr = tl.where(rotate, turned_mtr, mtr) - mti = tl.where(rotate, turned_mti, mti) - zvr = tl.where(rotate, turned_zvr, zvr) - zvi = tl.where(rotate, turned_zvi, zvi) - ztr = tl.where(rotate, turned_ztr, ztr) - zti = tl.where(rotate, turned_zti, zti) - - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - svr, svi, wvr, wvi = _shift( - pvr, pvi, mvr, mvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - ptr, pti, mtr, mti, state, state_mask, state_count - ) - pvr = tl.where(do_shift, svr, pvr) - pvi = tl.where(do_shift, svi, pvi) - ptr = tl.where(do_shift, str_, ptr) - pti = tl.where(do_shift, sti, pti) - mvr = tl.where(do_shift, wvr, mvr) - mvi = tl.where(do_shift, wvi, mvi) - mtr = tl.where(do_shift, wtr, mtr) - mti = tl.where(do_shift, wti, mti) - spoil = (event_action & 8) != 0 - pvr = tl.where(spoil, 0.0, pvr) - pvi = tl.where(spoil, 0.0, pvi) - ptr = tl.where(spoil, 0.0, ptr) - pti = tl.where(spoil, 0.0, pti) - mvr = tl.where(spoil, 0.0, mvr) - mvi = tl.where(spoil, 0.0, mvi) - mtr = tl.where(spoil, 0.0, mtr) - mti = tl.where(spoil, 0.0, mti) - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - bpvr, bpvi, bmvr, bmvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - bptr, bpti, bmtr, bmti, state, state_mask, state_count - ) - bpvr = tl.where(spoil, 0.0, tl.where(do_shift, svr, bpvr)) - bpvi = tl.where(spoil, 0.0, tl.where(do_shift, svi, bpvi)) - bptr = tl.where(spoil, 0.0, tl.where(do_shift, str_, bptr)) - bpti = tl.where(spoil, 0.0, tl.where(do_shift, sti, bpti)) - bmvr = tl.where(spoil, 0.0, tl.where(do_shift, wvr, bmvr)) - bmvi = tl.where(spoil, 0.0, tl.where(do_shift, wvi, bmvi)) - bmtr = tl.where(spoil, 0.0, tl.where(do_shift, wtr, bmtr)) - bmti = tl.where(spoil, 0.0, tl.where(do_shift, wti, bmti)) - return - - # ---- reverse ---- - pbvr = empty - pbvi = empty - pbtr = empty - pbti = empty - mbvr = empty - mbvi = empty - mbtr = empty - mbti = empty - zbvr = empty - zbvi = empty - zbtr = empty - zbti = empty - bbvr = empty - bbvi = empty - bbtr = empty - bbti = empty - ubvr = empty - ubvi = empty - ubtr = empty - ubti = empty - wbvr = empty - wbvi = empty - wbtr = empty - wbti = empty - cbvr = empty - cbvi = empty - cbtr = empty - cbti = empty - zero = tl.zeros((problems, 1), tl.float32) - g_boundv = zero - g_boundt = zero - g_exchv = zero - g_excht = zero - g_t1bv = zero - g_t1bt = zero - g_t2bv = zero - g_t2bt = zero - g_shiftv = zero - g_shiftt = zero - g_semiv = zero - g_semit = zero - g_sexchv = zero - g_sexcht = zero - g_t1cv = zero - g_t1ct = zero - g_diffv = zero - g_difft = zero - g_flowv = zero - g_flowt = zero - g_washv = zero - g_washt = zero - g_t1v = zero - g_t1t = zero - g_t2v = zero - g_t2t = zero - g_m0v = zero - g_m0t = zero - g_b1v = zero - g_b1t = zero - g_b1pv = zero - g_b1pt = zero - g_b0v = zero - g_b0t = zero - g_invv = zero - g_invt = zero - - for reverse in range(0, event_count): - event = event_count - 1 - reverse - slot = trajectory + event * record_stride - xpvr = tl.load(trajectory_vr + slot, mask=state_mask, other=0.0) - xpvi = tl.load(trajectory_vi + slot, mask=state_mask, other=0.0) - xptr = tl.load(trajectory_tr + slot, mask=state_mask, other=0.0) - xpti = tl.load(trajectory_ti + slot, mask=state_mask, other=0.0) - xmvr = tl.load(trajectory_vr + slot + minus_plane, mask=state_mask, other=0.0) - xmvi = tl.load(trajectory_vi + slot + minus_plane, mask=state_mask, other=0.0) - xmtr = tl.load(trajectory_tr + slot + minus_plane, mask=state_mask, other=0.0) - xmti = tl.load(trajectory_ti + slot + minus_plane, mask=state_mask, other=0.0) - xzvr = tl.load(trajectory_vr + slot + long_plane, mask=state_mask, other=0.0) - xzvi = tl.load(trajectory_vi + slot + long_plane, mask=state_mask, other=0.0) - xztr = tl.load(trajectory_tr + slot + long_plane, mask=state_mask, other=0.0) - xzti = tl.load(trajectory_ti + slot + long_plane, mask=state_mask, other=0.0) - xbvr = empty - xbvi = empty - xbtr = empty - xbti = empty - xbpvr = empty - xbpvi = empty - xbptr = empty - xbpti = empty - xbmvr = empty - xbmvi = empty - xbmtr = empty - xbmti = empty - xcvr = empty - xcvi = empty - xctr = empty - xcti = empty - if pools > 0: - xbvr = tl.load( - trajectory_vr + slot + bound_plane, mask=state_mask, other=0.0 - ) - xbvi = tl.load( - trajectory_vi + slot + bound_plane, mask=state_mask, other=0.0 - ) - xbtr = tl.load( - trajectory_tr + slot + bound_plane, mask=state_mask, other=0.0 - ) - xbti = tl.load( - trajectory_ti + slot + bound_plane, mask=state_mask, other=0.0 - ) - if pools == 2 or pools == 3: - xbpvr = tl.load( - trajectory_vr + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbpvi = tl.load( - trajectory_vi + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbptr = tl.load( - trajectory_tr + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbpti = tl.load( - trajectory_ti + slot + bplus_plane, mask=state_mask, other=0.0 - ) - xbmvr = tl.load( - trajectory_vr + slot + bminus_plane, mask=state_mask, other=0.0 - ) - xbmvi = tl.load( - trajectory_vi + slot + bminus_plane, mask=state_mask, other=0.0 - ) - xbmtr = tl.load( - trajectory_tr + slot + bminus_plane, mask=state_mask, other=0.0 - ) - xbmti = tl.load( - trajectory_ti + slot + bminus_plane, mask=state_mask, other=0.0 - ) - if pools == 3: - xcvr = tl.load( - trajectory_vr + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - xcvi = tl.load( - trajectory_vi + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - xctr = tl.load( - trajectory_tr + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - xcti = tl.load( - trajectory_ti + slot + semisolid_plane, mask=state_mask, other=0.0 - ) - - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event) - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - dt_tangent = _event_value( - dot_duration, event_base, event, active_atom, single_train - ) - wout_value = 1.0 - wout_tangent = 0.0 - if moving: - wout_value, wout_tangent = _washout_jvp( - atom_washout, d_washout, dt_value, dt_tangent - ) - dry1_value = tl.exp(-r1_value * dt_value) - dry1_tangent = -dry1_value * (r1_value * dt_tangent + r1_tangent * dt_value) - dry2_value = tl.exp(-r2_value * dt_value) - dry2_tangent = -dry2_value * (r2_value * dt_tangent + r2_tangent * dt_value) - e1_value = dry1_value * wout_value - e1_tangent = dry1_tangent * wout_value + dry1_value * wout_tangent - e2_value = dry2_value * wout_value - e2_tangent = dry2_tangent * wout_value + dry2_value * wout_tangent - damp_z = 1.0 - damp_z_tangent = 0.0 - damp_t = 1.0 - damp_t_tangent = 0.0 - if diffusing: - damp_z, damp_z_tangent, damp_t, damp_t_tangent = _damping_jvp( - atom_damping, d_damping, dt_value, dt_tangent, order - ) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value, recovery_tangent = 1.0 - e1_value, -e1_tangent - bare1_value, bare1_tangent = e1_value, e1_tangent - bare2_value, bare2_tangent = e2_value, e2_tangent - e1_tangent = e1_tangent * damp_z + bare1_value * damp_z_tangent - e1_value = bare1_value * damp_z - e2_tangent = e2_tangent * damp_t + bare2_value * damp_t_tangent - e2_value = bare2_value * damp_t - turn_t = 0.0 - dturn_t = 0.0 - szr, szi, sztr, szti = 1.0, 0.0, 0.0, 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, dt_value, order) - d_turn = d_flow * dt_value + atom_flow * dt_tangent - dturn_z = -order * d_turn - dturn_t = -(order + 0.5) * d_turn - szr, szi, sztr, szti = _dual_polar(turn_z, dturn_z) - qr, qi, qtr, qti = 1.0, 0.0, 0.0, 0.0 - if off_axis or moving: - angle_value = -2.0 * 3.141592653589793 * (atom_b0 * dt_value) + turn_t - angle_tangent = ( - -2.0 * 3.141592653589793 * (d_b0 * dt_value + atom_b0 * dt_tangent) - + dturn_t - ) - qr, qi, qtr, qti = _dual_polar(angle_value, angle_tangent) - ovr, ovi, otr, oti = _dual_scale(e2_value, e2_tangent, qr, qi, qtr, qti) - lvr, lvi, ltr, lti = _dual_scale(e1_value, e1_tangent, szr, szi, sztr, szti) - - # Replay the intra-event stages from the recorded entry state. - carried = _dual_scale(damp_t, damp_t_tangent, qr, qi, qtr, qti) - rbpvr = empty - rbpvi = empty - rbptr = empty - rbpti = empty - rbmvr = empty - rbmvi = empty - rbmtr = empty - rbmti = empty - a11 = (empty, empty, empty, empty) - a12 = (empty, empty, empty, empty) - a21 = (empty, empty, empty, empty) - a22 = (empty, empty, empty, empty) - if pools == 2 or pools == 3: - across = _two_pool_transverse_step_jvp( - r2_value, - r2_tangent, - r2b_value, - r2b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - atom_free, - d_free, - atom_shift, - d_shift, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - ) - a11 = (across[0], across[1], across[8], across[9]) - a12 = (across[2], across[3], across[10], across[11]) - a21 = (across[4], across[5], across[12], across[13]) - a22 = (across[6], across[7], across[14], across[15]) - free_plus = (xpvr, xpvi, xptr, xpti) - pool_plus = (xbpvr, xbpvi, xbptr, xbpti) - free_minus = (xmvr, xmvi, xmtr, xmti) - pool_minus = (xbmvr, xbmvi, xbmtr, xbmti) - conjugated = _dual_conj(carried) - rpvr, rpvi, rptr, rpti = _dual_product( - _dual_add(_dual_product(a11, free_plus), _dual_product(a12, pool_plus)), - carried, - ) - rbpvr, rbpvi, rbptr, rbpti = _dual_product( - _dual_add(_dual_product(a21, free_plus), _dual_product(a22, pool_plus)), - carried, - ) - rmvr, rmvi, rmtr, rmti = _dual_product( - _dual_add( - _dual_product(_dual_conj(a11), free_minus), - _dual_product(_dual_conj(a12), pool_minus), - ), - conjugated, - ) - rbmvr, rbmvi, rbmtr, rbmti = _dual_product( - _dual_add( - _dual_product(_dual_conj(a21), free_minus), - _dual_product(_dual_conj(a22), pool_minus), - ), - conjugated, - ) - else: - rpvr, rpvi, rptr, rpti = _dual_mul( - ovr, ovi, otr, oti, xpvr, xpvi, xptr, xpti - ) - rmvr, rmvi, rmtr, rmti = _dual_mul( - ovr, -ovi, otr, -oti, xmvr, xmvi, xmtr, xmti - ) - rbvr = empty - rbvi = empty - rbtr = empty - rbti = empty - rcvr = empty - rcvi = empty - rctr = empty - rcti = empty - if pools == 3: - if tabulated: - # The walk back needs the operator and the direction - # through it, which the row already holds -- and pooling - # the cotangents took what the eigenvalues were formed - # for, so nothing here reads them. - pool_row = tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ) - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - d_w11, - d_w12, - d_w13, - d_w21, - d_w22, - d_w23, - d_w31, - d_w32, - d_w33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_from_table_jvp( - pool_table, - pool_row, - atom, - atom_count, - active_atom, - r1_value, - r1b_value, - r1c_value, - atom_exchange, - atom_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_tangent, - wout_value, - wout_tangent, - ) - else: - # The pieces and the bare operator are kept rather than - # the step alone: the walk back pushes the cotangents - # through them, and forming them once for the interval - # is what keeps this kernel a size a compiler will take. - ( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - ) = _three_pool_pieces_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_value, - dt_tangent, - narrow, - ) - ( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - ) = _three_pool_assemble_jvp( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - narrow, - ) - ( - w11, - w12, - w13, - w21, - w22, - w23, - w31, - w32, - w33, - grow_free, - grow_pool_b, - grow_semisolid, - d_w11, - d_w12, - d_w13, - d_w21, - d_w22, - d_w23, - d_w31, - d_w32, - d_w33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_weigh_jvp( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - wout_value, - wout_tangent, - narrow, - ) - # The operator is O(1) once formed, so the per-order loop below - # takes it at the width the states are carried in. - w11 = w11.to(tl.float32) - w12 = w12.to(tl.float32) - w13 = w13.to(tl.float32) - w21 = w21.to(tl.float32) - w22 = w22.to(tl.float32) - w23 = w23.to(tl.float32) - w31 = w31.to(tl.float32) - w32 = w32.to(tl.float32) - w33 = w33.to(tl.float32) - grow_free = grow_free.to(tl.float32) - grow_pool_b = grow_pool_b.to(tl.float32) - grow_semisolid = grow_semisolid.to(tl.float32) - d_w11 = d_w11.to(tl.float32) - d_w12 = d_w12.to(tl.float32) - d_w13 = d_w13.to(tl.float32) - d_w21 = d_w21.to(tl.float32) - d_w22 = d_w22.to(tl.float32) - d_w23 = d_w23.to(tl.float32) - d_w31 = d_w31.to(tl.float32) - d_w32 = d_w32.to(tl.float32) - d_w33 = d_w33.to(tl.float32) - d_grow_free = d_grow_free.to(tl.float32) - d_grow_pool_b = d_grow_pool_b.to(tl.float32) - d_grow_semisolid = d_grow_semisolid.to(tl.float32) - spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti) - mixed_free = _dual_add( - _dual_add( - _dual_scale(w11, d_w11, xzvr, xzvi, xztr, xzti), - _dual_scale(w12, d_w12, xbvr, xbvi, xbtr, xbti), - ), - _dual_scale(w13, d_w13, xcvr, xcvi, xctr, xcti), - ) - mixed_bound = _dual_add( - _dual_add( - _dual_scale(w21, d_w21, xzvr, xzvi, xztr, xzti), - _dual_scale(w22, d_w22, xbvr, xbvi, xbtr, xbti), - ), - _dual_scale(w23, d_w23, xcvr, xcvi, xctr, xcti), - ) - mixed_semisolid = _dual_add( - _dual_add( - _dual_scale(w31, d_w31, xzvr, xzvi, xztr, xzti), - _dual_scale(w32, d_w32, xbvr, xbvi, xbtr, xbti), - ), - _dual_scale(w33, d_w33, xcvr, xcvi, xctr, xcti), - ) - rzvr, rzvi, rztr, rzti = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_free - ) - rbvr, rbvi, rbtr, rbti = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_bound - ) - rcvr, rcvi, rctr, rcti = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_semisolid - ) - rzvr += tl.where(state == 0, grow_free, 0.0) - rztr += tl.where(state == 0, d_grow_free, 0.0) - rbvr += tl.where(state == 0, grow_pool_b, 0.0) - rbtr += tl.where(state == 0, d_grow_pool_b, 0.0) - rcvr += tl.where(state == 0, grow_semisolid, 0.0) - rctr += tl.where(state == 0, d_grow_semisolid, 0.0) - elif pools > 0: - ( - pe11, - pe12, - pe21, - pe22, - prec_f, - prec_b, - de11, - de12, - de21, - de22, - drec_f, - drec_b, - ) = _two_pool_step_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - ) - spin = _dual_scale(damp_z, damp_z_tangent, szr, szi, sztr, szti) - free_part = _dual_scale(pe11, de11, xzvr, xzvi, xztr, xzti) - cross_in = _dual_scale(pe12, de12, xbvr, xbvi, xbtr, xbti) - cross_out = _dual_scale(pe21, de21, xzvr, xzvi, xztr, xzti) - bound_part = _dual_scale(pe22, de22, xbvr, xbvi, xbtr, xbti) - mixed_free = ( - free_part[0] + cross_in[0], - free_part[1] + cross_in[1], - free_part[2] + cross_in[2], - free_part[3] + cross_in[3], - ) - mixed_bound = ( - cross_out[0] + bound_part[0], - cross_out[1] + bound_part[1], - cross_out[2] + bound_part[2], - cross_out[3] + bound_part[3], - ) - rzvr, rzvi, rztr, rzti = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_free - ) - rbvr, rbvi, rbtr, rbti = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_bound - ) - rzvr += tl.where(state == 0, prec_f, 0.0) - rztr += tl.where(state == 0, drec_f, 0.0) - rbvr += tl.where(state == 0, prec_b, 0.0) - rbtr += tl.where(state == 0, drec_b, 0.0) - else: - rzvr, rzvi, rztr, rzti = _dual_mul( - lvr, lvi, ltr, lti, xzvr, xzvi, xztr, xzti - ) - rzvr += tl.where(state == 0, recovery_value, 0.0) - rztr += tl.where(state == 0, recovery_tangent, 0.0) - - pre_shift = (event_action & 1) != 0 - svr, svi, wvr, wvi = _shift( - rpvr, rpvi, rmvr, rmvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - rptr, rpti, rmtr, rmti, state, state_mask, state_count - ) - spvr = tl.where(pre_shift, svr, rpvr) - spvi = tl.where(pre_shift, svi, rpvi) - sptr = tl.where(pre_shift, str_, rptr) - spti = tl.where(pre_shift, sti, rpti) - smvr = tl.where(pre_shift, wvr, rmvr) - smvi = tl.where(pre_shift, wvi, rmvi) - smtr = tl.where(pre_shift, wtr, rmtr) - smti = tl.where(pre_shift, wti, rmti) - sbpvr = rbpvr - sbpvi = rbpvi - sbptr = rbptr - sbpti = rbpti - sbmvr = rbmvr - sbmvi = rbmvi - sbmtr = rbmtr - sbmti = rbmti - if pools == 2 or pools == 3: - svr, svi, wvr, wvi = _shift( - rbpvr, rbpvi, rbmvr, rbmvi, state, state_mask, state_count - ) - str_, sti, wtr, wti = _shift( - rbptr, rbpti, rbmtr, rbmti, state, state_mask, state_count - ) - sbpvr = tl.where(pre_shift, svr, rbpvr) - sbpvi = tl.where(pre_shift, svi, rbpvi) - sbptr = tl.where(pre_shift, str_, rbptr) - sbpti = tl.where(pre_shift, sti, rbpti) - sbmvr = tl.where(pre_shift, wvr, rbmvr) - sbmvi = tl.where(pre_shift, wvi, rbmvi) - sbmtr = tl.where(pre_shift, wtr, rbmtr) - sbmti = tl.where(pre_shift, wti, rbmti) - - # Undo the trailing spoil or shift. - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - spoil = (event_action & 8) != 0 - avr, avi, bvr, bvi = _shift_adjoint( - pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count - ) - atr, ati, btr, bti = _shift_adjoint( - pbtr, pbti, mbtr, mbti, state, state_mask, state_count - ) - trailing = do_shift & ~spoil - pbvr = tl.where(spoil, 0.0, tl.where(trailing, avr, pbvr)) - pbvi = tl.where(spoil, 0.0, tl.where(trailing, avi, pbvi)) - pbtr = tl.where(spoil, 0.0, tl.where(trailing, atr, pbtr)) - pbti = tl.where(spoil, 0.0, tl.where(trailing, ati, pbti)) - mbvr = tl.where(spoil, 0.0, tl.where(trailing, bvr, mbvr)) - mbvi = tl.where(spoil, 0.0, tl.where(trailing, bvi, mbvi)) - mbtr = tl.where(spoil, 0.0, tl.where(trailing, btr, mbtr)) - mbti = tl.where(spoil, 0.0, tl.where(trailing, bti, mbti)) - if pools == 2 or pools == 3: - avr, avi, bvr, bvi = _shift_adjoint( - ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count - ) - atr, ati, btr, bti = _shift_adjoint( - ubtr, ubti, wbtr, wbti, state, state_mask, state_count - ) - ubvr = tl.where(spoil, 0.0, tl.where(trailing, avr, ubvr)) - ubvi = tl.where(spoil, 0.0, tl.where(trailing, avi, ubvi)) - ubtr = tl.where(spoil, 0.0, tl.where(trailing, atr, ubtr)) - ubti = tl.where(spoil, 0.0, tl.where(trailing, ati, ubti)) - wbvr = tl.where(spoil, 0.0, tl.where(trailing, bvr, wbvr)) - wbvi = tl.where(spoil, 0.0, tl.where(trailing, bvi, wbvi)) - wbtr = tl.where(spoil, 0.0, tl.where(trailing, btr, wbtr)) - wbti = tl.where(spoil, 0.0, tl.where(trailing, bti, wbti)) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - event_dot_flip = _event_value( - dot_flip, event_base, event, active_atom, single_train - ) - event_phase = _event_value(phase, event_base, event, active_atom, single_train) - event_dot_phase = _event_value( - dot_phase, event_base, event, active_atom, single_train - ) - - # ---- recorded sample ---- - record = ((event_action & 32) != 0) & (event_kind == 2) - out = tl.load(output_index + event) - seed_mask = active_atom & record & (out >= 0) - seed_real = tl.load( - grad_output_real + problem * output_count + out, mask=seed_mask, other=0.0 - ) - seed_imag = tl.load( - grad_output_imag + problem * output_count + out, mask=seed_mask, other=0.0 - ) - dvr, dvi, dtr, dti = _dual_polar(-event_phase, -event_dot_phase) - # A coil sees the whole voxel, so what it records is the sum over pools. - recorded = (spvr, spvi, sptr, spti) - if pools == 2 or pools == 3: - recorded = _dual_add(recorded, (sbpvr, sbpvi, sbptr, sbpti)) - # grad_m0 = Re(conj(seed) * recorded * demodulation) - wr, wi, wtr_, wti_ = _dual_mul(*recorded, dvr, dvi, dtr, dti) - m0_value, m0_tangent = _dual_real_conj_mul( - seed_real, - seed_imag, - 0.0 * seed_real, - 0.0 * seed_imag, - wr, - wi, - wtr_, - wti_, - ) - g_m0v += tl.sum(tl.where(state == 0, m0_value, 0.0), axis=1)[:, None] - g_m0t += tl.sum(tl.where(state == 0, m0_tangent, 0.0), axis=1)[:, None] - # grad_phase = Re(conj(seed) * m0 * recorded * (-i) * demodulation) - yr, yi, ytr, yti = _dual_scale(atom_m0, d_m0, *recorded) - yr, yi, ytr, yti = _dual_times_i(yr, yi, ytr, yti) - yr, yi, ytr, yti = -yr, -yi, -ytr, -yti - yr, yi, ytr, yti = _dual_mul(yr, yi, ytr, yti, dvr, dvi, dtr, dti) - phase_value, phase_tangent = _dual_real_conj_mul( - seed_real, seed_imag, 0.0 * seed_real, 0.0 * seed_imag, yr, yi, ytr, yti - ) - tl.atomic_add( - grad_phase_value + event_base + event, - tl.sum(tl.where(state == 0, phase_value, 0.0), axis=1)[:, None], - mask=seed_mask, - ) - tl.atomic_add( - grad_phase_tangent + event_base + event, - tl.sum(tl.where(state == 0, phase_tangent, 0.0), axis=1)[:, None], - mask=seed_mask, - ) - # fplus_bar[0] += conj(m0 * demodulation) * seed - kr, ki, ktr, kti = _dual_scale(atom_m0, d_m0, dvr, dvi, dtr, dti) - sr, si, stg_r, stg_i = _dual_mul( - kr, - -ki, - ktr, - -kti, - seed_real, - seed_imag, - 0.0 * seed_real, - 0.0 * seed_imag, - ) - pbvr += tl.where(state == 0, sr, 0.0) - pbvi += tl.where(state == 0, si, 0.0) - pbtr += tl.where(state == 0, stg_r, 0.0) - pbti += tl.where(state == 0, stg_i, 0.0) - if pools == 2 or pools == 3: - ubvr += tl.where(state == 0, sr, 0.0) - ubvi += tl.where(state == 0, si, 0.0) - ubtr += tl.where(state == 0, stg_r, 0.0) - ubti += tl.where(state == 0, stg_i, 0.0) - - # ---- RF adjoint ---- - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - inv_value, inv_tangent = _dual_real_conj_mul( - zbvr, zbvi, zbtr, zbti, -rzvr, -rzvi, -rztr, -rzti - ) - g_invv += tl.sum(tl.where(invert, inv_value, 0.0), axis=1)[:, None] - g_invt += tl.sum(tl.where(invert, inv_tangent, 0.0), axis=1)[:, None] - ivr, ivi, itr, iti = _dual_scale(-atom_inv, -d_inv, zbvr, zbvi, zbtr, zbti) - zbvr = tl.where(invert, ivr, zbvr) - zbvi = tl.where(invert, ivi, zbvi) - zbtr = tl.where(invert, itr, zbtr) - zbti = tl.where(invert, iti, zbti) - if pools == 2 or pools == 3: - pool_v, pool_t = _dual_real_conj_mul( - bbvr, bbvi, bbtr, bbti, -rbvr, -rbvi, -rbtr, -rbti - ) - g_invv += tl.sum(tl.where(invert, pool_v, 0.0), axis=1)[:, None] - g_invt += tl.sum(tl.where(invert, pool_t, 0.0), axis=1)[:, None] - ivr, ivi, itr, iti = _dual_scale(-atom_inv, -d_inv, bbvr, bbvi, bbtr, bbti) - bbvr = tl.where(invert, ivr, bbvr) - bbvi = tl.where(invert, ivi, bbvi) - bbtr = tl.where(invert, itr, bbtr) - bbti = tl.where(invert, iti, bbti) - - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse a row of its own. - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - if off_axis: - atom_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - d_b1 = tl.load(dot_b1 + row + atom, mask=active_atom, other=0.0) - if off_axis: - d_b1_phase = tl.load( - dot_b1_phase + row + atom, mask=active_atom, other=0.0 - ) - alpha_value = event_flip * atom_b1 - alpha_tangent = event_dot_flip * atom_b1 + event_flip * d_b1 - phi_value = event_phase + atom_b1_phase - phi_tangent = event_dot_phase + d_b1_phase - sat_alpha_v = zero - sat_alpha_t = zero - sat_b0_v = zero - sat_b0_t = zero - if broadened: - # The pulse scales every order of the bound pool by one real - # number, so its cotangent is a single sum over the states it - # multiplied. The lineshape's own slope is differentiated too, - # which is what the curvature the reader returns is for. - offset_value = tl.load(rf_frequency + event) - atom_b0 - shape_value, shape_slope, shape_curve = _lineshape_at_curve( - lineshape, offset_value, lineshape_bins, lineshape_step - ) - shape_tangent = shape_slope * -d_b0 - slope_tangent = shape_curve * -d_b0 - event_saturation = tl.load(saturation + event) - power_value = event_saturation * alpha_value * alpha_value - power_tangent = event_saturation * 2.0 * alpha_value * alpha_tangent - absorbed_value = tl.exp(power_value * shape_value) - absorbed_tangent = absorbed_value * ( - power_tangent * shape_value + power_value * shape_tangent - ) - if pools == 1: - held_bar = (bbvr, bbvi, bbtr, bbti) - held_state = (rbvr, rbvi, rbtr, rbti) - else: - held_bar = (cbvr, cbvi, cbtr, cbti) - held_state = (rcvr, rcvi, rctr, rcti) - per_state_v, per_state_t = _dual_real_conj_mul(*held_bar, *held_state) - grad_absorbed_v = tl.sum(per_state_v, axis=1)[:, None] - grad_absorbed_t = tl.sum(per_state_t, axis=1)[:, None] - grad_exponent_v = grad_absorbed_v * absorbed_value - grad_exponent_t = ( - grad_absorbed_t * absorbed_value + grad_absorbed_v * absorbed_tangent - ) - twice = event_saturation * 2.0 - sat_alpha_v = grad_exponent_v * (twice * alpha_value * shape_value) - sat_alpha_t = grad_exponent_t * ( - twice * alpha_value * shape_value - ) + grad_exponent_v * twice * ( - alpha_tangent * shape_value + alpha_value * shape_tangent - ) - # The lineshape is read at the pulse's offset from the voxel, so a - # step in the voxel's own off-resonance moves the read the other - # way. - sat_b0_v = -grad_exponent_v * (power_value * shape_slope) - sat_b0_t = -( - grad_exponent_t * (power_value * shape_slope) - + grad_exponent_v - * (power_tangent * shape_slope + power_value * slope_tangent) - ) - damped = _dual_scale(absorbed_value, absorbed_tangent, *held_bar) - saturating = is_rf & ~is_inversion - if pools == 1: - bbvr = tl.where(saturating, damped[0], bbvr) - bbvi = tl.where(saturating, damped[1], bbvi) - bbtr = tl.where(saturating, damped[2], bbtr) - bbti = tl.where(saturating, damped[3], bbti) - else: - cbvr = tl.where(saturating, damped[0], cbvr) - cbvi = tl.where(saturating, damped[1], cbvi) - cbtr = tl.where(saturating, damped[2], cbtr) - cbti = tl.where(saturating, damped[3], cbti) - cos_value = tl.cos(alpha_value) - sin_value = tl.sin(alpha_value) - cos_tangent = -sin_value * alpha_tangent - sin_tangent = cos_value * alpha_tangent - p1r, p1i, p1tr, p1ti = _dual_polar(phi_value, phi_tangent) - p2r, p2i, p2tr, p2ti = _dual_mul(p1r, p1i, p1tr, p1ti, p1r, p1i, p1tr, p1ti) - t00, t01, t02, t12, t20, t21, t22 = _rotation_block( - 0.5 * (1.0 + cos_value), - 0.5 * cos_tangent, - 0.5 * (1.0 - cos_value), - -0.5 * cos_tangent, - sin_value, - sin_tangent, - cos_value, - cos_tangent, - p1r, - p1i, - p1tr, - p1ti, - p2r, - p2i, - p2tr, - p2ti, - p1r, - -p1i, - p1tr, - -p1ti, - ) - # The flip angle reaches a shaped pulse's rotation through the slope - # stored beside it, or not at all when the rotation is read per voxel, - # so the operator's derivative in the flip is only built where the - # pulse is a flip and a phase. - alpha_v = empty - alpha_t = empty - phi_v = empty - phi_t = empty - alpha_b_v = empty - alpha_b_t = empty - phi_b_v = empty - phi_b_t = empty - if not profiled and not dynamic: - d00, d01, d02, d12, d20, d21, d22 = _rotation_block( - -0.5 * sin_value, - -0.5 * sin_tangent, - 0.5 * sin_value, - 0.5 * sin_tangent, - cos_value, - cos_tangent, - -sin_value, - -sin_tangent, - p1r, - p1i, - p1tr, - p1ti, - p2r, - p2i, - p2tr, - p2ti, - p1r, - -p1i, - p1tr, - -p1ti, - ) - - # d/dalpha, contracted with the adjoint. - row0 = _dual_mul(d00[0], d00[1], d00[2], d00[3], spvr, spvi, sptr, spti) - add1 = _dual_mul(d01[0], d01[1], d01[2], d01[3], smvr, smvi, smtr, smti) - add2 = _dual_mul(d02[0], d02[1], d02[2], d02[3], rzvr, rzvi, rztr, rzti) - alpha_v, alpha_t = _dual_real_conj_mul( - pbvr, - pbvi, - pbtr, - pbti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - row0 = _dual_mul(d01[0], -d01[1], d01[2], -d01[3], spvr, spvi, sptr, spti) - add1 = _dual_mul(d00[0], d00[1], d00[2], d00[3], smvr, smvi, smtr, smti) - add2 = _dual_mul(d12[0], d12[1], d12[2], d12[3], rzvr, rzvi, rztr, rzti) - part_v, part_t = _dual_real_conj_mul( - mbvr, - mbvi, - mbtr, - mbti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - alpha_v += part_v - alpha_t += part_t - row0 = _dual_mul(d20[0], d20[1], d20[2], d20[3], spvr, spvi, sptr, spti) - add1 = _dual_mul(d21[0], d21[1], d21[2], d21[3], smvr, smvi, smtr, smti) - add2 = _dual_mul(d22[0], d22[1], d22[2], d22[3], rzvr, rzvi, rztr, rzti) - part_v, part_t = _dual_real_conj_mul( - zbvr, - zbvi, - zbtr, - zbti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - alpha_v += part_v - alpha_t += part_t - - # d/dphi, where only the phase factors carry the dependence. - u1 = _dual_mul(t01[0], t01[1], t01[2], t01[3], smvr, smvi, smtr, smti) - u2 = _dual_mul(t02[0], t02[1], t02[2], t02[3], rzvr, rzvi, rztr, rzti) - ur, ui, utr, uti = _dual_times_i( - 2.0 * u1[0] + u2[0], - 2.0 * u1[1] + u2[1], - 2.0 * u1[2] + u2[2], - 2.0 * u1[3] + u2[3], - ) - phi_v, phi_t = _dual_real_conj_mul(pbvr, pbvi, pbtr, pbti, ur, ui, utr, uti) - u1 = _dual_mul(t01[0], -t01[1], t01[2], -t01[3], spvr, spvi, sptr, spti) - u2 = _dual_mul(t12[0], t12[1], t12[2], t12[3], rzvr, rzvi, rztr, rzti) - ur, ui, utr, uti = _dual_times_i( - -2.0 * u1[0] - u2[0], - -2.0 * u1[1] - u2[1], - -2.0 * u1[2] - u2[2], - -2.0 * u1[3] - u2[3], - ) - part_v, part_t = _dual_real_conj_mul( - mbvr, mbvi, mbtr, mbti, ur, ui, utr, uti - ) - phi_v += part_v - phi_t += part_t - u1 = _dual_mul(t20[0], t20[1], t20[2], t20[3], spvr, spvi, sptr, spti) - u2 = _dual_mul(t21[0], t21[1], t21[2], t21[3], smvr, smvi, smtr, smti) - ur, ui, utr, uti = _dual_times_i( - u2[0] - u1[0], u2[1] - u1[1], u2[2] - u1[2], u2[3] - u1[3] - ) - part_v, part_t = _dual_real_conj_mul( - zbvr, zbvi, zbtr, zbti, ur, ui, utr, uti - ) - phi_v += part_v - phi_t += part_t - - # The same pulse turns the exchanging pool, so its cotangent adds to - # the flip and phase the free pool already left. - alpha_b_v = 0.0 * alpha_v - alpha_b_t = 0.0 * alpha_v - phi_b_v = 0.0 * alpha_v - phi_b_t = 0.0 * alpha_v - if pools == 2 or pools == 3: - row0 = _dual_mul( - d00[0], d00[1], d00[2], d00[3], sbpvr, sbpvi, sbptr, sbpti - ) - add1 = _dual_mul( - d01[0], d01[1], d01[2], d01[3], sbmvr, sbmvi, sbmtr, sbmti - ) - add2 = _dual_mul(d02[0], d02[1], d02[2], d02[3], rbvr, rbvi, rbtr, rbti) - alpha_b_v, alpha_b_t = _dual_real_conj_mul( - ubvr, - ubvi, - ubtr, - ubti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - row0 = _dual_mul( - d01[0], -d01[1], d01[2], -d01[3], sbpvr, sbpvi, sbptr, sbpti - ) - add1 = _dual_mul( - d00[0], d00[1], d00[2], d00[3], sbmvr, sbmvi, sbmtr, sbmti - ) - add2 = _dual_mul(d12[0], d12[1], d12[2], d12[3], rbvr, rbvi, rbtr, rbti) - part_v, part_t = _dual_real_conj_mul( - wbvr, - wbvi, - wbtr, - wbti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - alpha_b_v += part_v - alpha_b_t += part_t - row0 = _dual_mul( - d20[0], d20[1], d20[2], d20[3], sbpvr, sbpvi, sbptr, sbpti - ) - add1 = _dual_mul( - d21[0], d21[1], d21[2], d21[3], sbmvr, sbmvi, sbmtr, sbmti - ) - add2 = _dual_mul(d22[0], d22[1], d22[2], d22[3], rbvr, rbvi, rbtr, rbti) - part_v, part_t = _dual_real_conj_mul( - bbvr, - bbvi, - bbtr, - bbti, - row0[0] + add1[0] + add2[0], - row0[1] + add1[1] + add2[1], - row0[2] + add1[2] + add2[2], - row0[3] + add1[3] + add2[3], - ) - alpha_b_v += part_v - alpha_b_t += part_t - - u1 = _dual_mul( - t01[0], t01[1], t01[2], t01[3], sbmvr, sbmvi, sbmtr, sbmti - ) - u2 = _dual_mul(t02[0], t02[1], t02[2], t02[3], rbvr, rbvi, rbtr, rbti) - ur, ui, utr, uti = _dual_times_i( - 2.0 * u1[0] + u2[0], - 2.0 * u1[1] + u2[1], - 2.0 * u1[2] + u2[2], - 2.0 * u1[3] + u2[3], - ) - phi_b_v, phi_b_t = _dual_real_conj_mul( - ubvr, ubvi, ubtr, ubti, ur, ui, utr, uti - ) - u1 = _dual_mul( - t01[0], -t01[1], t01[2], -t01[3], sbpvr, sbpvi, sbptr, sbpti - ) - u2 = _dual_mul(t12[0], t12[1], t12[2], t12[3], rbvr, rbvi, rbtr, rbti) - ur, ui, utr, uti = _dual_times_i( - -2.0 * u1[0] - u2[0], - -2.0 * u1[1] - u2[1], - -2.0 * u1[2] - u2[2], - -2.0 * u1[3] - u2[3], - ) - part_v, part_t = _dual_real_conj_mul( - wbvr, wbvi, wbtr, wbti, ur, ui, utr, uti - ) - phi_b_v += part_v - phi_b_t += part_t - u1 = _dual_mul( - t20[0], t20[1], t20[2], t20[3], sbpvr, sbpvi, sbptr, sbpti - ) - u2 = _dual_mul( - t21[0], t21[1], t21[2], t21[3], sbmvr, sbmvi, sbmtr, sbmti - ) - ur, ui, utr, uti = _dual_times_i( - u2[0] - u1[0], u2[1] - u1[1], u2[2] - u1[2], u2[3] - u1[3] - ) - part_v, part_t = _dual_real_conj_mul( - bbvr, bbvi, bbtr, bbti, ur, ui, utr, uti - ) - phi_b_v += part_v - phi_b_t += part_t - - if profiled or dynamic: - shaped_slope_a = (empty, empty, empty, empty) - shaped_slope_b = (empty, empty, empty, empty) - if dynamic: - shaped_a, shaped_b = _dynamic_pair_dual_at( - pairs, - pair_direction, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - phi_value, - phi_tangent, - directed, - ) - else: - shaped_a, shaped_b, shaped_slope_a, shaped_slope_b = ( - _profiled_pair_dual( - profile, - _table_row(profile_index, event, location, locations), - alpha_value, - alpha_tangent, - phi_value, - phi_tangent, - profile_bins, - profile_step, - ) - ) - grad_a, grad_b, shaped_pb, shaped_mb, shaped_zb = _spinor_adjoint_dual( - shaped_a, - shaped_b, - (spvr, spvi, sptr, spti), - (smvr, smvi, smtr, smti), - (rzvr, rzvi, rztr, rzti), - (pbvr, pbvi, pbtr, pbti), - (mbvr, mbvi, mbtr, mbti), - (zbvr, zbvi, zbtr, zbti), - ) - if dynamic: - # The flip is inside the pair rather than read against it, so - # it has no gradient here: the cotangent goes out on the - # rotation and whatever integrated it carries the rest. ``b`` - # was turned by the phase after the pair came out, so the - # cotangent turns back the other way. - alpha_v = empty - alpha_t = empty - back = _dual_product( - grad_b, _dual_conj(_dual_polar(-phi_value, -phi_tangent)) - ) - _store_pair_cotangent( - grad_pair_value, - grad_pair_tangent, - pair_index, - event_base, - event, - atom, - atom_count, - is_rf & ~is_inversion, - active_atom, - state_mask, - grad_a, - back, - ) - else: - alpha_v, alpha_t = _dual_real_conj_mul( - grad_a[0], - grad_a[1], - grad_a[2], - grad_a[3], - shaped_slope_a[0], - shaped_slope_a[1], - shaped_slope_a[2], - shaped_slope_a[3], - ) - part_v, part_t = _dual_real_conj_mul( - grad_b[0], - grad_b[1], - grad_b[2], - grad_b[3], - shaped_slope_b[0], - shaped_slope_b[1], - shaped_slope_b[2], - shaped_slope_b[3], - ) - alpha_v += part_v - alpha_t += part_t - # d(b e^{-i phi})/dphi is -i times it, and nothing else moves. - turn_r, turn_i, turn_tr, turn_ti = _dual_times_i( - shaped_b[0], shaped_b[1], shaped_b[2], shaped_b[3] - ) - phi_v, phi_t = _dual_real_conj_mul( - grad_b[0], - grad_b[1], - grad_b[2], - grad_b[3], - -turn_r, - -turn_i, - -turn_tr, - -turn_ti, - ) - if pools == 2 or pools == 3: - pool_a, pool_b_pair, shaped_ub, shaped_wb, shaped_bb = ( - _spinor_adjoint_dual( - shaped_a, - shaped_b, - (sbpvr, sbpvi, sbptr, sbpti), - (sbmvr, sbmvi, sbmtr, sbmti), - (rbvr, rbvi, rbtr, rbti), - (ubvr, ubvi, ubtr, ubti), - (wbvr, wbvi, wbtr, wbti), - (bbvr, bbvi, bbtr, bbti), - ) - ) - if dynamic: - # The same pulse turned this pool, so its cotangent lands - # on the same row. - pool_back = _dual_product( - pool_b_pair, - _dual_conj(_dual_polar(-phi_value, -phi_tangent)), - ) - _store_pair_cotangent( - grad_pair_value, - grad_pair_tangent, - pair_index, - event_base, - event, - atom, - atom_count, - is_rf & ~is_inversion, - active_atom, - state_mask, - pool_a, - pool_back, - ) - else: - alpha_b_v, alpha_b_t = _dual_real_conj_mul( - pool_a[0], - pool_a[1], - pool_a[2], - pool_a[3], - shaped_slope_a[0], - shaped_slope_a[1], - shaped_slope_a[2], - shaped_slope_a[3], - ) - part_v, part_t = _dual_real_conj_mul( - pool_b_pair[0], - pool_b_pair[1], - pool_b_pair[2], - pool_b_pair[3], - shaped_slope_b[0], - shaped_slope_b[1], - shaped_slope_b[2], - shaped_slope_b[3], - ) - alpha_b_v += part_v - alpha_b_t += part_t - phi_b_v, phi_b_t = _dual_real_conj_mul( - pool_b_pair[0], - pool_b_pair[1], - pool_b_pair[2], - pool_b_pair[3], - -turn_r, - -turn_i, - -turn_tr, - -turn_ti, - ) - alpha_v += alpha_b_v - alpha_t += alpha_b_t - phi_v += phi_b_v - phi_t += phi_b_t - - rotate = is_rf & ~is_inversion - grad_alpha_v = tl.sum(tl.where(rotate, alpha_v, 0.0), axis=1)[:, None] - grad_alpha_t = tl.sum(tl.where(rotate, alpha_t, 0.0), axis=1)[:, None] - grad_phi_v = tl.sum(tl.where(rotate, phi_v, 0.0), axis=1)[:, None] - grad_phi_t = tl.sum(tl.where(rotate, phi_t, 0.0), axis=1)[:, None] - if pools == 1 or pools == 3: - turning = tl.where(rotate, 1.0, 0.0) - grad_alpha_v += sat_alpha_v * turning - grad_alpha_t += sat_alpha_t * turning - g_b0v += sat_b0_v * turning - g_b0t += sat_b0_t * turning - - # Conjugate transpose of the rotation. - n0 = _dual_mul(t00[0], -t00[1], t00[2], -t00[3], pbvr, pbvi, pbtr, pbti) - n1 = _dual_mul(t01[0], t01[1], t01[2], t01[3], mbvr, mbvi, mbtr, mbti) - n2 = _dual_mul(t20[0], -t20[1], t20[2], -t20[3], zbvr, zbvi, zbtr, zbti) - q0 = _dual_mul(t01[0], -t01[1], t01[2], -t01[3], pbvr, pbvi, pbtr, pbti) - q1 = _dual_mul(t00[0], -t00[1], t00[2], -t00[3], mbvr, mbvi, mbtr, mbti) - q2 = _dual_mul(t21[0], -t21[1], t21[2], -t21[3], zbvr, zbvi, zbtr, zbti) - w0 = _dual_mul(t02[0], -t02[1], t02[2], -t02[3], pbvr, pbvi, pbtr, pbti) - w1 = _dual_mul(t12[0], -t12[1], t12[2], -t12[3], mbvr, mbvi, mbtr, mbti) - w2 = _dual_mul(t22[0], -t22[1], t22[2], -t22[3], zbvr, zbvi, zbtr, zbti) - - back_pb = ( - n0[0] + n1[0] + n2[0], - n0[1] + n1[1] + n2[1], - n0[2] + n1[2] + n2[2], - n0[3] + n1[3] + n2[3], - ) - back_mb = ( - q0[0] + q1[0] + q2[0], - q0[1] + q1[1] + q2[1], - q0[2] + q1[2] + q2[2], - q0[3] + q1[3] + q2[3], - ) - back_zb = ( - w0[0] + w1[0] + w2[0], - w0[1] + w1[1] + w2[1], - w0[2] + w1[2] + w2[2], - w0[3] + w1[3] + w2[3], - ) - if pools == 2 or pools == 3: - n0 = _dual_mul(t00[0], -t00[1], t00[2], -t00[3], ubvr, ubvi, ubtr, ubti) - n1 = _dual_mul(t01[0], t01[1], t01[2], t01[3], wbvr, wbvi, wbtr, wbti) - n2 = _dual_mul(t20[0], -t20[1], t20[2], -t20[3], bbvr, bbvi, bbtr, bbti) - q0 = _dual_mul(t01[0], -t01[1], t01[2], -t01[3], ubvr, ubvi, ubtr, ubti) - q1 = _dual_mul(t00[0], -t00[1], t00[2], -t00[3], wbvr, wbvi, wbtr, wbti) - q2 = _dual_mul(t21[0], -t21[1], t21[2], -t21[3], bbvr, bbvi, bbtr, bbti) - w0 = _dual_mul(t02[0], -t02[1], t02[2], -t02[3], ubvr, ubvi, ubtr, ubti) - w1 = _dual_mul(t12[0], -t12[1], t12[2], -t12[3], wbvr, wbvi, wbtr, wbti) - w2 = _dual_mul(t22[0], -t22[1], t22[2], -t22[3], bbvr, bbvi, bbtr, bbti) - back_ub = ( - n0[0] + n1[0] + n2[0], - n0[1] + n1[1] + n2[1], - n0[2] + n1[2] + n2[2], - n0[3] + n1[3] + n2[3], - ) - back_wb = ( - q0[0] + q1[0] + q2[0], - q0[1] + q1[1] + q2[1], - q0[2] + q1[2] + q2[2], - q0[3] + q1[3] + q2[3], - ) - back_bb = ( - w0[0] + w1[0] + w2[0], - w0[1] + w1[1] + w2[1], - w0[2] + w1[2] + w2[2], - w0[3] + w1[3] + w2[3], - ) - if profiled or dynamic: - back_ub = shaped_ub - back_wb = shaped_wb - back_bb = shaped_bb - ubvr = tl.where(rotate, back_ub[0], ubvr) - ubvi = tl.where(rotate, back_ub[1], ubvi) - ubtr = tl.where(rotate, back_ub[2], ubtr) - ubti = tl.where(rotate, back_ub[3], ubti) - wbvr = tl.where(rotate, back_wb[0], wbvr) - wbvi = tl.where(rotate, back_wb[1], wbvi) - wbtr = tl.where(rotate, back_wb[2], wbtr) - wbti = tl.where(rotate, back_wb[3], wbti) - bbvr = tl.where(rotate, back_bb[0], bbvr) - bbvi = tl.where(rotate, back_bb[1], bbvi) - bbtr = tl.where(rotate, back_bb[2], bbtr) - bbti = tl.where(rotate, back_bb[3], bbti) - if profiled or dynamic: - back_pb = shaped_pb - back_mb = shaped_mb - back_zb = shaped_zb - - pbvr = tl.where(rotate, back_pb[0], pbvr) - pbvi = tl.where(rotate, back_pb[1], pbvi) - pbtr = tl.where(rotate, back_pb[2], pbtr) - pbti = tl.where(rotate, back_pb[3], pbti) - mbvr = tl.where(rotate, back_mb[0], mbvr) - mbvi = tl.where(rotate, back_mb[1], mbvi) - mbtr = tl.where(rotate, back_mb[2], mbtr) - mbti = tl.where(rotate, back_mb[3], mbti) - zbvr = tl.where(rotate, back_zb[0], zbvr) - zbvi = tl.where(rotate, back_zb[1], zbvi) - zbtr = tl.where(rotate, back_zb[2], zbtr) - zbti = tl.where(rotate, back_zb[3], zbti) - - writes_flip = active_atom & rotate - tl.atomic_add( - grad_flip_value + event_base + event, - grad_alpha_v * atom_b1, - mask=writes_flip, - ) - tl.atomic_add( - grad_flip_tangent + event_base + event, - grad_alpha_t * atom_b1 + grad_alpha_v * d_b1, - mask=writes_flip, - ) - tl.atomic_add( - grad_phase_value + event_base + event, grad_phi_v, mask=writes_flip - ) - tl.atomic_add( - grad_phase_tangent + event_base + event, grad_phi_t, mask=writes_flip - ) - # A pulse's transmit gradient belongs to the shim it drives, so with - # several it lands in that shim's row here rather than in a register - # summed over the whole train. ``row`` is the offset of the row the - # replay above read. - if shimmed: - tl.atomic_add( - grad_tissue_value + _B1_ROW * atom_count + row + atom, - grad_alpha_v * event_flip, - mask=writes_flip, - ) - tl.atomic_add( - grad_tissue_tangent + _B1_ROW * atom_count + row + atom, - grad_alpha_t * event_flip + grad_alpha_v * event_dot_flip, - mask=writes_flip, - ) - tl.atomic_add( - grad_tissue_value - + (_B1_PHASE_ROW + shim_rows - 1) * atom_count - + row - + atom, - grad_phi_v, - mask=writes_flip, - ) - tl.atomic_add( - grad_tissue_tangent - + (_B1_PHASE_ROW + shim_rows - 1) * atom_count - + row - + atom, - grad_phi_t, - mask=writes_flip, - ) - else: - g_b1v += grad_alpha_v * event_flip - g_b1t += grad_alpha_t * event_flip + grad_alpha_v * event_dot_flip - g_b1pv += grad_phi_v - g_b1pt += grad_phi_t - - avr, avi, bvr, bvi = _shift_adjoint( - pbvr, pbvi, mbvr, mbvi, state, state_mask, state_count - ) - atr, ati, btr, bti = _shift_adjoint( - pbtr, pbti, mbtr, mbti, state, state_mask, state_count - ) - pbvr = tl.where(pre_shift, avr, pbvr) - pbvi = tl.where(pre_shift, avi, pbvi) - pbtr = tl.where(pre_shift, atr, pbtr) - pbti = tl.where(pre_shift, ati, pbti) - mbvr = tl.where(pre_shift, bvr, mbvr) - mbvi = tl.where(pre_shift, bvi, mbvi) - mbtr = tl.where(pre_shift, btr, mbtr) - mbti = tl.where(pre_shift, bti, mbti) - if pools == 2 or pools == 3: - avr, avi, bvr, bvi = _shift_adjoint( - ubvr, ubvi, wbvr, wbvi, state, state_mask, state_count - ) - atr, ati, btr, bti = _shift_adjoint( - ubtr, ubti, wbtr, wbti, state, state_mask, state_count - ) - ubvr = tl.where(pre_shift, avr, ubvr) - ubvi = tl.where(pre_shift, avi, ubvi) - ubtr = tl.where(pre_shift, atr, ubtr) - ubti = tl.where(pre_shift, ati, ubti) - wbvr = tl.where(pre_shift, bvr, wbvr) - wbvi = tl.where(pre_shift, bvi, wbvi) - wbtr = tl.where(pre_shift, btr, wbtr) - wbti = tl.where(pre_shift, bti, wbti) - - # ---- relaxation and off-resonance adjoint ---- - grad_e2_v = zero - grad_e2_t = zero - attenuation_v = zero - attenuation_t = zero - two_pool_dt_v = zero - two_pool_dt_t = zero - # The damping is homogeneous of degree one in every transverse state it - # acts on, so its gradient times the damping itself is the cotangent - # taken against the states the interval leaves. With one pool that is - # the same thing as the relaxation factor's own gradient, scaled. - if pools == 2 or pools == 3: - plus_side = _dual_add( - _dual_product( - _dual_conj((pbvr, pbvi, pbtr, pbti)), - (rpvr, rpvi, rptr, rpti), - ), - _dual_product( - _dual_conj((ubvr, ubvi, ubtr, ubti)), - (rbpvr, rbpvi, rbptr, rbpti), - ), - ) - minus_side = _dual_add( - _dual_product( - _dual_conj((mbvr, mbvi, mbtr, mbti)), - (rmvr, rmvi, rmtr, rmti), - ), - _dual_product( - _dual_conj((wbvr, wbvi, wbtr, wbti)), - (rbmvr, rbmvi, rbmtr, rbmti), - ), - ) - damped = _dual_add(plus_side, minus_side) - wound = _dual_times_i(*_dual_subtract(plus_side, minus_side)) - cot2_v = damped[0] - cot2_t = damped[2] - per_angle_v = wound[0] - per_angle_t = wound[2] - else: - pq = _dual_mul(qr, qi, qtr, qti, xpvr, xpvi, xptr, xpti) - mq = _dual_mul(qr, -qi, qtr, -qti, xmvr, xmvi, xmtr, xmti) - e2_v, e2_t = _dual_real_conj_mul( - pbvr, pbvi, pbtr, pbti, pq[0], pq[1], pq[2], pq[3] - ) - part_v, part_t = _dual_real_conj_mul( - mbvr, mbvi, mbtr, mbti, mq[0], mq[1], mq[2], mq[3] - ) - bare_cot_v = e2_v + part_v - bare_cot_t = e2_t + part_t - grad_e2_v = tl.sum(bare_cot_v * damp_t, axis=1)[:, None] - grad_e2_t = tl.sum( - bare_cot_v * damp_t_tangent + bare_cot_t * damp_t, axis=1 - )[:, None] - - per_angle_v = empty - per_angle_t = empty - if off_axis or moving: - po = _dual_mul(ovr, ovi, otr, oti, xpvr, xpvi, xptr, xpti) - po = _dual_times_i(po[0], po[1], po[2], po[3]) - mo = _dual_mul(ovr, -ovi, otr, -oti, xmvr, xmvi, xmtr, xmti) - mo = _dual_times_i(mo[0], mo[1], mo[2], mo[3]) - angle_v, angle_t = _dual_real_conj_mul( - pbvr, pbvi, pbtr, pbti, po[0], po[1], po[2], po[3] - ) - part_v, part_t = _dual_real_conj_mul( - mbvr, mbvi, mbtr, mbti, mo[0], mo[1], mo[2], mo[3] - ) - per_angle_v = angle_v - part_v - per_angle_t = angle_t - part_t - cot2_v = bare_cot_v * bare2_value * damp_t - cot2_t = ( - bare_cot_t * bare2_value * damp_t - + bare_cot_v * bare2_tangent * damp_t - + bare_cot_v * bare2_value * damp_t_tangent - ) - # A turn of the transverse states and the off-resonance angle are the - # same derivative; only the weight each order carries differs. - grad_angle_v = zero - grad_angle_t = zero - if off_axis or moving: - grad_angle_v = tl.sum(per_angle_v, axis=1)[:, None] - grad_angle_t = tl.sum(per_angle_t, axis=1)[:, None] - - grad_e1_v = zero - grad_e1_t = zero - if pools == 3: - # The nine entries of the mixing operator and the three recoveries, - # summed over the orders that share them, then pushed back through - # the closed form once for the whole interval and in double. - free_bar = (zbvr, zbvi, zbtr, zbti) - bound_bar = (bbvr, bbvi, bbtr, bbti) - semi_bar = (cbvr, cbvi, cbtr, cbti) - spun_free = _dual_mul( - spin[0], spin[1], spin[2], spin[3], xzvr, xzvi, xztr, xzti - ) - spun_bound = _dual_mul( - spin[0], spin[1], spin[2], spin[3], xbvr, xbvi, xbtr, xbti - ) - spun_semi = _dual_mul( - spin[0], spin[1], spin[2], spin[3], xcvr, xcvi, xctr, xcti - ) - e11_v, e11_t = _dual_real_conj_mul(*free_bar, *spun_free) - e12_v, e12_t = _dual_real_conj_mul(*free_bar, *spun_bound) - e13_v, e13_t = _dual_real_conj_mul(*free_bar, *spun_semi) - e21_v, e21_t = _dual_real_conj_mul(*bound_bar, *spun_free) - e22_v, e22_t = _dual_real_conj_mul(*bound_bar, *spun_bound) - e23_v, e23_t = _dual_real_conj_mul(*bound_bar, *spun_semi) - e31_v, e31_t = _dual_real_conj_mul(*semi_bar, *spun_free) - e32_v, e32_t = _dual_real_conj_mul(*semi_bar, *spun_bound) - e33_v, e33_t = _dual_real_conj_mul(*semi_bar, *spun_semi) - if tabulated: - # Every gradient but the interval's own and the - # attenuation's is linear in these cotangents, so the - # events sharing a length pool them here and pay the - # closed form once each after the walk back. The tangent - # gradient carries a third term in the event's own - # interval direction, which pools as the value cotangents - # weighted by it. - bar11_v = tl.sum(e11_v, axis=1)[:, None] - bar12_v = tl.sum(e12_v, axis=1)[:, None] - bar13_v = tl.sum(e13_v, axis=1)[:, None] - bar21_v = tl.sum(e21_v, axis=1)[:, None] - bar22_v = tl.sum(e22_v, axis=1)[:, None] - bar23_v = tl.sum(e23_v, axis=1)[:, None] - bar31_v = tl.sum(e31_v, axis=1)[:, None] - bar32_v = tl.sum(e32_v, axis=1)[:, None] - bar33_v = tl.sum(e33_v, axis=1)[:, None] - barfree_v = tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None] - barpool_v = tl.sum(tl.where(state == 0, bbvr, 0.0), axis=1)[:, None] - barbound_v = tl.sum(tl.where(state == 0, cbvr, 0.0), axis=1)[:, None] - bar11_t = tl.sum(e11_t, axis=1)[:, None] - bar12_t = tl.sum(e12_t, axis=1)[:, None] - bar13_t = tl.sum(e13_t, axis=1)[:, None] - bar21_t = tl.sum(e21_t, axis=1)[:, None] - bar22_t = tl.sum(e22_t, axis=1)[:, None] - bar23_t = tl.sum(e23_t, axis=1)[:, None] - bar31_t = tl.sum(e31_t, axis=1)[:, None] - bar32_t = tl.sum(e32_t, axis=1)[:, None] - bar33_t = tl.sum(e33_t, axis=1)[:, None] - barfree_t = tl.sum(tl.where(state == 0, zbtr, 0.0), axis=1)[:, None] - barpool_t = tl.sum(tl.where(state == 0, bbtr, 0.0), axis=1)[:, None] - barbound_t = tl.sum(tl.where(state == 0, cbtr, 0.0), axis=1)[:, None] - held = pool_bars + (local * row_count + pool_row) * 36 - tl.store( - held + 0, - tl.load(held + 0, mask=active_atom, other=0.0) + bar11_v, - mask=active_atom, - ) - tl.store( - held + 1, - tl.load(held + 1, mask=active_atom, other=0.0) + bar12_v, - mask=active_atom, - ) - tl.store( - held + 2, - tl.load(held + 2, mask=active_atom, other=0.0) + bar13_v, - mask=active_atom, - ) - tl.store( - held + 3, - tl.load(held + 3, mask=active_atom, other=0.0) + bar21_v, - mask=active_atom, - ) - tl.store( - held + 4, - tl.load(held + 4, mask=active_atom, other=0.0) + bar22_v, - mask=active_atom, - ) - tl.store( - held + 5, - tl.load(held + 5, mask=active_atom, other=0.0) + bar23_v, - mask=active_atom, - ) - tl.store( - held + 6, - tl.load(held + 6, mask=active_atom, other=0.0) + bar31_v, - mask=active_atom, - ) - tl.store( - held + 7, - tl.load(held + 7, mask=active_atom, other=0.0) + bar32_v, - mask=active_atom, - ) - tl.store( - held + 8, - tl.load(held + 8, mask=active_atom, other=0.0) + bar33_v, - mask=active_atom, - ) - tl.store( - held + 9, - tl.load(held + 9, mask=active_atom, other=0.0) + barfree_v, - mask=active_atom, - ) - tl.store( - held + 10, - tl.load(held + 10, mask=active_atom, other=0.0) + barpool_v, - mask=active_atom, - ) - tl.store( - held + 11, - tl.load(held + 11, mask=active_atom, other=0.0) + barbound_v, - mask=active_atom, - ) - tl.store( - held + 12, - tl.load(held + 12, mask=active_atom, other=0.0) + bar11_t, - mask=active_atom, - ) - tl.store( - held + 13, - tl.load(held + 13, mask=active_atom, other=0.0) + bar12_t, - mask=active_atom, - ) - tl.store( - held + 14, - tl.load(held + 14, mask=active_atom, other=0.0) + bar13_t, - mask=active_atom, - ) - tl.store( - held + 15, - tl.load(held + 15, mask=active_atom, other=0.0) + bar21_t, - mask=active_atom, - ) - tl.store( - held + 16, - tl.load(held + 16, mask=active_atom, other=0.0) + bar22_t, - mask=active_atom, - ) - tl.store( - held + 17, - tl.load(held + 17, mask=active_atom, other=0.0) + bar23_t, - mask=active_atom, - ) - tl.store( - held + 18, - tl.load(held + 18, mask=active_atom, other=0.0) + bar31_t, - mask=active_atom, - ) - tl.store( - held + 19, - tl.load(held + 19, mask=active_atom, other=0.0) + bar32_t, - mask=active_atom, - ) - tl.store( - held + 20, - tl.load(held + 20, mask=active_atom, other=0.0) + bar33_t, - mask=active_atom, - ) - tl.store( - held + 21, - tl.load(held + 21, mask=active_atom, other=0.0) + barfree_t, - mask=active_atom, - ) - tl.store( - held + 22, - tl.load(held + 22, mask=active_atom, other=0.0) + barpool_t, - mask=active_atom, - ) - tl.store( - held + 23, - tl.load(held + 23, mask=active_atom, other=0.0) + barbound_t, - mask=active_atom, - ) - tl.store( - held + 24, - tl.load(held + 24, mask=active_atom, other=0.0) - + dt_tangent * bar11_v, - mask=active_atom, - ) - tl.store( - held + 25, - tl.load(held + 25, mask=active_atom, other=0.0) - + dt_tangent * bar12_v, - mask=active_atom, - ) - tl.store( - held + 26, - tl.load(held + 26, mask=active_atom, other=0.0) - + dt_tangent * bar13_v, - mask=active_atom, - ) - tl.store( - held + 27, - tl.load(held + 27, mask=active_atom, other=0.0) - + dt_tangent * bar21_v, - mask=active_atom, - ) - tl.store( - held + 28, - tl.load(held + 28, mask=active_atom, other=0.0) - + dt_tangent * bar22_v, - mask=active_atom, - ) - tl.store( - held + 29, - tl.load(held + 29, mask=active_atom, other=0.0) - + dt_tangent * bar23_v, - mask=active_atom, - ) - tl.store( - held + 30, - tl.load(held + 30, mask=active_atom, other=0.0) - + dt_tangent * bar31_v, - mask=active_atom, - ) - tl.store( - held + 31, - tl.load(held + 31, mask=active_atom, other=0.0) - + dt_tangent * bar32_v, - mask=active_atom, - ) - tl.store( - held + 32, - tl.load(held + 32, mask=active_atom, other=0.0) - + dt_tangent * bar33_v, - mask=active_atom, - ) - tl.store( - held + 33, - tl.load(held + 33, mask=active_atom, other=0.0) - + dt_tangent * barfree_v, - mask=active_atom, - ) - tl.store( - held + 34, - tl.load(held + 34, mask=active_atom, other=0.0) - + dt_tangent * barpool_v, - mask=active_atom, - ) - tl.store( - held + 35, - tl.load(held + 35, mask=active_atom, other=0.0) - + dt_tangent * barbound_v, - mask=active_atom, - ) - ( - back_dt_v, - back_att_v, - back_dt_t, - back_att_t, - ) = _three_pool_interval_adjoint_jvp( - pool_table, - pool_row, - atom, - atom_count, - active_atom, - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_tangent, - wout_value, - wout_tangent, - bar11_v, - bar12_v, - bar13_v, - bar21_v, - bar22_v, - bar23_v, - bar31_v, - bar32_v, - bar33_v, - barfree_v, - barpool_v, - barbound_v, - bar11_t, - bar12_t, - bar13_t, - bar21_t, - bar22_t, - bar23_t, - bar31_t, - bar32_t, - bar33_t, - barfree_t, - barpool_t, - barbound_t, - ) - attenuation_v = back_att_v - attenuation_t = back_att_t - two_pool_dt_v = back_dt_v - two_pool_dt_t = back_dt_t - else: - ( - back_r1_v, - back_r1b_v, - back_r1c_v, - back_exch_v, - back_sexch_v, - back_bound_v, - back_semi_v, - back_dt_v, - back_att_v, - back_r1_t, - back_r1b_t, - back_r1c_t, - back_exch_t, - back_sexch_t, - back_bound_t, - back_semi_t, - back_dt_t, - back_att_t, - ) = _three_pool_step_adjoint_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - tl.sum(e11_v, axis=1)[:, None], - tl.sum(e11_t, axis=1)[:, None], - tl.sum(e12_v, axis=1)[:, None], - tl.sum(e12_t, axis=1)[:, None], - tl.sum(e13_v, axis=1)[:, None], - tl.sum(e13_t, axis=1)[:, None], - tl.sum(e21_v, axis=1)[:, None], - tl.sum(e21_t, axis=1)[:, None], - tl.sum(e22_v, axis=1)[:, None], - tl.sum(e22_t, axis=1)[:, None], - tl.sum(e23_v, axis=1)[:, None], - tl.sum(e23_t, axis=1)[:, None], - tl.sum(e31_v, axis=1)[:, None], - tl.sum(e31_t, axis=1)[:, None], - tl.sum(e32_v, axis=1)[:, None], - tl.sum(e32_t, axis=1)[:, None], - tl.sum(e33_v, axis=1)[:, None], - tl.sum(e33_t, axis=1)[:, None], - tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None], - tl.sum(tl.where(state == 0, zbtr, 0.0), axis=1)[:, None], - tl.sum(tl.where(state == 0, bbvr, 0.0), axis=1)[:, None], - tl.sum(tl.where(state == 0, bbtr, 0.0), axis=1)[:, None], - tl.sum(tl.where(state == 0, cbvr, 0.0), axis=1)[:, None], - tl.sum(tl.where(state == 0, cbtr, 0.0), axis=1)[:, None], - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - narrow, - ) - slope1_v = -1000.0 / (atom_t1 * atom_t1) - slope1_t = 2000.0 * d_t1 / (atom_t1 * atom_t1 * atom_t1) - slope1b_v = -1000.0 / (atom_t1b * atom_t1b) - slope1b_t = 2000.0 * d_t1b / (atom_t1b * atom_t1b * atom_t1b) - slope1c_v = -1000.0 / (held_semisolid * held_semisolid) - slope1c_t = ( - 2000.0 - * d_semisolid_t1 - / (held_semisolid * held_semisolid * held_semisolid) - ) - g_t1v += back_r1_v * slope1_v - g_t1t += back_r1_t * slope1_v + back_r1_v * slope1_t - g_t1bv += back_r1b_v * slope1b_v - g_t1bt += back_r1b_t * slope1b_v + back_r1b_v * slope1b_t - g_t1cv += back_r1c_v * slope1c_v - g_t1ct += back_r1c_t * slope1c_v + back_r1c_v * slope1c_t - g_exchv += back_exch_v - g_excht += back_exch_t - g_sexchv += back_sexch_v - g_sexcht += back_sexch_t - g_boundv += back_bound_v - g_boundt += back_bound_t - g_semiv += back_semi_v - g_semit += back_semi_t - attenuation_v = back_att_v - attenuation_t = back_att_t - two_pool_dt_v = back_dt_v - two_pool_dt_t = back_dt_t - turned_free = _dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_free) - turned_bound = _dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_bound) - turned_semi = _dual_mul( - spin[0], spin[1], spin[2], spin[3], *mixed_semisolid - ) - damp_pair_v, damp_pair_t = _dual_real_conj_mul(*free_bar, *turned_free) - other_v, other_t = _dual_real_conj_mul(*bound_bar, *turned_bound) - stuck_v, stuck_t = _dual_real_conj_mul(*semi_bar, *turned_semi) - long_damp_v = damp_pair_v + other_v + stuck_v - long_damp_t = damp_pair_t + other_t + stuck_t - zangle_v, zangle_t = _dual_real_conj_mul( - *free_bar, *_dual_times_i(*turned_free) - ) - part_v, part_t = _dual_real_conj_mul( - *bound_bar, *_dual_times_i(*turned_bound) - ) - zangle_v += part_v - zangle_t += part_t - part_v, part_t = _dual_real_conj_mul( - *semi_bar, *_dual_times_i(*turned_semi) - ) - zangle_v += part_v - zangle_t += part_t - col_free = _dual_add( - _dual_add( - _dual_back(w11, d_w11, *spin, *free_bar), - _dual_back(w21, d_w21, *spin, *bound_bar), - ), - _dual_back(w31, d_w31, *spin, *semi_bar), - ) - col_bound = _dual_add( - _dual_add( - _dual_back(w12, d_w12, *spin, *free_bar), - _dual_back(w22, d_w22, *spin, *bound_bar), - ), - _dual_back(w32, d_w32, *spin, *semi_bar), - ) - col_semi = _dual_add( - _dual_add( - _dual_back(w13, d_w13, *spin, *free_bar), - _dual_back(w23, d_w23, *spin, *bound_bar), - ), - _dual_back(w33, d_w33, *spin, *semi_bar), - ) - zbvr, zbvi, zbtr, zbti = col_free - bbvr, bbvi, bbtr, bbti = col_bound - cbvr, cbvi, cbtr, cbti = col_semi - elif pools > 0: - # The four entries of the exchange operator and the two recoveries, - # summed over the orders that share them, then pushed back through - # the closed form once for the whole interval. - free_bar = (zbvr, zbvi, zbtr, zbti) - bound_bar = (bbvr, bbvi, bbtr, bbti) - spun_free = _dual_mul( - spin[0], spin[1], spin[2], spin[3], xzvr, xzvi, xztr, xzti - ) - spun_bound = _dual_mul( - spin[0], spin[1], spin[2], spin[3], xbvr, xbvi, xbtr, xbti - ) - e11_v, e11_t = _dual_real_conj_mul(*free_bar, *spun_free) - e12_v, e12_t = _dual_real_conj_mul(*free_bar, *spun_bound) - e21_v, e21_t = _dual_real_conj_mul(*bound_bar, *spun_free) - e22_v, e22_t = _dual_real_conj_mul(*bound_bar, *spun_bound) - bar_e11_v = tl.sum(e11_v, axis=1)[:, None] - bar_e11_t = tl.sum(e11_t, axis=1)[:, None] - bar_e12_v = tl.sum(e12_v, axis=1)[:, None] - bar_e12_t = tl.sum(e12_t, axis=1)[:, None] - bar_e21_v = tl.sum(e21_v, axis=1)[:, None] - bar_e21_t = tl.sum(e21_t, axis=1)[:, None] - bar_e22_v = tl.sum(e22_v, axis=1)[:, None] - bar_e22_t = tl.sum(e22_t, axis=1)[:, None] - rec_f_v = tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None] - rec_f_t = tl.sum(tl.where(state == 0, zbtr, 0.0), axis=1)[:, None] - rec_b_v = tl.sum(tl.where(state == 0, bbvr, 0.0), axis=1)[:, None] - rec_b_t = tl.sum(tl.where(state == 0, bbtr, 0.0), axis=1)[:, None] - ( - back_r1_v, - back_r1b_v, - back_exch_v, - back_bound_v, - back_dt_v, - back_att_v, - back_r1_t, - back_r1b_t, - back_exch_t, - back_bound_t, - back_dt_t, - back_att_t, - ) = _two_pool_step_adjoint_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - bar_e11_v, - bar_e11_t, - bar_e12_v, - bar_e12_t, - bar_e21_v, - bar_e21_t, - bar_e22_v, - bar_e22_t, - rec_f_v, - rec_f_t, - rec_b_v, - rec_b_t, - ) - # r1 = 1000/t1, so a rate gradient reaches the time through the - # square of it. - slope1_v = -1000.0 / (atom_t1 * atom_t1) - slope1_t = 2000.0 * d_t1 / (atom_t1 * atom_t1 * atom_t1) - slope1b_v = -1000.0 / (atom_t1b * atom_t1b) - slope1b_t = 2000.0 * d_t1b / (atom_t1b * atom_t1b * atom_t1b) - g_t1v += back_r1_v * slope1_v - g_t1t += back_r1_t * slope1_v + back_r1_v * slope1_t - g_t1bv += back_r1b_v * slope1b_v - g_t1bt += back_r1b_t * slope1b_v + back_r1b_v * slope1b_t - g_exchv += back_exch_v - g_excht += back_exch_t - g_boundv += back_bound_v - g_boundt += back_bound_t - attenuation_v = back_att_v - attenuation_t = back_att_t - two_pool_dt_v = back_dt_v - two_pool_dt_t = back_dt_t - # Both pools take the same per-order damping and turn, so each - # collects the cotangent of the mixture that reached it. - damp_pair_v, damp_pair_t = _dual_real_conj_mul( - *free_bar, *_dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_free) - ) - other_v, other_t = _dual_real_conj_mul( - *bound_bar, *_dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_bound) - ) - long_damp_v = damp_pair_v + other_v - long_damp_t = damp_pair_t + other_t - spun_mix_free = _dual_times_i( - *_dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_free) - ) - spun_mix_bound = _dual_times_i( - *_dual_mul(spin[0], spin[1], spin[2], spin[3], *mixed_bound) - ) - zangle_v, zangle_t = _dual_real_conj_mul(*free_bar, *spun_mix_free) - part_v, part_t = _dual_real_conj_mul(*bound_bar, *spun_mix_bound) - zangle_v += part_v - zangle_t += part_t - back_z = _dual_mul( - pe11 * spin[0], - -(pe11 * spin[1]), - de11 * spin[0] + pe11 * spin[2], - -(de11 * spin[1] + pe11 * spin[3]), - *free_bar, - ) - cross_z = _dual_mul( - pe21 * spin[0], - -(pe21 * spin[1]), - de21 * spin[0] + pe21 * spin[2], - -(de21 * spin[1] + pe21 * spin[3]), - *bound_bar, - ) - back_b = _dual_mul( - pe12 * spin[0], - -(pe12 * spin[1]), - de12 * spin[0] + pe12 * spin[2], - -(de12 * spin[1] + pe12 * spin[3]), - *free_bar, - ) - cross_b = _dual_mul( - pe22 * spin[0], - -(pe22 * spin[1]), - de22 * spin[0] + pe22 * spin[2], - -(de22 * spin[1] + pe22 * spin[3]), - *bound_bar, - ) - next_zbvr = back_z[0] + cross_z[0] - next_zbvi = back_z[1] + cross_z[1] - next_zbtr = back_z[2] + cross_z[2] - next_zbti = back_z[3] + cross_z[3] - bbvr = back_b[0] + cross_b[0] - bbvi = back_b[1] + cross_b[1] - bbtr = back_b[2] + cross_b[2] - bbti = back_b[3] + cross_b[3] - zbvr = next_zbvr - zbvi = next_zbvi - zbtr = next_zbtr - zbti = next_zbti - else: - spun = _dual_mul(szr, szi, sztr, szti, xzvr, xzvi, xztr, xzti) - e1_v, e1_t = _dual_real_conj_mul( - zbvr, zbvi, zbtr, zbti, spun[0], spun[1], spun[2], spun[3] - ) - grad_e1_v = tl.sum(e1_v * damp_z, axis=1)[:, None] - grad_e1_t = tl.sum(e1_v * damp_z_tangent + e1_t * damp_z, axis=1)[:, None] - grad_e1_v -= tl.sum(tl.where(state == 0, zbvr, 0.0), axis=1)[:, None] - grad_e1_t -= tl.sum(tl.where(state == 0, zbtr, 0.0), axis=1)[:, None] - long_damp_v = e1_v * bare1_value * damp_z - long_damp_t = ( - e1_t * bare1_value * damp_z - + e1_v * bare1_tangent * damp_z - + e1_v * bare1_value * damp_z_tangent - ) - # The longitudinal states turn too, and by a whole order rather - # than the transverse half-order more. - zo = _dual_mul(lvr, lvi, ltr, lti, xzvr, xzvi, xztr, xzti) - zo = _dual_times_i(zo[0], zo[1], zo[2], zo[3]) - zangle_v, zangle_t = _dual_real_conj_mul( - zbvr, zbvi, zbtr, zbti, zo[0], zo[1], zo[2], zo[3] - ) - zbvr, zbvi, zbtr, zbti = _dual_mul( - lvr, -lvi, ltr, -lti, zbvr, zbvi, zbtr, zbti - ) - - next_pb = (pbvr, pbvi, pbtr, pbti) - next_mb = (mbvr, mbvi, mbtr, mbti) - next_ub = (ubvr, ubvi, ubtr, ubti) - next_wb = (wbvr, wbvi, wbtr, wbti) - if pools == 2 or pools == 3: - # The four entries of the transverse operator, summed over the - # orders that share them, then pushed back through the closed form - # once for the whole interval. ``F-`` follows the conjugate of the - # operator, so its cotangent lands on the entry itself rather than - # on the conjugate of it. - ap = (pbvr, pbvi, pbtr, pbti) - am = (mbvr, mbvi, mbtr, mbti) - aub = (ubvr, ubvi, ubtr, ubti) - awb = (wbvr, wbvi, wbtr, wbti) - fp = (xpvr, xpvi, xptr, xpti) - fm = (xmvr, xmvi, xmtr, xmti) - bp = (xbpvr, xbpvi, xbptr, xbpti) - bm = (xbmvr, xbmvi, xbmtr, xbmti) - term11 = _dual_product( - _dual_add( - _dual_product(_dual_conj(ap), fp), - _dual_product(am, _dual_conj(fm)), - ), - carried, - ) - term12 = _dual_product( - _dual_add( - _dual_product(_dual_conj(ap), bp), - _dual_product(am, _dual_conj(bm)), - ), - carried, - ) - term21 = _dual_product( - _dual_add( - _dual_product(_dual_conj(aub), fp), - _dual_product(awb, _dual_conj(fm)), - ), - carried, - ) - term22 = _dual_product( - _dual_add( - _dual_product(_dual_conj(aub), bp), - _dual_product(awb, _dual_conj(bm)), - ), - carried, - ) - bar11 = ( - tl.sum(term11[0], axis=1)[:, None], - tl.sum(term11[1], axis=1)[:, None], - tl.sum(term11[2], axis=1)[:, None], - tl.sum(term11[3], axis=1)[:, None], - ) - bar12 = ( - tl.sum(term12[0], axis=1)[:, None], - tl.sum(term12[1], axis=1)[:, None], - tl.sum(term12[2], axis=1)[:, None], - tl.sum(term12[3], axis=1)[:, None], - ) - bar21 = ( - tl.sum(term21[0], axis=1)[:, None], - tl.sum(term21[1], axis=1)[:, None], - tl.sum(term21[2], axis=1)[:, None], - tl.sum(term21[3], axis=1)[:, None], - ) - bar22 = ( - tl.sum(term22[0], axis=1)[:, None], - tl.sum(term22[1], axis=1)[:, None], - tl.sum(term22[2], axis=1)[:, None], - tl.sum(term22[3], axis=1)[:, None], - ) - ( - back_r2_v, - back_r2_t, - back_r2b_v, - back_r2b_t, - back_xexch_v, - back_xexch_t, - back_xbound_v, - back_xbound_t, - back_xfree_v, - back_xfree_t, - back_shift_v, - back_shift_t, - back_xdt_v, - back_xdt_t, - back_xatt_v, - back_xatt_t, - ) = _two_pool_transverse_adjoint_jvp( - r2_value, - r2_tangent, - r2b_value, - r2b_tangent, - atom_exchange, - d_exchange, - atom_bound, - d_boundf, - atom_free, - d_free, - atom_shift, - d_shift, - dt_value, - dt_tangent, - wout_value, - wout_tangent, - bar11, - bar12, - bar21, - bar22, - ) - slope2_v = -1000.0 / (atom_t2 * atom_t2) - slope2_t = 2000.0 * d_t2 / (atom_t2 * atom_t2 * atom_t2) - slope2b_v = -1000.0 / (atom_t2b * atom_t2b) - slope2b_t = 2000.0 * d_t2b / (atom_t2b * atom_t2b * atom_t2b) - g_t2v += back_r2_v * slope2_v - g_t2t += back_r2_t * slope2_v + back_r2_v * slope2_t - g_t2bv += back_r2b_v * slope2b_v - g_t2bt += back_r2b_t * slope2b_v + back_r2b_v * slope2b_t - g_exchv += back_xexch_v - g_excht += back_xexch_t - # The free water is what both second pools leave, so a cotangent - # on it reaches each of their fractions turned over. - g_boundv += back_xbound_v - back_xfree_v - g_boundt += back_xbound_t - back_xfree_t - if pools == 3: - g_semiv -= back_xfree_v - g_semit -= back_xfree_t - g_shiftv += back_shift_v - g_shiftt += back_shift_t - attenuation_v += back_xatt_v - attenuation_t += back_xatt_t - two_pool_dt_v += back_xdt_v - two_pool_dt_t += back_xdt_t - step11 = _dual_product(a11, carried) - step12 = _dual_product(a12, carried) - step21 = _dual_product(a21, carried) - step22 = _dual_product(a22, carried) - next_pb = _dual_add( - _dual_product(_dual_conj(step11), ap), - _dual_product(_dual_conj(step21), aub), - ) - next_ub = _dual_add( - _dual_product(_dual_conj(step12), ap), - _dual_product(_dual_conj(step22), aub), - ) - next_mb = _dual_add(_dual_product(step11, am), _dual_product(step21, awb)) - next_wb = _dual_add(_dual_product(step12, am), _dual_product(step22, awb)) - - # The rate and the interval multiply every order's b-weight, so both - # take a weighted sum rather than one scalar. Order zero carries no - # longitudinal weight, which keeps recovery out of this. - spread_v = zero - spread_t = zero - if diffusing: - weighted_v = long_damp_v * longitudinal_weight + cot2_v * transverse_weight - weighted_t = long_damp_t * longitudinal_weight + cot2_t * transverse_weight - spread_v = tl.sum(weighted_v, axis=1)[:, None] - spread_t = tl.sum(weighted_t, axis=1)[:, None] - g_diffv += -spread_v * dt_value - g_difft += -(spread_v * dt_tangent + spread_t * dt_value) - - wound_v = zero - wound_t = zero - if moving: - wound_v = tl.sum(per_angle_v * (order + 0.5) + zangle_v * order, axis=1)[ - :, None - ] - wound_t = tl.sum(per_angle_t * (order + 0.5) + zangle_t * order, axis=1)[ - :, None - ] - g_flowv += -wound_v * dt_value - g_flowt += -(wound_v * dt_tangent + wound_t * dt_value) - - # Washout scales both relaxation factors, so its gradient is the one - # they already carry, taken against the factors before that scaling. - # Past the clamp the interval has replaced the voxel outright and - # nothing further depends on the rate. - wash_v = zero - wash_t = zero - if moving: - live = (atom_washout * dt_value < 1.0).to(tl.float32) - wash_v = -live * ( - grad_e1_v * dry1_value + grad_e2_v * dry2_value + attenuation_v - ) - wash_t = -live * ( - grad_e1_v * dry1_tangent - + grad_e1_t * dry1_value - + grad_e2_v * dry2_tangent - + grad_e2_t * dry2_value - + attenuation_t - ) - g_washv += wash_v * dt_value - g_washt += wash_v * dt_tangent + wash_t * dt_value - - if pools == 2 or pools == 3: - pbvr, pbvi, pbtr, pbti = next_pb - mbvr, mbvi, mbtr, mbti = next_mb - ubvr, ubvi, ubtr, ubti = next_ub - wbvr, wbvi, wbtr, wbti = next_wb - else: - pbvr, pbvi, pbtr, pbti = _dual_mul( - ovr, -ovi, otr, -oti, pbvr, pbvi, pbtr, pbti - ) - mbvr, mbvi, mbtr, mbti = _dual_mul( - ovr, ovi, otr, oti, mbvr, mbvi, mbtr, mbti - ) - - inverse1_value = 1000.0 / (atom_t1 * atom_t1) - inverse1_tangent = -2000.0 * d_t1 / (atom_t1 * atom_t1 * atom_t1) - inverse2_value = 1000.0 / (atom_t2 * atom_t2) - inverse2_tangent = -2000.0 * d_t2 / (atom_t2 * atom_t2 * atom_t2) - scale1_value = bare1_value * dt_value * inverse1_value - scale1_tangent = bare1_tangent * dt_value * inverse1_value - scale1_tangent += bare1_value * dt_tangent * inverse1_value - scale1_tangent += bare1_value * dt_value * inverse1_tangent - scale2_value = bare2_value * dt_value * inverse2_value - scale2_tangent = bare2_tangent * dt_value * inverse2_value - scale2_tangent += bare2_value * dt_tangent * inverse2_value - scale2_tangent += bare2_value * dt_value * inverse2_tangent - g_t1v += grad_e1_v * scale1_value - g_t1t += grad_e1_v * scale1_tangent + grad_e1_t * scale1_value - g_t2v += grad_e2_v * scale2_value - g_t2t += grad_e2_v * scale2_tangent + grad_e2_t * scale2_value - - turn = -2.0 * 3.141592653589793 - g_b0v += grad_angle_v * (turn * dt_value) - g_b0t += grad_angle_v * (turn * dt_tangent) + grad_angle_t * (turn * dt_value) - - decay1_value = r1_value * bare1_value - decay1_tangent = r1_value * bare1_tangent + r1_tangent * bare1_value - decay2_value = r2_value * bare2_value - decay2_tangent = r2_value * bare2_tangent + r2_tangent * bare2_value - duration_v = -grad_e1_v * decay1_value - grad_e2_v * decay2_value - duration_v += grad_angle_v * (turn * atom_b0) + two_pool_dt_v - duration_t = -(grad_e1_v * decay1_tangent + grad_e1_t * decay1_value) - duration_t -= grad_e2_v * decay2_tangent + grad_e2_t * decay2_value - duration_t += grad_angle_v * (turn * d_b0) + grad_angle_t * (turn * atom_b0) - duration_t += two_pool_dt_t - duration_v += -spread_v * atom_damping - wound_v * atom_flow - duration_t += -(spread_v * d_damping + spread_t * atom_damping) - duration_t += -(wound_v * d_flow + wound_t * atom_flow) - duration_v += wash_v * atom_washout - duration_t += wash_v * d_washout + wash_t * atom_washout - tl.atomic_add( - grad_duration_value + event_base + event, duration_v, mask=active_atom - ) - tl.atomic_add( - grad_duration_tangent + event_base + event, duration_t, mask=active_atom - ) - - if pools == 3 and tabulated: - # One closed form per distinct length rather than one per event, - # run twice. The walk back pooled the cotangents the eigenvalues - # are pushed through and the closed form is linear in them, so the - # pieces of the sum are the sum of the pieces. A gradient's own - # direction depends on the interval as well, and a row is shared - # by events whose interval directions differ -- so the second pass - # takes that dependence alone, driven by the cotangents the walk - # back weighted by each event's direction and read at a unit one. - for row in range(0, row_count): - held = pool_bars + (local * row_count + row) * 36 - row_dt = tl.load(pool_durations + row) + zero - nil = 0.0 * row_dt - unit = 1.0 + nil - one_att = unit - att_rate = nil - att_span = nil - if moving: - one_att, att_rate = _washout_jvp(atom_washout, d_washout, row_dt, nil) - _held_att, att_span = _washout_jvp(atom_washout, nil, row_dt, unit) - ( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - ) = _three_pool_pieces_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - row_dt, - nil, - narrow, - ) - ( - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - ) = _three_pool_assemble_jvp( - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - narrow, - ) - ( - back_r1_v, - back_r1b_v, - back_r1c_v, - back_exch_v, - back_sexch_v, - back_bound_v, - back_semi_v, - _row_dt_v, - _row_att_v, - back_r1_t, - back_r1b_t, - back_r1c_t, - back_exch_t, - back_sexch_t, - back_bound_t, - back_semi_t, - _row_dt_t, - _row_att_t, - ) = _three_pool_step_adjoint_jvp( - r1_value, - r1_tangent, - r1b_value, - r1b_tangent, - r1c_value, - r1c_tangent, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_boundf, - atom_semisolid, - d_semisolidf, - row_dt, - nil, - one_att, - att_rate, - tl.load(held + 0, mask=active_atom, other=0.0), - tl.load(held + 12, mask=active_atom, other=0.0), - tl.load(held + 1, mask=active_atom, other=0.0), - tl.load(held + 13, mask=active_atom, other=0.0), - tl.load(held + 2, mask=active_atom, other=0.0), - tl.load(held + 14, mask=active_atom, other=0.0), - tl.load(held + 3, mask=active_atom, other=0.0), - tl.load(held + 15, mask=active_atom, other=0.0), - tl.load(held + 4, mask=active_atom, other=0.0), - tl.load(held + 16, mask=active_atom, other=0.0), - tl.load(held + 5, mask=active_atom, other=0.0), - tl.load(held + 17, mask=active_atom, other=0.0), - tl.load(held + 6, mask=active_atom, other=0.0), - tl.load(held + 18, mask=active_atom, other=0.0), - tl.load(held + 7, mask=active_atom, other=0.0), - tl.load(held + 19, mask=active_atom, other=0.0), - tl.load(held + 8, mask=active_atom, other=0.0), - tl.load(held + 20, mask=active_atom, other=0.0), - tl.load(held + 9, mask=active_atom, other=0.0), - tl.load(held + 21, mask=active_atom, other=0.0), - tl.load(held + 10, mask=active_atom, other=0.0), - tl.load(held + 22, mask=active_atom, other=0.0), - tl.load(held + 11, mask=active_atom, other=0.0), - tl.load(held + 23, mask=active_atom, other=0.0), - three_free, - three_d_free, - three_pool_b, - three_d_pool_b, - three_pool_c, - three_d_pool_c, - three_a00, - three_d_a00, - three_a01, - three_d_a01, - three_a02, - three_d_a02, - three_a10, - three_d_a10, - three_a11, - three_d_a11, - three_a20, - three_d_a20, - three_a22, - three_d_a22, - three_s00, - three_d_s00, - three_s11, - three_d_s11, - three_s22, - three_d_s22, - three_minors, - three_d_minors, - three_sum_flat, - three_sum_linear, - three_sum_square, - three_d_sum_flat, - three_d_sum_linear, - three_d_sum_square, - three_lift, - three_d_lift, - three_low, - three_middle, - three_d_low, - three_d_middle, - three_leading, - three_d_leading, - three_first, - three_d_first, - three_second, - three_d_second, - three_determinant, - three_d_determinant, - three_high, - three_d_high, - three_radius, - three_d_radius, - three_cube, - three_raw, - three_d_raw, - three_argument, - three_inside_limit, - three_angle, - three_d_angle, - three_centre, - three_d_centre, - three_trailing, - three_d_trailing, - three_guarded, - three_d_guarded, - three_q00, - three_d_q00, - three_q01, - three_d_q01, - three_q02, - three_d_q02, - three_q10, - three_d_q10, - three_q11, - three_d_q11, - three_q12, - three_d_q12, - three_q20, - three_d_q20, - three_q21, - three_d_q21, - three_q22, - three_d_q22, - three_def_00, - three_dif_00, - three_def_01, - three_dif_01, - three_def_02, - three_dif_02, - three_def_10, - three_dif_10, - three_def_11, - three_dif_11, - three_def_12, - three_dif_12, - three_def_20, - three_dif_20, - three_def_21, - three_dif_21, - three_def_22, - three_dif_22, - narrow, - ) - ( - alt_free, - alt_d_free, - alt_pool_b, - alt_d_pool_b, - alt_pool_c, - alt_d_pool_c, - alt_a00, - alt_d_a00, - alt_a01, - alt_d_a01, - alt_a02, - alt_d_a02, - alt_a10, - alt_d_a10, - alt_a11, - alt_d_a11, - alt_a20, - alt_d_a20, - alt_a22, - alt_d_a22, - alt_s00, - alt_d_s00, - alt_s11, - alt_d_s11, - alt_s22, - alt_d_s22, - alt_minors, - alt_d_minors, - alt_sum_flat, - alt_sum_linear, - alt_sum_square, - alt_d_sum_flat, - alt_d_sum_linear, - alt_d_sum_square, - alt_lift, - alt_d_lift, - alt_low, - alt_middle, - alt_d_low, - alt_d_middle, - alt_leading, - alt_d_leading, - alt_first, - alt_d_first, - alt_second, - alt_d_second, - alt_determinant, - alt_d_determinant, - alt_high, - alt_d_high, - alt_radius, - alt_d_radius, - alt_cube, - alt_raw, - alt_d_raw, - alt_argument, - alt_inside_limit, - alt_angle, - alt_d_angle, - alt_centre, - alt_d_centre, - alt_trailing, - alt_d_trailing, - alt_guarded, - alt_d_guarded, - alt_q00, - alt_d_q00, - alt_q01, - alt_d_q01, - alt_q02, - alt_d_q02, - alt_q10, - alt_d_q10, - alt_q11, - alt_d_q11, - alt_q12, - alt_d_q12, - alt_q20, - alt_d_q20, - alt_q21, - alt_d_q21, - alt_q22, - alt_d_q22, - ) = _three_pool_pieces_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - row_dt, - unit, - narrow, - ) - ( - alt_def_00, - alt_dif_00, - alt_def_01, - alt_dif_01, - alt_def_02, - alt_dif_02, - alt_def_10, - alt_dif_10, - alt_def_11, - alt_dif_11, - alt_def_12, - alt_dif_12, - alt_def_20, - alt_dif_20, - alt_def_21, - alt_dif_21, - alt_def_22, - alt_dif_22, - ) = _three_pool_assemble_jvp( - alt_free, - alt_d_free, - alt_pool_b, - alt_d_pool_b, - alt_pool_c, - alt_d_pool_c, - alt_a00, - alt_d_a00, - alt_a01, - alt_d_a01, - alt_a02, - alt_d_a02, - alt_a10, - alt_d_a10, - alt_a11, - alt_d_a11, - alt_a20, - alt_d_a20, - alt_a22, - alt_d_a22, - alt_s00, - alt_d_s00, - alt_s11, - alt_d_s11, - alt_s22, - alt_d_s22, - alt_minors, - alt_d_minors, - alt_sum_flat, - alt_sum_linear, - alt_sum_square, - alt_d_sum_flat, - alt_d_sum_linear, - alt_d_sum_square, - alt_lift, - alt_d_lift, - alt_low, - alt_middle, - alt_d_low, - alt_d_middle, - alt_leading, - alt_d_leading, - alt_first, - alt_d_first, - alt_second, - alt_d_second, - alt_determinant, - alt_d_determinant, - alt_high, - alt_d_high, - alt_radius, - alt_d_radius, - alt_cube, - alt_raw, - alt_d_raw, - alt_argument, - alt_inside_limit, - alt_angle, - alt_d_angle, - alt_centre, - alt_d_centre, - alt_trailing, - alt_d_trailing, - alt_guarded, - alt_d_guarded, - alt_q00, - alt_d_q00, - alt_q01, - alt_d_q01, - alt_q02, - alt_d_q02, - alt_q10, - alt_d_q10, - alt_q11, - alt_d_q11, - alt_q12, - alt_d_q12, - alt_q20, - alt_d_q20, - alt_q21, - alt_d_q21, - alt_q22, - alt_d_q22, - narrow, - ) - ( - _span_r1_v, - _span_r1b_v, - _span_r1c_v, - _span_exch_v, - _span_sexch_v, - _span_bound_v, - _span_semi_v, - _span_dt_v, - _span_att_v, - span_r1_t, - span_r1b_t, - span_r1c_t, - span_exch_t, - span_sexch_t, - span_bound_t, - span_semi_t, - _span_dt_t, - _span_att_t, - ) = _three_pool_step_adjoint_jvp( - r1_value, - nil, - r1b_value, - nil, - r1c_value, - nil, - atom_exchange, - nil, - atom_semisolid_exchange, - nil, - atom_bound, - nil, - atom_semisolid, - nil, - row_dt, - unit, - one_att, - att_span, - tl.load(held + 24, mask=active_atom, other=0.0), - nil, - tl.load(held + 25, mask=active_atom, other=0.0), - nil, - tl.load(held + 26, mask=active_atom, other=0.0), - nil, - tl.load(held + 27, mask=active_atom, other=0.0), - nil, - tl.load(held + 28, mask=active_atom, other=0.0), - nil, - tl.load(held + 29, mask=active_atom, other=0.0), - nil, - tl.load(held + 30, mask=active_atom, other=0.0), - nil, - tl.load(held + 31, mask=active_atom, other=0.0), - nil, - tl.load(held + 32, mask=active_atom, other=0.0), - nil, - tl.load(held + 33, mask=active_atom, other=0.0), - nil, - tl.load(held + 34, mask=active_atom, other=0.0), - nil, - tl.load(held + 35, mask=active_atom, other=0.0), - nil, - alt_free, - alt_d_free, - alt_pool_b, - alt_d_pool_b, - alt_pool_c, - alt_d_pool_c, - alt_a00, - alt_d_a00, - alt_a01, - alt_d_a01, - alt_a02, - alt_d_a02, - alt_a10, - alt_d_a10, - alt_a11, - alt_d_a11, - alt_a20, - alt_d_a20, - alt_a22, - alt_d_a22, - alt_s00, - alt_d_s00, - alt_s11, - alt_d_s11, - alt_s22, - alt_d_s22, - alt_minors, - alt_d_minors, - alt_sum_flat, - alt_sum_linear, - alt_sum_square, - alt_d_sum_flat, - alt_d_sum_linear, - alt_d_sum_square, - alt_lift, - alt_d_lift, - alt_low, - alt_middle, - alt_d_low, - alt_d_middle, - alt_leading, - alt_d_leading, - alt_first, - alt_d_first, - alt_second, - alt_d_second, - alt_determinant, - alt_d_determinant, - alt_high, - alt_d_high, - alt_radius, - alt_d_radius, - alt_cube, - alt_raw, - alt_d_raw, - alt_argument, - alt_inside_limit, - alt_angle, - alt_d_angle, - alt_centre, - alt_d_centre, - alt_trailing, - alt_d_trailing, - alt_guarded, - alt_d_guarded, - alt_q00, - alt_d_q00, - alt_q01, - alt_d_q01, - alt_q02, - alt_d_q02, - alt_q10, - alt_d_q10, - alt_q11, - alt_d_q11, - alt_q12, - alt_d_q12, - alt_q20, - alt_d_q20, - alt_q21, - alt_d_q21, - alt_q22, - alt_d_q22, - alt_def_00, - alt_dif_00, - alt_def_01, - alt_dif_01, - alt_def_02, - alt_dif_02, - alt_def_10, - alt_dif_10, - alt_def_11, - alt_dif_11, - alt_def_12, - alt_dif_12, - alt_def_20, - alt_dif_20, - alt_def_21, - alt_dif_21, - alt_def_22, - alt_dif_22, - narrow, - ) - slope1_v = -1000.0 / (atom_t1 * atom_t1) - slope1_t = 2000.0 * d_t1 / (atom_t1 * atom_t1 * atom_t1) - slope1b_v = -1000.0 / (atom_t1b * atom_t1b) - slope1b_t = 2000.0 * d_t1b / (atom_t1b * atom_t1b * atom_t1b) - slope1c_v = -1000.0 / (held_semisolid * held_semisolid) - slope1c_t = ( - 2000.0 - * d_semisolid_t1 - / (held_semisolid * held_semisolid * held_semisolid) - ) - row_r1_t = back_r1_t + span_r1_t - row_r1b_t = back_r1b_t + span_r1b_t - row_r1c_t = back_r1c_t + span_r1c_t - g_t1v += back_r1_v * slope1_v - g_t1t += row_r1_t * slope1_v + back_r1_v * slope1_t - g_t1bv += back_r1b_v * slope1b_v - g_t1bt += row_r1b_t * slope1b_v + back_r1b_v * slope1b_t - g_t1cv += back_r1c_v * slope1c_v - g_t1ct += row_r1c_t * slope1c_v + back_r1c_v * slope1c_t - g_exchv += back_exch_v - g_excht += back_exch_t + span_exch_t - g_sexchv += back_sexch_v - g_sexcht += back_sexch_t + span_sexch_t - g_boundv += back_bound_v - g_boundt += back_bound_t + span_bound_t - g_semiv += back_semi_v - g_semit += back_semi_t + span_semi_t - - velocity_v = g_flowv * flow_scale + g_washv * direction * washout_scale - velocity_t = g_flowt * flow_scale + g_washt * direction * washout_scale - if pools > 0: - # The fraction also sets where each pool starts, which the walk back - # reaches last. - g_boundv += tl.sum(tl.where(state == 0, bbvr - zbvr, 0.0), axis=1)[:, None] - g_boundt += tl.sum(tl.where(state == 0, bbtr - zbtr, 0.0), axis=1)[:, None] - if pools == 3: - g_semiv += tl.sum(tl.where(state == 0, cbvr - zbvr, 0.0), axis=1)[:, None] - g_semit += tl.sum(tl.where(state == 0, cbtr - zbtr, 0.0), axis=1)[:, None] - if pools == 1: - base_row = _BOUND_ROW + 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue_value + base_row * atom_count + atom, - g_boundv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + base_row * atom_count + atom, - g_boundt, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_value + (base_row + 1) * atom_count + atom, - g_exchv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (base_row + 1) * atom_count + atom, - g_excht, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_value + (base_row + 2) * atom_count + atom, - g_t1bv, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (base_row + 2) * atom_count + atom, - g_t1bt, - mask=active_atom, - ) - if pools == 3: - semisolid_row = _BOUND_ROW + 2 * (shim_rows - 1) - stuck = (g_semiv, g_sexchv, g_t1cv) - stuck_tangents = (g_semit, g_sexcht, g_t1ct) - for offset in tl.static_range(3): - tl.atomic_add( - grad_tissue_value + (semisolid_row + offset) * atom_count + atom, - stuck[offset], - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (semisolid_row + offset) * atom_count + atom, - stuck_tangents[offset], - mask=active_atom, - ) - if pools == 2 or pools == 3: - base_row = _POOL_B_ROW + 2 * (shim_rows - 1) - rows = (g_boundv, g_exchv, g_t1bv, g_t2bv, g_shiftv) - tangent_rows = (g_boundt, g_excht, g_t1bt, g_t2bt, g_shiftt) - for offset in tl.static_range(5): - tl.atomic_add( - grad_tissue_value + (base_row + offset) * atom_count + atom, - rows[offset], - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (base_row + offset) * atom_count + atom, - tangent_rows[offset], - mask=active_atom, - ) - values = ( - g_t1v, - g_t2v, - g_m0v, - g_b1v, - g_b1pv, - g_b0v, - g_invv, - g_diffv, - velocity_v, - ) - tangents = ( - g_t1t, - g_t2t, - g_m0t, - g_b1t, - g_b1pt, - g_b0t, - g_invt, - g_difft, - velocity_t, - ) - for parameter in tl.static_range(_FREE_POOL_COUNT): - # The transmit pair went to its shim's row above when there is more - # than one; the rest sit past whatever rows that pair took. - if not shimmed or (parameter != _B1_ROW and parameter != _B1_PHASE_ROW): - plane = ( - parameter if parameter < _B1_ROW else parameter + 2 * (shim_rows - 1) - ) - tl.atomic_add( - grad_tissue_value + plane * atom_count + atom, - values[parameter], - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + plane * atom_count + atom, - tangents[parameter], - mask=active_atom, - ) - - -@triton.jit(do_not_specialize=["state_count"]) -def _epg_real_vjp_jvp_kernel( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - dot_t1, - dot_t2, - dot_m0, - dot_b1, - dot_inversion_efficiency, - dot_diffusion, - dot_duration, - dot_flip, - grad_output_imag, - grad_tissue_value, - grad_tissue_tangent, - grad_flip_value, - grad_flip_tangent, - grad_duration_value, - grad_duration_tangent, - trajectory_value, - trajectory_tangent, - problem_base, - problem_end, - atom_count, - train_count, - event_count, - output_count, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shim_rows, - shimmed: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = problem_base + tl.program_id(0) * problems - problem = problem + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - # The grid rounds up to whole tiles, so the last program of a wave reaches - # past it. Those problems are real, but their trajectory rows belong to a - # later launch and do not exist yet. - active_atom = problem < problem_end - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - # The trajectory holds the state entering every event: three planes of - # configuration orders, for the value and the tangent alike. - record_stride = 3 * state_count - trajectory = (problem - problem_base) * event_count * record_stride + state - minus_plane = state_count - long_plane = 2 * state_count - - empty = tl.zeros((problems, block_states), tl.float32) - plus_value = empty - plus_tangent = empty - minus_value = empty - minus_tangent = empty - long_value = empty + tl.where(state == 0, 1.0, 0.0) - long_tangent = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - atom_dot_t1 = tl.load(dot_t1 + atom, mask=active_atom, other=0.0) - atom_dot_t2 = tl.load(dot_t2 + atom, mask=active_atom, other=0.0) - atom_dot_m0 = 0.0 - if density: - atom_dot_m0 = tl.load(dot_m0 + scalar_atom, mask=active_atom, other=0.0) - atom_dot_b1 = 0.0 - if transmit: - atom_dot_b1 = tl.load(dot_b1 + scalar_atom, mask=active_atom, other=0.0) - atom_dot_inversion = 0.0 - if inverting: - atom_dot_inversion = tl.load( - dot_inversion_efficiency + scalar_atom, mask=active_atom, other=0.0 - ) - atom_damping = 0.0 - atom_dot_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - atom_dot_damping = tl.load( - dot_diffusion + scalar_atom, mask=active_atom, other=0.0 - ) - order = state.to(tl.float32) - longitudinal_weight = order * order - transverse_weight = longitudinal_weight + order + 0.3333333333333333 - rate1_value = 1000.0 / atom_t1 - rate1_tangent = -1000.0 * atom_dot_t1 / (atom_t1 * atom_t1) - rate2_value = 1000.0 / atom_t2 - rate2_tangent = -1000.0 * atom_dot_t2 / (atom_t2 * atom_t2) - - event_base = train * event_count - for event in range(0, event_count): - slot = trajectory + event * record_stride - tl.store(trajectory_value + slot, plus_value, mask=state_mask) - tl.store(trajectory_value + slot + minus_plane, minus_value, mask=state_mask) - tl.store(trajectory_value + slot + long_plane, long_value, mask=state_mask) - tl.store(trajectory_tangent + slot, plus_tangent, mask=state_mask) - tl.store( - trajectory_tangent + slot + minus_plane, minus_tangent, mask=state_mask - ) - tl.store(trajectory_tangent + slot + long_plane, long_tangent, mask=state_mask) - - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - dt_tangent = _event_value( - dot_duration, event_base, event, active_atom, single_train - ) - e1_value = tl.exp(-rate1_value * dt_value) - e1_tangent = -e1_value * (rate1_value * dt_tangent + rate1_tangent * dt_value) - e2_value = tl.exp(-rate2_value * dt_value) - e2_tangent = -e2_value * (rate2_value * dt_tangent + rate2_tangent * dt_value) - damp_z = 1.0 - damp_z_tangent = 0.0 - damp_t = 1.0 - damp_t_tangent = 0.0 - if diffusing: - damp_z, damp_z_tangent, damp_t, damp_t_tangent = _damping_jvp( - atom_damping, atom_dot_damping, dt_value, dt_tangent, order - ) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value, recovery_tangent = 1.0 - e1_value, -e1_tangent - bare1_value, bare1_tangent = e1_value, e1_tangent - bare2_value, bare2_tangent = e2_value, e2_tangent - e1_tangent = e1_tangent * damp_z + bare1_value * damp_z_tangent - e1_value = bare1_value * damp_z - e2_tangent = e2_tangent * damp_t + bare2_value * damp_t_tangent - e2_value = bare2_value * damp_t - - plus_tangent = plus_value * e2_tangent + plus_tangent * e2_value - plus_value = plus_value * e2_value - minus_tangent = minus_value * e2_tangent + minus_tangent * e2_value - minus_value = minus_value * e2_value - long_tangent = long_value * e1_tangent + long_tangent * e1_value - long_value = long_value * e1_value - long_value += tl.where(state == 0, recovery_value, 0.0) - long_tangent += tl.where(state == 0, recovery_tangent, 0.0) - - event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - shifted_pv, shifted_mv = _shift_real( - plus_value, minus_value, state, state_mask, state_count - ) - shifted_pt, shifted_mt = _shift_real( - plus_tangent, minus_tangent, state, state_mask, state_count - ) - plus_value = tl.where(pre_shift, shifted_pv, plus_value) - minus_value = tl.where(pre_shift, shifted_mv, minus_value) - plus_tangent = tl.where(pre_shift, shifted_pt, plus_tangent) - minus_tangent = tl.where(pre_shift, shifted_mt, minus_tangent) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - inverted_value = -atom_inversion * long_value - inverted_tangent = -atom_inversion * long_tangent - inverted_tangent -= atom_dot_inversion * long_value - long_value = tl.where(invert, inverted_value, long_value) - long_tangent = tl.where(invert, inverted_tangent, long_tangent) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - event_dot_flip = _event_value( - dot_flip, event_base, event, active_atom, single_train - ) - pulse_b1 = atom_b1 - pulse_dot_b1 = atom_dot_b1 - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse the row of the shim it drives. - if shimmed: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + shim_row + atom, mask=active_atom, other=1.0) - pulse_dot_b1 = tl.load( - dot_b1 + shim_row + atom, mask=active_atom, other=0.0 - ) - alpha_value = event_flip * pulse_b1 - alpha_tangent = event_dot_flip * pulse_b1 + event_flip * pulse_dot_b1 - cosine_value = tl.cos(alpha_value) - sine_value = tl.sin(alpha_value) - cosine_tangent = -sine_value * alpha_tangent - sine_tangent = cosine_value * alpha_tangent - chs_value = 0.5 * (1.0 + cosine_value) - chs_tangent = 0.5 * cosine_tangent - shs_value = 0.5 * (1.0 - cosine_value) - shs_tangent = -0.5 * cosine_tangent - half_sine_value = 0.5 * sine_value - half_sine_tangent = 0.5 * sine_tangent - - rotated_pv = chs_value * plus_value + shs_value * minus_value - rotated_pv -= sine_value * long_value - rotated_pt = chs_value * plus_tangent + chs_tangent * plus_value - rotated_pt += shs_value * minus_tangent + shs_tangent * minus_value - rotated_pt -= sine_value * long_tangent + sine_tangent * long_value - rotated_mv = shs_value * plus_value + chs_value * minus_value - rotated_mv += sine_value * long_value - rotated_mt = shs_value * plus_tangent + shs_tangent * plus_value - rotated_mt += chs_value * minus_tangent + chs_tangent * minus_value - rotated_mt += sine_value * long_tangent + sine_tangent * long_value - rotated_zv = half_sine_value * plus_value - half_sine_value * minus_value - rotated_zv += cosine_value * long_value - rotated_zt = half_sine_value * plus_tangent + half_sine_tangent * plus_value - rotated_zt -= half_sine_value * minus_tangent + half_sine_tangent * minus_value - rotated_zt += cosine_value * long_tangent + cosine_tangent * long_value - - rotate = is_rf & ~is_inversion - plus_value = tl.where(rotate, rotated_pv, plus_value) - plus_tangent = tl.where(rotate, rotated_pt, plus_tangent) - minus_value = tl.where(rotate, rotated_mv, minus_value) - minus_tangent = tl.where(rotate, rotated_mt, minus_tangent) - long_value = tl.where(rotate, rotated_zv, long_value) - long_tangent = tl.where(rotate, rotated_zt, long_tangent) - - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - shifted_pv, shifted_mv = _shift_real( - plus_value, minus_value, state, state_mask, state_count - ) - shifted_pt, shifted_mt = _shift_real( - plus_tangent, minus_tangent, state, state_mask, state_count - ) - plus_value = tl.where(do_shift, shifted_pv, plus_value) - minus_value = tl.where(do_shift, shifted_mv, minus_value) - plus_tangent = tl.where(do_shift, shifted_pt, plus_tangent) - minus_tangent = tl.where(do_shift, shifted_mt, minus_tangent) - spoil = (event_action & 8) != 0 - plus_value = tl.where(spoil, 0.0, plus_value) - minus_value = tl.where(spoil, 0.0, minus_value) - plus_tangent = tl.where(spoil, 0.0, plus_tangent) - minus_tangent = tl.where(spoil, 0.0, minus_tangent) - - plus_bar_value = empty - plus_bar_tangent = empty - minus_bar_value = empty - minus_bar_tangent = empty - long_bar_value = empty - long_bar_tangent = empty - zero = tl.zeros((problems, 1), tl.float32) - grad_t1_value = zero - grad_t1_tangent = zero - grad_t2_value = zero - grad_t2_tangent = zero - grad_m0_value = zero - grad_m0_tangent = zero - grad_b1_value = zero - grad_b1_tangent = zero - grad_inversion_value = zero - grad_inversion_tangent = zero - grad_damping_value = zero - grad_damping_tangent = zero - - for reverse in range(0, event_count): - event = event_count - 1 - reverse - slot = trajectory + event * record_stride - entry_pv = tl.load(trajectory_value + slot, mask=state_mask, other=0.0) - entry_mv = tl.load( - trajectory_value + slot + minus_plane, mask=state_mask, other=0.0 - ) - entry_zv = tl.load( - trajectory_value + slot + long_plane, mask=state_mask, other=0.0 - ) - entry_pt = tl.load(trajectory_tangent + slot, mask=state_mask, other=0.0) - entry_mt = tl.load( - trajectory_tangent + slot + minus_plane, mask=state_mask, other=0.0 - ) - entry_zt = tl.load( - trajectory_tangent + slot + long_plane, mask=state_mask, other=0.0 - ) - - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event) - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - dt_tangent = _event_value( - dot_duration, event_base, event, active_atom, single_train - ) - e1_value = tl.exp(-rate1_value * dt_value) - e1_tangent = -e1_value * (rate1_value * dt_tangent + rate1_tangent * dt_value) - e2_value = tl.exp(-rate2_value * dt_value) - e2_tangent = -e2_value * (rate2_value * dt_tangent + rate2_tangent * dt_value) - damp_z = 1.0 - damp_z_tangent = 0.0 - damp_t = 1.0 - damp_t_tangent = 0.0 - if diffusing: - damp_z, damp_z_tangent, damp_t, damp_t_tangent = _damping_jvp( - atom_damping, atom_dot_damping, dt_value, dt_tangent, order - ) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value, recovery_tangent = 1.0 - e1_value, -e1_tangent - bare1_value, bare1_tangent = e1_value, e1_tangent - bare2_value, bare2_tangent = e2_value, e2_tangent - e1_tangent = e1_tangent * damp_z + bare1_value * damp_z_tangent - e1_value = bare1_value * damp_z - e2_tangent = e2_tangent * damp_t + bare2_value * damp_t_tangent - e2_value = bare2_value * damp_t - - # Replay the intra-event stages from the recorded entry state. - stage_pv = entry_pv * e2_value - stage_pt = entry_pv * e2_tangent + entry_pt * e2_value - stage_mv = entry_mv * e2_value - stage_mt = entry_mv * e2_tangent + entry_mt * e2_value - stage_zv = entry_zv * e1_value + tl.where(state == 0, recovery_value, 0.0) - stage_zt = entry_zv * e1_tangent + entry_zt * e1_value - stage_zt += tl.where(state == 0, recovery_tangent, 0.0) - - pre_shift = (event_action & 1) != 0 - shifted_pv, shifted_mv = _shift_real( - stage_pv, stage_mv, state, state_mask, state_count - ) - shifted_pt, shifted_mt = _shift_real( - stage_pt, stage_mt, state, state_mask, state_count - ) - stage_pv = tl.where(pre_shift, shifted_pv, stage_pv) - stage_mv = tl.where(pre_shift, shifted_mv, stage_mv) - stage_pt = tl.where(pre_shift, shifted_pt, stage_pt) - stage_mt = tl.where(pre_shift, shifted_mt, stage_mt) - - # Undo the trailing spoil or shift. - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - spoil = (event_action & 8) != 0 - adjoint_pv, adjoint_mv = _shift_real_adjoint( - plus_bar_value, minus_bar_value, state, state_mask, state_count - ) - adjoint_pt, adjoint_mt = _shift_real_adjoint( - plus_bar_tangent, minus_bar_tangent, state, state_mask, state_count - ) - trailing = do_shift & ~spoil - plus_bar_value = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_pv, plus_bar_value) - ) - minus_bar_value = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_mv, minus_bar_value) - ) - plus_bar_tangent = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_pt, plus_bar_tangent) - ) - minus_bar_tangent = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_mt, minus_bar_tangent) - ) - - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - inversion_gain = -tl.sum( - tl.where(invert, long_bar_value * stage_zv, 0.0), axis=1 - )[:, None] - inversion_gain_tangent = -tl.sum( - tl.where( - invert, - long_bar_value * stage_zt + long_bar_tangent * stage_zv, - 0.0, - ), - axis=1, - )[:, None] - grad_inversion_value += inversion_gain - grad_inversion_tangent += inversion_gain_tangent - inverted_bar_value = -atom_inversion * long_bar_value - inverted_bar_tangent = ( - -atom_inversion * long_bar_tangent - atom_dot_inversion * long_bar_value - ) - long_bar_value = tl.where(invert, inverted_bar_value, long_bar_value) - long_bar_tangent = tl.where(invert, inverted_bar_tangent, long_bar_tangent) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - event_dot_flip = _event_value( - dot_flip, event_base, event, active_atom, single_train - ) - pulse_b1 = atom_b1 - pulse_dot_b1 = atom_dot_b1 - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse the row of the shim it drives. - if shimmed: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + shim_row + atom, mask=active_atom, other=1.0) - pulse_dot_b1 = tl.load( - dot_b1 + shim_row + atom, mask=active_atom, other=0.0 - ) - alpha_value = event_flip * pulse_b1 - alpha_tangent = event_dot_flip * pulse_b1 + event_flip * pulse_dot_b1 - cosine_value = tl.cos(alpha_value) - sine_value = tl.sin(alpha_value) - cosine_tangent = -sine_value * alpha_tangent - sine_tangent = cosine_value * alpha_tangent - chs_value = 0.5 * (1.0 + cosine_value) - chs_tangent = 0.5 * cosine_tangent - shs_value = 0.5 * (1.0 - cosine_value) - shs_tangent = -0.5 * cosine_tangent - half_sine_value = 0.5 * sine_value - half_sine_tangent = 0.5 * sine_tangent - - # d/dalpha of each output row, contracted with the adjoint. - row_p_value = half_sine_value * stage_mv - half_sine_value * stage_pv - row_p_value -= cosine_value * stage_zv - row_p_tangent = half_sine_value * stage_mt + half_sine_tangent * stage_mv - row_p_tangent -= half_sine_value * stage_pt + half_sine_tangent * stage_pv - row_p_tangent -= cosine_value * stage_zt + cosine_tangent * stage_zv - row_m_value = half_sine_value * stage_pv - half_sine_value * stage_mv - row_m_value += cosine_value * stage_zv - row_m_tangent = half_sine_value * stage_pt + half_sine_tangent * stage_pv - row_m_tangent -= half_sine_value * stage_mt + half_sine_tangent * stage_mv - row_m_tangent += cosine_value * stage_zt + cosine_tangent * stage_zv - row_z_value = 0.5 * cosine_value * stage_pv - 0.5 * cosine_value * stage_mv - row_z_value -= sine_value * stage_zv - row_z_tangent = 0.5 * (cosine_value * stage_pt + cosine_tangent * stage_pv) - row_z_tangent -= 0.5 * (cosine_value * stage_mt + cosine_tangent * stage_mv) - row_z_tangent -= sine_value * stage_zt + sine_tangent * stage_zv - - alpha_bar_terms_value = plus_bar_value * row_p_value - alpha_bar_terms_value += minus_bar_value * row_m_value - alpha_bar_terms_value += long_bar_value * row_z_value - alpha_bar_terms_tangent = plus_bar_value * row_p_tangent - alpha_bar_terms_tangent += plus_bar_tangent * row_p_value - alpha_bar_terms_tangent += minus_bar_value * row_m_tangent - alpha_bar_terms_tangent += minus_bar_tangent * row_m_value - alpha_bar_terms_tangent += long_bar_value * row_z_tangent - alpha_bar_terms_tangent += long_bar_tangent * row_z_value - rotate = is_rf & ~is_inversion - grad_alpha_value = tl.sum(tl.where(rotate, alpha_bar_terms_value, 0.0), axis=1)[ - :, None - ] - grad_alpha_tangent = tl.sum( - tl.where(rotate, alpha_bar_terms_tangent, 0.0), axis=1 - )[:, None] - - # Transpose of the rotation. - rotated_pbv = chs_value * plus_bar_value + shs_value * minus_bar_value - rotated_pbv += half_sine_value * long_bar_value - rotated_pbt = chs_value * plus_bar_tangent + chs_tangent * plus_bar_value - rotated_pbt += shs_value * minus_bar_tangent + shs_tangent * minus_bar_value - rotated_pbt += half_sine_value * long_bar_tangent - rotated_pbt += half_sine_tangent * long_bar_value - rotated_mbv = shs_value * plus_bar_value + chs_value * minus_bar_value - rotated_mbv -= half_sine_value * long_bar_value - rotated_mbt = shs_value * plus_bar_tangent + shs_tangent * plus_bar_value - rotated_mbt += chs_value * minus_bar_tangent + chs_tangent * minus_bar_value - rotated_mbt -= half_sine_value * long_bar_tangent - rotated_mbt -= half_sine_tangent * long_bar_value - rotated_zbv = -sine_value * plus_bar_value + sine_value * minus_bar_value - rotated_zbv += cosine_value * long_bar_value - rotated_zbt = -sine_value * plus_bar_tangent - sine_tangent * plus_bar_value - rotated_zbt += sine_value * minus_bar_tangent + sine_tangent * minus_bar_value - rotated_zbt += cosine_value * long_bar_tangent + cosine_tangent * long_bar_value - - plus_bar_value = tl.where(rotate, rotated_pbv, plus_bar_value) - plus_bar_tangent = tl.where(rotate, rotated_pbt, plus_bar_tangent) - minus_bar_value = tl.where(rotate, rotated_mbv, minus_bar_value) - minus_bar_tangent = tl.where(rotate, rotated_mbt, minus_bar_tangent) - long_bar_value = tl.where(rotate, rotated_zbv, long_bar_value) - long_bar_tangent = tl.where(rotate, rotated_zbt, long_bar_tangent) - - flip_gain_value = grad_alpha_value * pulse_b1 - flip_gain_tangent = grad_alpha_tangent * pulse_b1 - flip_gain_tangent += grad_alpha_value * pulse_dot_b1 - writes_flip = active_atom & rotate - tl.atomic_add( - grad_flip_value + event_base + event, flip_gain_value, mask=writes_flip - ) - tl.atomic_add( - grad_flip_tangent + event_base + event, flip_gain_tangent, mask=writes_flip - ) - if shimmed: - # A pulse's transmit gradient belongs to the shim it drives, so - # with several it lands in that shim's row rather than in a - # register summed over the whole train. - tl.atomic_add( - grad_tissue_value + 3 * atom_count + shim_row + atom, - grad_alpha_value * event_flip, - mask=writes_flip, - ) - tl.atomic_add( - grad_tissue_tangent + 3 * atom_count + shim_row + atom, - grad_alpha_tangent * event_flip + grad_alpha_value * event_dot_flip, - mask=writes_flip, - ) - else: - grad_b1_value += tl.where(rotate, grad_alpha_value * event_flip, 0.0) - grad_b1_tangent += tl.where( - rotate, - grad_alpha_tangent * event_flip + grad_alpha_value * event_dot_flip, - 0.0, - ) - - # The sample is i * m0 * plus[0]; only the imaginary seed acts. - record = ((event_action & 32) != 0) & (event_kind == 2) - out = tl.load(output_index + event) - seed = tl.load( - grad_output_imag + problem * output_count + out, - mask=active_atom & record & (out >= 0), - other=0.0, - ) - grad_m0_value += tl.sum(tl.where(state == 0, seed * stage_pv, 0.0), axis=1)[ - :, None - ] - grad_m0_tangent += tl.sum(tl.where(state == 0, seed * stage_pt, 0.0), axis=1)[ - :, None - ] - plus_bar_value += tl.where(state == 0, seed * atom_m0, 0.0) - plus_bar_tangent += tl.where(state == 0, seed * atom_dot_m0, 0.0) - - adjoint_pv, adjoint_mv = _shift_real_adjoint( - plus_bar_value, minus_bar_value, state, state_mask, state_count - ) - adjoint_pt, adjoint_mt = _shift_real_adjoint( - plus_bar_tangent, minus_bar_tangent, state, state_mask, state_count - ) - plus_bar_value = tl.where(pre_shift, adjoint_pv, plus_bar_value) - minus_bar_value = tl.where(pre_shift, adjoint_mv, minus_bar_value) - plus_bar_tangent = tl.where(pre_shift, adjoint_pt, plus_bar_tangent) - minus_bar_tangent = tl.where(pre_shift, adjoint_mt, minus_bar_tangent) - - cot2_value = plus_bar_value * entry_pv + minus_bar_value * entry_mv - cot2_tangent = ( - plus_bar_value * entry_pt - + plus_bar_tangent * entry_pv - + minus_bar_value * entry_mt - + minus_bar_tangent * entry_mv - ) - cot1_value = long_bar_value * entry_zv - cot1_tangent = long_bar_value * entry_zt + long_bar_tangent * entry_zv - grad_e2_value = tl.sum(cot2_value * damp_t, axis=1)[:, None] - grad_e2_tangent = tl.sum( - cot2_value * damp_t_tangent + cot2_tangent * damp_t, axis=1 - )[:, None] - grad_e1_value = tl.sum(cot1_value * damp_z, axis=1)[:, None] - grad_e1_value -= tl.sum(tl.where(state == 0, long_bar_value, 0.0), axis=1)[ - :, None - ] - grad_e1_tangent = tl.sum( - cot1_value * damp_z_tangent + cot1_tangent * damp_z, axis=1 - )[:, None] - grad_e1_tangent -= tl.sum(tl.where(state == 0, long_bar_tangent, 0.0), axis=1)[ - :, None - ] - - # The rate and the interval multiply every order's b-weight, so both - # take a weighted sum. Order zero has no longitudinal weight, which - # keeps recovery out of this. - spread_value = zero - spread_tangent = zero - if diffusing: - weighted_value = ( - cot1_value * bare1_value * damp_z * longitudinal_weight - + cot2_value * bare2_value * damp_t * transverse_weight - ) - weighted_tangent = ( - cot1_tangent * bare1_value * damp_z - + cot1_value * bare1_tangent * damp_z - + cot1_value * bare1_value * damp_z_tangent - ) * longitudinal_weight + ( - cot2_tangent * bare2_value * damp_t - + cot2_value * bare2_tangent * damp_t - + cot2_value * bare2_value * damp_t_tangent - ) * transverse_weight - spread_value = tl.sum(weighted_value, axis=1)[:, None] - spread_tangent = tl.sum(weighted_tangent, axis=1)[:, None] - grad_damping_value += -spread_value * dt_value - grad_damping_tangent += -( - spread_value * dt_tangent + spread_tangent * dt_value - ) - - plus_bar_tangent = plus_bar_value * e2_tangent + plus_bar_tangent * e2_value - plus_bar_value = plus_bar_value * e2_value - minus_bar_tangent = minus_bar_value * e2_tangent + minus_bar_tangent * e2_value - minus_bar_value = minus_bar_value * e2_value - long_bar_tangent = long_bar_value * e1_tangent + long_bar_tangent * e1_value - long_bar_value = long_bar_value * e1_value - - inverse1_value = 1000.0 / (atom_t1 * atom_t1) - inverse1_tangent = -2000.0 * atom_dot_t1 / (atom_t1 * atom_t1 * atom_t1) - inverse2_value = 1000.0 / (atom_t2 * atom_t2) - inverse2_tangent = -2000.0 * atom_dot_t2 / (atom_t2 * atom_t2 * atom_t2) - scale1_value = bare1_value * dt_value * inverse1_value - scale1_tangent = bare1_tangent * dt_value * inverse1_value - scale1_tangent += bare1_value * dt_tangent * inverse1_value - scale1_tangent += bare1_value * dt_value * inverse1_tangent - scale2_value = bare2_value * dt_value * inverse2_value - scale2_tangent = bare2_tangent * dt_value * inverse2_value - scale2_tangent += bare2_value * dt_tangent * inverse2_value - scale2_tangent += bare2_value * dt_value * inverse2_tangent - grad_t1_value += grad_e1_value * scale1_value - grad_t1_tangent += grad_e1_value * scale1_tangent - grad_t1_tangent += grad_e1_tangent * scale1_value - grad_t2_value += grad_e2_value * scale2_value - grad_t2_tangent += grad_e2_value * scale2_tangent - grad_t2_tangent += grad_e2_tangent * scale2_value - - decay1_value = rate1_value * bare1_value - decay1_tangent = rate1_value * bare1_tangent + rate1_tangent * bare1_value - decay2_value = rate2_value * bare2_value - decay2_tangent = rate2_value * bare2_tangent + rate2_tangent * bare2_value - duration_gain_value = -grad_e1_value * decay1_value - duration_gain_value -= grad_e2_value * decay2_value - duration_gain_tangent = -( - grad_e1_value * decay1_tangent + grad_e1_tangent * decay1_value - ) - duration_gain_tangent -= ( - grad_e2_value * decay2_tangent + grad_e2_tangent * decay2_value - ) - duration_gain_value += -spread_value * atom_damping - duration_gain_tangent += -( - spread_value * atom_dot_damping + spread_tangent * atom_damping - ) - tl.atomic_add( - grad_duration_value + event_base + event, - duration_gain_value, - mask=active_atom, - ) - tl.atomic_add( - grad_duration_tangent + event_base + event, - duration_gain_tangent, - mask=active_atom, - ) - - tl.atomic_add(grad_tissue_value + atom, grad_t1_value, mask=active_atom) - tl.atomic_add(grad_tissue_tangent + atom, grad_t1_tangent, mask=active_atom) - tl.atomic_add( - grad_tissue_value + atom_count + atom, grad_t2_value, mask=active_atom - ) - tl.atomic_add( - grad_tissue_tangent + atom_count + atom, grad_t2_tangent, mask=active_atom - ) - tl.atomic_add( - grad_tissue_value + 2 * atom_count + atom, grad_m0_value, mask=active_atom - ) - tl.atomic_add( - grad_tissue_tangent + 2 * atom_count + atom, grad_m0_tangent, mask=active_atom - ) - if not shimmed: - tl.atomic_add( - grad_tissue_value + 3 * atom_count + atom, - grad_b1_value, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + 3 * atom_count + atom, - grad_b1_tangent, - mask=active_atom, - ) - # The transmit pair takes a row per shim each in the plane the complex - # path allocates, so the rows past it move even though this kernel leaves - # the transmit phase at zero throughout. - past_transmit = 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue_value + (6 + past_transmit) * atom_count + atom, - grad_inversion_value, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (6 + past_transmit) * atom_count + atom, - grad_inversion_tangent, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_value + (7 + past_transmit) * atom_count + atom, - grad_damping_value, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue_tangent + (7 + past_transmit) * atom_count + atom, - grad_damping_tangent, - mask=active_atom, - ) - - -@triton.jit(do_not_specialize=["state_count"]) -def _epg_real_vjp_kernel( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - grad_output_imag, - grad_tissue, - grad_flip, - grad_duration, - trajectory_value, - problem_base, - problem_end, - atom_count, - train_count, - event_count, - output_count, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shim_rows, - shimmed: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = problem_base + tl.program_id(0) * problems - problem = problem + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - # The grid rounds up to whole tiles, so the last program of a wave reaches - # past it. Those problems are real, but their trajectory rows belong to a - # later launch and do not exist yet. - active_atom = problem < problem_end - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - # The trajectory holds the state entering every event: three planes of - # configuration orders. - record_stride = 3 * state_count - trajectory = (problem - problem_base) * event_count * record_stride + state - minus_plane = state_count - long_plane = 2 * state_count - - empty = tl.zeros((problems, block_states), tl.float32) - plus_value = empty - minus_value = empty - long_value = empty + tl.where(state == 0, 1.0, 0.0) - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - atom_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - order = state.to(tl.float32) - longitudinal_weight = order * order - transverse_weight = longitudinal_weight + order + 0.3333333333333333 - rate1_value = 1000.0 / atom_t1 - rate2_value = 1000.0 / atom_t2 - - event_base = train * event_count - for event in range(0, event_count): - slot = trajectory + event * record_stride - tl.store(trajectory_value + slot, plus_value, mask=state_mask) - tl.store(trajectory_value + slot + minus_plane, minus_value, mask=state_mask) - tl.store(trajectory_value + slot + long_plane, long_value, mask=state_mask) - - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - e1_value = tl.exp(-rate1_value * dt_value) - e2_value = tl.exp(-rate2_value * dt_value) - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt_value, order) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value = 1.0 - e1_value - bare1_value = e1_value - bare2_value = e2_value - e1_value = bare1_value * damp_z - e2_value = bare2_value * damp_t - - plus_value = plus_value * e2_value - minus_value = minus_value * e2_value - long_value = long_value * e1_value - long_value += tl.where(state == 0, recovery_value, 0.0) - - event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - shifted_pv, shifted_mv = _shift_real( - plus_value, minus_value, state, state_mask, state_count - ) - plus_value = tl.where(pre_shift, shifted_pv, plus_value) - minus_value = tl.where(pre_shift, shifted_mv, minus_value) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - long_value = tl.where(invert, -atom_inversion * long_value, long_value) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - pulse_b1 = atom_b1 - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse the row of the shim it drives. - if shimmed: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + shim_row + atom, mask=active_atom, other=1.0) - alpha_value = event_flip * pulse_b1 - cosine_value = tl.cos(alpha_value) - sine_value = tl.sin(alpha_value) - chs_value = 0.5 * (1.0 + cosine_value) - shs_value = 0.5 * (1.0 - cosine_value) - half_sine_value = 0.5 * sine_value - - rotated_pv = chs_value * plus_value + shs_value * minus_value - rotated_pv -= sine_value * long_value - rotated_mv = shs_value * plus_value + chs_value * minus_value - rotated_mv += sine_value * long_value - rotated_zv = half_sine_value * plus_value - half_sine_value * minus_value - rotated_zv += cosine_value * long_value - - rotate = is_rf & ~is_inversion - plus_value = tl.where(rotate, rotated_pv, plus_value) - minus_value = tl.where(rotate, rotated_mv, minus_value) - long_value = tl.where(rotate, rotated_zv, long_value) - - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - shifted_pv, shifted_mv = _shift_real( - plus_value, minus_value, state, state_mask, state_count - ) - plus_value = tl.where(do_shift, shifted_pv, plus_value) - minus_value = tl.where(do_shift, shifted_mv, minus_value) - spoil = (event_action & 8) != 0 - plus_value = tl.where(spoil, 0.0, plus_value) - minus_value = tl.where(spoil, 0.0, minus_value) - - plus_bar_value = empty - minus_bar_value = empty - long_bar_value = empty - zero = tl.zeros((problems, 1), tl.float32) - grad_t1_value = zero - grad_t2_value = zero - grad_m0_value = zero - grad_b1_value = zero - grad_inversion_value = zero - grad_damping_value = zero - - for reverse in range(0, event_count): - event = event_count - 1 - reverse - slot = trajectory + event * record_stride - entry_pv = tl.load(trajectory_value + slot, mask=state_mask, other=0.0) - entry_mv = tl.load( - trajectory_value + slot + minus_plane, mask=state_mask, other=0.0 - ) - entry_zv = tl.load( - trajectory_value + slot + long_plane, mask=state_mask, other=0.0 - ) - - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event) - dt_value = _event_value(duration, event_base, event, active_atom, single_train) - e1_value = tl.exp(-rate1_value * dt_value) - e2_value = tl.exp(-rate2_value * dt_value) - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt_value, order) - # Order zero is undamped, so recovery keeps the bare longitudinal factor. - recovery_value = 1.0 - e1_value - bare1_value = e1_value - bare2_value = e2_value - e1_value = bare1_value * damp_z - e2_value = bare2_value * damp_t - - # Replay the intra-event stages from the recorded entry state. - stage_pv = entry_pv * e2_value - stage_mv = entry_mv * e2_value - stage_zv = entry_zv * e1_value + tl.where(state == 0, recovery_value, 0.0) - - pre_shift = (event_action & 1) != 0 - shifted_pv, shifted_mv = _shift_real( - stage_pv, stage_mv, state, state_mask, state_count - ) - stage_pv = tl.where(pre_shift, shifted_pv, stage_pv) - stage_mv = tl.where(pre_shift, shifted_mv, stage_mv) - - # Undo the trailing spoil or shift. - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - spoil = (event_action & 8) != 0 - adjoint_pv, adjoint_mv = _shift_real_adjoint( - plus_bar_value, minus_bar_value, state, state_mask, state_count - ) - trailing = do_shift & ~spoil - plus_bar_value = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_pv, plus_bar_value) - ) - minus_bar_value = tl.where( - spoil, 0.0, tl.where(trailing, adjoint_mv, minus_bar_value) - ) - - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - grad_inversion_value += -tl.sum( - tl.where(invert, long_bar_value * stage_zv, 0.0), axis=1 - )[:, None] - long_bar_value = tl.where( - invert, -atom_inversion * long_bar_value, long_bar_value - ) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - pulse_b1 = atom_b1 - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse the row of the shim it drives. - if shimmed: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load(b1 + shim_row + atom, mask=active_atom, other=1.0) - alpha_value = event_flip * pulse_b1 - cosine_value = tl.cos(alpha_value) - sine_value = tl.sin(alpha_value) - chs_value = 0.5 * (1.0 + cosine_value) - shs_value = 0.5 * (1.0 - cosine_value) - half_sine_value = 0.5 * sine_value - - # d/dalpha of each output row, contracted with the adjoint. - row_p_value = half_sine_value * stage_mv - half_sine_value * stage_pv - row_p_value -= cosine_value * stage_zv - row_m_value = half_sine_value * stage_pv - half_sine_value * stage_mv - row_m_value += cosine_value * stage_zv - row_z_value = 0.5 * cosine_value * stage_pv - 0.5 * cosine_value * stage_mv - row_z_value -= sine_value * stage_zv - - alpha_bar_terms_value = plus_bar_value * row_p_value - alpha_bar_terms_value += minus_bar_value * row_m_value - alpha_bar_terms_value += long_bar_value * row_z_value - rotate = is_rf & ~is_inversion - grad_alpha_value = tl.sum(tl.where(rotate, alpha_bar_terms_value, 0.0), axis=1)[ - :, None - ] - - # Transpose of the rotation. - rotated_pbv = chs_value * plus_bar_value + shs_value * minus_bar_value - rotated_pbv += half_sine_value * long_bar_value - rotated_mbv = shs_value * plus_bar_value + chs_value * minus_bar_value - rotated_mbv -= half_sine_value * long_bar_value - rotated_zbv = -sine_value * plus_bar_value + sine_value * minus_bar_value - rotated_zbv += cosine_value * long_bar_value - - plus_bar_value = tl.where(rotate, rotated_pbv, plus_bar_value) - minus_bar_value = tl.where(rotate, rotated_mbv, minus_bar_value) - long_bar_value = tl.where(rotate, rotated_zbv, long_bar_value) - - writes_flip = active_atom & rotate - tl.atomic_add( - grad_flip + event_base + event, - grad_alpha_value * pulse_b1, - mask=writes_flip, - ) - if shimmed: - # A pulse's transmit gradient belongs to the shim it drives, so - # with several it lands in that shim's row rather than in a - # register summed over the whole train. - tl.atomic_add( - grad_tissue + 3 * atom_count + shim_row + atom, - grad_alpha_value * event_flip, - mask=writes_flip, - ) - else: - grad_b1_value += tl.where(rotate, grad_alpha_value * event_flip, 0.0) - - # The sample is i * m0 * plus[0]; only the imaginary seed acts. - record = ((event_action & 32) != 0) & (event_kind == 2) - out = tl.load(output_index + event) - seed = tl.load( - grad_output_imag + problem * output_count + out, - mask=active_atom & record & (out >= 0), - other=0.0, - ) - grad_m0_value += tl.sum(tl.where(state == 0, seed * stage_pv, 0.0), axis=1)[ - :, None - ] - plus_bar_value += tl.where(state == 0, seed * atom_m0, 0.0) - - adjoint_pv, adjoint_mv = _shift_real_adjoint( - plus_bar_value, minus_bar_value, state, state_mask, state_count - ) - plus_bar_value = tl.where(pre_shift, adjoint_pv, plus_bar_value) - minus_bar_value = tl.where(pre_shift, adjoint_mv, minus_bar_value) - - cot2_value = plus_bar_value * entry_pv + minus_bar_value * entry_mv - cot1_value = long_bar_value * entry_zv - grad_e2_value = tl.sum(cot2_value * damp_t, axis=1)[:, None] - grad_e1_value = tl.sum(cot1_value * damp_z, axis=1)[:, None] - grad_e1_value -= tl.sum(tl.where(state == 0, long_bar_value, 0.0), axis=1)[ - :, None - ] - - # The rate and the interval multiply every order's b-weight, so both - # take a weighted sum. Order zero has no longitudinal weight, which - # keeps recovery out of this. - spread_value = zero - if diffusing: - weighted_value = ( - cot1_value * bare1_value * damp_z * longitudinal_weight - + cot2_value * bare2_value * damp_t * transverse_weight - ) - spread_value = tl.sum(weighted_value, axis=1)[:, None] - grad_damping_value += -spread_value * dt_value - - plus_bar_value = plus_bar_value * e2_value - minus_bar_value = minus_bar_value * e2_value - long_bar_value = long_bar_value * e1_value - - inverse1_value = 1000.0 / (atom_t1 * atom_t1) - inverse2_value = 1000.0 / (atom_t2 * atom_t2) - grad_t1_value += grad_e1_value * (bare1_value * dt_value * inverse1_value) - grad_t2_value += grad_e2_value * (bare2_value * dt_value * inverse2_value) - - duration_gain_value = -grad_e1_value * (rate1_value * bare1_value) - duration_gain_value -= grad_e2_value * (rate2_value * bare2_value) - duration_gain_value += -spread_value * atom_damping - tl.atomic_add( - grad_duration + event_base + event, - duration_gain_value, - mask=active_atom, - ) - - tl.atomic_add(grad_tissue + atom, grad_t1_value, mask=active_atom) - tl.atomic_add(grad_tissue + atom_count + atom, grad_t2_value, mask=active_atom) - tl.atomic_add(grad_tissue + 2 * atom_count + atom, grad_m0_value, mask=active_atom) - if not shimmed: - tl.atomic_add( - grad_tissue + 3 * atom_count + atom, grad_b1_value, mask=active_atom - ) - # The transmit pair takes a row per shim each in the plane the complex - # path allocates, so the rows past it move even though this kernel leaves - # the transmit phase at zero throughout. - past_transmit = 2 * (shim_rows - 1) - tl.atomic_add( - grad_tissue + (6 + past_transmit) * atom_count + atom, - grad_inversion_value, - mask=active_atom, - ) - tl.atomic_add( - grad_tissue + (7 + past_transmit) * atom_count + atom, - grad_damping_value, - mask=active_atom, - ) - - -@triton.jit(do_not_specialize=["state_count"]) -def _epg_real_kernel( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - output_real, - output_imag, - atom_count, - train_count, - event_count, - output_count, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shimmed: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = tl.program_id(0) * problems + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < train_count * atom_count - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - - empty = tl.zeros((problems, block_states), tl.float32) - plus = empty - minus = empty - longitudinal = empty + tl.where(state == 0, 1.0, 0.0) - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - rate1 = 1000.0 / atom_t1 - rate2 = 1000.0 / atom_t2 - atom_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - order = state.to(tl.float32) - - # The relaxation factors depend on the event only through its duration, and - # a train repeats its intervals: an interval as long as the last one reuses - # the factors rather than taking the two exponentials again. Where several - # trains share the program the durations differ across its lanes and there - # is nothing uniform to compare, so only a single-train launch memoizes. - last_dt = -1.0 - e1 = rate1 * 0.0 + 1.0 - e2 = rate2 * 0.0 + 1.0 - - event_base = train * event_count - # Two events to an iteration. A repetition is several events -- a pulse, - # a sample, an interval -- so the loop runs longer than the sequence is - # repetitions, and unrolling lets one back-edge and one set of event - # bookkeeping serve two of them. Two is where it stops paying: four was - # measured slower, and the body is already large enough that widening it - # costs registers. - for event in tl.range(0, event_count, loop_unroll_factor=2): - # Read here rather than through the helper: one train gives a duration - # the whole program shares, and the skip and the memo below both want - # to compare it as the single number it is. - if single_train: - dt = tl.load(duration + event) - else: - dt = tl.load(duration + event_base + event, mask=active_atom, other=0.0) - # An event of no duration relaxes nothing: both factors are one and the - # recovery term is zero. Half the events of a spoiled repetition are - # instantaneous, and reducing over the trains this program carries makes - # that a branch the whole program agrees on rather than a tile of - # multiplies by one. - if single_train: - relaxes = dt != 0.0 - else: - relaxes = tl.max(dt) != 0.0 - if relaxes: - if single_train: - if dt != last_dt: - e1 = tl.exp(-rate1 * dt) - e2 = tl.exp(-rate2 * dt) - last_dt = dt - else: - e1 = tl.exp(-rate1 * dt) - e2 = tl.exp(-rate2 * dt) - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt, order) - recovery = 1.0 - e1 - plus *= e2 * damp_t - minus *= e2 * damp_t - longitudinal = longitudinal * (e1 * damp_z) + tl.where( - state == 0, recovery, 0.0 - ) - - # Every flag below is read from a per-event array with no atom index, so - # it is uniform across the program and can steer real control flow. A - # `tl.where` would make every event pay for every operator: a spoiled - # repetition is four events and needs one rotation and one shift. - event_action = tl.load(action + event).to(tl.int32) - if (event_action & 1) != 0: - plus, minus = _shift_real(plus, minus, state, state_mask, state_count) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - if is_rf and is_inversion: - longitudinal = -atom_inversion * longitudinal - elif is_rf: - alpha = _event_value(flip, event_base, event, active_atom, single_train) - pulse_b1 = atom_b1 - # One shim is the whole sequence's transmit field, loaded once - # above; several give each pulse the row of the shim it drives. - if shimmed and transmit: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - pulse_b1 = tl.load(b1 + shim_row + atom, mask=active_atom, other=1.0) - alpha *= pulse_b1 - cosine = tl.cos(alpha) - sine = tl.sin(alpha) - cosine_half_sq = 0.5 * (1.0 + cosine) - sine_half_sq = 0.5 * (1.0 - cosine) - half_sine = 0.5 * sine - rotated_p = ( - cosine_half_sq * plus + sine_half_sq * minus - sine * longitudinal - ) - rotated_m = ( - sine_half_sq * plus + cosine_half_sq * minus + sine * longitudinal - ) - longitudinal = half_sine * plus - half_sine * minus + cosine * longitudinal - plus = rotated_p - minus = rotated_m - - if ((event_action & 32) != 0) and event_kind == 2: - out = tl.load(output_index + event) - output_offset = problem * output_count + out - output_mask = active_atom & (state == 0) & (out >= 0) - tl.store(output_real + output_offset + state, empty, mask=output_mask) - tl.store( - output_imag + output_offset + state, atom_m0 * plus, mask=output_mask - ) - - if ((event_action & 2) != 0) or ((event_action & 16) != 0): - plus, minus = _shift_real(plus, minus, state, state_mask, state_count) - if (event_action & 8) != 0: - plus = empty - minus = empty - - -@triton.jit(do_not_specialize=["state_count"]) -def _epg_real_jvp_kernel( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - tangent_t1, - tangent_t2, - tangent_m0, - tangent_b1, - tangent_inversion_efficiency, - tangent_diffusion, - tangent_duration, - tangent_flip, - output_real, - output_imag, - atom_count, - train_count, - event_count, - output_count, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shimmed: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = tl.program_id(0) * problems + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < train_count * atom_count - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - - empty = tl.zeros((problems, block_states), tl.float32) - plus = empty - minus = empty - longitudinal = empty + tl.where(state == 0, 1.0, 0.0) - dot_plus = empty - dot_minus = empty - dot_longitudinal = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - dot_t1 = tl.load(tangent_t1 + atom, mask=active_atom, other=0.0) - dot_t2 = tl.load(tangent_t2 + atom, mask=active_atom, other=0.0) - dot_m0 = 0.0 - if density: - dot_m0 = tl.load(tangent_m0 + scalar_atom, mask=active_atom, other=0.0) - dot_b1 = 0.0 - if transmit: - dot_b1 = tl.load(tangent_b1 + scalar_atom, mask=active_atom, other=0.0) - dot_inversion = 0.0 - if inverting: - dot_inversion = tl.load( - tangent_inversion_efficiency + scalar_atom, mask=active_atom, other=0.0 - ) - rate1 = 1000.0 / atom_t1 - rate2 = 1000.0 / atom_t2 - atom_damping = 0.0 - dot_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - dot_damping = tl.load( - tangent_diffusion + scalar_atom, mask=active_atom, other=0.0 - ) - order = state.to(tl.float32) - - event_base = train * event_count - # Two events to an iteration. A repetition is several events -- a pulse, - # a sample, an interval -- so the loop runs longer than the sequence is - # repetitions, and unrolling lets one back-edge and one set of event - # bookkeeping serve two of them. Two is where it stops paying: four was - # measured slower, and the body is already large enough that widening it - # costs registers. - for event in tl.range(0, event_count, loop_unroll_factor=2): - dt = _event_value(duration, event_base, event, active_atom, single_train) - dot_dt = _event_value( - tangent_duration, event_base, event, active_atom, single_train - ) - # An event of no duration relaxes nothing, and carries no tangent along - # the relaxation either: both factors are one and both their derivatives - # are zero. - if tl.max(dt) != 0.0 or tl.max(dot_dt) != 0.0: - e1 = tl.exp(-rate1 * dt) - e2 = tl.exp(-rate2 * dt) - dot_e1 = e1 * (1000.0 * dt * dot_t1 / (atom_t1 * atom_t1) - rate1 * dot_dt) - dot_e2 = e2 * (1000.0 * dt * dot_t2 / (atom_t2 * atom_t2) - rate2 * dot_dt) - damp_z = 1.0 - ddamp_z = 0.0 - damp_t = 1.0 - ddamp_t = 0.0 - if diffusing: - damp_z, ddamp_z, damp_t, ddamp_t = _damping_jvp( - atom_damping, dot_damping, dt, dot_dt, order - ) - # Order zero is undamped, so the recovery keeps the bare factor. - recovery, dot_recovery = 1.0 - e1, -dot_e1 - dot_e1 = dot_e1 * damp_z + e1 * ddamp_z - e1 = e1 * damp_z - dot_e2 = dot_e2 * damp_t + e2 * ddamp_t - e2 = e2 * damp_t - dot_plus = dot_plus * e2 + plus * dot_e2 - dot_minus = dot_minus * e2 + minus * dot_e2 - dot_longitudinal = dot_longitudinal * e1 + longitudinal * dot_e1 - dot_longitudinal += tl.where(state == 0, dot_recovery, 0.0) - plus *= e2 - minus *= e2 - longitudinal = longitudinal * e1 + tl.where(state == 0, recovery, 0.0) - - # Every flag below is read from a per-event array with no atom index, so - # it is uniform across the program and can steer real control flow. - event_action = tl.load(action + event).to(tl.int32) - if (event_action & 1) != 0: - plus, minus = _shift_real(plus, minus, state, state_mask, state_count) - dot_plus, dot_minus = _shift_real( - dot_plus, dot_minus, state, state_mask, state_count - ) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - if is_rf and is_inversion: - dot_longitudinal = ( - -atom_inversion * dot_longitudinal - dot_inversion * longitudinal - ) - longitudinal = -atom_inversion * longitudinal - elif is_rf: - event_flip = _event_value( - flip, event_base, event, active_atom, single_train - ) - dot_flip = _event_value( - tangent_flip, event_base, event, active_atom, single_train - ) - pulse_b1 = atom_b1 - pulse_dot_b1 = dot_b1 - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse the row of the shim it drives. - if shimmed: - shim_row = tl.load(shim_index + event).to(tl.int64) * atom_count - if transmit: - pulse_b1 = tl.load( - b1 + shim_row + atom, mask=active_atom, other=1.0 - ) - pulse_dot_b1 = tl.load( - tangent_b1 + shim_row + atom, mask=active_atom, other=0.0 - ) - alpha = event_flip * pulse_b1 - dot_alpha = dot_flip * pulse_b1 + event_flip * pulse_dot_b1 - cosine = tl.cos(alpha) - sine = tl.sin(alpha) - cosine_half_sq = 0.5 * (1.0 + cosine) - sine_half_sq = 0.5 * (1.0 - cosine) - half_sine = 0.5 * sine - dot_cosine = -sine * dot_alpha - dot_sine = cosine * dot_alpha - dot_cosine_half_sq = -0.5 * sine * dot_alpha - dot_sine_half_sq = 0.5 * sine * dot_alpha - dot_half_sine = 0.5 * cosine * dot_alpha - - rotated_dp = cosine_half_sq * dot_plus + dot_cosine_half_sq * plus - rotated_dp += sine_half_sq * dot_minus + dot_sine_half_sq * minus - rotated_dp -= sine * dot_longitudinal + dot_sine * longitudinal - rotated_dm = sine_half_sq * dot_plus + dot_sine_half_sq * plus - rotated_dm += cosine_half_sq * dot_minus + dot_cosine_half_sq * minus - rotated_dm += sine * dot_longitudinal + dot_sine * longitudinal - rotated_dz = half_sine * dot_plus + dot_half_sine * plus - rotated_dz -= half_sine * dot_minus + dot_half_sine * minus - rotated_dz += cosine * dot_longitudinal + dot_cosine * longitudinal - rotated_p = ( - cosine_half_sq * plus + sine_half_sq * minus - sine * longitudinal - ) - rotated_m = ( - sine_half_sq * plus + cosine_half_sq * minus + sine * longitudinal - ) - rotated_z = half_sine * plus - half_sine * minus + cosine * longitudinal - - plus = rotated_p - minus = rotated_m - longitudinal = rotated_z - dot_plus = rotated_dp - dot_minus = rotated_dm - dot_longitudinal = rotated_dz - - if ((event_action & 32) != 0) and event_kind == 2: - out = tl.load(output_index + event) - output_offset = problem * output_count + out - output_mask = active_atom & (state == 0) & (out >= 0) - signal_imag = dot_m0 * plus + atom_m0 * dot_plus - tl.store(output_real + output_offset + state, empty, mask=output_mask) - tl.store(output_imag + output_offset + state, signal_imag, mask=output_mask) - - if ((event_action & 2) != 0) or ((event_action & 16) != 0): - plus, minus = _shift_real(plus, minus, state, state_mask, state_count) - dot_plus, dot_minus = _shift_real( - dot_plus, dot_minus, state, state_mask, state_count - ) - if (event_action & 8) != 0: - plus = empty - minus = empty - dot_plus = empty - dot_minus = empty - - -@triton.jit -def _rotate_flip_phase( - cosine, - sine, - cos_phi, - sin_phi, - cos_2phi, - sin_2phi, - fp_r, - fp_i, - fm_r, - fm_i, - z_r, - z_i, -): - """One pool through a hard pulse named by its flip angle and phase. - - Pulled out of the kernel body so a second pool can take the same rotation: - a chemical shift moves where a pool precesses, not what a pulse does to it. - """ - cosine_half_sq = 0.5 * (1.0 + cosine) - sine_half_sq = 0.5 * (1.0 - cosine) - half_sine = 0.5 * sine - - # Every sum of products is one fused multiply-add in a fixed order, so a - # kernel compiled for fewer terms rounds the rotation as the full one does. - minus_2phi_r = tl.fma(cos_2phi, fm_r, -(sin_2phi * fm_i)) - minus_2phi_i = tl.fma(sin_2phi, fm_r, cos_2phi * fm_i) - plus_2phi_r = tl.fma(cos_2phi, fp_r, sin_2phi * fp_i) - plus_2phi_i = tl.fma(cos_2phi, fp_i, -(sin_2phi * fp_r)) - z_turn_a = tl.fma(sin_phi, z_r, cos_phi * z_i) - z_turn_b = tl.fma(sin_phi, z_i, -(cos_phi * z_r)) - z_turn_c = tl.fma(sin_phi, z_r, -(cos_phi * z_i)) - z_turn_d = tl.fma(cos_phi, z_r, sin_phi * z_i) - - rotated_pr = tl.fma( - sine, z_turn_a, tl.fma(sine_half_sq, minus_2phi_r, cosine_half_sq * fp_r) - ) - rotated_pi = tl.fma( - sine, z_turn_b, tl.fma(sine_half_sq, minus_2phi_i, cosine_half_sq * fp_i) - ) - rotated_mr = tl.fma( - sine, z_turn_c, tl.fma(cosine_half_sq, fm_r, sine_half_sq * plus_2phi_r) - ) - rotated_mi = tl.fma( - sine, z_turn_d, tl.fma(cosine_half_sq, fm_i, sine_half_sq * plus_2phi_i) - ) - - plus_turn_r = tl.fma(sin_phi, fp_r, -(cos_phi * fp_i)) - minus_turn_r = tl.fma(sin_phi, fm_r, cos_phi * fm_i) - plus_turn_i = tl.fma(cos_phi, fp_r, sin_phi * fp_i) - minus_turn_i = tl.fma(cos_phi, fm_r, -(sin_phi * fm_i)) - rotated_zr = tl.fma( - cosine, z_r, tl.fma(-half_sine, minus_turn_r, -half_sine * plus_turn_r) - ) - rotated_zi = tl.fma( - cosine, z_i, tl.fma(half_sine, minus_turn_i, -half_sine * plus_turn_i) - ) - return ( - rotated_pr, - rotated_pi, - rotated_mr, - rotated_mi, - rotated_zr, - rotated_zi, - ) - - -@triton.jit( - do_not_specialize=["state_count", "locations", "profile_bins", "lineshape_bins"] -) -def _epg_kernel( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - bound_exchange, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - phase_cos, - phase_sin, - action, - output_index, - shim_index, - saturation, - rf_frequency, - profile, - profile_index, - lineshape, - pairs, - pair_index, - duration_row, - pool_table, - output_real, - output_imag, - atom_count, - train_count, - event_count, - output_count, - flow_scale, - washout_scale, - profile_step, - lineshape_step, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shim_rows, - shimmed: tl.constexpr, - locations, - profiled: tl.constexpr, - profile_bins, - dynamic: tl.constexpr, - broadened: tl.constexpr, - lineshape_bins, - pools: tl.constexpr, - narrow: tl.constexpr, - tabulated: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = tl.program_id(0) * problems + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < train_count * atom_count - # A partial block carries lanes with no problem behind them, and they must - # take no part in a reduction or a store. - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - # Voxels are spread over the slice voxel-major, so a voxel's place along - # the slice is its index modulo the profile's width. One pulse shape holds - # that many consecutive rows, and the event says which shape it drives. - location = atom % locations - - empty = tl.zeros((problems, block_states), tl.float32) - fplus_real = empty - fplus_imag = empty - fminus_real = empty - fminus_imag = empty - # A second pool holds its own share of the equilibrium. The semisolid one - # carries longitudinal states alone -- nothing dephases it, so it reaches - # the higher orders only through exchange with the free pool's -- while the - # chemically exchanging one carries a transverse pair of its own. - # - # ``bound`` is whichever second pool the longitudinal step pairs the free - # water with -- the semisolid one when it is the only one, the exchanging - # one otherwise -- and ``semisolid`` is the third, which only a three-pool - # run carries. - atom_bound = 0.0 - atom_exchange = 0.0 - atom_r1_bound = 0.0 - atom_r2_bound = 0.0 - atom_shift = 0.0 - atom_semisolid = 0.0 - atom_semisolid_exchange = 0.0 - atom_r1_semisolid = 0.0 - if pools == 1: - atom_bound = tl.load(bound_fraction + scalar_atom, mask=active_atom, other=0.0) - atom_exchange = tl.load( - bound_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - atom_r1_bound = 1000.0 / tl.load( - t1_bound + scalar_atom, mask=active_atom, other=1.0 - ) - if pools == 2 or pools == 3: - atom_bound = tl.load(pool_b_fraction + scalar_atom, mask=active_atom, other=0.0) - atom_exchange = tl.load( - pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - atom_r1_bound = 1000.0 / tl.load( - t1_pool_b + scalar_atom, mask=active_atom, other=1.0 - ) - atom_r2_bound = 1000.0 / tl.load( - t2_pool_b + scalar_atom, mask=active_atom, other=1.0 - ) - atom_shift = tl.load(pool_b_shift + scalar_atom, mask=active_atom, other=0.0) - if pools == 3: - atom_semisolid = tl.load( - bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_semisolid_exchange = tl.load( - bound_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - atom_r1_semisolid = 1000.0 / tl.load( - t1_bound + scalar_atom, mask=active_atom, other=1.0 - ) - # A semisolid pool holds a share of the voxel without carrying any - # transverse magnetization, so the 2x2 below is blind to it and the - # exchange inside that 2x2 is not. - atom_free = 1.0 - atom_bound - atom_semisolid - longitudinal_real = empty + tl.where(state == 0, atom_free, 0.0) - longitudinal_imag = empty - bound_real = empty + tl.where(state == 0, atom_bound + 0.0, 0.0) - bound_imag = empty - semisolid_real = empty + tl.where(state == 0, atom_semisolid + 0.0, 0.0) - semisolid_imag = empty - bplus_real = empty - bplus_imag = empty - bminus_real = empty - bminus_imag = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_b1_phase = 0.0 - atom_b0 = 0.0 - if off_axis: - atom_b1_phase = tl.load(b1_phase + scalar_atom, mask=active_atom, other=0.0) - atom_b0 = tl.load(b0 + scalar_atom, mask=active_atom, other=0.0) - b1_cos = tl.cos(atom_b1_phase) - b1_sin = tl.sin(atom_b1_phase) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - atom_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - atom_flow = 0.0 - atom_washout = 0.0 - if moving: - atom_velocity = tl.load(velocity + scalar_atom, mask=active_atom, other=0.0) - atom_flow = atom_velocity * flow_scale - atom_washout = tl.abs(atom_velocity) * washout_scale - order = state.to(tl.float32) - - event_base = train * event_count - for event in range(0, event_count): - dt = _event_value(duration, event_base, event, active_atom, single_train) - wout = 1.0 - if moving: - wout = _washout(atom_washout, dt) - e1 = tl.exp(-(1000.0 / atom_t1) * dt) * wout - e2 = tl.exp(-(1000.0 / atom_t2) * dt) * wout - damp_z = 1.0 - damp_t = 1.0 - if diffusing: - damp_z, damp_t = _damping(atom_damping, dt, order) - turn_z = 0.0 - turn_t = 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, dt, order) - recovery = 1.0 - e1 - e1 = e1 * damp_z - e2 = e2 * damp_t - off_cos = 1.0 - off_sin = 0.0 - if off_axis or moving: - # Flow winds the transverse states through the same rotation - # off-resonance does, so the two phases add before either is taken. - off_phase = -2.0 * 3.141592653589793 * atom_b0 * dt + turn_t - off_cos = tl.cos(off_phase) - off_sin = tl.sin(off_phase) - if pools == 2 or pools == 3: - # Both pools take the same off-resonance and the same per-order - # damping; what separates them is the chemical shift, which the - # exchange operator already carries. - ( - x11r, - x11i, - x12r, - x12i, - x21r, - x21i, - x22r, - x22i, - ) = _two_pool_transverse_step( - 1000.0 / atom_t2, - atom_r2_bound, - atom_exchange, - atom_bound, - atom_free, - atom_shift, - dt, - wout, - ) - mixed_pr = ( - x11r * fplus_real - - x11i * fplus_imag - + x12r * bplus_real - - x12i * bplus_imag - ) - mixed_pi = ( - x11r * fplus_imag - + x11i * fplus_real - + x12r * bplus_imag - + x12i * bplus_real - ) - mixed_br = ( - x21r * fplus_real - - x21i * fplus_imag - + x22r * bplus_real - - x22i * bplus_imag - ) - mixed_bi = ( - x21r * fplus_imag - + x21i * fplus_real - + x22r * bplus_imag - + x22i * bplus_real - ) - # ``F-`` follows the conjugate of the operator entry by entry, not - # its transpose: it is the conjugate state, and the map it takes is - # the conjugate map. - mixed_mr = ( - x11r * fminus_real - + x11i * fminus_imag - + x12r * bminus_real - + x12i * bminus_imag - ) - mixed_mi = ( - x11r * fminus_imag - - x11i * fminus_real - + x12r * bminus_imag - - x12i * bminus_real - ) - mixed_nr = ( - x21r * fminus_real - + x21i * fminus_imag - + x22r * bminus_real - + x22i * bminus_imag - ) - mixed_ni = ( - x21r * fminus_imag - - x21i * fminus_real - + x22r * bminus_imag - - x22i * bminus_real - ) - fplus_real = damp_t * (mixed_pr * off_cos - mixed_pi * off_sin) - fplus_imag = damp_t * (mixed_pr * off_sin + mixed_pi * off_cos) - bplus_real = damp_t * (mixed_br * off_cos - mixed_bi * off_sin) - bplus_imag = damp_t * (mixed_br * off_sin + mixed_bi * off_cos) - fminus_real = damp_t * (mixed_mr * off_cos + mixed_mi * off_sin) - fminus_imag = damp_t * (-mixed_mr * off_sin + mixed_mi * off_cos) - bminus_real = damp_t * (mixed_nr * off_cos + mixed_ni * off_sin) - bminus_imag = damp_t * (-mixed_nr * off_sin + mixed_ni * off_cos) - else: - old_real = fplus_real - fplus_real = e2 * (old_real * off_cos - fplus_imag * off_sin) - fplus_imag = e2 * (old_real * off_sin + fplus_imag * off_cos) - old_real = fminus_real - fminus_real = e2 * (old_real * off_cos + fminus_imag * off_sin) - fminus_imag = e2 * (-old_real * off_sin + fminus_imag * off_cos) - # The longitudinal states carry a phase of their own, which nothing - # else in the state machine gives them. - turn_cos = 1.0 - turn_sin = 0.0 - if moving: - turn_cos = tl.cos(turn_z) - turn_sin = tl.sin(turn_z) - if pools == 3: - # Three pools mix through a 3x3 formed in double; every pool takes - # the same per-order damping and flow phase, their order-n states - # describing one dephasing configuration. - if tabulated: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - ) = _three_pool_from_table( - pool_table, - tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ), - atom, - atom_count, - active_atom, - wout, - atom_free, - atom_bound, - atom_semisolid, - ) - else: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - ) = _three_pool_step( - 1000.0 / atom_t1, - atom_r1_bound, - atom_r1_semisolid, - atom_exchange, - atom_semisolid_exchange, - atom_bound, - atom_semisolid, - dt, - wout, - narrow, - ) - free_real = ( - t11 * longitudinal_real + t12 * bound_real + t13 * semisolid_real - ) - free_imag = ( - t11 * longitudinal_imag + t12 * bound_imag + t13 * semisolid_imag - ) - held_real = ( - t21 * longitudinal_real + t22 * bound_real + t23 * semisolid_real - ) - held_imag = ( - t21 * longitudinal_imag + t22 * bound_imag + t23 * semisolid_imag - ) - stuck_real = ( - t31 * longitudinal_real + t32 * bound_real + t33 * semisolid_real - ) - stuck_imag = ( - t31 * longitudinal_imag + t32 * bound_imag + t33 * semisolid_imag - ) - longitudinal_real = damp_z * (free_real * turn_cos - free_imag * turn_sin) - longitudinal_imag = damp_z * (free_real * turn_sin + free_imag * turn_cos) - bound_real = damp_z * (held_real * turn_cos - held_imag * turn_sin) - bound_imag = damp_z * (held_real * turn_sin + held_imag * turn_cos) - semisolid_real = damp_z * (stuck_real * turn_cos - stuck_imag * turn_sin) - semisolid_imag = damp_z * (stuck_real * turn_sin + stuck_imag * turn_cos) - longitudinal_real += tl.where(state == 0, grow_free, 0.0) - bound_real += tl.where(state == 0, grow_pool_b, 0.0) - semisolid_real += tl.where(state == 0, grow_semisolid, 0.0) - elif pools > 0: - # The exchange operator is a property of the interval, not of a - # dephasing order, so it is formed once and the per-order damping - # multiplies it. Both pools take that damping and the flow phase: - # their order-n states describe one dephasing configuration, and a - # second pool has no diffusion coefficient of its own to damp by. - e11, e12, e21, e22, grow_free, grow_bound = _two_pool_step( - 1000.0 / atom_t1, atom_r1_bound, atom_exchange, atom_bound, dt, wout - ) - free_real = e11 * longitudinal_real + e12 * bound_real - free_imag = e11 * longitudinal_imag + e12 * bound_imag - held_real = e21 * longitudinal_real + e22 * bound_real - held_imag = e21 * longitudinal_imag + e22 * bound_imag - longitudinal_real = damp_z * (free_real * turn_cos - free_imag * turn_sin) - longitudinal_imag = damp_z * (free_real * turn_sin + free_imag * turn_cos) - bound_real = damp_z * (held_real * turn_cos - held_imag * turn_sin) - bound_imag = damp_z * (held_real * turn_sin + held_imag * turn_cos) - longitudinal_real += tl.where(state == 0, grow_free, 0.0) - bound_real += tl.where(state == 0, grow_bound, 0.0) - else: - old_real = longitudinal_real - longitudinal_real = e1 * ( - old_real * turn_cos - longitudinal_imag * turn_sin - ) - longitudinal_imag = e1 * ( - old_real * turn_sin + longitudinal_imag * turn_cos - ) - longitudinal_real += tl.where(state == 0, recovery, 0.0) - - # Every program reads the same event, so a branch on what it does is - # taken by all of them alike: an event pays only for what it does. - event_action = tl.load(action + event).to(tl.int32) - if (event_action & 1) != 0: - fplus_real, fplus_imag, fminus_real, fminus_imag = _shift( - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - state, - state_mask, - state_count, - ) - if pools == 2 or pools == 3: - bplus_real, bplus_imag, bminus_real, bminus_imag = _shift( - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, - state, - state_mask, - state_count, - ) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - longitudinal_real = tl.where( - invert, -atom_inversion * longitudinal_real, longitudinal_real - ) - longitudinal_imag = tl.where( - invert, -atom_inversion * longitudinal_imag, longitudinal_imag - ) - if pools == 2 or pools == 3: - # A chemically exchanging pool is free water and turns over like - # any other; a semisolid one is saturated instead, which its own - # saturation term already carries. - bound_real = tl.where(invert, -atom_inversion * bound_real, bound_real) - bound_imag = tl.where(invert, -atom_inversion * bound_imag, bound_imag) - - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse a row of its own. - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - if off_axis: - atom_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - b1_cos = tl.cos(atom_b1_phase) - b1_sin = tl.sin(atom_b1_phase) - if (event_kind == 1) & ((event_action & 4) == 0): - alpha = ( - _event_value(flip, event_base, event, active_atom, single_train) - * atom_b1 - ) - # The pulse's phase, read off the cosine and sine the launch took - # of it, turned by the transmit field's own. - cos_event = _event_value( - phase_cos, event_base, event, active_atom, single_train - ) - sin_event = _event_value( - phase_sin, event_base, event, active_atom, single_train - ) - cos_phi = tl.fma(cos_event, b1_cos, -(sin_event * b1_sin)) - sin_phi = tl.fma(sin_event, b1_cos, cos_event * b1_sin) - if profiled or dynamic: - # Either pair is built at zero RF phase, which turns the rotation - # axis and so reaches ``b`` alone. - if dynamic: - # Already integrated at this pulse's own flip, so the flip is - # inside the pair rather than read against it. - pair = _dynamic_pair_at( - pairs, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - else: - pair = _profile_pair( - profile, - _table_row(profile_index, event, location, locations), - alpha, - profile_bins, - profile_step, - ) - turn_r = cos_phi - turn_i = -sin_phi - spun_br = pair[2] * turn_r - pair[3] * turn_i - spun_bi = pair[2] * turn_i + pair[3] * turn_r - (shaped_pr, shaped_pi, shaped_mr, shaped_mi, shaped_zr, shaped_zi) = ( - _rotate_spinor( - pair[0], - pair[1], - spun_br, - spun_bi, - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - longitudinal_real, - longitudinal_imag, - ) - ) - sine, cosine = _sincos(alpha) - cos_2phi = tl.fma(cos_phi, cos_phi, -(sin_phi * sin_phi)) - sin_2phi = 2.0 * sin_phi * cos_phi - - ( - rotated_pr, - rotated_pi, - rotated_mr, - rotated_mi, - rotated_zr, - rotated_zi, - ) = _rotate_flip_phase( - cosine, - sine, - cos_phi, - sin_phi, - cos_2phi, - sin_2phi, - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - longitudinal_real, - longitudinal_imag, - ) - - if pools == 2 or pools == 3: - ( - b_rot_pr, - b_rot_pi, - b_rot_mr, - b_rot_mi, - b_rot_zr, - b_rot_zi, - ) = _rotate_flip_phase( - cosine, - sine, - cos_phi, - sin_phi, - cos_2phi, - sin_2phi, - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, - bound_real, - bound_imag, - ) - if profiled or dynamic: - # The same pulse, the same rotation: a chemical shift moves - # where a pool precesses, not what a pulse does to it. - ( - b_rot_pr, - b_rot_pi, - b_rot_mr, - b_rot_mi, - b_rot_zr, - b_rot_zi, - ) = _rotate_spinor( - pair[0], - pair[1], - spun_br, - spun_bi, - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, - bound_real, - bound_imag, - ) - if profiled or dynamic: - rotated_pr = shaped_pr - rotated_pi = shaped_pi - rotated_mr = shaped_mr - rotated_mi = shaped_mi - rotated_zr = shaped_zr - rotated_zi = shaped_zi - - if pools == 2 or pools == 3: - bplus_real = b_rot_pr - bplus_imag = b_rot_pi - bminus_real = b_rot_mr - bminus_imag = b_rot_mi - bound_real = b_rot_zr - bound_imag = b_rot_zi - if pools == 1 or pools == 3: - # The semisolid pool absorbs the power the pulse deposits, so it - # reads the bare flip the transmit field gives the voxel -- not the - # slice-shaped rotation the free pool takes from the table. - offset = tl.load(rf_frequency + event) - atom_b0 - absorbed = tl.exp( - tl.load(saturation + event) - * alpha - * alpha - * _lineshape_at(lineshape, offset, lineshape_bins, lineshape_step) - ) - if pools == 1: - bound_real = absorbed * bound_real - bound_imag = absorbed * bound_imag - else: - semisolid_real = absorbed * semisolid_real - semisolid_imag = absorbed * semisolid_imag - fplus_real = rotated_pr - fplus_imag = rotated_pi - fminus_real = rotated_mr - fminus_imag = rotated_mi - longitudinal_real = rotated_zr - longitudinal_imag = rotated_zi - - if ((event_action & 32) != 0) & (event_kind == 2): - adc_cos = _event_value( - phase_cos, event_base, event, active_atom, single_train - ) - adc_sin = _event_value( - phase_sin, event_base, event, active_atom, single_train - ) - # A coil sees the whole voxel, so what it records is the sum over - # pools; each pool's share is already in its own state. - read_real = fplus_real - read_imag = fplus_imag - if pools == 2 or pools == 3: - read_real = fplus_real + bplus_real - read_imag = fplus_imag + bplus_imag - signal_real = atom_m0 * (read_real * adc_cos + read_imag * adc_sin) - signal_imag = atom_m0 * (read_imag * adc_cos - read_real * adc_sin) - out = tl.load(output_index + event) - output_offset = problem * output_count + out - output_mask = active_atom & (state == 0) & (out >= 0) - tl.store(output_real + output_offset + state, signal_real, mask=output_mask) - tl.store(output_imag + output_offset + state, signal_imag, mask=output_mask) - - if (event_action & 18) != 0: - fplus_real, fplus_imag, fminus_real, fminus_imag = _shift( - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - state, - state_mask, - state_count, - ) - if pools == 2 or pools == 3: - bplus_real, bplus_imag, bminus_real, bminus_imag = _shift( - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, - state, - state_mask, - state_count, - ) - if (event_action & 8) != 0: - fplus_real = empty - fplus_imag = empty - fminus_real = empty - fminus_imag = empty - if pools == 2 or pools == 3: - bplus_real = empty - bplus_imag = empty - bminus_real = empty - bminus_imag = empty - - -@triton.jit -def _rotate_flip_phase_jvp( - cosine, - dcosine, - sine, - dsine, - cos_phi, - dcos_phi, - sin_phi, - dsin_phi, - cos_2phi, - dcos_2phi, - sin_2phi, - dsin_2phi, - fpr, - fpi, - fmr, - fmi, - zr, - zi, - dfpr, - dfpi, - dfmr, - dfmi, - dzr, - dzi, -): - """One pool through a hard pulse, carried alongside a tangent. - - Pulled out of the kernel body so a second pool can take the same - rotation: a chemical shift moves where a pool precesses, not what a - pulse does to it. - """ - ch = 0.5 * (1.0 + cosine) - sh = 0.5 * (1.0 - cosine) - dch = 0.5 * dcosine - dsh = -0.5 * dcosine - pr_a = cos_2phi * fmr - sin_2phi * fmi - dpr_a = dcos_2phi * fmr + cos_2phi * dfmr - dpr_a -= dsin_2phi * fmi + sin_2phi * dfmi - pr_b = sin_phi * zr + cos_phi * zi - dpr_b = dsin_phi * zr + sin_phi * dzr + dcos_phi * zi + cos_phi * dzi - rotated_pr = ch * fpr + sh * pr_a + sine * pr_b - rotated_dpr = dch * fpr + ch * dfpr + dsh * pr_a + sh * dpr_a - rotated_dpr += dsine * pr_b + sine * dpr_b - - pi_a = sin_2phi * fmr + cos_2phi * fmi - dpi_a = dsin_2phi * fmr + sin_2phi * dfmr - dpi_a += dcos_2phi * fmi + cos_2phi * dfmi - pi_b = sin_phi * zi - cos_phi * zr - dpi_b = dsin_phi * zi + sin_phi * dzi - dcos_phi * zr - cos_phi * dzr - rotated_pi = ch * fpi + sh * pi_a + sine * pi_b - rotated_dpi = dch * fpi + ch * dfpi + dsh * pi_a + sh * dpi_a - rotated_dpi += dsine * pi_b + sine * dpi_b - - mr_a = cos_2phi * fpr + sin_2phi * fpi - dmr_a = dcos_2phi * fpr + cos_2phi * dfpr - dmr_a += dsin_2phi * fpi + sin_2phi * dfpi - mr_b = sin_phi * zr - cos_phi * zi - dmr_b = dsin_phi * zr + sin_phi * dzr - dcos_phi * zi - cos_phi * dzi - rotated_mr = sh * mr_a + ch * fmr + sine * mr_b - rotated_dmr = dsh * mr_a + sh * dmr_a + dch * fmr + ch * dfmr - rotated_dmr += dsine * mr_b + sine * dmr_b - - mi_a = -sin_2phi * fpr + cos_2phi * fpi - dmi_a = -dsin_2phi * fpr - sin_2phi * dfpr - dmi_a += dcos_2phi * fpi + cos_2phi * dfpi - mi_b = cos_phi * zr + sin_phi * zi - dmi_b = dcos_phi * zr + cos_phi * dzr + dsin_phi * zi + sin_phi * dzi - rotated_mi = sh * mi_a + ch * fmi + sine * mi_b - rotated_dmi = dsh * mi_a + sh * dmi_a + dch * fmi + ch * dfmi - rotated_dmi += dsine * mi_b + sine * dmi_b - - zr_a = sin_phi * fpr - cos_phi * fpi - dzr_a = dsin_phi * fpr + sin_phi * dfpr - dcos_phi * fpi - cos_phi * dfpi - zr_b = sin_phi * fmr + cos_phi * fmi - dzr_b = dsin_phi * fmr + sin_phi * dfmr + dcos_phi * fmi + cos_phi * dfmi - rotated_zr = -0.5 * sine * zr_a - 0.5 * sine * zr_b + cosine * zr - rotated_dzr = -0.5 * (dsine * zr_a + sine * dzr_a) - rotated_dzr -= 0.5 * (dsine * zr_b + sine * dzr_b) - rotated_dzr += dcosine * zr + cosine * dzr - - zi_a = cos_phi * fpr + sin_phi * fpi - dzi_a = dcos_phi * fpr + cos_phi * dfpr + dsin_phi * fpi + sin_phi * dfpi - zi_b = cos_phi * fmr - sin_phi * fmi - dzi_b = dcos_phi * fmr + cos_phi * dfmr - dsin_phi * fmi - sin_phi * dfmi - rotated_zi = -0.5 * sine * zi_a + 0.5 * sine * zi_b + cosine * zi - rotated_dzi = -0.5 * (dsine * zi_a + sine * dzi_a) - rotated_dzi += 0.5 * (dsine * zi_b + sine * dzi_b) - rotated_dzi += dcosine * zi + cosine * dzi - - return ( - rotated_pr, - rotated_pi, - rotated_mr, - rotated_mi, - rotated_zr, - rotated_zi, - rotated_dpr, - rotated_dpi, - rotated_dmr, - rotated_dmi, - rotated_dzr, - rotated_dzi, - ) - - -@triton.jit( - do_not_specialize=["state_count", "locations", "profile_bins", "lineshape_bins"] -) -def _epg_jvp_kernel( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - # The bound pool's three. A bound pool is outside the real subspace this - # kernel stands for, so the dispatch never sends one here; the pointers - # hold the ABI's shape. - bound_fraction, - exchange_rate, - t1_bound, - # The chemically exchanging pool's five, on the same terms. - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - tangent_t1, - tangent_t2, - tangent_m0, - tangent_b1, - tangent_b1_phase, - tangent_b0, - tangent_inversion_efficiency, - tangent_diffusion, - tangent_velocity, - tangent_bound_fraction, - tangent_exchange_rate, - tangent_t1_bound, - tangent_pool_b_fraction, - tangent_pool_b_exchange, - tangent_t1_pool_b, - tangent_t2_pool_b, - tangent_pool_b_shift, - tangent_duration, - tangent_flip, - tangent_phase, - saturation, - rf_frequency, - profile, - profile_index, - lineshape, - pairs, - pair_index, - pair_direction, - duration_row, - pool_table, - output_real, - output_imag, - atom_count, - train_count, - event_count, - output_count, - flow_scale, - washout_scale, - profile_step, - lineshape_step, - state_count, - single_train: tl.constexpr, - atom_stride: tl.constexpr, - shim_rows, - shimmed: tl.constexpr, - locations, - profiled: tl.constexpr, - profile_bins, - dynamic: tl.constexpr, - broadened: tl.constexpr, - lineshape_bins, - pools: tl.constexpr, - narrow: tl.constexpr, - tabulated: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - block_states: tl.constexpr, - problems: tl.constexpr, -): - problem = tl.program_id(0) * problems + tl.arange(0, problems)[:, None] - state = tl.arange(0, block_states)[None, :] - active_atom = problem < train_count * atom_count - # A partial block carries lanes with no problem behind them, and they must - # take no part in a reduction or a store. - state_mask = (state < state_count) & active_atom - atom = problem % atom_count - # A property given as one value for the whole tissue is read at one - # address by every voxel, which is a stride of zero through it. - scalar_atom = atom * atom_stride - train = problem // atom_count - # Voxels are spread over the slice voxel-major, so a voxel's place along - # the slice is its index modulo the profile's width. One pulse shape holds - # that many consecutive rows, and the event says which shape it drives. - location = atom % locations - - empty = tl.zeros((problems, block_states), tl.float32) - fpr = empty - fpi = empty - fmr = empty - fmi = empty - # Equilibrium is split between the pools, so a direction along the bound - # fraction moves magnetization from one to the other before a single event - # has run. - atom_bound = 0.0 - d_bound = 0.0 - atom_exchange = 0.0 - d_exchange = 0.0 - atom_r1_bound = 0.0 - d_r1_bound = 0.0 - atom_r2_bound = 0.0 - d_r2_bound = 0.0 - atom_shift = 0.0 - d_shift = 0.0 - atom_semisolid = 0.0 - d_semisolid = 0.0 - atom_semisolid_exchange = 0.0 - d_semisolid_exchange = 0.0 - atom_r1_semisolid = 0.0 - d_r1_semisolid = 0.0 - if pools == 1: - atom_bound = tl.load(bound_fraction + scalar_atom, mask=active_atom, other=0.0) - d_bound = tl.load( - tangent_bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - d_exchange = tl.load( - tangent_exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - held_t1 = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - atom_r1_bound = 1000.0 / held_t1 - d_r1_bound = ( - -1000.0 - * tl.load(tangent_t1_bound + scalar_atom, mask=active_atom, other=0.0) - / (held_t1 * held_t1) - ) - if pools == 2 or pools == 3: - atom_bound = tl.load(pool_b_fraction + scalar_atom, mask=active_atom, other=0.0) - d_bound = tl.load( - tangent_pool_b_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_exchange = tl.load( - pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - d_exchange = tl.load( - tangent_pool_b_exchange + scalar_atom, mask=active_atom, other=0.0 - ) - held_t1 = tl.load(t1_pool_b + scalar_atom, mask=active_atom, other=1.0) - atom_r1_bound = 1000.0 / held_t1 - d_r1_bound = ( - -1000.0 - * tl.load(tangent_t1_pool_b + scalar_atom, mask=active_atom, other=0.0) - / (held_t1 * held_t1) - ) - held_t2 = tl.load(t2_pool_b + scalar_atom, mask=active_atom, other=1.0) - atom_r2_bound = 1000.0 / held_t2 - d_r2_bound = ( - -1000.0 - * tl.load(tangent_t2_pool_b + scalar_atom, mask=active_atom, other=0.0) - / (held_t2 * held_t2) - ) - atom_shift = tl.load(pool_b_shift + scalar_atom, mask=active_atom, other=0.0) - d_shift = tl.load( - tangent_pool_b_shift + scalar_atom, mask=active_atom, other=0.0 - ) - if pools == 3: - atom_semisolid = tl.load( - bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - d_semisolid = tl.load( - tangent_bound_fraction + scalar_atom, mask=active_atom, other=0.0 - ) - atom_semisolid_exchange = tl.load( - exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - d_semisolid_exchange = tl.load( - tangent_exchange_rate + scalar_atom, mask=active_atom, other=0.0 - ) - held_semisolid = tl.load(t1_bound + scalar_atom, mask=active_atom, other=1.0) - atom_r1_semisolid = 1000.0 / held_semisolid - d_r1_semisolid = ( - -1000.0 - * tl.load(tangent_t1_bound + scalar_atom, mask=active_atom, other=0.0) - / (held_semisolid * held_semisolid) - ) - atom_free = 1.0 - atom_bound - atom_semisolid - d_free = -d_bound - d_semisolid - zr = empty + tl.where(state == 0, atom_free, 0.0) - zi = empty - br = empty + tl.where(state == 0, atom_bound + 0.0, 0.0) - bi = empty - cr = empty + tl.where(state == 0, atom_semisolid + 0.0, 0.0) - ci = empty - dcr = empty + tl.where(state == 0, d_semisolid + 0.0, 0.0) - dci = empty - dfpr = empty - dfpi = empty - dfmr = empty - dfmi = empty - dzr = empty + tl.where(state == 0, -d_bound - d_semisolid, 0.0) - dzi = empty - dbr = empty + tl.where(state == 0, d_bound + 0.0, 0.0) - dbi = empty - bpr = empty - bpi = empty - bmr = empty - bmi = empty - dbpr = empty - dbpi = empty - dbmr = empty - dbmi = empty - - atom_t1 = tl.load(t1 + atom, mask=active_atom, other=1.0) - atom_t2 = tl.load(t2 + atom, mask=active_atom, other=1.0) - atom_m0 = 1.0 - if density: - atom_m0 = tl.load(m0 + scalar_atom, mask=active_atom, other=0.0) - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + scalar_atom, mask=active_atom, other=1.0) - atom_b1_phase = 0.0 - atom_b0 = 0.0 - if off_axis: - atom_b1_phase = tl.load(b1_phase + scalar_atom, mask=active_atom, other=0.0) - atom_b0 = tl.load(b0 + scalar_atom, mask=active_atom, other=0.0) - atom_inversion = 1.0 - if inverting: - atom_inversion = tl.load( - inversion_efficiency + scalar_atom, mask=active_atom, other=1.0 - ) - atom_damping = 0.0 - d_damping = 0.0 - if diffusing: - atom_damping = tl.load(diffusion + scalar_atom, mask=active_atom, other=0.0) - d_damping = tl.load( - tangent_diffusion + scalar_atom, mask=active_atom, other=0.0 - ) - atom_flow = 0.0 - d_flow = 0.0 - atom_washout = 0.0 - d_washout = 0.0 - if moving: - atom_velocity = tl.load(velocity + scalar_atom, mask=active_atom, other=0.0) - d_velocity = tl.load( - tangent_velocity + scalar_atom, mask=active_atom, other=0.0 - ) - atom_flow = atom_velocity * flow_scale - d_flow = d_velocity * flow_scale - # |v| has no derivative at the origin, so a still voxel contributes - # none. - direction = (atom_velocity > 0.0).to(tl.float32) - (atom_velocity < 0.0).to( - tl.float32 - ) - atom_washout = tl.abs(atom_velocity) * washout_scale - d_washout = direction * d_velocity * washout_scale - order = state.to(tl.float32) - dt1 = tl.load(tangent_t1 + atom, mask=active_atom, other=0.0) - dt2 = tl.load(tangent_t2 + atom, mask=active_atom, other=0.0) - dm0 = 0.0 - if density: - dm0 = tl.load(tangent_m0 + scalar_atom, mask=active_atom, other=0.0) - db1 = 0.0 - if transmit: - db1 = tl.load(tangent_b1 + scalar_atom, mask=active_atom, other=0.0) - db1_phase = 0.0 - db0 = 0.0 - if off_axis: - db1_phase = tl.load(tangent_b1_phase + scalar_atom, mask=active_atom, other=0.0) - db0 = tl.load(tangent_b0 + scalar_atom, mask=active_atom, other=0.0) - dinversion = 0.0 - if inverting: - dinversion = tl.load( - tangent_inversion_efficiency + scalar_atom, mask=active_atom, other=0.0 - ) - - event_base = train * event_count - for event in range(0, event_count): - event_dt = _event_value(duration, event_base, event, active_atom, single_train) - ddt = _event_value( - tangent_duration, event_base, event, active_atom, single_train - ) - r1 = 1000.0 / atom_t1 - r2 = 1000.0 / atom_t2 - wout = 1.0 - dwout = 0.0 - if moving: - wout, dwout = _washout_jvp(atom_washout, d_washout, event_dt, ddt) - dry1 = tl.exp(-r1 * event_dt) - dry2 = tl.exp(-r2 * event_dt) - e1 = dry1 * wout - e2 = dry2 * wout - de1 = ( - e1 * (1000.0 * event_dt * dt1 / (atom_t1 * atom_t1) - r1 * ddt) - + dry1 * dwout - ) - de2 = ( - e2 * (1000.0 * event_dt * dt2 / (atom_t2 * atom_t2) - r2 * ddt) - + dry2 * dwout - ) - damp_z = 1.0 - ddamp_z = 0.0 - damp_t = 1.0 - ddamp_t = 0.0 - if diffusing: - damp_z, ddamp_z, damp_t, ddamp_t = _damping_jvp( - atom_damping, d_damping, event_dt, ddt, order - ) - # Order zero is undamped, so the recovery term keeps the bare factor. - recovery, drecovery = 1.0 - e1, -de1 - de1 = de1 * damp_z + e1 * ddamp_z - e1 = e1 * damp_z - de2 = de2 * damp_t + e2 * ddamp_t - e2 = e2 * damp_t - turn_z = 0.0 - turn_t = 0.0 - dturn_z = 0.0 - dturn_t = 0.0 - if moving: - turn_z, turn_t = _flow(atom_flow, event_dt, order) - d_turn = d_flow * event_dt + atom_flow * ddt - dturn_z = -order * d_turn - dturn_t = -(order + 0.5) * d_turn - off_cos = 1.0 - off_sin = 0.0 - doff_cos = 0.0 - doff_sin = 0.0 - if off_axis or moving: - # Flow winds the transverse states through the same rotation - # off-resonance does, so the two phases add before either is taken. - off_phase = -2.0 * 3.141592653589793 * atom_b0 * event_dt + turn_t - doff_phase = ( - -2.0 * 3.141592653589793 * (db0 * event_dt + atom_b0 * ddt) + dturn_t - ) - off_cos = tl.cos(off_phase) - off_sin = tl.sin(off_phase) - doff_cos = -off_sin * doff_phase - doff_sin = off_cos * doff_phase - - if pools == 2 or pools == 3: - # Both pools take the same off-resonance and per-order damping; - # what separates them is the chemical shift, which the exchange - # operator already carries. - ( - x11r, - x11i, - x12r, - x12i, - x21r, - x21i, - x22r, - x22i, - d11r, - d11i, - d12r, - d12i, - d21r, - d21i, - d22r, - d22i, - ) = _two_pool_transverse_step_jvp( - r2, - -1000.0 * dt2 / (atom_t2 * atom_t2), - atom_r2_bound, - d_r2_bound, - atom_exchange, - d_exchange, - atom_bound, - d_bound, - atom_free, - d_free, - atom_shift, - d_shift, - event_dt, - ddt, - wout, - dwout, - ) - mix_pr = x11r * fpr - x11i * fpi + x12r * bpr - x12i * bpi - mix_pi = x11r * fpi + x11i * fpr + x12r * bpi + x12i * bpr - dmix_pr = ( - d11r * fpr - + x11r * dfpr - - d11i * fpi - - x11i * dfpi - + d12r * bpr - + x12r * dbpr - - d12i * bpi - - x12i * dbpi - ) - dmix_pi = ( - d11r * fpi - + x11r * dfpi - + d11i * fpr - + x11i * dfpr - + d12r * bpi - + x12r * dbpi - + d12i * bpr - + x12i * dbpr - ) - mix_br = x21r * fpr - x21i * fpi + x22r * bpr - x22i * bpi - mix_bi = x21r * fpi + x21i * fpr + x22r * bpi + x22i * bpr - dmix_br = ( - d21r * fpr - + x21r * dfpr - - d21i * fpi - - x21i * dfpi - + d22r * bpr - + x22r * dbpr - - d22i * bpi - - x22i * dbpi - ) - dmix_bi = ( - d21r * fpi - + x21r * dfpi - + d21i * fpr - + x21i * dfpr - + d22r * bpi - + x22r * dbpi - + d22i * bpr - + x22i * dbpr - ) - # ``F-`` follows the conjugate of the operator entry by entry. - mix_mr = x11r * fmr + x11i * fmi + x12r * bmr + x12i * bmi - mix_mi = x11r * fmi - x11i * fmr + x12r * bmi - x12i * bmr - dmix_mr = ( - d11r * fmr - + x11r * dfmr - + d11i * fmi - + x11i * dfmi - + d12r * bmr - + x12r * dbmr - + d12i * bmi - + x12i * dbmi - ) - dmix_mi = ( - d11r * fmi - + x11r * dfmi - - d11i * fmr - - x11i * dfmr - + d12r * bmi - + x12r * dbmi - - d12i * bmr - - x12i * dbmr - ) - mix_nr = x21r * fmr + x21i * fmi + x22r * bmr + x22i * bmi - mix_ni = x21r * fmi - x21i * fmr + x22r * bmi - x22i * bmr - dmix_nr = ( - d21r * fmr - + x21r * dfmr - + d21i * fmi - + x21i * dfmi - + d22r * bmr - + x22r * dbmr - + d22i * bmi - + x22i * dbmi - ) - dmix_ni = ( - d21r * fmi - + x21r * dfmi - - d21i * fmr - - x21i * dfmr - + d22r * bmi - + x22r * dbmi - - d22i * bmr - - x22i * dbmr - ) - # The damping and off-resonance both pools share, applied after. - carry_r = damp_t * off_cos - carry_i = damp_t * off_sin - dcarry_r = ddamp_t * off_cos + damp_t * doff_cos - dcarry_i = ddamp_t * off_sin + damp_t * doff_sin - fpr = mix_pr * carry_r - mix_pi * carry_i - fpi = mix_pr * carry_i + mix_pi * carry_r - dfpr = ( - dmix_pr * carry_r - + mix_pr * dcarry_r - - dmix_pi * carry_i - - mix_pi * dcarry_i - ) - dfpi = ( - dmix_pr * carry_i - + mix_pr * dcarry_i - + dmix_pi * carry_r - + mix_pi * dcarry_r - ) - bpr = mix_br * carry_r - mix_bi * carry_i - bpi = mix_br * carry_i + mix_bi * carry_r - dbpr = ( - dmix_br * carry_r - + mix_br * dcarry_r - - dmix_bi * carry_i - - mix_bi * dcarry_i - ) - dbpi = ( - dmix_br * carry_i - + mix_br * dcarry_i - + dmix_bi * carry_r - + mix_bi * dcarry_r - ) - fmr = mix_mr * carry_r + mix_mi * carry_i - fmi = -mix_mr * carry_i + mix_mi * carry_r - dfmr = ( - dmix_mr * carry_r - + mix_mr * dcarry_r - + dmix_mi * carry_i - + mix_mi * dcarry_i - ) - dfmi = ( - -dmix_mr * carry_i - - mix_mr * dcarry_i - + dmix_mi * carry_r - + mix_mi * dcarry_r - ) - bmr = mix_nr * carry_r + mix_ni * carry_i - bmi = -mix_nr * carry_i + mix_ni * carry_r - dbmr = ( - dmix_nr * carry_r - + mix_nr * dcarry_r - + dmix_ni * carry_i - + mix_ni * dcarry_i - ) - dbmi = ( - -dmix_nr * carry_i - - mix_nr * dcarry_i - + dmix_ni * carry_r - + mix_ni * dcarry_r - ) - else: - old_fpr = fpr - old_fpi = fpi - old_dfpr = dfpr - old_dfpi = dfpi - fpr = e2 * (old_fpr * off_cos - old_fpi * off_sin) - fpi = e2 * (old_fpr * off_sin + old_fpi * off_cos) - dfpr = de2 * (old_fpr * off_cos - old_fpi * off_sin) - dfpr += e2 * ( - old_dfpr * off_cos - + old_fpr * doff_cos - - old_dfpi * off_sin - - old_fpi * doff_sin - ) - dfpi = de2 * (old_fpr * off_sin + old_fpi * off_cos) - dfpi += e2 * ( - old_dfpr * off_sin - + old_fpr * doff_sin - + old_dfpi * off_cos - + old_fpi * doff_cos - ) - - old_fmr = fmr - old_fmi = fmi - old_dfmr = dfmr - old_dfmi = dfmi - fmr = e2 * (old_fmr * off_cos + old_fmi * off_sin) - fmi = e2 * (-old_fmr * off_sin + old_fmi * off_cos) - dfmr = de2 * (old_fmr * off_cos + old_fmi * off_sin) - dfmr += e2 * ( - old_dfmr * off_cos - + old_fmr * doff_cos - + old_dfmi * off_sin - + old_fmi * doff_sin - ) - dfmi = de2 * (-old_fmr * off_sin + old_fmi * off_cos) - dfmi += e2 * ( - -old_dfmr * off_sin - - old_fmr * doff_sin - + old_dfmi * off_cos - + old_fmi * doff_cos - ) - - # The longitudinal states carry a phase of their own, which nothing - # else in the state machine gives them. - turn_cos = 1.0 - turn_sin = 0.0 - dturn_cos = 0.0 - dturn_sin = 0.0 - if moving: - turn_cos = tl.cos(turn_z) - turn_sin = tl.sin(turn_z) - dturn_cos = -turn_sin * dturn_z - dturn_sin = turn_cos * dturn_z - old_zr = zr - old_zi = zi - old_dzr = dzr - old_dzi = dzi - spun_zr = old_zr * turn_cos - old_zi * turn_sin - spun_zi = old_zr * turn_sin + old_zi * turn_cos - dspun_zr = ( - old_dzr * turn_cos - + old_zr * dturn_cos - - old_dzi * turn_sin - - old_zi * dturn_sin - ) - dspun_zi = ( - old_dzr * turn_sin - + old_zr * dturn_sin - + old_dzi * turn_cos - + old_zi * dturn_cos - ) - if pools == 3: - # Three pools mix through a 3x3 formed in double, tangent and all: - # a direction through an operator this ill-conditioned needs the - # width as much as the value does. - if tabulated: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - d_t11, - d_t12, - d_t13, - d_t21, - d_t22, - d_t23, - d_t31, - d_t32, - d_t33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_from_table_jvp( - pool_table, - tl.load( - duration_row + event_base + event, - mask=active_atom, - other=0, - ), - atom, - atom_count, - active_atom, - r1, - atom_r1_bound, - atom_r1_semisolid, - atom_exchange, - atom_semisolid_exchange, - atom_bound, - d_bound, - atom_semisolid, - d_semisolid, - ddt, - wout, - dwout, - ) - else: - ( - t11, - t12, - t13, - t21, - t22, - t23, - t31, - t32, - t33, - grow_free, - grow_pool_b, - grow_semisolid, - d_t11, - d_t12, - d_t13, - d_t21, - d_t22, - d_t23, - d_t31, - d_t32, - d_t33, - d_grow_free, - d_grow_pool_b, - d_grow_semisolid, - ) = _three_pool_step_jvp( - r1, - -1000.0 * dt1 / (atom_t1 * atom_t1), - atom_r1_bound, - d_r1_bound, - atom_r1_semisolid, - d_r1_semisolid, - atom_exchange, - d_exchange, - atom_semisolid_exchange, - d_semisolid_exchange, - atom_bound, - d_bound, - atom_semisolid, - d_semisolid, - event_dt, - ddt, - wout, - dwout, - narrow, - ) - spun_hr = br * turn_cos - bi * turn_sin - spun_hi = br * turn_sin + bi * turn_cos - dspun_hr = dbr * turn_cos + br * dturn_cos - dbi * turn_sin - bi * dturn_sin - dspun_hi = dbr * turn_sin + br * dturn_sin + dbi * turn_cos + bi * dturn_cos - spun_cr = cr * turn_cos - ci * turn_sin - spun_ci = cr * turn_sin + ci * turn_cos - dspun_cr = dcr * turn_cos + cr * dturn_cos - dci * turn_sin - ci * dturn_sin - dspun_ci = dcr * turn_sin + cr * dturn_sin + dci * turn_cos + ci * dturn_cos - free_r = t11 * spun_zr + t12 * spun_hr + t13 * spun_cr - free_i = t11 * spun_zi + t12 * spun_hi + t13 * spun_ci - held_r = t21 * spun_zr + t22 * spun_hr + t23 * spun_cr - held_i = t21 * spun_zi + t22 * spun_hi + t23 * spun_ci - stuck_r = t31 * spun_zr + t32 * spun_hr + t33 * spun_cr - stuck_i = t31 * spun_zi + t32 * spun_hi + t33 * spun_ci - d_free_r = ( - d_t11 * spun_zr - + t11 * dspun_zr - + d_t12 * spun_hr - + t12 * dspun_hr - + d_t13 * spun_cr - + t13 * dspun_cr - ) - d_free_i = ( - d_t11 * spun_zi - + t11 * dspun_zi - + d_t12 * spun_hi - + t12 * dspun_hi - + d_t13 * spun_ci - + t13 * dspun_ci - ) - d_held_r = ( - d_t21 * spun_zr - + t21 * dspun_zr - + d_t22 * spun_hr - + t22 * dspun_hr - + d_t23 * spun_cr - + t23 * dspun_cr - ) - d_held_i = ( - d_t21 * spun_zi - + t21 * dspun_zi - + d_t22 * spun_hi - + t22 * dspun_hi - + d_t23 * spun_ci - + t23 * dspun_ci - ) - d_stuck_r = ( - d_t31 * spun_zr - + t31 * dspun_zr - + d_t32 * spun_hr - + t32 * dspun_hr - + d_t33 * spun_cr - + t33 * dspun_cr - ) - d_stuck_i = ( - d_t31 * spun_zi - + t31 * dspun_zi - + d_t32 * spun_hi - + t32 * dspun_hi - + d_t33 * spun_ci - + t33 * dspun_ci - ) - zr = damp_z * free_r + tl.where(state == 0, grow_free, 0.0) - zi = damp_z * free_i - dzr = ( - ddamp_z * free_r - + damp_z * d_free_r - + tl.where(state == 0, d_grow_free, 0.0) - ) - dzi = ddamp_z * free_i + damp_z * d_free_i - br = damp_z * held_r + tl.where(state == 0, grow_pool_b, 0.0) - bi = damp_z * held_i - dbr = ( - ddamp_z * held_r - + damp_z * d_held_r - + tl.where(state == 0, d_grow_pool_b, 0.0) - ) - dbi = ddamp_z * held_i + damp_z * d_held_i - cr = damp_z * stuck_r + tl.where(state == 0, grow_semisolid, 0.0) - ci = damp_z * stuck_i - dcr = ( - ddamp_z * stuck_r - + damp_z * d_stuck_r - + tl.where(state == 0, d_grow_semisolid, 0.0) - ) - dci = ddamp_z * stuck_i + damp_z * d_stuck_i - elif pools > 0: - # The exchange operator belongs to the interval, not to a dephasing - # order, so it is formed once and carries its own tangent; the - # per-order damping multiplies both pools, whose order-n states - # describe one dephasing configuration. - ( - e11, - e12, - e21, - e22, - grow_free, - grow_bound, - d_e11, - d_e12, - d_e21, - d_e22, - d_grow_free, - d_grow_bound, - ) = _two_pool_step_jvp( - r1, - -1000.0 * dt1 / (atom_t1 * atom_t1), - atom_r1_bound, - d_r1_bound, - atom_exchange, - d_exchange, - atom_bound, - d_bound, - event_dt, - ddt, - wout, - dwout, - ) - old_br = br - old_bi = bi - old_dbr = dbr - old_dbi = dbi - spun_hr = old_br * turn_cos - old_bi * turn_sin - spun_hi = old_br * turn_sin + old_bi * turn_cos - dspun_hr = ( - old_dbr * turn_cos - + old_br * dturn_cos - - old_dbi * turn_sin - - old_bi * dturn_sin - ) - dspun_hi = ( - old_dbr * turn_sin - + old_br * dturn_sin - + old_dbi * turn_cos - + old_bi * dturn_cos - ) - free_r = e11 * spun_zr + e12 * spun_hr - free_i = e11 * spun_zi + e12 * spun_hi - held_r = e21 * spun_zr + e22 * spun_hr - held_i = e21 * spun_zi + e22 * spun_hi - d_free_r = ( - d_e11 * spun_zr + e11 * dspun_zr + d_e12 * spun_hr + e12 * dspun_hr - ) - d_free_i = ( - d_e11 * spun_zi + e11 * dspun_zi + d_e12 * spun_hi + e12 * dspun_hi - ) - d_held_r = ( - d_e21 * spun_zr + e21 * dspun_zr + d_e22 * spun_hr + e22 * dspun_hr - ) - d_held_i = ( - d_e21 * spun_zi + e21 * dspun_zi + d_e22 * spun_hi + e22 * dspun_hi - ) - zr = damp_z * free_r + tl.where(state == 0, grow_free, 0.0) - zi = damp_z * free_i - dzr = ( - ddamp_z * free_r - + damp_z * d_free_r - + tl.where(state == 0, d_grow_free, 0.0) - ) - dzi = ddamp_z * free_i + damp_z * d_free_i - br = damp_z * held_r + tl.where(state == 0, grow_bound, 0.0) - bi = damp_z * held_i - dbr = ( - ddamp_z * held_r - + damp_z * d_held_r - + tl.where(state == 0, d_grow_bound, 0.0) - ) - dbi = ddamp_z * held_i + damp_z * d_held_i - else: - dzr = dspun_zr * e1 + spun_zr * de1 + tl.where(state == 0, drecovery, 0.0) - dzi = dspun_zi * e1 + spun_zi * de1 - zr = spun_zr * e1 + tl.where(state == 0, recovery, 0.0) - zi = spun_zi * e1 - - event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - shifted_pr, shifted_pi, shifted_mr, shifted_mi = _shift( - fpr, fpi, fmr, fmi, state, state_mask, state_count - ) - shifted_dpr, shifted_dpi, shifted_dmr, shifted_dmi = _shift( - dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count - ) - fpr = tl.where(pre_shift, shifted_pr, fpr) - fpi = tl.where(pre_shift, shifted_pi, fpi) - fmr = tl.where(pre_shift, shifted_mr, fmr) - fmi = tl.where(pre_shift, shifted_mi, fmi) - dfpr = tl.where(pre_shift, shifted_dpr, dfpr) - dfpi = tl.where(pre_shift, shifted_dpi, dfpi) - dfmr = tl.where(pre_shift, shifted_dmr, dfmr) - dfmi = tl.where(pre_shift, shifted_dmi, dfmi) - if pools == 2 or pools == 3: - s_bpr, s_bpi, s_bmr, s_bmi = _shift( - bpr, bpi, bmr, bmi, state, state_mask, state_count - ) - s_dbpr, s_dbpi, s_dbmr, s_dbmi = _shift( - dbpr, dbpi, dbmr, dbmi, state, state_mask, state_count - ) - bpr = tl.where(pre_shift, s_bpr, bpr) - bpi = tl.where(pre_shift, s_bpi, bpi) - bmr = tl.where(pre_shift, s_bmr, bmr) - bmi = tl.where(pre_shift, s_bmi, bmi) - dbpr = tl.where(pre_shift, s_dbpr, dbpr) - dbpi = tl.where(pre_shift, s_dbpi, dbpi) - dbmr = tl.where(pre_shift, s_dbmr, dbmr) - dbmi = tl.where(pre_shift, s_dbmi, dbmi) - - event_kind = tl.load(kind + event) - is_rf = event_kind == 1 - is_inversion = (event_action & 4) != 0 - invert = is_rf & is_inversion - dzr = tl.where(invert, -dinversion * zr - atom_inversion * dzr, dzr) - dzi = tl.where(invert, -dinversion * zi - atom_inversion * dzi, dzi) - zr = tl.where(invert, -atom_inversion * zr, zr) - zi = tl.where(invert, -atom_inversion * zi, zi) - if pools == 2 or pools == 3: - # A chemically exchanging pool is free water and turns over like - # any other; a semisolid one is saturated instead. - dbr = tl.where(invert, -dinversion * br - atom_inversion * dbr, dbr) - dbi = tl.where(invert, -dinversion * bi - atom_inversion * dbi, dbi) - br = tl.where(invert, -atom_inversion * br, br) - bi = tl.where(invert, -atom_inversion * bi, bi) - - event_flip = _event_value(flip, event_base, event, active_atom, single_train) - event_phase = _event_value(phase, event_base, event, active_atom, single_train) - # One shim is the whole sequence's transmit field, loaded once above; - # several give each pulse a row of its own. - if shimmed: - row = tl.load(shim_index + event).to(tl.int64) * atom_count - atom_b1 = 1.0 - if transmit: - atom_b1 = tl.load(b1 + row + atom, mask=active_atom, other=1.0) - db1 = tl.load(tangent_b1 + row + atom, mask=active_atom, other=0.0) - if off_axis: - atom_b1_phase = tl.load( - b1_phase + row + atom, mask=active_atom, other=0.0 - ) - db1_phase = tl.load( - tangent_b1_phase + row + atom, mask=active_atom, other=0.0 - ) - alpha = event_flip * atom_b1 - dalpha = ( - _event_value(tangent_flip, event_base, event, active_atom, single_train) - * atom_b1 - + event_flip * db1 - ) - phi = event_phase + atom_b1_phase - dphi = ( - _event_value(tangent_phase, event_base, event, active_atom, single_train) - + db1_phase - ) - if pools == 1 or pools == 3: - # The semisolid pool absorbs the power the pulse deposits, so it reads - # the bare flip the transmit field gives the voxel. The offset - # reaches it through the voxel's own off-resonance, which is where - # the lineshape's slope enters a forward direction. - offset = tl.load(rf_frequency + event) - atom_b0 - shape, shape_slope = _lineshape_at_slope( - lineshape, offset, lineshape_bins, lineshape_step - ) - deposited = tl.load(saturation + event) - absorbed = tl.exp(deposited * alpha * alpha * shape) - d_exponent = deposited * ( - 2.0 * alpha * dalpha * shape - alpha * alpha * shape_slope * db0 - ) - saturating = is_rf & ~is_inversion - if pools == 1: - dbr = tl.where(saturating, absorbed * (dbr + br * d_exponent), dbr) - dbi = tl.where(saturating, absorbed * (dbi + bi * d_exponent), dbi) - br = tl.where(saturating, absorbed * br, br) - bi = tl.where(saturating, absorbed * bi, bi) - else: - dcr = tl.where(saturating, absorbed * (dcr + cr * d_exponent), dcr) - dci = tl.where(saturating, absorbed * (dci + ci * d_exponent), dci) - cr = tl.where(saturating, absorbed * cr, cr) - ci = tl.where(saturating, absorbed * ci, ci) - if profiled or dynamic: - if dynamic: - # The array was resolved outside the kernel, so a direction - # along it arrives already carried through the pulse integral. - held = _dynamic_pair_at( - pairs, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - moved = _dynamic_pair_at( - pair_direction, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - pair_ar, pair_ai, pair_br, pair_bi = held - dot_ar, dot_ai, dot_br, dot_bi = moved - else: - read = _profile_pair_slope( - profile, - _table_row(profile_index, event, location, locations), - alpha, - profile_bins, - profile_step, - ) - # The flip angle carries the tangent into the table. - pair_ar, pair_ai = read[0], read[2] - pair_br, pair_bi = read[4], read[6] - dot_ar, dot_ai = read[1] * dalpha, read[3] * dalpha - dot_br, dot_bi = read[5] * dalpha, read[7] * dalpha - # The RF phase turns the axis after the pair comes out, and so - # reaches ``b`` alone. - turn_r = tl.cos(phi) - turn_i = -tl.sin(phi) - spun_br = pair_br * turn_r - pair_bi * turn_i - spun_bi = pair_br * turn_i + pair_bi * turn_r - slope_br = dot_br - slope_bi = dot_bi - ( - shaped_pr, - shaped_pi, - shaped_mr, - shaped_mi, - shaped_zr, - shaped_zi, - shaped_dpr, - shaped_dpi, - shaped_dmr, - shaped_dmi, - shaped_dzr, - shaped_dzi, - ) = _rotate_spinor_dual( - pair_ar, - pair_ai, - spun_br, - spun_bi, - dot_ar, - dot_ai, - slope_br * turn_r - slope_bi * turn_i + dphi * spun_bi, - slope_br * turn_i + slope_bi * turn_r - dphi * spun_br, - fpr, - fpi, - fmr, - fmi, - zr, - zi, - dfpr, - dfpi, - dfmr, - dfmi, - dzr, - dzi, - ) - if pools == 2 or pools == 3: - # The same pulse, the same rotation. - ( - held_pr, - held_pi, - held_mr, - held_mi, - held_zr, - held_zi, - held_dpr, - held_dpi, - held_dmr, - held_dmi, - held_dzr, - held_dzi, - ) = _rotate_spinor_dual( - pair_ar, - pair_ai, - spun_br, - spun_bi, - dot_ar, - dot_ai, - slope_br * turn_r - slope_bi * turn_i + dphi * spun_bi, - slope_br * turn_i + slope_bi * turn_r - dphi * spun_br, - bpr, - bpi, - bmr, - bmi, - br, - bi, - dbpr, - dbpi, - dbmr, - dbmi, - dbr, - dbi, - ) - cosine = tl.cos(alpha) - sine = tl.sin(alpha) - dcosine = -sine * dalpha - dsine = cosine * dalpha - cos_phi = tl.cos(phi) - sin_phi = tl.sin(phi) - cos_2phi = tl.cos(2.0 * phi) - sin_2phi = tl.sin(2.0 * phi) - dcos_phi = -sin_phi * dphi - dsin_phi = cos_phi * dphi - dcos_2phi = -2.0 * sin_2phi * dphi - dsin_2phi = 2.0 * cos_2phi * dphi - - ( - rotated_pr, - rotated_pi, - rotated_mr, - rotated_mi, - rotated_zr, - rotated_zi, - rotated_dpr, - rotated_dpi, - rotated_dmr, - rotated_dmi, - rotated_dzr, - rotated_dzi, - ) = _rotate_flip_phase_jvp( - cosine, - dcosine, - sine, - dsine, - cos_phi, - dcos_phi, - sin_phi, - dsin_phi, - cos_2phi, - dcos_2phi, - sin_2phi, - dsin_2phi, - fpr, - fpi, - fmr, - fmi, - zr, - zi, - dfpr, - dfpi, - dfmr, - dfmi, - dzr, - dzi, - ) - if pools == 2 or pools == 3: - ( - b_pr, - b_pi, - b_mr, - b_mi, - b_zr, - b_zi, - b_dpr, - b_dpi, - b_dmr, - b_dmi, - b_dzr, - b_dzi, - ) = _rotate_flip_phase_jvp( - cosine, - dcosine, - sine, - dsine, - cos_phi, - dcos_phi, - sin_phi, - dsin_phi, - cos_2phi, - dcos_2phi, - sin_2phi, - dsin_2phi, - bpr, - bpi, - bmr, - bmi, - br, - bi, - dbpr, - dbpi, - dbmr, - dbmi, - dbr, - dbi, - ) - if profiled or dynamic: - rotated_pr = shaped_pr - rotated_pi = shaped_pi - rotated_mr = shaped_mr - rotated_mi = shaped_mi - rotated_zr = shaped_zr - rotated_zi = shaped_zi - rotated_dpr = shaped_dpr - rotated_dpi = shaped_dpi - rotated_dmr = shaped_dmr - rotated_dmi = shaped_dmi - rotated_dzr = shaped_dzr - rotated_dzi = shaped_dzi - if pools == 2 or pools == 3: - b_pr = held_pr - b_pi = held_pi - b_mr = held_mr - b_mi = held_mi - b_zr = held_zr - b_zi = held_zi - b_dpr = held_dpr - b_dpi = held_dpi - b_dmr = held_dmr - b_dmi = held_dmi - b_dzr = held_dzr - b_dzi = held_dzi - - rotate = is_rf & ~is_inversion - fpr = tl.where(rotate, rotated_pr, fpr) - fpi = tl.where(rotate, rotated_pi, fpi) - fmr = tl.where(rotate, rotated_mr, fmr) - fmi = tl.where(rotate, rotated_mi, fmi) - zr = tl.where(rotate, rotated_zr, zr) - zi = tl.where(rotate, rotated_zi, zi) - dfpr = tl.where(rotate, rotated_dpr, dfpr) - dfpi = tl.where(rotate, rotated_dpi, dfpi) - dfmr = tl.where(rotate, rotated_dmr, dfmr) - dfmi = tl.where(rotate, rotated_dmi, dfmi) - dzr = tl.where(rotate, rotated_dzr, dzr) - dzi = tl.where(rotate, rotated_dzi, dzi) - if pools == 2 or pools == 3: - bpr = tl.where(rotate, b_pr, bpr) - bpi = tl.where(rotate, b_pi, bpi) - bmr = tl.where(rotate, b_mr, bmr) - bmi = tl.where(rotate, b_mi, bmi) - br = tl.where(rotate, b_zr, br) - bi = tl.where(rotate, b_zi, bi) - dbpr = tl.where(rotate, b_dpr, dbpr) - dbpi = tl.where(rotate, b_dpi, dbpi) - dbmr = tl.where(rotate, b_dmr, dbmr) - dbmi = tl.where(rotate, b_dmi, dbmi) - dbr = tl.where(rotate, b_dzr, dbr) - dbi = tl.where(rotate, b_dzi, dbi) - - record = ((event_action & 32) != 0) & (event_kind == 2) - adc_cos = tl.cos(event_phase) - adc_sin = tl.sin(event_phase) - dadc_phase = _event_value( - tangent_phase, event_base, event, active_atom, single_train - ) - dadc_cos = -adc_sin * dadc_phase - dadc_sin = adc_cos * dadc_phase - read_r = fpr - read_i = fpi - dread_r = dfpr - dread_i = dfpi - if pools == 2 or pools == 3: - read_r = fpr + bpr - read_i = fpi + bpi - dread_r = dfpr + dbpr - dread_i = dfpi + dbpi - signal_real = dm0 * (read_r * adc_cos + read_i * adc_sin) - signal_real += atom_m0 * ( - dread_r * adc_cos - + read_r * dadc_cos - + dread_i * adc_sin - + read_i * dadc_sin - ) - signal_imag = dm0 * (read_i * adc_cos - read_r * adc_sin) - signal_imag += atom_m0 * ( - dread_i * adc_cos - + read_i * dadc_cos - - dread_r * adc_sin - - read_r * dadc_sin - ) - out = tl.load(output_index + event) - output_offset = problem * output_count + out - output_mask = active_atom & (state == 0) & record & (out >= 0) - tl.store(output_real + output_offset + state, signal_real, mask=output_mask) - tl.store(output_imag + output_offset + state, signal_imag, mask=output_mask) - - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - shifted_pr, shifted_pi, shifted_mr, shifted_mi = _shift( - fpr, fpi, fmr, fmi, state, state_mask, state_count - ) - shifted_dpr, shifted_dpi, shifted_dmr, shifted_dmi = _shift( - dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count - ) - fpr = tl.where(do_shift, shifted_pr, fpr) - fpi = tl.where(do_shift, shifted_pi, fpi) - fmr = tl.where(do_shift, shifted_mr, fmr) - fmi = tl.where(do_shift, shifted_mi, fmi) - dfpr = tl.where(do_shift, shifted_dpr, dfpr) - dfpi = tl.where(do_shift, shifted_dpi, dfpi) - dfmr = tl.where(do_shift, shifted_dmr, dfmr) - dfmi = tl.where(do_shift, shifted_dmi, dfmi) - spoil = (event_action & 8) != 0 - fpr = tl.where(spoil, 0.0, fpr) - fpi = tl.where(spoil, 0.0, fpi) - fmr = tl.where(spoil, 0.0, fmr) - fmi = tl.where(spoil, 0.0, fmi) - dfpr = tl.where(spoil, 0.0, dfpr) - dfpi = tl.where(spoil, 0.0, dfpi) - dfmr = tl.where(spoil, 0.0, dfmr) - dfmi = tl.where(spoil, 0.0, dfmi) - if pools == 2 or pools == 3: - s_bpr, s_bpi, s_bmr, s_bmi = _shift( - bpr, bpi, bmr, bmi, state, state_mask, state_count - ) - s_dbpr, s_dbpi, s_dbmr, s_dbmi = _shift( - dbpr, dbpi, dbmr, dbmi, state, state_mask, state_count - ) - bpr = tl.where(spoil, 0.0, tl.where(do_shift, s_bpr, bpr)) - bpi = tl.where(spoil, 0.0, tl.where(do_shift, s_bpi, bpi)) - bmr = tl.where(spoil, 0.0, tl.where(do_shift, s_bmr, bmr)) - bmi = tl.where(spoil, 0.0, tl.where(do_shift, s_bmi, bmi)) - dbpr = tl.where(spoil, 0.0, tl.where(do_shift, s_dbpr, dbpr)) - dbpi = tl.where(spoil, 0.0, tl.where(do_shift, s_dbpi, dbpi)) - dbmr = tl.where(spoil, 0.0, tl.where(do_shift, s_dbmr, dbmr)) - dbmi = tl.where(spoil, 0.0, tl.where(do_shift, s_dbmi, dbmi)) - - -def _pool_flag(lineshape: Any, exchanging: bool) -> int: - """Which pools a launch is to carry, as the kernels' own constexpr reads it. - - Kept in one place so a launcher cannot describe the tissue one way and the - kernel read it another. - """ - if lineshape is not None and exchanging: - return 3 - if exchanging: - return 2 - return 1 if lineshape is not None else 0 - - -# The operator table holds nine entries per voxel per distinct interval and -# the adjoint's cotangent table twelve, both float32. -_TABLE_FLOATS_PER_ROW = 9 -_BAR_FLOATS_PER_ROW = 12 - -# What the two tables may take of what the card can spare. The trajectory is -# the larger claim on the same memory and is allocated after them. -_TABLE_SHARE = 0.25 -_TABLE_FLOOR_BYTES = 64 << 20 - - -def _three_pool_table_bytes( - tissue: tuple[torch.Tensor, ...], - rows: int, - *, - problems: int | None, - dual: bool, -) -> int: - """What the tables would take -- the operator's, and the adjoint's bars. - - The operator table holds a row of voxels; the cotangent table holds a row - of problems, which is voxels times trains cut to what one chunk carries. - ``problems`` of ``None`` is a caller that builds no cotangent table. - - A ``dual`` launch stores the operator twice over, value and direction, and - pools its cotangents three times over: the value bars, the tangent bars, - and the value bars weighted by each event's own interval direction. - """ - entries = _TABLE_FLOATS_PER_ROW * (2 if dual else 1) - total = int(tissue[0].numel()) * int(rows) * entries - if problems is not None: - pooled = _BAR_FLOATS_PER_ROW * (3 if dual else 1) - total += int(problems) * int(rows) * pooled - return total * 4 - - -def _table_budget(device: torch.device) -> int: - """How many bytes the three-pool tables may claim on this device.""" - if device.type != "cuda": - return _TABLE_FLOOR_BYTES - free, _total = torch.cuda.mem_get_info(device) - return max(_TABLE_FLOOR_BYTES, int(free * _TABLE_SHARE)) - - -def _three_pool_table_jvp( - tissue: tuple[torch.Tensor, ...], - tangents: tuple[torch.Tensor, ...], - durations: torch.Tensor, -) -> torch.Tensor: - """The three-pool operator and a direction through it, per distinct length. - - Parameters - ---------- - tissue: - The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. - tangents: - The directions along them, in the same order. - durations: - The distinct interval lengths, in seconds. - - Returns - ------- - torch.Tensor - ``(rows, 18, voxels)`` float32, undamped and at ``d_dt`` of zero. - """ - voxels = int(tissue[0].numel()) - rows = int(durations.numel()) - table = torch.empty( - (rows, 18, voxels), dtype=torch.float32, device=tissue[0].device - ) - order = (0, 14, 11, 13, 10, 12, 9) - block = min(1024, triton.next_power_of_2(max(voxels, 1))) - spread = durations.abs().to(torch.float64) * three_pool_spread_rate(tissue) - for picked, narrow in ( - (torch.nonzero(spread <= NARROW_SPREAD).flatten(), True), - (torch.nonzero(spread > NARROW_SPREAD).flatten(), False), - ): - if picked.numel() == 0: - continue - _three_pool_table_jvp_kernel[(picked.numel(), triton.cdiv(voxels, block))]( - *(tissue[index] for index in order), - *(tangents[index] for index in order), - durations.to(torch.float32), - picked.to(torch.int32), - table, - voxels, - BLOCK=block, - narrow=narrow, - num_warps=4, - ) - return table - - -def _tabulate_three_pool( - tissue: tuple[torch.Tensor, ...], - duration: torch.Tensor, - *, - pools: int, - narrow: bool, - problems: int | None = None, - tangents: tuple[torch.Tensor, ...] | None = None, -) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: - """The operator table an event loop should read, or ``None`` to form it. - - Only a wide launch has anything to gain: under ``narrow`` the operator is - already 504 float32 instructions and forming it per event costs less than - a round trip through memory. - - A wide launch is wide because of its longest interval, and one preparation - delay is enough -- so the events that pay the roots in double are mostly - events whose own length would have taken the series. Splitting the table by - row is what lets each interval take the branch its own spread asks for, - which is worth more than sharing a row between events and does not need a - row to be shared at all. - - Parameters - ---------- - tissue: - The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. - duration: - The packed event durations, in seconds. - pools: - Which pool model the launch carries. - narrow: - Whether every interval keeps the eigenvalues close together. - problems: - How many problems a chunk of the adjoint carries, which is the height - of the cotangent table it allocates. ``None`` for a caller that builds - no such table. - tangents: - The directions along the tissue, for a caller that follows one. The - table then carries the direction through the operator beside its - value, at twice the width. - - Returns - ------- - tuple - The per-event row index, the table and the distinct lengths, or - ``(None, None, None)``. - """ - if pools != 3 or narrow: - return None, None, None - distinct, inverse = torch.unique(duration.detach(), return_inverse=True) - # A row costs a formation, an event costs one too, so a train whose - # lengths are all different has nothing to gain and a table to write. - if distinct.numel() >= duration.numel(): - return None, None, None - if _three_pool_table_bytes( - tissue, distinct.numel(), problems=problems, dual=tangents is not None - ) > _table_budget(tissue[0].device): - # A pathological train has as many lengths as events, and the tables - # grow with their product. Forming the operator per event is slower - # and always fits, so that is what an unbounded one falls back to. - return None, None, None - lengths = distinct.to(torch.float32).contiguous() - if tangents is not None: - built = _three_pool_table_jvp(tissue, tangents, lengths) - else: - built = _three_pool_table(tissue, lengths) - return inverse.reshape(duration.shape).to(torch.int32), built, lengths - - -def _three_pool_table( - tissue: tuple[torch.Tensor, ...], durations: torch.Tensor -) -> torch.Tensor: - """The three-pool operator for each distinct interval, over every voxel. - - Parameters - ---------- - tissue: - The prepared per-voxel buffers, in ``TISSUE_NAMES`` order. - durations: - The distinct interval lengths, in seconds. - - Returns - ------- - torch.Tensor - ``(rows, 9, voxels)`` float32, undamped -- the reading event applies - its own washout. - """ - ( - t1, - _t2, - _m0, - _b1, - _b1_phase, - _b0, - _inversion, - _diffusion, - _velocity, - bound_fraction, - bound_exchange, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - _t2_pool_b, - _pool_b_shift, - ) = tissue - voxels = t1.numel() - rows = durations.numel() - table = torch.empty((rows, 9, voxels), dtype=torch.float32, device=t1.device) - # The spread a row reaches is its own length times the rate, so the split - # is exact per row rather than one verdict for the whole table. - spread = durations.abs().to(torch.float64) * three_pool_spread_rate(tissue) - block = min(1024, triton.next_power_of_2(max(voxels, 1))) - narrow_rows = torch.nonzero(spread <= NARROW_SPREAD, as_tuple=False).flatten() - wide_rows = torch.nonzero(spread > NARROW_SPREAD, as_tuple=False).flatten() - for picked, narrow in ((narrow_rows, True), (wide_rows, False)): - if picked.numel() == 0: - continue - _three_pool_table_kernel[(picked.numel(), triton.cdiv(voxels, block))]( - t1, - t1_pool_b, - t1_bound, - pool_b_exchange, - bound_exchange, - pool_b_fraction, - bound_fraction, - durations.to(torch.float32), - picked.to(torch.int32), - table, - voxels, - BLOCK=block, - narrow=narrow, - num_warps=4, - ) - return table - - -# Elements of the state tile one program carries. Four to a lane keeps enough -# independent arithmetic in flight to cover the latency of the chain the event -# loop is, and a wider tile spends registers without buying more of it; both -# halves of that were measured over 8 to 64 configuration orders. -_TILE_ELEMENTS = 64 - - -def _atom_stride(*tuples: tuple[torch.Tensor, ...]) -> int: - """How far to step through a property to reach one voxel's value. - - Zero where every optional property was given as one value for the whole - tissue: each is then read at one address by every voxel and needs no room - per voxel. The relaxation times lead each tuple and are stepped by one - whatever this says, since a tissue is its two relaxation times before it is - anything else. - - One stride serves the values and the directions followed beside them, so a - pass carrying tangents is asked about both: a direction laid out per voxel - has to be stepped through even where the value it follows is one number. - """ - return ( - 0 if all(value.numel() <= 1 for values in tuples for value in values[2:]) else 1 - ) - - -def _problems_per_program(block_states: int) -> int: - """How many independent problems to carry on one program's lane axis. - - A warp's lanes cost about the same whether they are used or not, so packing - several problems into one program is close to free. - - It depends on the state count alone, and deliberately not on how many - problems the launch has. A run cut into chunks would otherwise compile a - different tile from the same run whole, and the two tiles reassociate their - arithmetic differently -- so a streamed volume would answer a little - differently from an unstreamed one, which is a difference a caller has no - way to account for. ``tests/sequence/test_both_pools.py`` pins that. - - The result indexes a ``tl.arange``, so it must be a power of two. - """ - widest = max(1, _TILE_ELEMENTS // block_states) - return 1 << (widest.bit_length() - 1) - - -def _output_shape( - train_count: int, atom_count: int, output_count: int -) -> tuple[int, ...]: - """Signal shape, matching what the CPU kernels return.""" - if train_count == 1: - return (atom_count, output_count) - return (train_count, atom_count, output_count) - - -def _only_scalars(flags: dict) -> dict: - """The switches a real-subspace kernel takes. - - Off-resonance and flow are not in its representation to begin with -- it - carries three real planes where the complex kernels carry four -- so it is - given the terms that survive that reduction and no others. - """ - return { - name: flags[name] for name in ("diffusing", "transmit", "density", "inverting") - } - - -def simulate( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - features: frozenset[str] | None = None, - pools: Any = None, -) -> torch.Tensor: - """Run a packed state machine on CUDA and return complex signals. - - ``real_axis`` of 1 selects the real-subspace kernel; see - ``real_subspace_axis`` for when that is legitimate. - """ - if pools is not None: - from . import _pools_triton - - return _pools_triton.simulate( - tissue, - events, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - features=features, - pools=pools, - ) - train_count = _train_count(events) - atom_count = tissue[0].numel() - output_real = torch.empty( - _output_shape(train_count, atom_count, output_count), - dtype=torch.float32, - device=tissue[0].device, - ) - output_imag = torch.empty_like(output_real) - simulate_into( - tissue, - events, - output_real, - output_imag, - state_count=state_count, - output_count=output_count, - real_axis=real_axis, - atom_count=atom_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - exchanging=exchanging, - dynamic=dynamic, - features=features, - ) - return torch.complex(output_real, output_imag) - - -def simulate_into( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - output_real: torch.Tensor, - output_imag: torch.Tensor, - *, - state_count: int, - output_count: int, - real_axis: int | None, - atom_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - features: frozenset[str] | None = None, -) -> None: - """Run the forward machine into buffers the caller owns. - - Streaming reuses one set of buffers per chunk, so allocating here would put - an allocation in the loop -- and an allocation that reaches ``cudaMalloc`` - synchronizes the device, which is exactly what the streams exist to avoid. - - ``atom_count`` is given rather than taken from ``tissue`` because a chunk's - buffers are sized for the largest chunk and the last one is shorter. - """ - ( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - bound_exchange, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - ) = tissue - ( - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - ) = events - train_count = _train_count(events) - shims = _shim_count(tissue) - pools = _pool_flag(lineshape, exchanging) - block_states = triton.next_power_of_2(state_count) - total = train_count * atom_count - problems = _problems_per_program(block_states) - grid = (triton.cdiv(total, problems),) - # A kernel argument has to be a tensor even where the branch reading it is - # compiled out, so an unprofiled launch passes one it already has. - table = None if profile is None else profile.packed(t1.device) - pairs = None if dynamic is None else dynamic.packed(t1.device) - pair_rows = ( - None - if dynamic is None - else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) - ) - table_rows = None if profile is None else profile.rows(kind.device) - absorption = None if lineshape is None else lineshape.packed(t1.device) - narrow = narrow_three_pool(tissue, duration, pools=pools) - duration_row, pool_table, _lengths = _tabulate_three_pool( - tissue, duration, pools=pools, narrow=narrow - ) - - # Phases grow without bound under RF spoiling, so their cosines and sines - # are taken once here, in double precision, rather than in every program. - phase_cos = torch.cos(phase.double()).to(torch.float32) - phase_sin = torch.sin(phase.double()).to(torch.float32) - - if real_axis == 1: - _epg_real_kernel[grid]( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - output_real, - output_imag, - atom_count, - train_count, - kind.numel(), - output_count, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - shimmed=_shim_count(tissue) > 1, - **_only_scalars(_feature_flags(features, geometry)), - block_states=block_states, - problems=problems, - num_warps=1, - ) - return - - _epg_kernel[grid]( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - bound_exchange, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - phase_cos, - phase_sin, - action, - output_index, - shim_index, - saturation, - rf_frequency, - t1 if table is None else table, - kind if table_rows is None else table_rows, - t1 if absorption is None else absorption, - t1 if pairs is None else pairs, - kind if pair_rows is None else pair_rows, - kind if duration_row is None else duration_row, - t1 if pool_table is None else pool_table, - output_real, - output_imag, - atom_count, - train_count, - kind.numel(), - output_count, - geometry.flow_scale, - geometry.washout_scale, - 1.0 if profile is None else profile.step, - 1.0 if lineshape is None else lineshape.step, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - shim_rows=shims, - shimmed=shims > 1, - locations=1 if profile is None else profile.points, - profiled=profile is not None and profile.bins > 0, - profile_bins=0 if profile is None else profile.bins, - dynamic=dynamic is not None, - broadened=lineshape is not None and lineshape.bins > 0, - lineshape_bins=0 if lineshape is None else lineshape.bins, - pools=pools, - narrow=narrow, - tabulated=pool_table is not None, - **_feature_flags(features, geometry), - block_states=block_states, - problems=problems, - num_warps=1, - ) - - -def simulate_jvp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tissue_tangents: tuple[torch.Tensor, ...], - event_tangents: tuple[torch.Tensor, torch.Tensor, torch.Tensor], - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - dynamic_direction: Any = None, - features: frozenset[str] | None = None, - pools: Any = None, -) -> torch.Tensor: - """Run one fused state-machine Jacobian-vector product on CUDA. - - ``real_axis`` of 1 selects the real-subspace kernel, which produces no - derivative along ``b1_phase``, ``b0`` or the RF phase -- seeds along those - directions leave the subspace, so the caller must rule them out. - """ - if pools is not None: - from . import _pools_triton - - return _pools_triton.simulate_jvp( - tissue, - events, - tissue_tangents, - event_tangents, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - features=features, - pools=pools, - ) - train_count = _train_count(events) - atom_count = tissue[0].numel() - output_real = torch.empty( - _output_shape(train_count, atom_count, output_count), - dtype=torch.float32, - device=tissue[0].device, - ) - output_imag = torch.empty_like(output_real) - simulate_jvp_into( - tissue, - events, - tissue_tangents, - event_tangents, - output_real, - output_imag, - state_count=state_count, - output_count=output_count, - real_axis=real_axis, - atom_count=atom_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - exchanging=exchanging, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - features=features, - ) - return torch.complex(output_real, output_imag) - - -def simulate_jvp_into( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tissue_tangents: tuple[torch.Tensor, ...], - event_tangents: tuple[torch.Tensor, torch.Tensor, torch.Tensor], - output_real: torch.Tensor, - output_imag: torch.Tensor, - *, - state_count: int, - output_count: int, - real_axis: int | None, - atom_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - dynamic_direction: Any = None, - features: frozenset[str] | None = None, -) -> None: - """Run one Jacobian-vector product into buffers the caller owns. - - See ``simulate_into`` for why the streaming path needs this. - """ - ( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - bound_exchange, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - ) = tissue - ( - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - ) = events - tangent_duration, tangent_flip, tangent_phase = event_tangents - train_count = _train_count(events) - pools = _pool_flag(lineshape, exchanging) - shims = _shim_count(tissue) - block_states = triton.next_power_of_2(state_count) - total = train_count * atom_count - problems = _problems_per_program(block_states) - grid = (triton.cdiv(total, problems),) - - if real_axis == 1: - _epg_real_jvp_kernel[grid]( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - tissue_tangents[0], - tissue_tangents[1], - tissue_tangents[2], - tissue_tangents[3], - tissue_tangents[6], - tissue_tangents[7], - tangent_duration, - tangent_flip, - output_real, - output_imag, - atom_count, - train_count, - kind.numel(), - output_count, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue, tissue_tangents), - shimmed=shims > 1, - **_only_scalars(_feature_flags(features, geometry)), - block_states=block_states, - problems=problems, - num_warps=1, - ) - return - - table = None if profile is None else profile.packed(t1.device) - pairs = None if dynamic is None else dynamic.packed(t1.device) - pair_rows = ( - None - if dynamic is None - else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) - ) - pair_direction = ( - None if dynamic_direction is None else dynamic_direction.to(t1.device) - ) - table_rows = None if profile is None else profile.rows(kind.device) - absorption = None if lineshape is None else lineshape.packed(t1.device) - narrow = narrow_three_pool(tissue, duration, pools=pools) - duration_row, pool_table, _lengths = _tabulate_three_pool( - tissue, duration, pools=pools, narrow=narrow, tangents=tissue_tangents - ) - _epg_jvp_kernel[grid]( - *tissue, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - *tissue_tangents, - tangent_duration, - tangent_flip, - tangent_phase, - saturation, - rf_frequency, - t1 if table is None else table, - kind if table_rows is None else table_rows, - t1 if absorption is None else absorption, - t1 if pairs is None else pairs, - kind if pair_rows is None else pair_rows, - t1 if pair_direction is None else pair_direction, - kind if duration_row is None else duration_row, - t1 if pool_table is None else pool_table, - output_real, - output_imag, - atom_count, - train_count, - kind.numel(), - output_count, - geometry.flow_scale, - geometry.washout_scale, - 1.0 if profile is None else profile.step, - 1.0 if lineshape is None else lineshape.step, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue, tissue_tangents), - shim_rows=shims, - shimmed=shims > 1, - locations=1 if profile is None else profile.points, - profiled=profile is not None and profile.bins > 0, - profile_bins=0 if profile is None else profile.bins, - dynamic=dynamic is not None, - broadened=lineshape is not None and lineshape.bins > 0, - lineshape_bins=0 if lineshape is None else lineshape.bins, - pools=pools, - narrow=narrow, - tabulated=pool_table is not None, - **_feature_flags(features, geometry), - block_states=block_states, - problems=problems, - num_warps=1, - ) - - -# How much device memory the recorded trajectory may hold at once. Beyond this -# the problems are run in waves, which the gradient buffers absorb because they -# accumulate rather than being written. -_TRAJECTORY_BUDGET_BYTES = 256 << 20 - - -def _trajectory_wave( - event_count: int, state_count: int, total: int, planes: int, blocks: int = 3 -) -> int: - """How many problems can record their trajectory in one launch.""" - per_problem = event_count * blocks * state_count * planes * 4 - return max(1, min(total, _TRAJECTORY_BUDGET_BYTES // max(1, per_problem))) - - -def simulate_vjp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - dynamic: Any = None, - lineshape: Any = None, - exchanging: bool = False, - features: frozenset[str] | None = None, - pools: Any = None, -) -> tuple[torch.Tensor, ...]: - """The first-order adjoint on CUDA, for a whole volume on one device. - - Returns the gradients in the differentiable-input order -- every tissue - property, then event duration, flip and phase, and the pair's cotangent - where one is given. A shard takes this same kernel a level down and a - streamed volume has chunked launchers of its own; - :func:`blochsim.sequence._accelerators` decides which route a run takes. - - Carrying no forward direction, this records two trajectory planes per - recorded state where that pass records four, and holds one state where it - holds a dual. - """ - if pools is not None: - from . import _pools_triton - - return _pools_triton.simulate_vjp( - tissue, - events, - grad_output, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - dynamic=dynamic, - lineshape=lineshape, - features=features, - pools=pools, - ) - ( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - ) = tissue - ( - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - ) = events[:9] - atom_count = t1.numel() - train_count = _train_count(events) - event_count = kind.numel() - total = train_count * atom_count - block_states = triton.next_power_of_2(state_count) - device = t1.device - shims = max(1, b1.numel() // atom_count) if atom_count else 1 - table = None if profile is None else profile.packed(device) - table_rows = None if profile is None else profile.rows().to(device) - pairs = None if dynamic is None else dynamic.packed(device) - pair_rows = ( - None - if dynamic is None - else dynamic.rows_per_event(train_count, event_count).to(device) - ) - grad_pair = None if dynamic is None else torch.zeros_like(pairs) - locations = 1 if profile is None else profile.points - absorption = None if lineshape is None else lineshape.packed(device) - - grad_tissue = torch.zeros( - tissue_gradient_height(shims) * atom_count, - dtype=torch.float32, - device=device, - ) - grad_flip = torch.zeros_like(flip) - grad_phase = torch.zeros_like(phase) - grad_duration = torch.zeros_like(duration) - grad_output = grad_output.resolve_conj() - grad_real = grad_output.real.contiguous() - grad_imag = grad_output.imag.contiguous() - - # A semisolid pool records a plane of its own beside the three the free - # water keeps; a chemically exchanging one three, and the two together - # four. - pools = _pool_flag(lineshape, exchanging) - narrow = narrow_three_pool(tissue, duration, pools=pools) - blocks = 7 if pools == 3 else (6 if pools == 2 else (4 if pools == 1 else 3)) - wave = _trajectory_wave(event_count, state_count, total, 2, blocks) - duration_row, pool_table, pool_durations = _tabulate_three_pool( - tissue, duration, pools=pools, narrow=narrow, problems=wave - ) - row_count = 0 if pool_durations is None else pool_durations.numel() - pool_bars = None - if pool_table is not None: - # A slot per problem the chunk carries, so the walk back accumulates - # into memory it owns and no two programs contend for a row. The - # chunks run one after another, so one chunk's worth is enough. - pool_bars = torch.zeros( - wave * row_count * 12, dtype=torch.float32, device=device - ) - trajectory = [ - torch.empty( - (wave, event_count * blocks * state_count), - dtype=torch.float32, - device=device, - ) - for _ in range(2) - ] - - problems = _problems_per_program(block_states) - for base in range(0, total, wave): - span = min(wave, total - base) - if pool_bars is not None: - # The slots are per chunk, so each chunk starts from nothing. - pool_bars.zero_() - # The trajectory is written by one launch and walked back by the - # next, so each compiles one sweep instead of both. - for recording in (True, False): - _epg_vjp_kernel[(triton.cdiv(span, problems),)]( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - absorption, - table, - table_rows, - pairs, - pair_rows, - duration_row, - pool_table, - pool_bars, - pool_durations, - row_count, - grad_pair, - grad_real, - grad_imag, - grad_tissue, - grad_flip, - grad_phase, - grad_duration, - *trajectory, - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - geometry.flow_scale, - geometry.washout_scale, - shims, - 1.0 if profile is None else profile.step, - 1.0 if lineshape is None else lineshape.step, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - shimmed=shims > 1, - locations=locations, - profiled=profile is not None and profile.bins > 0, - profile_bins=0 if profile is None else profile.bins, - dynamic=dynamic is not None, - broadened=lineshape is not None and lineshape.bins > 0, - lineshape_bins=0 if lineshape is None else lineshape.bins, - pools=pools, - narrow=narrow, - tabulated=pool_table is not None, - recording=recording, - block_states=block_states, - problems=problems, - num_warps=1, - **_feature_flags(features, geometry), - ) - voxel = tuple( - grad_tissue[base * atom_count : (base + rows) * atom_count] - for base, rows in zip( - tissue_gradient_bases(shims), tissue_gradient_rows(shims), strict=True - ) - ) - if dynamic is not None: - return (*voxel, grad_duration, grad_flip, grad_phase, grad_pair) - return (*voxel, grad_duration, grad_flip, grad_phase) - - -def simulate_real_vjp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - features: frozenset[str] | None = None, -) -> tuple[torch.Tensor, ...]: - """The first-order adjoint through the real subspace, on CUDA. - - Returns the gradients in the differentiable-input order -- every tissue - property, then event duration, flip and phase. The representation divides - the RF phase out, so transmit phase, off-resonance, velocity and RF phase - come back at zero and callers must not ask for those. - - Carrying no forward direction, this records one trajectory plane where the - forward-over-reverse pass records two, and holds one state where it holds a - dual. - """ - ( - t1, - t2, - m0, - b1, - _b1_phase, - _b0, - inversion_efficiency, - diffusion, - *_rest, - ) = tissue - duration, kind, flip, phase, action, output_index, shim_index = events[:7] - atom_count = t1.numel() - train_count = _train_count(events) - event_count = kind.numel() - total = train_count * atom_count - block_states = triton.next_power_of_2(state_count) - device = t1.device - shims = _shim_count(tissue) - - grad_tissue = torch.zeros( - tissue_gradient_height(shims) * atom_count, - dtype=torch.float32, - device=device, - ) - grad_flip = torch.zeros_like(flip) - grad_duration = torch.zeros_like(duration) - grad_phase = torch.zeros_like(phase) - grad_imag = grad_output.resolve_conj().imag.contiguous() - - wave = _trajectory_wave(event_count, state_count, total, 1) - trajectory = torch.empty( - (wave, event_count * 3 * state_count), dtype=torch.float32, device=device - ) - - problems = _problems_per_program(block_states) - for base in range(0, total, wave): - span = min(wave, total - base) - _epg_real_vjp_kernel[(triton.cdiv(span, problems),)]( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - grad_imag, - grad_tissue, - grad_flip, - grad_duration, - trajectory, - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - shim_rows=shims, - shimmed=shims > 1, - **_only_scalars(_feature_flags(features, NO_GEOMETRY)), - block_states=block_states, - problems=problems, - num_warps=1, - ) - voxel = tuple( - grad_tissue[base * atom_count : (base + rows) * atom_count] - for base, rows in zip( - tissue_gradient_bases(shims), tissue_gradient_rows(shims), strict=True - ) - ) - return (*voxel, grad_duration, grad_flip, grad_phase) - - -class AdjointBuffers: - """Device memory a forward-over-reverse pass writes into. - - Sized for ``chunk`` voxels and reusable for any narrower one. Per-voxel - gradients are cleared before each pass; per-event gradients accumulate over - every pass the buffers serve and are read out with ``event_gradients``. - - ``real_axis`` of 1 halves the state planes, so buffers built for one - representation cannot be handed to the other. - """ - - def __init__( - self, - events: tuple[torch.Tensor, ...], - chunk: int, - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - shims: int = 1, - pools: int = 0, - ) -> None: - ( - duration, - kind, - flip, - phase, - _action, - _output_index, - _shim, - _saturation, - _rf_frequency, - ) = events - device = kind.device - train_count = _train_count(events) - event_count = kind.numel() - self.planes = 2 if real_axis == 1 else 4 - # A bound pool records a fourth block of states per event: the RF - # operator scales it, so the reverse sweep cannot replay it from the - # free pool's. - self.blocks = 3 + (4 if pools == 3 else (3 if pools == 2 else pools)) - self.chunk = chunk - self.shims = shims - self.rows = tissue_gradient_height(shims) - self.state_count = state_count - self.output_count = output_count - self.train_count = train_count - # One dual accumulator per plane: value is the gradient w.r.t. the - # tangent inputs, tangent the gradient w.r.t. the primal ones. - self.tissue = [ - torch.zeros(self.rows * chunk, dtype=torch.float32, device=device) - for _ in range(2) - ] - self.flip = [torch.zeros_like(flip) for _ in range(2)] - self.duration = [torch.zeros_like(duration) for _ in range(2)] - self.phase = [torch.zeros_like(phase) for _ in range(2)] - self.cotangent = [ - torch.empty( - train_count * chunk * output_count, - dtype=torch.float32, - device=device, - ) - for _ in range(2) - ] - self.wave = _trajectory_wave( - event_count, - state_count, - train_count * chunk, - self.planes, - self.blocks, - ) - self.trajectory = [ - torch.empty( - (self.wave, event_count * self.blocks * state_count), - dtype=torch.float32, - device=device, - ) - for _ in range(self.planes) - ] - - def tissue_gradients(self, atom_count: int) -> tuple[tuple[torch.Tensor, ...], ...]: - """The per-voxel gradients of the last pass, one entry per parameter. - - Each is flat and as wide as the buffer it belongs to, so the transmit - pair spans every shim. Ordered to match ``event_gradients``: tangent - plane first. - """ - return tuple( - tuple( - self.tissue[plane][base * atom_count : (base + rows) * atom_count] - for base, rows in zip( - tissue_gradient_bases(self.shims), - tissue_gradient_rows(self.shims), - strict=True, - ) - ) - for plane in (1, 0) - ) - - def event_gradients(self) -> tuple[tuple[torch.Tensor, ...], ...]: - """The per-event gradients summed over every pass so far. - - Ordered ``(duration, flip, phase)`` to match the tail of the - differentiable-input order, tangent plane first. - """ - return tuple( - (self.duration[plane], self.flip[plane], self.phase[plane]) - for plane in (1, 0) - ) - - -class GradientBuffers: - """Device memory a first-order adjoint writes into. - - Half of what the forward-over-reverse pass needs: one accumulator per - gradient rather than a dual, and one trajectory plane per real state rather - than a plane per component of one. Sized for ``chunk`` voxels and reusable - for any narrower one; per-event gradients accumulate over every pass the - buffers serve. - - ``real_axis`` of 1 halves the planes again, so buffers built for one - representation cannot be handed to the other. - """ - - def __init__( - self, - events: tuple[torch.Tensor, ...], - chunk: int, - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - ) -> None: - duration, kind, flip, phase = events[:4] - device = kind.device - train_count = _train_count(events) - event_count = kind.numel() - self.real_axis = real_axis - self.planes = 1 if real_axis == 1 else 2 - self.chunk = chunk - self.rows = tissue_gradient_height(1) - self.state_count = state_count - self.output_count = output_count - self.train_count = train_count - self.tissue = torch.zeros(self.rows * chunk, dtype=torch.float32, device=device) - self.flip = torch.zeros_like(flip) - self.duration = torch.zeros_like(duration) - self.phase = torch.zeros_like(phase) - self.cotangent = [ - torch.empty( - train_count * chunk * output_count, - dtype=torch.float32, - device=device, - ) - for _ in range(2) - ] - self.wave = _trajectory_wave( - event_count, state_count, train_count * chunk, self.planes - ) - self.trajectory = [ - torch.empty( - (self.wave, event_count * 3 * state_count), - dtype=torch.float32, - device=device, - ) - for _ in range(self.planes) - ] - - def tissue_gradients(self, atom_count: int) -> tuple[torch.Tensor, ...]: - """The per-voxel gradients of the last pass, one entry per parameter.""" - return tuple( - self.tissue[base * atom_count : (base + rows) * atom_count] - for base, rows in zip( - tissue_gradient_bases(1), tissue_gradient_rows(1), strict=True - ) - ) - - def event_gradients(self) -> tuple[torch.Tensor, ...]: - """``(duration, flip, phase)``, summed over every pass so far.""" - return (self.duration, self.flip, self.phase) - - -def simulate_vjp_into( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - buffers: GradientBuffers, - *, - state_count: int, - output_count: int, - atom_count: int, - geometry: Geometry = NO_GEOMETRY, - features: frozenset[str] | None = None, -) -> tuple[torch.Tensor, ...]: - """One chunk of a first-order adjoint, into buffers the caller owns. - - ``grad_output`` is already on the device. Returns the per-voxel gradients - of this chunk; the per-event ones accumulate in ``buffers``. - """ - ( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - ) = tissue - ( - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - ) = events[:9] - train_count = _train_count(events) - event_count = kind.numel() - total = train_count * atom_count - block_states = triton.next_power_of_2(state_count) - - buffers.tissue.zero_() - grad_output = grad_output.resolve_conj() - size = total * output_count - grad_real = buffers.cotangent[0][:size] - grad_imag = buffers.cotangent[1][:size] - grad_real.copy_(grad_output.real.reshape(-1)) - grad_imag.copy_(grad_output.imag.reshape(-1)) - - problems = _problems_per_program(block_states) - for base in range(0, total, buffers.wave): - span = min(buffers.wave, total - base) - # The trajectory is written by one launch and walked back by the - # next, so each compiles one sweep instead of both. - for recording in (True, False): - _epg_vjp_kernel[(triton.cdiv(span, problems),)]( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - bound_fraction, - exchange_rate, - t1_bound, - pool_b_fraction, - pool_b_exchange, - t1_pool_b, - t2_pool_b, - pool_b_shift, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - 0, - grad_real, - grad_imag, - buffers.tissue, - buffers.flip, - buffers.phase, - buffers.duration, - *buffers.trajectory, - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - geometry.flow_scale, - geometry.washout_scale, - 1, - 1.0, - 1.0, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - shimmed=False, - locations=1, - profiled=False, - profile_bins=0, - dynamic=False, - broadened=False, - lineshape_bins=0, - pools=0, - narrow=False, - tabulated=False, - recording=recording, - block_states=block_states, - problems=problems, - num_warps=1, - **_feature_flags(features, geometry), - ) - return buffers.tissue_gradients(atom_count) - - -def simulate_real_vjp_into( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - buffers: GradientBuffers, - *, - state_count: int, - output_count: int, - atom_count: int, - features: frozenset[str] | None = None, -) -> tuple[torch.Tensor, ...]: - """The same, for a train the real subspace covers.""" - ( - t1, - t2, - m0, - b1, - _b1_phase, - _b0, - inversion_efficiency, - diffusion, - *_rest, - ) = tissue - duration, kind, flip, _phase, action, output_index, shim_index = events[:7] - train_count = _train_count(events) - event_count = kind.numel() - total = train_count * atom_count - block_states = triton.next_power_of_2(state_count) - - buffers.tissue.zero_() - size = total * output_count - grad_imag = buffers.cotangent[1][:size] - grad_imag.copy_(grad_output.resolve_conj().imag.reshape(-1)) - - problems = _problems_per_program(block_states) - for base in range(0, total, buffers.wave): - span = min(buffers.wave, total - base) - _epg_real_vjp_kernel[(triton.cdiv(span, problems),)]( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - grad_imag, - buffers.tissue, - buffers.flip, - buffers.duration, - buffers.trajectory[0], - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue), - # The streamed route carries one shim, as the complex one does: - # ``GradientBuffers`` sizes its gradient plane for a single row. - shim_rows=1, - shimmed=False, - block_states=block_states, - problems=problems, - num_warps=1, - **_only_scalars(_feature_flags(features, NO_GEOMETRY)), - ) - return buffers.tissue_gradients(atom_count) - - -def simulate_vjp_jvp_into( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tangents: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - buffers: AdjointBuffers, - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - atom_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - dynamic_direction: Any = None, - dynamic_gradients: tuple[torch.Tensor, torch.Tensor] | None = None, - features: frozenset[str] | None = None, -) -> tuple[tuple[torch.Tensor, ...], ...]: - """Forward-over-reverse for one chunk of voxels, into caller-owned buffers. - - Returns the per-voxel gradients of this chunk -- tangent plane first, one - entry per tissue parameter, views into ``buffers`` that the next call - overwrites. The per-event gradients accumulate inside ``buffers`` instead, - because every chunk contributes to all of them. - - ``atom_count`` is this chunk's width, which may be narrower than the one - the buffers were built for. - """ - ( - t1, - t2, - m0, - b1, - b1_phase, - b0, - inversion_efficiency, - diffusion, - velocity, - _bound_fraction, - _bound_exchange, - _t1_bound, - _pool_b_fraction, - _pool_b_exchange, - _t1_pool_b, - _t2_pool_b, - _pool_b_shift, - ) = tissue - ( - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - _saturation, - _rf_frequency, - ) = events - train_count = _train_count(events) - pools = _pool_flag(lineshape, exchanging) - event_count = kind.numel() - total = train_count * atom_count - block_states = triton.next_power_of_2(state_count) - real = real_axis == 1 - - grad_output = grad_output.resolve_conj() - size = total * output_count - grad_real, grad_imag = ( - plane[:size].view(grad_output.shape) for plane in buffers.cotangent - ) - grad_real.copy_(grad_output.real) - grad_imag.copy_(grad_output.imag) - grad_tissue = [plane[: buffers.rows * atom_count] for plane in buffers.tissue] - for plane in grad_tissue: - plane.zero_() - grad_flip, grad_duration, grad_phase = buffers.flip, buffers.duration, buffers.phase - trajectory = buffers.trajectory - table = None if profile is None else profile.packed(t1.device) - pairs = None if dynamic is None else dynamic.packed(t1.device) - pair_rows = ( - None - if dynamic is None - else dynamic.rows_per_event(train_count, kind.numel()).to(t1.device) - ) - pair_direction = ( - None if dynamic_direction is None else dynamic_direction.to(t1.device) - ) - grad_pair_value = None if dynamic_gradients is None else dynamic_gradients[0] - grad_pair_tangent = None if dynamic_gradients is None else dynamic_gradients[1] - table_rows = None if profile is None else profile.rows(kind.device) - absorption = None if lineshape is None else lineshape.packed(t1.device) - narrow = narrow_three_pool(tissue, duration, pools=pools) - wave = buffers.wave - duration_row, pool_table, pool_durations = _tabulate_three_pool( - tissue, - duration, - pools=pools, - narrow=narrow, - tangents=tangents, - problems=wave, - ) - row_count = 0 if pool_durations is None else pool_durations.numel() - pool_bars = None - if pool_table is not None: - # A slot per problem the chunk carries, so the walk back accumulates - # into memory it owns and no two programs contend for a row. Three - # sets of twelve: the value cotangents, their directions, and the - # value cotangents weighted by each event's own interval direction. - pool_bars = torch.zeros( - wave * row_count * 36, dtype=torch.float32, device=t1.device - ) - - problems = _problems_per_program(block_states) - for base in range(0, total, wave): - span = min(wave, total - base) - if pool_bars is not None: - # The slots are per chunk, so each chunk starts from nothing. - pool_bars.zero_() - grid = (triton.cdiv(span, problems),) - shape = dict( - state_count=state_count, - single_train=train_count == 1, - atom_stride=_atom_stride(tissue, tangents), - block_states=block_states, - problems=problems, - num_warps=1, - ) - if real: - _epg_real_vjp_jvp_kernel[grid]( - t1, - t2, - m0, - b1, - inversion_efficiency, - diffusion, - duration, - kind, - flip, - action, - output_index, - shim_index, - tangents[0], - tangents[1], - tangents[2], - tangents[3], - tangents[6], - tangents[7], - tangents[_DURATION_SEED], - tangents[_FLIP_SEED], - grad_imag, - *grad_tissue, - *grad_flip, - *grad_duration, - *trajectory, - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - shim_rows=_shim_count(tissue), - shimmed=_shim_count(tissue) > 1, - **_only_scalars(_feature_flags(features, geometry)), - **shape, - ) - else: - # The trajectory is written by one launch and walked back by the - # next, so each compiles one sweep instead of both. - for recording in (True, False): - _epg_vjp_jvp_kernel[grid]( - *tissue, - *events, - t1 if table is None else table, - kind if table_rows is None else table_rows, - t1 if absorption is None else absorption, - t1 if pairs is None else pairs, - kind if pair_rows is None else pair_rows, - t1 if pair_direction is None else pair_direction, - t1 if grad_pair_value is None else grad_pair_value, - t1 if grad_pair_tangent is None else grad_pair_tangent, - *tangents, - kind if duration_row is None else duration_row, - t1 if pool_table is None else pool_table, - t1 if pool_bars is None else pool_bars, - t1 if pool_durations is None else pool_durations, - row_count, - grad_real, - grad_imag, - *grad_tissue, - *grad_flip, - *grad_phase, - *grad_duration, - *trajectory, - base, - base + span, - atom_count, - train_count, - event_count, - output_count, - geometry.flow_scale, - geometry.washout_scale, - 1.0 if profile is None else profile.step, - 1.0 if lineshape is None else lineshape.step, - shim_rows=_shim_count(tissue), - shimmed=_shim_count(tissue) > 1, - locations=1 if profile is None else profile.points, - profiled=profile is not None and profile.bins > 0, - profile_bins=0 if profile is None else profile.bins, - dynamic=dynamic is not None, - directed=dynamic_direction is not None, - broadened=lineshape is not None and lineshape.bins > 0, - lineshape_bins=0 if lineshape is None else lineshape.bins, - pools=pools, - narrow=narrow, - tabulated=pool_table is not None, - recording=recording, - **_feature_flags(features, geometry), - **shape, - ) - - # Plane 1 is the tangent part -> d/d(primal inputs); plane 0 the value part. - return buffers.tissue_gradients(atom_count) - - -def simulate_vjp_jvp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tangents: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - real_axis: int | None = None, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - exchanging: bool = False, - dynamic: Any = None, - dynamic_direction: Any = None, - features: frozenset[str] | None = None, - pools: Any = None, -) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: - """Forward-over-reverse through the state machine on CUDA. - - ``tangents`` follows the differentiable-input order -- every tissue - property, then event duration, flip and phase -- and the two returned - tuples, gradients with respect to the primal inputs then to the tangent - inputs, follow it too. - - ``real_axis`` of 1 selects the real-subspace adjoint. That representation - divides the RF phase out, so it leaves ``b1_phase``, ``b0`` and ``phase`` at - zero and callers must not ask for those; the complex adjoint produces every - one of them. - - Gradients land through atomic accumulation, so repeated runs agree to - floating-point tolerance rather than bit for bit. - """ - if pools is not None: - from . import _pools_triton - - return _pools_triton.simulate_vjp_jvp( - tissue, - events, - tangents, - grad_output, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - features=features, - pools=pools, - ) - atom_count = tissue[0].numel() - gradients = None - if dynamic is not None: - held = dynamic.packed(tissue[0].device) - gradients = (torch.zeros_like(held), torch.zeros_like(held)) - buffers = AdjointBuffers( - events, - atom_count, - state_count=state_count, - output_count=output_count, - real_axis=real_axis, - shims=_shim_count(tissue), - pools=_pool_flag(lineshape, exchanging), - ) - voxel_grads = simulate_vjp_jvp_into( - tissue, - events, - tangents, - grad_output, - buffers, - state_count=state_count, - output_count=output_count, - real_axis=real_axis, - atom_count=atom_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - exchanging=exchanging, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - dynamic_gradients=gradients, - features=features, - ) - sides = tuple( - (*voxels, *per_event) - for voxels, per_event in zip( - voxel_grads, buffers.event_gradients(), strict=True - ) - ) - if gradients is None: - return sides - # The value plane is the adjoint and the tangent plane its own derivative, - # which is the split the tissue gradients take; the sides come back in the - # order the caller reads them, curvature first. - return (*sides[0], gradients[1]), (*sides[1], gradients[0]) diff --git a/src/blochsim/sequence/_parameters.py b/src/blochsim/sequence/_parameters.py index dadc978b..947a5465 100644 --- a/src/blochsim/sequence/_parameters.py +++ b/src/blochsim/sequence/_parameters.py @@ -2,7 +2,7 @@ The kernels take a flat list of buffers: the tissue properties, then the packed per-event buffers. Their order is an ABI shared by the Python dispatch, the -CPU extension and the Triton kernels, and the counts derived from it -- how +CPU extension and the GPU kernels, and the counts derived from it -- how many buffers are saved for backward, which of them carry gradients, how wide the raw pointer array is -- appear in all three. This module is where that order is written down, so those counts are read from one place instead of @@ -361,9 +361,8 @@ def feature_flags(features: Any, geometry: Geometry) -> dict[str, bool]: reads off the tissue; ``None`` is a caller who did not declare, and every term stays. - Fewer switches than properties, because each Triton flag multiplies how - many kernels the cache holds and these groups are what the arithmetic - actually splits into. ``off_axis`` is the static phase a tissue puts on the + Fewer switches than properties, because these groups are what the + arithmetic actually splits into. ``off_axis`` is the static phase a tissue puts on the states -- off-resonance and transmit phase reach the interval and the pulse through the same turn. ``moving`` is what a voxel's velocity drives, which it does only through the sequence geometry: flow winding and washout are @@ -411,9 +410,9 @@ def feature_flags(features: Any, geometry: Geometry) -> dict[str, bool]: def feature_mask(features: Any, geometry: Geometry) -> int: """The same answer as :func:`feature_flags`, as the host kernels read it. - Triton takes a flag per term because each one compiles a kernel of its own; - the host kernels take one integer and branch on it at run time, which is - the same choice the pool count already makes on each side. + The GPU kernels take a flag per term; the host kernels take one integer + and branch on it at run time, which is the same choice the pool count + already makes on each side. """ flags = feature_flags(features, geometry) return sum(1 << bit for bit, name in enumerate(FEATURE_BITS) if flags[name]) diff --git a/src/blochsim/sequence/_pools_gpu.py b/src/blochsim/sequence/_pools_gpu.py new file mode 100644 index 00000000..434999e6 --- /dev/null +++ b/src/blochsim/sequence/_pools_gpu.py @@ -0,0 +1,525 @@ +"""The GPU kernels for a tissue whose pools are tabulated per interval. + +The device side of ``simulate_pooled`` and ``simulate_pooled_adjoint`` in +``_epg_cpu.cpp``, reading the tables ``blochsim.sequence._pools`` builds. One +program carries one (train, voxel) problem, and its states are tiles of pools +by dephasing orders, so an interval's relaxation and exchange is a product of +the tabulated operator with the tile and a pulse turns every exchanging pool at +once. + +Each body is written once over dual numbers. A complex dual is the tuple +``(real, imag, tangent real, tangent imag)`` and a real one ``(value, +tangent)``; with ``following`` off the helpers leave the tangents at zero and +compute none of them, so the pass that follows no direction pays for none. +""" + +from __future__ import annotations + +__all__: list[str] = [] + +from typing import Any + +import torch + +from .._gpu_launch import Kernel, next_power_of_2 +from ._accelerators import _shim_count, _train_count +from ._epg_gpu import _TRAJECTORY_BUDGET_BYTES, _atom_stride, _output_shape +from ._parameters import ( + FLOAT_NAMES, + NO_GEOMETRY, + TISSUE_NAMES, + Geometry, + tissue_gradient_bases, + tissue_gradient_height, + tissue_gradient_rows, +) +from ._parameters import feature_flags as _feature_flags + +# The seven per-voxel properties a tabulated tissue still reads, in packing +# order. Relaxation, exchange and the pools' shares are inside the tables. +_VOXEL = ( + "m0", + "b1", + "b1_phase_rad", + "b0_hz", + "inversion_efficiency", + "diffusion_um2_per_ms", + "velocity_m_per_s", +) +_VOXEL_INDEX = tuple(TISSUE_NAMES.index(name) for name in _VOXEL) + +_pooled_kernel = Kernel("_pooled_kernel") +_pooled_adjoint_kernel = Kernel("_pooled_adjoint_kernel") + + +# --------------------------------------------------------------------------- +# Launchers. +# --------------------------------------------------------------------------- + + +def _tiles(layout: Any, state_count: int) -> dict[str, int]: + """The tile shape a launch runs: pools along ``P`` and orders along ``S``.""" + pools = next_power_of_2(layout.longitudinal) + states = next_power_of_2(state_count) + return { + "n": layout.longitudinal, + "m": layout.transverse, + "blocks": layout.blocks, + "P": pools, + "S": states, + } + + +def _inputs( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + pools: Any, + profile: Any, + lineshape: Any, + dynamic: Any, + tissue_tangents: tuple[torch.Tensor, ...] | None, + event_tangents: tuple[torch.Tensor, ...] | None, + dynamic_direction: Any, +) -> tuple[torch.Tensor, ...]: + """The buffers both kernels open with, in their order. + + A buffer whose branch is compiled out is still an argument, so one the + launch already holds stands in for it. + """ + device = tissue[0].device + voxel = tuple(tissue[index] for index in _VOXEL_INDEX) + moving = ( + voxel + if tissue_tangents is None + else tuple(tissue_tangents[index] for index in _VOXEL_INDEX) + ) + duration, kind, flip, phase = events[:4] + stepping = (duration, flip, phase) if event_tangents is None else event_tangents + pairs = duration if dynamic is None else dynamic.packed(device) + return ( + *voxel, + *moving, + *events[:9], + *stepping, + pools.values, + pools.values if pools.direction is None else pools.direction, + pools.index, + duration if profile is None else profile.packed(device), + kind if profile is None else profile.rows(device), + duration if lineshape is None else lineshape.packed(device), + pairs, + kind + if dynamic is None + else dynamic.rows_per_event(_train_count(events), kind.numel()).to(device), + pairs if dynamic_direction is None else dynamic_direction.to(device), + ) + + +def _scalars( + base: int, + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + output_count: int, + state_count: int, + pools: Any, + geometry: Geometry, + profile: Any, + lineshape: Any, +) -> tuple[Any, ...]: + """The scalar arguments both kernels take after their buffers.""" + return ( + base, + tissue[0].numel(), + events[1].numel(), + output_count, + state_count, + pools.layout.rows, + geometry.flow_scale, + geometry.washout_scale, + 1.0 if profile is None else profile.step, + 1.0 if lineshape is None else lineshape.step, + 1 if profile is None else profile.points, + 0 if profile is None else profile.bins, + 0 if lineshape is None else lineshape.bins, + ) + + +def _switches( + tissue: tuple[torch.Tensor, ...], + moving: tuple[torch.Tensor, ...] | None, + pools: Any, + profile: Any, + dynamic: Any, + dynamic_direction: Any, + features: frozenset[str] | None, + geometry: Geometry, + following: bool, +) -> dict[str, Any]: + """The constexpr switches both kernels take.""" + return { + "atom_stride": _atom_stride(tissue) + if moving is None + else _atom_stride(tissue, moving), + "shimmed": _shim_count(tissue) > 1, + "profiled": profile is not None and profile.bins > 0, + "dynamic": dynamic is not None, + "directed_pairs": following and dynamic_direction is not None, + "directed_table": following and pools.direction is not None, + "following": following, + **_feature_flags(features, geometry), + } + + +def _forward( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tissue_tangents: tuple[torch.Tensor, ...] | None, + event_tangents: tuple[torch.Tensor, ...] | None, + *, + state_count: int, + output_count: int, + geometry: Geometry, + profile: Any, + lineshape: Any, + dynamic: Any, + dynamic_direction: Any, + features: frozenset[str] | None, + pools: Any, +) -> torch.Tensor: + following = tissue_tangents is not None + atoms = tissue[0].numel() + total = _train_count(events) * atoms + output_real = torch.zeros( + _output_shape(_train_count(events), atoms, output_count), + dtype=torch.float32, + device=tissue[0].device, + ) + output_imag = torch.zeros_like(output_real) + if total: + _pooled_kernel[(total,)]( + *_inputs( + tissue, + events, + pools, + profile, + lineshape, + dynamic, + tissue_tangents, + event_tangents, + dynamic_direction, + ), + output_real, + output_imag, + output_real, + *_scalars( + 0, + tissue, + events, + output_count, + state_count, + pools, + geometry, + profile, + lineshape, + ), # fmt: skip + planes=12 if following else 6, + keep=False, + **_switches( + tissue, + tissue_tangents, + pools, + profile, + dynamic, + dynamic_direction, + features, + geometry, + following, + ), + **_tiles(pools.layout, state_count), + ) + return torch.complex(output_real, output_imag) + + +def simulate( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + *, + state_count: int, + output_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + dynamic: Any = None, + features: frozenset[str] | None = None, + pools: Any, +) -> torch.Tensor: + """The signals of a tissue whose pools are tabulated, on CUDA.""" + return _forward( + tissue, + events, + None, + None, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=None, + features=features, + pools=pools, + ) + + +def simulate_jvp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tissue_tangents: tuple[torch.Tensor, ...], + event_tangents: tuple[torch.Tensor, ...], + *, + state_count: int, + output_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + dynamic: Any = None, + dynamic_direction: Any = None, + features: frozenset[str] | None = None, + pools: Any, +) -> torch.Tensor: + """The signals' derivative along a direction, for a tabulated tissue. + + ``pools.direction`` is the direction along the tables, if they move. + """ + return _forward( + tissue, + events, + tissue_tangents, + tuple(event_tangents), + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + features=features, + pools=pools, + ) + + +def _adjoint( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tangents: tuple[torch.Tensor, ...] | None, + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + geometry: Geometry, + profile: Any, + lineshape: Any, + dynamic: Any, + dynamic_direction: Any, + features: frozenset[str] | None, + pools: Any, +) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: + """Both planes of the adjoint: the cotangents, then their derivative. + + Each plane holds the tissue's gradient rows, then duration, flip and + phase, then the pair's and the tables' cotangents where there are any. + The derivative plane is zero unless ``tangents`` gives a direction. + """ + following = tangents is not None + device = tissue[0].device + atoms = tissue[0].numel() + trains = _train_count(events) + event_count = events[1].numel() + total = trains * atoms + shims = _shim_count(tissue) + tissue_tangents = None if tangents is None else tangents[: len(TISSUE_NAMES)] + event_tangents = ( + None + if tangents is None + else tuple(tangents[FLOAT_NAMES.index(name)] for name in _EVENT) + ) + inputs = _inputs( + tissue, + events, + pools, + profile, + lineshape, + dynamic, + tissue_tangents, + event_tangents, + dynamic_direction, + ) + switches = _switches( + tissue, + tissue_tangents, + pools, + profile, + dynamic, + dynamic_direction, + features, + geometry, + following, + ) + tiles = _tiles(pools.layout, state_count) + planes = 12 if following else 6 + held = planes * tiles["P"] * tiles["S"] * max(1, event_count) + wave = max(1, min(total, _TRAJECTORY_BUDGET_BYTES // (4 * held))) + trajectory = torch.empty(wave * held, dtype=torch.float32, device=device) + + grad_output = grad_output.resolve_conj() + grad_real = grad_output.real.contiguous() + grad_imag = grad_output.imag.contiguous() + duration, _kind, flip, phase = events[:4] + pairs = inputs[_PAIRS] + + def plane() -> tuple[torch.Tensor, ...]: + return ( + torch.zeros( + tissue_gradient_height(shims) * atoms, + dtype=torch.float32, + device=device, + ), + torch.zeros_like(duration), + torch.zeros_like(flip), + torch.zeros_like(phase), + torch.zeros( + (trains, *pools.values.shape), dtype=torch.float32, device=device + ), + torch.zeros_like(pairs), + ) + + value, tangent = plane(), plane() + rows = tissue_gradient_bases(shims) + for base in range(0, total, wave): + span = min(wave, total - base) + scalars = _scalars( + base, tissue, events, output_count, state_count, pools, geometry, + profile, lineshape, + ) # fmt: skip + _pooled_kernel[(span,)]( + *inputs, + grad_real, + grad_imag, + trajectory, + *scalars, + planes=planes, + keep=True, + **switches, + **tiles, + ) + _pooled_adjoint_kernel[(span,)]( + *inputs, + grad_real, + grad_imag, + *(entry for pair in zip(value, tangent, strict=True) for entry in pair), + trajectory, + *scalars, + *(rows[TISSUE_NAMES.index(name)] for name in _VOXEL), + planes=planes, + **switches, + **tiles, + ) + + def gathered(side: tuple[torch.Tensor, ...]) -> tuple[torch.Tensor, ...]: + voxel = tuple( + side[0][start * atoms : (start + count) * atoms] + for start, count in zip(rows, tissue_gradient_rows(shims), strict=True) + ) + return ( + *voxel, + side[1], + side[2], + side[3], + *(() if dynamic is None else (side[5],)), + side[4].sum(0), + ) + + return gathered(value), gathered(tangent) + + +# Where the pair sits among the buffers ``_inputs`` returns. +_PAIRS = 2 * len(_VOXEL) + 9 + 3 + 3 + 3 + +_EVENT = ("duration", "flip", "phase") + + +def simulate_vjp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + dynamic: Any = None, + lineshape: Any = None, + features: frozenset[str] | None = None, + pools: Any, +) -> tuple[torch.Tensor, ...]: + """The first-order adjoint of a tabulated tissue, on CUDA. + + Returns the tissue's gradient rows, then duration, flip and phase, then the + pair's cotangent where there is a pair and the tables' last. + """ + value, _ = _adjoint( + tissue, + events, + None, + grad_output, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=None, + features=features, + pools=pools, + ) + return value + + +def simulate_vjp_jvp( + tissue: tuple[torch.Tensor, ...], + events: tuple[torch.Tensor, ...], + tangents: tuple[torch.Tensor, ...], + grad_output: torch.Tensor, + *, + state_count: int, + output_count: int, + geometry: Geometry = NO_GEOMETRY, + profile: Any = None, + lineshape: Any = None, + dynamic: Any = None, + dynamic_direction: Any = None, + features: frozenset[str] | None = None, + pools: Any, +) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: + """Forward-over-reverse through a tabulated tissue, on CUDA. + + Returns the gradients with respect to the primal inputs -- the derivative + of the adjoint along ``tangents`` -- and then with respect to the tangent + inputs, which is the adjoint itself. + """ + value, tangent = _adjoint( + tissue, + events, + tangents, + grad_output, + state_count=state_count, + output_count=output_count, + geometry=geometry, + profile=profile, + lineshape=lineshape, + dynamic=dynamic, + dynamic_direction=dynamic_direction, + features=features, + pools=pools, + ) + return tangent, value diff --git a/src/blochsim/sequence/_pools_triton.py b/src/blochsim/sequence/_pools_triton.py deleted file mode 100644 index b5fa16ee..00000000 --- a/src/blochsim/sequence/_pools_triton.py +++ /dev/null @@ -1,2226 +0,0 @@ -"""Triton kernels for a tissue whose pools are tabulated per interval. - -The device side of ``simulate_pooled`` and ``simulate_pooled_adjoint`` in -``_epg_cpu.cpp``, reading the tables ``blochsim.sequence._pools`` builds. One -program carries one (train, voxel) problem, and its states are tiles of pools -by dephasing orders, so an interval's relaxation and exchange is a product of -the tabulated operator with the tile and a pulse turns every exchanging pool at -once. - -Each body is written once over dual numbers. A complex dual is the tuple -``(real, imag, tangent real, tangent imag)`` and a real one ``(value, -tangent)``; with ``following`` off the helpers leave the tangents at zero and -compute none of them, so the pass that follows no direction pays for none. -""" - -from __future__ import annotations - -__all__: list[str] = [] - -from typing import Any - -import torch -import triton -import triton.language as tl - -from ._accelerators import _shim_count, _train_count -from ._epg_triton import ( - _TRAJECTORY_BUDGET_BYTES, - _atom_stride, - _dynamic_pair_at, - _dynamic_pair_dual_at, - _lineshape_at_curve, - _lineshape_at_slope, - _output_shape, - _profile_pair, - _profile_pair_slope, - _profiled_pair_dual, - _rotate_spinor, - _rotate_spinor_dual, - _shift, - _shift_adjoint, - _spinor_adjoint, - _spinor_adjoint_dual, - _table_row, -) -from ._parameters import ( - FLOAT_NAMES, - NO_GEOMETRY, - TISSUE_NAMES, - Geometry, - tissue_gradient_bases, - tissue_gradient_height, - tissue_gradient_rows, -) -from ._parameters import feature_flags as _feature_flags - -# The seven per-voxel properties a tabulated tissue still reads, in packing -# order. Relaxation, exchange and the pools' shares are inside the tables. -_VOXEL = ( - "m0", - "b1", - "b1_phase_rad", - "b0_hz", - "inversion_efficiency", - "diffusion_um2_per_ms", - "velocity_m_per_s", -) -_VOXEL_INDEX = tuple(TISSUE_NAMES.index(name) for name in _VOXEL) - - -# --------------------------------------------------------------------------- -# Dual arithmetic. -# --------------------------------------------------------------------------- - - -@triton.jit -def _rmul(x, y, following: tl.constexpr): - """Two real duals multiplied.""" - if following: - return (x[0] * y[0], x[1] * y[0] + x[0] * y[1]) - else: - return (x[0] * y[0], 0.0) - - -@triton.jit -def _cmul(x, y, following: tl.constexpr): - """Two complex duals multiplied.""" - real = x[0] * y[0] - x[1] * y[1] - imag = x[0] * y[1] + x[1] * y[0] - if following: - return ( - real, - imag, - x[2] * y[0] - x[3] * y[1] + x[0] * y[2] - x[1] * y[3], - x[2] * y[1] + x[3] * y[0] + x[0] * y[3] + x[1] * y[2], - ) - else: - return (real, imag, 0.0, 0.0) - - -@triton.jit -def _cscale(r, x, following: tl.constexpr): - """A real dual times a complex one.""" - if following: - return ( - r[0] * x[0], - r[0] * x[1], - r[1] * x[0] + r[0] * x[2], - r[1] * x[1] + r[0] * x[3], - ) - else: - return (r[0] * x[0], r[0] * x[1], 0.0, 0.0) - - -@triton.jit -def _conj(x): - return (x[0], -x[1], x[2], -x[3]) - - -@triton.jit -def _cadd(x, y): - return (x[0] + y[0], x[1] + y[1], x[2] + y[2], x[3] + y[3]) - - -@triton.jit -def _polar(angle, following: tl.constexpr): - """``exp(i angle)`` for a real dual angle.""" - cosine = tl.cos(angle[0]) - sine = tl.sin(angle[0]) - if following: - return (cosine, sine, -sine * angle[1], cosine * angle[1]) - else: - return (cosine, sine, 0.0, 0.0) - - -@triton.jit -def _re_dot(x, y, following: tl.constexpr): - """``Re(conj(x) y)`` as a real dual.""" - value = x[0] * y[0] + x[1] * y[1] - if following: - return (value, x[2] * y[0] + x[3] * y[1] + x[0] * y[2] + x[1] * y[3]) - else: - return (value, 0.0) - - -@triton.jit -def _total(value, mask): - """The sum of a tile over the entries ``mask`` keeps.""" - return tl.sum(tl.sum(tl.where(mask, value, 0.0), axis=1), axis=0) - - -@triton.jit -def _rtotal(x, mask, following: tl.constexpr): - """A real dual tile summed to a real dual scalar.""" - if following: - return (_total(x[0], mask), _total(x[1], mask)) - else: - return (_total(x[0], mask), 0.0) - - -# --------------------------------------------------------------------------- -# Operators on pool tiles. -# --------------------------------------------------------------------------- - - -@triton.jit -def _times(operator, planes, transposed: tl.constexpr): - """``operator @ planes`` over the pools, or its transpose's.""" - if transposed: - return tl.sum(operator[:, :, None] * planes[:, None, :], axis=0) - else: - return tl.sum(operator[:, :, None] * planes[None, :, :], axis=1) - - -@triton.jit -def _outer(left, right): - """``sum_k left[i, k] right[j, k]``: the operator a pair of tiles makes.""" - return tl.sum(left[:, None, :] * right[None, :, :], axis=2) - - -@triton.jit -def _apply(operator, planes, conjugate: tl.constexpr, transposed: tl.constexpr, - following: tl.constexpr): # fmt: skip - """A complex dual operator applied to complex dual pool tiles.""" - er = operator[0] - ei = operator[1] - if conjugate: - ei = -ei - real = _times(er, planes[0], transposed) - _times(ei, planes[1], transposed) - imag = _times(er, planes[1], transposed) + _times(ei, planes[0], transposed) - if following: - etr = operator[2] - eti = operator[3] - if conjugate: - eti = -eti - tangent_real = ( - _times(etr, planes[0], transposed) - - _times(eti, planes[1], transposed) - + _times(er, planes[2], transposed) - - _times(ei, planes[3], transposed) - ) - tangent_imag = ( - _times(etr, planes[1], transposed) - + _times(eti, planes[0], transposed) - + _times(er, planes[3], transposed) - + _times(ei, planes[2], transposed) - ) - return (real, imag, tangent_real, tangent_imag) - else: - return (real, imag, 0.0, 0.0) - - -@triton.jit -def _apply_real(operator, planes, transposed: tl.constexpr, following: tl.constexpr): - """A real dual operator applied to complex dual pool tiles.""" - real = _times(operator[0], planes[0], transposed) - imag = _times(operator[0], planes[1], transposed) - if following: - return ( - real, - imag, - _times(operator[1], planes[0], transposed) - + _times(operator[0], planes[2], transposed), - _times(operator[1], planes[1], transposed) - + _times(operator[0], planes[3], transposed), - ) - else: - return (real, imag, 0.0, 0.0) - - -@triton.jit -def _couter(left, right, following: tl.constexpr): - """``sum_k left[i, k] right[j, k]`` of two complex dual tiles.""" - real = _outer(left[0], right[0]) - _outer(left[1], right[1]) - imag = _outer(left[0], right[1]) + _outer(left[1], right[0]) - if following: - return ( - real, - imag, - _outer(left[2], right[0]) - - _outer(left[3], right[1]) - + _outer(left[0], right[2]) - - _outer(left[1], right[3]), - _outer(left[2], right[1]) - + _outer(left[3], right[0]) - + _outer(left[0], right[3]) - + _outer(left[1], right[2]), - ) - else: - return (real, imag, 0.0, 0.0) - - -@triton.jit -def _read(values, directions, at, live: tl.constexpr, identity, - following: tl.constexpr): # fmt: skip - """A buffer entry and the direction along it, or its identity when absent.""" - if live: - value = tl.load(values + at) - if following: - return (value, tl.load(directions + at)) - else: - return (value, 0.0) - else: - return (identity, 0.0) - - -@triton.jit -def _entries(slot, directions, offset, mask, along, sloped_at, - directed: tl.constexpr, sloped: tl.constexpr, - following: tl.constexpr): # fmt: skip - """Table entries at ``offset``, moving with the table's direction and slope. - - An event reads the row of its own interval length; a pass following a - direction in that length moves every entry along the row's slope, which - sits ``sloped_at`` further into the table. - """ - value = tl.load(slot + offset, mask=mask, other=0.0) - if following: - tangent = value * 0.0 - if directed: - tangent += tl.load(directions + offset, mask=mask, other=0.0) - if sloped: - tangent += tl.load(slot + sloped_at + offset, mask=mask, other=0.0) * along - return (value, tangent) - else: - return (value, 0.0) - - -@triton.jit -def _operators(slot, directions, row_offset, along, slope_offset, pool, column, - n: tl.constexpr, m: tl.constexpr, directed: tl.constexpr, - sloped: tl.constexpr, following: tl.constexpr): # fmt: skip - """One interval's longitudinal operator, its recovery and transverse operator.""" - longitudinal = _entries( - slot, - directions, - row_offset + pool * n + column, - (pool < n) & (column < n), - along, - slope_offset, - directed, - sloped, - following, - ) - restored = _entries( - slot, - directions, - row_offset + n * n + pool, - pool < n, - along, - slope_offset, - directed, - sloped, - following, - ) - across = row_offset + n * n + n + 2 * (pool * m + column) - carried = (pool < m) & (column < m) - real = _entries( - slot, directions, across, carried, along, slope_offset, directed, sloped, - following, - ) # fmt: skip - imag = _entries( - slot, directions, across + 1, carried, along, slope_offset, directed, sloped, - following, - ) # fmt: skip - return longitudinal, restored, (real[0], imag[0], real[1], imag[1]) - - -@triton.jit -def _factors(dt, damping_rate, b0, flow_rate, washout_rate, order, - off_axis: tl.constexpr, moving: tl.constexpr, - diffusing: tl.constexpr, following: tl.constexpr): # fmt: skip - """What an interval does to every pool alike, per dephasing order. - - Returns the fraction washout leaves, the transverse and longitudinal - factors before washout, and the per-order damping weights. - """ - squared = order * order - transverse_weight = squared + order + 0.3333333333333333 - damp_z = (order * 0.0 + 1.0, 0.0) - damp_t = (order * 0.0 + 1.0, 0.0) - if diffusing: - b_factor = _rmul(damping_rate, dt, following) - z = tl.exp(-squared * b_factor[0]) - t = tl.exp(-transverse_weight * b_factor[0]) - if following: - damp_z = (z, z * (-squared * b_factor[1])) - damp_t = (t, t * (-transverse_weight * b_factor[1])) - else: - damp_z = (z, 0.0) - damp_t = (t, 0.0) - wout = (1.0, 0.0) - if moving: - fraction = washout_rate[0] * dt[0] - left = 1.0 - tl.minimum(fraction, 1.0) - if following: - wout = ( - left, - tl.where( - fraction < 1.0, -(washout_rate[1] * dt[0] + washout_rate[0] * dt[1]), - 0.0, - ), - ) # fmt: skip - else: - wout = (left, 0.0) - unit_t = (damp_t[0], damp_t[0] * 0.0, damp_t[1], 0.0) - unit_z = (damp_z[0], damp_z[0] * 0.0, damp_z[1], 0.0) - if off_axis or moving: - angle = _rmul((-6.283185307179586, 0.0), _rmul(b0, dt, following), following) - turn = (0.0, 0.0) - if moving: - turn = _rmul(flow_rate, dt, following) - half = -(order + 0.5) - theta = (angle[0] + half * turn[0], angle[1] + half * turn[1]) - unit_t = _cscale(damp_t, _polar(theta, following), following) - if moving: - unit_z = _cscale( - damp_z, _polar((-order * turn[0], -order * turn[1]), following), - following, - ) # fmt: skip - if following: - # Every tangent a tile, so the operator products can take them. - unit_t = ( - unit_t[0], - unit_t[1], - unit_t[2] + order * 0.0, - unit_t[3] + order * 0.0, - ) - unit_z = ( - unit_z[0], - unit_z[1], - unit_z[2] + order * 0.0, - unit_z[3] + order * 0.0, - ) - return wout, unit_t, unit_z, squared, transverse_weight - - -@triton.jit -def _relax(plus, minus, longitudinal, transverse_op, longitudinal_op, restored, - equilibrium, wout, carried, spin, state, following: tl.constexpr): # fmt: skip - """One interval's relaxation and exchange over every order. - - Returns the three states it leaves and the operator products before the - per-order factors, which the adjoint reuses. - """ - mixed_plus = _apply(transverse_op, plus, False, False, following) - mixed_minus = _apply(transverse_op, minus, True, False, following) - mixed_z = _apply_real(longitudinal_op, longitudinal, False, following) - out_plus = _cmul(carried, mixed_plus, following) - out_minus = _cmul(_conj(carried), mixed_minus, following) - out_z = _cmul(spin, mixed_z, following) - # Inflowing spins arrive at equilibrium, so washout scales what the pools - # held and not what they recover towards. - origin = state == 0 - grown = equilibrium[0] - wout[0] * restored[0] - if following: - grown_tangent = equilibrium[1] - (wout[1] * restored[0] + wout[0] * restored[1]) - out_z = ( - out_z[0] + tl.where(origin, grown, 0.0), - out_z[1], - out_z[2] + tl.where(origin, grown_tangent, 0.0), - out_z[3], - ) - else: - out_z = (out_z[0] + tl.where(origin, grown, 0.0), out_z[1], 0.0, 0.0) - return out_plus, out_minus, out_z, mixed_plus, mixed_minus, mixed_z - - -@triton.jit -def _hard_pair(alpha, phi, following: tl.constexpr): - """A hard pulse as its Cayley-Klein pair, with the pair's slope in the flip. - - ``a = cos(alpha / 2)`` and ``b = -i sin(alpha / 2) exp(-i phi)``, the - rotation ``_rotate_flip_phase`` performs. - """ - half = 0.5 * alpha[0] - cosine = tl.cos(half) - sine = tl.sin(half) - nothing = cosine * 0.0 - turn = _polar((-phi[0], -phi[1]), following) - if following: - a = (cosine, nothing, -0.5 * sine * alpha[1], nothing) - b = (nothing, -sine, nothing, -0.5 * cosine * alpha[1]) - slope_a = (-0.5 * sine, nothing, -0.25 * cosine * alpha[1], nothing) - slope_b = (nothing, -0.5 * cosine, nothing, 0.25 * sine * alpha[1]) - else: - a = (cosine, nothing, 0.0, 0.0) - b = (nothing, -sine, 0.0, 0.0) - slope_a = (-0.5 * sine, nothing, 0.0, 0.0) - slope_b = (nothing, -0.5 * cosine, 0.0, 0.0) - return a, _cmul(b, turn, following), slope_a, _cmul(slope_b, turn, following), turn - - -@triton.jit -def _absorption(lineshape, rf_frequency, saturation, event, alpha, b0, - lineshape_bins, lineshape_step, following: tl.constexpr): # fmt: skip - """The semisolid pool's saturation by a pulse, and the lineshape it read. - - Returns ``exp(saturation * alpha^2 * G(offset))``, ``G`` and its slope in - the offset, each a real dual. - """ - offset = tl.load(rf_frequency + event) - b0[0] - deposited = tl.load(saturation + event) - if following: - shape, slope, curve = _lineshape_at_curve( - lineshape, offset, lineshape_bins, lineshape_step - ) - # The lineshape is read at the pulse's offset from the voxel, so a - # step in the voxel's own off-resonance moves the read the other way. - shape_dual = (shape, slope * -b0[1]) - slope_dual = (slope, curve * -b0[1]) - else: - shape, slope = _lineshape_at_slope( - lineshape, offset, lineshape_bins, lineshape_step - ) - shape_dual = (shape, 0.0) - slope_dual = (slope, 0.0) - exponent = _rmul( - (deposited, 0.0), - _rmul(_rmul(alpha, alpha, following), shape_dual, following), - following, - ) - absorbed = tl.exp(exponent[0]) - return (absorbed, absorbed * exponent[1]), shape_dual, slope_dual, deposited - - -# --------------------------------------------------------------------------- -# The forward state machine, and with ``following`` its derivative along a -# direction. With ``keep`` it records the state entering every event for the -# adjoint instead of writing the signal. -# --------------------------------------------------------------------------- - - -@triton.jit( - do_not_specialize=[ - "state_count", - "rows", - "locations", - "profile_bins", - "lineshape_bins", - ] -) -def _pooled_kernel( - m0, - b1, - b1_phase, - b0, - efficiency, - diffusion, - velocity, - dm0, - db1, - db1_phase, - db0, - defficiency, - ddiffusion, - dvelocity, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - dduration, - dflip, - dphase, - table, - dtable, - pool_index, - profile, - profile_index, - lineshape, - pairs, - pair_index, - dpairs, - output_real, - output_imag, - trajectory, - base, - atom_count, - event_count, - output_count, - state_count, - rows, - flow_scale, - washout_scale, - profile_step, - lineshape_step, - locations, - profile_bins, - lineshape_bins, - n: tl.constexpr, - m: tl.constexpr, - blocks: tl.constexpr, - planes: tl.constexpr, - atom_stride: tl.constexpr, - shimmed: tl.constexpr, - profiled: tl.constexpr, - dynamic: tl.constexpr, - directed_pairs: tl.constexpr, - directed_table: tl.constexpr, - following: tl.constexpr, - keep: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - P: tl.constexpr, - S: tl.constexpr, -): - problem = tl.program_id(0).to(tl.int64) + base - atom = problem % atom_count - train = problem // atom_count - event_base = train * event_count - voxel_at = atom * atom_stride - location = atom % locations - live = problem >= 0 - - pool = tl.arange(0, P)[:, None] - column = tl.arange(0, P)[None, :] - state = tl.arange(0, S)[None, :] - state_mask = state < state_count - order = state.to(tl.float32) - exchanging = pool < m - semisolid = pool == n - 1 - row_width = n * n + n + 2 * m * m - width = n + rows * row_width * blocks - slot = table + atom * width - directions = dtable + atom * width - - density_of = _read(m0, dm0, voxel_at, density, 1.0, following) - voxel_b1 = _read(b1, db1, voxel_at, transmit, 1.0, following) - voxel_b1_phase = _read(b1_phase, db1_phase, voxel_at, off_axis, 0.0, following) - voxel_b0 = _read(b0, db0, voxel_at, off_axis, 0.0, following) - inversion = _read(efficiency, defficiency, voxel_at, inverting, 1.0, following) - damping_rate = _read(diffusion, ddiffusion, voxel_at, diffusing, 0.0, following) - moved = _read(velocity, dvelocity, voxel_at, moving, 0.0, following) - flow_rate = (flow_scale * moved[0], flow_scale * moved[1]) - washout_rate = (0.0, 0.0) - if moving: - heading = tl.where(moved[0] > 0.0, 1.0, 0.0) - tl.where( - moved[0] < 0.0, 1.0, 0.0 - ) - washout_rate = ( - washout_scale * tl.abs(moved[0]), - washout_scale * heading * moved[1], - ) - - equilibrium = _entries( - slot, directions, pool, pool < n, 0.0, 0, directed_table, False, following - ) - zero = tl.zeros((P, S), tl.float32) - fpr = zero - fpi = zero - fmr = zero - fmi = zero - zr = tl.where(state == 0, equilibrium[0], 0.0) - zi = zero - dfpr = zero - dfpi = zero - dfmr = zero - dfmi = zero - dzr = zero - dzi = zero - if following: - dzr = tl.where(state == 0, equilibrium[1], 0.0) - tile = pool * S + state - - for event in range(0, event_count): - if keep: - at = trajectory + ((problem - base) * event_count + event) * ( - planes * P * S - ) - tl.store(at + 0 * P * S + tile, fpr) - tl.store(at + 1 * P * S + tile, fpi) - tl.store(at + 2 * P * S + tile, fmr) - tl.store(at + 3 * P * S + tile, fmi) - tl.store(at + 4 * P * S + tile, zr) - tl.store(at + 5 * P * S + tile, zi) - if following: - tl.store(at + 6 * P * S + tile, dfpr) - tl.store(at + 7 * P * S + tile, dfpi) - tl.store(at + 8 * P * S + tile, dfmr) - tl.store(at + 9 * P * S + tile, dfmi) - tl.store(at + 10 * P * S + tile, dzr) - tl.store(at + 11 * P * S + tile, dzi) - - dt = _read( - duration + event_base, dduration + event_base, event, True, 0.0, following - ) - wout, unit_t, unit_z, _squared, _weight = _factors( - dt, damping_rate, voxel_b0, flow_rate, washout_rate, order, - off_axis, moving, diffusing, following, - ) # fmt: skip - carried = _cscale(wout, unit_t, following) - spin = _cscale(wout, unit_z, following) - row = tl.load(pool_index + event_base + event).to(tl.int64) - longitudinal_op, restored, transverse_op = _operators( - slot, directions, n + row * row_width, dt[1], rows * row_width, pool, - column, n, m, directed_table, blocks > 1, following, - ) # fmt: skip - plus, minus, longitudinal, _mp, _mm, _mz = _relax( - (fpr, fpi, dfpr, dfpi), - (fmr, fmi, dfmr, dfmi), - (zr, zi, dzr, dzi), - transverse_op, - longitudinal_op, - restored, - equilibrium, - wout, - carried, - spin, - state, - following, - ) - fpr = plus[0] - fpi = plus[1] - fmr = minus[0] - fmi = minus[1] - zr = longitudinal[0] - zi = longitudinal[1] - if following: - dfpr = plus[2] - dfpi = plus[3] - dfmr = minus[2] - dfmi = minus[3] - dzr = longitudinal[2] - dzi = longitudinal[3] - - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event).to(tl.int32) - if (event_action & 1) != 0: - fpr, fpi, fmr, fmi = _shift( - fpr, fpi, fmr, fmi, state, state_mask, state_count - ) - if following: - dfpr, dfpi, dfmr, dfmi = _shift( - dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count - ) - if event_kind == 1: - if (event_action & 4) != 0: - # Every exchanging pool is free water and inverts like it; a - # semisolid one is saturated by the pulse's own term. - if following: - dzr = tl.where( - exchanging, -(inversion[1] * zr + inversion[0] * dzr), dzr - ) - dzi = tl.where( - exchanging, -(inversion[1] * zi + inversion[0] * dzi), dzi - ) - zr = tl.where(exchanging, -inversion[0] * zr, zr) - zi = tl.where(exchanging, -inversion[0] * zi, zi) - else: - pulse_b1 = voxel_b1 - pulse_b1_phase = voxel_b1_phase - if shimmed: - transmit_at = ( - tl.load(shim_index + event).to(tl.int64) * atom_count + atom - ) - pulse_b1 = _read(b1, db1, transmit_at, transmit, 1.0, following) - pulse_b1_phase = _read( - b1_phase, db1_phase, transmit_at, True, 0.0, following - ) - nominal = _read( - flip + event_base, dflip + event_base, event, True, 0.0, following - ) - played = _read( - phase + event_base, dphase + event_base, event, True, 0.0, following - ) - alpha = _rmul(nominal, pulse_b1, following) - phi = (played[0] + pulse_b1_phase[0], played[1] + pulse_b1_phase[1]) - if n > m: - absorbed, _shape, _slope, _deposited = _absorption( - lineshape, rf_frequency, saturation, event, alpha, voxel_b0, - lineshape_bins, lineshape_step, following, - ) # fmt: skip - if following: - dzr = tl.where( - semisolid, absorbed[1] * zr + absorbed[0] * dzr, dzr - ) - dzi = tl.where( - semisolid, absorbed[1] * zi + absorbed[0] * dzi, dzi - ) - zr = tl.where(semisolid, absorbed[0] * zr, zr) - zi = tl.where(semisolid, absorbed[0] * zi, zi) - if dynamic: - if following: - a, spun = _dynamic_pair_dual_at( - pairs, dpairs, pair_index, event_base, event, atom, - atom_count, live, phi[0], phi[1], directed_pairs, - ) # fmt: skip - else: - held = _dynamic_pair_at( - pairs, pair_index, event_base, event, atom, atom_count, live - ) - a = (held[0], held[1], 0.0, 0.0) - spun = _cmul( - (held[2], held[3], 0.0, 0.0), - _polar((-phi[0], 0.0), following), - following, - ) - elif profiled: - at_row = _table_row(profile_index, event, location, locations) - turn = _polar((-phi[0], -phi[1]), following) - if following: - read = _profile_pair_slope( - profile, at_row, alpha[0], profile_bins, profile_step - ) - a = (read[0], read[2], read[1] * alpha[1], read[3] * alpha[1]) - b = (read[4], read[6], read[5] * alpha[1], read[7] * alpha[1]) - else: - read = _profile_pair( - profile, at_row, alpha[0], profile_bins, profile_step - ) - a = (read[0], read[1], 0.0, 0.0) - b = (read[2], read[3], 0.0, 0.0) - spun = _cmul(b, turn, following) - else: - a, spun, _sa, _sb, _turn = _hard_pair(alpha, phi, following) - if following: - turned = _rotate_spinor_dual( - a[0], a[1], spun[0], spun[1], a[2], a[3], spun[2], spun[3], - fpr, fpi, fmr, fmi, zr, zi, dfpr, dfpi, dfmr, dfmi, dzr, dzi, - ) # fmt: skip - dfpr = tl.where(exchanging, turned[6], dfpr) - dfpi = tl.where(exchanging, turned[7], dfpi) - dfmr = tl.where(exchanging, turned[8], dfmr) - dfmi = tl.where(exchanging, turned[9], dfmi) - dzr = tl.where(exchanging, turned[10], dzr) - dzi = tl.where(exchanging, turned[11], dzi) - else: - turned = _rotate_spinor( - a[0], a[1], spun[0], spun[1], fpr, fpi, fmr, fmi, zr, zi - ) - fpr = tl.where(exchanging, turned[0], fpr) - fpi = tl.where(exchanging, turned[1], fpi) - fmr = tl.where(exchanging, turned[2], fmr) - fmi = tl.where(exchanging, turned[3], fmi) - zr = tl.where(exchanging, turned[4], zr) - zi = tl.where(exchanging, turned[5], zi) - if not keep: - if (event_kind == 2) & ((event_action & 32) != 0): - origin = state == 0 - recorded = ( - _total(fpr, origin), - _total(fpi, origin), - _total(dfpr, origin), - _total(dfpi, origin), - ) - read_phase = _read( - phase + event_base, dphase + event_base, event, True, 0.0, following - ) - signal = _cscale( - density_of, - _cmul( - recorded, - _polar((-read_phase[0], -read_phase[1]), following), - following, - ), - following, - ) - out = tl.load(output_index + event) - written = problem * output_count + out - if following: - tl.store(output_real + written, signal[2], mask=out >= 0) - tl.store(output_imag + written, signal[3], mask=out >= 0) - else: - tl.store(output_real + written, signal[0], mask=out >= 0) - tl.store(output_imag + written, signal[1], mask=out >= 0) - if (event_action & 2) != 0: - fpr, fpi, fmr, fmi = _shift( - fpr, fpi, fmr, fmi, state, state_mask, state_count - ) - if following: - dfpr, dfpi, dfmr, dfmi = _shift( - dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count - ) - if (event_action & 8) != 0: - fpr = zero - fpi = zero - fmr = zero - fmi = zero - if following: - dfpr = zero - dfpi = zero - dfmr = zero - dfmi = zero - elif (event_action & 16) != 0: - fpr, fpi, fmr, fmi = _shift( - fpr, fpi, fmr, fmi, state, state_mask, state_count - ) - if following: - dfpr, dfpi, dfmr, dfmi = _shift( - dfpr, dfpi, dfmr, dfmi, state, state_mask, state_count - ) - - -# --------------------------------------------------------------------------- -# The adjoint, walking back from the states ``_pooled_kernel`` recorded. With -# ``following`` it is the forward-over-reverse pass: the value planes carry the -# first-order cotangents and the tangent planes their derivative. -# --------------------------------------------------------------------------- - - -@triton.jit( - do_not_specialize=[ - "state_count", - "rows", - "locations", - "profile_bins", - "lineshape_bins", - ] -) -def _pooled_adjoint_kernel( - m0, - b1, - b1_phase, - b0, - efficiency, - diffusion, - velocity, - dm0, - db1, - db1_phase, - db0, - defficiency, - ddiffusion, - dvelocity, - duration, - kind, - flip, - phase, - action, - output_index, - shim_index, - saturation, - rf_frequency, - dduration, - dflip, - dphase, - table, - dtable, - pool_index, - profile, - profile_index, - lineshape, - pairs, - pair_index, - dpairs, - grad_real, - grad_imag, - grad_tissue, - dgrad_tissue, - grad_duration, - dgrad_duration, - grad_flip, - dgrad_flip, - grad_phase, - dgrad_phase, - grad_table, - dgrad_table, - grad_pairs, - dgrad_pairs, - trajectory, - base, - atom_count, - event_count, - output_count, - state_count, - rows, - flow_scale, - washout_scale, - profile_step, - lineshape_step, - locations, - profile_bins, - lineshape_bins, - m0_row, - b1_row, - b1_phase_row, - b0_row, - efficiency_row, - diffusion_row, - velocity_row, - n: tl.constexpr, - m: tl.constexpr, - blocks: tl.constexpr, - planes: tl.constexpr, - atom_stride: tl.constexpr, - shimmed: tl.constexpr, - profiled: tl.constexpr, - dynamic: tl.constexpr, - directed_pairs: tl.constexpr, - directed_table: tl.constexpr, - following: tl.constexpr, - off_axis: tl.constexpr, - moving: tl.constexpr, - diffusing: tl.constexpr, - transmit: tl.constexpr, - density: tl.constexpr, - inverting: tl.constexpr, - P: tl.constexpr, - S: tl.constexpr, -): - problem = tl.program_id(0).to(tl.int64) + base - atom = problem % atom_count - train = problem // atom_count - event_base = train * event_count - voxel_at = atom * atom_stride - location = atom % locations - live = problem >= 0 - - pool = tl.arange(0, P)[:, None] - column = tl.arange(0, P)[None, :] - state = tl.arange(0, S)[None, :] - state_mask = state < state_count - order = state.to(tl.float32) - origin = state == 0 - exchanging = pool < m - semisolid = pool == n - 1 - held_rows = pool < n - square = (pool < n) & (column < n) - across_mask = (pool < m) & (column < m) - live_exchanging = exchanging & state_mask - row_width = n * n + n + 2 * m * m - width = n + rows * row_width * blocks - slot = table + atom * width - directions = dtable + atom * width - slot_grad = grad_table + problem * width - slot_curve = dgrad_table + problem * width - - density_of = _read(m0, dm0, voxel_at, density, 1.0, following) - voxel_b1 = _read(b1, db1, voxel_at, transmit, 1.0, following) - voxel_b1_phase = _read(b1_phase, db1_phase, voxel_at, off_axis, 0.0, following) - voxel_b0 = _read(b0, db0, voxel_at, off_axis, 0.0, following) - inversion = _read(efficiency, defficiency, voxel_at, inverting, 1.0, following) - damping_rate = _read(diffusion, ddiffusion, voxel_at, diffusing, 0.0, following) - moved = _read(velocity, dvelocity, voxel_at, moving, 0.0, following) - flow_rate = (flow_scale * moved[0], flow_scale * moved[1]) - washout_rate = (0.0, 0.0) - heading = 0.0 - if moving: - heading = tl.where(moved[0] > 0.0, 1.0, 0.0) - tl.where( - moved[0] < 0.0, 1.0, 0.0 - ) - washout_rate = ( - washout_scale * tl.abs(moved[0]), - washout_scale * heading * moved[1], - ) - equilibrium = _entries( - slot, directions, pool, held_rows, 0.0, 0, directed_table, False, following - ) - - zero = tl.zeros((P, S), tl.float32) - pbr = zero - pbi = zero - mbr = zero - mbi = zero - zbr = zero - zbi = zero - dpbr = zero - dpbi = zero - dmbr = zero - dmbi = zero - dzbr = zero - dzbi = zero - grad_eq = equilibrium[0] * 0.0 - curve_eq = equilibrium[0] * 0.0 - grad_m0 = 0.0 - grad_b1 = 0.0 - grad_b1_phase = 0.0 - grad_b0 = 0.0 - grad_efficiency = 0.0 - grad_damping = 0.0 - grad_flow = 0.0 - grad_washout = 0.0 - # Only what every interval adds to is carried in the derivative plane; - # what a pulse or a readout adds is stored as the branch reaches it. A - # carried scalar updated inside one of those branches crashes the layout - # pass of Triton 3.8 in this kernel. - curve_b0 = 0.0 - curve_damping = 0.0 - curve_flow = 0.0 - curve_washout = 0.0 - # Transmit gradients are summed per shim: the running pair is flushed to - # its row whenever the walk back reaches a pulse on a different one. - held = 0 - tile = pool * S + state - - for step in range(0, event_count): - event = event_count - 1 - step - at = trajectory + ((problem - base) * event_count + event) * (planes * P * S) - if following: - plus_in = ( - tl.load(at + 0 * P * S + tile), - tl.load(at + 1 * P * S + tile), - tl.load(at + 6 * P * S + tile), - tl.load(at + 7 * P * S + tile), - ) - minus_in = ( - tl.load(at + 2 * P * S + tile), - tl.load(at + 3 * P * S + tile), - tl.load(at + 8 * P * S + tile), - tl.load(at + 9 * P * S + tile), - ) - z_in = ( - tl.load(at + 4 * P * S + tile), - tl.load(at + 5 * P * S + tile), - tl.load(at + 10 * P * S + tile), - tl.load(at + 11 * P * S + tile), - ) - else: - plus_in = (tl.load(at + tile), tl.load(at + P * S + tile), 0.0, 0.0) - minus_in = ( - tl.load(at + 2 * P * S + tile), - tl.load(at + 3 * P * S + tile), - 0.0, - 0.0, - ) - z_in = ( - tl.load(at + 4 * P * S + tile), - tl.load(at + 5 * P * S + tile), - 0.0, - 0.0, - ) - - dt = _read( - duration + event_base, dduration + event_base, event, True, 0.0, following - ) - wout, unit_t, unit_z, squared, weight = _factors( - dt, damping_rate, voxel_b0, flow_rate, washout_rate, order, - off_axis, moving, diffusing, following, - ) # fmt: skip - carried = _cscale(wout, unit_t, following) - spin = _cscale(wout, unit_z, following) - row = tl.load(pool_index + event_base + event).to(tl.int64) - row_offset = n + row * row_width - longitudinal_op, restored, transverse_op = _operators( - slot, directions, row_offset, dt[1], rows * row_width, pool, column, - n, m, directed_table, blocks > 1, following, - ) # fmt: skip - - # Replay the interval to recover the states the event acted on. - relaxed_plus, relaxed_minus, relaxed_z, mixed_plus, mixed_minus, mixed_z = ( - _relax( - plus_in, - minus_in, - z_in, - transverse_op, - longitudinal_op, - restored, - equilibrium, - wout, - carried, - spin, - state, - following, - ) - ) - spr = relaxed_plus[0] - spi = relaxed_plus[1] - smr = relaxed_minus[0] - smi = relaxed_minus[1] - dspr = zero - dspi = zero - dsmr = zero - dsmi = zero - if following: - dspr = relaxed_plus[2] - dspi = relaxed_plus[3] - dsmr = relaxed_minus[2] - dsmi = relaxed_minus[3] - event_action = tl.load(action + event).to(tl.int32) - event_kind = tl.load(kind + event).to(tl.int32) - if (event_action & 1) != 0: - spr, spi, smr, smi = _shift( - spr, spi, smr, smi, state, state_mask, state_count - ) - if following: - dspr, dspi, dsmr, dsmi = _shift( - dspr, dspi, dsmr, dsmi, state, state_mask, state_count - ) - - # The trailing shift or spoil. - if (event_action & 8) != 0: - pbr = zero - pbi = zero - mbr = zero - mbi = zero - if following: - dpbr = zero - dpbi = zero - dmbr = zero - dmbi = zero - elif (event_action & 16) != 0: - pbr, pbi, mbr, mbi = _shift_adjoint( - pbr, pbi, mbr, mbi, state, state_mask, state_count - ) - if following: - dpbr, dpbi, dmbr, dmbi = _shift_adjoint( - dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count - ) - if (event_action & 2) != 0: - pbr, pbi, mbr, mbi = _shift_adjoint( - pbr, pbi, mbr, mbi, state, state_mask, state_count - ) - if following: - dpbr, dpbi, dmbr, dmbi = _shift_adjoint( - dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count - ) - - # The readout. It carries no pulse, so what it recorded is the state - # the pre-shift left. - out = tl.load(output_index + event) - if (event_kind == 2) & ((event_action & 32) != 0) & (out >= 0): - index = problem * output_count + out - seed = (tl.load(grad_real + index), tl.load(grad_imag + index), 0.0, 0.0) - read_phase = _read( - phase + event_base, dphase + event_base, event, True, 0.0, following - ) - demodulation = _polar((-read_phase[0], -read_phase[1]), following) - recorded = ( - _total(spr, origin), - _total(spi, origin), - _total(dspr, origin), - _total(dspi, origin), - ) - density_grad = _re_dot( - seed, _cmul(recorded, demodulation, following), following - ) - turned = ( - demodulation[1], - -demodulation[0], - demodulation[3], - -demodulation[2], - ) - phase_grad = _re_dot( - seed, - _cscale(density_of, _cmul(recorded, turned, following), following), - following, - ) - grad_m0 += density_grad[0] - tl.atomic_add(grad_phase + event_base + event, phase_grad[0]) - weighted = _cmul( - _conj(_cscale(density_of, demodulation, following)), seed, following - ) - put = origin & exchanging - pbr += tl.where(put, weighted[0], 0.0) - pbi += tl.where(put, weighted[1], 0.0) - if following: - tl.atomic_add( - dgrad_tissue + m0_row * atom_count + atom, density_grad[1] - ) - tl.atomic_add(dgrad_phase + event_base + event, phase_grad[1]) - dpbr += tl.where(put, weighted[2], 0.0) - dpbi += tl.where(put, weighted[3], 0.0) - - # The pulse. - if event_kind == 1: - if (event_action & 4) != 0: - bar = (zbr, zbi, dzbr, dzbi) - taken = _rtotal( - _re_dot(bar, relaxed_z, following), live_exchanging, following - ) - grad_efficiency -= taken[0] - if following: - tl.atomic_add( - dgrad_tissue + efficiency_row * atom_count + atom, -taken[1] - ) - dzbr = tl.where( - exchanging, -(inversion[1] * zbr + inversion[0] * dzbr), dzbr - ) - dzbi = tl.where( - exchanging, -(inversion[1] * zbi + inversion[0] * dzbi), dzbi - ) - zbr = tl.where(exchanging, -inversion[0] * zbr, zbr) - zbi = tl.where(exchanging, -inversion[0] * zbi, zbi) - else: - pulse_b1 = voxel_b1 - pulse_b1_phase = voxel_b1_phase - transmit_row = 0 - if shimmed: - shim = tl.load(shim_index + event).to(tl.int32) - changed = shim != held - b1_at = (b1_row + held).to(tl.int64) * atom_count + atom - b1_phase_at = (b1_phase_row + held).to(tl.int64) * atom_count + atom - tl.atomic_add(grad_tissue + b1_at, grad_b1, mask=changed) - tl.atomic_add( - grad_tissue + b1_phase_at, grad_b1_phase, mask=changed - ) - grad_b1 = tl.where(changed, 0.0, grad_b1) - grad_b1_phase = tl.where(changed, 0.0, grad_b1_phase) - held = shim - transmit_row = shim - transmit_at = shim.to(tl.int64) * atom_count + atom - pulse_b1 = _read(b1, db1, transmit_at, transmit, 1.0, following) - pulse_b1_phase = _read( - b1_phase, db1_phase, transmit_at, True, 0.0, following - ) - nominal = _read( - flip + event_base, dflip + event_base, event, True, 0.0, following - ) - played = _read( - phase + event_base, dphase + event_base, event, True, 0.0, following - ) - alpha = _rmul(nominal, pulse_b1, following) - phi = (played[0] + pulse_b1_phase[0], played[1] + pulse_b1_phase[1]) - turn = _polar((-phi[0], -phi[1]), following) - if dynamic: - a, spun = _dynamic_pair_dual_at( - pairs, dpairs, pair_index, event_base, event, atom, atom_count, - live, phi[0], phi[1], directed_pairs, - ) # fmt: skip - slope_a = a - slope_b = spun - elif profiled: - a, spun, slope_a, slope_b = _profiled_pair_dual( - profile, - _table_row(profile_index, event, location, locations), - alpha[0], - alpha[1], - phi[0], - phi[1], - profile_bins, - profile_step, - ) - else: - a, spun, slope_a, slope_b, _turn = _hard_pair(alpha, phi, following) - if following: - pair_a, pair_b, back_p, back_m, back_z = _spinor_adjoint_dual( - a, - spun, - (spr, spi, dspr, dspi), - (smr, smi, dsmr, dsmi), - relaxed_z, - (pbr, pbi, dpbr, dpbi), - (mbr, mbi, dmbr, dmbi), - (zbr, zbi, dzbr, dzbi), - ) - grad_a = ( - _total(pair_a[0], live_exchanging), - _total(pair_a[1], live_exchanging), - _total(pair_a[2], live_exchanging), - _total(pair_a[3], live_exchanging), - ) - grad_b = ( - _total(pair_b[0], live_exchanging), - _total(pair_b[1], live_exchanging), - _total(pair_b[2], live_exchanging), - _total(pair_b[3], live_exchanging), - ) - dpbr = tl.where(exchanging, back_p[2], dpbr) - dpbi = tl.where(exchanging, back_p[3], dpbi) - dmbr = tl.where(exchanging, back_m[2], dmbr) - dmbi = tl.where(exchanging, back_m[3], dmbi) - dzbr = tl.where(exchanging, back_z[2], dzbr) - dzbi = tl.where(exchanging, back_z[3], dzbi) - pbr = tl.where(exchanging, back_p[0], pbr) - pbi = tl.where(exchanging, back_p[1], pbi) - mbr = tl.where(exchanging, back_m[0], mbr) - mbi = tl.where(exchanging, back_m[1], mbi) - zbr = tl.where(exchanging, back_z[0], zbr) - zbi = tl.where(exchanging, back_z[1], zbi) - else: - back = _spinor_adjoint( - a[0], a[1], spun[0], spun[1], spr, spi, smr, smi, - relaxed_z[0], relaxed_z[1], pbr, pbi, mbr, mbi, zbr, zbi, - ) # fmt: skip - grad_a = ( - _total(back[0], live_exchanging), - _total(back[1], live_exchanging), - 0.0, - 0.0, - ) - grad_b = ( - _total(back[2], live_exchanging), - _total(back[3], live_exchanging), - 0.0, - 0.0, - ) - pbr = tl.where(exchanging, back[4], pbr) - pbi = tl.where(exchanging, back[5], pbi) - mbr = tl.where(exchanging, back[6], mbr) - mbi = tl.where(exchanging, back[7], mbi) - zbr = tl.where(exchanging, back[8], zbr) - zbi = tl.where(exchanging, back[9], zbi) - # The RF phase turns the axis once the pair is out, so it - # reaches ``b`` alone -- under every mode. - grad_phi = _re_dot( - grad_b, (spun[1], -spun[0], spun[3], -spun[2]), following - ) - grad_alpha = (0.0, 0.0) - if dynamic: - # The flip is inside the pair rather than read against it, - # so the cotangent goes out on the pair; ``b`` was turned - # by the phase after the pair came out, so it turns back. - unturned = _cmul(grad_b, _conj(turn), following) - entry = ( - tl.load(pair_index + event_base + event).to(tl.int64) - * atom_count - + atom - ) * 4 - tl.atomic_add(grad_pairs + entry + 0, grad_a[0]) - tl.atomic_add(grad_pairs + entry + 1, grad_a[1]) - tl.atomic_add(grad_pairs + entry + 2, unturned[0]) - tl.atomic_add(grad_pairs + entry + 3, unturned[1]) - if following: - tl.atomic_add(dgrad_pairs + entry + 0, grad_a[2]) - tl.atomic_add(dgrad_pairs + entry + 1, grad_a[3]) - tl.atomic_add(dgrad_pairs + entry + 2, unturned[2]) - tl.atomic_add(dgrad_pairs + entry + 3, unturned[3]) - else: - along_a = _re_dot(grad_a, slope_a, following) - along_b = _re_dot(grad_b, slope_b, following) - grad_alpha = (along_a[0] + along_b[0], along_a[1] + along_b[1]) - if n > m: - # The pulse scales every order of the semisolid pool by one - # real number, so its cotangent is one sum over the states. - absorbed, shape, slope, deposited = _absorption( - lineshape, rf_frequency, saturation, event, alpha, voxel_b0, - lineshape_bins, lineshape_step, following, - ) # fmt: skip - taken = _rtotal( - _re_dot((zbr, zbi, dzbr, dzbi), relaxed_z, following), - semisolid & state_mask, - following, - ) - if following: - dzbr = tl.where( - semisolid, absorbed[1] * zbr + absorbed[0] * dzbr, dzbr - ) - dzbi = tl.where( - semisolid, absorbed[1] * zbi + absorbed[0] * dzbi, dzbi - ) - zbr = tl.where(semisolid, absorbed[0] * zbr, zbr) - zbi = tl.where(semisolid, absorbed[0] * zbi, zbi) - exponent = _rmul(taken, absorbed, following) - swing = _rmul( - (2.0 * deposited, 0.0), - _rmul(_rmul(exponent, alpha, following), shape, following), - following, - ) - grad_alpha = (grad_alpha[0] + swing[0], grad_alpha[1] + swing[1]) - shifted = _rmul( - (deposited, 0.0), - _rmul( - _rmul(_rmul(exponent, alpha, following), alpha, following), - slope, - following, - ), - following, - ) - grad_b0 -= shifted[0] - if following: - tl.atomic_add( - dgrad_tissue + b0_row * atom_count + atom, -shifted[1] - ) - flip_grad = _rmul(grad_alpha, pulse_b1, following) - b1_grad = _rmul(grad_alpha, nominal, following) - tl.atomic_add(grad_flip + event_base + event, flip_grad[0]) - tl.atomic_add(grad_phase + event_base + event, grad_phi[0]) - grad_b1 += b1_grad[0] - grad_b1_phase += grad_phi[0] - if following: - tl.atomic_add(dgrad_flip + event_base + event, flip_grad[1]) - tl.atomic_add(dgrad_phase + event_base + event, grad_phi[1]) - tl.atomic_add( - dgrad_tissue - + (b1_row + transmit_row).to(tl.int64) * atom_count - + atom, - b1_grad[1], - ) - tl.atomic_add( - dgrad_tissue - + (b1_phase_row + transmit_row).to(tl.int64) * atom_count - + atom, - grad_phi[1], - ) - - if (event_action & 1) != 0: - pbr, pbi, mbr, mbi = _shift_adjoint( - pbr, pbi, mbr, mbi, state, state_mask, state_count - ) - if following: - dpbr, dpbi, dmbr, dmbi = _shift_adjoint( - dpbr, dpbi, dmbr, dmbi, state, state_mask, state_count - ) - - # The interval. Order zero also carries the recovery, which is the - # equilibrium less what washout leaves of the operator applied to it. - plus_bar = (pbr, pbi, dpbr, dpbi) - minus_bar = (mbr, mbi, dmbr, dmbi) - z_bar = (zbr, zbi, dzbr, dzbi) - if not following: - plus_bar = (pbr, pbi, 0.0, 0.0) - minus_bar = (mbr, mbi, 0.0, 0.0) - z_bar = (zbr, zbi, 0.0, 0.0) - seed = ( - tl.sum(tl.where(origin, zbr, 0.0), axis=1)[:, None], - tl.sum(tl.where(origin, dzbr, 0.0), axis=1)[:, None], - ) - grad_eq += seed[0] - restored_grad = (-(wout[0] * seed[0]), -(wout[1] * seed[0] + wout[0] * seed[1])) - grad_wout = ( - -_total(seed[0] * restored[0], held_rows), - -_total(seed[1] * restored[0] + seed[0] * restored[1], held_rows), - ) - if following: - curve_eq += seed[1] - - out_plus = _cmul( - _conj(plus_bar), _cmul(carried, mixed_plus, following), following - ) - out_minus = _cmul( - _conj(minus_bar), _cmul(_conj(carried), mixed_minus, following), following - ) - out_z = _cmul(_conj(z_bar), _cmul(spin, mixed_z, following), following) - # The damping is homogeneous of degree one in every state it acts on, - # so its gradient times the damping itself is the cotangent taken - # against the states the interval leaves; the turns are the same - # derivatives with an imaginary weight. - transverse_scaled = tl.sum(out_plus[0] + out_minus[0], axis=0)[None, :] - transverse_angle = tl.sum(out_minus[1] - out_plus[1], axis=0)[None, :] - longitudinal_scaled = tl.sum(out_z[0], axis=0)[None, :] - longitudinal_angle = tl.sum(-out_z[1], axis=0)[None, :] - wout_plus = _re_dot(plus_bar, _cmul(unit_t, mixed_plus, following), following) - wout_minus = _re_dot( - minus_bar, _cmul(_conj(unit_t), mixed_minus, following), following - ) - wout_z = _re_dot(z_bar, _cmul(unit_z, mixed_z, following), following) - grad_wout = ( - grad_wout[0] + _total(wout_plus[0] + wout_minus[0] + wout_z[0], state_mask), - grad_wout[1], - ) - half = order + 0.5 - grad_angle = (_total(transverse_angle, state_mask), 0.0) - grad_b_factor = ( - -_total( - weight * transverse_scaled + squared * longitudinal_scaled, state_mask - ), - 0.0, - ) - grad_turn = ( - -_total(half * transverse_angle + order * longitudinal_angle, state_mask), - 0.0, - ) - if following: - grad_wout = ( - grad_wout[0], - grad_wout[1] - + _total(wout_plus[1] + wout_minus[1] + wout_z[1], state_mask), - ) - transverse_scaled_t = tl.sum(out_plus[2] + out_minus[2], axis=0)[None, :] - transverse_angle_t = tl.sum(out_minus[3] - out_plus[3], axis=0)[None, :] - longitudinal_scaled_t = tl.sum(out_z[2], axis=0)[None, :] - longitudinal_angle_t = tl.sum(-out_z[3], axis=0)[None, :] - grad_angle = (grad_angle[0], _total(transverse_angle_t, state_mask)) - grad_b_factor = ( - grad_b_factor[0], - -_total( - weight * transverse_scaled_t + squared * longitudinal_scaled_t, - state_mask, - ), - ) - grad_turn = ( - grad_turn[0], - -_total( - half * transverse_angle_t + order * longitudinal_angle_t, state_mask - ), - ) - - # The operators' own entries. ``F-`` follows the conjugate of the - # transverse operator, so its cotangent lands on the entry itself. - released = _conj(carried) - transverse_grad = _cadd( - _couter(_cmul(plus_bar, released, following), _conj(plus_in), following), - _couter(_cmul(_conj(minus_bar), released, following), minus_in, following), - ) - weighed = _cmul(_conj(z_bar), spin, following) - longitudinal_grad = ( - _outer(weighed[0], z_in[0]) - _outer(weighed[1], z_in[1]), - 0.0, - ) - if following: - longitudinal_grad = ( - longitudinal_grad[0], - _outer(weighed[2], z_in[0]) - - _outer(weighed[3], z_in[1]) - + _outer(weighed[0], z_in[2]) - - _outer(weighed[1], z_in[3]), - ) - - # The cotangents back through the interval. - back_plus = _cmul( - released, _apply(transverse_op, plus_bar, True, True, following), following - ) - back_minus = _cmul( - carried, _apply(transverse_op, minus_bar, False, True, following), following - ) - back_z = _apply_real( - longitudinal_op, _cmul(_conj(spin), z_bar, following), True, following - ) - pbr = back_plus[0] - pbi = back_plus[1] - mbr = back_minus[0] - mbi = back_minus[1] - zbr = back_z[0] - zbi = back_z[1] - if following: - dpbr = back_plus[2] - dpbi = back_plus[3] - dmbr = back_minus[2] - dmbi = back_minus[3] - dzbr = back_z[2] - dzbi = back_z[3] - - # The row this event read is shared with every event of its length; - # its cotangent is summed into it, and reaches the event's own length - # through the row's slope. - longitudinal_at = row_offset + pool * n + column - restored_at = row_offset + n * n + pool - across = row_offset + n * n + n + 2 * (pool * m + column) - tl.atomic_add(slot_grad + longitudinal_at, longitudinal_grad[0], mask=square) - tl.atomic_add(slot_grad + restored_at, restored_grad[0], mask=held_rows) - tl.atomic_add(slot_grad + across, transverse_grad[0], mask=across_mask) - tl.atomic_add(slot_grad + across + 1, transverse_grad[1], mask=across_mask) - if following: - tl.atomic_add( - slot_curve + longitudinal_at, longitudinal_grad[1], mask=square - ) - tl.atomic_add(slot_curve + restored_at, restored_grad[1], mask=held_rows) - tl.atomic_add(slot_curve + across, transverse_grad[2], mask=across_mask) - tl.atomic_add(slot_curve + across + 1, transverse_grad[3], mask=across_mask) - table_duration = (0.0, 0.0) - if blocks > 1: - further = rows * row_width - slope_z = _entries( - slot, directions, further + longitudinal_at, square, dt[1], further, - directed_table, True, following, - ) # fmt: skip - slope_restored = _entries( - slot, directions, further + restored_at, held_rows, dt[1], further, - directed_table, True, following, - ) # fmt: skip - slope_real = _entries( - slot, directions, further + across, across_mask, dt[1], further, - directed_table, True, following, - ) # fmt: skip - slope_imag = _entries( - slot, directions, further + across + 1, across_mask, dt[1], further, - directed_table, True, following, - ) # fmt: skip - table_duration = ( - _total(longitudinal_grad[0] * slope_z[0], square) - + _total(restored_grad[0] * slope_restored[0], held_rows) - + _total( - transverse_grad[0] * slope_real[0] - + transverse_grad[1] * slope_imag[0], - across_mask, - ), - 0.0, - ) - if following: - table_duration = ( - table_duration[0], - _total( - longitudinal_grad[1] * slope_z[0] - + longitudinal_grad[0] * slope_z[1], - square, - ) - + _total( - restored_grad[1] * slope_restored[0] - + restored_grad[0] * slope_restored[1], - held_rows, - ) - + _total( - transverse_grad[2] * slope_real[0] - + transverse_grad[0] * slope_real[1] - + transverse_grad[3] * slope_imag[0] - + transverse_grad[1] * slope_imag[1], - across_mask, - ), - ) - # At the row's own length the slope reaches the output only - # through a direction in that length, so only the tangent - # plane takes this. - tl.atomic_add( - slot_curve + further + longitudinal_at, - longitudinal_grad[0] * dt[1], - mask=square, - ) - tl.atomic_add( - slot_curve + further + restored_at, - restored_grad[0] * dt[1], - mask=held_rows, - ) - tl.atomic_add( - slot_curve + further + across, - transverse_grad[0] * dt[1], - mask=across_mask, - ) - tl.atomic_add( - slot_curve + further + across + 1, - transverse_grad[1] * dt[1], - mask=across_mask, - ) - - # Washout scales every factor the interval applies and the recovery it - # leaves; past the clamp nothing depends on the rate. - fraction_grad = (0.0, 0.0) - if moving: - inside = washout_rate[0] * dt[0] < 1.0 - fraction_grad = ( - tl.where(inside, -grad_wout[0], 0.0), - tl.where(inside, -grad_wout[1], 0.0), - ) - angle_rate = _rmul((-6.283185307179586, 0.0), dt, following) - b0_grad = _rmul(grad_angle, angle_rate, following) - damping_grad = _rmul(grad_b_factor, dt, following) - flow_grad = _rmul(grad_turn, dt, following) - washout_grad = _rmul(fraction_grad, dt, following) - duration_grad = _rmul( - grad_angle, _rmul((-6.283185307179586, 0.0), voxel_b0, following), following - ) - through_damping = _rmul(grad_b_factor, damping_rate, following) - through_flow = _rmul(grad_turn, flow_rate, following) - through_washout = _rmul(fraction_grad, washout_rate, following) - grad_b0 += b0_grad[0] - grad_damping += damping_grad[0] - grad_flow += flow_grad[0] - grad_washout += washout_grad[0] - tl.atomic_add( - grad_duration + event_base + event, - duration_grad[0] - + through_damping[0] - + through_flow[0] - + through_washout[0] - + table_duration[0], - ) - if following: - curve_b0 += b0_grad[1] - curve_damping += damping_grad[1] - curve_flow += flow_grad[1] - curve_washout += washout_grad[1] - tl.atomic_add( - dgrad_duration + event_base + event, - duration_grad[1] - + through_damping[1] - + through_flow[1] - + through_washout[1] - + table_duration[1], - ) - - # The equilibrium is also where every pool starts, which the walk back - # reaches last. - grad_eq += tl.sum(tl.where(origin, zbr, 0.0), axis=1)[:, None] - tl.atomic_add(slot_grad + pool, grad_eq, mask=held_rows) - tl.atomic_add(grad_tissue + m0_row * atom_count + atom, grad_m0) - tl.atomic_add( - grad_tissue + (b1_row + held).to(tl.int64) * atom_count + atom, grad_b1 - ) - tl.atomic_add( - grad_tissue + (b1_phase_row + held).to(tl.int64) * atom_count + atom, - grad_b1_phase, - ) - tl.atomic_add(grad_tissue + b0_row * atom_count + atom, grad_b0) - tl.atomic_add(grad_tissue + efficiency_row * atom_count + atom, grad_efficiency) - tl.atomic_add(grad_tissue + diffusion_row * atom_count + atom, grad_damping) - # One buffer drives two rates, so the velocity gradient is the sum of what - # each geometry carries back. - tl.atomic_add( - grad_tissue + velocity_row * atom_count + atom, - flow_scale * grad_flow + heading * washout_scale * grad_washout, - ) - if following: - curve_eq += tl.sum(tl.where(origin, dzbr, 0.0), axis=1)[:, None] - tl.atomic_add(slot_curve + pool, curve_eq, mask=held_rows) - tl.atomic_add(dgrad_tissue + b0_row * atom_count + atom, curve_b0) - tl.atomic_add(dgrad_tissue + diffusion_row * atom_count + atom, curve_damping) - tl.atomic_add( - dgrad_tissue + velocity_row * atom_count + atom, - flow_scale * curve_flow + heading * washout_scale * curve_washout, - ) - - -# --------------------------------------------------------------------------- -# Launchers. -# --------------------------------------------------------------------------- - - -def _tiles(layout: Any, state_count: int) -> dict[str, int]: - """The tile shape and warps a launch compiles for. - - The widest intermediate is an operator times a tile of states, pools by - pools by orders, and the warps are sized to hold it. - """ - pools = triton.next_power_of_2(layout.longitudinal) - states = triton.next_power_of_2(state_count) - return { - "n": layout.longitudinal, - "m": layout.transverse, - "blocks": layout.blocks, - "P": pools, - "S": states, - "num_warps": max(1, min(8, pools * pools * states // 256)), - } - - -def _inputs( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - pools: Any, - profile: Any, - lineshape: Any, - dynamic: Any, - tissue_tangents: tuple[torch.Tensor, ...] | None, - event_tangents: tuple[torch.Tensor, ...] | None, - dynamic_direction: Any, -) -> tuple[torch.Tensor, ...]: - """The buffers both kernels open with, in their order. - - A buffer whose branch is compiled out is still an argument, so one the - launch already holds stands in for it. - """ - device = tissue[0].device - voxel = tuple(tissue[index] for index in _VOXEL_INDEX) - moving = ( - voxel - if tissue_tangents is None - else tuple(tissue_tangents[index] for index in _VOXEL_INDEX) - ) - duration, kind, flip, phase = events[:4] - stepping = (duration, flip, phase) if event_tangents is None else event_tangents - pairs = duration if dynamic is None else dynamic.packed(device) - return ( - *voxel, - *moving, - *events[:9], - *stepping, - pools.values, - pools.values if pools.direction is None else pools.direction, - pools.index, - duration if profile is None else profile.packed(device), - kind if profile is None else profile.rows(device), - duration if lineshape is None else lineshape.packed(device), - pairs, - kind - if dynamic is None - else dynamic.rows_per_event(_train_count(events), kind.numel()).to(device), - pairs if dynamic_direction is None else dynamic_direction.to(device), - ) - - -def _scalars( - base: int, - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - output_count: int, - state_count: int, - pools: Any, - geometry: Geometry, - profile: Any, - lineshape: Any, -) -> tuple[Any, ...]: - """The scalar arguments both kernels take after their buffers.""" - return ( - base, - tissue[0].numel(), - events[1].numel(), - output_count, - state_count, - pools.layout.rows, - geometry.flow_scale, - geometry.washout_scale, - 1.0 if profile is None else profile.step, - 1.0 if lineshape is None else lineshape.step, - 1 if profile is None else profile.points, - 0 if profile is None else profile.bins, - 0 if lineshape is None else lineshape.bins, - ) - - -def _switches( - tissue: tuple[torch.Tensor, ...], - moving: tuple[torch.Tensor, ...] | None, - pools: Any, - profile: Any, - dynamic: Any, - dynamic_direction: Any, - features: frozenset[str] | None, - geometry: Geometry, - following: bool, -) -> dict[str, Any]: - """The constexpr switches both kernels take.""" - return { - "atom_stride": _atom_stride(tissue) - if moving is None - else _atom_stride(tissue, moving), - "shimmed": _shim_count(tissue) > 1, - "profiled": profile is not None and profile.bins > 0, - "dynamic": dynamic is not None, - "directed_pairs": following and dynamic_direction is not None, - "directed_table": following and pools.direction is not None, - "following": following, - **_feature_flags(features, geometry), - } - - -def _forward( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tissue_tangents: tuple[torch.Tensor, ...] | None, - event_tangents: tuple[torch.Tensor, ...] | None, - *, - state_count: int, - output_count: int, - geometry: Geometry, - profile: Any, - lineshape: Any, - dynamic: Any, - dynamic_direction: Any, - features: frozenset[str] | None, - pools: Any, -) -> torch.Tensor: - following = tissue_tangents is not None - atoms = tissue[0].numel() - total = _train_count(events) * atoms - output_real = torch.zeros( - _output_shape(_train_count(events), atoms, output_count), - dtype=torch.float32, - device=tissue[0].device, - ) - output_imag = torch.zeros_like(output_real) - if total: - _pooled_kernel[(total,)]( - *_inputs( - tissue, - events, - pools, - profile, - lineshape, - dynamic, - tissue_tangents, - event_tangents, - dynamic_direction, - ), - output_real, - output_imag, - output_real, - *_scalars( - 0, - tissue, - events, - output_count, - state_count, - pools, - geometry, - profile, - lineshape, - ), # fmt: skip - planes=12 if following else 6, - keep=False, - **_switches( - tissue, - tissue_tangents, - pools, - profile, - dynamic, - dynamic_direction, - features, - geometry, - following, - ), - **_tiles(pools.layout, state_count), - ) - return torch.complex(output_real, output_imag) - - -def simulate( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - *, - state_count: int, - output_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - dynamic: Any = None, - features: frozenset[str] | None = None, - pools: Any, -) -> torch.Tensor: - """The signals of a tissue whose pools are tabulated, on CUDA.""" - return _forward( - tissue, - events, - None, - None, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=None, - features=features, - pools=pools, - ) - - -def simulate_jvp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tissue_tangents: tuple[torch.Tensor, ...], - event_tangents: tuple[torch.Tensor, ...], - *, - state_count: int, - output_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - dynamic: Any = None, - dynamic_direction: Any = None, - features: frozenset[str] | None = None, - pools: Any, -) -> torch.Tensor: - """The signals' derivative along a direction, for a tabulated tissue. - - ``pools.direction`` is the direction along the tables, if they move. - """ - return _forward( - tissue, - events, - tissue_tangents, - tuple(event_tangents), - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - features=features, - pools=pools, - ) - - -def _adjoint( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tangents: tuple[torch.Tensor, ...] | None, - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - geometry: Geometry, - profile: Any, - lineshape: Any, - dynamic: Any, - dynamic_direction: Any, - features: frozenset[str] | None, - pools: Any, -) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: - """Both planes of the adjoint: the cotangents, then their derivative. - - Each plane holds the tissue's gradient rows, then duration, flip and - phase, then the pair's and the tables' cotangents where there are any. - The derivative plane is zero unless ``tangents`` gives a direction. - """ - following = tangents is not None - device = tissue[0].device - atoms = tissue[0].numel() - trains = _train_count(events) - event_count = events[1].numel() - total = trains * atoms - shims = _shim_count(tissue) - tissue_tangents = None if tangents is None else tangents[: len(TISSUE_NAMES)] - event_tangents = ( - None - if tangents is None - else tuple(tangents[FLOAT_NAMES.index(name)] for name in _EVENT) - ) - inputs = _inputs( - tissue, - events, - pools, - profile, - lineshape, - dynamic, - tissue_tangents, - event_tangents, - dynamic_direction, - ) - switches = _switches( - tissue, - tissue_tangents, - pools, - profile, - dynamic, - dynamic_direction, - features, - geometry, - following, - ) - tiles = _tiles(pools.layout, state_count) - planes = 12 if following else 6 - held = planes * tiles["P"] * tiles["S"] * max(1, event_count) - wave = max(1, min(total, _TRAJECTORY_BUDGET_BYTES // (4 * held))) - trajectory = torch.empty(wave * held, dtype=torch.float32, device=device) - - grad_output = grad_output.resolve_conj() - grad_real = grad_output.real.contiguous() - grad_imag = grad_output.imag.contiguous() - duration, _kind, flip, phase = events[:4] - pairs = inputs[_PAIRS] - - def plane() -> tuple[torch.Tensor, ...]: - return ( - torch.zeros( - tissue_gradient_height(shims) * atoms, - dtype=torch.float32, - device=device, - ), - torch.zeros_like(duration), - torch.zeros_like(flip), - torch.zeros_like(phase), - torch.zeros( - (trains, *pools.values.shape), dtype=torch.float32, device=device - ), - torch.zeros_like(pairs), - ) - - value, tangent = plane(), plane() - rows = tissue_gradient_bases(shims) - for base in range(0, total, wave): - span = min(wave, total - base) - scalars = _scalars( - base, tissue, events, output_count, state_count, pools, geometry, - profile, lineshape, - ) # fmt: skip - _pooled_kernel[(span,)]( - *inputs, - grad_real, - grad_imag, - trajectory, - *scalars, - planes=planes, - keep=True, - **switches, - **tiles, - ) - _pooled_adjoint_kernel[(span,)]( - *inputs, - grad_real, - grad_imag, - *(entry for pair in zip(value, tangent, strict=True) for entry in pair), - trajectory, - *scalars, - *(rows[TISSUE_NAMES.index(name)] for name in _VOXEL), - planes=planes, - **switches, - **tiles, - ) - - def gathered(side: tuple[torch.Tensor, ...]) -> tuple[torch.Tensor, ...]: - voxel = tuple( - side[0][start * atoms : (start + count) * atoms] - for start, count in zip(rows, tissue_gradient_rows(shims), strict=True) - ) - return ( - *voxel, - side[1], - side[2], - side[3], - *(() if dynamic is None else (side[5],)), - side[4].sum(0), - ) - - return gathered(value), gathered(tangent) - - -# Where the pair sits among the buffers ``_inputs`` returns. -_PAIRS = 2 * len(_VOXEL) + 9 + 3 + 3 + 3 - -_EVENT = ("duration", "flip", "phase") - - -def simulate_vjp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - dynamic: Any = None, - lineshape: Any = None, - features: frozenset[str] | None = None, - pools: Any, -) -> tuple[torch.Tensor, ...]: - """The first-order adjoint of a tabulated tissue, on CUDA. - - Returns the tissue's gradient rows, then duration, flip and phase, then the - pair's cotangent where there is a pair and the tables' last. - """ - value, _ = _adjoint( - tissue, - events, - None, - grad_output, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=None, - features=features, - pools=pools, - ) - return value - - -def simulate_vjp_jvp( - tissue: tuple[torch.Tensor, ...], - events: tuple[torch.Tensor, ...], - tangents: tuple[torch.Tensor, ...], - grad_output: torch.Tensor, - *, - state_count: int, - output_count: int, - geometry: Geometry = NO_GEOMETRY, - profile: Any = None, - lineshape: Any = None, - dynamic: Any = None, - dynamic_direction: Any = None, - features: frozenset[str] | None = None, - pools: Any, -) -> tuple[tuple[torch.Tensor, ...], tuple[torch.Tensor, ...]]: - """Forward-over-reverse through a tabulated tissue, on CUDA. - - Returns the gradients with respect to the primal inputs -- the derivative - of the adjoint along ``tangents`` -- and then with respect to the tangent - inputs, which is the adjoint itself. - """ - value, tangent = _adjoint( - tissue, - events, - tangents, - grad_output, - state_count=state_count, - output_count=output_count, - geometry=geometry, - profile=profile, - lineshape=lineshape, - dynamic=dynamic, - dynamic_direction=dynamic_direction, - features=features, - pools=pools, - ) - return tangent, value diff --git a/tests/estimators/test_perk_kernel.py b/tests/estimators/test_perk_kernel.py index 666a8397..5c1bbaaf 100644 --- a/tests/estimators/test_perk_kernel.py +++ b/tests/estimators/test_perk_kernel.py @@ -42,7 +42,7 @@ def composed(monkeypatch): """Force the plain Torch path, whatever backends are loaded.""" def only_composed(): - monkeypatch.setattr(_perk, "_TRITON", None) + monkeypatch.setattr(_perk, "_GPU", None) monkeypatch.setattr(_perk, "_NATIVE", None) return only_composed @@ -94,7 +94,7 @@ def test_the_fused_path_is_the_one_that_ran(device, monkeypatch) -> None: """Agreement cannot tell a fused kernel from a fallback, so ask directly.""" estimator, measured = _fitted(device) reached: list[str] = [] - backend = _perk._TRITON if device == "cuda" else _perk._NATIVE + backend = _perk._GPU if device == "cuda" else _perk._NATIVE original = backend.regress monkeypatch.setattr( backend, @@ -133,7 +133,7 @@ def test_the_adjoint_is_the_composed_gradient(device, monkeypatch) -> None: cotangent = torch.randn_like(estimator(fused_input)) estimator(fused_input).backward(cotangent) - monkeypatch.setattr(_perk, "_TRITON", None) + monkeypatch.setattr(_perk, "_GPU", None) monkeypatch.setattr(_perk, "_NATIVE", None) plain_input = measured.clone().requires_grad_() estimator(plain_input).backward(cotangent) @@ -147,7 +147,7 @@ def test_a_gradient_wanted_for_a_fitted_tensor_falls_back(device, monkeypatch) - path that differentiates everything -- and the route is what is asserted.""" estimator, measured = _fitted(device) estimator.weight.requires_grad_(True) - backend = _perk._TRITON if device == "cuda" else _perk._NATIVE + backend = _perk._GPU if device == "cuda" else _perk._NATIVE monkeypatch.setattr(backend, "regress", lambda *args: pytest.fail("the kernel ran")) values = estimator(measured) @@ -205,3 +205,47 @@ def test_a_single_voxel_goes_through_a_kernel_tiled_for_many(device) -> None: estimator, measured = _fitted(device) assert estimator(measured[:1]).shape == (1, 2) + + +def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint() -> None: + """The kernels the card runs, compiled for the host, against Torch. + + Feature, parameter and contrast counts that are not multiples of the + blocks the kernels walk them in, so every edge of the tiling is read. + """ + gpu = pytest.importorskip("blochsim.estimators._perk_gpu") + pytest.importorskip("blochsim._gpu_host") + generator = torch.Generator().manual_seed(0) + voxels, contrasts, features, parameters = 300, 37, 45, 17 + signals = torch.randn(voxels, contrasts, generator=generator) + frequency = torch.randn(features, contrasts, generator=generator) / 4 + phase = 2 * math.pi * torch.rand(features, generator=generator) + feature_mean = torch.randn(features, generator=generator) + weight = torch.randn(parameters, features, generator=generator) + parameter_mean = torch.randn(parameters, generator=generator) + cotangent = torch.randn(voxels, parameters, generator=generator) + scale = math.sqrt(2.0 / features) + + def line(x: torch.Tensor) -> torch.Tensor: + mapped = scale * torch.cos(x @ frequency.T + phase) - feature_mean + return parameter_mean + mapped @ weight.T + + x = signals.double().requires_grad_() + with torch.no_grad(): + frequency, phase = frequency.double(), phase.double() + feature_mean, weight = feature_mean.double(), weight.double() + parameter_mean = parameter_mean.double() + expected = line(x) + (expected_gradient,) = torch.autograd.grad(expected, x, cotangent.double()) + + estimated = gpu.regress( + signals, frequency, phase, feature_mean, weight, parameter_mean + ) + gradient = gpu.regress_vjp(cotangent, signals, frequency, phase, weight) + + torch.testing.assert_close( + estimated.double(), expected.detach(), atol=1e-5, rtol=1e-5 + ) + torch.testing.assert_close( + gradient.double(), expected_gradient, atol=1e-5, rtol=1e-5 + ) diff --git a/tests/sequence/test_both_pools.py b/tests/sequence/test_both_pools.py index 80a9a77c..b38f10e9 100644 --- a/tests/sequence/test_both_pools.py +++ b/tests/sequence/test_both_pools.py @@ -1462,54 +1462,22 @@ def _spread_over(columns): return torch.sqrt(torch.clamp(-2.0 * minors, min=0.0)) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") def test_the_series_carries_the_answer_up_to_the_spread_it_is_trusted_to() -> None: """``NARROW_SPREAD`` is a measurement, so it is pinned as one. The series branch alone, in float32, against ``matrix_exp`` in double, at the widest spread the gate will let through. Above the bound it is expected - to fail -- which is what makes the bound a choice rather than a hope. + to fail -- which is what makes the bound a choice rather than a hope. The + table kernel forms the undamped operator, and is run on a card where there + is one and through the host build of the same source where there is not. """ - import triton - import triton.language as tl - - from blochsim.sequence._epg_triton import _three_pool_step + from blochsim._gpu_launch import Kernel, available from blochsim.sequence._parameters import NARROW_SPREAD - @triton.jit - def only_the_series( - a, - b, - c, - d, - e, - f, - g, - h, - i, - out, - n, - NARROW: tl.constexpr, - BLOCK: tl.constexpr, - ): - j = tl.arange(0, BLOCK) - mask = j < n - entries = _three_pool_step( - tl.load(a + j, mask=mask, other=1.0), - tl.load(b + j, mask=mask, other=1.0), - tl.load(c + j, mask=mask, other=1.0), - tl.load(d + j, mask=mask, other=1.0), - tl.load(e + j, mask=mask, other=1.0), - tl.load(f + j, mask=mask, other=0.1), - tl.load(g + j, mask=mask, other=0.1), - tl.load(h + j, mask=mask, other=0.005), - tl.load(i + j, mask=mask, other=1.0), - NARROW, - ) - for entry in tl.static_range(12): - tl.store(out + entry * n + j, entries[entry], mask=mask) - + device = "cuda" if torch.cuda.is_available() and available() else "cpu" + table_kernel = Kernel("_three_pool_table_kernel") voxels = 4096 + block = 1024 worst = {} for label, seconds in (("inside", 20e-3), ("at the bound", None), ("past", 1.0)): columns = _three_pool_columns(voxels, seconds=seconds or 1.0, seed=5) @@ -1518,17 +1486,31 @@ def only_the_series( rate = float(_spread_over(columns).max()) columns = _three_pool_columns(voxels, seconds=NARROW_SPREAD / rate, seed=5) reached = float(_spread_over(columns).max()) - expected = _matrix_exp_oracle(columns) + r1a, r1b, r1c, exb, exc, fb, fc, dt, _att = columns + undamped = (r1a, r1b, r1c, exb, exc, fb, fc, dt, torch.ones_like(dt)) + expected = _matrix_exp_oracle(undamped)[:, :9] scale = expected.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) - out = torch.zeros(12 * voxels, device="cuda", dtype=torch.float32) - only_the_series[(1,)]( - *[value.to("cuda", torch.float32) for value in columns], + out = torch.zeros((1, 9, voxels), device=device, dtype=torch.float32) + + def ready(value): + return value.to(device, torch.float32).contiguous() + + table_kernel[(1, -(-voxels // block))]( + ready(1000.0 / r1a), + ready(1000.0 / r1b), + ready(1000.0 / r1c), + ready(exb), + ready(exc), + ready(fb), + ready(fc), + ready(dt[:1]), + torch.zeros(1, dtype=torch.int32, device=device), out, voxels, - NARROW=True, - BLOCK=triton.next_power_of_2(voxels), + BLOCK=block, + narrow=True, ) - got = out.reshape(12, voxels).T.double().cpu() + got = out.reshape(9, voxels).T.double().cpu() worst[label] = ( reached, float(((got - expected).abs() / scale).amax(dim=1).max()), @@ -1597,7 +1579,7 @@ def test_the_series_branch_gives_the_answer_the_roots_give(state_count) -> None: The gate cannot be measured against the host: that comparison moves for reasons of its own. What it has to be held to is the branch it replaces. """ - from blochsim.sequence import _accelerators, _epg_triton + from blochsim.sequence import _accelerators, _epg_gpu voxels = 256 tissue, _, _ = _prepare_tissue( @@ -1640,14 +1622,14 @@ def test_the_series_branch_gives_the_answer_the_roots_give(state_count) -> None: ) def both_ways(run): - settled = _epg_triton.narrow_three_pool + settled = _epg_gpu.narrow_three_pool sides = [] for forced in (True, False): - _epg_triton.narrow_three_pool = lambda *a, taken=forced, **k: taken + _epg_gpu.narrow_three_pool = lambda *a, taken=forced, **k: taken try: sides.append(run()) finally: - _epg_triton.narrow_three_pool = settled + _epg_gpu.narrow_three_pool = settled return sides def leaves(value): diff --git a/tests/sequence/test_cuda_parity.py b/tests/sequence/test_cuda_parity.py index f2be8b6e..cdfa6f7e 100644 --- a/tests/sequence/test_cuda_parity.py +++ b/tests/sequence/test_cuda_parity.py @@ -1,7 +1,7 @@ """The CUDA kernels against the CPU kernels they stand in for. The two share no code, so agreement across a batch of echo trains is what keeps -the Triton grid indexing honest: a train axis dropped there reads one train's +the GPU grid indexing honest: a train axis dropped there reads one train's flip angles for every train, which is a wrong answer rather than an error. """ @@ -231,7 +231,7 @@ def test_an_off_resonance_seed_keeps_the_complex_kernel_on_cuda(): @pytest.mark.parametrize("block_states", [4, 16]) def test_the_packing_width_is_a_power_of_two(block_states): """It indexes a ``tl.arange``, which rejects anything else.""" - from blochsim.sequence._epg_triton import _problems_per_program + from blochsim.sequence._epg_gpu import _problems_per_program width = _problems_per_program(block_states) assert width >= 1 @@ -246,7 +246,7 @@ def test_the_packing_width_ignores_how_many_problems_there_are(block_states): off the launch size would make a streamed volume answer differently from the same volume run whole. """ - from blochsim.sequence._epg_triton import _problems_per_program + from blochsim.sequence._epg_gpu import _problems_per_program assert _problems_per_program(block_states) >= 1 @@ -362,10 +362,10 @@ def test_an_inversion_pulse_reaches_the_same_gradients(inversion): def test_a_trajectory_too_large_for_one_launch_is_split(monkeypatch): """The grid rounds up past a wave, onto rows the next launch owns.""" - from blochsim.sequence import _epg_triton + from blochsim.sequence import _epg_gpu expected = _second_order("cpu", 17, 5) - monkeypatch.setattr(_epg_triton, "_TRAJECTORY_BUDGET_BYTES", 40_000) + monkeypatch.setattr(_epg_gpu, "_TRAJECTORY_BUDGET_BYTES", 40_000) actual = _second_order("cuda", 17, 5) assert _worst_disagreement(expected, actual) < _tolerance(17) @@ -496,10 +496,10 @@ def test_the_complex_second_order_kernel_matches_across_shapes(trains, atoms): def test_the_complex_trajectory_splits_into_waves(monkeypatch): """Twice the planes of the real one, so it reaches the budget sooner.""" - from blochsim.sequence import _epg_triton + from blochsim.sequence import _epg_gpu expected = _complex_second_order("cpu", 17, 5) - monkeypatch.setattr(_epg_triton, "_TRAJECTORY_BUDGET_BYTES", 40_000) + monkeypatch.setattr(_epg_gpu, "_TRAJECTORY_BUDGET_BYTES", 40_000) actual = _complex_second_order("cuda", 17, 5) assert _worst_disagreement(expected, actual) < _tolerance(17) diff --git a/tests/sequence/test_dynamic_transmit.py b/tests/sequence/test_dynamic_transmit.py index 4771a0ad..52636a53 100644 --- a/tests/sequence/test_dynamic_transmit.py +++ b/tests/sequence/test_dynamic_transmit.py @@ -725,7 +725,7 @@ def test_the_cuda_forward_reads_the_pair_the_host_does(): produce a plausible train. """ from blochsim.sequence._accelerators import _run_packed - from blochsim.sequence._epg_triton import simulate + from blochsim.sequence._epg_gpu import simulate from blochsim.sequence._transition import DynamicPairs _, prepared, events, pairs = _train() diff --git a/tests/sequence/test_feature_gates.py b/tests/sequence/test_feature_gates.py index 558550cf..713196fa 100644 --- a/tests/sequence/test_feature_gates.py +++ b/tests/sequence/test_feature_gates.py @@ -17,8 +17,10 @@ from blochsim.sequence import ( EpgEngine, TissueProperties, + _epg_gpu, fse_description, ) +from blochsim.sequence._epg_gpu import _feature_flags from blochsim.sequence._parameters import ( TISSUE_NAMES, TISSUE_PARAMETERS, @@ -27,14 +29,6 @@ features_of, ) -# The kernel module imports Triton at its top, and Triton is a dependency of -# the CUDA build of PyTorch rather than of BlochSim. Everything below the guard -# reaches it, so the two imports stay here rather than moving up with the rest. -pytest.importorskip("triton") - -from blochsim.sequence import _epg_triton # noqa: E402 -from blochsim.sequence._epg_triton import _feature_flags # noqa: E402 - ECHOES = 8 STATES = 8 @@ -235,7 +229,7 @@ def test_the_answer_does_not_depend_on_the_gate( """ gated = _adjoint(crusher_rad, **properties) monkeypatch.setattr( - _epg_triton, + _epg_gpu, "_feature_flags", lambda features, geometry: _feature_flags(None, geometry), ) diff --git a/tests/sequence/test_host_feature_mask.py b/tests/sequence/test_host_feature_mask.py index ddd85c37..40eda005 100644 --- a/tests/sequence/test_host_feature_mask.py +++ b/tests/sequence/test_host_feature_mask.py @@ -1,7 +1,6 @@ """The mask the host kernels read, against the same launch carrying everything. -The Triton backend takes a flag per term because each one compiles a kernel of -its own; the host kernels take one integer and branch on it at run time. These +The GPU kernels take a flag per term; the host kernels take one integer and branch on it at run time. These pin the two ends together -- that the mask says what :func:`blochsim.sequence._parameters.feature_flags` says, and that a host kernel told to drop a term gives the answer it gives when told to keep it. diff --git a/tests/sequence/test_host_kernels.py b/tests/sequence/test_host_kernels.py new file mode 100644 index 00000000..632d1286 --- /dev/null +++ b/tests/sequence/test_host_kernels.py @@ -0,0 +1,66 @@ +"""The GPU kernels, run on the host and held to what they specialize or to an +oracle. + +The host build of the kernels runs the source the CUDA build compiles one +program at a time over host tensors, so these reach the plumbing without a +card: the argument alignment of the launchers, the trajectory planes a pool +model claims, the operator table's row indexing, and the washout and +recoveries the reading event applies. +""" + +from __future__ import annotations + +import pytest + +from utils import host_kernels + + +@pytest.mark.parametrize( + "case", + [ + "narrow", + "wide", + "chunked", + "streamed", + "washed", + "shimmed", + "profiled", + "one_pool", + "two_pools", + "real", + "real_shimmed", + "spoiled", + "narrowed", + ], +) +def test_the_kernels_agree_with_what_they_specialize(case: str) -> None: + """Each case runs one launch two ways and holds the two to each other. + + The table cases differ only in where the operator comes from; the pool + cases hold a kernel to an oracle sharing no code with it, and to the pass + it specializes. + + ``narrow`` forces a table onto a train the launch-wide gate calls narrow, + so every row takes the series and the two arms agree to the bit. ``wide`` + is the case that ships -- an inversion makes the launch decline the gate, + and the table carries series rows beside a roots row. ``chunked`` is the + same launch cut into chunks, which is what tells a chunk-local index for + the cotangent table from a global one. ``streamed`` runs the chunked + launcher, whose fixed positional list is what a grown kernel signature + misaligns first. ``washed`` gives the interval a washout, so a pooled row + has to carry its own attenuation rather than one. ``shimmed`` drives a + transmit row per shim, which is what tells the three-pool row index from + the shim row it sits beside. ``profiled`` turns a shaped pulse through its + own table, which is the one launch that answers whether a table is read + separately from how many knots it holds. + ``one_pool`` and ``two_pools`` reach the pool models the table cases + never do, against the packed reference and against the + forward-over-reverse pass. ``real`` reaches the real-subspace kernels, + which carry three real planes where every other case here carries four + components -- against the reference, against the complex adjoint, and + with the gradients the representation cannot hold held to exactly zero. + ``real_shimmed`` gives those kernels a transmit row per shim and leaves + the last row undriven, so a layout that is merely wide enough cannot pass + for one that reads the index. + """ + host_kernels.run(case) diff --git a/tests/sequence/test_interpreted.py b/tests/sequence/test_interpreted.py deleted file mode 100644 index 0c22cc39..00000000 --- a/tests/sequence/test_interpreted.py +++ /dev/null @@ -1,101 +0,0 @@ -"""The Triton kernels, held to what they specialize or to an oracle. - -Triton's CPU interpreter runs the kernels in Python over host tensors, so -these reach the plumbing without a GPU and without invalidating a compile -cache: the argument alignment of the launchers, the trajectory planes a pool -model claims, the operator table's row indexing, and the washout and -recoveries the reading event applies. - -They are slow -- a minute each, since the interpreter walks every element in -Python -- so they are opt-in: ``pytest -m interpreted``. -""" - -from __future__ import annotations - -import os -import subprocess -import sys -from pathlib import Path - -import pytest - -TESTS = Path(__file__).resolve().parent.parent -ROOT = TESTS.parent -TIMEOUT_S = 900 - - -def _run(case: str) -> subprocess.CompletedProcess[str]: - """One case, in a process that sets the interpreter flag before Triton.""" - environment = dict(os.environ) - # Triton reads this at import, so it cannot be set from inside a test. - environment["TRITON_INTERPRET"] = "1" - environment["CUDA_VISIBLE_DEVICES"] = "" - environment["PYTHONPATH"] = os.pathsep.join([str(ROOT / "src"), str(TESTS)]) - return subprocess.run( - [sys.executable, "-m", "utils.interpreted", case], - capture_output=True, - text=True, - timeout=TIMEOUT_S, - cwd=ROOT, - env=environment, - ) - - -@pytest.mark.interpreted -@pytest.mark.parametrize( - "case", - [ - "narrow", - "wide", - "chunked", - "unread", - "streamed", - "washed", - "shimmed", - "profiled", - "one_pool", - "two_pools", - "real", - "real_shimmed", - "spoiled", - "narrowed", - ], -) -def test_the_kernels_agree_with_what_they_specialize(case: str) -> None: - """Each case runs one launch two ways and holds the two to each other. - - The table cases differ only in where the operator comes from; the pool - cases hold a kernel to an oracle sharing no code with it, and to the pass - it specializes. - - ``narrow`` forces a table onto a train the launch-wide gate calls narrow, - so every row takes the series and the two arms agree to the bit. ``wide`` - is the case that ships -- an inversion makes the launch decline the gate, - and the table carries series rows beside a roots row. ``chunked`` is the - same launch cut into chunks, which is what tells a chunk-local index for - the cotangent table from a global one. ``unread`` poisons ``acos`` so a - narrow launch that touched the three roots could not come back finite. - ``streamed`` runs the chunked launcher, whose fixed positional list is - what a grown kernel signature misaligns first. ``washed`` gives the - interval a washout, so a pooled row has to carry its own attenuation - rather than one. ``shimmed`` drives a transmit row per shim, which is - what tells the three-pool row index from the shim row it sits beside. - ``profiled`` turns a shaped pulse through its own table, which is the - one launch that answers whether a table is read separately from how - many knots it holds. - ``one_pool`` and ``two_pools`` reach the pool models the table cases - never do, against the packed reference and against the - forward-over-reverse pass. ``real`` reaches the real-subspace kernels, - which carry three real planes where every other case here carries four - components -- against the reference, against the complex adjoint, and - with the gradients the representation cannot hold held to exactly zero. - ``real_shimmed`` gives those kernels a transmit row per shim and leaves - the last row undriven, so a layout that is merely wide enough cannot pass - for one that reads the index. - """ - finished = _run(case) - assert finished.returncode == 0, ( - f"{case} case failed:\n{finished.stdout}\n{finished.stderr}" - ) - # A case that fell out before its checks prints no terminator. - assert finished.stdout.rstrip().endswith("checked"), finished.stdout diff --git a/tests/sequence/test_many_pools_host.py b/tests/sequence/test_many_pools_host.py new file mode 100644 index 00000000..761ea93f --- /dev/null +++ b/tests/sequence/test_many_pools_host.py @@ -0,0 +1,42 @@ +"""The GPU kernels for tabulated pools, run on the host and held to the C++ +ones. + +Each case runs one pass of both backends on the same tables. The C++ kernels +are held to the state machine written out in torch by ``test_many_pools.py``, +so agreeing with them is agreeing with that. +""" + +from __future__ import annotations + +import pytest + +from utils.host_pools import check + +PASSES = ("forward", "jvp", "vjp", "vjp_jvp") + + +@pytest.mark.parametrize("pass_name", PASSES) +@pytest.mark.parametrize( + ("pools", "rotation"), + [("2f", "instant"), ("3s", "profile"), ("4s", "dynamic"), ("4s", "shimmed")], +) +def test_the_gpu_kernels_agree_with_the_cpp_ones(pass_name, pools, rotation): + """Two to four exchanging pools, with and without a semisolid one, turned + by a hard pulse, a tabulated one, a rotation per voxel and a transmit row + per shim.""" + check(pass_name, pools, rotation) + + +@pytest.mark.parametrize("pass_name", ("vjp", "vjp_jvp")) +def test_an_adjoint_in_waves_over_two_trains_agrees_with_the_cpp_one(pass_name): + """Every problem recorded in a wave of its own, over two trains of + different lengths: the trajectory and the per-problem table cotangents are + indexed from the wave's base, and summed over the trains.""" + check(pass_name, "4s", "dynamic", "trains", "chunked") + + +@pytest.mark.parametrize("pass_name", PASSES) +def test_the_kernels_agree_with_every_optional_term_off(pass_name): + """A tissue declaring nothing and a sequence moving nothing turns every + optional term off in both kernels.""" + check(pass_name, "3s", "instant", "undeclared") diff --git a/tests/sequence/test_many_pools_interpreted.py b/tests/sequence/test_many_pools_interpreted.py deleted file mode 100644 index 31f82bc3..00000000 --- a/tests/sequence/test_many_pools_interpreted.py +++ /dev/null @@ -1,79 +0,0 @@ -"""The Triton kernels for tabulated pools, held to the C++ ones. - -Each case runs one pass of both backends on the same tables in Triton's CPU -interpreter, in a process of its own because Triton reads the interpreter flag -at import. The C++ kernels are held to the state machine written out in torch -by ``test_many_pools.py``, so agreeing with them is agreeing with that. - -They take seconds each rather than milliseconds, so they are opt-in: -``pytest -m interpreted``. -""" - -from __future__ import annotations - -import os -import subprocess -import sys -from pathlib import Path - -import pytest - -TESTS = Path(__file__).resolve().parent.parent -ROOT = TESTS.parent -TIMEOUT_S = 900 - -PASSES = ("forward", "jvp", "vjp", "vjp_jvp") - - -def _run(*arguments: str) -> subprocess.CompletedProcess[str]: - environment = dict(os.environ) - environment["TRITON_INTERPRET"] = "1" - environment["CUDA_VISIBLE_DEVICES"] = "" - environment["PYTHONPATH"] = os.pathsep.join([str(ROOT / "src"), str(TESTS)]) - return subprocess.run( - [sys.executable, "-m", "utils.interpreted_pools", *arguments], - capture_output=True, - text=True, - timeout=TIMEOUT_S, - cwd=ROOT, - env=environment, - ) - - -def _held(*arguments: str) -> None: - finished = _run(*arguments) - assert finished.returncode == 0, ( - f"{' '.join(arguments)} failed:\n{finished.stdout}\n{finished.stderr}" - ) - # A case that fell out before its comparison prints no terminator. - assert finished.stdout.rstrip().endswith("checked"), finished.stdout - - -@pytest.mark.interpreted -@pytest.mark.parametrize("pass_name", PASSES) -@pytest.mark.parametrize( - ("pools", "rotation"), - [("2f", "instant"), ("3s", "profile"), ("4s", "dynamic"), ("4s", "shimmed")], -) -def test_the_triton_kernels_agree_with_the_cpp_ones(pass_name, pools, rotation): - """Two to four exchanging pools, with and without a semisolid one, turned - by a hard pulse, a tabulated one, a rotation per voxel and a transmit row - per shim.""" - _held(pass_name, pools, rotation) - - -@pytest.mark.interpreted -@pytest.mark.parametrize("pass_name", ("vjp", "vjp_jvp")) -def test_an_adjoint_in_waves_over_two_trains_agrees_with_the_cpp_one(pass_name): - """Every problem recorded in a wave of its own, over two trains of - different lengths: the trajectory and the per-problem table cotangents are - indexed from the wave's base, and summed over the trains.""" - _held(pass_name, "4s", "dynamic", "trains", "chunked") - - -@pytest.mark.interpreted -@pytest.mark.parametrize("pass_name", PASSES) -def test_the_kernels_agree_with_every_optional_term_off(pass_name): - """A tissue declaring nothing and a sequence moving nothing compiles every - optional term out of both kernels.""" - _held(pass_name, "3s", "instant", "undeclared") diff --git a/tests/sequence/test_parameters.py b/tests/sequence/test_parameters.py index 442d9969..6a7b1b70 100644 --- a/tests/sequence/test_parameters.py +++ b/tests/sequence/test_parameters.py @@ -1,7 +1,7 @@ """The parameter table has to describe the buffers the kernels actually take. Its whole purpose is that a count appears once rather than in the Python -dispatch, the CPU extension and the Triton kernels separately. That only holds +dispatch, the CPU extension and the GPU kernels separately. That only holds while the table and the things it describes agree, which is what these check -- a new parameter added to the table but not to the dataclass, or to the kernels' pointer list, fails here rather than by reading past the end of a buffer. diff --git a/tests/sequence/test_subspace_streams.py b/tests/sequence/test_subspace_streams.py index 2046c9e8..1b78acf9 100644 --- a/tests/sequence/test_subspace_streams.py +++ b/tests/sequence/test_subspace_streams.py @@ -111,12 +111,12 @@ def _pair(backend: str, tissue, events, recorded): if backend == "host": shape = (tissue, events, STATES, recorded) return _run_packed(*shape, 1), _run_packed(*shape, 1, real_axis=1) - from blochsim.sequence import _epg_triton + from blochsim.sequence import _epg_gpu shape = dict(state_count=STATES, output_count=recorded) return ( - _epg_triton.simulate(tissue, events, **shape), - _epg_triton.simulate(tissue, events, real_axis=1, **shape), + _epg_gpu.simulate(tissue, events, **shape), + _epg_gpu.simulate(tissue, events, real_axis=1, **shape), ) diff --git a/tests/utils/interpreted.py b/tests/utils/host_kernels.py similarity index 81% rename from tests/utils/interpreted.py rename to tests/utils/host_kernels.py index 43622be3..914617aa 100644 --- a/tests/utils/interpreted.py +++ b/tests/utils/host_kernels.py @@ -1,25 +1,20 @@ -"""The three-pool operator table, checked in Triton's CPU interpreter. +"""The GPU kernels run on the host, held to what they specialize. -`TRITON_INTERPRET=1` runs a kernel in Python over host tensors, which reaches -the plumbing a compiled launch hides: a launcher whose positional list has -drifted from the kernel it calls, a trajectory plane without its buffer, a -replay disagreeing with the reverse sweep. Those are structural and show at -three voxels. +The same kernel source the CUDA build compiles, run one program at a time over +host tensors, reaches the plumbing a card hides: a launcher whose positional +list has drifted from the kernel it calls, a trajectory plane without its +buffer, a replay disagreeing with the reverse sweep. Those are structural and +show at three voxels. -Run as ``python -m utils.interpreted `` with ``TRITON_INTERPRET=1`` set -before the process starts -- Triton reads it at import. Each case exits -non-zero when a comparison drifts, so the pytest wrapper only has to run it. +Each case raises when a comparison drifts. """ from __future__ import annotations import math -import sys from typing import Any import torch -import triton -import triton.language as tl # What a narrow row and a roots row are each held to against the arm that # forms its operator per event. The roots row is looser because the table @@ -34,50 +29,6 @@ ORACLE_TOLERANCE = 5e-5 -@triton.jit -def _acos(x: Any) -> Any: - """``acos`` for the interpreter, which has no inverse trigonometry. - - ``libdevice`` is a CUDA extern and ``tl.math`` carries no acos, asin or - atan2, so a kernel forming the three roots cannot run on the host at all. - Newton on ``cos(theta) - x`` reaches 1.8e-6 over the range the callers - clamp to, which is ample when both arms of a comparison use it. - """ - theta = 1.5707963267948966 - x * (1.0 + 0.16666667 * x * x) - for _ in tl.static_range(0, 12): - sine = tl.sin(theta) - theta = theta + (tl.cos(theta) - x) / tl.where( - tl.abs(sine) > 1e-12, sine, 1e-12 - ) - return theta - - -@triton.jit -def _poisoned_acos(x: Any) -> Any: - """An ``acos`` whose result cannot be used without showing. - - A narrow launch is supposed to reach only the series. Nothing that passes - through this can stay finite, so a finite answer is proof the roots were - never read. - """ - return x * float("nan") - - -class _Interpretable: - acos = _acos - - -class _Poisoned: - acos = _poisoned_acos - - -def install(poison: bool = False) -> None: - """Point the kernels at an ``acos`` the interpreter can evaluate.""" - from blochsim.sequence import _epg_triton - - _epg_triton.libdevice = _Poisoned if poison else _Interpretable - - def _tissue(voxels: int) -> tuple[torch.Tensor, ...]: from blochsim.sequence import ( TissueProperties, @@ -131,9 +82,9 @@ def _both(run: Any, force_narrow: bool) -> tuple[Any, Any]: narrow, so every row takes the series and the two arms should agree to the bit. """ - from blochsim.sequence import _epg_triton + from blochsim.sequence import _epg_gpu - original = _epg_triton._tabulate_three_pool + original = _epg_gpu._tabulate_three_pool built: list[bool] = [] def patched( @@ -159,7 +110,7 @@ def patched( return rows, table, lengths built_wanted = [False] - _epg_triton._tabulate_three_pool = patched + _epg_gpu._tabulate_three_pool = patched try: without = run() built_wanted[0] = True @@ -169,7 +120,7 @@ def patched( # nothing at all. assert built and all(built), "the operator table was never built" finally: - _epg_triton._tabulate_three_pool = original + _epg_gpu._tabulate_three_pool = original return without, with_table @@ -188,7 +139,7 @@ def _chunked(voxels: int) -> Any: The cotangent table is sized and indexed per chunk, so a single-chunk run cannot tell a chunk-local index from a global one. """ - from blochsim.sequence import _epg_triton + from blochsim.sequence import _epg_gpu cut: list[int] = [] @@ -197,58 +148,10 @@ def narrow_wave(*arguments: Any, **keywords: Any) -> int: cut.append(wave) return wave - _epg_triton._trajectory_wave = narrow_wave + _epg_gpu._trajectory_wave = narrow_wave return cut -def _unread(voxels: int, states: int) -> None: - """Check that a narrow launch does not read the three roots. - - The `close` select sits in the consumers of `_three_pool_pieces_jvp`, and - each of them takes the series outright when the caller has bounded the - spread -- so the roots the pieces still form are unread. Poisoning `acos` - is what says so rather than reading the branches and believing it. - """ - install(poison=True) - from blochsim.sequence import _builders, _epg_triton - from blochsim.sequence._lineshape import lineshape_table - from blochsim.sequence._parameters import narrow_three_pool - - echoes = 6 - tissue = _tissue(voxels) - events, outputs = _events( - _builders.fse_description(torch.full((echoes,), math.radians(150.0)), 8e-3) - ) - assert narrow_three_pool(tissue, events[0].reshape(-1), pools=3), ( - "this train has to be narrow for the check to mean anything" - ) - options: dict[str, Any] = dict(lineshape=lineshape_table(), exchanging=True) - signal = _epg_triton.simulate( - tissue, events, state_count=states, output_count=outputs, **options - ) - assert bool(torch.isfinite(torch.view_as_real(signal)).all()), ( - "the forward read the roots" - ) - seed = ( - torch.rand(voxels, outputs, generator=torch.Generator().manual_seed(7)) * 2.0 - - 1.0 - ).to(torch.complex64) - for index, gradient in enumerate( - _epg_triton.simulate_vjp( - tissue, - events, - seed, - state_count=states, - output_count=outputs, - **options, - ) - ): - assert not gradient.numel() or bool(torch.isfinite(gradient).all()), ( - f"gradient {index} read the roots" - ) - print(" the roots are unread") - - def _streamed(voxels: int, states: int) -> None: """The chunked adjoint launcher, against the whole-volume one. @@ -256,8 +159,7 @@ def _streamed(voxels: int, states: int) -> None: takes, so it is the launcher a grown kernel signature misaligns first -- and nothing else here reaches it. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu echoes = 6 tissue = _tissue(voxels) @@ -268,13 +170,13 @@ def _streamed(voxels: int, states: int) -> None: torch.rand(voxels, outputs, generator=torch.Generator().manual_seed(3)) * 2.0 - 1.0 ).to(torch.complex64) - whole = _epg_triton.simulate_vjp( + whole = _epg_gpu.simulate_vjp( tissue, events, seed, state_count=states, output_count=outputs ) - buffers = _epg_triton.GradientBuffers( + buffers = _epg_gpu.GradientBuffers( events, voxels, state_count=states, output_count=outputs ) - chunked = _epg_triton.simulate_vjp_into( + chunked = _epg_gpu.simulate_vjp_into( tissue, events, seed, @@ -300,8 +202,7 @@ def _washed(voxels: int, states: int) -> None: ``1 - rate dt`` -- but the row has to be given that attenuation rather than one, because the gradients it pools are scaled by it. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._lineshape import lineshape_table from blochsim.sequence._parameters import TISSUE_NAMES, Geometry @@ -327,7 +228,7 @@ def _washed(voxels: int, states: int) -> None: - 1.0 ).to(torch.complex64) without, with_table = _both( - lambda: _epg_triton.simulate_vjp( + lambda: _epg_gpu.simulate_vjp( tissue, events, seed, @@ -352,8 +253,7 @@ def _real(voxels: int, states: int, shims: int = 1) -> None: every plane count, trajectory stride and gradient row differs -- and none of it is reached by the other cases here, which are all complex. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._accelerators import real_subspace_axis from blochsim.sequence._parameters import FLOAT_NAMES, OUTSIDE_THE_SUBSPACE from utils.packed_reference import simulate_packed @@ -392,16 +292,14 @@ def _real(voxels: int, states: int, shims: int = 1) -> None: torch.arange(events[6].numel(), dtype=torch.int32) % driven ).contiguous() events = tuple(events) - assert _epg_triton._shim_count(tissue) == shims, ( - "the array did not reach the kernel" - ) + assert _epg_gpu._shim_count(tissue) == shims, "the array did not reach the kernel" # Forcing the verdict rather than earning it would compare the real kernel # against a train it was never contracted for. assert real_subspace_axis(events, tissue) == 1, "this train is not in the subspace" shape = dict(state_count=states, output_count=outputs) expected = simulate_packed(tissue, events, **shape) - signal = _epg_triton.simulate(tissue, events, real_axis=1, **shape) + signal = _epg_gpu.simulate(tissue, events, real_axis=1, **shape) forward = float((signal - expected).abs().max() / expected.abs().max()) print(f" real forward {forward:.2e}") assert forward <= ORACLE_TOLERANCE, f"real forward drifted: {forward:.2e}" @@ -415,8 +313,8 @@ def _real(voxels: int, states: int, shims: int = 1) -> None: torch.rand(voxels, outputs, generator=generator) * 2.0 - 1.0, torch.rand(voxels, outputs, generator=generator) * 2.0 - 1.0, ) - real = _epg_triton.simulate_real_vjp(tissue, events, seed, **shape) - complex_side = _epg_triton.simulate_vjp(tissue, events, seed, **shape) + real = _epg_gpu.simulate_real_vjp(tissue, events, seed, **shape) + complex_side = _epg_gpu.simulate_vjp(tissue, events, seed, **shape) # Drawn from the gradients actually compared: a scale taken from one the # loop skips would let two nothings agree perfectly against a large number # that neither of them carries. @@ -464,8 +362,7 @@ def _spoiled(voxels: int, states: int) -> None: forced, and the forward pass, the forward-mode pass and the adjoint are each held to the complex kernels that carry the same train. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._accelerators import real_subspace_axis from blochsim.sequence._parameters import FLOAT_NAMES, OUTSIDE_THE_SUBSPACE from utils.packed_reference import simulate_packed @@ -483,7 +380,7 @@ def _spoiled(voxels: int, states: int) -> None: shape = dict(state_count=states, output_count=outputs) expected = simulate_packed(tissue, events, **shape) - signal = _epg_triton.simulate(tissue, events, real_axis=1, **shape) + signal = _epg_gpu.simulate(tissue, events, real_axis=1, **shape) forward = float((signal - expected).abs().max() / expected.abs().max()) print(f" spoiled forward {forward:.2e}") assert forward <= ORACLE_TOLERANCE, f"spoiled forward drifted: {forward:.2e}" @@ -498,12 +395,10 @@ def _spoiled(voxels: int, states: int) -> None: torch.zeros_like(events[2]), torch.zeros_like(events[3]), ) - real_jvp = _epg_triton.simulate_jvp( + real_jvp = _epg_gpu.simulate_jvp( tissue, events, tuple(tangents), zeros, real_axis=1, **shape ) - complex_jvp = _epg_triton.simulate_jvp( - tissue, events, tuple(tangents), zeros, **shape - ) + complex_jvp = _epg_gpu.simulate_jvp(tissue, events, tuple(tangents), zeros, **shape) scale = float(complex_jvp[1].abs().max()) assert scale > 1e-6, "the forward-mode pass returned nothing" drift = float((complex_jvp[1] - real_jvp[1]).abs().max()) / scale @@ -515,8 +410,8 @@ def _spoiled(voxels: int, states: int) -> None: torch.rand(voxels, outputs, generator=generator) * 2.0 - 1.0, torch.rand(voxels, outputs, generator=generator) * 2.0 - 1.0, ) - real = _epg_triton.simulate_real_vjp(tissue, events, seed, **shape) - complex_side = _epg_triton.simulate_vjp(tissue, events, seed, **shape) + real = _epg_gpu.simulate_real_vjp(tissue, events, seed, **shape) + complex_side = _epg_gpu.simulate_vjp(tissue, events, seed, **shape) scale = max( float(value.abs().max()) for index, value in enumerate(complex_side) @@ -543,11 +438,10 @@ def _pooled(voxels: int, states: int, pools: int) -> None: The table cases reach ``pools == 3`` only. What a second pool costs structurally is its own trajectory planes and its own place in every - launcher's positional list -- the two things this interpreter is for -- + launcher's positional list -- the two things the host build is for -- and neither is exercised at one or two pools without a card. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._lineshape import lineshape_table from utils.packed_reference import simulate_packed @@ -560,10 +454,10 @@ def _pooled(voxels: int, states: int, pools: int) -> None: lineshape=lineshape_table() if pools in (1, 3) else None, exchanging=pools in (2, 3), ) - assert _epg_triton._pool_flag(**options) == pools, "not the pool model asked for" + assert _epg_gpu._pool_flag(**options) == pools, "not the pool model asked for" shape = dict(state_count=states, output_count=outputs) - signal = _epg_triton.simulate(tissue, events, **shape, **options) + signal = _epg_gpu.simulate(tissue, events, **shape, **options) expected = simulate_packed(tissue, events, **shape, **options) forward = float((signal - expected).abs().max() / expected.abs().max()) print(f" {pools}-pool forward {forward:.2e}") @@ -571,7 +465,7 @@ def _pooled(voxels: int, states: int, pools: int) -> None: # Agreement between two ways of carrying nothing would say nothing, so the # pool has to be shown to move the answer it is being checked against. - bare = _epg_triton.simulate(tissue, events, **shape) + bare = _epg_gpu.simulate(tissue, events, **shape) moved = float((signal - bare).abs().max() / bare.abs().max()) print(f" {pools}-pool moves it {moved:.2e}") assert moved > 1e-3, f"the {pools}-pool model changed nothing: {moved:.2e}" @@ -580,13 +474,13 @@ def _pooled(voxels: int, states: int, pools: int) -> None: torch.rand(voxels, outputs, generator=torch.Generator().manual_seed(13)) * 2.0 - 1.0 ).to(torch.complex64) - first = _epg_triton.simulate_vjp(tissue, events, seed, **shape, **options) + first = _epg_gpu.simulate_vjp(tissue, events, seed, **shape, **options) # The pass the first-order kernel specializes: zero directions in, the # adjoint out as the gradient with respect to the tangent inputs. still = tuple( torch.zeros_like(value) for value in (*tissue, events[0], events[2], events[3]) ) - _curve, adjoint = _epg_triton.simulate_vjp_jvp( + _curve, adjoint = _epg_gpu.simulate_vjp_jvp( tissue, events, still, seed, **shape, **options ) scale = max(float(value.abs().max()) for value in first if value.numel()) @@ -610,8 +504,7 @@ def _shimmed(voxels: int, states: int) -> None: pooling. Nothing else here drives a transmit array, so this is the only case that would show one standing in for the other. """ - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._lineshape import lineshape_table shims, echoes = 2, 6 @@ -644,7 +537,7 @@ def _shimmed(voxels: int, states: int) -> None: torch.arange(events[6].numel(), dtype=torch.int32) % shims ).contiguous() events = tuple(events) - assert _epg_triton._shim_count(tissue) == shims, ( + assert _epg_gpu._shim_count(tissue) == shims, ( "the transmit array did not reach the kernel" ) @@ -654,7 +547,7 @@ def _shimmed(voxels: int, states: int) -> None: - 1.0 ).to(torch.complex64) without, with_table = _both( - lambda: _epg_triton.simulate_vjp( + lambda: _epg_gpu.simulate_vjp( tissue, events, seed, @@ -669,7 +562,7 @@ def _shimmed(voxels: int, states: int) -> None: assert worst <= WIDE_TOLERANCE, f"shimmed adjoint drifted: {worst:.2e}" without, with_table = _both( - lambda: _epg_triton.simulate_vjp_jvp( + lambda: _epg_gpu.simulate_vjp_jvp( tissue, events, ( @@ -707,13 +600,12 @@ def _profiled(voxels: int, states: int) -> None: across the slice positions the table samples, which is a layout only this case exercises. """ - install() import math import numpy as np from blochsim import rf_definition - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._accelerators import _across_the_table from blochsim.sequence._transition import ( SliceTables, @@ -744,7 +636,7 @@ def _profiled(voxels: int, states: int) -> None: tissue = _across_the_table(_tissue(voxels), profile.points) shape = dict(state_count=states, output_count=outputs) - signal = _epg_triton.simulate(tissue, events, profile=profile, **shape) + signal = _epg_gpu.simulate(tissue, events, profile=profile, **shape) expected = simulate_packed( tissue, events, profile=table, locations=profile.points, **shape ) @@ -754,7 +646,7 @@ def _profiled(voxels: int, states: int) -> None: # A table that moved nothing would agree with the oracle and prove neither # of them read it. - plain = _epg_triton.simulate(tissue, events, **shape) + plain = _epg_gpu.simulate(tissue, events, **shape) moved = float((signal - plain).abs().max() / plain.abs().max()) print(f" profiled moves it {moved:.2e}") assert moved > 1e-3, f"the profile changed nothing: {moved:.2e}" @@ -770,13 +662,12 @@ def _narrowed(voxels: int, states: int) -> None: the end of a buffer, which is silent -- so the two launches are held to each other bit for bit rather than to a tolerance. """ - install() import math from blochsim.sequence import ( TissueProperties, _builders, - _epg_triton, + _epg_gpu, ) from blochsim.sequence._accelerators import real_subspace_axis from blochsim.sequence._parameters import TISSUE_NAMES, features_of @@ -815,10 +706,10 @@ def _narrowed(voxels: int, states: int) -> None: shape = dict(state_count=states, output_count=outputs) for label, axis in (("real", 1), ("complex", -1)): - whole = _epg_triton.simulate( + whole = _epg_gpu.simulate( full, events, real_axis=axis, features=features, **shape ) - narrow = _epg_triton.simulate( + narrow = _epg_gpu.simulate( thin, events, real_axis=axis, features=features, **shape ) drift = float((whole - narrow).abs().max()) @@ -826,7 +717,7 @@ def _narrowed(voxels: int, states: int) -> None: assert drift == 0.0, f"the {label} kernel read a term it was told to drop" -def _case(name: str) -> None: +def _run(name: str) -> None: if name == "narrowed": _narrowed(3, 4) return @@ -857,11 +748,7 @@ def _case(name: str) -> None: if name == "streamed": _streamed(3, 4) return - if name == "unread": - _unread(3, 4) - return - install() - from blochsim.sequence import _builders, _epg_triton + from blochsim.sequence import _builders, _epg_gpu from blochsim.sequence._lineshape import lineshape_table voxels, states = 3, 4 @@ -890,7 +777,7 @@ def _case(name: str) -> None: print(f"{name}: {events[0].numel()} events over {lengths} lengths") without, with_table = _both( - lambda: _epg_triton.simulate( + lambda: _epg_gpu.simulate( tissue, events, state_count=states, output_count=outputs, **options ), force_narrow, @@ -904,7 +791,7 @@ def _case(name: str) -> None: - 1.0 ).to(torch.complex64) without, with_table = _both( - lambda: _epg_triton.simulate_vjp( + lambda: _epg_gpu.simulate_vjp( tissue, events, seed, @@ -933,7 +820,7 @@ def _case(name: str) -> None: torch.zeros_like(events[3]), ) without, with_table = _both( - lambda: _epg_triton.simulate_jvp( + lambda: _epg_gpu.simulate_jvp( tissue, events, directions, @@ -952,7 +839,7 @@ def _case(name: str) -> None: seeded = (*directions, *event_directions) for half, label in ((0, "curvature"), (1, "gradient")): without, with_table = _both( - lambda h=half: _epg_triton.simulate_vjp_jvp( + lambda h=half: _epg_gpu.simulate_vjp_jvp( tissue, events, seeded, @@ -973,7 +860,12 @@ def _case(name: str) -> None: print(f" chunks {-(-voxels // max(cut))}") -if __name__ == "__main__": - _case(sys.argv[1] if len(sys.argv) > 1 else "narrow") - # The wrapper reads this: a case that fell out early prints no such line. - print("checked") +def run(name: str) -> None: + """One case, with whatever it patched into the launcher put back.""" + from blochsim.sequence import _epg_gpu + + wave = _epg_gpu._trajectory_wave + try: + _run(name) + finally: + _epg_gpu._trajectory_wave = wave diff --git a/tests/utils/interpreted_pools.py b/tests/utils/host_pools.py similarity index 86% rename from tests/utils/interpreted_pools.py rename to tests/utils/host_pools.py index d3ffb917..6823e709 100644 --- a/tests/utils/interpreted_pools.py +++ b/tests/utils/host_pools.py @@ -1,15 +1,12 @@ -"""The Triton kernels for tabulated pools, held to the C++ ones in the interpreter. +"""The GPU kernels for tabulated pools, run on the host against the C++ ones. -``TRITON_INTERPRET=1`` runs a kernel in Python over host tensors, so the -device kernels meet the host ones on the same buffers: the tables, the +The device kernels meet the host ones on the same buffers: the tables, the trajectory the adjoint records and walks back, the per-problem cotangent slots and every rotation mode. The C++ kernels are held to the state machine written -out in torch by ``tests/sequence/test_many_pools.py``; these hold the Triton +out in torch by ``tests/sequence/test_many_pools.py``; these hold the GPU kernels to them, to float32 round-off. -Run as ``python -m utils.interpreted_pools -[variant ...]`` with ``TRITON_INTERPRET=1`` set before the process starts. -```` is the exchanging pool count followed by ``s`` for a semisolid pool +``pools`` is the exchanging pool count followed by ``s`` for a semisolid pool or ``f`` for none. ``trains`` packs two trains of different lengths and flips, ``chunked`` records one problem at a time, and ``undeclared`` turns off every optional term the tissue does not declare. @@ -17,11 +14,9 @@ from __future__ import annotations -import sys - import torch -from blochsim.sequence import _accelerators, _pools, _pools_triton +from blochsim.sequence import _accelerators, _pools, _pools_gpu from blochsim.sequence._lineshape import lineshape_table from blochsim.sequence._parameters import NO_GEOMETRY from blochsim.sequence._transition import DynamicPairs @@ -59,7 +54,7 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl ) if "chunked" in variants: # Less than one problem's trajectory, so every problem is a wave. - _pools_triton._TRAJECTORY_BUDGET_BYTES = 1 + _pools_gpu._TRAJECTORY_BUDGET_BYTES = 1 features = frozenset() if "undeclared" in variants else None geometry = NO_GEOMETRY if "undeclared" in variants else GEOMETRY @@ -127,7 +122,7 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl ), ) device = ( - _pools_triton.simulate( + _pools_gpu.simulate( tissue, events, state_count=STATES, output_count=OUTPUTS, **shared ), ) @@ -147,7 +142,7 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl ), ) device = ( - _pools_triton.simulate_jvp( + _pools_gpu.simulate_jvp( tissue, events, tissue_tangents, @@ -162,7 +157,7 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl host = _accelerators._run_packed_vjp( tissue, events, seed, STATES, OUTPUTS, 1, exchanging=layout, **shared ) - device = _pools_triton.simulate_vjp( + device = _pools_gpu.simulate_vjp( tissue, events, seed, state_count=STATES, output_count=OUTPUTS, **shared ) else: @@ -183,7 +178,7 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl (), ) device = sum( - _pools_triton.simulate_vjp_jvp( + _pools_gpu.simulate_vjp_jvp( tissue, events, tangents, @@ -213,9 +208,11 @@ def _launch(pass_name: str, pools: str, rotation: str, variants: set[str]) -> fl return worst -if __name__ == "__main__": - pass_name, pools, rotation, *rest = sys.argv[1:] - worst = _launch(pass_name, pools, rotation, set(rest)) - print(f"{pass_name} {pools} {rotation} {' '.join(rest)}: {worst:.2e}") +def check(pass_name: str, pools: str, rotation: str, *variants: str) -> None: + """One pass of both backends, held to each other.""" + budget = _pools_gpu._TRAJECTORY_BUDGET_BYTES + try: + worst = _launch(pass_name, pools, rotation, set(variants)) + finally: + _pools_gpu._TRAJECTORY_BUDGET_BYTES = budget assert worst <= TOLERANCE, f"{worst:.2e} past {TOLERANCE:.0e}" - print("checked") From 2b253fc616bf3b1cc0896628c0f27364c53c9f5f Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 17:58:23 +0000 Subject: [PATCH 02/16] Compile the host build of the GPU kernels with /bigobj under MSVC Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H --- CMakeLists.txt | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/CMakeLists.txt b/CMakeLists.txt index 57eff498..ab5d0adf 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -63,6 +63,11 @@ blochsim_add_kernel(_perk_cpu src/blochsim/_perk_cpu.cpp) option(BLOCHSIM_HOST_KERNELS "Compile the GPU kernels for the host, for the tests" ON) if(BLOCHSIM_HOST_KERNELS) blochsim_add_kernel(_gpu_host src/blochsim/_gpu_host.cpp) + if(MSVC) + # Every kernel's tile instantiations in one object, and ``#pragma + # unroll``, which only nvcc reads. + target_compile_options(_gpu_host PRIVATE /bigobj /wd4068) + endif() endif() # The GPU kernels compiled ahead of time for the card, wherever a CUDA From 8f508bf40dd0584735aa4c17bed812d71e2947f0 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:53:47 +0000 Subject: [PATCH 03/16] Build and test the host build of the GPU kernels on Linux only Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H --- .github/skills/build-and-test/SKILL.md | 3 ++- CLAUDE.md | 4 ++-- CMakeLists.txt | 17 +++++++++-------- docs/developer_guide.md | 4 ++-- tests/estimators/test_perk_kernel.py | 2 +- tests/sequence/test_both_pools.py | 2 ++ tests/sequence/test_host_kernels.py | 2 ++ tests/sequence/test_many_pools_host.py | 2 ++ 8 files changed, 22 insertions(+), 14 deletions(-) diff --git a/.github/skills/build-and-test/SKILL.md b/.github/skills/build-and-test/SKILL.md index 6eb5e8a5..fe22f87e 100644 --- a/.github/skills/build-and-test/SKILL.md +++ b/.github/skills/build-and-test/SKILL.md @@ -44,7 +44,8 @@ spoiling, the two-pool and three-pool longitudinal steps — while `sequence/`, ## The GPU kernels, without a GPU -The install compiles the GPU kernels for the host as `blochsim._gpu_host`, and +On Linux the install compiles the GPU kernels for the host as +`blochsim._gpu_host`, and for the card as `blochsim._gpu` wherever CMake finds `nvcc` (`--config-settings=cmake.define.BLOCHSIM_CUDA=ON` insists on it). `tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` run the diff --git a/CLAUDE.md b/CLAUDE.md index c8014256..b4596aba 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -63,8 +63,8 @@ for `nvcc`, and `BLOCHSIM_CUDA=ON` or `OFF` (`--config-settings=cmake.define.BLO overrides what it finds; `CMAKE_CUDA_ARCHITECTURES` picks the cards. A kernel compiles in a file of its own, twice -- bounded to 256 threads and to 1024 -- so a build is minutes of `nvcc` spread over as many cores as Ninja is given. -The same kernels compiled for the host (`_gpu_host`) run one program at a -time over host tensors; `tests/sequence/test_host_kernels.py` and +On Linux the same kernels are also compiled for the host (`_gpu_host`), one +program at a time over host tensors; `tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` hold them to the C++ kernels, which is how the GPU path is verified on a machine with no card. diff --git a/CMakeLists.txt b/CMakeLists.txt index ab5d0adf..d5046787 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -59,22 +59,23 @@ blochsim_add_kernel(_perk_cpu src/blochsim/_perk_cpu.cpp) # The GPU kernels compiled for the host, one program at a time. Nothing in the # package dispatches to them; they are how the suite checks the GPU kernels on -# a machine with no card. -option(BLOCHSIM_HOST_KERNELS "Compile the GPU kernels for the host, for the tests" ON) +# a machine with no card. Linux only, like the card's build. +if(CMAKE_SYSTEM_NAME STREQUAL "Linux") + set(_blochsim_host_default ON) +else() + set(_blochsim_host_default OFF) +endif() +option(BLOCHSIM_HOST_KERNELS "Compile the GPU kernels for the host, for the tests" + ${_blochsim_host_default}) if(BLOCHSIM_HOST_KERNELS) blochsim_add_kernel(_gpu_host src/blochsim/_gpu_host.cpp) - if(MSVC) - # Every kernel's tile instantiations in one object, and ``#pragma - # unroll``, which only nvcc reads. - target_compile_options(_gpu_host PRIVATE /bigobj /wd4068) - endif() endif() # The GPU kernels compiled ahead of time for the card, wherever a CUDA # compiler is found unless BLOCHSIM_CUDA says otherwise. include(CheckLanguage) check_language(CUDA) -if(CMAKE_CUDA_COMPILER) +if(CMAKE_CUDA_COMPILER AND CMAKE_SYSTEM_NAME STREQUAL "Linux") set(_blochsim_cuda_default ON) else() set(_blochsim_cuda_default OFF) diff --git a/docs/developer_guide.md b/docs/developer_guide.md index 0914dea3..7f7b76c0 100644 --- a/docs/developer_guide.md +++ b/docs/developer_guide.md @@ -251,8 +251,8 @@ diffusion, flow, spoiling, the two-pool and three-pool longitudinal steps -- while `sequence/`, `model/`, `estimators/`, `recon/` and `optim/` cover the layers above. -**The GPU kernels run without a card.** The install compiles them for the -host as well, one program at a time over host tensors, and +**The GPU kernels run without a card.** On Linux the install compiles them +for the host as well, one program at a time over host tensors, and `tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` hold that build to the C++ kernels. That is how the GPU kernels are verified on a machine with no card; the tests that need one skip themselves. diff --git a/tests/estimators/test_perk_kernel.py b/tests/estimators/test_perk_kernel.py index 5c1bbaaf..fee8e26e 100644 --- a/tests/estimators/test_perk_kernel.py +++ b/tests/estimators/test_perk_kernel.py @@ -214,7 +214,7 @@ def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint() -> None: blocks the kernels walk them in, so every edge of the tiling is read. """ gpu = pytest.importorskip("blochsim.estimators._perk_gpu") - pytest.importorskip("blochsim._gpu_host") + pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") generator = torch.Generator().manual_seed(0) voxels, contrasts, features, parameters = 300, 37, 45, 17 signals = torch.randn(voxels, contrasts, generator=generator) diff --git a/tests/sequence/test_both_pools.py b/tests/sequence/test_both_pools.py index b38f10e9..fafe9213 100644 --- a/tests/sequence/test_both_pools.py +++ b/tests/sequence/test_both_pools.py @@ -1475,6 +1475,8 @@ def test_the_series_carries_the_answer_up_to_the_spread_it_is_trusted_to() -> No from blochsim.sequence._parameters import NARROW_SPREAD device = "cuda" if torch.cuda.is_available() and available() else "cpu" + if device == "cpu": + pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") table_kernel = Kernel("_three_pool_table_kernel") voxels = 4096 block = 1024 diff --git a/tests/sequence/test_host_kernels.py b/tests/sequence/test_host_kernels.py index 632d1286..898ba922 100644 --- a/tests/sequence/test_host_kernels.py +++ b/tests/sequence/test_host_kernels.py @@ -12,6 +12,8 @@ import pytest +pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + from utils import host_kernels diff --git a/tests/sequence/test_many_pools_host.py b/tests/sequence/test_many_pools_host.py index 761ea93f..939bda43 100644 --- a/tests/sequence/test_many_pools_host.py +++ b/tests/sequence/test_many_pools_host.py @@ -10,6 +10,8 @@ import pytest +pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + from utils.host_pools import check PASSES = ("forward", "jvp", "vjp", "vjp_jvp") From 562a6cbf6a1d66bc22997ec2285f2ffe61471065 Mon Sep 17 00:00:00 2001 From: mcencini Date: Wed, 7 Oct 2026 13:17:56 +0200 Subject: [PATCH 04/16] Tile the PERK kernels as matrix products The ported PERK kernels formed each voxel's features one thread at a time, reading every frequency from global memory once per multiply. They ran 2.5x slower than the Triton kernels they replaced. Each program now stages a block of signals and frequencies in shared memory, and each thread forms two voxels by a block of features in registers, reading weights four at a time. The accumulators are stored through unrolled loops, so they stay in registers and nothing spills. On an RTX 4060 (sm_89), 200k voxels, 24 contrasts, 256 features: forward 2.4 ms against 5.1 ms for Triton and 12.9 ms before; adjoint 4.6 ms against 5.4 ms and 14.0 ms. Below about 4k voxels a program's serial chain over the features is longer than Triton's, by up to 45 us. On the host the same source runs a program's threads one after another between barriers, so the host build still checks the card's code. The parity test now runs on the card too. Co-Authored-By: Claude Opus 5.5 --- src/blochsim/_kernels.hpp | 4 +- src/blochsim/_perk_kernels.hpp | 390 ++++++++++++++++++++++----- src/blochsim/estimators/_perk_gpu.py | 8 +- tests/estimators/test_perk_kernel.py | 21 +- 4 files changed, 339 insertions(+), 84 deletions(-) diff --git a/src/blochsim/_kernels.hpp b/src/blochsim/_kernels.hpp index 50ce00f1..b76cca32 100644 --- a/src/blochsim/_kernels.hpp +++ b/src/blochsim/_kernels.hpp @@ -819,8 +819,8 @@ inline constexpr KernelInfo KERNELS[] = { {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1}, {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69}, {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87}, - {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "pppppppiiiifi", 12, -1, -1}, - {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "ppppppiiiifi", 11, -1, -1}, + {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,threads", "pppppppiiiifi", 12, -1, -1}, + {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,threads", "ppppppiiiifi", 11, -1, -1}, }; #define BLOCHSIM_FOR_EACH_KERNEL(X) \ diff --git a/src/blochsim/_perk_kernels.hpp b/src/blochsim/_perk_kernels.hpp index b77ffb4b..d03622a6 100644 --- a/src/blochsim/_perk_kernels.hpp +++ b/src/blochsim/_perk_kernels.hpp @@ -1,112 +1,358 @@ -// The PERK feature map and its regression, fused, one voxel per thread. +// The PERK feature map and its regression, fused. // // y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T // -// A block of features is formed and consumed into the output accumulator in -// registers, so the ``(voxels, features)`` matrix never exists. The adjoint -// does the same and forms the angles again rather than keeping them. Every -// array is contiguous and row-major. +// Two matrix products with a cosine between, tiled the way a matrix product +// is: a program stages a block of signals and a block of frequencies in shared +// memory, and each thread forms the angles of VOXELS voxels by a block of +// features in registers, so every value read from shared memory feeds several +// multiplies. The cosine and the second product consume the block while it is +// still in registers, so the ``(voxels, features)`` matrix never exists. The +// adjoint forms the angles again rather than keeping them. Every array is +// contiguous and row-major. +// +// Staged blocks are zero past the last feature, parameter or voxel, and a zero +// weight is what removes a padded feature from both products, so the inner +// loops carry no bounds. On a card a program is THREADS threads; on the host +// the same source runs them one after another between barriers. + +// Threads per program, and the voxels each holds. ``_perk_gpu`` launches +// programs of ``THREADS`` threads over ``THREADS * VOXELS`` voxels. +constexpr int THREADS = 64; +constexpr int VOXELS = 2; +constexpr int BLOCK_VOXELS = THREADS * VOXELS; +// Contrasts staged at once. +constexpr int CONTRASTS = 32; +// Features formed at once by the forward pass, and parameters it accumulates. +constexpr int FEATURES = 32; +constexpr int PARAMETERS = 8; +// Features formed at once by the adjoint, and contrasts of the gradient it +// accumulates: it holds the angles, their cotangents and the gradient at once. +constexpr int ADJOINT_FEATURES = 16; +constexpr int GRADIENT = 32; +// A row of staged weights is read four at a time, so it is padded to keep each +// row 16-byte aligned and the rows on different banks. +constexpr int PAD = 4; + +#if defined(BLOCHSIM_SIMT) +#define PERK_SHARED __shared__ +#define PERK_SHARED_ROWS __shared__ __align__(16) +#define PERK_SYNC() __syncthreads() +// The body runs once, as this thread. +#define PERK_EACH_THREAD(t) \ + for (int t = static_cast(threadIdx.x), t##_once = 1; t##_once; t##_once = 0) +// What a thread keeps across a barrier: its own registers. +template +struct Own { + T value; + BSK_HD T& operator[](int) { return value; } +}; +#else +#define PERK_SHARED static thread_local +#define PERK_SHARED_ROWS alignas(16) static thread_local +#define PERK_SYNC() static_cast(0) +// The body runs as each thread in turn, so a barrier is the end of the loop. +#define PERK_EACH_THREAD(t) for (int t = 0; t < THREADS; ++t) +template +struct Own { + T value[THREADS]; + T& operator[](int t) { return value[t]; } +}; +#endif -// Features formed at once, which is how often a voxel's signal is read. -constexpr int FEATURE_BLOCK = 32; -// Parameters accumulated at once by the forward pass. -constexpr int PARAMETER_BLOCK = 16; -// Contrasts accumulated at once by the adjoint. -constexpr int CONTRAST_BLOCK = 32; +struct alignas(16) Four { + float x, y, z, w; +}; -// The angles ``W @ x + b`` of features ``first`` onward, as many as there are. -template -BSK_HD void _angles(const float* signal, const float* frequency, const float* phase, - const Voxel& voxel, const Live& live, std::int64_t contrasts, - std::int64_t features, std::int64_t first, bsk::V* angle) { - for (int j = 0; j < FEATURE_BLOCK; ++j) { - angle[j] = bsk::V(first + j < features ? phase[first + j] : 0.0f); +BSK_HD const Four& four(const float* address) { + return *reinterpret_cast(address); +} + +BSK_HD int clamp_width(std::int64_t left, int block) { + return left < block ? static_cast(left) : block; +} + +// Signals of voxels ``first`` onward, contrasts ``c0`` onward, ``width`` of +// them, into ``staged[contrast][voxel]``. Consecutive threads copy consecutive +// elements of the rows, which are contiguous in memory. +BSK_HD void _stage_signals(const float* signal, std::int64_t voxels, std::int64_t contrasts, + std::int64_t first, std::int64_t c0, int width, + float (*staged)[BLOCK_VOXELS + 1], int t) { + for (int i = t; i < BLOCK_VOXELS * width; i += THREADS) { + const int v = i / width; + const int c = i - v * width; + const std::int64_t voxel = first + v; + staged[c][v] = voxel < voxels ? signal[voxel * contrasts + c0 + c] : 0.0f; } - for (std::int64_t contrast = 0; contrast < contrasts; ++contrast) { - const auto value = bsk::ld(signal + voxel * contrasts + contrast, live, 0.0f); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j < features) { - angle[j] = angle[j] + value * frequency[(first + j) * contrasts + contrast]; +} + +// Frequencies of ``width`` features ``f0`` onward against contrasts ``c0`` +// onward, into ``staged[contrast][feature]``, zero past the last feature. +template +BSK_HD void _stage_frequencies(const float* frequency, std::int64_t contrasts, + std::int64_t f0, int features, std::int64_t c0, int width, + float (*staged)[BLOCK + PAD], int t) { + for (int i = t; i < BLOCK * width; i += THREADS) { + const int j = i / width; + const int c = i - j * width; + staged[c][j] = j < features ? frequency[(f0 + j) * contrasts + c0 + c] : 0.0f; + } +} + +// The angles ``W @ x + b`` of a thread's voxels over a block of features, +// accumulated over the staged contrasts. +template +BSK_HD void _accumulate_angles(const float (*signals)[BLOCK_VOXELS + 1], + const float (*frequencies)[BLOCK + PAD], int width, + float (&angle)[VOXELS][BLOCK], int t) { + for (int c = 0; c < width; ++c) { + float x[VOXELS]; +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + x[r] = signals[c][t + r * THREADS]; + } +#pragma unroll + for (int j = 0; j < BLOCK; j += 4) { + const Four w = four(&frequencies[c][j]); +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + angle[r][j] += x[r] * w.x; + angle[r][j + 1] += x[r] * w.y; + angle[r][j + 2] += x[r] * w.z; + angle[r][j + 3] += x[r] * w.w; + } + } + } +} + +// The angles of every voxel of the program over ``width`` features ``f0`` +// onward, the phases already staged in ``offset``. Ends at a barrier. +template +BSK_HD void _angles(const float* signal, const float* frequency, std::int64_t voxels, + std::int64_t contrasts, std::int64_t first, std::int64_t f0, int width, + const float* offset, float (*signals)[BLOCK_VOXELS + 1], + float (*frequencies)[BLOCK + PAD], Own& angle) { + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { + angle[t][r][j] = offset[j]; } } } + for (std::int64_t c0 = 0; c0 < contrasts; c0 += CONTRASTS) { + const int staged = clamp_width(contrasts - c0, CONTRASTS); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + _stage_signals(signal, voxels, contrasts, first, c0, staged, signals, t); + _stage_frequencies(frequency, contrasts, f0, width, c0, staged, frequencies, + t); + } + PERK_SYNC(); + PERK_EACH_THREAD(t) { + _accumulate_angles(signals, frequencies, staged, angle[t], t); + } + } + PERK_SYNC(); } -// One block of voxels, from signal to parameters. +// One program of voxels, from signal to parameters. BSK_HD void _regress_kernel(const float* signal, const float* frequency, const float* phase, const float* feature_mean, const float* weight, const float* parameter_mean, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, - std::int64_t parameters, float scale, std::int64_t BLOCK_VOXELS) { - const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); - const auto live = voxel < voxels; - bsk::V angle[FEATURE_BLOCK]; - for (std::int64_t base = 0; base < parameters; base += PARAMETER_BLOCK) { - bsk::V total[PARAMETER_BLOCK]; - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - total[k] = bsk::V(0.0f); + std::int64_t parameters, float scale, std::int64_t threads) { + static_cast(threads); + PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; + PERK_SHARED_ROWS float frequencies[CONTRASTS][FEATURES + PAD]; + PERK_SHARED_ROWS float weights[FEATURES][PARAMETERS]; + PERK_SHARED float offset[FEATURES]; + PERK_SHARED float mean[FEATURES]; + const std::int64_t first = bsk::program_id(0) * BLOCK_VOXELS; + Own angle; + Own total; + for (std::int64_t p0 = 0; p0 < parameters; p0 += PARAMETERS) { + const int held = clamp_width(parameters - p0, PARAMETERS); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int k = 0; k < PARAMETERS; ++k) { + total[t][r][k] = 0.0f; + } + } } - for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { - _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j >= features) { - break; + for (std::int64_t f0 = 0; f0 < features; f0 += FEATURES) { + const int width = clamp_width(features - f0, FEATURES); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int j = t; j < FEATURES; j += THREADS) { + offset[j] = j < width ? phase[f0 + j] : 0.0f; + mean[j] = j < width ? feature_mean[f0 + j] : 0.0f; + } + for (int i = t; i < FEATURES * PARAMETERS; i += THREADS) { + const int j = i / PARAMETERS; + const int k = i - j * PARAMETERS; + weights[j][k] = + j < width && k < held ? weight[(p0 + k) * features + f0 + j] : 0.0f; } - const auto mapped = scale * bsk::cos(angle[j]) - feature_mean[first + j]; - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - if (base + k < parameters) { - total[k] = total[k] + mapped * weight[(base + k) * features + first + j]; + } + PERK_SYNC(); + _angles(signal, frequency, voxels, contrasts, first, f0, width, offset, + signals, frequencies, angle); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int j = 0; j < FEATURES; ++j) { + const float m = mean[j]; +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const float mapped = scale * cosf(angle[t][r][j]) - m; +#pragma unroll + for (int k = 0; k < PARAMETERS; k += 4) { + const Four w = four(&weights[j][k]); + total[t][r][k] += mapped * w.x; + total[t][r][k + 1] += mapped * w.y; + total[t][r][k + 2] += mapped * w.z; + total[t][r][k + 3] += mapped * w.w; + } } } } } - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - if (base + k < parameters) { - bsk::st(output + voxel * parameters + base + k, - total[k] + parameter_mean[base + k], live); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; +#pragma unroll + for (int k = 0; k < PARAMETERS; ++k) { + if (voxel < voxels && k < held) { + output[voxel * parameters + p0 + k] = + total[t][r][k] + parameter_mean[p0 + k]; + } + } } } } } -// The derivative of one block of voxels with respect to their signals. +// The derivative of one program of voxels with respect to their signals. BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, const float* phase, const float* weight, const float* cotangent, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, std::int64_t parameters, float scale, - std::int64_t BLOCK_VOXELS) { - const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); - const auto live = voxel < voxels; - bsk::V angle[FEATURE_BLOCK]; - for (std::int64_t base = 0; base < contrasts; base += CONTRAST_BLOCK) { - bsk::V gradient[CONTRAST_BLOCK]; - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - gradient[c] = bsk::V(0.0f); + std::int64_t threads) { + static_cast(threads); + constexpr int BLOCK = ADJOINT_FEATURES; + PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; + PERK_SHARED_ROWS float frequencies[CONTRASTS][BLOCK + PAD]; + PERK_SHARED_ROWS float back[BLOCK][GRADIENT + PAD]; + PERK_SHARED_ROWS float weights[PARAMETERS][BLOCK]; + PERK_SHARED float offset[BLOCK]; + const std::int64_t first = bsk::program_id(0) * BLOCK_VOXELS; + Own angle; + Own through; + Own gradient; + for (std::int64_t g0 = 0; g0 < contrasts; g0 += GRADIENT) { + const int held = clamp_width(contrasts - g0, GRADIENT); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int c = 0; c < GRADIENT; ++c) { + gradient[t][r][c] = 0.0f; + } + } } - for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { - _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j >= features) { - break; + for (std::int64_t f0 = 0; f0 < features; f0 += BLOCK) { + const int width = clamp_width(features - f0, BLOCK); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int j = t; j < BLOCK; j += THREADS) { + offset[j] = j < width ? phase[f0 + j] : 0.0f; + } + for (int i = t; i < BLOCK * held; i += THREADS) { + const int j = i / held; + const int c = i - j * held; + back[j][c] = j < width ? frequency[(f0 + j) * contrasts + g0 + c] : 0.0f; + } + for (int i = t; i < BLOCK * (GRADIENT - held); i += THREADS) { + const int j = i / (GRADIENT - held); + back[j][held + i - j * (GRADIENT - held)] = 0.0f; } - bsk::V through(0.0f); - for (std::int64_t p = 0; p < parameters; ++p) { - through = through - + bsk::ld(cotangent + voxel * parameters + p, live, 0.0f) - * weight[p * features + first + j]; + } + PERK_SYNC(); + _angles(signal, frequency, voxels, contrasts, first, f0, width, offset, + signals, frequencies, angle); + // What reaches each feature from the parameters: cotangent @ weight. + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { + through[t][r][j] = 0.0f; + } + } + } + for (std::int64_t p0 = 0; p0 < parameters; p0 += PARAMETERS) { + const int count = clamp_width(parameters - p0, PARAMETERS); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int i = t; i < PARAMETERS * BLOCK; i += THREADS) { + const int k = i / BLOCK; + const int j = i - k * BLOCK; + weights[k][j] = + k < count && j < width ? weight[(p0 + k) * features + f0 + j] : 0.0f; + } } - through = through * (-scale) * bsk::sin(angle[j]); - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - if (base + c < contrasts) { - gradient[c] = gradient[c] - + through * frequency[(first + j) * contrasts + base + c]; + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int k = 0; k < count; ++k) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; + const float pulled = + voxel < voxels ? cotangent[voxel * parameters + p0 + k] : 0.0f; +#pragma unroll + for (int j = 0; j < BLOCK; j += 4) { + const Four w = four(&weights[k][j]); + through[t][r][j] += pulled * w.x; + through[t][r][j + 1] += pulled * w.y; + through[t][r][j + 2] += pulled * w.z; + through[t][r][j + 3] += pulled * w.w; + } + } + } + } + } + PERK_EACH_THREAD(t) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const float slope = -scale * sinf(angle[t][r][j]) * through[t][r][j]; +#pragma unroll + for (int c = 0; c < GRADIENT; c += 4) { + const Four w = four(&back[j][c]); + gradient[t][r][c] += slope * w.x; + gradient[t][r][c + 1] += slope * w.y; + gradient[t][r][c + 2] += slope * w.z; + gradient[t][r][c + 3] += slope * w.w; + } } } } } - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - if (base + c < contrasts) { - bsk::st(output + voxel * contrasts + base + c, gradient[c], live); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; +#pragma unroll + for (int c = 0; c < GRADIENT; ++c) { + if (voxel < voxels && c < held) { + output[voxel * contrasts + g0 + c] = gradient[t][r][c]; + } + } } } } diff --git a/src/blochsim/estimators/_perk_gpu.py b/src/blochsim/estimators/_perk_gpu.py index b68088b7..909ace13 100644 --- a/src/blochsim/estimators/_perk_gpu.py +++ b/src/blochsim/estimators/_perk_gpu.py @@ -23,7 +23,9 @@ from .._gpu_launch import Kernel, cdiv -#: Voxels per program, one per thread. +#: Threads per program, and the voxels a program holds: ``THREADS`` and +#: ``THREADS * VOXELS`` in ``_perk_kernels.hpp``. +_THREADS = 64 _BLOCK_VOXELS = 128 _regress_kernel = Kernel("_regress_kernel") @@ -71,7 +73,7 @@ def regress( features, parameters, math.sqrt(2.0 / features), - _BLOCK_VOXELS, + _THREADS, ) return output @@ -108,6 +110,6 @@ def regress_vjp( features, parameters, math.sqrt(2.0 / features), - _BLOCK_VOXELS, + _THREADS, ) return output diff --git a/tests/estimators/test_perk_kernel.py b/tests/estimators/test_perk_kernel.py index fee8e26e..1f160011 100644 --- a/tests/estimators/test_perk_kernel.py +++ b/tests/estimators/test_perk_kernel.py @@ -207,14 +207,18 @@ def test_a_single_voxel_goes_through_a_kernel_tiled_for_many(device) -> None: assert estimator(measured[:1]).shape == (1, 2) -def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint() -> None: - """The kernels the card runs, compiled for the host, against Torch. +@pytest.mark.parametrize("device", DEVICES) +def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint(device) -> None: + """The kernels the card runs, on the card and compiled for the host, + against Torch. Feature, parameter and contrast counts that are not multiples of the - blocks the kernels walk them in, so every edge of the tiling is read. + blocks the kernels walk them in, and more voxels than one program holds, + so every edge of the tiling is read. """ gpu = pytest.importorskip("blochsim.estimators._perk_gpu") - pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + if device == "cpu": + pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") generator = torch.Generator().manual_seed(0) voxels, contrasts, features, parameters = 300, 37, 45, 17 signals = torch.randn(voxels, contrasts, generator=generator) @@ -238,10 +242,13 @@ def line(x: torch.Tensor) -> torch.Tensor: expected = line(x) (expected_gradient,) = torch.autograd.grad(expected, x, cotangent.double()) + def on(*tensors: torch.Tensor) -> list[torch.Tensor]: + return [tensor.float().to(device) for tensor in tensors] + estimated = gpu.regress( - signals, frequency, phase, feature_mean, weight, parameter_mean - ) - gradient = gpu.regress_vjp(cotangent, signals, frequency, phase, weight) + *on(signals, frequency, phase, feature_mean, weight, parameter_mean) + ).cpu() + gradient = gpu.regress_vjp(*on(cotangent, signals, frequency, phase, weight)).cpu() torch.testing.assert_close( estimated.double(), expected.detach(), atol=1e-5, rtol=1e-5 From 8f0c08d31850c52516c9ad41f39c623402fb0429 Mon Sep 17 00:00:00 2001 From: mcencini Date: Wed, 7 Oct 2026 15:24:21 +0200 Subject: [PATCH 05/16] WIP: hold several problems per thread in the tile runtime On a card a tile held one element per thread. A value varying along y now holds Y_LANES rows per thread in registers, set per kernel in _lanes.hpp; the launcher divides a program's rows by it to size the block, and Python sizes problems per program from the kernel's lanes. One-row values compile to the same straight-line code as before. The real forward kernel holds 8 lanes, takes the cosine and sine of a uniform flip once, and indexes in 32 bits. Measured against Triton on one captured launch (100k atoms, 32 states), it is 1.48x slower at 4 lanes, against 2.25x for the port: the lanes help, the generated code caps them. Not for merging as is: the lane counts other than the real forward kernel are unswept, and the 32-bit indices need a bound on problems x outputs. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 9 +- src/blochsim/_epg_kernels.hpp | 48 +++++-- src/blochsim/_gpu.cu | 8 +- src/blochsim/_gpu_kernel.cu.in | 3 + src/blochsim/_gpu_launch.py | 14 +- src/blochsim/_kernels.hpp | 31 +++-- src/blochsim/_lanes.hpp | 23 ++++ src/blochsim/_launch.hpp | 9 +- src/blochsim/_tile.hpp | 209 +++++++++++++++++++++-------- src/blochsim/sequence/_epg_gpu.py | 29 ++-- tests/sequence/test_cuda_parity.py | 16 +-- 11 files changed, 286 insertions(+), 113 deletions(-) create mode 100644 src/blochsim/_lanes.hpp diff --git a/CLAUDE.md b/CLAUDE.md index b4596aba..b09d4413 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -68,9 +68,12 @@ program at a time over host tensors; `tests/sequence/test_host_kernels.py` and `test_many_pools_host.py` hold them to the C++ kernels, which is how the GPU path is verified on a machine with no card. -**A block is at most 1024 threads.** A tile is an element per thread, so an -EPG launch with more than 1024 state orders, or a pooled one whose orders by -pools pass 1024, is refused with the kernel's name rather than launched. +**A block is at most 1024 threads.** A tile holds one state order per thread, +so an EPG launch with more than 1024 state orders, or a pooled one whose orders +by pools pass 1024, is refused with the kernel's name rather than launched. The +problems a program carries are another matter: a thread holds `Y_LANES` of +them in registers, set per kernel in `src/blochsim/_lanes.hpp`, so that one +reading of each event serves all of them. **`--cov` is on by default** through `addopts`, so a bare `pytest` writes `coverage.xml`. It is ignored, not tracked. diff --git a/src/blochsim/_epg_kernels.hpp b/src/blochsim/_epg_kernels.hpp index 395558de..4ee56798 100644 --- a/src/blochsim/_epg_kernels.hpp +++ b/src/blochsim/_epg_kernels.hpp @@ -9154,7 +9154,21 @@ BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, bsk::atomic_add(((grad_tissue_tangent + ((7 + past_transmit) * atom_count)) + atom), grad_damping_tangent, active_atom); } -BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count_wide, std::int64_t train_count_wide, std::int64_t event_count_wide, std::int64_t output_count_wide, std::int64_t state_count_wide, std::int64_t single_train_wide, std::int64_t atom_stride_wide, std::int64_t shimmed_wide, std::int64_t diffusing_wide, std::int64_t transmit_wide, std::int64_t density_wide, std::int64_t inverting_wide, std::int64_t block_states_wide, std::int64_t problems_wide) { + const std::int32_t atom_count = static_cast(atom_count_wide); + const std::int32_t train_count = static_cast(train_count_wide); + const std::int32_t event_count = static_cast(event_count_wide); + const std::int32_t output_count = static_cast(output_count_wide); + const std::int32_t state_count = static_cast(state_count_wide); + const std::int32_t single_train = static_cast(single_train_wide); + const std::int32_t atom_stride = static_cast(atom_stride_wide); + const std::int32_t shimmed = static_cast(shimmed_wide); + const std::int32_t diffusing = static_cast(diffusing_wide); + const std::int32_t transmit = static_cast(transmit_wide); + const std::int32_t density = static_cast(density_wide); + const std::int32_t inverting = static_cast(inverting_wide); + const std::int32_t block_states = static_cast(block_states_wide); + const std::int32_t problems = static_cast(problems_wide); bsk::V alpha{}; bsk::V atom_b1{}; bsk::V atom_damping{}; @@ -9171,7 +9185,7 @@ BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* bsk::V plus{}; bsk::V pulse_b1{}; bsk::V relaxes{}; - auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); + auto problem = ((static_cast(bsk::program_id(0)) * problems) + bsk::arange_y()); auto state = bsk::arange_x(); auto active_atom = (problem < (train_count * atom_count)); auto state_mask = bsk::band((state < state_count), active_atom); @@ -9278,17 +9292,27 @@ BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* if (bsk::truth((bsk::truth(is_rf) && bsk::truth(is_inversion)))) { longitudinal = ((-atom_inversion) * longitudinal); } else if (bsk::truth(is_rf)) { - alpha = _event_value(flip, event_base, event, active_atom, single_train); - pulse_b1 = atom_b1; - // One shim is the whole sequence's transmit field, loaded once - // above; several give each pulse the row of the shim it drives. - if (bsk::truth((bsk::truth(shimmed) && bsk::truth(transmit)))) { - auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); - pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + decltype(bsk::cos(alpha)) cosine; + decltype(bsk::sin(alpha)) sine; + if (bsk::truth(single_train) && !bsk::truth(transmit)) { + // One train and no transmit field: every problem's pulse is + // the same angle, so its cosine and sine are taken once. + const float angle = bsk::ld(flip + event); + cosine = cosf(angle); + sine = sinf(angle); + } else { + alpha = _event_value(flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + // One shim is the whole sequence's transmit field, loaded once + // above; several give each pulse the row of the shim it drives. + if (bsk::truth((bsk::truth(shimmed) && bsk::truth(transmit)))) { + auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); + } + alpha = (alpha * pulse_b1); + cosine = bsk::cos(alpha); + sine = bsk::sin(alpha); } - alpha = (alpha * pulse_b1); - auto cosine = bsk::cos(alpha); - auto sine = bsk::sin(alpha); auto cosine_half_sq = (0.5f * (1.0f + cosine)); auto sine_half_sq = (0.5f * (1.0f - cosine)); auto half_sine = (0.5f * sine); diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index 476c3747..acfb639a 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -46,6 +46,12 @@ PyObject* launch(PyObject*, PyObject* args) { if (!blochsim_launch::read_launch(name, grid, values, request)) { return nullptr; } + // A thread holds ``lanes`` of a program's rows, so the rows fill whole threads. + if (request.block[1] % request.lanes != 0) { + PyErr_Format(PyExc_ValueError, "%s: %d rows do not fill threads of %d", + bsk::KERNELS[request.kernel].name, request.block[1], request.lanes); + return nullptr; + } if (request.grid[0] > 2147483647LL || request.grid[1] > 65535) { PyErr_SetString(PyExc_ValueError, "the grid is larger than a launch can hold"); return nullptr; @@ -63,7 +69,7 @@ PyObject* launch(PyObject*, PyObject* args) { return cuda_error(status, "selecting the device"); } const dim3 blocks(static_cast(request.grid[0]), static_cast(request.grid[1])); - const dim3 threads(static_cast(request.block[0]), static_cast(request.block[1])); + const dim3 threads(static_cast(request.block[0]), static_cast(request.block[1] / request.lanes)); // A row's reduction and gather go through one word per thread, and a // product with an operator over pools through as many again. const std::size_t shared = diff --git a/src/blochsim/_gpu_kernel.cu.in b/src/blochsim/_gpu_kernel.cu.in index 233418f2..d70864fc 100644 --- a/src/blochsim/_gpu_kernel.cu.in +++ b/src/blochsim/_gpu_kernel.cu.in @@ -7,6 +7,9 @@ // whenever a block fits it. #define BLOCHSIM_SIMT 1 +#include "_lanes.hpp" +#define BLOCHSIM_Y_LANES BLOCHSIM_LANES_@NAME@ + #include "_kernels.hpp" __global__ void __launch_bounds__(256) kernel@NAME@_256(bsk::Arguments arguments, int z) { diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index 0235a3fe..03073dc9 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -39,9 +39,14 @@ def _module(device_type: str) -> Any: @cache +def _entry(name: str) -> tuple[tuple[str, ...], str, int]: + params, kinds, lanes = _module("cuda" if available() else "cpu").kernels()[name] + return tuple(params.split(",")), kinds, lanes + + def _signature(name: str) -> tuple[tuple[str, ...], str]: - params, kinds = _module("cuda" if available() else "cpu").kernels()[name] - return tuple(params.split(",")), kinds + names, kinds, _ = _entry(name) + return names, kinds @cache @@ -60,6 +65,11 @@ class Kernel: def __init__(self, name: str) -> None: self.name = name + @property + def lanes(self) -> int: + """Rows of a program's y axis each thread holds on a card.""" + return _entry(self.name)[2] + def __getitem__(self, grid: tuple[int, ...]) -> Any: def run(*args: Any, **kwargs: Any) -> None: self.launch(tuple(int(count) for count in grid), args, kwargs) diff --git a/src/blochsim/_kernels.hpp b/src/blochsim/_kernels.hpp index b76cca32..0a307d30 100644 --- a/src/blochsim/_kernels.hpp +++ b/src/blochsim/_kernels.hpp @@ -5,6 +5,7 @@ #include +#include "_lanes.hpp" #include "_tile.hpp" #ifndef BLOCHSIM_TABLE_ONLY namespace epg { @@ -45,6 +46,8 @@ struct KernelInfo { int y; // The parameter whose value is the length of z, or -1. int z; + // Rows of y each thread holds on a card. + int lanes; }; #ifndef BLOCHSIM_TABLE_ONLY @@ -807,20 +810,20 @@ BSK_HD void call_regress_vjp_kernel(const Arg* a) { #endif inline constexpr KernelInfo KERNELS[] = { - {"_three_pool_table_jvp_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,d_t1,d_t1_pool_b,d_t1_bound,d_pool_b_exchange,d_bound_exchange,d_pool_b_fraction,d_bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "pppppppppppppppppiii", 18, -1, -1}, - {"_three_pool_table_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "ppppppppppiii", 11, -1, -1}, - {"_epg_vjp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,lineshape,profile,profile_index,pairs,pair_index,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_pair,grad_output_real,grad_output_imag,grad_tissue,grad_flip,grad_phase,grad_duration,trajectory_r,trajectory_i,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,shim_rows,profile_step,lineshape_step,state_count,single_train,atom_stride,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,recording,block_states,problems", "pppppppppppppppppppppppppppppppppppipppppppppiiiiiiffiffiiiiiiiiiiiiiiiiiiiiii", 76, 77, -1}, - {"_epg_vjp_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,grad_pair_value,grad_pair_tangent,dot_t1,dot_t2,dot_m0,dot_b1,dot_b1_phase,dot_b0,dot_inversion_efficiency,dot_diffusion,dot_velocity,dot_bound_fraction,dot_exchange_rate,dot_t1_bound,dot_pool_b_fraction,dot_pool_b_exchange,dot_t1_pool_b,dot_t2_pool_b,dot_pool_b_shift,dot_duration,dot_flip,dot_phase,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_output_real,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_phase_value,grad_phase_tangent,grad_duration_value,grad_duration_tangent,trajectory_vr,trajectory_vi,trajectory_tr,trajectory_ti,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,directed,off_axis,moving,diffusing,transmit,density,inverting,broadened,lineshape_bins,pools,narrow,tabulated,recording,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppppippppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiii", 105, 106, -1}, - {"_epg_real_vjp_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,dot_t1,dot_t2,dot_m0,dot_b1,dot_inversion_efficiency,dot_diffusion,dot_duration,dot_flip,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_duration_value,grad_duration_tangent,trajectory_value,trajectory_tangent,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppiiiiiiiiiiiiiiiii", 44, 45, -1}, - {"_epg_real_vjp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,grad_output_imag,grad_tissue,grad_flip,grad_duration,trajectory_value,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppiiiiiiiiiiiiiiiii", 32, 33, -1}, - {"_epg_real_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppiiiiiiiiiiiiii", 26, 27, -1}, - {"_epg_real_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_inversion_efficiency,tangent_diffusion,tangent_duration,tangent_flip,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppiiiiiiiiiiiiii", 34, 35, -1}, - {"_epg_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,bound_exchange,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,phase_cos,phase_sin,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 65, 66, -1}, - {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1}, - {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69}, - {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87}, - {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,threads", "pppppppiiiifi", 12, -1, -1}, - {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,threads", "ppppppiiiifi", 11, -1, -1}, + {"_three_pool_table_jvp_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,d_t1,d_t1_pool_b,d_t1_bound,d_pool_b_exchange,d_bound_exchange,d_pool_b_fraction,d_bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "pppppppppppppppppiii", 18, -1, -1, BLOCHSIM_LANES__three_pool_table_jvp_kernel}, + {"_three_pool_table_kernel", "t1,t1_pool_b,t1_bound,pool_b_exchange,bound_exchange,pool_b_fraction,bound_fraction,durations,rows,table,voxel_count,BLOCK,narrow", "ppppppppppiii", 11, -1, -1, BLOCHSIM_LANES__three_pool_table_kernel}, + {"_epg_vjp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,lineshape,profile,profile_index,pairs,pair_index,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_pair,grad_output_real,grad_output_imag,grad_tissue,grad_flip,grad_phase,grad_duration,trajectory_r,trajectory_i,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,shim_rows,profile_step,lineshape_step,state_count,single_train,atom_stride,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,recording,block_states,problems", "pppppppppppppppppppppppppppppppppppipppppppppiiiiiiffiffiiiiiiiiiiiiiiiiiiiiii", 76, 77, -1, BLOCHSIM_LANES__epg_vjp_kernel}, + {"_epg_vjp_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,grad_pair_value,grad_pair_tangent,dot_t1,dot_t2,dot_m0,dot_b1,dot_b1_phase,dot_b0,dot_inversion_efficiency,dot_diffusion,dot_velocity,dot_bound_fraction,dot_exchange_rate,dot_t1_bound,dot_pool_b_fraction,dot_pool_b_exchange,dot_t1_pool_b,dot_t2_pool_b,dot_pool_b_shift,dot_duration,dot_flip,dot_phase,duration_row,pool_table,pool_bars,pool_durations,row_count,grad_output_real,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_phase_value,grad_phase_tangent,grad_duration_value,grad_duration_tangent,trajectory_vr,trajectory_vi,trajectory_tr,trajectory_ti,problem_base,problem_end,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,directed,off_axis,moving,diffusing,transmit,density,inverting,broadened,lineshape_bins,pools,narrow,tabulated,recording,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppppippppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiii", 105, 106, -1, BLOCHSIM_LANES__epg_vjp_jvp_kernel}, + {"_epg_real_vjp_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,dot_t1,dot_t2,dot_m0,dot_b1,dot_inversion_efficiency,dot_diffusion,dot_duration,dot_flip,grad_output_imag,grad_tissue_value,grad_tissue_tangent,grad_flip_value,grad_flip_tangent,grad_duration_value,grad_duration_tangent,trajectory_value,trajectory_tangent,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppiiiiiiiiiiiiiiiii", 44, 45, -1, BLOCHSIM_LANES__epg_real_vjp_jvp_kernel}, + {"_epg_real_vjp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,grad_output_imag,grad_tissue,grad_flip,grad_duration,trajectory_value,problem_base,problem_end,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shim_rows,shimmed,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppiiiiiiiiiiiiiiiii", 32, 33, -1, BLOCHSIM_LANES__epg_real_vjp_kernel}, + {"_epg_real_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppiiiiiiiiiiiiii", 26, 27, -1, BLOCHSIM_LANES__epg_real_kernel}, + {"_epg_real_jvp_kernel", "t1,t2,m0,b1,inversion_efficiency,diffusion,duration,kind,flip,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_inversion_efficiency,tangent_diffusion,tangent_duration,tangent_flip,output_real,output_imag,atom_count,train_count,event_count,output_count,state_count,single_train,atom_stride,shimmed,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppiiiiiiiiiiiiii", 34, 35, -1, BLOCHSIM_LANES__epg_real_jvp_kernel}, + {"_epg_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,bound_exchange,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,phase_cos,phase_sin,action,output_index,shim_index,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "pppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 65, 66, -1, BLOCHSIM_LANES__epg_kernel}, + {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1, BLOCHSIM_LANES__epg_jvp_kernel}, + {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69, BLOCHSIM_LANES__pooled_kernel}, + {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87, BLOCHSIM_LANES__pooled_adjoint_kernel}, + {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,threads", "pppppppiiiifi", 12, -1, -1, BLOCHSIM_LANES__regress_kernel}, + {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,threads", "ppppppiiiifi", 11, -1, -1, BLOCHSIM_LANES__regress_vjp_kernel}, }; #define BLOCHSIM_FOR_EACH_KERNEL(X) \ diff --git a/src/blochsim/_lanes.hpp b/src/blochsim/_lanes.hpp new file mode 100644 index 00000000..6924c5c8 --- /dev/null +++ b/src/blochsim/_lanes.hpp @@ -0,0 +1,23 @@ +// Rows of a program's y axis each thread holds on a card, per kernel: the +// ``Y_LANES`` of _tile.hpp. A kernel's file is compiled with its own, and the +// launcher divides a program's rows by it to size the block. +// +// In the EPG kernels y is the problems a program carries, so rows held in a +// thread share its reading of every event. A kernel whose y is something else +// -- the pools of the pooled kernels -- or that has no y holds one. +#pragma once + +#define BLOCHSIM_LANES__three_pool_table_jvp_kernel 1 +#define BLOCHSIM_LANES__three_pool_table_kernel 1 +#define BLOCHSIM_LANES__epg_vjp_kernel 1 +#define BLOCHSIM_LANES__epg_vjp_jvp_kernel 1 +#define BLOCHSIM_LANES__epg_real_vjp_jvp_kernel 1 +#define BLOCHSIM_LANES__epg_real_vjp_kernel 1 +#define BLOCHSIM_LANES__epg_real_kernel 8 +#define BLOCHSIM_LANES__epg_real_jvp_kernel 1 +#define BLOCHSIM_LANES__epg_kernel 1 +#define BLOCHSIM_LANES__epg_jvp_kernel 1 +#define BLOCHSIM_LANES__pooled_kernel 1 +#define BLOCHSIM_LANES__pooled_adjoint_kernel 1 +#define BLOCHSIM_LANES__regress_kernel 1 +#define BLOCHSIM_LANES__regress_vjp_kernel 1 diff --git a/src/blochsim/_launch.hpp b/src/blochsim/_launch.hpp index 57035c0e..e18b3411 100644 --- a/src/blochsim/_launch.hpp +++ b/src/blochsim/_launch.hpp @@ -30,7 +30,7 @@ inline PyObject* kernel_table(PyObject*, PyObject*) { return nullptr; } for (const auto& info : bsk::KERNELS) { - PyObject* entry = Py_BuildValue("(ss)", info.params, info.kinds); + PyObject* entry = Py_BuildValue("(ssi)", info.params, info.kinds, info.lanes); if (entry == nullptr || PyDict_SetItemString(table, info.name, entry) < 0) { Py_XDECREF(entry); Py_DECREF(table); @@ -47,6 +47,8 @@ struct Launch { int block[2] = {1, 1}; // The length of z, which a program holds in each thread. int z = 1; + // The rows each thread holds on a card. + int lanes = 1; bsk::Arguments arguments{}; }; @@ -101,9 +103,10 @@ inline bool read_launch(PyObject* name, PyObject* grid, PyObject* args, Launch& launch.block[0] = info.x < 0 ? 1 : static_cast(launch.arguments.a[info.x].i); launch.block[1] = info.y < 0 ? 1 : static_cast(launch.arguments.a[info.y].i); launch.z = info.z < 0 ? 1 : static_cast(launch.arguments.a[info.z].i); - if (launch.block[0] < 1 || launch.block[1] < 1 || launch.block[0] * launch.block[1] > 1024) { + launch.lanes = info.lanes; + if (launch.block[0] < 1 || launch.block[1] < 1 || launch.block[0] * launch.block[1] > 1024 * launch.lanes) { PyErr_Format(PyExc_ValueError, "%s: a block of %d by %d threads is more than a card runs", - info.name, launch.block[0], launch.block[1]); + info.name, launch.block[0], launch.block[1] / launch.lanes); return false; } if (launch.z < 1 || launch.z > bsk::MAX_Z) { diff --git a/src/blochsim/_tile.hpp b/src/blochsim/_tile.hpp index 272aec61..99f76e39 100644 --- a/src/blochsim/_tile.hpp +++ b/src/blochsim/_tile.hpp @@ -12,8 +12,11 @@ // and a store or an atomic of it is made once rather than once per thread. // Values that vary along neither are plain C++ scalars, uniform over the block. // -// Under ``BLOCHSIM_SIMT`` (the CUDA build) a tile is one element per thread, -// and the operations across a row are warp shuffles or shared memory. Without +// Under ``BLOCHSIM_SIMT`` (the CUDA build) a tile is one element per thread +// along x and z, and ``Y_LANES`` rows of y per thread: a thread holds rows +// ``threadIdx.y * Y_LANES`` onward in registers, so the work a program does +// once per row -- reading an event, branching on it -- is shared by those rows. +// The operations across a row are warp shuffles or shared memory. Without // it, a tile is the whole array, so the same kernel source runs on the host one // program at a time: that is the build the tests run without a card. #pragma once @@ -44,11 +47,17 @@ namespace bsk { extern __shared__ unsigned long long shared_words[]; +// Rows of y each thread holds, set per kernel by _lanes.hpp. +#ifndef BLOCHSIM_Y_LANES +#define BLOCHSIM_Y_LANES 1 +#endif +constexpr int Y_LANES = BLOCHSIM_Y_LANES; + // The length of z, set once by a kernel's entry. __shared__ int z_width; BSK_HD int width_x() { return static_cast(blockDim.x); } -BSK_HD int width_y() { return static_cast(blockDim.y); } +BSK_HD int width_y() { return static_cast(blockDim.y) * Y_LANES; } BSK_HD int width_z() { return z_width; } __device__ __forceinline__ void enter(int nz) { @@ -116,35 +125,48 @@ constexpr bool any_tile = ((axes_of != 0) || ...); template struct V { static_assert(AX > 0 && AX < 8 && AX != 5 && AX != 7, "a tile varies along x or z, y, or both"); - static constexpr int lanes = (AX & 4) ? MAX_Z : 1; + // A thread's rows of y, each holding its lanes of z. + static constexpr int rows = (AX & 2) ? Y_LANES : 1; + static constexpr int depth = (AX & 4) ? MAX_Z : 1; + static constexpr int lanes = rows * depth; T v[lanes]; V() = default; template || std::is_pointer_v, int> = 0> BSK_HD V(U scalar) { - for (int z = 0; z < lanes; ++z) v[z] = static_cast(scalar); + for (int i = 0; i < lanes; ++i) v[i] = static_cast(scalar); } template = 0> BSK_HD V(const V& other) { - for (int z = 0; z < lanes; ++z) v[z] = static_cast(other.v[(BX & 4) ? z : 0]); + if constexpr (rows == 1) { + for (int z = 0; z < depth; ++z) v[z] = static_cast(other.v[(BX & 4) ? z : 0]); + } else { +#pragma unroll + for (int y = 0; y < rows; ++y) { + for (int z = 0; z < depth; ++z) { + v[y * depth + z] = static_cast( + other.v[((BX & 2) ? y : 0) * V::depth + ((BX & 4) ? z : 0)]); + } + } + } } }; -// The element a thread holds at ``z``. +// The element a thread holds in its row ``y`` at ``z``. template -BSK_HD decltype(auto) element(const A& a, int z = 0) { +BSK_HD decltype(auto) element(const A& a, int y = 0, int z = 0) { if constexpr (axes_of == 0) { return a; - } else if constexpr ((axes_of & 4) != 0) { - return (a.v[z]); } else { - return (a.v[0]); + using Tile = std::decay_t; + return (a.v[((axes_of & 2) ? y : 0) * Tile::depth + ((axes_of & 4) ? z : 0)]); } } -// Lanes of z past its length are left unset, so nothing reads memory for them. +// Every row a thread holds; lanes of z past its length are left unset, so +// nothing reads memory for them. template BSK_HD auto zip(F f, const A&... a) { constexpr int AX = (0 | ... | axes_of); @@ -153,15 +175,33 @@ BSK_HD auto zip(F f, const A&... a) { } else { using R = decltype(f(element(a)...)); V out; - if constexpr ((AX & 4) != 0) { - const int nz = width_z(); - for (int z = 0; z < MAX_Z; ++z) { - if (z < nz) { - out.v[z] = f(element(a, z)...); + if constexpr (V::rows == 1) { + // One row: straight-line code, which is what most of a kernel is + // and what the compiler should not have to unroll its way back to. + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + out.v[z] = f(element(a, 0, z)...); + } } + } else { + out.v[0] = f(element(a)...); } } else { - out.v[0] = f(element(a)...); +#pragma unroll + for (int y = 0; y < V::rows; ++y) { + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + out.v[y * MAX_Z + z] = f(element(a, y, z)...); + } + } + } else { + out.v[y] = f(element(a, y)...); + } + } } return out; } @@ -501,7 +541,10 @@ BSK_HD V arange_x() { BSK_HD V arange_y() { #if defined(BLOCHSIM_SIMT) V out; - out.v[0] = static_cast(threadIdx.y); +#pragma unroll + for (int y = 0; y < Y_LANES; ++y) { + out.v[y] = static_cast(threadIdx.y) * Y_LANES + y; + } return out; #else V out; @@ -572,25 +615,42 @@ BSK_HD void each_written(F f) { if (!writes()) { return; } - if constexpr ((AX & 4) != 0) { - const int nz = width_z(); - for (int z = 0; z < MAX_Z; ++z) { - if (z < nz) { - f(z); + constexpr int rows = (AX & 2) ? Y_LANES : 1; + if constexpr (rows == 1) { + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + f(0, z); + } } + } else { + f(0, 0); } } else { - f(0); +#pragma unroll + for (int y = 0; y < rows; ++y) { + if constexpr ((AX & 4) != 0) { + const int nz = width_z(); + for (int z = 0; z < MAX_Z; ++z) { + if (z < nz) { + f(y, z); + } + } + } else { + f(y, 0); + } + } } } template BSK_HD void st(const P& pointer, const T& value, const M& mask) { constexpr int AX = axes_of

| axes_of; - each_written([&](int z) { - if (element(mask, z)) { - auto p = element(pointer, z); - *p = static_cast>(element(value, z)); + each_written([&](int y, int z) { + if (element(mask, y, z)) { + auto p = element(pointer, y, z); + *p = static_cast>(element(value, y, z)); } }); } @@ -598,10 +658,10 @@ BSK_HD void st(const P& pointer, const T& value, const M& mask) { template BSK_HD void atomic_add(const P& pointer, const T& value, const M& mask) { constexpr int AX = axes_of

| axes_of; - each_written([&](int z) { - if (element(mask, z)) { - auto p = element(pointer, z); - atomicAdd(p, static_cast>(element(value, z))); + each_written([&](int y, int z) { + if (element(mask, y, z)) { + auto p = element(pointer, y, z); + atomicAdd(p, static_cast>(element(value, y, z))); } }); } @@ -737,18 +797,6 @@ struct Max { BSK_HD X operator()(X a, X b) const { return s_max(a, b); } }; -// A value with an axis taken out of AX, held where the reduction left it. -template -BSK_HD auto reduced(const T* lanes) { - if constexpr (AX == 0) { - return lanes[0]; - } else { - V out; - for (int z = 0; z < V::lanes; ++z) out.v[z] = lanes[z]; - return out; - } -} - template BSK_HD auto reduce_x(const V& a, Op op) { if constexpr ((AX & 1) == 0) { @@ -756,8 +804,15 @@ BSK_HD auto reduce_x(const V& a, Op op) { } else { constexpr int RX = AX & ~1; #if defined(BLOCHSIM_SIMT) - T total = reduce_row(a.v[0], op); - return reduced(&total); + // x and z are never both in a tile, so each row is one lane. + if constexpr (RX == 0) { + return reduce_row(a.v[0], op); + } else { + V out; +#pragma unroll + for (int y = 0; y < V::rows; ++y) out.v[y] = reduce_row(a.v[y], op); + return out; + } #else if constexpr (RX == 0) { T total = a.at(0, 0); @@ -783,9 +838,22 @@ BSK_HD auto reduce_y(const V& a, Op op) { } else { constexpr int RX = AX & ~2; #if defined(BLOCHSIM_SIMT) - T totals[V::lanes]; - for (int z = 0; z < V::lanes; ++z) totals[z] = reduce_column(a.v[z], op); - return reduced(totals); + // A thread's own rows first, then across the rows of threads. + constexpr int depth = V::depth; + T totals[depth]; + for (int z = 0; z < depth; ++z) { + T total = a.v[z]; +#pragma unroll + for (int y = 1; y < V::rows; ++y) total = op(total, a.v[y * depth + z]); + totals[z] = reduce_column(total, op); + } + if constexpr (RX == 0) { + return totals[0]; + } else { + V out; + for (int z = 0; z < depth; ++z) out.v[z] = totals[z]; + return out; + } #else if constexpr (RX == 0) { T total = a.at(0, 0); @@ -813,14 +881,26 @@ BSK_HD auto reduce_z(const V& a, Op op) { } else { constexpr int RX = AX & ~4; #if defined(BLOCHSIM_SIMT) - T total = a.v[0]; const int nz = width_z(); - for (int z = 1; z < MAX_Z; ++z) { - if (z < nz) { - total = op(total, a.v[z]); + T totals[V::rows]; +#pragma unroll + for (int y = 0; y < V::rows; ++y) { + T total = a.v[y * MAX_Z]; + for (int z = 1; z < MAX_Z; ++z) { + if (z < nz) { + total = op(total, a.v[y * MAX_Z + z]); + } } + totals[y] = total; + } + if constexpr (RX == 0) { + return totals[0]; + } else { + V out; +#pragma unroll + for (int y = 0; y < V::rows; ++y) out.v[y] = totals[y]; + return out; } - return reduced(&total); #else if constexpr (RX == 0) { T total = a.at(0, 0, 0); @@ -883,8 +963,12 @@ template BSK_HD auto gather_x(const V& values, const I& index) { static_assert(AX & 1, "a gather along x reads a value that varies along x"); #if defined(BLOCHSIM_SIMT) - V> out; - out.v[0] = gather_row(values.v[0], static_cast(element(index))); + using Out = V>; + Out out; +#pragma unroll + for (int y = 0; y < Out::rows; ++y) { + out.v[y] = gather_row(element(values, y), static_cast(element(index, y))); + } return out; #else constexpr int RX = AX | axes_of; @@ -911,6 +995,12 @@ BSK_HD auto times(const O& op_in, const P& planes_in, bool transposed) { const V planes(planes_in); V out(T(0)); #if defined(BLOCHSIM_SIMT) + // Only the pooled kernels take this product, and they hold one row to a + // thread; every kernel's file compiles their bodies, none other runs them. + if constexpr (Y_LANES != 1) { + __trap(); + return out; + } // The planes, then the operator, through shared memory: a thread reads the // column of planes beneath its state and the operator's entries it needs. T* words = reinterpret_cast(shared_words); @@ -957,6 +1047,11 @@ BSK_HD auto outer(const L& left_in, const R& right_in) { const V right(right_in); V out(T(0)); #if defined(BLOCHSIM_SIMT) + // As ``times``: the pooled kernels alone, one row to a thread. + if constexpr (Y_LANES != 1) { + __trap(); + return out; + } T* words = reinterpret_cast(shared_words); const int nx = width_x(); const int ny = width_y(); diff --git a/src/blochsim/sequence/_epg_gpu.py b/src/blochsim/sequence/_epg_gpu.py index 0299fad3..acfc9bbe 100644 --- a/src/blochsim/sequence/_epg_gpu.py +++ b/src/blochsim/sequence/_epg_gpu.py @@ -284,8 +284,9 @@ def _three_pool_table( return table -# Elements of the state tile one program carries, one to a thread. -_TILE_ELEMENTS = 64 +# Threads of one program: a warp, its lanes along the states and, where the +# states are fewer, across rows of problems. +_PROGRAM_THREADS = 32 def _atom_stride(*tuples: tuple[torch.Tensor, ...]) -> int: @@ -306,11 +307,13 @@ def _atom_stride(*tuples: tuple[torch.Tensor, ...]) -> int: ) -def _problems_per_program(block_states: int) -> int: +def _problems_per_program(block_states: int, kernel: Kernel) -> int: """How many independent problems to carry on one program's lane axis. A warp's lanes cost about the same whether they are used or not, so packing - several problems into one program is close to free. + several problems into one program is close to free. Each thread then holds + ``kernel.lanes`` problems of its row in registers, and the work it does once + an event -- reading it, branching on it -- serves all of them. It depends on the state count alone, and deliberately not on how many problems the launch has. A run cut into chunks would otherwise compile a @@ -321,8 +324,8 @@ def _problems_per_program(block_states: int) -> int: The result sizes a block of threads, so it must be a power of two. """ - widest = max(1, _TILE_ELEMENTS // block_states) - return 1 << (widest.bit_length() - 1) + rows = max(1, _PROGRAM_THREADS // block_states) + return (1 << (rows.bit_length() - 1)) * kernel.lanes def _output_shape( @@ -469,7 +472,7 @@ def simulate_into( pools = _pool_flag(lineshape, exchanging) block_states = next_power_of_2(state_count) total = train_count * atom_count - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_real_kernel if real_axis == 1 else _epg_kernel) grid = (cdiv(total, problems),) # A kernel argument has to be a tensor even where the branch reading it is # compiled out, so an unprofiled launch passes one it already has. @@ -720,7 +723,7 @@ def simulate_jvp_into( shims = _shim_count(tissue) block_states = next_power_of_2(state_count) total = train_count * atom_count - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_real_jvp_kernel if real_axis == 1 else _epg_jvp_kernel) grid = (cdiv(total, problems),) if real_axis == 1: @@ -976,7 +979,7 @@ def simulate_vjp( for _ in range(2) ] - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_vjp_kernel) for base in range(0, total, wave): span = min(wave, total - base) if pool_bars is not None: @@ -1125,7 +1128,7 @@ def simulate_real_vjp( (wave, event_count * 3 * state_count), dtype=torch.float32, device=device ) - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_real_vjp_kernel) for base in range(0, total, wave): span = min(wave, total - base) _epg_real_vjp_kernel[(cdiv(span, problems),)]( @@ -1412,7 +1415,7 @@ def simulate_vjp_into( grad_real.copy_(grad_output.real.reshape(-1)) grad_imag.copy_(grad_output.imag.reshape(-1)) - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_vjp_kernel) for base in range(0, total, buffers.wave): span = min(buffers.wave, total - base) # The trajectory is written by one launch and walked back by the @@ -1529,7 +1532,7 @@ def simulate_real_vjp_into( grad_imag = buffers.cotangent[1][:size] grad_imag.copy_(grad_output.resolve_conj().imag.reshape(-1)) - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_real_vjp_kernel) for base in range(0, total, buffers.wave): span = min(buffers.wave, total - base) _epg_real_vjp_kernel[(cdiv(span, problems),)]( @@ -1684,7 +1687,7 @@ def simulate_vjp_jvp_into( wave * row_count * 36, dtype=torch.float32, device=t1.device ) - problems = _problems_per_program(block_states) + problems = _problems_per_program(block_states, _epg_real_vjp_jvp_kernel if real else _epg_vjp_jvp_kernel) for base in range(0, total, wave): span = min(wave, total - base) if pool_bars is not None: diff --git a/tests/sequence/test_cuda_parity.py b/tests/sequence/test_cuda_parity.py index cdfa6f7e..2a680a6a 100644 --- a/tests/sequence/test_cuda_parity.py +++ b/tests/sequence/test_cuda_parity.py @@ -228,13 +228,13 @@ def test_an_off_resonance_seed_keeps_the_complex_kernel_on_cuda(): assert torch.equal(automatic, complex_kernel) -@pytest.mark.parametrize("block_states", [4, 16]) -def test_the_packing_width_is_a_power_of_two(block_states): - """It indexes a ``tl.arange``, which rejects anything else.""" - from blochsim.sequence._epg_gpu import _problems_per_program +@pytest.mark.parametrize("block_states", [4, 16, 64]) +def test_the_packing_width_fills_whole_threads(block_states): + """A program's rows are threads of ``lanes`` rows each, in a power of two.""" + from blochsim.sequence._epg_gpu import _epg_real_kernel, _problems_per_program - width = _problems_per_program(block_states) - assert width >= 1 + width = _problems_per_program(block_states, _epg_real_kernel) + assert width % _epg_real_kernel.lanes == 0 assert width & (width - 1) == 0 @@ -246,9 +246,9 @@ def test_the_packing_width_ignores_how_many_problems_there_are(block_states): off the launch size would make a streamed volume answer differently from the same volume run whole. """ - from blochsim.sequence._epg_gpu import _problems_per_program + from blochsim.sequence._epg_gpu import _epg_kernel, _problems_per_program - assert _problems_per_program(block_states) >= 1 + assert _problems_per_program(block_states, _epg_kernel) >= 1 @pytest.mark.parametrize("atoms", [3, 16, 21]) From 15dab4e07666aa9f7ba8e52b9006728e5430ae90 Mon Sep 17 00:00:00 2001 From: mcencini Date: Wed, 7 Oct 2026 19:03:19 +0200 Subject: [PATCH 06/16] Compile each EPG kernel again for the switch combinations listed Every EPG kernel takes its feature switches as arguments, so as compiled once it carries the registers and code of every term it might evaluate. Each entry of _specializations.json compiles the same kernel with those switches as constants: every function is inlined, a constant switch folds, and the terms it turns off are never generated. The launcher runs the entry whose switches a launch matches exactly and the general kernel where none does, so the list decides speed and never correctness. The 124 entries are the combinations the test suite and the benchmark cases launch (scripts/kernel_census.py), less the second-order kernel's combinations with slice profile, motion or pools, whose compile took as long as a dozen others together. The real forward kernel holds 4 problems per thread and the complex forward kernel 2. The kernel source is the port's, unchanged. On an RTX 4060, device time against Triton (main): complex forward 1.32x, complex adjoint 0.78-0.82x, complex second order 1.01x, real forward 1.44x, real adjoint 1.27x; the port was 1.4-4.1x on the same cases. The sm_89 module grows from 7.9 to 15.7 MB. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 15 + CMakeLists.txt | 44 +- scripts/kernel_census.py | 130 ++ src/blochsim/_epg_kernels.hpp | 48 +- src/blochsim/_gpu.cu | 130 +- src/blochsim/_gpu_launch.py | 41 + src/blochsim/_gpu_special.cu.in | 27 + src/blochsim/_lanes.hpp | 4 +- src/blochsim/_special.hpp | 98 + src/blochsim/_specializations.json | 2406 ++++++++++++++++++++ tests/sequence/test_specialized_kernels.py | 105 + 11 files changed, 3007 insertions(+), 41 deletions(-) create mode 100644 scripts/kernel_census.py create mode 100644 src/blochsim/_gpu_special.cu.in create mode 100644 src/blochsim/_special.hpp create mode 100644 src/blochsim/_specializations.json create mode 100644 tests/sequence/test_specialized_kernels.py diff --git a/CLAUDE.md b/CLAUDE.md index b09d4413..f1edb35f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -75,6 +75,21 @@ problems a program carries are another matter: a thread holds `Y_LANES` of them in registers, set per kernel in `src/blochsim/_lanes.hpp`, so that one reading of each event serves all of them. +**A kernel runs compiled for its own switches where one is listed.** Every +EPG kernel takes its feature switches as arguments, so as compiled once it +carries the registers and code of every term it might evaluate. Each entry of +`src/blochsim/_specializations.json` compiles the same kernel again with those +switches as constants (`_special.hpp`): every function is inlined, a constant +switch folds, and the terms it turns off are never generated. The launcher runs +the entry a launch's switches match exactly, and the kernel compiled for all of +them where none does, so the list decides speed and never correctness. +`scripts/kernel_census.py` records the switches launches use and writes the +list; `_gpu_launch.generic_kernels()` runs a block on the general kernels, which +is how `tests/sequence/test_specialized_kernels.py` holds the two to each other. +Each entry is another compile of its kernel, so the list is most of a CUDA +build's time; `--config-settings=cmake.define.BLOCHSIM_SPECIALIZE=OFF` builds +the general kernels alone. + **`--cov` is on by default** through `addopts`, so a bare `pytest` writes `coverage.xml`. It is ignored, not tracked. diff --git a/CMakeLists.txt b/CMakeLists.txt index d5046787..93f40643 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -109,12 +109,54 @@ if(BLOCHSIM_CUDA) list(APPEND _blochsim_gpu_sources "${_source}") endforeach() + # The kernels again, each for the combinations of its feature switches + # listed in _specializations.json: a switch fixed at compile time folds, + # and the terms it turns off are never generated. The launcher runs one + # where a launch's switches match it and the kernel above where none does. + # OFF builds the kernels above alone, which is the quick build for work on + # them. + option(BLOCHSIM_SPECIALIZE "Compile the kernels for the switch combinations listed" ON) + set(_blochsim_special_dir "${CMAKE_CURRENT_BINARY_DIR}/special") + set(_blochsim_special_table "") + if(BLOCHSIM_SPECIALIZE) + set(_list "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim/_specializations.json") + set_property(DIRECTORY APPEND PROPERTY CMAKE_CONFIGURE_DEPENDS "${_list}") + file(READ "${_list}" _json) + string(JSON _count LENGTH "${_json}" specializations) + if(_count GREATER 0) + math(EXPR _last "${_count} - 1") + foreach(INDEX RANGE ${_last}) + string(JSON KERNEL GET "${_json}" specializations ${INDEX} kernel) + string(JSON _fixed GET "${_json}" specializations ${INDEX} fixed) + string(JSON _switches LENGTH "${_fixed}") + set(FIXED "") + set(_text "") + if(_switches GREATER 0) + math(EXPR _last_switch "${_switches} - 1") + foreach(_j RANGE ${_last_switch}) + string(JSON _name MEMBER "${_fixed}" ${_j}) + string(JSON _value GET "${_fixed}" ${_name}) + string(APPEND FIXED ", BLOCHSIM_PARAM(${_name}), ${_value}") + string(APPEND _text "${_name}=${_value},") + endforeach() + endif() + set(_source "${_blochsim_special_dir}/special${INDEX}.cu") + configure_file(src/blochsim/_gpu_special.cu.in "${_source}" @ONLY) + list(APPEND _blochsim_gpu_sources "${_source}") + string(APPEND _blochsim_special_table " X(${INDEX}, ${KERNEL}, \"${_text}\") \\\n") + endforeach() + endif() + endif() + file(CONFIGURE OUTPUT "${_blochsim_special_dir}/_special_table.hpp" + CONTENT "// Written by CMake from _specializations.json: X(index, kernel, \"switch=value,...\").\n#define BLOCHSIM_FOR_EACH_SPECIAL(X) \\\n${_blochsim_special_table}\n" + @ONLY) + python_add_library(_gpu MODULE USE_SABI ${BLOCHSIM_ABI3_VERSION} WITH_SOABI ${_blochsim_gpu_sources} ) - target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim") + target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim" "${_blochsim_special_dir}") # The runtime is linked in, so a machine needs the driver and nothing else. target_link_libraries(_gpu PRIVATE CUDA::cudart_static) # Warning 221 is a double constant that rounds to zero in float, which the diff --git a/scripts/kernel_census.py b/scripts/kernel_census.py new file mode 100644 index 00000000..bb97bb1b --- /dev/null +++ b/scripts/kernel_census.py @@ -0,0 +1,130 @@ +"""Which feature switches the GPU kernels are launched with, and the list to compile. + +As a pytest plugin it records, for every launch on a card, the kernel and the +values of its feature switches, and writes them out when the session ends:: + + PYTHONPATH=scripts pytest -p kernel_census tests/ --census census.json + +Run as a script it turns one or more such records into the entries of +``src/blochsim/_specializations.json``: one per distinct combination of a +kernel's switches, whatever else differed between the launches:: + + python scripts/kernel_census.py census.json > src/blochsim/_specializations.json +""" + +from __future__ import annotations + +import collections +import json +import sys + +#: The switches a kernel is compiled for when it is specialized: every argument +#: that only turns terms on and off. Tile sizes are left out; they shape the +#: launch, not the code. +SWITCHES = ( + "single_train", + "atom_stride", + "shimmed", + "profiled", + "dynamic", + "broadened", + "pools", + "narrow", + "tabulated", + "off_axis", + "moving", + "diffusing", + "transmit", + "density", + "inverting", + "recording", + "directed", +) + +#: The kernels specialized. The pooled kernels' tiles are their pools, and the +#: three-pool tables and PERK have no switches worth a kernel of their own. +SPECIALIZED = ( + "_epg_kernel", + "_epg_jvp_kernel", + "_epg_vjp_kernel", + "_epg_vjp_jvp_kernel", + "_epg_real_kernel", + "_epg_real_jvp_kernel", + "_epg_real_vjp_kernel", + "_epg_real_vjp_jvp_kernel", +) + +#: Switches with which a kernel is left to its general build. Fixing them +#: does not shrink the second-order kernel's compile but multiplies it: one such +#: entry took as long to compile as a dozen without, and they are rare. +COSTLY = { + "_epg_vjp_jvp_kernel": ( + "profiled", + "moving", + "pools", + "dynamic", + "broadened", + "narrow", + "tabulated", + ), +} + +_seen: dict[str, collections.Counter] = collections.defaultdict(collections.Counter) + + +def pytest_addoption(parser) -> None: + parser.addoption( + "--census", default="census.json", help="where to write the census" + ) + + +def pytest_configure(config) -> None: + from blochsim import _gpu_launch + + original = _gpu_launch.Kernel.launch + + def recorded(self, grid, args, kwargs): + names, _ = _gpu_launch._signature(self.name) + values = dict( + zip( + names, + list(args) + [kwargs.get(n) for n in names[len(args) :]], + strict=True, + ) + ) + fixed = {n: int(values[n]) for n in SWITCHES if n in values} + _seen[self.name][json.dumps(fixed, sort_keys=True)] += 1 + return original(self, grid, args, kwargs) + + _gpu_launch.Kernel.launch = recorded + + +def pytest_unconfigure(config) -> None: + census = {name: dict(counts) for name, counts in sorted(_seen.items())} + with open(config.getoption("--census"), "w") as stream: + json.dump(census, stream, indent=1, sort_keys=True) + + +def entries(*censuses: dict) -> list[dict]: + """One entry per distinct combination of a specialized kernel's switches. + + A combination turning on a switch :data:`COSTLY` names for its kernel is + left out. + """ + combinations = { + (name, key) + for census in censuses + for name, counts in census.items() + if name in SPECIALIZED + for key in counts + if not any(json.loads(key).get(switch) for switch in COSTLY.get(name, ())) + } + return [ + {"kernel": name, "fixed": json.loads(key)} for name, key in sorted(combinations) + ] + + +if __name__ == "__main__": + loaded = [json.load(open(path)) for path in sys.argv[1:]] + json.dump({"specializations": entries(*loaded)}, sys.stdout, indent=1) + sys.stdout.write("\n") diff --git a/src/blochsim/_epg_kernels.hpp b/src/blochsim/_epg_kernels.hpp index 4ee56798..395558de 100644 --- a/src/blochsim/_epg_kernels.hpp +++ b/src/blochsim/_epg_kernels.hpp @@ -9154,21 +9154,7 @@ BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, bsk::atomic_add(((grad_tissue_tangent + ((7 + past_transmit) * atom_count)) + atom), grad_damping_tangent, active_atom); } -BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count_wide, std::int64_t train_count_wide, std::int64_t event_count_wide, std::int64_t output_count_wide, std::int64_t state_count_wide, std::int64_t single_train_wide, std::int64_t atom_stride_wide, std::int64_t shimmed_wide, std::int64_t diffusing_wide, std::int64_t transmit_wide, std::int64_t density_wide, std::int64_t inverting_wide, std::int64_t block_states_wide, std::int64_t problems_wide) { - const std::int32_t atom_count = static_cast(atom_count_wide); - const std::int32_t train_count = static_cast(train_count_wide); - const std::int32_t event_count = static_cast(event_count_wide); - const std::int32_t output_count = static_cast(output_count_wide); - const std::int32_t state_count = static_cast(state_count_wide); - const std::int32_t single_train = static_cast(single_train_wide); - const std::int32_t atom_stride = static_cast(atom_stride_wide); - const std::int32_t shimmed = static_cast(shimmed_wide); - const std::int32_t diffusing = static_cast(diffusing_wide); - const std::int32_t transmit = static_cast(transmit_wide); - const std::int32_t density = static_cast(density_wide); - const std::int32_t inverting = static_cast(inverting_wide); - const std::int32_t block_states = static_cast(block_states_wide); - const std::int32_t problems = static_cast(problems_wide); +BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { bsk::V alpha{}; bsk::V atom_b1{}; bsk::V atom_damping{}; @@ -9185,7 +9171,7 @@ BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* bsk::V plus{}; bsk::V pulse_b1{}; bsk::V relaxes{}; - auto problem = ((static_cast(bsk::program_id(0)) * problems) + bsk::arange_y()); + auto problem = ((bsk::program_id(0) * problems) + bsk::arange_y()); auto state = bsk::arange_x(); auto active_atom = (problem < (train_count * atom_count)); auto state_mask = bsk::band((state < state_count), active_atom); @@ -9292,27 +9278,17 @@ BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* if (bsk::truth((bsk::truth(is_rf) && bsk::truth(is_inversion)))) { longitudinal = ((-atom_inversion) * longitudinal); } else if (bsk::truth(is_rf)) { - decltype(bsk::cos(alpha)) cosine; - decltype(bsk::sin(alpha)) sine; - if (bsk::truth(single_train) && !bsk::truth(transmit)) { - // One train and no transmit field: every problem's pulse is - // the same angle, so its cosine and sine are taken once. - const float angle = bsk::ld(flip + event); - cosine = cosf(angle); - sine = sinf(angle); - } else { - alpha = _event_value(flip, event_base, event, active_atom, single_train); - pulse_b1 = atom_b1; - // One shim is the whole sequence's transmit field, loaded once - // above; several give each pulse the row of the shim it drives. - if (bsk::truth((bsk::truth(shimmed) && bsk::truth(transmit)))) { - auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); - pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); - } - alpha = (alpha * pulse_b1); - cosine = bsk::cos(alpha); - sine = bsk::sin(alpha); + alpha = _event_value(flip, event_base, event, active_atom, single_train); + pulse_b1 = atom_b1; + // One shim is the whole sequence's transmit field, loaded once + // above; several give each pulse the row of the shim it drives. + if (bsk::truth((bsk::truth(shimmed) && bsk::truth(transmit)))) { + auto shim_row = (bsk::cast(bsk::ld((shim_index + event))) * atom_count); + pulse_b1 = bsk::ld(((b1 + shim_row) + atom), active_atom, 1.0f); } + alpha = (alpha * pulse_b1); + auto cosine = bsk::cos(alpha); + auto sine = bsk::sin(alpha); auto cosine_half_sq = (0.5f * (1.0f + cosine)); auto sine_half_sq = (0.5f * (1.0f - cosine)); auto half_sine = (0.5f * sine); diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index acfb639a..b02e2275 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -10,15 +10,26 @@ #define BLOCHSIM_TABLE_ONLY 1 #include "_launch.hpp" +#include "_special.hpp" +#include "_special_table.hpp" #include +#include +#include +#include + #define BLOCHSIM_DEVICE_ENTRY(name) \ __global__ void kernel##name##_256(bsk::Arguments arguments, int z); \ __global__ void kernel##name##_1024(bsk::Arguments arguments, int z); BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_ENTRY) #undef BLOCHSIM_DEVICE_ENTRY +#define BLOCHSIM_SPECIAL_DECLARATION(index, kernel, fixed) \ + __global__ void special##index(bsk::Arguments arguments, int z); +BLOCHSIM_FOR_EACH_SPECIAL(BLOCHSIM_SPECIAL_DECLARATION) +#undef BLOCHSIM_SPECIAL_DECLARATION + namespace { // Each kernel bounded to 256 threads, then to 1024. @@ -28,6 +39,72 @@ namespace { const void* const KERNEL_FUNCTIONS[][2] = {BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_POINTER)}; #undef BLOCHSIM_DEVICE_POINTER +// The kernels compiled for one combination of their switches, as CMake listed +// them, and parsed once into the arguments each fixes. +struct SpecialSource { + const char* kernel; + const char* fixed; + const void* function; +}; + +#define BLOCHSIM_SPECIAL_SOURCE(index, kernel, fixed) \ + {#kernel, fixed, reinterpret_cast(&special##index)}, +const SpecialSource SPECIAL_SOURCES[] = { + BLOCHSIM_FOR_EACH_SPECIAL(BLOCHSIM_SPECIAL_SOURCE){nullptr, nullptr, nullptr}}; +#undef BLOCHSIM_SPECIAL_SOURCE + +struct Special { + std::vector> fixed; + const void* function; +}; + +const std::vector>& specials() { + static const std::vector> table = [] { + std::vector> out(sizeof(bsk::KERNELS) / sizeof(bsk::KERNELS[0])); + for (const SpecialSource& source : SPECIAL_SOURCES) { + if (source.kernel == nullptr) { + break; + } + const int kernel = blochsim_launch::find_kernel(source.kernel); + Special special{{}, source.function}; + const std::string text = source.fixed; + std::size_t start = 0; + while (start < text.size()) { + const std::size_t end = text.find(',', start); + const std::string pair = text.substr(start, end - start); + const std::size_t equals = pair.find('='); + special.fixed.emplace_back(bsk::param_index(kernel, pair.substr(0, equals).c_str()), + std::stoll(pair.substr(equals + 1))); + start = end + 1; + } + out[kernel].push_back(std::move(special)); + } + return out; + }(); + return table; +} + +// Whether a launch may run a specialized kernel, and how many have. +bool specializing = true; +unsigned long long specialized_launches = 0; + +// The specialized kernel whose fixed switches this launch matches, if any. +const void* matching(const blochsim_launch::Launch& request) { + for (const Special& special : specials()[request.kernel]) { + bool match = true; + for (const auto& [index, value] : special.fixed) { + if (request.arguments.a[index].i != value) { + match = false; + break; + } + } + if (match) { + return special.function; + } + } + return nullptr; +} + PyObject* cuda_error(cudaError_t status, const char* what) { PyErr_Format(PyExc_RuntimeError, "%s: %s", what, cudaGetErrorString(status)); return nullptr; @@ -76,8 +153,17 @@ PyObject* launch(PyObject*, PyObject* args) { sizeof(unsigned long long) * (threads.x * threads.y + bsk::MAX_Z * bsk::MAX_Z); void* parameters[] = {&request.arguments, &request.z}; const int bounded = threads.x * threads.y > 256 ? 1 : 0; - status = cudaLaunchKernel(KERNEL_FUNCTIONS[request.kernel][bounded], blocks, threads, parameters, - shared, reinterpret_cast(stream)); + // A specialized kernel is compiled for 256 threads; a wider block runs the + // kernel compiled for every combination. + const void* function = KERNEL_FUNCTIONS[request.kernel][bounded]; + if (specializing && !bounded) { + if (const void* special = matching(request)) { + function = special; + ++specialized_launches; + } + } + status = cudaLaunchKernel(function, blocks, threads, parameters, shared, + reinterpret_cast(stream)); if (previous != device) { cudaSetDevice(previous); } @@ -87,11 +173,51 @@ PyObject* launch(PyObject*, PyObject* args) { Py_RETURN_NONE; } +PyObject* specializations(PyObject*, PyObject*) { + PyObject* out = PyList_New(0); + if (out == nullptr) { + return nullptr; + } + for (const SpecialSource& source : SPECIAL_SOURCES) { + if (source.kernel == nullptr) { + break; + } + PyObject* entry = Py_BuildValue("(ss)", source.kernel, source.fixed); + if (entry == nullptr || PyList_Append(out, entry) < 0) { + Py_XDECREF(entry); + Py_DECREF(out); + return nullptr; + } + Py_DECREF(entry); + } + return out; +} + +PyObject* use_specializations(PyObject*, PyObject* args) { + int on = 1; + if (!PyArg_ParseTuple(args, "p", &on)) { + return nullptr; + } + const bool previous = specializing; + specializing = on != 0; + return PyBool_FromLong(previous); +} + +PyObject* specialized_launch_count(PyObject*, PyObject*) { + return PyLong_FromUnsignedLongLong(specialized_launches); +} + PyMethodDef METHODS[] = { {"kernels", blochsim_launch::kernel_table, METH_NOARGS, "Each kernel's parameter names and kinds."}, {"launch", launch, METH_VARARGS, "Queue a kernel over a grid of programs on a device's stream."}, + {"specializations", specializations, METH_NOARGS, + "Each specialized kernel, and the switches it was compiled for."}, + {"use_specializations", use_specializations, METH_VARARGS, + "Whether launches may run specialized kernels; returns the previous setting."}, + {"specialized_launches", specialized_launch_count, METH_NOARGS, + "How many launches have run a specialized kernel."}, {nullptr, nullptr, 0, nullptr}, }; diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index 03073dc9..874a5569 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -11,6 +11,8 @@ __all__: list[str] = [] +from collections.abc import Iterator +from contextlib import contextmanager from functools import cache from typing import Any @@ -59,6 +61,45 @@ def available() -> bool: return True +def specializations() -> list[tuple[str, dict[str, int]]]: + """Each kernel compiled for one combination of its switches, and the switches. + + Empty where this installation carries no kernels for a card. + """ + if not available(): + return [] + return [ + ( + kernel, + { + name: int(value) + for name, value in ( + pair.split("=") for pair in fixed.split(",") if pair + ) + }, + ) + for kernel, fixed in _module("cuda").specializations() + ] + + +def specialized_launches() -> int: + """How many launches on a card have run a specialized kernel.""" + return _module("cuda").specialized_launches() if available() else 0 + + +@contextmanager +def generic_kernels() -> Iterator[None]: + """Run every launch inside on the kernels compiled for all combinations.""" + if not available(): + yield + return + previous = _module("cuda").use_specializations(False) + try: + yield + finally: + _module("cuda").use_specializations(previous) + + class Kernel: """A compiled kernel, launched as ``kernel[grid](*arguments)``.""" diff --git a/src/blochsim/_gpu_special.cu.in b/src/blochsim/_gpu_special.cu.in new file mode 100644 index 00000000..f74a7d8d --- /dev/null +++ b/src/blochsim/_gpu_special.cu.in @@ -0,0 +1,27 @@ +// One kernel compiled for one combination of its feature switches. CMake +// writes one of these per entry of _specializations.json: @INDEX@ is the +// entry, @KERNEL@ the kernel, and the switches it fixes follow ``Call``. +#define BLOCHSIM_SIMT 1 + +#include "_lanes.hpp" +#define BLOCHSIM_Y_LANES BLOCHSIM_LANES_@KERNEL@ + +#include "_special.hpp" + +namespace { + +constexpr int KERNEL = bsk::kernel_index("@KERNEL@"); +#define BLOCHSIM_PARAM(name) bsk::param_index(KERNEL, #name) + +struct Call { + BSK_HD static void run(const bsk::Arg* a) { bsk::call@KERNEL@(a); } +}; + +using Special = bsk::Fixing; + +} // namespace + +__global__ void __launch_bounds__(256) special@INDEX@(bsk::Arguments arguments, int z) { + bsk::enter(z); + Special::run(arguments.a); +} diff --git a/src/blochsim/_lanes.hpp b/src/blochsim/_lanes.hpp index 6924c5c8..437a0832 100644 --- a/src/blochsim/_lanes.hpp +++ b/src/blochsim/_lanes.hpp @@ -13,9 +13,9 @@ #define BLOCHSIM_LANES__epg_vjp_jvp_kernel 1 #define BLOCHSIM_LANES__epg_real_vjp_jvp_kernel 1 #define BLOCHSIM_LANES__epg_real_vjp_kernel 1 -#define BLOCHSIM_LANES__epg_real_kernel 8 +#define BLOCHSIM_LANES__epg_real_kernel 4 #define BLOCHSIM_LANES__epg_real_jvp_kernel 1 -#define BLOCHSIM_LANES__epg_kernel 1 +#define BLOCHSIM_LANES__epg_kernel 2 #define BLOCHSIM_LANES__epg_jvp_kernel 1 #define BLOCHSIM_LANES__pooled_kernel 1 #define BLOCHSIM_LANES__pooled_adjoint_kernel 1 diff --git a/src/blochsim/_special.hpp b/src/blochsim/_special.hpp new file mode 100644 index 00000000..854f1437 --- /dev/null +++ b/src/blochsim/_special.hpp @@ -0,0 +1,98 @@ +// A kernel with some of its arguments fixed at compile time. +// +// The kernels take every feature switch as an argument and branch on it, so +// one kernel serves every combination and is compiled for the worst of them: +// the registers and code of every term it might evaluate. Called with those +// switches as constants, the same kernel is compiled for one combination +// alone -- every function it calls is inlined, so a constant switch folds and +// the terms it turns off are never generated. _specializations.json lists the +// combinations compiled this way; the launcher runs one where a launch's +// switches match it exactly and the kernel as compiled for all of them where +// none does. +#pragma once + +#include "_kernels.hpp" + +namespace bsk { + +constexpr bool same_name(const char* a, const char* b) { + while (*a != '\0' && *a == *b) { + ++a; + ++b; + } + return *a == *b; +} + +constexpr int kernel_index(const char* name) { + int index = 0; + for (const auto& info : KERNELS) { + if (same_name(info.name, name)) { + return index; + } + ++index; + } + return -1; +} + +// Where ``param`` sits in the comma-separated parameter list of ``kernel``. +constexpr int param_index(int kernel, const char* param) { + const char* p = KERNELS[kernel].params; + int index = 0; + while (*p != '\0') { + const char* q = param; + const char* s = p; + while (*q != '\0' && *s == *q) { + ++q; + ++s; + } + if (*q == '\0' && (*s == ',' || *s == '\0')) { + return index; + } + while (*p != '\0' && *p != ',') { + ++p; + } + if (*p == ',') { + ++p; + } + ++index; + } + return -1; +} + +// ``Call::run`` with the arguments at ``Fixed``'s even entries replaced by the +// integers that follow them; the copy and the replaced reads fold away. The +// indices are worked out where a constant expression is host code -- an alias +// at namespace scope -- and ``Call`` is a type, so no device function's +// address is taken there. +template +constexpr bool every_index_named() { + constexpr long long pairs[sizeof...(Fixed) + 1] = {Fixed..., 0}; + for (int k = 0; k + 1 < static_cast(sizeof...(Fixed)); k += 2) { + if (pairs[k] < 0) { + return false; + } + } + return true; +} + +template +struct Fixing { + static_assert(sizeof...(Fixed) % 2 == 0, "fixed arguments come as index, value pairs"); + static_assert(every_index_named(), "a fixed argument names no parameter"); + + BSK_HD static void run(const Arg* in) { + constexpr long long pairs[sizeof...(Fixed) + 1] = {Fixed..., 0}; + Arg a[MAX_ARGUMENTS]; +#pragma unroll + for (int i = 0; i < MAX_ARGUMENTS; ++i) { + a[i] = in[i]; + } +#pragma unroll + for (int k = 0; k + 1 < static_cast(sizeof...(Fixed)); k += 2) { + a[pairs[k]].i = pairs[k + 1]; + } + Call::run(a); + } +}; + +} // namespace bsk diff --git a/src/blochsim/_specializations.json b/src/blochsim/_specializations.json new file mode 100644 index 00000000..bf14439d --- /dev/null +++ b/src/blochsim/_specializations.json @@ -0,0 +1,2406 @@ +{ + "specializations": [ + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 1, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 1, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 1, + "off_axis": 0, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 1, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 1, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 1, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 1, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 1, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 1, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 1, + "off_axis": 0, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 1, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_jvp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_jvp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 0, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 0 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 0, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 0, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 1, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 0 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 1, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 1, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 0 + } + }, + { + "kernel": "_epg_real_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_vjp_kernel", + "fixed": { + "atom_stride": 0, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 0 + } + }, + { + "kernel": "_epg_real_vjp_kernel", + "fixed": { + "atom_stride": 1, + "density": 0, + "diffusing": 0, + "inverting": 0, + "shimmed": 0, + "single_train": 1, + "transmit": 0 + } + }, + { + "kernel": "_epg_real_vjp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_vjp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 0, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_real_vjp_kernel", + "fixed": { + "atom_stride": 1, + "density": 1, + "diffusing": 1, + "inverting": 1, + "shimmed": 1, + "single_train": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "directed": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 0, + "diffusing": 0, + "directed": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_jvp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "directed": 0, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 1, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 0, + "diffusing": 1, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 0, + "off_axis": 0, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 0, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 1, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 1, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 1, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 2, + "profiled": 1, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 0, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 1, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 0, + "density": 1, + "diffusing": 1, + "dynamic": 1, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 0, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 1, + "off_axis": 0, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 0, + "diffusing": 0, + "dynamic": 0, + "inverting": 0, + "moving": 0, + "narrow": 1, + "off_axis": 0, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 0 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 1, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 1, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 1, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 1, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 1, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 0, + "narrow": 1, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 0, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 0, + "transmit": 1 + } + }, + { + "kernel": "_epg_vjp_kernel", + "fixed": { + "atom_stride": 1, + "broadened": 1, + "density": 1, + "diffusing": 1, + "dynamic": 0, + "inverting": 1, + "moving": 1, + "narrow": 0, + "off_axis": 1, + "pools": 3, + "profiled": 0, + "recording": 1, + "shimmed": 0, + "single_train": 1, + "tabulated": 1, + "transmit": 1 + } + } + ] +} diff --git a/tests/sequence/test_specialized_kernels.py b/tests/sequence/test_specialized_kernels.py new file mode 100644 index 00000000..f283cbff --- /dev/null +++ b/tests/sequence/test_specialized_kernels.py @@ -0,0 +1,105 @@ +"""Whether a kernel compiled for its switches computes what the general one does. + +A specialized kernel that quietly did not run agrees perfectly, so every case +also asserts that one did. +""" + +from __future__ import annotations + +from dataclasses import replace + +import numpy as np +import pytest +import torch + +from blochsim import _gpu_launch, rf_definition +from blochsim.sequence import EpgEngine, exact_slice_profile, fse_description +from blochsim.sequence._simulation import TissueProperties + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or not _gpu_launch.specializations(), + reason="needs a card and the specialized kernels", +) + + +def _tissue(atoms: int = 300) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(0) + return { + "t1_ms": (200 + 2800 * torch.rand(atoms, generator=generator)).cuda(), + "t2_ms": (10 + 290 * torch.rand(atoms, generator=generator)).cuda(), + } + + +def _sinc_pulse(description): + grid = np.linspace(-2.0, 2.0, 128) + envelope = np.sinc(grid) * (0.54 + 0.46 * np.cos(np.pi * grid / 2.0)) + definition = rf_definition( + envelope.astype(np.complex128), + dwell_s=1e-5, + bandwidth_hz=2000.0, + definition_id=0, + ) + return replace(description, rf_definitions={definition.id: definition}) + + +def _forward(phase: float, profiled: bool = False) -> torch.Tensor: + description = fse_description( + torch.deg2rad(torch.full((32,), 150.0)), + echo_spacing_s=5e-3, + phases_rad=phase, + excitation_phase_rad=torch.pi / 2, + ) + across = None + if profiled: + description = _sinc_pulse(description) + across = exact_slice_profile(9) + return ( + EpgEngine() + .simulate( + description, TissueProperties(**_tissue()), across_slice=across, nstates=32 + ) + .signal + ) + + +def _gradient(phase: float) -> torch.Tensor: + description = fse_description( + torch.deg2rad(torch.full((32,), 150.0)), echo_spacing_s=5e-3, phases_rad=phase + ) + tissue = _tissue() + tissue["t2_ms"] = tissue["t2_ms"].clone().requires_grad_() + signal = ( + EpgEngine().simulate(description, TissueProperties(**tissue), nstates=32).signal + ) + signal.abs().sum().backward() + return tissue["t2_ms"].grad + + +CASES = { + "real forward": lambda: _forward(torch.pi / 2), + "complex forward": lambda: _forward(0.0), + "slice profile": lambda: _forward(torch.pi / 2, profiled=True), + "real gradient": lambda: _gradient(torch.pi / 2), + "complex gradient": lambda: _gradient(0.0), +} + + +@pytest.mark.parametrize("case", CASES) +def test_a_specialized_kernel_computes_what_the_general_one_does(case) -> None: + with _gpu_launch.generic_kernels(): + general = CASES[case]() + before = _gpu_launch.specialized_launches() + special = CASES[case]() + + assert _gpu_launch.specialized_launches() > before + error = (special - general).abs().max() + scale = general.abs().max() + assert float(error / scale) < 1e-5, f"{float(error):.3e} against {float(scale):.3e}" + + +def test_the_general_kernels_run_where_asked() -> None: + before = _gpu_launch.specialized_launches() + with _gpu_launch.generic_kernels(): + _forward(torch.pi / 2) + + assert _gpu_launch.specialized_launches() == before From b0ab84dbb85cced65581fc69127ae07546758b00 Mon Sep 17 00:00:00 2001 From: mcencini Date: Wed, 7 Oct 2026 20:50:15 +0200 Subject: [PATCH 07/16] Cut the instructions the GPU kernels execute that Triton's did not Profiled against Triton on an RTX 4060, the kernels issued as well as Triton's and simply executed more instructions, almost all of them integer bookkeeping. In the tile runtime and the kernels' types: - A costly function of a thread's rows is taken once where the rows share the argument, as a pulse's flip does; Triton sinks a broadcast below the function and the tile runtime could not. - The EPG kernels index in 32 bits, as Triton passed every integer argument that fit; the launcher refuses one that does not. - A row wider than a warp is gathered and reduced in a function of its own: inlined, its barriers and shared accesses were predicated and issued on every shift. - A column of one row is not reduced, and one warp's rows are reduced by shuffles. - Specialized kernels are compiled for rows a warp wide and launched only on them; a program is two warps. - The values the adjoint kernels narrow to float before the per-order loop are declared float, as Triton rebinds them; declared double, the loop ran in fp64. The series-against-roots test allows 5e-5: one fused multiply-add moves the roots branch by 1e-5, and Triton's differs from the CUDA build's by 2e-5. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 9 ++- src/blochsim/_epg_kernels.hpp | 126 +++++++++++++++--------------- src/blochsim/_gpu.cu | 6 +- src/blochsim/_gpu_special.cu.in | 2 + src/blochsim/_launch.hpp | 8 ++ src/blochsim/_tile.hpp | 97 +++++++++++++++++++++-- src/blochsim/sequence/_epg_gpu.py | 7 +- tests/sequence/test_both_pools.py | 5 +- 8 files changed, 182 insertions(+), 78 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index f1edb35f..92d5b765 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -88,7 +88,14 @@ list; `_gpu_launch.generic_kernels()` runs a block on the general kernels, which is how `tests/sequence/test_specialized_kernels.py` holds the two to each other. Each entry is another compile of its kernel, so the list is most of a CUDA build's time; `--config-settings=cmake.define.BLOCHSIM_SPECIALIZE=OFF` builds -the general kernels alone. +the general kernels alone. A specialized kernel is compiled for rows of at most +32 state orders, so its shifts are shuffles with no test; a wider launch runs +the general one. + +**The EPG kernels index in 32 bits** (`bsk::index_t`), as Triton did for every +integer argument that fit. An offset that can pass 2^31 is cast to 64 bits +where it is formed, as the Triton source cast it, and the launcher refuses an +integer argument that does not fit rather than truncating it. **`--cov` is on by default** through `addopts`, so a bare `pytest` writes `coverage.xml`. It is ignored, not tracked. diff --git a/src/blochsim/_epg_kernels.hpp b/src/blochsim/_epg_kernels.hpp index 395558de..e3f434d9 100644 --- a/src/blochsim/_epg_kernels.hpp +++ b/src/blochsim/_epg_kernels.hpp @@ -983,7 +983,7 @@ BSK_HD auto _three_pool_pieces_jvp_in_precision(const T0& r1_free, const T1& d_r d_sum_square = d_square; factorial = Work(1.0); #pragma unroll - for (std::int64_t order = 1; order < 16; order += 1) { + for (bsk::index_t order = 1; order < 16; order += 1) { auto next_flat = (square * determinant); auto d_next_flat = ((d_square * determinant) + (square * d_determinant)); auto next_linear = (flat - (square * minors)); @@ -1359,7 +1359,7 @@ BSK_HD auto _three_pool_step_adjoint_jvp_in_precision(const T0& r1_free, const T d_slope_v_square = (Work(0.0) * a00); factorial = Work(1.0); #pragma unroll - for (std::int64_t order = 1; order < 16; order += 1) { + for (bsk::index_t order = 1; order < 16; order += 1) { auto next_flat = (square * determinant); auto d_next_flat = ((d_square * determinant) + (square * d_determinant)); auto next_linear = (flat - (square * minors)); @@ -2384,7 +2384,7 @@ BSK_HD auto _washout_jvp(const T0& rate, const T1& rate_tangent, const T2& dt, c return bsk::make_tup(bsk::where(live, (1.0f - fraction), 0.0f), bsk::where(live, (-((rate_tangent * dt) + (rate * dt_tangent))), 0.0f)); } -BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, float* grad_pair_value, float* grad_pair_tangent, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_b1_phase, float* dot_b0, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_velocity, float* dot_bound_fraction, float* dot_exchange_rate, float* dot_t1_bound, float* dot_pool_b_fraction, float* dot_pool_b_exchange, float* dot_t1_pool_b, float* dot_t2_pool_b, float* dot_pool_b_shift, float* dot_duration, float* dot_flip, float* dot_phase, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, std::int64_t row_count, float* grad_output_real, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_phase_value, float* grad_phase_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_vr, float* trajectory_vi, float* trajectory_tr, float* trajectory_ti, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t directed, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t recording, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, float* grad_pair_value, float* grad_pair_tangent, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_b1_phase, float* dot_b0, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_velocity, float* dot_bound_fraction, float* dot_exchange_rate, float* dot_t1_bound, float* dot_pool_b_fraction, float* dot_pool_b_exchange, float* dot_t1_pool_b, float* dot_t2_pool_b, float* dot_pool_b_shift, float* dot_duration, float* dot_flip, float* dot_phase, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, bsk::index_t row_count, float* grad_output_real, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_phase_value, float* grad_phase_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_vr, float* trajectory_vi, float* trajectory_tr, float* trajectory_ti, bsk::index_t problem_base, bsk::index_t problem_end, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shim_rows, bsk::index_t shimmed, bsk::index_t locations, bsk::index_t profiled, bsk::index_t profile_bins, bsk::index_t dynamic, bsk::index_t directed, bsk::index_t off_axis, bsk::index_t moving, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t broadened, bsk::index_t lineshape_bins, bsk::index_t pools, bsk::index_t narrow, bsk::index_t tabulated, bsk::index_t recording, bsk::index_t block_states, bsk::index_t problems) { bsk::tup, bsk::V, bsk::V, bsk::V> a11{}; bsk::tup, bsk::V, bsk::V, bsk::V> a12{}; bsk::tup, bsk::V, bsk::V, bsk::V> a21{}; @@ -2497,9 +2497,9 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa bsk::V d_exchange{}; bsk::V d_flow{}; bsk::V d_free{}; - bsk::V d_grow_free{}; - bsk::V d_grow_pool_b{}; - bsk::V d_grow_semisolid{}; + bsk::V d_grow_free{}; + bsk::V d_grow_pool_b{}; + bsk::V d_grow_semisolid{}; bsk::V d_inv{}; bsk::V d_m0{}; bsk::V d_semisolid_exchange{}; @@ -2518,15 +2518,15 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa bsk::V d_t32{}; bsk::V d_t33{}; bsk::V d_turn{}; - bsk::V d_w11{}; - bsk::V d_w12{}; - bsk::V d_w13{}; - bsk::V d_w21{}; - bsk::V d_w22{}; - bsk::V d_w23{}; - bsk::V d_w31{}; - bsk::V d_w32{}; - bsk::V d_w33{}; + bsk::V d_w11{}; + bsk::V d_w12{}; + bsk::V d_w13{}; + bsk::V d_w21{}; + bsk::V d_w22{}; + bsk::V d_w23{}; + bsk::V d_w31{}; + bsk::V d_w32{}; + bsk::V d_w33{}; bsk::V d_washout{}; bsk::V damp_pair_t{}; bsk::V damp_pair_v{}; @@ -2621,9 +2621,9 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa bsk::V grad_e1_v{}; bsk::V grad_e2_t{}; bsk::V grad_e2_v{}; - bsk::V grow_free{}; - bsk::V grow_pool_b{}; - bsk::V grow_semisolid{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; bsk::V held{}; bsk::tup, bsk::V, bsk::V, bsk::V> held_bar{}; bsk::V held_semisolid{}; @@ -2963,16 +2963,16 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa bsk::V utr{}; bsk::tup, bsk::V, bsk::V, bsk::V> w0{}; bsk::tup, bsk::V, bsk::V, bsk::V> w1{}; - bsk::V w11{}; - bsk::V w12{}; - bsk::V w13{}; + bsk::V w11{}; + bsk::V w12{}; + bsk::V w13{}; bsk::tup, bsk::V, bsk::V, bsk::V> w2{}; - bsk::V w21{}; - bsk::V w22{}; - bsk::V w23{}; - bsk::V w31{}; - bsk::V w32{}; - bsk::V w33{}; + bsk::V w21{}; + bsk::V w22{}; + bsk::V w23{}; + bsk::V w31{}; + bsk::V w32{}; + bsk::V w33{}; bsk::V wash_t{}; bsk::V wash_v{}; bsk::V wbti{}; @@ -3217,7 +3217,7 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa // and the two are launched separately: each compiles the sweep it is // asked for and no more. if (bsk::truth(recording)) { - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { slot = (trajectory + (event * record_stride)); bsk::st((trajectory_vr + slot), pvr, state_mask); bsk::st((trajectory_vi + slot), pvi, state_mask); @@ -3888,7 +3888,7 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa g_b0t = zero; g_invv = zero; g_invt = zero; - for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + for (bsk::index_t reverse = 0; reverse < event_count; reverse += 1) { event = ((event_count - 1) - reverse); slot = (trajectory + (event * record_stride)); auto xpvr = bsk::ld((trajectory_vr + slot), state_mask, 0.0f); @@ -5585,7 +5585,7 @@ BSK_HD void _epg_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, floa // by events whose interval directions differ -- so the second pass // takes that dependence alone, driven by the cotangents the walk // back weighted by each event's direction and read at a unit one. - for (std::int64_t row = 0; row < row_count; row += 1) { + for (bsk::index_t row = 0; row < row_count; row += 1) { held = (pool_bars + (((local * row_count) + row) * 36)); auto row_dt = (bsk::ld((pool_durations + row)) + zero); auto nil = (0.0f * row_dt); @@ -6208,7 +6208,7 @@ BSK_HD auto _three_pool_step_in_precision(const T0& r1_free, const T1& r1_pool_b sum_square = square; factorial = Work(1.0); #pragma unroll - for (std::int64_t order = 1; order < terms; order += 1) { + for (bsk::index_t order = 1; order < terms; order += 1) { auto next_flat = (square * determinant); auto next_linear = (flat - (square * minors)); auto next_square = linear; @@ -6436,7 +6436,7 @@ BSK_HD auto _washout(const T0& rate, const T1& dt) { return (1.0f - bsk::minimum((rate * dt), 1.0f)); } -BSK_HD void _epg_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* bound_exchange, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, float* phase_cos, float* phase_sin, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* bound_exchange, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, float* phase_cos, float* phase_sin, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shim_rows, bsk::index_t shimmed, bsk::index_t locations, bsk::index_t profiled, bsk::index_t profile_bins, bsk::index_t dynamic, bsk::index_t broadened, bsk::index_t lineshape_bins, bsk::index_t pools, bsk::index_t narrow, bsk::index_t tabulated, bsk::index_t off_axis, bsk::index_t moving, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V atom_b0{}; bsk::V atom_b1{}; bsk::V atom_b1_phase{}; @@ -6623,7 +6623,7 @@ BSK_HD void _epg_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_ph } auto order = bsk::cast(state); auto event_base = (train * event_count); - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { auto dt = _event_value(duration, event_base, event, active_atom, single_train); wout = 1.0f; if (bsk::truth(moving)) { @@ -6979,7 +6979,7 @@ BSK_HD void _epg_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_ph // event that reads it supplies its own attenuation. Laid out // ``(rows, 9, voxels)`` -- entry-major over the voxel axis -- so the nine // loads an event makes are each coalesced. -BSK_HD void _three_pool_table_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* durations, std::int32_t* rows, float* table, std::int64_t voxel_count, std::int64_t BLOCK, std::int64_t narrow) { +BSK_HD void _three_pool_table_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* durations, std::int32_t* rows, float* table, bsk::index_t voxel_count, bsk::index_t BLOCK, bsk::index_t narrow) { auto row = bsk::ld((rows + bsk::program_id(0))); auto atom = ((bsk::program_id(1) * BLOCK) + bsk::arange_x()); auto live = (atom < voxel_count); @@ -7015,7 +7015,7 @@ BSK_HD void _three_pool_table_kernel(float* t1, float* t1_pool_b, float* t1_boun // the direction is ``A1 C d_dt``, which the reading event adds because // ``d_dt`` is its own and the row's is not. Laid out ``(rows, 18, voxels)``, // the tangent following the value. -BSK_HD void _three_pool_table_jvp_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* d_t1, float* d_t1_pool_b, float* d_t1_bound, float* d_pool_b_exchange, float* d_bound_exchange, float* d_pool_b_fraction, float* d_bound_fraction, float* durations, std::int32_t* rows, float* table, std::int64_t voxel_count, std::int64_t BLOCK, std::int64_t narrow) { +BSK_HD void _three_pool_table_jvp_kernel(float* t1, float* t1_pool_b, float* t1_bound, float* pool_b_exchange, float* bound_exchange, float* pool_b_fraction, float* bound_fraction, float* d_t1, float* d_t1_pool_b, float* d_t1_bound, float* d_pool_b_exchange, float* d_bound_exchange, float* d_pool_b_fraction, float* d_bound_fraction, float* durations, std::int32_t* rows, float* table, bsk::index_t voxel_count, bsk::index_t BLOCK, bsk::index_t narrow) { auto row = bsk::ld((rows + bsk::program_id(0))); auto atom = ((bsk::program_id(1) * BLOCK) + bsk::arange_x()); auto live = (atom < voxel_count); @@ -7085,7 +7085,7 @@ BSK_HD auto _shift_real_adjoint(const T0& plus_bar, const T1& minus_bar, const T return bsk::make_tup(shifted_plus, shifted_minus); } -BSK_HD void _epg_real_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_duration, float* trajectory_value, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_real_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_duration, float* trajectory_value, bsk::index_t problem_base, bsk::index_t problem_end, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shim_rows, bsk::index_t shimmed, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V adjoint_mv{}; bsk::V adjoint_pv{}; bsk::V alpha_bar_terms_value{}; @@ -7197,7 +7197,7 @@ BSK_HD void _epg_real_vjp_kernel(float* t1, float* t2, float* m0, float* b1, flo auto rate1_value = bsk::truediv(1000.0f, atom_t1); auto rate2_value = bsk::truediv(1000.0f, atom_t2); auto event_base = (train * event_count); - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { slot = (trajectory + (event * record_stride)); bsk::st((trajectory_value + slot), plus_value, state_mask); bsk::st(((trajectory_value + slot) + minus_plane), minus_value, state_mask); @@ -7280,7 +7280,7 @@ BSK_HD void _epg_real_vjp_kernel(float* t1, float* t2, float* m0, float* b1, flo grad_b1_value = zero; grad_inversion_value = zero; grad_damping_value = zero; - for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + for (bsk::index_t reverse = 0; reverse < event_count; reverse += 1) { event = ((event_count - 1) - reverse); slot = (trajectory + (event * record_stride)); auto entry_pv = bsk::ld((trajectory_value + slot), state_mask, 0.0f); @@ -7536,7 +7536,7 @@ BSK_HD auto _rotate_flip_phase_jvp(const T0& cosine, const T1& dcosine, const T2 return bsk::make_tup(rotated_pr, rotated_pi, rotated_mr, rotated_mi, rotated_zr, rotated_zi, rotated_dpr, rotated_dpi, rotated_dmr, rotated_dmi, rotated_dzr, rotated_dzi); } -BSK_HD void _epg_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_b1_phase, float* tangent_b0, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_velocity, float* tangent_bound_fraction, float* tangent_exchange_rate, float* tangent_t1_bound, float* tangent_pool_b_fraction, float* tangent_pool_b_exchange, float* tangent_t1_pool_b, float* tangent_t2_pool_b, float* tangent_pool_b_shift, float* tangent_duration, float* tangent_flip, float* tangent_phase, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_b1_phase, float* tangent_b0, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_velocity, float* tangent_bound_fraction, float* tangent_exchange_rate, float* tangent_t1_bound, float* tangent_pool_b_fraction, float* tangent_pool_b_exchange, float* tangent_t1_pool_b, float* tangent_t2_pool_b, float* tangent_pool_b_shift, float* tangent_duration, float* tangent_flip, float* tangent_phase, float* saturation, float* rf_frequency, float* profile, std::int32_t* profile_index, float* lineshape, float* pairs, std::int32_t* pair_index, float* pair_direction, std::int32_t* duration_row, float* pool_table, float* output_real, float* output_imag, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, float flow_scale, float washout_scale, float profile_step, float lineshape_step, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shim_rows, bsk::index_t shimmed, bsk::index_t locations, bsk::index_t profiled, bsk::index_t profile_bins, bsk::index_t dynamic, bsk::index_t broadened, bsk::index_t lineshape_bins, bsk::index_t pools, bsk::index_t narrow, bsk::index_t tabulated, bsk::index_t off_axis, bsk::index_t moving, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V atom_b0{}; bsk::V atom_b1{}; bsk::V atom_b1_phase{}; @@ -7894,7 +7894,7 @@ BSK_HD void _epg_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b dinversion = bsk::ld((tangent_inversion_efficiency + scalar_atom), active_atom, 0.0f); } auto event_base = (train * event_count); - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { auto event_dt = _event_value(duration, event_base, event, active_atom, single_train); auto ddt = _event_value(tangent_duration, event_base, event, active_atom, single_train); auto r1 = bsk::truediv(1000.0f, atom_t1); @@ -8537,7 +8537,7 @@ BSK_HD void _epg_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* b } } -BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_duration, float* dot_flip, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_value, float* trajectory_tangent, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shim_rows, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* dot_t1, float* dot_t2, float* dot_m0, float* dot_b1, float* dot_inversion_efficiency, float* dot_diffusion, float* dot_duration, float* dot_flip, float* grad_output_imag, float* grad_tissue_value, float* grad_tissue_tangent, float* grad_flip_value, float* grad_flip_tangent, float* grad_duration_value, float* grad_duration_tangent, float* trajectory_value, float* trajectory_tangent, bsk::index_t problem_base, bsk::index_t problem_end, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shim_rows, bsk::index_t shimmed, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V adjoint_mt{}; bsk::V adjoint_mv{}; bsk::V adjoint_pt{}; @@ -8726,7 +8726,7 @@ BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, auto rate2_value = bsk::truediv(1000.0f, atom_t2); auto rate2_tangent = bsk::truediv((-1000.0f * atom_dot_t2), (atom_t2 * atom_t2)); auto event_base = (train * event_count); - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { slot = (trajectory + (event * record_stride)); bsk::st((trajectory_value + slot), plus_value, state_mask); bsk::st(((trajectory_value + slot) + minus_plane), minus_value, state_mask); @@ -8877,7 +8877,7 @@ BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, grad_inversion_tangent = zero; grad_damping_value = zero; grad_damping_tangent = zero; - for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + for (bsk::index_t reverse = 0; reverse < event_count; reverse += 1) { event = ((event_count - 1) - reverse); slot = (trajectory + (event * record_stride)); auto entry_pv = bsk::ld((trajectory_value + slot), state_mask, 0.0f); @@ -9154,7 +9154,7 @@ BSK_HD void _epg_real_vjp_jvp_kernel(float* t1, float* t2, float* m0, float* b1, bsk::atomic_add(((grad_tissue_tangent + ((7 + past_transmit) * atom_count)) + atom), grad_damping_tangent, active_atom); } -BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* output_real, float* output_imag, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shimmed, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V alpha{}; bsk::V atom_b1{}; bsk::V atom_damping{}; @@ -9220,7 +9220,7 @@ BSK_HD void _epg_real_kernel(float* t1, float* t2, float* m0, float* b1, float* // bookkeeping serve two of them. Two is where it stops paying: four was // measured slower, and the body is already large enough that widening it // costs registers. - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { // Read here rather than through the helper: one train gives a duration // the whole program shares, and the skip and the memo below both want // to compare it as the single number it is. @@ -9500,7 +9500,7 @@ BSK_HD auto _three_pool_interval_adjoint(const T0& table, const T1& row, const T return bsk::make_tup(grad_dt, grad_att); } -BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* lineshape, float* profile, std::int32_t* profile_index, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, std::int64_t row_count, float* grad_pair, float* grad_output_real, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_phase, float* grad_duration, float* trajectory_r, float* trajectory_i, std::int64_t problem_base, std::int64_t problem_end, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, float flow_scale, float washout_scale, std::int64_t shim_rows, float profile_step, float lineshape_step, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t locations, std::int64_t profiled, std::int64_t profile_bins, std::int64_t dynamic, std::int64_t broadened, std::int64_t lineshape_bins, std::int64_t pools, std::int64_t narrow, std::int64_t tabulated, std::int64_t off_axis, std::int64_t moving, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t recording, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b1_phase, float* b0, float* inversion_efficiency, float* diffusion, float* velocity, float* bound_fraction, float* exchange_rate, float* t1_bound, float* pool_b_fraction, float* pool_b_exchange, float* t1_pool_b, float* t2_pool_b, float* pool_b_shift, float* duration, std::int32_t* kind, float* flip, float* phase, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* saturation, float* rf_frequency, float* lineshape, float* profile, std::int32_t* profile_index, float* pairs, std::int32_t* pair_index, std::int32_t* duration_row, float* pool_table, float* pool_bars, float* pool_durations, bsk::index_t row_count, float* grad_pair, float* grad_output_real, float* grad_output_imag, float* grad_tissue, float* grad_flip, float* grad_phase, float* grad_duration, float* trajectory_r, float* trajectory_i, bsk::index_t problem_base, bsk::index_t problem_end, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, float flow_scale, float washout_scale, bsk::index_t shim_rows, float profile_step, float lineshape_step, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shimmed, bsk::index_t locations, bsk::index_t profiled, bsk::index_t profile_bins, bsk::index_t dynamic, bsk::index_t broadened, bsk::index_t lineshape_bins, bsk::index_t pools, bsk::index_t narrow, bsk::index_t tabulated, bsk::index_t off_axis, bsk::index_t moving, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t recording, bsk::index_t block_states, bsk::index_t problems) { bsk::V _d11{}; bsk::V _d12{}; bsk::V _d21{}; @@ -9644,9 +9644,9 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b bsk::V grad_alpha_v{}; bsk::V grad_angle_v{}; bsk::V grad_e1_v{}; - bsk::V grow_free{}; - bsk::V grow_pool_b{}; - bsk::V grow_semisolid{}; + bsk::V grow_free{}; + bsk::V grow_pool_b{}; + bsk::V grow_semisolid{}; bsk::V h11i{}; bsk::V h11r{}; bsk::V h12i{}; @@ -9929,16 +9929,16 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b bsk::V vr_{}; bsk::tup, bsk::V> w0{}; bsk::tup, bsk::V> w1{}; - bsk::V w11{}; - bsk::V w12{}; - bsk::V w13{}; + bsk::V w11{}; + bsk::V w12{}; + bsk::V w13{}; bsk::tup, bsk::V> w2{}; - bsk::V w21{}; - bsk::V w22{}; - bsk::V w23{}; - bsk::V w31{}; - bsk::V w32{}; - bsk::V w33{}; + bsk::V w21{}; + bsk::V w22{}; + bsk::V w23{}; + bsk::V w31{}; + bsk::V w32{}; + bsk::V w33{}; bsk::V wash_v{}; bsk::V wbvi{}; bsk::V wbvr{}; @@ -10088,7 +10088,7 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b // and the two are launched separately: each compiles the sweep it is // asked for and no more. if (bsk::truth(recording)) { - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { slot = (trajectory + (event * record_stride)); bsk::st((trajectory_r + slot), pvr, state_mask); bsk::st((trajectory_i + slot), pvi, state_mask); @@ -10567,7 +10567,7 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b ubvi = empty; wbvr = empty; wbvi = empty; - for (std::int64_t reverse = 0; reverse < event_count; reverse += 1) { + for (bsk::index_t reverse = 0; reverse < event_count; reverse += 1) { event = ((event_count - 1) - reverse); slot = (trajectory + (event * record_stride)); auto xpvr = bsk::ld((trajectory_r + slot), state_mask, 0.0f); @@ -11833,7 +11833,7 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b // walk back pooled the cotangents the eigenvalues are pushed through, // and the closed form is linear in them, so the pieces of the sum are // the sum of the pieces. - for (std::int64_t row = 0; row < row_count; row += 1) { + for (bsk::index_t row = 0; row < row_count; row += 1) { held = (pool_bars + (((local * row_count) + row) * 12)); auto row_dt = (bsk::ld((pool_durations + row)) + zero); auto one_att = bsk::select(bsk::truth(moving), _washout(atom_washout, row_dt), (1.0f + (0.0f * row_dt))); @@ -12008,7 +12008,7 @@ BSK_HD void _epg_vjp_kernel(float* t1, float* t2, float* m0, float* b1, float* b }); } -BSK_HD void _epg_real_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_duration, float* tangent_flip, float* output_real, float* output_imag, std::int64_t atom_count, std::int64_t train_count, std::int64_t event_count, std::int64_t output_count, std::int64_t state_count, std::int64_t single_train, std::int64_t atom_stride, std::int64_t shimmed, std::int64_t diffusing, std::int64_t transmit, std::int64_t density, std::int64_t inverting, std::int64_t block_states, std::int64_t problems) { +BSK_HD void _epg_real_jvp_kernel(float* t1, float* t2, float* m0, float* b1, float* inversion_efficiency, float* diffusion, float* duration, std::int32_t* kind, float* flip, std::uint8_t* action, std::int32_t* output_index, std::int32_t* shim_index, float* tangent_t1, float* tangent_t2, float* tangent_m0, float* tangent_b1, float* tangent_inversion_efficiency, float* tangent_diffusion, float* tangent_duration, float* tangent_flip, float* output_real, float* output_imag, bsk::index_t atom_count, bsk::index_t train_count, bsk::index_t event_count, bsk::index_t output_count, bsk::index_t state_count, bsk::index_t single_train, bsk::index_t atom_stride, bsk::index_t shimmed, bsk::index_t diffusing, bsk::index_t transmit, bsk::index_t density, bsk::index_t inverting, bsk::index_t block_states, bsk::index_t problems) { bsk::V atom_b1{}; bsk::V atom_damping{}; bsk::V atom_inversion{}; @@ -12096,7 +12096,7 @@ BSK_HD void _epg_real_jvp_kernel(float* t1, float* t2, float* m0, float* b1, flo // bookkeeping serve two of them. Two is where it stops paying: four was // measured slower, and the body is already large enough that widening it // costs registers. - for (std::int64_t event = 0; event < event_count; event += 1) { + for (bsk::index_t event = 0; event < event_count; event += 1) { auto dt = _event_value(duration, event_base, event, active_atom, single_train); auto dot_dt = _event_value(tangent_duration, event_base, event, active_atom, single_train); // An event of no duration relaxes nothing, and carries no tangent along diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index b02e2275..4eda0426 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -153,10 +153,10 @@ PyObject* launch(PyObject*, PyObject* args) { sizeof(unsigned long long) * (threads.x * threads.y + bsk::MAX_Z * bsk::MAX_Z); void* parameters[] = {&request.arguments, &request.z}; const int bounded = threads.x * threads.y > 256 ? 1 : 0; - // A specialized kernel is compiled for 256 threads; a wider block runs the - // kernel compiled for every combination. + // A specialized kernel is compiled for 256 threads and for rows a warp + // wide; a wider block or row runs the kernel compiled for every combination. const void* function = KERNEL_FUNCTIONS[request.kernel][bounded]; - if (specializing && !bounded) { + if (specializing && !bounded && threads.x <= 32) { if (const void* special = matching(request)) { function = special; ++specialized_launches; diff --git a/src/blochsim/_gpu_special.cu.in b/src/blochsim/_gpu_special.cu.in index f74a7d8d..a707a042 100644 --- a/src/blochsim/_gpu_special.cu.in +++ b/src/blochsim/_gpu_special.cu.in @@ -2,6 +2,8 @@ // writes one of these per entry of _specializations.json: @INDEX@ is the // entry, @KERNEL@ the kernel, and the switches it fixes follow ``Call``. #define BLOCHSIM_SIMT 1 +// Launched only on rows a warp wide or narrower; see _tile.hpp. +#define BLOCHSIM_ROWS_IN_A_WARP 1 #include "_lanes.hpp" #define BLOCHSIM_Y_LANES BLOCHSIM_LANES_@KERNEL@ diff --git a/src/blochsim/_launch.hpp b/src/blochsim/_launch.hpp index e18b3411..37a60496 100644 --- a/src/blochsim/_launch.hpp +++ b/src/blochsim/_launch.hpp @@ -93,6 +93,14 @@ inline bool read_launch(PyObject* name, PyObject* grid, PyObject* args, Launch& break; default: arg.i = PyLong_AsLongLong(item); + // The EPG kernels index in 32 bits, as Triton did for an + // argument that fit; one that does not is refused, not cut. + if (!PyErr_Occurred() && (arg.i > 2147483647LL || arg.i < -2147483648LL)) { + PyErr_Format(PyExc_ValueError, + "%s: argument %zd is %lld, past what a kernel indexes with", + info.name, i, static_cast(arg.i)); + return false; + } break; } if (PyErr_Occurred()) { diff --git a/src/blochsim/_tile.hpp b/src/blochsim/_tile.hpp index 99f76e39..cf2f5079 100644 --- a/src/blochsim/_tile.hpp +++ b/src/blochsim/_tile.hpp @@ -39,6 +39,11 @@ namespace bsk { +// The integers the EPG kernels index with: 32 bits, as Triton passed every +// integer argument that fit. An offset that can pass 2^31 is cast to 64 bits +// where it is formed, and the launcher refuses an argument that does not fit. +using index_t = std::int32_t; + // --------------------------------------------------------------------------- // The launch a program belongs to. // --------------------------------------------------------------------------- @@ -66,9 +71,8 @@ __device__ __forceinline__ void enter(int nz) { } __syncthreads(); } -BSK_HD std::int64_t program_id(int axis) { - return axis == 0 ? static_cast(blockIdx.x) - : static_cast(blockIdx.y); +BSK_HD index_t program_id(int axis) { + return axis == 0 ? static_cast(blockIdx.x) : static_cast(blockIdx.y); } #else @@ -85,7 +89,7 @@ inline thread_local HostProgram program; inline int width_x() { return program.nx; } inline int width_y() { return program.ny; } inline int width_z() { return program.nz; } -inline std::int64_t program_id(int axis) { return program.pid[axis]; } +inline index_t program_id(int axis) { return static_cast(program.pid[axis]); } #endif @@ -489,14 +493,46 @@ BSK_HD auto s_max(X x, Y y) { } } +// A function too costly to take once per row where a thread's rows all hold +// the same argument -- a pulse's flip, a shared interval -- as they do where +// a value read once per event is spread over the problems. The test is a few +// compares; the result is the same either way, bit for bit. +template +BSK_HD auto once_if_shared(F f, const A& a) { +#if defined(BLOCHSIM_SIMT) + if constexpr (axes_of != 0) { + using Tile = std::decay_t; + if constexpr (Tile::rows > 1 && Tile::depth == 1) { + bool shared = true; +#pragma unroll + for (int y = 1; y < Tile::rows; ++y) { + shared = shared && a.v[y] == a.v[0]; + } + if (shared) { + using R = decltype(f(a.v[0])); + return V>(f(a.v[0])); + } + } + } +#endif + return zip(f, a); +} + +#define BSK_COSTLY_UNARY(name) \ + template \ + BSK_HD auto name(const A& a) { \ + return once_if_shared([](auto x) { return s_##name(x); }, a); \ + } +BSK_COSTLY_UNARY(exp) +BSK_COSTLY_UNARY(cos) +BSK_COSTLY_UNARY(sin) +#undef BSK_COSTLY_UNARY + #define BSK_UNARY(name) \ template \ BSK_HD auto name(const A& a) { \ return zip([](auto x) { return s_##name(x); }, a); \ } -BSK_UNARY(exp) -BSK_UNARY(cos) -BSK_UNARY(sin) BSK_UNARY(sqrt) BSK_UNARY(floor) BSK_UNARY(rint) @@ -736,6 +772,20 @@ BSK_HD T shuffle(T value, int lane, int width) { // Every thread of the block runs these together: a kernel's control flow // depends on nothing a single thread holds. +// +// A row wider than a warp goes through shared memory, in a function of its +// own: inlined, the compiler predicates its barriers and accesses rather than +// branching past them, and every row a warp wide pays for them. +template +__device__ __noinline__ T reduce_wide_row(T value, Op op); + +template +__device__ __noinline__ T gather_wide_row(T value, int lane); + +// A kernel compiled with BLOCHSIM_ROWS_IN_A_WARP is launched only on rows a +// warp wide or narrower, so its gathers and reductions are shuffles with no +// test: the test splits the code between them, and the indices and masks each +// recomputes can no longer be shared. template BSK_HD T reduce_row(T value, Op op) { const int nx = width_x(); @@ -743,9 +793,19 @@ BSK_HD T reduce_row(T value, Op op) { for (int offset = width >> 1; offset > 0; offset >>= 1) { value = op(value, shuffle_xor(value, offset, width)); } +#if defined(BLOCHSIM_ROWS_IN_A_WARP) + return value; +#else if (nx <= 32) { return value; } + return reduce_wide_row(value, op); +#endif +} + +template +__device__ __noinline__ T reduce_wide_row(T value, Op op) { + const int nx = width_x(); T* words = reinterpret_cast(shared_words); const int warps = nx >> 5; __syncthreads(); @@ -762,6 +822,19 @@ BSK_HD T reduce_row(T value, Op op) { template BSK_HD T reduce_column(T value, Op op) { + const int nx = width_x(); + const int ny = static_cast(blockDim.y); + if (ny == 1) { + return value; + } + if (nx * ny <= 32) { + // One warp: a column's rows sit nx lanes apart, and nx and ny are + // powers of two. + for (int offset = nx; offset < nx * ny; offset <<= 1) { + value = op(value, shuffle_xor(value, offset, 32)); + } + return value; + } T* words = reinterpret_cast(shared_words); __syncthreads(); words[threadIdx.y * blockDim.x + threadIdx.x] = value; @@ -776,9 +849,19 @@ BSK_HD T reduce_column(T value, Op op) { template BSK_HD T gather_row(T value, int lane) { const int nx = width_x(); +#if defined(BLOCHSIM_ROWS_IN_A_WARP) + return shuffle(value, lane, nx); +#else if (nx <= 32) { return shuffle(value, lane, nx); } + return gather_wide_row(value, lane); +#endif +} + +template +__device__ __noinline__ T gather_wide_row(T value, int lane) { + const int nx = width_x(); T* words = reinterpret_cast(shared_words); __syncthreads(); words[threadIdx.y * nx + threadIdx.x] = value; diff --git a/src/blochsim/sequence/_epg_gpu.py b/src/blochsim/sequence/_epg_gpu.py index acfc9bbe..b618d465 100644 --- a/src/blochsim/sequence/_epg_gpu.py +++ b/src/blochsim/sequence/_epg_gpu.py @@ -284,9 +284,10 @@ def _three_pool_table( return table -# Threads of one program: a warp, its lanes along the states and, where the -# states are fewer, across rows of problems. -_PROGRAM_THREADS = 32 +# Threads of one program: two warps, their lanes along the states and, where +# the states are fewer, across rows of problems. A program of one warp caps a +# card at as many warps as it runs blocks, short of what the registers allow. +_PROGRAM_THREADS = 64 def _atom_stride(*tuples: tuple[torch.Tensor, ...]) -> int: diff --git a/tests/sequence/test_both_pools.py b/tests/sequence/test_both_pools.py index fafe9213..55a33fbe 100644 --- a/tests/sequence/test_both_pools.py +++ b/tests/sequence/test_both_pools.py @@ -1677,4 +1677,7 @@ def leaves(value): float((left - right).abs().max()) for left, right in zip(narrow, roots, strict=True) ) - assert worst / largest < 1e-5, (name, worst / largest) + # The roots branch moves by about 1e-5 of the largest value between + # correct compilations of itself -- one fused multiply-add more or + # less -- and Triton's differs from the CUDA build's by 2e-5. + assert worst / largest < 5e-5, (name, worst / largest) From 43a692669562030e224a074b2c4d62f82bc058d6 Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 00:16:15 +0200 Subject: [PATCH 08/16] Run the forward EPG kernels and their JVPs written for their layouts The complex and real forward kernels and their Jacobian-vector products are written once over a number type: a float gives the simulation and a dual number its derivative along a direction. Each is compiled for its layout -- pools, how a pulse is formed, per-voxel maps, problems per thread, one train or several -- and reads every other switch at run time, branching around whole blocks once per event. A launch whose rows fit a warp runs the layout ahead of any tile kernel. The two-pool transverse step takes its complex root in a form that does not cancel just above the negative real axis, where its discriminant usually sits. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 14 + CMakeLists.txt | 16 + src/blochsim/_gpu.cu | 36 + src/blochsim/_gpu_launch.py | 14 +- src/blochsim/_layout.cu | 195 ++++ src/blochsim/_layout.hpp | 39 + src/blochsim/_layout_complex.cu.in | 52 + src/blochsim/_layout_complex.hpp | 1190 ++++++++++++++++++++ src/blochsim/_layout_numbers.hpp | 271 +++++ src/blochsim/_layout_real.cu | 45 + src/blochsim/_layout_real.hpp | 213 ++++ tests/sequence/test_specialized_kernels.py | 18 +- 12 files changed, 2093 insertions(+), 10 deletions(-) create mode 100644 src/blochsim/_layout.cu create mode 100644 src/blochsim/_layout.hpp create mode 100644 src/blochsim/_layout_complex.cu.in create mode 100644 src/blochsim/_layout_complex.hpp create mode 100644 src/blochsim/_layout_numbers.hpp create mode 100644 src/blochsim/_layout_real.cu create mode 100644 src/blochsim/_layout_real.hpp diff --git a/CLAUDE.md b/CLAUDE.md index 92d5b765..f05b1a47 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -92,6 +92,20 @@ the general kernels alone. A specialized kernel is compiled for rows of at most 32 state orders, so its shifts are shuffles with no test; a wider launch runs the general one. +**The forward EPG kernels and their Jacobian-vector products are written for +their layouts** (`_layout.hpp`), and a launch whose rows fit a warp runs them +ahead of any tile kernel. A layout is what is compiled: the pools, how a +pulse is formed, whether the tissue has per-voxel maps, the problems a thread +holds, one train or several; every other switch is read at run time and +steers whole blocks once per event, so one compile serves every combination +of them. The loops are written once over a number type +(`_layout_numbers.hpp`): at `float` they are the forward simulation, at +`num::Dual` -- a value and its derivative along the direction -- the +Jacobian-vector product. `_gpu_launch.generic_kernels()` turns layouts off +with the specializations, and `layout_launches()` counts them. They exist only +on the card: `_gpu_host` compiles the tile kernels, so the host lane holds +those, not these, to the C++ kernels. + **The EPG kernels index in 32 bits** (`bsk::index_t`), as Triton did for every integer argument that fit. An offset that can pass 2^31 is cast to 64 bits where it is formed, as the Triton source cast it, and the launcher refuses an diff --git a/CMakeLists.txt b/CMakeLists.txt index 93f40643..3723e973 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -151,6 +151,22 @@ if(BLOCHSIM_CUDA) CONTENT "// Written by CMake from _specializations.json: X(index, kernel, \"switch=value,...\").\n#define BLOCHSIM_FOR_EACH_SPECIAL(X) \\\n${_blochsim_special_table}\n" @ONLY) + # The EPG kernels written for their layouts (_layout.hpp): a unit per + # kernel and pool layout, so they compile side by side. + list(APPEND _blochsim_gpu_sources src/blochsim/_layout.cu src/blochsim/_layout_real.cu) + foreach(KIND IN ITEMS forward jvp) + if(KIND STREQUAL "forward") + set(TYPE "float") + else() + set(TYPE "num::Dual") + endif() + foreach(POOLS RANGE 3) + set(_source "${CMAKE_CURRENT_BINARY_DIR}/layout/complex_${KIND}_${POOLS}.cu") + configure_file(src/blochsim/_layout_complex.cu.in "${_source}" @ONLY) + list(APPEND _blochsim_gpu_sources "${_source}") + endforeach() + endforeach() + python_add_library(_gpu MODULE USE_SABI ${BLOCHSIM_ABI3_VERSION} WITH_SOABI diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index 4eda0426..eca13e5d 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -10,6 +10,7 @@ #define BLOCHSIM_TABLE_ONLY 1 #include "_launch.hpp" +#include "_layout.hpp" #include "_special.hpp" #include "_special_table.hpp" @@ -87,6 +88,9 @@ const std::vector>& specials() { // Whether a launch may run a specialized kernel, and how many have. bool specializing = true; unsigned long long specialized_launches = 0; +// Whether a launch may run its kernel's layout (_layout.hpp), and how many have. +bool laying_out = true; +unsigned long long layout_launches = 0; // The specialized kernel whose fixed switches this launch matches, if any. const void* matching(const blochsim_launch::Launch& request) { @@ -145,6 +149,20 @@ PyObject* launch(PyObject*, PyObject* args) { if (status != cudaSuccess) { return cuda_error(status, "selecting the device"); } + if (laying_out) { + const int laid = blochsim_layout::launch(request.kernel, request.arguments, + reinterpret_cast(stream)); + if (laid >= 0) { + ++layout_launches; + if (previous != device) { + cudaSetDevice(previous); + } + if (laid != cudaSuccess) { + return cuda_error(static_cast(laid), bsk::KERNELS[request.kernel].name); + } + Py_RETURN_NONE; + } + } const dim3 blocks(static_cast(request.grid[0]), static_cast(request.grid[1])); const dim3 threads(static_cast(request.block[0]), static_cast(request.block[1] / request.lanes)); // A row's reduction and gather go through one word per thread, and a @@ -207,6 +225,20 @@ PyObject* specialized_launch_count(PyObject*, PyObject*) { return PyLong_FromUnsignedLongLong(specialized_launches); } +PyObject* use_layouts(PyObject*, PyObject* args) { + int on = 1; + if (!PyArg_ParseTuple(args, "p", &on)) { + return nullptr; + } + const bool previous = laying_out; + laying_out = on != 0; + return PyBool_FromLong(previous); +} + +PyObject* layout_launch_count(PyObject*, PyObject*) { + return PyLong_FromUnsignedLongLong(layout_launches); +} + PyMethodDef METHODS[] = { {"kernels", blochsim_launch::kernel_table, METH_NOARGS, "Each kernel's parameter names and kinds."}, @@ -218,6 +250,10 @@ PyMethodDef METHODS[] = { "Whether launches may run specialized kernels; returns the previous setting."}, {"specialized_launches", specialized_launch_count, METH_NOARGS, "How many launches have run a specialized kernel."}, + {"use_layouts", use_layouts, METH_VARARGS, + "Whether launches may run their kernel's layout; returns the previous setting."}, + {"layout_launches", layout_launch_count, METH_NOARGS, + "How many launches have run a kernel's layout."}, {nullptr, nullptr, 0, nullptr}, }; diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index 874a5569..19b12a68 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -87,17 +87,25 @@ def specialized_launches() -> int: return _module("cuda").specialized_launches() if available() else 0 +def layout_launches() -> int: + """How many launches on a card have run a kernel written for its layout.""" + return _module("cuda").layout_launches() if available() else 0 + + @contextmanager def generic_kernels() -> Iterator[None]: - """Run every launch inside on the kernels compiled for all combinations.""" + """Run every launch inside on the tile kernels compiled for all combinations.""" if not available(): yield return - previous = _module("cuda").use_specializations(False) + module = _module("cuda") + special = module.use_specializations(False) + layouts = module.use_layouts(False) try: yield finally: - _module("cuda").use_specializations(previous) + module.use_specializations(special) + module.use_layouts(layouts) class Kernel: diff --git a/src/blochsim/_layout.cu b/src/blochsim/_layout.cu new file mode 100644 index 00000000..6cb80c94 --- /dev/null +++ b/src/blochsim/_layout.cu @@ -0,0 +1,195 @@ +// Which kernels have a layout, and their arguments read by name. +#define BLOCHSIM_TABLE_ONLY 1 + +#include "_layout.hpp" +#include "_special.hpp" + +namespace blochsim_layout { +namespace { + +// An argument by name, or nothing where this kernel has no such parameter. +struct Reader { + int kernel; + const bsk::Arg* a; + + int at(const char* name) const { return bsk::param_index(kernel, name); } + const float* floats(const char* name) const { + const int index = at(name); + return index < 0 ? nullptr : static_cast(a[index].p); + } + const int* ints(const char* name) const { + const int index = at(name); + return index < 0 ? nullptr : static_cast(a[index].p); + } + float* outputs(const char* name) const { return static_cast(a[at(name)].p); } + long long integer(const char* name) const { return a[at(name)].i; } + float real(const char* name) const { return static_cast(a[at(name)].f); } + bool flag(const char* name) const { return a[at(name)].i != 0; } +}; + +epg::Params complex_params(const Reader& r) { + epg::Params p{}; + p.t1 = r.floats("t1"); + p.t2 = r.floats("t2"); + p.m0 = r.floats("m0"); + p.b1 = r.floats("b1"); + p.b1_phase = r.floats("b1_phase"); + p.b0 = r.floats("b0"); + p.inversion_efficiency = r.floats("inversion_efficiency"); + p.diffusion = r.floats("diffusion"); + p.velocity = r.floats("velocity"); + p.bound_fraction = r.floats("bound_fraction"); + p.bound_exchange = r.at("bound_exchange") >= 0 ? r.floats("bound_exchange") : r.floats("exchange_rate"); + p.t1_bound = r.floats("t1_bound"); + p.pool_b_fraction = r.floats("pool_b_fraction"); + p.pool_b_exchange = r.floats("pool_b_exchange"); + p.t1_pool_b = r.floats("t1_pool_b"); + p.t2_pool_b = r.floats("t2_pool_b"); + p.pool_b_shift = r.floats("pool_b_shift"); + p.duration = r.floats("duration"); + p.flip = r.floats("flip"); + p.phase = r.floats("phase"); + p.phase_cos = r.floats("phase_cos"); + p.phase_sin = r.floats("phase_sin"); + p.profile = r.floats("profile"); + p.saturation = r.floats("saturation"); + p.rf_frequency = r.floats("rf_frequency"); + p.lineshape = r.floats("lineshape"); + p.pairs = r.floats("pairs"); + p.pool_table = r.floats("pool_table"); + p.kind = r.ints("kind"); + p.output_index = r.ints("output_index"); + p.shim_index = r.ints("shim_index"); + p.profile_index = r.ints("profile_index"); + p.pair_index = r.ints("pair_index"); + p.duration_row = r.ints("duration_row"); + p.action = static_cast(r.a[r.at("action")].p); + p.d_t1 = r.floats("tangent_t1"); + p.d_t2 = r.floats("tangent_t2"); + p.d_m0 = r.floats("tangent_m0"); + p.d_b1 = r.floats("tangent_b1"); + p.d_b1_phase = r.floats("tangent_b1_phase"); + p.d_b0 = r.floats("tangent_b0"); + p.d_inversion_efficiency = r.floats("tangent_inversion_efficiency"); + p.d_diffusion = r.floats("tangent_diffusion"); + p.d_velocity = r.floats("tangent_velocity"); + p.d_bound_fraction = r.floats("tangent_bound_fraction"); + p.d_bound_exchange = r.floats("tangent_exchange_rate"); + p.d_t1_bound = r.floats("tangent_t1_bound"); + p.d_pool_b_fraction = r.floats("tangent_pool_b_fraction"); + p.d_pool_b_exchange = r.floats("tangent_pool_b_exchange"); + p.d_t1_pool_b = r.floats("tangent_t1_pool_b"); + p.d_t2_pool_b = r.floats("tangent_t2_pool_b"); + p.d_pool_b_shift = r.floats("tangent_pool_b_shift"); + p.d_duration = r.floats("tangent_duration"); + p.d_flip = r.floats("tangent_flip"); + p.d_phase = r.floats("tangent_phase"); + p.pair_direction = r.floats("pair_direction"); + p.output_real = r.outputs("output_real"); + p.output_imag = r.outputs("output_imag"); + p.atom_count = static_cast(r.integer("atom_count")); + p.train_count = static_cast(r.integer("train_count")); + p.event_count = static_cast(r.integer("event_count")); + p.output_count = static_cast(r.integer("output_count")); + p.state_count = static_cast(r.integer("state_count")); + p.width = static_cast(r.integer("block_states")); + p.locations = static_cast(r.integer("locations")); + p.profile_bins = static_cast(r.integer("profile_bins")); + p.lineshape_bins = static_cast(r.integer("lineshape_bins")); + p.flow_scale = r.real("flow_scale"); + p.washout_scale = r.real("washout_scale"); + p.profile_step = r.real("profile_step"); + p.lineshape_step = r.real("lineshape_step"); + p.single_train = r.flag("single_train"); + p.atom_stride = r.flag("atom_stride"); + p.shimmed = r.flag("shimmed"); + p.off_axis = r.flag("off_axis"); + p.moving = r.flag("moving"); + p.diffusing = r.flag("diffusing"); + p.transmit = r.flag("transmit"); + p.density = r.flag("density"); + p.inverting = r.flag("inverting"); + return p; +} + +layout_real::Params real_params(const Reader& r) { + layout_real::Params p{}; + p.t1 = r.floats("t1"); + p.t2 = r.floats("t2"); + p.m0 = r.floats("m0"); + p.b1 = r.floats("b1"); + p.inversion_efficiency = r.floats("inversion_efficiency"); + p.diffusion = r.floats("diffusion"); + p.duration = r.floats("duration"); + p.flip = r.floats("flip"); + p.d_t1 = r.floats("tangent_t1"); + p.d_t2 = r.floats("tangent_t2"); + p.d_m0 = r.floats("tangent_m0"); + p.d_b1 = r.floats("tangent_b1"); + p.d_inversion_efficiency = r.floats("tangent_inversion_efficiency"); + p.d_diffusion = r.floats("tangent_diffusion"); + p.d_duration = r.floats("tangent_duration"); + p.d_flip = r.floats("tangent_flip"); + p.kind = r.ints("kind"); + p.output_index = r.ints("output_index"); + p.shim_index = r.ints("shim_index"); + p.action = static_cast(r.a[r.at("action")].p); + p.output_real = r.outputs("output_real"); + p.output_imag = r.outputs("output_imag"); + p.atom_count = static_cast(r.integer("atom_count")); + p.train_count = static_cast(r.integer("train_count")); + p.event_count = static_cast(r.integer("event_count")); + p.output_count = static_cast(r.integer("output_count")); + p.state_count = static_cast(r.integer("state_count")); + p.width = static_cast(r.integer("block_states")); + p.single_train = r.flag("single_train"); + p.atom_stride = r.flag("atom_stride"); + p.shimmed = r.flag("shimmed"); + p.diffusing = r.flag("diffusing"); + p.transmit = r.flag("transmit"); + p.density = r.flag("density"); + p.inverting = r.flag("inverting"); + return p; +} + +int complex_launch(bool jvp, const Reader& r, cudaStream_t stream) { + const epg::Params p = complex_params(r); + const int rf = r.flag("dynamic") ? epg::DYNAMIC : (r.flag("profiled") ? epg::PROFILE : epg::HARD); + const int mode = r.flag("tabulated") ? epg::TABLE : (r.flag("narrow") ? epg::NARROW : epg::ROOTS); + switch (static_cast(r.integer("pools")) + (jvp ? 4 : 0)) { + case 0: return complex_forward_0(p, rf, mode, stream); + case 1: return complex_forward_1(p, rf, mode, stream); + case 2: return complex_forward_2(p, rf, mode, stream); + case 3: return complex_forward_3(p, rf, mode, stream); + case 4: return complex_jvp_0(p, rf, mode, stream); + case 5: return complex_jvp_1(p, rf, mode, stream); + case 6: return complex_jvp_2(p, rf, mode, stream); + case 7: return complex_jvp_3(p, rf, mode, stream); + default: return -1; + } +} + +const int COMPLEX = bsk::kernel_index("_epg_kernel"); +const int COMPLEX_JVP = bsk::kernel_index("_epg_jvp_kernel"); +const int REAL = bsk::kernel_index("_epg_real_kernel"); +const int REAL_JVP = bsk::kernel_index("_epg_real_jvp_kernel"); + +} // namespace + +int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream) { + const Reader r{kernel, arguments.a}; + if (kernel != COMPLEX && kernel != COMPLEX_JVP && kernel != REAL && kernel != REAL_JVP) { + return -1; + } + // A layout's rows are at most a warp wide. + if (r.integer("block_states") > 32) { + return -1; + } + if (kernel == COMPLEX || kernel == COMPLEX_JVP) { + return complex_launch(kernel == COMPLEX_JVP, r, stream); + } + const layout_real::Params p = real_params(r); + return kernel == REAL ? real_forward(p, stream) : real_jvp(p, stream); +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout.hpp b/src/blochsim/_layout.hpp new file mode 100644 index 00000000..b998f560 --- /dev/null +++ b/src/blochsim/_layout.hpp @@ -0,0 +1,39 @@ +// The EPG kernels written for their layouts rather than over tiles. +// +// A layout is what a kernel is compiled for: the pools it carries, how a +// pulse is formed, whether the tissue has per-voxel maps, how many problems a +// thread holds and whether there is one train. Every other switch is read at +// run time and steers whole blocks once per event, so one compile serves +// every combination of them. The forward kernels and their Jacobian-vector +// products are one source over a number type: a float, or a dual number that +// carries a direction beside each value (_layout_numbers.hpp). +#pragma once + +#include + +#include "_kernels.hpp" +#include "_layout_complex.hpp" +#include "_layout_real.hpp" + +namespace blochsim_layout { + +// Queue ``kernel`` on ``stream`` and return its launch status, or -1 where it +// has no layout for these arguments and the tile kernel is to run instead. +int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream); + +// One translation unit per kernel and pool layout, so a build compiles them +// side by side. ``rf`` and ``mode`` are epg::Rf and epg::Mode. +#define BLOCHSIM_LAYOUT_COMPLEX(X) \ + X(forward, 0) X(forward, 1) X(forward, 2) X(forward, 3) X(jvp, 0) X(jvp, 1) X(jvp, 2) X(jvp, 3) +#define BLOCHSIM_LAYOUT_COMPLEX_DECLARATION(kind, pools) \ + int complex_##kind##_##pools(const epg::Params& p, int rf, int mode, cudaStream_t stream); +BLOCHSIM_LAYOUT_COMPLEX(BLOCHSIM_LAYOUT_COMPLEX_DECLARATION) +#undef BLOCHSIM_LAYOUT_COMPLEX_DECLARATION + +int real_forward(const layout_real::Params& p, cudaStream_t stream); +int real_jvp(const layout_real::Params& p, cudaStream_t stream); + +// Programs of two warps each. +constexpr int WARPS = 2; + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_complex.cu.in b/src/blochsim/_layout_complex.cu.in new file mode 100644 index 00000000..fefc1fc5 --- /dev/null +++ b/src/blochsim/_layout_complex.cu.in @@ -0,0 +1,52 @@ +// Written by CMake from _layout_complex.cu.in: the complex @KIND@ kernel +// for @POOLS@ in its layouts. +#include "_layout.hpp" + +namespace blochsim_layout { +namespace { + +template +__global__ void __launch_bounds__(32 * WARPS) complex_kernel(epg::Params p) { + epg::complex_loop(p); +} + +template +int launch_mode(const epg::Params& p, int rf, cudaStream_t stream) { + const int groups = 32 / p.width; + const long long problems = static_cast(p.train_count) * p.atom_count; + const long long per_block = static_cast(WARPS) * groups * Y; + const unsigned grid = static_cast((problems + per_block - 1) / per_block); + const bool maps = p.transmit || p.off_axis || p.density || p.inverting; +#define BLOCHSIM_GO(R, MP) \ + if (p.single_train) complex_kernel<<>>(p); \ + else complex_kernel<<>>(p) + if (rf == epg::DYNAMIC) { + if (maps) { BLOCHSIM_GO(epg::DYNAMIC, true); } else { BLOCHSIM_GO(epg::DYNAMIC, false); } + } else if (rf == epg::PROFILE) { + if (maps) { BLOCHSIM_GO(epg::PROFILE, true); } else { BLOCHSIM_GO(epg::PROFILE, false); } + } else { + if (maps) { BLOCHSIM_GO(epg::HARD, true); } else { BLOCHSIM_GO(epg::HARD, false); } + } +#undef BLOCHSIM_GO + return static_cast(cudaGetLastError()); +} + +} // namespace + +int complex_@KIND@_@POOLS@(const epg::Params& p, int rf, int mode, cudaStream_t stream) { + using T = @TYPE@; + constexpr int POOLS = @POOLS@; + // Problems a thread holds: a dual number is two registers, and a pool + // layout's operators are held beside the states. + constexpr int Y = num::is_dual::value ? (POOLS ? 1 : 2) : (POOLS ? 2 : 4); + if constexpr (POOLS == 3) { + if (mode == epg::TABLE) return launch_mode(p, rf, stream); + if (mode == epg::NARROW) return launch_mode(p, rf, stream); + return launch_mode(p, rf, stream); + } else { + (void)mode; + return launch_mode(p, rf, stream); + } +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_complex.hpp b/src/blochsim/_layout_complex.hpp new file mode 100644 index 00000000..7af7ae88 --- /dev/null +++ b/src/blochsim/_layout_complex.hpp @@ -0,0 +1,1190 @@ +// The complex EPG event loop over a number type ``T``: the forward simulation +// at ``float``, its Jacobian-vector product at ``num::Dual``. +// +// Layout (compile time): POOLS (0 free water alone, 1 beside a semisolid +// pool, 2 beside an exchanging pool b, 3 beside both), MODE (how a three-pool +// interval is formed: series, roots or a table), RF (hard pulse, slice- +// profile table or a pair per pulse per voxel), MAPS (per-voxel maps in +// registers), Y problems per thread, ONE_TRAIN. Run time: atom_stride, +// transmit, density, inverting, diffusing, off_axis, moving, shimmed. +// One warp per program; lanes along the states, groups of lanes across +// problems, Y problems of a group in each thread's registers. +#pragma once + +#include + +#include + +#include "_layout_numbers.hpp" + +namespace epg { + +using num::abs_; +using num::div_; +using num::exp_; +using num::fma_; +using num::max_; +using num::min_; +using num::same; +using num::sincos_; +using num::sqrt_; +using num::value; + +constexpr float TWO_PI = 6.283185307179586f; + +enum Rf { HARD = 0, PROFILE = 1, DYNAMIC = 2 }; +enum Mode { NARROW = 0, ROOTS = 1, TABLE = 2 }; + +// What a launch hands the loop: the forward kernel's arguments, and for a +// Jacobian-vector product the direction along every input that has one. +struct Params { + const float *t1, *t2, *m0, *b1, *b1_phase, *b0, *inversion_efficiency, *diffusion, *velocity; + const float *bound_fraction, *bound_exchange, *t1_bound, *pool_b_fraction, *pool_b_exchange; + const float *t1_pool_b, *t2_pool_b, *pool_b_shift; + const float *duration, *flip, *phase, *phase_cos, *phase_sin, *profile; + const float *saturation, *rf_frequency, *lineshape, *pairs, *pool_table; + const int *kind, *output_index, *shim_index, *profile_index, *pair_index, *duration_row; + const unsigned char* action; + // Directions, read only by a dual loop. + const float *d_t1, *d_t2, *d_m0, *d_b1, *d_b1_phase, *d_b0, *d_inversion_efficiency; + const float *d_diffusion, *d_velocity, *d_bound_fraction, *d_bound_exchange, *d_t1_bound; + const float *d_pool_b_fraction, *d_pool_b_exchange, *d_t1_pool_b, *d_t2_pool_b, *d_pool_b_shift; + const float *d_duration, *d_flip, *d_phase, *pair_direction; + float *output_real, *output_imag; + int atom_count, train_count, event_count, output_count, state_count, width; + int locations, profile_bins, lineshape_bins; + float flow_scale, washout_scale, profile_step, lineshape_step; + bool single_train, atom_stride, shimmed, off_axis, moving, diffusing, transmit, density, + inverting; +}; + +template +constexpr bool DUAL = num::is_dual::value; + +template +__device__ __forceinline__ T read(const float* values, const float* directions, long long at) { + return num::load(values, directions, at); +} + +// An event's phase as its cosine and sine: read where the launch took them, +// formed where a direction moves the phase. +template +__device__ __forceinline__ void event_phase(const Params& p, int at, T& c, T& s) { + if constexpr (DUAL) { + sincos_(T{__ldg(p.phase + at), __ldg(p.d_phase + at)}, s, c); + } else { + c = p.phase_cos[at]; + s = p.phase_sin[at]; + } +} + +// What a readout stores: the value, or along a direction its derivative. +template +__device__ __forceinline__ float stored(T x) { + if constexpr (DUAL) { + return x.d; + } else { + return x; + } +} + +// One pool through a hard pulse, in Triton's fused order. +template +__device__ __forceinline__ void rotate_flip_phase(T cosine, T sine, T cos_phi, T sin_phi, T cos_2phi, + T sin_2phi, T& fp_r, T& fp_i, T& fm_r, T& fm_i, + T& z_r, T& z_i) { + const T chs = 0.5f * (1.0f + cosine), shs = 0.5f * (1.0f - cosine), hs = 0.5f * sine; + const T m2r = fma_(cos_2phi, fm_r, -(sin_2phi * fm_i)); + const T m2i = fma_(sin_2phi, fm_r, cos_2phi * fm_i); + const T p2r = fma_(cos_2phi, fp_r, sin_2phi * fp_i); + const T p2i = fma_(cos_2phi, fp_i, -(sin_2phi * fp_r)); + const T za = fma_(sin_phi, z_r, cos_phi * z_i); + const T zb = fma_(sin_phi, z_i, -(cos_phi * z_r)); + const T zc = fma_(sin_phi, z_r, -(cos_phi * z_i)); + const T zd = fma_(cos_phi, z_r, sin_phi * z_i); + const T pr = fma_(sine, za, fma_(shs, m2r, chs * fp_r)); + const T pi = fma_(sine, zb, fma_(shs, m2i, chs * fp_i)); + const T mr = fma_(sine, zc, fma_(chs, fm_r, shs * p2r)); + const T mi = fma_(sine, zd, fma_(chs, fm_i, shs * p2i)); + const T ptr = fma_(sin_phi, fp_r, -(cos_phi * fp_i)); + const T mtr = fma_(sin_phi, fm_r, cos_phi * fm_i); + const T pti = fma_(cos_phi, fp_r, sin_phi * fp_i); + const T mti = fma_(cos_phi, fm_r, -(sin_phi * fm_i)); + const T zr = fma_(cosine, z_r, fma_(-hs, mtr, -hs * ptr)); + const T zi = fma_(cosine, z_i, fma_(hs, mti, -hs * pti)); + fp_r = pr; fp_i = pi; fm_r = mr; fm_i = mi; z_r = zr; z_i = zi; +} + +// The rotation named by its Cayley-Klein pair. +template +__device__ __forceinline__ void rotate_spinor(T ar, T ai, T br, T bi, T& fp_r, T& fp_i, T& fm_r, T& fm_i, + T& z_r, T& z_i) { + const T aa_r = ar * ar - ai * ai, aa_i = 2.0f * ar * ai; + const T bb_r = br * br - bi * bi, bb_i = 2.0f * br * bi; + const T ab_r = ar * br - ai * bi, ab_i = ar * bi + ai * br; + const T t00_r = aa_r, t00_i = -aa_i, t01_r = -bb_r, t01_i = bb_i; + const T t02_r = -2.0f * ab_r, t02_i = 2.0f * ab_i; + const T t10_r = -bb_r, t10_i = -bb_i, t11_r = aa_r, t11_i = aa_i; + const T t12_r = -2.0f * ab_r, t12_i = -2.0f * ab_i; + const T cross_r = ar * br + ai * bi, cross_i = ar * bi - ai * br; + const T t20_r = cross_r, t20_i = cross_i, t21_r = cross_r, t21_i = -cross_i; + const T t22 = ar * ar + ai * ai - br * br - bi * bi; + const T pr = t00_r * fp_r - t00_i * fp_i + t01_r * fm_r - t01_i * fm_i + t02_r * z_r - t02_i * z_i; + const T pi = t00_r * fp_i + t00_i * fp_r + t01_r * fm_i + t01_i * fm_r + t02_r * z_i + t02_i * z_r; + const T mr = t10_r * fp_r - t10_i * fp_i + t11_r * fm_r - t11_i * fm_i + t12_r * z_r - t12_i * z_i; + const T mi = t10_r * fp_i + t10_i * fp_r + t11_r * fm_i + t11_i * fm_r + t12_r * z_i + t12_i * z_r; + const T zr = t20_r * fp_r - t20_i * fp_i + t21_r * fm_r - t21_i * fm_i + t22 * z_r; + const T zi = t20_r * fp_i + t20_i * fp_r + t21_r * fm_i + t21_i * fm_r + t22 * z_i; + fp_r = pr; fp_i = pi; fm_r = mr; fm_i = mi; z_r = zr; z_i = zi; +} + +// expm((K - diag(R1)) dt) for free water beside one second pool, times the +// attenuation, as e[0..3]; what each pool recovers as grow[0..1]. Where the +// root is on its series branch it carries no direction of its own. +template +__device__ __forceinline__ void two_pool_step(T r1_free, T r1_bound, T exchange, T bound, T dt, + T attenuation, T* e, T* grow) { + const T free = 1.0f - bound; + const T kab = exchange * bound, kba = exchange * free; + const T l11 = (-kab - r1_free) * dt, l12 = kba * dt, l21 = kab * dt; + const T l22 = (-kba - r1_bound) * dt; + const T half_trace = 0.5f * (l11 + l22), half_gap = 0.5f * (l11 - l22); + const T square = half_gap * half_gap + l12 * l21; + const bool turning = value(square) > 1e-12f; + const T root = turning ? sqrt_(square) : T(num::sqrt_approx(fmaxf(value(square), 0.0f))); + const T upper = exp_(half_trace + root), lower = exp_(half_trace - root); + const T cosine = 0.5f * (upper + lower); + const T scale = turning ? div_(0.5f * (upper - lower), root) + : exp_(half_trace) * + (1.0f + square * (1.0f / 6.0f) + square * square * (1.0f / 120.0f)); + e[0] = attenuation * (cosine + scale * half_gap); + e[1] = attenuation * scale * l12; + e[2] = attenuation * scale * l21; + e[3] = attenuation * (cosine - scale * half_gap); + grow[0] = free - (e[0] * free + e[1] * bound); + grow[1] = bound - (e[2] * free + e[3] * bound); +} + +// The principal square root of a complex number; along a direction +// dz / (2 w), and none at the origin. The larger part is taken by a root +// and the smaller as im / (2 x): (|z| + re) / 2 alone cancels to nothing +// just above the negative real axis. +__device__ __forceinline__ void complex_sqrt(float re, float im, float& rr, float& ri) { + const float magnitude = num::sqrt_approx(re * re + im * im); + if (re >= 0.0f) { + rr = num::sqrt_approx(0.5f * (magnitude + re)); + ri = rr > 0.0f ? __fdividef(im, 2.0f * rr) : 0.0f; + } else { + const float root_imag = num::sqrt_approx(0.5f * (magnitude - re)); + rr = __fdividef(fabsf(im), 2.0f * root_imag); + ri = im < 0.0f ? -root_imag : root_imag; + } +} +__device__ __forceinline__ void complex_sqrt(num::Dual re, num::Dual im, num::Dual& rr, num::Dual& ri) { + float vr, vi; + complex_sqrt(re.v, im.v, vr, vi); + const float guard = 2.0f * (vr * vr + vi * vi); + float tr = 0.0f, ti = 0.0f; + if (guard > 0.0f) { + tr = __fdividef(re.d * vr + im.d * vi, guard); + ti = __fdividef(im.d * vr - re.d * vi, guard); + } + rr = {vr, tr}; + ri = {vi, ti}; +} + +template +__device__ __forceinline__ void complex_exp(T re, T im, T& er, T& ei) { + const T scale = exp_(re); + T s, c; + sincos_(im, s, c); + er = scale * c; + ei = scale * s; +} + +// expm((K - diag(R2) - 2 pi i diag(0, df)) dt) for two exchanging pools' +// transverse states, times the attenuation: four complex entries, x[0..7]. +template +__device__ __forceinline__ void transverse_step(T r2_free, T r2_bound, T exchange, T bound, T free, + T shift_hz, T dt, T attenuation, T* x) { + const T kab = exchange * bound, kba = exchange * free; + const T l11 = (-kab - r2_free) * dt, l12 = kba * dt, l21 = kab * dt; + const T l22 = (-kba - r2_bound) * dt; + const T l22_imag = -TWO_PI * shift_hz * dt; + const T trace_r = 0.5f * (l11 + l22), trace_i = 0.5f * l22_imag; + const T gap_r = 0.5f * (l11 - l22), gap_i = -0.5f * l22_imag; + const T square_r = gap_r * gap_r - gap_i * gap_i + l12 * l21; + const T square_i = 2.0f * gap_r * gap_i; + T root_r, root_i, upper_r, upper_i, lower_r, lower_i; + complex_sqrt(square_r, square_i, root_r, root_i); + complex_exp(trace_r + root_r, trace_i + root_i, upper_r, upper_i); + complex_exp(trace_r - root_r, trace_i - root_i, lower_r, lower_i); + const T cos_r = 0.5f * (upper_r + lower_r), cos_i = 0.5f * (upper_i + lower_i); + T scale_r, scale_i; + if (value(square_r) * value(square_r) + value(square_i) * value(square_i) > 1e-24f) { + const T half_r = 0.5f * (upper_r - lower_r), half_i = 0.5f * (upper_i - lower_i); + const T inverse = div_(T(1.0f), root_r * root_r + root_i * root_i); + scale_r = (half_r * root_r + half_i * root_i) * inverse; + scale_i = (half_i * root_r - half_r * root_i) * inverse; + } else { + T plain_r, plain_i; + complex_exp(trace_r, trace_i, plain_r, plain_i); + const T square2_r = square_r * square_r - square_i * square_i; + const T square2_i = 2.0f * square_r * square_i; + const T poly_r = 1.0f + square_r * (1.0f / 6.0f) + square2_r * (1.0f / 120.0f); + const T poly_i = square_i * (1.0f / 6.0f) + square2_i * (1.0f / 120.0f); + scale_r = plain_r * poly_r - plain_i * poly_i; + scale_i = plain_r * poly_i + plain_i * poly_r; + } + const T off_r = scale_r * gap_r - scale_i * gap_i; + const T off_i = scale_r * gap_i + scale_i * gap_r; + x[0] = attenuation * (cos_r + off_r); + x[1] = attenuation * (cos_i + off_i); + x[2] = attenuation * scale_r * l12; + x[3] = attenuation * scale_i * l12; + x[4] = attenuation * scale_r * l21; + x[5] = attenuation * scale_i * l21; + x[6] = attenuation * (cos_r - off_r); + x[7] = attenuation * (cos_i - off_i); +} + +// [a, b] exp from exponentials already taken; a series near coalescence. +template +__device__ __forceinline__ W exp_difference(W lower, W upper, W exp_lower, W exp_upper) { + const W half = 0.5 * (upper - lower); + if (fabs(value(half)) < 1e-4) { + const W square = half * half; + return exp_lower * (1.0 + half + 0.5 * square) * (1.0 + square / 6.0); + } + return (exp_upper - exp_lower) / (upper - lower); +} + +__host__ __device__ constexpr double inverse_factorial(int k) { + double f = 1.0; + for (int i = 2; i <= k; ++i) f *= i; + return 1.0 / f; +} + +// Terms K.. of the exponential's series reduced modulo the characteristic +// polynomial, each weight a constant. +template +__device__ __forceinline__ void series_terms(W determinant, W minors, W& flat, W& linear, W& square, + W& sum_flat, W& sum_linear, W& sum_square) { + if constexpr (K < TERMS) { + const W next_flat = square * determinant; + const W next_linear = flat - square * minors; + const W next_square = linear; + flat = next_flat; + linear = next_linear; + square = next_square; + using V = decltype(value(determinant)); + constexpr double weight = inverse_factorial(K); + sum_flat = sum_flat + V(weight) * flat; + sum_linear = sum_linear + V(weight) * linear; + sum_square = sum_square + V(weight) * square; + series_terms(determinant, minors, flat, linear, square, sum_flat, sum_linear, + sum_square); + } +} + +// The precision a three-pool interval is formed in. +template +struct WorkOf { + using type = typename std::conditional::type; +}; +template +struct WorkOf { + using type = typename std::conditional::type; +}; + +// expm((K - diag(R1)) dt) for free water (a), pool b and the semisolid pool +// (c), which exchange with a and not with each other, times the attenuation: +// e[0..8] row-major, and the recoveries grow[0..2]. NARROW: the series alone, +// in float; otherwise in double, the series where the roots are close and the +// Newton form at the three roots where they are not. +template +__device__ __forceinline__ void three_pool_step(T r1_free, T r1_b, T r1_c, T exchange_b, T exchange_c, + T fraction_b, T fraction_c, T dt, T attenuation, T* e, + T* grow) { + using W = typename WorkOf::type; + using V = decltype(value(W())); + constexpr int TERMS = NARROW ? 24 : 16; + const W step = W(dt); + const W free = W(1.0f - fraction_b - fraction_c); + const W pool_b = W(fraction_b), pool_c = W(fraction_c); + const W kab = W(exchange_b) * pool_b, kba = W(exchange_b) * free; + const W kac = W(exchange_c) * pool_c, kca = W(exchange_c) * free; + const W a00 = (-kab - kac - W(r1_free)) * step, a01 = kba * step, a02 = kca * step; + const W a10 = kab * step, a11 = (-kba - W(r1_b)) * step; + const W a20 = kac * step, a22 = (-kca - W(r1_c)) * step; + W third; + if constexpr (NARROW) third = (a00 + a11 + a22) * V(1.0 / 3.0); + else third = (a00 + a11 + a22) / V(3.0); + const W s00 = a00 - third, s11 = a11 - third, s22 = a22 - third; + const W minors = s00 * s11 - a01 * a10 + s00 * s22 - a02 * a20 + s11 * s22; + const W determinant = s00 * s11 * s22 - a01 * (a10 * s22) + a02 * (-s11 * a20); + W c[9]; + bool close = true; + if constexpr (!NARROW) close = -2.0 * value(minors) < 1.0; + if (close) { + W flat = V(1), linear = V(0), square = V(0), sum_flat = V(1), sum_linear = V(0), sum_square = V(0); + series_terms(determinant, minors, flat, linear, square, sum_flat, sum_linear, + sum_square); + const W q00 = s00 * s00 + a01 * a10 + a02 * a20, q01 = s00 * a01 + a01 * s11; + const W q02 = s00 * a02 + a02 * s22, q10 = a10 * s00 + s11 * a10; + const W q11 = a10 * a01 + s11 * s11, q12 = a10 * a02; + const W q20 = a20 * s00 + s22 * a20, q21 = a20 * a01, q22 = a20 * a02 + s22 * s22; + const W lift = exp_(third); + c[0] = lift * (sum_flat + sum_linear * s00 + sum_square * q00); + c[1] = lift * (sum_linear * a01 + sum_square * q01); + c[2] = lift * (sum_linear * a02 + sum_square * q02); + c[3] = lift * (sum_linear * a10 + sum_square * q10); + c[4] = lift * (sum_flat + sum_linear * s11 + sum_square * q11); + c[5] = lift * (sum_square * q12); + c[6] = lift * (sum_linear * a20 + sum_square * q20); + c[7] = lift * (sum_square * q21); + c[8] = lift * (sum_flat + sum_linear * s22 + sum_square * q22); + } else { + if constexpr (!NARROW) { + const W radius = sqrt_(max_(-minors * (1.0 / 3.0), W(1e-300))); + const double limit = 1.0 - 1e-16; + const W argument = min_(max_(0.5 * determinant / (radius * radius * radius), W(-limit)), W(limit)); + const W angle = num::acos_(argument) / 3.0; + const double turn = 2.09439510239319549231; + const W root_a = 2.0 * radius * num::cos_(angle) + third; + const W root_b = 2.0 * radius * num::cos_(angle - turn) + third; + const W root_c = 2.0 * radius * num::cos_(angle - 2.0 * turn) + third; + const W low = min_(min_(root_a, root_b), root_c); + const W high = max_(max_(root_a, root_b), root_c); + const W middle = max_(min_(root_a, root_b), min_(max_(root_a, root_b), root_c)); + const W leading = exp_(low), centre = exp_(middle), trailing = exp_(high); + const W first = exp_difference(low, middle, leading, centre); + const W span = high - low; + const W second = + (exp_difference(middle, high, centre, trailing) - first) / (value(span) > 0.0 ? span : W(1.0)); + const W m00 = a00 - low, m11 = a11 - low, m22 = a22 - low; + const W n00 = a00 - middle, n11 = a11 - middle, n22 = a22 - middle; + const W p00 = m00 * n00 + a01 * a10 + a02 * a20, p01 = m00 * a01 + a01 * n11; + const W p02 = m00 * a02 + a02 * n22, p10 = a10 * n00 + m11 * a10; + const W p11 = a10 * a01 + m11 * n11, p12 = a10 * a02; + const W p20 = a20 * n00 + m22 * a20, p21 = a20 * a01; + const W p22 = a20 * a02 + m22 * n22; + c[0] = leading + first * m00 + second * p00; + c[1] = first * a01 + second * p01; + c[2] = first * a02 + second * p02; + c[3] = first * a10 + second * p10; + c[4] = leading + first * m11 + second * p11; + c[5] = second * p12; + c[6] = first * a20 + second * p20; + c[7] = second * p21; + c[8] = leading + first * m22 + second * p22; + } + } + const W damp = W(attenuation); + W d[9]; +#pragma unroll + for (int k = 0; k < 9; ++k) d[k] = damp * c[k]; + grow[0] = T(free - (d[0] * free + d[1] * pool_b + d[2] * pool_c)); + grow[1] = T(pool_b - (d[3] * free + d[4] * pool_b + d[5] * pool_c)); + grow[2] = T(pool_c - (d[6] * free + d[7] * pool_b + d[8] * pool_c)); +#pragma unroll + for (int k = 0; k < 9; ++k) e[k] = T(d[k]); +} + +// One interval's three-pool operator read from the table, undamped there. +// Along a direction the row carries the tissue's share and the interval's +// own is the generator times the operator. +template +__device__ __forceinline__ void three_pool_from_table(const float* table, int row, int atom, int voxels, + T dt, T attenuation, T r1_free, T r1_b, T r1_c, + T exchange_b, T exchange_c, T free, T pool_b, + T pool_c, T* e, T* grow) { + if constexpr (DUAL) { + const float* base = table + static_cast(row) * (18 * voxels) + atom; + float c[9]; +#pragma unroll + for (int k = 0; k < 9; ++k) c[k] = __ldg(base + k * voxels); + const float xb = exchange_b.v, xc = exchange_c.v, fb = pool_b.v, fc = pool_c.v; + const float fa = 1.0f - fb - fc; + const float a00 = -xb * fb - xc * fc - r1_free.v, a01 = xb * fa, a02 = xc * fa; + const float a10 = xb * fb, a11 = -xb * fa - r1_b.v, a20 = xc * fc, a22 = -xc * fa - r1_c.v; + const float g[9] = {a00 * c[0] + a01 * c[3] + a02 * c[6], a00 * c[1] + a01 * c[4] + a02 * c[7], + a00 * c[2] + a01 * c[5] + a02 * c[8], a10 * c[0] + a11 * c[3], + a10 * c[1] + a11 * c[4], a10 * c[2] + a11 * c[5], + a20 * c[0] + a22 * c[6], a20 * c[1] + a22 * c[7], + a20 * c[2] + a22 * c[8]}; +#pragma unroll + for (int k = 0; k < 9; ++k) { + e[k] = attenuation * T{c[k], __ldg(base + (9 + k) * voxels) + dt.d * g[k]}; + } + } else { + const float* base = table + static_cast(row) * (9 * voxels) + atom; +#pragma unroll + for (int k = 0; k < 9; ++k) e[k] = attenuation * __ldg(base + k * voxels); + } + grow[0] = free - (e[0] * free + e[1] * pool_b + e[2] * pool_c); + grow[1] = pool_b - (e[3] * free + e[4] * pool_b + e[5] * pool_c); + grow[2] = pool_c - (e[6] * free + e[7] * pool_b + e[8] * pool_c); +} + +// Cubic Hermite through a table of knots (value, slope) spaced ``step``, +// clamped at the far end. +template +__device__ __forceinline__ T hermite_weights(T scaled, float lower, float step, T& h10, T& h01, T& h11) { + const T u = scaled - lower, u2 = u * u, u3 = u2 * u; + h10 = (u3 - 2.0f * u2 + u) * step; + h01 = -2.0f * u3 + 3.0f * u2; + h11 = (u3 - u2) * step; + return 2.0f * u3 - 3.0f * u2 + 1.0f; +} + +// How well the semisolid pool absorbs a pulse this far off its resonance. +template +__device__ __forceinline__ T lineshape_at(const float* lineshape, T offset_hz, int bins, float step) { + const int last = bins - 1; + const T scaled = min_(div_(abs_(offset_hz), T(step)), T(last + 0.0f)); + const float lower = fminf(floorf(value(scaled)), last - 1.0f); + T h10, h01, h11; + const T h00 = hermite_weights(scaled, lower, step, h10, h01, h11); + const float* base = lineshape + static_cast(lower) * 2; + return h00 * __ldg(base) + h10 * __ldg(base + 1) + h01 * __ldg(base + 2) + h11 * __ldg(base + 3); +} + +// Relaxation over one interval for one combination of its switches. +template +__device__ __forceinline__ void relax(const Params& p, int event, T dt_shared, const bool* active, + const int* train, const int* atom, float order, int state, + const T* r1, const T* r2, const T* b0, T* fpr, T* fpi, T* fmr, + T* fmi, T* zr, T* zi) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + T dt = dt_shared; + if constexpr (!SINGLE) { + dt = active[y] ? read(p.duration, p.d_duration, train[y] * p.event_count + event) : T(0.0f); + } + const int at = p.atom_stride ? atom[y] : 0; + T wout = 1.0f, velocity = 0.0f; + if constexpr (MOVING) { + velocity = active[y] ? read(p.velocity, p.d_velocity, at) : T(0.0f); + wout = 1.0f - min_(abs_(velocity) * p.washout_scale * dt, T(1.0f)); + } + T e1 = exp_(-r1[y] * dt) * wout; + T e2 = exp_(-r2[y] * dt) * wout; + const T recovery = 1.0f - e1; + if constexpr (DIFFUSING) { + const T b = (active[y] ? read(p.diffusion, p.d_diffusion, at) : T(0.0f)) * dt; + const float sq = order * order; + e1 *= exp_(-b * sq); + e2 *= exp_(-b * (sq + order + 0.3333333333333333f)); + } + T b0y = 0.0f; + if constexpr (OFF_AXIS) b0y = MAPS ? b0[MAPS ? y : 0] : T(0.0f); + if constexpr (MOVING || OFF_AXIS) { + T turn = 0.0f; + if constexpr (MOVING) turn = velocity * p.flow_scale * dt; + T os, oc; + sincos_(-2.0f * 3.141592653589793f * b0y * dt - (order + 0.5f) * turn, os, oc); + T old = fpr[y]; + fpr[y] = e2 * (old * oc - fpi[y] * os); + fpi[y] = e2 * (old * os + fpi[y] * oc); + old = fmr[y]; + fmr[y] = e2 * (old * oc + fmi[y] * os); + fmi[y] = e2 * (-old * os + fmi[y] * oc); + if constexpr (MOVING) { + T ts, tc; + sincos_(-order * turn, ts, tc); + old = zr[y]; + zr[y] = e1 * (old * tc - zi[y] * ts); + zi[y] = e1 * (old * ts + zi[y] * tc); + } else { + zr[y] = e1 * zr[y]; + zi[y] = e1 * zi[y]; + } + } else { + fpr[y] = e2 * fpr[y]; + fpi[y] = e2 * fpi[y]; + fmr[y] = e2 * fmr[y]; + fmi[y] = e2 * fmi[y]; + zr[y] = e1 * zr[y]; + zi[y] = e1 * zi[y]; + } + if (state == 0) zr[y] += recovery; + } +} + +// Off-resonance alone over one train: the interval repeats, and with it the +// two factors and the precession, which are then taken once per repeat. +template +__device__ __forceinline__ void relax_off_axis_repeat(T dt, T& last_dt, int state, const T* r1, const T* r2, + const T* b0, T* e1c, T* e2c, T* occ, T* osc, T* fpr, + T* fpi, T* fmr, T* fmi, T* zr, T* zi) { + if (!same(dt, last_dt)) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + e1c[y] = exp_(-r1[y] * dt); + e2c[y] = exp_(-r2[y] * dt); + sincos_(-2.0f * 3.141592653589793f * b0[y] * dt, osc[y], occ[y]); + } + last_dt = dt; + } +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T e2 = e2c[y], oc = occ[y], os = osc[y]; + T old = fpr[y]; + fpr[y] = e2 * (old * oc - fpi[y] * os); + fpi[y] = e2 * (old * os + fpi[y] * oc); + old = fmr[y]; + fmr[y] = e2 * (old * oc + fmi[y] * os); + fmi[y] = e2 * (-old * os + fmi[y] * oc); + zr[y] = e1c[y] * zr[y]; + zi[y] = e1c[y] * zi[y]; + if (state == 0) zr[y] += 1.0f - e1c[y]; + } +} + +// What a thread's problems are made of beyond free water, read once. +template +struct PoolTissue { + static constexpr int YL = POOLS >= 1 ? Y : 1, YB = POOLS >= 2 ? Y : 1, YC = POOLS == 3 ? Y : 1; + T fraction_b[YL], exchange_b[YL], r1_b[YL], free[YL]; + T r2_b[YB], shift_b[YB]; + T fraction_c[YC], exchange_c[YC], r1_c[YC]; +}; + +// One interval's operators for each problem a thread holds, with the +// damping, precession and flow of the thread's own order folded in; kept +// while the interval repeats. Transverse: a complex factor (POOLS 1) or a +// complex 2x2 (POOLS 2, 3). Longitudinal: the exchange matrix, its +// recoveries, and the complex factor flow and damping put on it. +template +struct PoolOperators { + static constexpr int NT = POOLS >= 2 ? 8 : 2, NE = POOLS == 3 ? 9 : 4, NG = POOLS == 3 ? 3 : 2; + T t[Y][NT], e[Y][NE], grow[Y][NG], factor_r[Y], factor_i[Y]; +}; + +template +__device__ __forceinline__ void pool_operators(const Params& p, const T* dts, const int* rows, + const bool* active, const int* atom, int state, float order, + const T* r1, const T* r2, const T* b0, + const PoolTissue& tissue, + PoolOperators& ops) { + // A row of at least Y states spreads the three-pool step over its lanes + // where it is formed in double; in float the shuffles cost what they save. + const bool SPREAD = POOLS == 3 && MODE == ROOTS && Y > 1 && p.width >= Y; + T attenuations[Y], damps[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T dt = dts[y]; + const int at = p.atom_stride ? atom[y] : 0; + T wout = 1.0f, turn = 0.0f; + if constexpr (MOVING) { + const T velocity = active[y] ? read(p.velocity, p.d_velocity, at) : T(0.0f); + wout = 1.0f - min_(abs_(velocity) * p.washout_scale * dt, T(1.0f)); + turn = velocity * p.flow_scale * dt; + } + T damp_z = 1.0f, damp_t = 1.0f; + if constexpr (DIFFUSING) { + const T b = (active[y] ? read(p.diffusion, p.d_diffusion, at) : T(0.0f)) * dt; + const float sq = order * order; + damp_z = exp_(-b * sq); + damp_t = exp_(-b * (sq + order + 0.3333333333333333f)); + } + T oc = 1.0f, os = 0.0f; + if constexpr (MOVING || OFF_AXIS) { + T b0y = 0.0f; + if constexpr (OFF_AXIS && MAPS) b0y = b0[y]; + sincos_(-TWO_PI * b0y * dt - (order + 0.5f) * turn, os, oc); + } + if constexpr (POOLS == 1) { + const T e2 = exp_(-r2[y] * dt) * wout * damp_t; + ops.t[y][0] = e2 * oc; + ops.t[y][1] = e2 * os; + } else { + T x[8]; + transverse_step(r2[y], tissue.r2_b[y], tissue.exchange_b[y], tissue.fraction_b[y], tissue.free[y], + tissue.shift_b[y], dt, wout, x); + const T rr = damp_t * oc, ri = damp_t * os; +#pragma unroll + for (int k = 0; k < 4; ++k) { + ops.t[y][2 * k] = rr * x[2 * k] - ri * x[2 * k + 1]; + ops.t[y][2 * k + 1] = rr * x[2 * k + 1] + ri * x[2 * k]; + } + } + if constexpr (POOLS == 3) { + if constexpr (MODE == TABLE) { + three_pool_from_table(p.pool_table, rows[y], atom[y], p.atom_count, dt, wout, r1[y], + tissue.r1_b[y], tissue.r1_c[y], tissue.exchange_b[y], + tissue.exchange_c[y], tissue.free[y], tissue.fraction_b[y], + tissue.fraction_c[y], ops.e[y], ops.grow[y]); + } else if (!SPREAD) { + three_pool_step(r1[y], tissue.r1_b[y], tissue.r1_c[y], tissue.exchange_b[y], + tissue.exchange_c[y], tissue.fraction_b[y], + tissue.fraction_c[y], dt, wout, ops.e[y], ops.grow[y]); + } else { + attenuations[y] = wout; + } + } else { + two_pool_step(r1[y], tissue.r1_b[y], tissue.exchange_b[y], tissue.fraction_b[y], dt, wout, ops.e[y], + ops.grow[y]); + } + damps[y] = damp_z; + if constexpr (MOVING) { + T ts, tc; + sincos_(-order * turn, ts, tc); + ops.factor_r[y] = damp_z * tc; + ops.factor_i[y] = damp_z * ts; + } + } + if constexpr (POOLS == 3 && MODE == ROOTS && Y > 1) { + if (SPREAD) { + // The group's lanes share their problems, so each forms one + // problem's three-pool operator and the others read it: a + // Y-th of the work, which in double is most of the interval's. + const int mine = state % Y; + auto pick = [&](const T* values) { + T picked = values[0]; +#pragma unroll + for (int y = 1; y < Y; ++y) picked = mine == y ? values[y] : picked; + return picked; + }; + const T e_r1 = pick(r1), e_r1b = pick(tissue.r1_b), e_r1c = pick(tissue.r1_c); + const T e_xb = pick(tissue.exchange_b), e_xc = pick(tissue.exchange_c); + const T e_fb = pick(tissue.fraction_b), e_fc = pick(tissue.fraction_c); + const T e_dt = pick(dts), e_wout = pick(attenuations); + T e[9], grow[3]; + three_pool_step(e_r1, e_r1b, e_r1c, e_xb, e_xc, e_fb, e_fc, e_dt, e_wout, e, grow); + const unsigned mask = __activemask(); + const int first_lane = (threadIdx.x & 31) - state; +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int k = 0; k < 9; ++k) ops.e[y][k] = num::shfl_from(e[k], first_lane + y, mask); +#pragma unroll + for (int k = 0; k < 3; ++k) ops.grow[y][k] = num::shfl_from(grow[k], first_lane + y, mask); + } + } + } + if constexpr (!MOVING) { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int k = 0; k < PoolOperators::NE; ++k) ops.e[y][k] *= damps[y]; + } + } +} + +template +__device__ __forceinline__ void apply_pools(const PoolOperators& ops, int state, T* fpr, T* fpi, + T* fmr, T* fmi, T* zr, T* zi, T* bpr, T* bpi, T* bmr, T* bmi, + T* lr, T* li, T* cr, T* ci) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T* t = ops.t[y]; + if constexpr (POOLS == 1) { + const T pr = t[0] * fpr[y] - t[1] * fpi[y], pi = t[0] * fpi[y] + t[1] * fpr[y]; + const T mr = t[0] * fmr[y] + t[1] * fmi[y], mi = t[0] * fmi[y] - t[1] * fmr[y]; + fpr[y] = pr; fpi[y] = pi; fmr[y] = mr; fmi[y] = mi; + } else { + const T pr = t[0] * fpr[y] - t[1] * fpi[y] + t[2] * bpr[y] - t[3] * bpi[y]; + const T pi = t[0] * fpi[y] + t[1] * fpr[y] + t[2] * bpi[y] + t[3] * bpr[y]; + const T qr = t[4] * fpr[y] - t[5] * fpi[y] + t[6] * bpr[y] - t[7] * bpi[y]; + const T qi = t[4] * fpi[y] + t[5] * fpr[y] + t[6] * bpi[y] + t[7] * bpr[y]; + const T mr = t[0] * fmr[y] + t[1] * fmi[y] + t[2] * bmr[y] + t[3] * bmi[y]; + const T mi = t[0] * fmi[y] - t[1] * fmr[y] + t[2] * bmi[y] - t[3] * bmr[y]; + const T nr = t[4] * fmr[y] + t[5] * fmi[y] + t[6] * bmr[y] + t[7] * bmi[y]; + const T ni = t[4] * fmi[y] - t[5] * fmr[y] + t[6] * bmi[y] - t[7] * bmr[y]; + fpr[y] = pr; fpi[y] = pi; bpr[y] = qr; bpi[y] = qi; + fmr[y] = mr; fmi[y] = mi; bmr[y] = nr; bmi[y] = ni; + } + const T* e = ops.e[y]; + T ar, ai, br, bi, cr2 = 0.0f, ci2 = 0.0f; + if constexpr (POOLS == 3) { + ar = e[0] * zr[y] + e[1] * lr[y] + e[2] * cr[y]; + ai = e[0] * zi[y] + e[1] * li[y] + e[2] * ci[y]; + br = e[3] * zr[y] + e[4] * lr[y] + e[5] * cr[y]; + bi = e[3] * zi[y] + e[4] * li[y] + e[5] * ci[y]; + cr2 = e[6] * zr[y] + e[7] * lr[y] + e[8] * cr[y]; + ci2 = e[6] * zi[y] + e[7] * li[y] + e[8] * ci[y]; + } else { + ar = e[0] * zr[y] + e[1] * lr[y]; + ai = e[0] * zi[y] + e[1] * li[y]; + br = e[2] * zr[y] + e[3] * lr[y]; + bi = e[2] * zi[y] + e[3] * li[y]; + } + if constexpr (MOVING) { + const T fr = ops.factor_r[y], fi = ops.factor_i[y]; + zr[y] = fr * ar - fi * ai; zi[y] = fr * ai + fi * ar; + lr[y] = fr * br - fi * bi; li[y] = fr * bi + fi * br; + if constexpr (POOLS == 3) { + cr[y] = fr * cr2 - fi * ci2; + ci[y] = fr * ci2 + fi * cr2; + } + } else { + zr[y] = ar; zi[y] = ai; lr[y] = br; li[y] = bi; + if constexpr (POOLS == 3) { + cr[y] = cr2; + ci[y] = ci2; + } + } + if (state == 0) { + zr[y] += ops.grow[y][0]; + lr[y] += ops.grow[y][1]; + if constexpr (POOLS == 3) cr[y] += ops.grow[y][2]; + } + } +} + +// One pulse for one combination of its switches: the thread's train row or +// a row per problem, one shim or several, and whether every problem a thread +// holds sees the same rotation. +template +__device__ __forceinline__ void pulse(const Params& p, int event, int base, const bool* active, const int* train, + const int* atom, const int* location, const T* b1, const T* b1c, + const T* b1s, const T* b0, T* fpr, T* fpi, T* fmr, T* fmi, T* zr, T* zi, + T* bpr, T* bpi, T* bmr, T* bmi, T* lr, T* li, T* cr, T* ci) { + constexpr bool SATURATES = POOLS == 1 || POOLS == 3; + T flip_u = 0.0f, ce_u = 0.0f, se_u = 0.0f; + if constexpr (SINGLE) { + flip_u = read(p.flip, p.d_flip, base + event); + event_phase(p, base + event, ce_u, se_u); + } + int profile_row = 0; + if constexpr (RF == PROFILE) profile_row = p.profile_index[event] * p.locations; + int shim_row = 0; + if constexpr (SHIMMED) shim_row = p.shim_index[event] * p.atom_count; + float saturation = 0.0f, rf_frequency = 0.0f; + T absorbed_shared = 1.0f, shape_shared = 0.0f; + if constexpr (SATURATES) { + saturation = p.saturation[event]; + rf_frequency = p.rf_frequency[event]; + // With no off-resonance map every voxel sits at the pulse's own + // offset, and the line shape is read once. + if (!MAPS || !p.off_axis) { + shape_shared = lineshape_at(p.lineshape, T(rf_frequency), p.lineshape_bins, p.lineshape_step); + } + } + T sine = 0.0f, cosine = 1.0f, cphi = 1.0f, sphi = 0.0f, c2 = 1.0f, s2 = 0.0f; + if constexpr (SHARED) { + cphi = ce_u; + sphi = se_u; + if constexpr (RF == HARD) sincos_(flip_u, sine, cosine); + c2 = fma_(cphi, cphi, -(sphi * sphi)); + s2 = 2.0f * sphi * cphi; + if constexpr (SATURATES) absorbed_shared = exp_(saturation * flip_u * flip_u * shape_shared); + } +#pragma unroll + for (int y = 0; y < Y; ++y) { + T alpha; + if constexpr (SHARED) { + alpha = flip_u; + } else { + T flip_v = flip_u, ce = ce_u, se = se_u; + if constexpr (!SINGLE) { + const int e = active[y] ? train[y] * p.event_count + event : event; + flip_v = read(p.flip, p.d_flip, e); + event_phase(p, e, ce, se); + } + T tb1 = 1.0f, tc = 1.0f, ts = 0.0f; + if constexpr (MAPS) { + tb1 = b1[y]; + tc = b1c[y]; + ts = b1s[y]; + } + if constexpr (SHIMMED) { + const int cell = shim_row + atom[y]; + tb1 = p.transmit ? (active[y] ? read(p.b1, p.d_b1, cell) : T(1.0f)) : T(1.0f); + if (p.off_axis) sincos_(active[y] ? read(p.b1_phase, p.d_b1_phase, cell) : T(0.0f), ts, tc); + } + alpha = flip_v * tb1; + cphi = fma_(ce, tc, -(se * ts)); + sphi = fma_(se, tc, ce * ts); + if constexpr (RF == HARD) sincos_(alpha, sine, cosine); + c2 = fma_(cphi, cphi, -(sphi * sphi)); + s2 = 2.0f * sphi * cphi; + } + if constexpr (RF == HARD) { + rotate_flip_phase(cosine, sine, cphi, sphi, c2, s2, fpr[y], fpi[y], fmr[y], fmi[y], zr[y], zi[y]); + if constexpr (POOLS >= 2) { + rotate_flip_phase(cosine, sine, cphi, sphi, c2, s2, bpr[y], bpi[y], bmr[y], bmi[y], lr[y], li[y]); + } + } else { + T pair[4]; + if constexpr (RF == PROFILE) { + const int row = profile_row + location[y]; + const int last = p.profile_bins - 1; + const T scaled = min_(max_(div_(alpha, T(p.profile_step)), T(0.0f)), T(last + 0.0f)); + const float lower = fminf(floorf(value(scaled)), last - 1.0f); + T h10, h01, h11; + const T h00 = hermite_weights(scaled, lower, p.profile_step, h10, h01, h11); + const float* base_row = p.profile + (row * p.profile_bins + static_cast(lower)) * 8; +#pragma unroll + for (int c = 0; c < 4; ++c) { + pair[c] = h00 * __ldg(base_row + c) + h10 * __ldg(base_row + 4 + c) + + h01 * __ldg(base_row + 8 + c) + h11 * __ldg(base_row + 12 + c); + } + } else { + // Integrated per pulse per voxel, so the read is the pair. + const int row = p.pair_index[(active[y] ? train[y] * p.event_count : 0) + event]; + const long long cell = static_cast(row) * p.atom_count + atom[y]; + const float4 entry = __ldg(reinterpret_cast(p.pairs) + cell); + if constexpr (DUAL) { + const float4 direction = __ldg(reinterpret_cast(p.pair_direction) + cell); + pair[0] = T{entry.x, direction.x}; + pair[1] = T{entry.y, direction.y}; + pair[2] = T{entry.z, direction.z}; + pair[3] = T{entry.w, direction.w}; + } else { + pair[0] = entry.x; + pair[1] = entry.y; + pair[2] = entry.z; + pair[3] = entry.w; + } + } + const T turn_r = cphi, turn_i = -sphi; + const T sbr = pair[2] * turn_r - pair[3] * turn_i; + const T sbi = pair[2] * turn_i + pair[3] * turn_r; + rotate_spinor(pair[0], pair[1], sbr, sbi, fpr[y], fpi[y], fmr[y], fmi[y], zr[y], zi[y]); + if constexpr (POOLS >= 2) { + rotate_spinor(pair[0], pair[1], sbr, sbi, bpr[y], bpi[y], bmr[y], bmi[y], lr[y], li[y]); + } + } + if constexpr (SATURATES) { + // The semisolid pool absorbs the power the pulse deposits at the + // bare flip the transmit field gives the voxel. + T absorbed = absorbed_shared; + if constexpr (!SHARED) { + T shape = shape_shared; + if constexpr (MAPS) { + if (p.off_axis) { + shape = lineshape_at(p.lineshape, rf_frequency - b0[y], p.lineshape_bins, p.lineshape_step); + } + } + absorbed = exp_(saturation * alpha * alpha * shape); + } + if constexpr (POOLS == 1) { + lr[y] *= absorbed; + li[y] *= absorbed; + } else { + cr[y] *= absorbed; + ci[y] *= absorbed; + } + } + } +} +#define EPG_PULSE_ARGS \ + p, event, base, active, train, atom, location, b1, b1c, b1s, b0, fpr, fpi, fmr, fmi, zr, zi, bpr, bpi, \ + bmr, bmi, lr, li, cr, ci +#define EPG_RELAX_ARGS p, event, dt_shared, active, train, atom, order, state, r1, r2, b0, fpr, fpi, fmr, fmi, zr, zi +#define EPG_POOL_ARGS p, dts, rows, active, atom, state, order, r1, r2, b0, tissue, ops + +template +__device__ __forceinline__ void complex_loop(const Params& p) { + constexpr int YL = POOLS >= 1 ? Y : 1, YB = POOLS >= 2 ? Y : 1, YC = POOLS == 3 ? Y : 1; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const int first = ((blockIdx.x * (blockDim.x >> 5) + warp) * groups + group) * Y; + const float order = static_cast(state); + const int total = p.train_count * p.atom_count; + + int problem[Y], atom[Y], train[Y], location[Y]; + bool active[Y], live[Y]; + T r1[Y], r2[Y]; + // A mapped tissue holds its maps, and the transmit phase's rotation, in + // registers from the start; a uniform one holds none of them. + T m0[MAPS ? Y : 1], b1[MAPS ? Y : 1], b1c[MAPS ? Y : 1], b1s[MAPS ? Y : 1]; + T b0[MAPS ? Y : 1], inversion[MAPS ? Y : 1]; + T fpr[Y], fpi[Y], fmr[Y], fmi[Y], zr[Y], zi[Y]; + // Pool b's transverse and longitudinal states, or the semisolid pool's + // longitudinal ones where it is the only second pool; then the semisolid + // pool's beside pool b. + T bpr[YB], bpi[YB], bmr[YB], bmi[YB], lr[YL], li[YL], cr[YC], ci[YC]; + PoolTissue tissue; +#pragma unroll + for (int y = 0; y < Y; ++y) { + problem[y] = first + y; + active[y] = problem[y] < total; + live[y] = active[y] && state < p.state_count; + atom[y] = problem[y] % p.atom_count; + train[y] = problem[y] / p.atom_count; + // A voxel's place along the slice, which picks its row of a table. + location[y] = RF == PROFILE ? atom[y] % p.locations : 0; + r1[y] = num::rate(active[y] ? read(p.t1, p.d_t1, atom[y]) : T(1.0f)); + r2[y] = num::rate(active[y] ? read(p.t2, p.d_t2, atom[y]) : T(1.0f)); + const int at = p.atom_stride ? atom[y] : 0; + if constexpr (MAPS) { + m0[y] = p.density ? (active[y] ? read(p.m0, p.d_m0, at) : T(0.0f)) : T(1.0f); + b1[y] = p.transmit ? (active[y] ? read(p.b1, p.d_b1, at) : T(1.0f)) : T(1.0f); + const T b1_phase = p.off_axis ? (active[y] ? read(p.b1_phase, p.d_b1_phase, at) : T(0.0f)) : T(0.0f); + if constexpr (DUAL) { + sincos_(b1_phase, b1s[y], b1c[y]); + } else { + sincosf(b1_phase, &b1s[y], &b1c[y]); + } + b0[y] = p.off_axis ? (active[y] ? read(p.b0, p.d_b0, at) : T(0.0f)) : T(0.0f); + inversion[y] = p.inverting ? (active[y] ? read(p.inversion_efficiency, p.d_inversion_efficiency, at) + : T(1.0f)) + : T(1.0f); + } + T free = 1.0f; + if constexpr (POOLS >= 1) { + const float* fraction = POOLS == 1 ? p.bound_fraction : p.pool_b_fraction; + const float* d_fraction = POOLS == 1 ? p.d_bound_fraction : p.d_pool_b_fraction; + const float* exchange = POOLS == 1 ? p.bound_exchange : p.pool_b_exchange; + const float* d_exchange = POOLS == 1 ? p.d_bound_exchange : p.d_pool_b_exchange; + const float* t1_second = POOLS == 1 ? p.t1_bound : p.t1_pool_b; + const float* d_t1_second = POOLS == 1 ? p.d_t1_bound : p.d_t1_pool_b; + tissue.fraction_b[y] = active[y] ? read(fraction, d_fraction, at) : T(0.0f); + tissue.exchange_b[y] = active[y] ? read(exchange, d_exchange, at) : T(0.0f); + tissue.r1_b[y] = num::rate(active[y] ? read(t1_second, d_t1_second, at) : T(1.0f)); + free = 1.0f - tissue.fraction_b[y]; + if constexpr (POOLS >= 2) { + tissue.r2_b[y] = num::rate(active[y] ? read(p.t2_pool_b, p.d_t2_pool_b, at) : T(1.0f)); + tissue.shift_b[y] = active[y] ? read(p.pool_b_shift, p.d_pool_b_shift, at) : T(0.0f); + } + if constexpr (POOLS == 3) { + tissue.fraction_c[y] = active[y] ? read(p.bound_fraction, p.d_bound_fraction, at) : T(0.0f); + tissue.exchange_c[y] = active[y] ? read(p.bound_exchange, p.d_bound_exchange, at) : T(0.0f); + tissue.r1_c[y] = num::rate(active[y] ? read(p.t1_bound, p.d_t1_bound, at) : T(1.0f)); + free = 1.0f - tissue.fraction_b[y] - tissue.fraction_c[y]; + } + tissue.free[y] = free; + lr[y] = state == 0 ? tissue.fraction_b[y] : T(0.0f); + li[y] = 0.0f; + } + if constexpr (POOLS >= 2) bpr[y] = bpi[y] = bmr[y] = bmi[y] = 0.0f; + if constexpr (POOLS == 3) { + cr[y] = state == 0 ? tissue.fraction_c[y] : T(0.0f); + ci[y] = 0.0f; + } + fpr[y] = fpi[y] = fmr[y] = fmi[y] = zi[y] = 0.0f; + zr[y] = state == 0 ? free : T(0.0f); + } + // A thread's problems are consecutive, so they are almost always voxels + // of one train, and then an event's duration, flip and phase are one + // read for all of them: ``base`` is that train's row. + const bool uniform = ONE_TRAIN || (active[0] && train[0] == train[Y - 1]); + const int base = ONE_TRAIN ? 0 : train[0] * p.event_count; + // The pulse's flip and phase are then the same for every problem a + // thread holds unless a transmit field or its phase moves them, and the + // trigonometry is taken once. + const bool shared_pulse = RF != DYNAMIC && uniform && !p.transmit && !p.off_axis && !p.shimmed; + const bool plain_relax = !p.off_axis && !p.moving && !p.diffusing; + const int relax_code = (p.moving ? 4 : 0) | (p.diffusing ? 2 : 0) | (p.off_axis ? 1 : 0); + T e1c[Y], e2c[Y]; + T occ[MAPS ? Y : 1], osc[MAPS ? Y : 1]; + T last_dt = -1.0f; + int last_row = -1; + PoolOperators ops; + + auto shift_pair = [&](T* pr_, T* pi_, T* mr_, T* mi_) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T up_r = num::shfl_up(pr_[y], 1, width); + const T up_i = num::shfl_up(pi_[y], 1, width); + const T down_r = num::shfl_down(mr_[y], 1, width); + const T down_i = num::shfl_down(mi_[y], 1, width); + const bool keep_up = state > 0 && live[y]; + const bool keep_down = state + 1 < p.state_count && live[y]; + const T pr = keep_up ? up_r : T(0.0f), pi = keep_up ? up_i : T(0.0f); + const T mr = keep_down ? down_r : T(0.0f), mi = keep_down ? down_i : T(0.0f); + pr_[y] = state == 0 ? mr : pr; + pi_[y] = state == 0 ? -mi : pi; + mr_[y] = mr; + mi_[y] = mi; + } + }; + auto shift = [&]() { + shift_pair(fpr, fpi, fmr, fmi); + if constexpr (POOLS >= 2) shift_pair(bpr, bpi, bmr, bmi); + }; + + float held_r[Y], held_i[Y]; + int run_start = 0, run_len = 0; + auto flush = [&]() { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (state < run_len && active[y]) { + const long at_out = static_cast(problem[y]) * p.output_count + run_start + state; + p.output_real[at_out] = held_r[y]; + p.output_imag[at_out] = held_i[y]; + } + } + run_len = 0; + }; +#pragma unroll((MAPS || POOLS || DUAL) ? 1 : 2) + for (int event = 0; event < p.event_count; ++event) { + const T dt_shared = uniform ? read(p.duration, p.d_duration, base + event) : T(0.0f); + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + // Relaxation over the interval; an interval of no length, along no + // direction, leaves every state exactly where it is. Each switch + // branches around a whole block once per event, so the common case + // runs only its own terms. + if (!(uniform && same(dt_shared, T(0.0f)))) { + if constexpr (POOLS > 0) { + // The operators are a property of the interval, so one train + // whose interval repeats forms them once per repeat. + int row_shared = 0; + if constexpr (MODE == TABLE) row_shared = uniform ? p.duration_row[base + event] : 0; + const bool fresh = + !uniform || !same(dt_shared, last_dt) || (MODE == TABLE && row_shared != last_row); + if (fresh) { + T dts[Y]; + int rows[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const int e = active[y] ? train[y] * p.event_count + event : event; + dts[y] = uniform ? dt_shared : (active[y] ? read(p.duration, p.d_duration, e) : T(0.0f)); + rows[y] = 0; + if constexpr (MODE == TABLE) rows[y] = uniform ? row_shared : p.duration_row[e]; + } + switch (relax_code) { + case 0: pool_operators(EPG_POOL_ARGS); break; + case 1: pool_operators(EPG_POOL_ARGS); break; + case 2: pool_operators(EPG_POOL_ARGS); break; + case 3: pool_operators(EPG_POOL_ARGS); break; + case 4: pool_operators(EPG_POOL_ARGS); break; + case 5: pool_operators(EPG_POOL_ARGS); break; + case 6: pool_operators(EPG_POOL_ARGS); break; + default: pool_operators(EPG_POOL_ARGS); break; + } + last_dt = uniform ? dt_shared : T(-1.0f); + last_row = row_shared; + } + if (p.moving) { + apply_pools(ops, state, fpr, fpi, fmr, fmi, zr, zi, bpr, bpi, bmr, bmi, lr, + li, cr, ci); + } else { + apply_pools(ops, state, fpr, fpi, fmr, fmi, zr, zi, bpr, bpi, bmr, bmi, lr, + li, cr, ci); + } + } else if (plain_relax) { + // No off-resonance, flow or diffusion: two factors per + // problem, reused while the interval repeats. + if (!uniform || !same(dt_shared, last_dt)) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T dt = + uniform ? dt_shared + : (active[y] ? read(p.duration, p.d_duration, train[y] * p.event_count + event) + : T(0.0f)); + e1c[y] = exp_(-r1[y] * dt); + e2c[y] = exp_(-r2[y] * dt); + } + last_dt = uniform ? dt_shared : T(-1.0f); + } +#pragma unroll + for (int y = 0; y < Y; ++y) { + fpr[y] = e2c[y] * fpr[y]; + fpi[y] = e2c[y] * fpi[y]; + fmr[y] = e2c[y] * fmr[y]; + fmi[y] = e2c[y] * fmi[y]; + zr[y] = e1c[y] * zr[y]; + zi[y] = e1c[y] * zi[y]; + if (state == 0) zr[y] += 1.0f - e1c[y]; + } + } else { + // One variant per combination of the interval's switches, + // chosen once per event: each runs only its own terms. + if (uniform) { + switch (relax_code) { + case 1: + if constexpr (MAPS) { + relax_off_axis_repeat(dt_shared, last_dt, state, r1, r2, b0, e1c, e2c, occ, + osc, fpr, fpi, fmr, fmi, zr, zi); + } else { + relax(EPG_RELAX_ARGS); + } + break; + case 2: relax(EPG_RELAX_ARGS); break; + case 3: relax(EPG_RELAX_ARGS); break; + case 4: relax(EPG_RELAX_ARGS); break; + case 5: relax(EPG_RELAX_ARGS); break; + case 6: relax(EPG_RELAX_ARGS); break; + default: relax(EPG_RELAX_ARGS); break; + } + } else { + switch (relax_code) { + case 1: relax(EPG_RELAX_ARGS); break; + case 2: relax(EPG_RELAX_ARGS); break; + case 3: relax(EPG_RELAX_ARGS); break; + case 4: relax(EPG_RELAX_ARGS); break; + case 5: relax(EPG_RELAX_ARGS); break; + case 6: relax(EPG_RELAX_ARGS); break; + default: relax(EPG_RELAX_ARGS); break; + } + } + } + } + if (act & 1) shift(); + if (kind == 1 && (act & 4)) { + // Pool b is free water and turns over like any other; a semisolid + // pool is saturated instead. +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T eff = MAPS ? inversion[MAPS ? y : 0] : T(1.0f); + zr[y] = -eff * zr[y]; + zi[y] = -eff * zi[y]; + if constexpr (POOLS >= 2) { + lr[y] = -eff * lr[y]; + li[y] = -eff * li[y]; + } + } + } + if (kind == 1 && !(act & 4)) { + if (shared_pulse) { + if constexpr (RF != DYNAMIC) pulse(EPG_PULSE_ARGS); + } else if (p.shimmed) { + if (uniform) pulse(EPG_PULSE_ARGS); + else pulse(EPG_PULSE_ARGS); + } else { + if (uniform) pulse(EPG_PULSE_ARGS); + else pulse(EPG_PULSE_ARGS); + } + } + if ((act & 32) && kind == 2) { + const int out = p.output_index[event]; + if (out >= 0) { + // A readout is kept by the lane whose state is its slot in the + // run, and a run of consecutive outputs goes out at once. + if (run_len > 0 && (out != run_start + run_len || run_len == width)) flush(); + if (run_len == 0) run_start = out; + T ac, as; + if (uniform) event_phase(p, base + event, ac, as); +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (!uniform) event_phase(p, active[y] ? train[y] * p.event_count + event : event, ac, as); + // A coil sees the whole voxel: the sum over the pools. + T read_r = fpr[y], read_i = fpi[y]; + if constexpr (POOLS >= 2) { + read_r += bpr[y]; + read_i += bpi[y]; + } + const T r0 = num::shfl(read_r, 0, width); + const T i0 = num::shfl(read_i, 0, width); + const T m0y = MAPS ? m0[MAPS ? y : 0] : T(1.0f); + if (state == run_len) { + held_r[y] = stored(m0y * (r0 * ac + i0 * as)); + held_i[y] = stored(m0y * (i0 * ac - r0 * as)); + } + } + ++run_len; + } + } + if (act & 18) shift(); + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + fpr[y] = fpi[y] = fmr[y] = fmi[y] = 0.0f; + if constexpr (POOLS >= 2) bpr[y] = bpi[y] = bmr[y] = bmi[y] = 0.0f; + } + } + } + if (run_len > 0) flush(); +} + +} // namespace epg diff --git a/src/blochsim/_layout_numbers.hpp b/src/blochsim/_layout_numbers.hpp new file mode 100644 index 00000000..28076aaf --- /dev/null +++ b/src/blochsim/_layout_numbers.hpp @@ -0,0 +1,271 @@ +// The arithmetic the kernels are written in, for a plain number and for a +// dual number: a value and its derivative along one direction. A kernel +// written once over ``T`` is the forward simulation at ``float`` and its +// Jacobian-vector product at ``Dual``. +#pragma once + +#include + +namespace num { + +template +struct DualT { + F v, d; + DualT() = default; + __host__ __device__ constexpr DualT(F value) : v(value), d(F(0)) {} + __host__ __device__ constexpr DualT(F value, F tangent) : v(value), d(tangent) {} + template + __host__ __device__ explicit constexpr DualT(const DualT& other) : v(F(other.v)), d(F(other.d)) {} +}; +using Dual = DualT; +using Dual64 = DualT; + +template +struct is_dual { + static constexpr bool value = false; +}; +template +struct is_dual> { + static constexpr bool value = true; +}; + +__device__ __forceinline__ float value(float x) { return x; } +__device__ __forceinline__ double value(double x) { return x; } +template +__device__ __forceinline__ F value(DualT x) { + return x.v; +} +__device__ __forceinline__ float tangent(float) { return 0.0f; } +template +__device__ __forceinline__ F tangent(DualT x) { + return x.d; +} + +// Equal as a cache key: the value and, for a dual, the direction too. +__device__ __forceinline__ bool same(float a, float b) { return a == b; } +template +__device__ __forceinline__ bool same(DualT a, DualT b) { + return a.v == b.v && a.d == b.d; +} + +template +__device__ __forceinline__ DualT operator+(DualT a, DualT b) { + return {a.v + b.v, a.d + b.d}; +} +template +__device__ __forceinline__ DualT operator-(DualT a, DualT b) { + return {a.v - b.v, a.d - b.d}; +} +template +__device__ __forceinline__ DualT operator-(DualT a) { + return {-a.v, -a.d}; +} +template +__device__ __forceinline__ DualT operator*(DualT a, DualT b) { + return {a.v * b.v, a.d * b.v + a.v * b.d}; +} +template +__device__ __forceinline__ DualT operator/(DualT a, DualT b) { + const F q = a.v / b.v; + return {q, (a.d - q * b.d) / b.v}; +} +template +__device__ __forceinline__ DualT operator+(DualT a, F b) { + return {a.v + b, a.d}; +} +template +__device__ __forceinline__ DualT operator+(F a, DualT b) { + return {a + b.v, b.d}; +} +template +__device__ __forceinline__ DualT operator-(DualT a, F b) { + return {a.v - b, a.d}; +} +template +__device__ __forceinline__ DualT operator-(F a, DualT b) { + return {a - b.v, -b.d}; +} +template +__device__ __forceinline__ DualT operator*(DualT a, F b) { + return {a.v * b, a.d * b}; +} +template +__device__ __forceinline__ DualT operator*(F a, DualT b) { + return {a * b.v, a * b.d}; +} +template +__device__ __forceinline__ DualT operator/(DualT a, F b) { + return {a.v / b, a.d / b}; +} +template +__device__ __forceinline__ DualT& operator+=(DualT& a, DualT b) { + return a = a + b; +} +template +__device__ __forceinline__ DualT& operator-=(DualT& a, DualT b) { + return a = a - b; +} +template +__device__ __forceinline__ DualT& operator*=(DualT& a, DualT b) { + return a = a * b; +} +template +__device__ __forceinline__ DualT& operator*=(DualT& a, F b) { + return a = a * b; +} +template +__device__ __forceinline__ bool operator<(DualT a, F b) { + return a.v < b; +} +template +__device__ __forceinline__ bool operator>(DualT a, F b) { + return a.v > b; +} +template +__device__ __forceinline__ bool operator<(DualT a, DualT b) { + return a.v < b.v; +} +template +__device__ __forceinline__ bool operator>(DualT a, DualT b) { + return a.v > b.v; +} + +__device__ __forceinline__ float fma_(float a, float b, float c) { return fmaf(a, b, c); } +template +__device__ __forceinline__ DualT fma_(DualT a, DualT b, DualT c) { + return {fma(a.v, b.v, c.v), fma(a.d, b.v, fma(a.v, b.d, c.d))}; +} + +__device__ __forceinline__ float exp_(float x) { return __expf(x); } +__device__ __forceinline__ double exp_(double x) { return exp(x); } +template +__device__ __forceinline__ DualT exp_(DualT x) { + const F e = exp_(x.v); + return {e, e * x.d}; +} + +// Division the way Triton's float32 ``/`` divides: approximate. +__device__ __forceinline__ float div_(float a, float b) { return __fdividef(a, b); } +__device__ __forceinline__ double div_(double a, double b) { return a / b; } +template +__device__ __forceinline__ DualT div_(DualT a, DualT b) { + const F q = div_(a.v, b.v); + return {q, div_(a.d - q * b.d, b.v)}; +} + +__device__ __forceinline__ float sqrt_approx(float x) { + float r; + asm("sqrt.approx.f32 %0, %1;" : "=f"(r) : "f"(x)); + return r; +} +__device__ __forceinline__ float sqrt_(float x) { return sqrt_approx(x); } +__device__ __forceinline__ double sqrt_(double x) { return sqrt(x); } +template +__device__ __forceinline__ DualT sqrt_(DualT x) { + const F r = sqrt_(x.v); + return {r, div_(F(0.5) * x.d, r)}; +} + +__device__ __forceinline__ float min_(float a, float b) { return fminf(a, b); } +__device__ __forceinline__ float max_(float a, float b) { return fmaxf(a, b); } +__device__ __forceinline__ double min_(double a, double b) { return fmin(a, b); } +__device__ __forceinline__ double max_(double a, double b) { return fmax(a, b); } +template +__device__ __forceinline__ DualT min_(DualT a, DualT b) { + return b.v < a.v ? b : a; +} +template +__device__ __forceinline__ DualT max_(DualT a, DualT b) { + return b.v > a.v ? b : a; +} +__device__ __forceinline__ float abs_(float a) { return fabsf(a); } +__device__ __forceinline__ double abs_(double a) { return fabs(a); } +template +__device__ __forceinline__ DualT abs_(DualT a) { + return a.v < F(0) ? -a : a; +} + +__device__ __forceinline__ double acos_(double x) { return acos(x); } +__device__ __forceinline__ Dual64 acos_(Dual64 x) { + return {acos(x.v), -x.d / sqrt(1.0 - x.v * x.v)}; +} +__device__ __forceinline__ double cos_(double x) { return cos(x); } +__device__ __forceinline__ Dual64 cos_(Dual64 x) { + double s, c; + sincos(x.v, &s, &c); + return {c, -s * x.d}; +} + +// Triton's _sincos: one Cody-Waite reduction by a quarter turn, the Cephes +// polynomials either side of zero. +__device__ __forceinline__ void sincos_cw(float x, float& s, float& c) { + const float quarter = rintf(x * 0.6366197723675814f); + float r = fmaf(-quarter, 1.5703125f, x); + r = fmaf(-quarter, 4.837512969970703125e-4f, r); + r = fmaf(-quarter, 7.54978995489188216e-8f, r); + const float r2 = r * r; + float sine = fmaf(-1.9515295891e-4f, r2, 8.3321608736e-3f); + sine = fmaf(sine, r2, -1.6666654611e-1f); + sine = fmaf(r * r2, sine, r); + float cosine = fmaf(2.443315711809948e-5f, r2, -1.388731625493765e-3f); + cosine = fmaf(cosine, r2, 4.166664568298827e-2f); + cosine = fmaf(r2 * r2, cosine, fmaf(-0.5f, r2, 1.0f)); + const int q = static_cast(quarter) & 3; + s = q == 0 ? sine : (q == 1 ? cosine : (q == 2 ? -sine : -cosine)); + c = q == 0 ? cosine : (q == 1 ? -sine : (q == 2 ? -cosine : sine)); +} +__device__ __forceinline__ void sincos_(float x, float& s, float& c) { sincos_cw(x, s, c); } +__device__ __forceinline__ void sincos_(Dual x, Dual& s, Dual& c) { + float sv, cv; + sincos_cw(x.v, sv, cv); + s = {sv, cv * x.d}; + c = {cv, -sv * x.d}; +} + +__device__ __forceinline__ float shfl_up(float x, int delta, int width) { + return __shfl_up_sync(0xffffffffu, x, delta, width); +} +__device__ __forceinline__ Dual shfl_up(Dual x, int delta, int width) { + return {__shfl_up_sync(0xffffffffu, x.v, delta, width), __shfl_up_sync(0xffffffffu, x.d, delta, width)}; +} +__device__ __forceinline__ float shfl_down(float x, int delta, int width) { + return __shfl_down_sync(0xffffffffu, x, delta, width); +} +__device__ __forceinline__ Dual shfl_down(Dual x, int delta, int width) { + return {__shfl_down_sync(0xffffffffu, x.v, delta, width), + __shfl_down_sync(0xffffffffu, x.d, delta, width)}; +} +__device__ __forceinline__ float shfl(float x, int lane, int width) { + return __shfl_sync(0xffffffffu, x, lane, width); +} +__device__ __forceinline__ Dual shfl(Dual x, int lane, int width) { + return {__shfl_sync(0xffffffffu, x.v, lane, width), __shfl_sync(0xffffffffu, x.d, lane, width)}; +} + +// A value from an absolute lane, among the lanes ``mask`` names. +__device__ __forceinline__ float shfl_from(float x, int lane, unsigned mask) { + return __shfl_sync(mask, x, lane); +} +__device__ __forceinline__ Dual shfl_from(Dual x, int lane, unsigned mask) { + return {__shfl_sync(mask, x.v, lane), __shfl_sync(mask, x.d, lane)}; +} + +// A value read with its direction where the type carries one. +template +__device__ __forceinline__ T load(const float* values, const float* tangents, long long at) { + if constexpr (is_dual::value) { + return T{__ldg(values + at), __ldg(tangents + at)}; + } else { + return __ldg(values + at); + } +} + +// A rate in 1/s from a time constant in ms, divided exactly as the launch +// formed it. +__device__ __forceinline__ float rate(float ms) { return 1000.0f / ms; } +__device__ __forceinline__ Dual rate(Dual ms) { + const float r = 1000.0f / ms.v; + return {r, -1000.0f * ms.d / (ms.v * ms.v)}; +} + +} // namespace num diff --git a/src/blochsim/_layout_real.cu b/src/blochsim/_layout_real.cu new file mode 100644 index 00000000..ad48a785 --- /dev/null +++ b/src/blochsim/_layout_real.cu @@ -0,0 +1,45 @@ +// The real kernel and its Jacobian-vector product in their layouts. +#include "_layout.hpp" + +namespace blochsim_layout { +namespace { + +template +__global__ void __launch_bounds__(32 * WARPS) real_kernel(layout_real::Params p) { + layout_real::real_loop(p); +} + +// A dual number doubles what a thread holds; capped at a quarter of the +// register file the eight programs an SM then holds hide each other's +// latency better than four larger ones do. +template +__global__ void __launch_bounds__(32 * WARPS, 8) real_kernel_capped(layout_real::Params p) { + layout_real::real_loop(p); +} + +template +int launch_real(const layout_real::Params& p, cudaStream_t stream) { + const int groups = 32 / p.width; + const long long problems = static_cast(p.train_count) * p.atom_count; + const long long per_block = static_cast(WARPS) * groups * Y; + const unsigned grid = static_cast((problems + per_block - 1) / per_block); + const bool maps = p.transmit || p.density || p.inverting; +#define BLOCHSIM_GO(MP, ONE) \ + if constexpr (CAPPED) real_kernel_capped<<>>(p); \ + else real_kernel<<>>(p) + if (p.single_train) { + if (maps) { BLOCHSIM_GO(true, true); } else { BLOCHSIM_GO(false, true); } + } else { + if (maps) { BLOCHSIM_GO(true, false); } else { BLOCHSIM_GO(false, false); } + } +#undef BLOCHSIM_GO + return static_cast(cudaGetLastError()); +} + +} // namespace + +int real_forward(const layout_real::Params& p, cudaStream_t stream) { return launch_real(p, stream); } + +int real_jvp(const layout_real::Params& p, cudaStream_t stream) { return launch_real(p, stream); } + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_real.hpp b/src/blochsim/_layout_real.hpp new file mode 100644 index 00000000..73fbb16f --- /dev/null +++ b/src/blochsim/_layout_real.hpp @@ -0,0 +1,213 @@ +// The real EPG event loop over a number type: the forward simulation at +// float, its Jacobian-vector product at num::Dual. One warp per program; +// lanes along the states, groups of lanes across problems, Y problems of a +// group in each thread's registers. +// +// Layout (compile time): MAPS (per-voxel transmit, density or inversion +// efficiency in registers), Y problems per thread, ONE_TRAIN. Run time: +// atom_stride, shimmed, diffusing, transmit, density, inverting. +#pragma once +#include +#include "_layout_numbers.hpp" + +namespace layout_real { + +template +__device__ __forceinline__ float stored(T x) { + if constexpr (num::is_dual::value) return x.d; + else return x; +} + +struct Params { + const float *t1, *t2, *m0, *b1, *inversion_efficiency, *diffusion, *duration, *flip; + const float *d_t1, *d_t2, *d_m0, *d_b1, *d_inversion_efficiency, *d_diffusion, *d_duration, *d_flip; + const int *kind, *output_index, *shim_index; + const unsigned char* action; + float *output_real, *output_imag; + int atom_count, train_count, event_count, output_count, state_count, width; + bool single_train, atom_stride, shimmed, diffusing, transmit, density, inverting; +}; + +template +__device__ __forceinline__ void rotate(T c, T s, T& plus, T& minus, T& z) { + const T chs = 0.5f * (1.0f + c), shs = 0.5f * (1.0f - c), hs = 0.5f * s; + const T p = chs * plus + shs * minus - s * z; + const T m = shs * plus + chs * minus + s * z; + z = hs * plus - hs * minus + c * z; + plus = p; + minus = m; +} + +template +__device__ __forceinline__ void real_loop(const Params& p) { + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const int first = ((blockIdx.x * (blockDim.x >> 5) + warp) * groups + group) * Y; + const float order = static_cast(state); + const int total = p.train_count * p.atom_count; + + int problem[Y], atom[Y], train[Y]; + bool active[Y], live[Y]; + T r1[Y], r2[Y], m0[MAPS ? Y : 1], b1[MAPS ? Y : 1], inversion[MAPS ? Y : 1]; + // Along a direction the interval is rarely the last one again, so its + // factors are formed every event and diffusion is held for them. + constexpr bool HOLD = num::is_dual::value; + T diffusion[HOLD ? Y : 1]; + T plus[Y], minus[Y], z[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + problem[y] = first + y; + active[y] = problem[y] < total; + live[y] = active[y] && state < p.state_count; + atom[y] = problem[y] % p.atom_count; + train[y] = problem[y] / p.atom_count; + r1[y] = num::rate(active[y] ? num::load(p.t1, p.d_t1, atom[y]) : T(1.0f)); + r2[y] = num::rate(active[y] ? num::load(p.t2, p.d_t2, atom[y]) : T(1.0f)); + if constexpr (MAPS) { + const int at = p.atom_stride ? atom[y] : 0; + m0[y] = p.density ? (active[y] ? num::load(p.m0, p.d_m0, at) : T(0.0f)) : T(1.0f); + b1[y] = p.transmit ? (active[y] ? num::load(p.b1, p.d_b1, at) : T(1.0f)) : T(1.0f); + inversion[y] = p.inverting ? (active[y] ? num::load(p.inversion_efficiency, p.d_inversion_efficiency, at) : T(1.0f)) : T(1.0f); + } + if constexpr (HOLD) { + diffusion[y] = p.diffusing && active[y] ? num::load(p.diffusion, p.d_diffusion, p.atom_stride ? atom[y] : 0) : T(0.0f); + } + plus[y] = minus[y] = 0.0f; + z[y] = state == 0 ? 1.0f : 0.0f; + } + // A thread's problems are almost always voxels of one train; then an + // event's duration and flip are one read for all of them. + const bool uniform = ONE_TRAIN || (active[0] && train[0] == train[Y - 1]); + const int base = ONE_TRAIN ? 0 : train[0] * p.event_count; + const bool shared_pulse = uniform && !p.transmit; + // The relaxation factors, with diffusion's damping of this thread's + // order folded in, kept while the interval repeats. + T e1c[Y], e2c[Y], recovery[Y]; + T last_dt = -1.0f; + + auto shift = [&]() { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T up = num::shfl_up(plus[y], 1, width); + const T down = num::shfl_down(minus[y], 1, width); + const T shifted_plus = (state > 0 && live[y]) ? up : T(0.0f); + const T shifted_minus = (state + 1 < p.state_count && live[y]) ? down : T(0.0f); + plus[y] = state == 0 ? -shifted_minus : shifted_plus; + minus[y] = shifted_minus; + } + }; + auto factors = [&](int y, T dt) { + T e1 = num::exp_(-r1[y] * dt), e2 = num::exp_(-r2[y] * dt); + recovery[y] = 1.0f - e1; + if (p.diffusing) { + T held_diffusion; + if constexpr (HOLD) { + held_diffusion = diffusion[y]; + } else { + const int at = p.atom_stride ? atom[y] : 0; + held_diffusion = active[y] ? num::load(p.diffusion, p.d_diffusion, at) : T(0.0f); + } + const T b = held_diffusion * dt; + const float sq = order * order; + e1 *= num::exp_(-b * sq); + e2 *= num::exp_(-b * (sq + order + 0.3333333333333333f)); + } + e1c[y] = e1; + e2c[y] = e2; + }; + + float held[Y]; + int run_start = 0, run_len = 0; + auto flush = [&]() { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (state < run_len && active[y]) { + const long at_out = static_cast(problem[y]) * p.output_count + run_start + state; + p.output_real[at_out] = 0.0f; + p.output_imag[at_out] = held[y]; + } + } + run_len = 0; + }; +#pragma unroll(num::is_dual::value ? 1 : 2) + for (int event = 0; event < p.event_count; ++event) { + const T dt_shared = uniform ? num::load(p.duration, p.d_duration, base + event) : T(0.0f); + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + if (uniform) { + if (!num::same(dt_shared, T(0.0f))) { + if (!num::same(dt_shared, last_dt)) { +#pragma unroll + for (int y = 0; y < Y; ++y) factors(y, dt_shared); + last_dt = dt_shared; + } +#pragma unroll + for (int y = 0; y < Y; ++y) { + plus[y] *= e2c[y]; + minus[y] *= e2c[y]; + z[y] = z[y] * e1c[y] + (state == 0 ? recovery[y] : T(0.0f)); + } + } + } else { +#pragma unroll + for (int y = 0; y < Y; ++y) { + factors(y, active[y] ? num::load(p.duration, p.d_duration, train[y] * p.event_count + event) : T(0.0f)); + plus[y] *= e2c[y]; + minus[y] *= e2c[y]; + z[y] = z[y] * e1c[y] + (state == 0 ? recovery[y] : T(0.0f)); + } + last_dt = -1.0f; + } + if (act & 1) shift(); + if (kind == 1 && (act & 4)) { +#pragma unroll + for (int y = 0; y < Y; ++y) z[y] = -(MAPS ? inversion[MAPS ? y : 0] : T(1.0f)) * z[y]; + } else if (kind == 1) { + if (shared_pulse) { + T s, c; + num::sincos_(num::load(p.flip, p.d_flip, base + event), s, c); +#pragma unroll + for (int y = 0; y < Y; ++y) rotate(c, s, plus[y], minus[y], z[y]); + } else { + const T flip_u = uniform ? num::load(p.flip, p.d_flip, base + event) : T(0.0f); + const bool shim = p.shimmed && p.transmit; + const int shim_row = shim ? p.shim_index[event] * p.atom_count : 0; +#pragma unroll + for (int y = 0; y < Y; ++y) { + T alpha = uniform ? flip_u + : (active[y] ? num::load(p.flip, p.d_flip, train[y] * p.event_count + event) : T(0.0f)); + T pulse_b1 = MAPS ? b1[MAPS ? y : 0] : T(1.0f); + if (shim) pulse_b1 = active[y] ? num::load(p.b1, p.d_b1, shim_row + atom[y]) : T(1.0f); + T s, c; + num::sincos_(alpha * pulse_b1, s, c); + rotate(c, s, plus[y], minus[y], z[y]); + } + } + } + if ((act & 32) && kind == 2) { + const int out = p.output_index[event]; + if (out >= 0) { + if (run_len > 0 && (out != run_start + run_len || run_len == width)) flush(); + if (run_len == 0) run_start = out; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T first_state = num::shfl(plus[y], 0, width); + if (state == run_len) held[y] = stored((MAPS ? m0[MAPS ? y : 0] : T(1.0f)) * first_state); + } + ++run_len; + } + } + if (act & 18) shift(); + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) plus[y] = minus[y] = 0.0f; + } + } + if (run_len > 0) flush(); +} + +} // namespace layout_real diff --git a/tests/sequence/test_specialized_kernels.py b/tests/sequence/test_specialized_kernels.py index f283cbff..70fc2157 100644 --- a/tests/sequence/test_specialized_kernels.py +++ b/tests/sequence/test_specialized_kernels.py @@ -1,7 +1,7 @@ -"""Whether a kernel compiled for its switches computes what the general one does. +"""Whether a kernel compiled for its switches or its layout computes what the general one does. -A specialized kernel that quietly did not run agrees perfectly, so every case -also asserts that one did. +A kernel that quietly did not run agrees perfectly, so every case also asserts +that one did. """ from __future__ import annotations @@ -75,6 +75,10 @@ def _gradient(phase: float) -> torch.Tensor: return tissue["t2_ms"].grad +def _fast_launches() -> int: + return _gpu_launch.specialized_launches() + _gpu_launch.layout_launches() + + CASES = { "real forward": lambda: _forward(torch.pi / 2), "complex forward": lambda: _forward(0.0), @@ -88,18 +92,18 @@ def _gradient(phase: float) -> torch.Tensor: def test_a_specialized_kernel_computes_what_the_general_one_does(case) -> None: with _gpu_launch.generic_kernels(): general = CASES[case]() - before = _gpu_launch.specialized_launches() + before = _fast_launches() special = CASES[case]() - assert _gpu_launch.specialized_launches() > before + assert _fast_launches() > before error = (special - general).abs().max() scale = general.abs().max() assert float(error / scale) < 1e-5, f"{float(error):.3e} against {float(scale):.3e}" def test_the_general_kernels_run_where_asked() -> None: - before = _gpu_launch.specialized_launches() + before = _fast_launches() with _gpu_launch.generic_kernels(): _forward(torch.pi / 2) - assert _gpu_launch.specialized_launches() == before + assert _fast_launches() == before From bf22e0f172fd793de447012faba4daa1a26103e6 Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 03:35:31 +0200 Subject: [PATCH 09/16] Run the EPG adjoints and their second-order sweeps written for their layouts The complex and real vector-Jacobian products are written once over a number type, so their derivative along a direction is the same source at a dual number. The forward sweep keeps the state every four events and the reverse sweep replays each stretch into shared memory from its checkpoint, which replaces a trajectory of every event in global memory. An interval's tissue gradient sums, per lane, the products its operator met while the interval repeats and contracts them against the operator's derivatives -- along every tissue input at once -- only when the interval changes. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 12 +- CMakeLists.txt | 15 +- src/blochsim/_layout.cu | 202 +++- src/blochsim/_layout.hpp | 16 + src/blochsim/_layout_complex.hpp | 45 +- src/blochsim/_layout_complex_vjp.cu.in | 50 + src/blochsim/_layout_complex_vjp.hpp | 1293 ++++++++++++++++++++++++ src/blochsim/_layout_numbers.hpp | 245 ++++- src/blochsim/_layout_real_vjp.cu | 44 + src/blochsim/_layout_real_vjp.hpp | 423 ++++++++ 10 files changed, 2303 insertions(+), 42 deletions(-) create mode 100644 src/blochsim/_layout_complex_vjp.cu.in create mode 100644 src/blochsim/_layout_complex_vjp.hpp create mode 100644 src/blochsim/_layout_real_vjp.cu create mode 100644 src/blochsim/_layout_real_vjp.hpp diff --git a/CLAUDE.md b/CLAUDE.md index f05b1a47..fa6342b5 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -92,16 +92,20 @@ the general kernels alone. A specialized kernel is compiled for rows of at most 32 state orders, so its shifts are shuffles with no test; a wider launch runs the general one. -**The forward EPG kernels and their Jacobian-vector products are written for -their layouts** (`_layout.hpp`), and a launch whose rows fit a warp runs them -ahead of any tile kernel. A layout is what is compiled: the pools, how a +**The EPG kernels are written for their layouts** (`_layout.hpp`), and a +launch whose rows fit a warp runs them ahead of any tile kernel. A layout is what is compiled: the pools, how a pulse is formed, whether the tissue has per-voxel maps, the problems a thread holds, one train or several; every other switch is read at run time and steers whole blocks once per event, so one compile serves every combination of them. The loops are written once over a number type (`_layout_numbers.hpp`): at `float` they are the forward simulation, at `num::Dual` -- a value and its derivative along the direction -- the -Jacobian-vector product. `_gpu_launch.generic_kernels()` turns layouts off +Jacobian-vector product. The adjoints are written the same way, so their +derivative along a direction is the same source at `num::Dual`: the forward +sweep keeps the state every few events and the reverse sweep replays each +stretch from it, and an interval's gradient is contracted against its +operator's derivatives -- taken along every tissue input at once by +`num::Multi` -- only when the interval changes. `_gpu_launch.generic_kernels()` turns layouts off with the specializations, and `layout_launches()` counts them. They exist only on the card: `_gpu_host` compiles the tile kernels, so the host lane holds those, not these, to the C++ kernels. diff --git a/CMakeLists.txt b/CMakeLists.txt index 3723e973..f2b78f6f 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -153,7 +153,8 @@ if(BLOCHSIM_CUDA) # The EPG kernels written for their layouts (_layout.hpp): a unit per # kernel and pool layout, so they compile side by side. - list(APPEND _blochsim_gpu_sources src/blochsim/_layout.cu src/blochsim/_layout_real.cu) + list(APPEND _blochsim_gpu_sources src/blochsim/_layout.cu src/blochsim/_layout_real.cu + src/blochsim/_layout_real_vjp.cu) foreach(KIND IN ITEMS forward jvp) if(KIND STREQUAL "forward") set(TYPE "float") @@ -166,6 +167,18 @@ if(BLOCHSIM_CUDA) list(APPEND _blochsim_gpu_sources "${_source}") endforeach() endforeach() + foreach(KIND IN ITEMS vjp vjp_jvp) + if(KIND STREQUAL "vjp") + set(TYPE "float") + else() + set(TYPE "num::Dual") + endif() + foreach(POOLS RANGE 3) + set(_source "${CMAKE_CURRENT_BINARY_DIR}/layout/complex_${KIND}_${POOLS}.cu") + configure_file(src/blochsim/_layout_complex_vjp.cu.in "${_source}" @ONLY) + list(APPEND _blochsim_gpu_sources "${_source}") + endforeach() + endforeach() python_add_library(_gpu MODULE USE_SABI ${BLOCHSIM_ABI3_VERSION} diff --git a/src/blochsim/_layout.cu b/src/blochsim/_layout.cu index 6cb80c94..891ac2e9 100644 --- a/src/blochsim/_layout.cu +++ b/src/blochsim/_layout.cu @@ -27,7 +27,8 @@ struct Reader { bool flag(const char* name) const { return a[at(name)].i != 0; } }; -epg::Params complex_params(const Reader& r) { +// What the forward kernel and the adjoint read alike. +epg::Params complex_params_common(const Reader& r) { epg::Params p{}; p.t1 = r.floats("t1"); p.t2 = r.floats("t2"); @@ -64,6 +65,33 @@ epg::Params complex_params(const Reader& r) { p.pair_index = r.ints("pair_index"); p.duration_row = r.ints("duration_row"); p.action = static_cast(r.a[r.at("action")].p); + p.atom_count = static_cast(r.integer("atom_count")); + p.train_count = static_cast(r.integer("train_count")); + p.event_count = static_cast(r.integer("event_count")); + p.output_count = static_cast(r.integer("output_count")); + p.state_count = static_cast(r.integer("state_count")); + p.width = static_cast(r.integer("block_states")); + p.locations = static_cast(r.integer("locations")); + p.profile_bins = static_cast(r.integer("profile_bins")); + p.lineshape_bins = static_cast(r.integer("lineshape_bins")); + p.flow_scale = r.real("flow_scale"); + p.washout_scale = r.real("washout_scale"); + p.profile_step = r.real("profile_step"); + p.lineshape_step = r.real("lineshape_step"); + p.single_train = r.flag("single_train"); + p.atom_stride = r.flag("atom_stride"); + p.shimmed = r.flag("shimmed"); + p.off_axis = r.flag("off_axis"); + p.moving = r.flag("moving"); + p.diffusing = r.flag("diffusing"); + p.transmit = r.flag("transmit"); + p.density = r.flag("density"); + p.inverting = r.flag("inverting"); + return p; +} + +epg::Params complex_params(const Reader& r) { + epg::Params p = complex_params_common(r); p.d_t1 = r.floats("tangent_t1"); p.d_t2 = r.floats("tangent_t2"); p.d_m0 = r.floats("tangent_m0"); @@ -87,28 +115,6 @@ epg::Params complex_params(const Reader& r) { p.pair_direction = r.floats("pair_direction"); p.output_real = r.outputs("output_real"); p.output_imag = r.outputs("output_imag"); - p.atom_count = static_cast(r.integer("atom_count")); - p.train_count = static_cast(r.integer("train_count")); - p.event_count = static_cast(r.integer("event_count")); - p.output_count = static_cast(r.integer("output_count")); - p.state_count = static_cast(r.integer("state_count")); - p.width = static_cast(r.integer("block_states")); - p.locations = static_cast(r.integer("locations")); - p.profile_bins = static_cast(r.integer("profile_bins")); - p.lineshape_bins = static_cast(r.integer("lineshape_bins")); - p.flow_scale = r.real("flow_scale"); - p.washout_scale = r.real("washout_scale"); - p.profile_step = r.real("profile_step"); - p.lineshape_step = r.real("lineshape_step"); - p.single_train = r.flag("single_train"); - p.atom_stride = r.flag("atom_stride"); - p.shimmed = r.flag("shimmed"); - p.off_axis = r.flag("off_axis"); - p.moving = r.flag("moving"); - p.diffusing = r.flag("diffusing"); - p.transmit = r.flag("transmit"); - p.density = r.flag("density"); - p.inverting = r.flag("inverting"); return p; } @@ -169,7 +175,147 @@ int complex_launch(bool jvp, const Reader& r, cudaStream_t stream) { } } +// The complex adjoint's arguments, for one sweep or its derivative. +epg_vjp::Params adjoint_params(bool dual, const Reader& r) { + epg_vjp::Params v{}; + epg::Params& p = v.f; + p = complex_params_common(r); + p.d_t1 = r.floats("dot_t1"); + p.d_t2 = r.floats("dot_t2"); + p.d_m0 = r.floats("dot_m0"); + p.d_b1 = r.floats("dot_b1"); + p.d_b1_phase = r.floats("dot_b1_phase"); + p.d_b0 = r.floats("dot_b0"); + p.d_inversion_efficiency = r.floats("dot_inversion_efficiency"); + p.d_diffusion = r.floats("dot_diffusion"); + p.d_velocity = r.floats("dot_velocity"); + p.d_bound_fraction = r.floats("dot_bound_fraction"); + p.d_bound_exchange = r.floats("dot_exchange_rate"); + p.d_t1_bound = r.floats("dot_t1_bound"); + p.d_pool_b_fraction = r.floats("dot_pool_b_fraction"); + p.d_pool_b_exchange = r.floats("dot_pool_b_exchange"); + p.d_t1_pool_b = r.floats("dot_t1_pool_b"); + p.d_t2_pool_b = r.floats("dot_t2_pool_b"); + p.d_pool_b_shift = r.floats("dot_pool_b_shift"); + p.d_duration = r.floats("dot_duration"); + p.d_flip = r.floats("dot_flip"); + p.d_phase = r.floats("dot_phase"); + // A pair with no direction of its own along this sweep. + p.pair_direction = r.at("directed") >= 0 && r.flag("directed") ? r.floats("pair_direction") : nullptr; + v.grad_output_real = r.floats("grad_output_real"); + v.grad_output_imag = r.floats("grad_output_imag"); + auto out = [&](const char* name) { return r.at(name) < 0 ? nullptr : r.outputs(name); }; + if (dual) { + v.grad_tissue = out("grad_tissue_value"); + v.grad_tissue_t = out("grad_tissue_tangent"); + v.grad_flip = out("grad_flip_value"); + v.grad_flip_t = out("grad_flip_tangent"); + v.grad_phase = out("grad_phase_value"); + v.grad_phase_t = out("grad_phase_tangent"); + v.grad_duration = out("grad_duration_value"); + v.grad_duration_t = out("grad_duration_tangent"); + v.grad_pair = out("grad_pair_value"); + v.grad_pair_t = out("grad_pair_tangent"); + v.trajectory_r = out("trajectory_vr"); + v.trajectory_i = out("trajectory_vi"); + v.trajectory_tr = out("trajectory_tr"); + v.trajectory_ti = out("trajectory_ti"); + } else { + v.grad_tissue = out("grad_tissue"); + v.grad_flip = out("grad_flip"); + v.grad_phase = out("grad_phase"); + v.grad_duration = out("grad_duration"); + v.grad_pair = out("grad_pair"); + v.trajectory_r = out("trajectory_r"); + v.trajectory_i = out("trajectory_i"); + } + v.problem_base = static_cast(r.integer("problem_base")); + v.problem_end = static_cast(r.integer("problem_end")); + v.shim_rows = static_cast(r.integer("shim_rows")); + v.mode = r.flag("tabulated") ? epg::TABLE : (r.flag("narrow") ? epg::NARROW : epg::ROOTS); + return v; +} + +int adjoint_launch(bool dual, const Reader& r, cudaStream_t stream) { + // One launch walks both ways; the recording launch has nothing to do. + if (r.flag("recording")) return cudaSuccess; + const epg_vjp::Params v = adjoint_params(dual, r); + const int rf = r.flag("dynamic") ? epg::DYNAMIC : (r.flag("profiled") ? epg::PROFILE : epg::HARD); + switch (static_cast(r.integer("pools")) + (dual ? 4 : 0)) { + case 0: return complex_vjp_0(v, rf, stream); + case 1: return complex_vjp_1(v, rf, stream); + case 2: return complex_vjp_2(v, rf, stream); + case 3: return complex_vjp_3(v, rf, stream); + case 4: return complex_vjp_jvp_0(v, rf, stream); + case 5: return complex_vjp_jvp_1(v, rf, stream); + case 6: return complex_vjp_jvp_2(v, rf, stream); + case 7: return complex_vjp_jvp_3(v, rf, stream); + default: return -1; + } +} + +layout_real_vjp::Params real_adjoint_params(bool dual, const Reader& r) { + layout_real_vjp::Params p{}; + p.t1 = r.floats("t1"); + p.t2 = r.floats("t2"); + p.m0 = r.floats("m0"); + p.b1 = r.floats("b1"); + p.inversion_efficiency = r.floats("inversion_efficiency"); + p.diffusion = r.floats("diffusion"); + p.duration = r.floats("duration"); + p.flip = r.floats("flip"); + p.d_t1 = r.floats("dot_t1"); + p.d_t2 = r.floats("dot_t2"); + p.d_m0 = r.floats("dot_m0"); + p.d_b1 = r.floats("dot_b1"); + p.d_inversion_efficiency = r.floats("dot_inversion_efficiency"); + p.d_diffusion = r.floats("dot_diffusion"); + p.d_duration = r.floats("dot_duration"); + p.d_flip = r.floats("dot_flip"); + p.kind = r.ints("kind"); + p.output_index = r.ints("output_index"); + p.shim_index = r.ints("shim_index"); + p.action = static_cast(r.a[r.at("action")].p); + p.grad_output_imag = r.floats("grad_output_imag"); + if (dual) { + p.grad_tissue = r.outputs("grad_tissue_value"); + p.grad_tissue_t = r.outputs("grad_tissue_tangent"); + p.grad_flip = r.outputs("grad_flip_value"); + p.grad_flip_t = r.outputs("grad_flip_tangent"); + p.grad_duration = r.outputs("grad_duration_value"); + p.grad_duration_t = r.outputs("grad_duration_tangent"); + p.trajectory = r.outputs("trajectory_value"); + p.trajectory_t = r.outputs("trajectory_tangent"); + } else { + p.grad_tissue = r.outputs("grad_tissue"); + p.grad_flip = r.outputs("grad_flip"); + p.grad_duration = r.outputs("grad_duration"); + p.trajectory = r.outputs("trajectory_value"); + } + p.problem_base = static_cast(r.integer("problem_base")); + p.problem_end = static_cast(r.integer("problem_end")); + p.atom_count = static_cast(r.integer("atom_count")); + p.train_count = static_cast(r.integer("train_count")); + p.event_count = static_cast(r.integer("event_count")); + p.output_count = static_cast(r.integer("output_count")); + p.state_count = static_cast(r.integer("state_count")); + p.width = static_cast(r.integer("block_states")); + p.shim_rows = static_cast(r.integer("shim_rows")); + p.single_train = r.flag("single_train"); + p.atom_stride = r.flag("atom_stride"); + p.shimmed = r.flag("shimmed"); + p.diffusing = r.flag("diffusing"); + p.transmit = r.flag("transmit"); + p.density = r.flag("density"); + p.inverting = r.flag("inverting"); + return p; +} + const int COMPLEX = bsk::kernel_index("_epg_kernel"); +const int COMPLEX_VJP = bsk::kernel_index("_epg_vjp_kernel"); +const int COMPLEX_VJP_JVP = bsk::kernel_index("_epg_vjp_jvp_kernel"); +const int REAL_VJP = bsk::kernel_index("_epg_real_vjp_kernel"); +const int REAL_VJP_JVP = bsk::kernel_index("_epg_real_vjp_jvp_kernel"); const int COMPLEX_JVP = bsk::kernel_index("_epg_jvp_kernel"); const int REAL = bsk::kernel_index("_epg_real_kernel"); const int REAL_JVP = bsk::kernel_index("_epg_real_jvp_kernel"); @@ -178,7 +324,8 @@ const int REAL_JVP = bsk::kernel_index("_epg_real_jvp_kernel"); int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream) { const Reader r{kernel, arguments.a}; - if (kernel != COMPLEX && kernel != COMPLEX_JVP && kernel != REAL && kernel != REAL_JVP) { + const bool adjoint = kernel == COMPLEX_VJP || kernel == COMPLEX_VJP_JVP || kernel == REAL_VJP || kernel == REAL_VJP_JVP; + if (kernel != COMPLEX && kernel != COMPLEX_JVP && kernel != REAL && kernel != REAL_JVP && !adjoint) { return -1; } // A layout's rows are at most a warp wide. @@ -188,6 +335,13 @@ int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream) { if (kernel == COMPLEX || kernel == COMPLEX_JVP) { return complex_launch(kernel == COMPLEX_JVP, r, stream); } + if (kernel == COMPLEX_VJP || kernel == COMPLEX_VJP_JVP) { + return adjoint_launch(kernel == COMPLEX_VJP_JVP, r, stream); + } + if (kernel == REAL_VJP || kernel == REAL_VJP_JVP) { + const layout_real_vjp::Params p = real_adjoint_params(kernel == REAL_VJP_JVP, r); + return kernel == REAL_VJP ? real_vjp(p, stream) : real_vjp_jvp(p, stream); + } const layout_real::Params p = real_params(r); return kernel == REAL ? real_forward(p, stream) : real_jvp(p, stream); } diff --git a/src/blochsim/_layout.hpp b/src/blochsim/_layout.hpp index b998f560..519a9e28 100644 --- a/src/blochsim/_layout.hpp +++ b/src/blochsim/_layout.hpp @@ -13,7 +13,9 @@ #include "_kernels.hpp" #include "_layout_complex.hpp" +#include "_layout_complex_vjp.hpp" #include "_layout_real.hpp" +#include "_layout_real_vjp.hpp" namespace blochsim_layout { @@ -33,6 +35,20 @@ BLOCHSIM_LAYOUT_COMPLEX(BLOCHSIM_LAYOUT_COMPLEX_DECLARATION) int real_forward(const layout_real::Params& p, cudaStream_t stream); int real_jvp(const layout_real::Params& p, cudaStream_t stream); +// The adjoints: the forward sweep keeps a checkpoint every few events and +// the reverse sweep replays each stretch from it, so a launch is one sweep +// each way and the recording launch of the tile kernels' protocol does +// nothing. +#define BLOCHSIM_LAYOUT_ADJOINT(X) \ + X(vjp, 0) X(vjp, 1) X(vjp, 2) X(vjp, 3) X(vjp_jvp, 0) X(vjp_jvp, 1) X(vjp_jvp, 2) X(vjp_jvp, 3) +#define BLOCHSIM_LAYOUT_ADJOINT_DECLARATION(kind, pools) \ + int complex_##kind##_##pools(const epg_vjp::Params& v, int rf, cudaStream_t stream); +BLOCHSIM_LAYOUT_ADJOINT(BLOCHSIM_LAYOUT_ADJOINT_DECLARATION) +#undef BLOCHSIM_LAYOUT_ADJOINT_DECLARATION + +int real_vjp(const layout_real_vjp::Params& p, cudaStream_t stream); +int real_vjp_jvp(const layout_real_vjp::Params& p, cudaStream_t stream); + // Programs of two warps each. constexpr int WARPS = 2; diff --git a/src/blochsim/_layout_complex.hpp b/src/blochsim/_layout_complex.hpp index 7af7ae88..277b1bc0 100644 --- a/src/blochsim/_layout_complex.hpp +++ b/src/blochsim/_layout_complex.hpp @@ -28,6 +28,7 @@ using num::min_; using num::same; using num::sincos_; using num::sqrt_; +using num::primal; using num::value; constexpr float TWO_PI = 6.283185307179586f; @@ -72,9 +73,11 @@ template __device__ __forceinline__ void event_phase(const Params& p, int at, T& c, T& s) { if constexpr (DUAL) { sincos_(T{__ldg(p.phase + at), __ldg(p.d_phase + at)}, s, c); - } else { + } else if (p.phase_cos != nullptr) { c = p.phase_cos[at]; s = p.phase_sin[at]; + } else { + sincos_(__ldg(p.phase + at), s, c); } } @@ -150,8 +153,8 @@ __device__ __forceinline__ void two_pool_step(T r1_free, T r1_bound, T exchange, const T l22 = (-kba - r1_bound) * dt; const T half_trace = 0.5f * (l11 + l22), half_gap = 0.5f * (l11 - l22); const T square = half_gap * half_gap + l12 * l21; - const bool turning = value(square) > 1e-12f; - const T root = turning ? sqrt_(square) : T(num::sqrt_approx(fmaxf(value(square), 0.0f))); + const bool turning = primal(square) > 1e-12f; + const T root = turning ? sqrt_(square) : T(num::sqrt_approx(fmaxf(primal(square), 0.0f))); const T upper = exp_(half_trace + root), lower = exp_(half_trace - root); const T cosine = 0.5f * (upper + lower); const T scale = turning ? div_(0.5f * (upper - lower), root) @@ -193,6 +196,22 @@ __device__ __forceinline__ void complex_sqrt(num::Dual re, num::Dual im, num::Du ri = {vi, ti}; } +template +__device__ __forceinline__ void complex_sqrt(const num::Multi& re, const num::Multi& im, num::Multi& rr, + num::Multi& ri) { + S vr, vi; + complex_sqrt(re.v, im.v, vr, vi); + const S guard = 2.0f * (vr * vr + vi * vi); + const bool live = primal(guard) > 0.0f; + rr.v = vr; + ri.v = vi; +#pragma unroll + for (int k = 0; k < K; ++k) { + rr.d[k] = live ? div_(re.d[k] * vr + im.d[k] * vi, guard) : S(0.0f); + ri.d[k] = live ? div_(im.d[k] * vr - re.d[k] * vi, guard) : S(0.0f); + } +} + template __device__ __forceinline__ void complex_exp(T re, T im, T& er, T& ei) { const T scale = exp_(re); @@ -221,7 +240,7 @@ __device__ __forceinline__ void transverse_step(T r2_free, T r2_bound, T exchang complex_exp(trace_r - root_r, trace_i - root_i, lower_r, lower_i); const T cos_r = 0.5f * (upper_r + lower_r), cos_i = 0.5f * (upper_i + lower_i); T scale_r, scale_i; - if (value(square_r) * value(square_r) + value(square_i) * value(square_i) > 1e-24f) { + if (primal(square_r) * primal(square_r) + primal(square_i) * primal(square_i) > 1e-24f) { const T half_r = 0.5f * (upper_r - lower_r), half_i = 0.5f * (upper_i - lower_i); const T inverse = div_(T(1.0f), root_r * root_r + root_i * root_i); scale_r = (half_r * root_r + half_i * root_i) * inverse; @@ -252,7 +271,7 @@ __device__ __forceinline__ void transverse_step(T r2_free, T r2_bound, T exchang template __device__ __forceinline__ W exp_difference(W lower, W upper, W exp_lower, W exp_upper) { const W half = 0.5 * (upper - lower); - if (fabs(value(half)) < 1e-4) { + if (fabs(primal(half)) < 1e-4) { const W square = half * half; return exp_lower * (1.0 + half + 0.5 * square) * (1.0 + square / 6.0); } @@ -277,7 +296,7 @@ __device__ __forceinline__ void series_terms(W determinant, W minors, W& flat, W flat = next_flat; linear = next_linear; square = next_square; - using V = decltype(value(determinant)); + using V = decltype(primal(determinant)); constexpr double weight = inverse_factorial(K); sum_flat = sum_flat + V(weight) * flat; sum_linear = sum_linear + V(weight) * linear; @@ -296,6 +315,10 @@ template struct WorkOf { using type = typename std::conditional::type; }; +template +struct WorkOf, NARROW> { + using type = num::Multi::type>; +}; // expm((K - diag(R1)) dt) for free water (a), pool b and the semisolid pool // (c), which exchange with a and not with each other, times the attenuation: @@ -307,7 +330,7 @@ __device__ __forceinline__ void three_pool_step(T r1_free, T r1_b, T r1_c, T exc T fraction_b, T fraction_c, T dt, T attenuation, T* e, T* grow) { using W = typename WorkOf::type; - using V = decltype(value(W())); + using V = decltype(primal(W())); constexpr int TERMS = NARROW ? 24 : 16; const W step = W(dt); const W free = W(1.0f - fraction_b - fraction_c); @@ -325,7 +348,7 @@ __device__ __forceinline__ void three_pool_step(T r1_free, T r1_b, T r1_c, T exc const W determinant = s00 * s11 * s22 - a01 * (a10 * s22) + a02 * (-s11 * a20); W c[9]; bool close = true; - if constexpr (!NARROW) close = -2.0 * value(minors) < 1.0; + if constexpr (!NARROW) close = -2.0 * primal(minors) < 1.0; if (close) { W flat = V(1), linear = V(0), square = V(0), sum_flat = V(1), sum_linear = V(0), sum_square = V(0); series_terms(determinant, minors, flat, linear, square, sum_flat, sum_linear, @@ -361,7 +384,7 @@ __device__ __forceinline__ void three_pool_step(T r1_free, T r1_b, T r1_c, T exc const W first = exp_difference(low, middle, leading, centre); const W span = high - low; const W second = - (exp_difference(middle, high, centre, trailing) - first) / (value(span) > 0.0 ? span : W(1.0)); + (exp_difference(middle, high, centre, trailing) - first) / (primal(span) > 0.0 ? span : W(1.0)); const W m00 = a00 - low, m11 = a11 - low, m22 = a22 - low; const W n00 = a00 - middle, n11 = a11 - middle, n22 = a22 - middle; const W p00 = m00 * n00 + a01 * a10 + a02 * a20, p01 = m00 * a01 + a01 * n11; @@ -443,7 +466,7 @@ template __device__ __forceinline__ T lineshape_at(const float* lineshape, T offset_hz, int bins, float step) { const int last = bins - 1; const T scaled = min_(div_(abs_(offset_hz), T(step)), T(last + 0.0f)); - const float lower = fminf(floorf(value(scaled)), last - 1.0f); + const float lower = fminf(floorf(primal(scaled)), last - 1.0f); T h10, h01, h11; const T h00 = hermite_weights(scaled, lower, step, h10, h01, h11); const float* base = lineshape + static_cast(lower) * 2; @@ -813,7 +836,7 @@ __device__ __forceinline__ void pulse(const Params& p, int event, int base, cons const int row = profile_row + location[y]; const int last = p.profile_bins - 1; const T scaled = min_(max_(div_(alpha, T(p.profile_step)), T(0.0f)), T(last + 0.0f)); - const float lower = fminf(floorf(value(scaled)), last - 1.0f); + const float lower = fminf(floorf(primal(scaled)), last - 1.0f); T h10, h01, h11; const T h00 = hermite_weights(scaled, lower, p.profile_step, h10, h01, h11); const float* base_row = p.profile + (row * p.profile_bins + static_cast(lower)) * 8; diff --git a/src/blochsim/_layout_complex_vjp.cu.in b/src/blochsim/_layout_complex_vjp.cu.in new file mode 100644 index 00000000..db320c05 --- /dev/null +++ b/src/blochsim/_layout_complex_vjp.cu.in @@ -0,0 +1,50 @@ +// Written by CMake from _layout_complex_vjp.cu.in: the complex @KIND@ kernel +// for @POOLS@ in its layouts. +#include "_layout.hpp" + +namespace blochsim_layout { +namespace { + +// Events a checkpoint covers: the reverse sweep replays that many from it. +constexpr int SEGMENT = 4; + +template +__global__ void __launch_bounds__(32 * WARPS) adjoint_kernel(epg_vjp::Params v) { + extern __shared__ unsigned char segment[]; + epg_vjp::complex_vjp_loop(v, reinterpret_cast(segment)); +} + +template +int launch_rf(const epg_vjp::Params& v, cudaStream_t stream) { + constexpr int POOLS = @POOLS@; + const int groups = 32 / v.f.width; + const long long problems = v.problem_end - v.problem_base; + const long long per_block = static_cast(WARPS) * groups * Y; + const unsigned grid = static_cast((problems + per_block - 1) / per_block); + constexpr int PLANES = 2 * (POOLS >= 2 ? 2 : 1) + (POOLS == 3 ? 3 : (POOLS ? 2 : 1)); + const size_t shared = static_cast(SEGMENT) * PLANES * Y * 2 * 32 * WARPS * sizeof(T); + if (v.f.single_train) { + cudaFuncSetAttribute(adjoint_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, + static_cast(shared)); + adjoint_kernel<<>>(v); + } else { + cudaFuncSetAttribute(adjoint_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, + static_cast(shared)); + adjoint_kernel<<>>(v); + } + return static_cast(cudaGetLastError()); +} + +} // namespace + +int complex_@KIND@_@POOLS@(const epg_vjp::Params& v, int rf, cudaStream_t stream) { + using T = @TYPE@; + // Problems a thread holds: two for free water alone at float, one where + // a dual or a pool layout already fills the registers. + constexpr int Y = (@POOLS@ == 0 && !num::is_dual::value) ? 2 : 1; + if (rf == epg::DYNAMIC) return launch_rf(v, stream); + if (rf == epg::PROFILE) return launch_rf(v, stream); + return launch_rf(v, stream); +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_complex_vjp.hpp b/src/blochsim/_layout_complex_vjp.hpp new file mode 100644 index 00000000..b2844adc --- /dev/null +++ b/src/blochsim/_layout_complex_vjp.hpp @@ -0,0 +1,1293 @@ +// The complex EPG adjoint over a number type: the vector-Jacobian product at +// float, its derivative along a direction at num::Dual. +// +// The forward sweep keeps the state every K events; the reverse sweep replays +// each stretch of K events from its checkpoint into shared memory and walks it +// back. An interval's operator repeats while its length does, so the walk +// back sums, per lane, the products of each state with the cotangent the +// operator met it with, and contracts that sum against the operator's +// derivatives -- taken along every tissue input at once by num::Multi -- only +// when the operator changes. Pulses and samples are explicit complex maps +// and contract against their own derivatives at once. +#pragma once +#include "_layout_complex.hpp" + +namespace epg_vjp { + +using namespace epg; +using num::Multi; + +struct Params { + epg::Params f; + const float *grad_output_real, *grad_output_imag; + // Value and, for a dual sweep, direction of each gradient. + float *grad_tissue, *grad_flip, *grad_phase, *grad_duration, *grad_pair; + float *grad_tissue_t, *grad_flip_t, *grad_phase_t, *grad_duration_t, *grad_pair_t; + float *trajectory_r, *trajectory_i, *trajectory_tr, *trajectory_ti; + int problem_base, problem_end, shim_rows; + // How a three-pool interval is formed: an epg::Mode. + int mode; +}; + +template +struct Cx { + T r, i; +}; +template +__device__ __forceinline__ Cx cmul(const Cx& a, const Cx& b) { + return {a.r * b.r - a.i * b.i, a.r * b.i + a.i * b.r}; +} +template +__device__ __forceinline__ Cx cconj(const Cx& a) { + return {a.r, -a.i}; +} +// Re(conj(a) b). +template +__device__ __forceinline__ T rdot(const Cx& a, const Cx& b) { + return a.r * b.r + a.i * b.i; +} + +template +__device__ __forceinline__ void atomic_add(float* value, float* tangent, long long at, T x) { + if constexpr (num::is_dual::value) { + atomicAdd(value + at, x.v); + atomicAdd(tangent + at, x.d); + } else { + atomicAdd(value + at, x); + } +} +template +__device__ __forceinline__ T shfl_xor(T x, int mask) { + if constexpr (num::is_dual::value) { + return T{__shfl_xor_sync(0xffffffffu, x.v, mask), __shfl_xor_sync(0xffffffffu, x.d, mask)}; + } else { + return __shfl_xor_sync(0xffffffffu, x, mask); + } +} +template +__device__ __forceinline__ T lanes_sum(T x, int from, int to) { + for (int offset = from; offset < to; offset <<= 1) x += shfl_xor(x, offset); + return x; +} +template +__device__ __forceinline__ void put(float* values, float* tangents, long long at, T x) { + if constexpr (num::is_dual::value) { + values[at] = x.v; + tangents[at] = x.d; + } else { + values[at] = x; + } +} +template +__device__ __forceinline__ T get(const float* values, const float* tangents, long long at) { + if constexpr (num::is_dual::value) { + return T{values[at], tangents[at]}; + } else { + return values[at]; + } +} + +// What an interval reads of the launch besides the tissue and its length. +struct Geometry { + float flow_scale, washout_scale; + const float* pool_table; + int atom_count; +}; + +// Tissue directions an interval's operator is differentiated along. +enum Direction { D_R1 = 0, D_R2, D_DIFFUSION, D_B0, D_VELOCITY, D_FB, D_XB, D_R1B, D_R2B, D_SHIFT, D_FC, D_XC, D_R1C }; +template +constexpr int DIRECTIONS = POOLS == 3 ? 13 : (POOLS == 2 ? 10 : (POOLS == 1 ? 8 : 5)); + +// One interval's operator for one problem at this lane's order: transverse +// entries (one, or a 2x2 for an exchanging pool), longitudinal ones (the +// exchange matrix times the flow and damping factor), and recoveries. +template +struct Interval { + static constexpr int NT = POOLS >= 2 ? 4 : 1, NL = POOLS == 3 ? 9 : (POOLS ? 4 : 1); + static constexpr int NG = POOLS == 3 ? 3 : (POOLS ? 2 : 1), N = POOLS == 3 ? 3 : (POOLS ? 2 : 1); + Cx t[NT], l[NL]; + T g[NG]; +}; + +// What one problem is made of, as the interval reads it. +template +struct Inputs { + T r1, r2, diffusion, b0, velocity, fb, xb, r1b, r2b, shift, fc, xc, r1c; +}; + +// Whether an interval keeps the three pools' eigenvalues within +// NARROW_SPREAD (4) of each other: the shifted roots sum to zero, so the sum +// of their squares is -2 * minors. +__device__ __forceinline__ bool three_pool_narrow(float r1_free, float r1_b, float r1_c, float exchange_b, + float exchange_c, float fraction_b, float fraction_c, float dt) { + const float free = 1.0f - fraction_b - fraction_c; + const float kab = exchange_b * fraction_b, kba = exchange_b * free; + const float kac = exchange_c * fraction_c, kca = exchange_c * free; + const float a00 = (-kab - kac - r1_free) * dt, a01 = kba * dt, a02 = kca * dt; + const float a10 = kab * dt, a11 = (-kba - r1_b) * dt, a20 = kac * dt, a22 = (-kca - r1_c) * dt; + const float third = (a00 + a11 + a22) * (1.0f / 3.0f); + const float s00 = a00 - third, s11 = a11 - third, s22 = a22 - third; + const float minors = s00 * s11 - a01 * a10 + s00 * s22 - a02 * a20 + s11 * s22; + return -2.0f * minors < 16.0f; +} + +// The switches are read at run time: an interval is formed only when it +// changes, and one body then serves all of them. +template +__device__ __forceinline__ void interval(const Geometry& p, bool MOVING, bool DIFFUSING, bool OFF_AXIS, T dt, + const Inputs& in, float order, int row, int atom, Interval& out) { + T wout = 1.0f, turn = 0.0f; + if (MOVING) { + wout = 1.0f - min_(abs_(in.velocity) * p.washout_scale * dt, T(1.0f)); + turn = in.velocity * p.flow_scale * dt; + } + T damp_z = 1.0f, damp_t = 1.0f; + if (DIFFUSING) { + const T b = in.diffusion * dt; + const float sq = order * order; + damp_z = exp_(-b * sq); + damp_t = exp_(-b * (sq + order + 0.3333333333333333f)); + } + T oc = 1.0f, os = 0.0f; + if (MOVING || OFF_AXIS) { + T phase = -(order + 0.5f) * turn; + if (OFF_AXIS) phase = phase - TWO_PI * in.b0 * dt; + sincos_(phase, os, oc); + } + if constexpr (POOLS < 2) { + const T e2 = exp_(-in.r2 * dt) * wout * damp_t; + out.t[0] = {e2 * oc, e2 * os}; + } else { + T x[8]; + const T free = 1.0f - in.fb - (POOLS == 3 ? in.fc : T(0.0f)); + transverse_step(in.r2, in.r2b, in.xb, in.fb, free, in.shift, dt, wout, x); + const T rr = damp_t * oc, ri = damp_t * os; +#pragma unroll + for (int k = 0; k < 4; ++k) out.t[k] = {rr * x[2 * k] - ri * x[2 * k + 1], rr * x[2 * k + 1] + ri * x[2 * k]}; + } + T e[Interval::NL]; + if constexpr (POOLS == 0) { + e[0] = exp_(-in.r1 * dt) * wout; + out.g[0] = 1.0f - e[0]; + } else if constexpr (POOLS == 3) { + if constexpr (MODE == TABLE && !num::is_dual::value) { + const float* base = p.pool_table + static_cast(row) * (9 * p.atom_count) + atom; +#pragma unroll + for (int k = 0; k < 9; ++k) e[k] = wout * __ldg(base + k * p.atom_count); + const T fa = 1.0f - in.fb - in.fc; + out.g[0] = fa - (e[0] * fa + e[1] * in.fb + e[2] * in.fc); + out.g[1] = in.fb - (e[3] * fa + e[4] * in.fb + e[5] * in.fc); + out.g[2] = in.fc - (e[6] * fa + e[7] * in.fb + e[8] * in.fc); + } else if constexpr (MODE == NARROW) { + three_pool_step(in.r1, in.r1b, in.r1c, in.xb, in.xc, in.fb, in.fc, dt, wout, e, out.g); + } else { + // The float series holds the operator to float32 while the + // eigenvalues stay within NARROW_SPREAD of each other; double is + // for an interval that spreads them further. + if (three_pool_narrow(primal(in.r1), primal(in.r1b), primal(in.r1c), primal(in.xb), primal(in.xc), + primal(in.fb), primal(in.fc), primal(dt))) { + three_pool_step(in.r1, in.r1b, in.r1c, in.xb, in.xc, in.fb, in.fc, dt, wout, e, out.g); + } else { + three_pool_step(in.r1, in.r1b, in.r1c, in.xb, in.xc, in.fb, in.fc, dt, wout, e, out.g); + } + } + } else { + two_pool_step(in.r1, in.r1b, in.xb, in.fb, dt, wout, e, out.g); + } + T fr = damp_z, fi = 0.0f; + if (MOVING) { + T ts, tc; + sincos_(-order * turn, ts, tc); + fr = damp_z * tc; + fi = damp_z * ts; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) out.l[k] = {fr * e[k], fi * e[k]}; +} + +// One interval for the launch's switches. +template +__device__ __forceinline__ void form_interval(const Geometry& p, int relax_code, U dt, const Inputs& in, + float order, int row, int atom, Interval& out) { + interval(p, (relax_code & 4) != 0, (relax_code & 2) != 0, (relax_code & 1) != 0, dt, in, order, row, + atom, out); +} + +// Products of state and cotangent an operator met, summed while it repeats. +template +struct Met { + Cx t[Interval::NT], l[Interval::NL]; + T g[Interval::NG]; +}; + +// Re of every entry's derivative times what it met: the interval's share of +// a gradient. +template +__device__ __forceinline__ U contract(const Interval& derivative, const Met& met) { + U sum = 0.0f; +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) sum += derivative.t[k].r * met.t[k].r - derivative.t[k].i * met.t[k].i; +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) sum += derivative.l[k].r * met.l[k].r - derivative.l[k].i * met.l[k].i; +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) sum += derivative.g[k] * met.g[k]; + return sum; +} + +// The derivative of each entry along direction ``k`` of a multi-valued interval. +template +__device__ __forceinline__ Interval along(const Interval, POOLS>& m, int k) { + Interval out; +#pragma unroll + for (int j = 0; j < Interval::NT; ++j) out.t[j] = {m.t[j].r.d[k], m.t[j].i.d[k]}; +#pragma unroll + for (int j = 0; j < Interval::NL; ++j) out.l[j] = {m.l[j].r.d[k], m.l[j].i.d[k]}; +#pragma unroll + for (int j = 0; j < Interval::NG; ++j) out.g[j] = m.g[j].d[k]; + return out; +} +template +__device__ __forceinline__ Interval values(const Interval, POOLS>& m) { + Interval out; +#pragma unroll + for (int j = 0; j < Interval::NT; ++j) out.t[j] = {m.t[j].r.v, m.t[j].i.v}; +#pragma unroll + for (int j = 0; j < Interval::NL; ++j) out.l[j] = {m.l[j].r.v, m.l[j].i.v}; +#pragma unroll + for (int j = 0; j < Interval::NG; ++j) out.g[j] = m.g[j].v; + return out; +} + +// A number with its direction dropped, and the direction alone. +template +__device__ __forceinline__ T flat(T x) { + if constexpr (num::is_dual::value) return T(x.v); + else return x; +} +template +__device__ __forceinline__ float direction(T x) { + if constexpr (num::is_dual::value) return x.d; + else return 0.0f; +} +template +__device__ __forceinline__ Inputs flat(const Inputs& in) { + return {flat(in.r1), flat(in.r2), flat(in.diffusion), flat(in.b0), flat(in.velocity), flat(in.fb), flat(in.xb), + flat(in.r1b), flat(in.r2b), flat(in.shift), flat(in.fc), flat(in.xc), flat(in.r1c)}; +} +// An operator moved along its length by ``d``: its direction gains d times +// the slope's value. A float operator has no direction to move. +template +__device__ __forceinline__ Interval moved(const Interval& a, const Interval& slope, float d) { + Interval out = a; + if (d == 0.0f) return out; + if constexpr (num::is_dual::value) { +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) { + out.t[k].r.d += d * slope.t[k].r.v; + out.t[k].i.d += d * slope.t[k].i.v; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) { + out.l[k].r.d += d * slope.l[k].r.v; + out.l[k].i.d += d * slope.l[k].i.v; + } +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) out.g[k].d += d * slope.g[k].v; + } + return out; +} + +// What an interval contributes to each tissue direction's gradient. +template +struct Gradients { + T g[DIRECTIONS]; +}; + +// Out of line: an interval changes rarely, and one copy then serves every +// kernel of the layout instead of one inlined per kernel. Everything goes in +// and comes out by value, so nothing in the caller is addressed. +template +__device__ __noinline__ Gradients contract_interval(Geometry g, int relax_code, T dt, Inputs in, float order, + int row, int atom, Met met) { + constexpr int KD = DIRECTIONS, CHUNK = 4; + Gradients out; +#pragma unroll + for (int k = 0; k < KD; ++k) out.g[k] = 0.0f; +#pragma unroll 1 + for (int c = 0; c * CHUNK < KD; ++c) { + using V = Multi; + auto seed = [&](T value, int direction) { + const int local = direction - c * CHUNK; + return V(value, direction < KD && local >= 0 && local < CHUNK ? local : -1); + }; + Inputs seeded; + seeded.r1 = seed(in.r1, D_R1); + seeded.r2 = seed(in.r2, D_R2); + seeded.diffusion = seed(in.diffusion, D_DIFFUSION); + seeded.b0 = seed(in.b0, D_B0); + seeded.velocity = seed(in.velocity, D_VELOCITY); + seeded.fb = seed(in.fb, D_FB); + seeded.xb = seed(in.xb, D_XB); + seeded.r1b = seed(in.r1b, D_R1B); + seeded.r2b = seed(in.r2b, D_R2B); + seeded.shift = seed(in.shift, D_SHIFT); + seeded.fc = seed(in.fc, D_FC); + seeded.xc = seed(in.xc, D_XC); + seeded.r1c = seed(in.r1c, D_R1C); + Interval m; + form_interval(g, relax_code, V(dt, -1), seeded, order, row, atom, m); +#pragma unroll + for (int k = 0; k < CHUNK; ++k) { + const T got = contract(along(m, k), met); +#pragma unroll + for (int j = 0; j < KD; ++j) { + if (j == c * CHUNK + k) out.g[j] = got; + } + } + } + return out; +} + +// An interval's operator and its derivative along its own length. +template +struct Opened { + Interval value, slope; +}; +template +__device__ __noinline__ Opened open_interval_at(Geometry g, int relax_code, T dt, Inputs in, float order, + int row, int atom) { + using V = Multi<1, T>; + Inputs seeded; + seeded.r1 = V(in.r1, -1); + seeded.r2 = V(in.r2, -1); + seeded.diffusion = V(in.diffusion, -1); + seeded.b0 = V(in.b0, -1); + seeded.velocity = V(in.velocity, -1); + seeded.fb = V(in.fb, -1); + seeded.xb = V(in.xb, -1); + seeded.r1b = V(in.r1b, -1); + seeded.r2b = V(in.r2b, -1); + seeded.shift = V(in.shift, -1); + seeded.fc = V(in.fc, -1); + seeded.xc = V(in.xc, -1); + seeded.r1c = V(in.r1c, -1); + Interval m; + form_interval(g, relax_code, V(dt, 0), seeded, order, row, atom, m); + return {values(m), along(m, 0)}; +} + +// An interval's operator alone, for the forward sweep. +template +__device__ __noinline__ Interval interval_at(Geometry g, int relax_code, T dt, Inputs in, float order, int row, + int atom) { + Interval out; + form_interval(g, relax_code, dt, in, order, row, atom, out); + return out; +} + +// The three out-of-line interval functions with the three-pool mode read at +// run time: only an interval's change pays for the choice. +template +__device__ __forceinline__ Interval interval_in(int mode, Geometry g, int relax_code, T dt, const Inputs& in, + float order, int row, int atom) { + if constexpr (POOLS == 3) { + if (mode == TABLE) return interval_at(g, relax_code, dt, in, order, row, atom); + if (mode == ROOTS) return interval_at(g, relax_code, dt, in, order, row, atom); + } + return interval_at(g, relax_code, dt, in, order, row, atom); +} +template +__device__ __forceinline__ Opened open_in(int mode, Geometry g, int relax_code, T dt, const Inputs& in, + float order, int row, int atom) { + if constexpr (POOLS == 3) { + if (mode == TABLE) return open_interval_at(g, relax_code, dt, in, order, row, atom); + if (mode == ROOTS) return open_interval_at(g, relax_code, dt, in, order, row, atom); + } + return open_interval_at(g, relax_code, dt, in, order, row, atom); +} +template +__device__ __forceinline__ Gradients contract_in(int mode, Geometry g, int relax_code, T dt, const Inputs& in, + float order, int row, int atom, const Met& met) { + if constexpr (POOLS == 3) { + if (mode == TABLE) return contract_interval(g, relax_code, dt, in, order, row, atom, met); + if (mode == ROOTS) return contract_interval(g, relax_code, dt, in, order, row, atom, met); + } + return contract_interval(g, relax_code, dt, in, order, row, atom, met); +} + +// The pulse as a complex 3x3 on (F+, F-, Z): a hard pulse of flip ``alpha`` +// at phase (c, s), or the rotation of a Cayley-Klein pair (a, b). +template +__device__ __forceinline__ void hard_rotation(T alpha, T cphi, T sphi, Cx* R) { + T s, c; + sincos_(alpha, s, c); + const T chs = 0.5f * (1.0f + c), shs = 0.5f * (1.0f - c), hs = 0.5f * s; + const T c2 = cphi * cphi - sphi * sphi, s2 = 2.0f * sphi * cphi; + R[0] = {chs, T(0.0f)}; + R[1] = {shs * c2, shs * s2}; + R[2] = {s * sphi, -(s * cphi)}; + R[3] = {shs * c2, -(shs * s2)}; + R[4] = {chs, T(0.0f)}; + R[5] = {s * sphi, s * cphi}; + R[6] = {-(hs * sphi), -(hs * cphi)}; + R[7] = {-(hs * sphi), hs * cphi}; + R[8] = {c, T(0.0f)}; +} +template +__device__ __forceinline__ void spinor_rotation(Cx a, Cx b, Cx* R) { + const Cx aa = cmul(a, a), bb = cmul(b, b), ab = cmul(a, b); + const Cx cross = cmul(cconj(a), b); + R[0] = cconj(aa); + R[1] = {-bb.r, bb.i}; + R[2] = {-2.0f * ab.r, 2.0f * ab.i}; + R[3] = {-bb.r, -bb.i}; + R[4] = aa; + R[5] = {-2.0f * ab.r, -2.0f * ab.i}; + R[6] = cross; + R[7] = cconj(cross); + R[8] = {a.r * a.r + a.i * a.i - b.r * b.r - b.i * b.i, T(0.0f)}; +} + +template +__device__ __forceinline__ void complex_vjp_loop(const Params& v, T* segment) { + const epg::Params& p = v.f; + const int mode = v.mode; + constexpr bool one_train = ONE_TRAIN; + constexpr int NS = POOLS >= 2 ? 2 : 1; // transverse pools + constexpr int NZ = Interval::N; // longitudinal pools + constexpr int PLANES = 2 * NS + NZ; // complex planes a state is + constexpr int KD = DIRECTIONS; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const int first = v.problem_base + ((blockIdx.x * (blockDim.x >> 5) + warp) * groups + group) * Y; + const float order = static_cast(state); + + int problem[Y], atom[Y], train[Y], location[Y]; + bool active[Y], live[Y]; + Inputs in[Y]; + T t1v[Y], t2v[Y], t1b[Y], t2b[Y], t1c[Y]; + T m0[Y], b1[Y], b1c[Y], b1s[Y], inversion[Y]; + // The state: per transverse pool F+ and F-, per longitudinal pool Z. + Cx plus[NS][Y], minus[NS][Y], z[NZ][Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + problem[y] = first + y; + active[y] = problem[y] < v.problem_end; + live[y] = active[y] && state < p.state_count; + atom[y] = problem[y] % p.atom_count; + train[y] = problem[y] / p.atom_count; + location[y] = RF == PROFILE ? atom[y] % p.locations : 0; + const int at = p.atom_stride ? atom[y] : 0; + auto read = [&](const float* values, const float* directions, int where, float otherwise) { + return active[y] ? num::load(values, directions, where) : T(otherwise); + }; + t1v[y] = read(p.t1, p.d_t1, atom[y], 1.0f); + t2v[y] = read(p.t2, p.d_t2, atom[y], 1.0f); + in[y].r1 = num::rate(t1v[y]); + in[y].r2 = num::rate(t2v[y]); + m0[y] = p.density ? read(p.m0, p.d_m0, at, 0.0f) : T(1.0f); + b1[y] = p.transmit ? read(p.b1, p.d_b1, at, 1.0f) : T(1.0f); + const T b1_phase = p.off_axis ? read(p.b1_phase, p.d_b1_phase, at, 0.0f) : T(0.0f); + sincos_(b1_phase, b1s[y], b1c[y]); + in[y].b0 = p.off_axis ? read(p.b0, p.d_b0, at, 0.0f) : T(0.0f); + inversion[y] = p.inverting ? read(p.inversion_efficiency, p.d_inversion_efficiency, at, 1.0f) : T(1.0f); + in[y].diffusion = p.diffusing ? read(p.diffusion, p.d_diffusion, at, 0.0f) : T(0.0f); + in[y].velocity = p.moving ? read(p.velocity, p.d_velocity, at, 0.0f) : T(0.0f); + in[y].fb = in[y].xb = in[y].r1b = in[y].r2b = in[y].shift = in[y].fc = in[y].xc = in[y].r1c = 0.0f; + t1b[y] = t2b[y] = t1c[y] = 1.0f; + if constexpr (POOLS == 1) { + in[y].fb = read(p.bound_fraction, p.d_bound_fraction, at, 0.0f); + in[y].xb = read(p.bound_exchange, p.d_bound_exchange, at, 0.0f); + t1b[y] = read(p.t1_bound, p.d_t1_bound, at, 1.0f); + } + if constexpr (POOLS >= 2) { + in[y].fb = read(p.pool_b_fraction, p.d_pool_b_fraction, at, 0.0f); + in[y].xb = read(p.pool_b_exchange, p.d_pool_b_exchange, at, 0.0f); + t1b[y] = read(p.t1_pool_b, p.d_t1_pool_b, at, 1.0f); + t2b[y] = read(p.t2_pool_b, p.d_t2_pool_b, at, 1.0f); + in[y].shift = read(p.pool_b_shift, p.d_pool_b_shift, at, 0.0f); + } + if constexpr (POOLS == 3) { + in[y].fc = read(p.bound_fraction, p.d_bound_fraction, at, 0.0f); + in[y].xc = read(p.bound_exchange, p.d_bound_exchange, at, 0.0f); + t1c[y] = read(p.t1_bound, p.d_t1_bound, at, 1.0f); + } + in[y].r1b = num::rate(t1b[y]); + in[y].r2b = num::rate(t2b[y]); + in[y].r1c = num::rate(t1c[y]); + } + const bool uniform = one_train || (active[0] && train[0] == train[Y - 1]); + const int base = one_train ? 0 : train[0] * p.event_count; + const int first_train = __shfl_sync(0xffffffffu, train[0], 0); + const bool warp_train = one_train || __all_sync(0xffffffffu, uniform && train[0] == first_train); + const int relax_code = (p.moving ? 4 : 0) | (p.diffusing ? 2 : 0) | (p.off_axis ? 1 : 0); + const Geometry geometry{p.flow_scale, p.washout_scale, p.pool_table, p.atom_count}; + + auto initial = [&]() { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int s = 0; s < NS; ++s) plus[s][y] = minus[s][y] = {T(0.0f), T(0.0f)}; + const T free = 1.0f - in[y].fb - in[y].fc; + z[0][y] = {state == 0 ? free : T(0.0f), T(0.0f)}; + if constexpr (NZ >= 2) z[1][y] = {state == 0 ? in[y].fb : T(0.0f), T(0.0f)}; + if constexpr (NZ >= 3) z[2][y] = {state == 0 ? in[y].fc : T(0.0f), T(0.0f)}; + } + }; + auto event_dt = [&](int y, int event) { + if (uniform) return num::load(p.duration, p.d_duration, base + event); + return active[y] ? num::load(p.duration, p.d_duration, train[y] * p.event_count + event) : T(0.0f); + }; + auto event_row = [&](int y, int event) { + if (mode != TABLE) return 0; + return p.duration_row[(uniform || !active[y] ? base : train[y] * p.event_count) + event]; + }; + // The forward step's interval, kept while it repeats. + Interval ahead[Y], ahead_slope[Y]; + T ahead_dt = -1.0f; + int ahead_row = -1; + auto apply = [&](const Interval* op) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const Interval& o = op[y]; + if constexpr (NS == 1) { + plus[0][y] = cmul(o.t[0], plus[0][y]); + minus[0][y] = cmul(cconj(o.t[0]), minus[0][y]); + } else { + const Cx p0 = plus[0][y], p1 = plus[1][y], m0v = minus[0][y], m1 = minus[1][y]; + const Cx a = cmul(o.t[0], p0), b = cmul(o.t[1], p1), c = cmul(o.t[2], p0), d = cmul(o.t[3], p1); + plus[0][y] = {a.r + b.r, a.i + b.i}; + plus[1][y] = {c.r + d.r, c.i + d.i}; + const Cx e = cmul(cconj(o.t[0]), m0v), f = cmul(cconj(o.t[1]), m1); + const Cx g = cmul(cconj(o.t[2]), m0v), h = cmul(cconj(o.t[3]), m1); + minus[0][y] = {e.r + f.r, e.i + f.i}; + minus[1][y] = {g.r + h.r, g.i + h.i}; + } + Cx next[NZ]; +#pragma unroll + for (int i = 0; i < NZ; ++i) { + next[i] = {state == 0 ? o.g[i] : T(0.0f), T(0.0f)}; +#pragma unroll + for (int j = 0; j < NZ; ++j) { + const Cx term = cmul(o.l[i * NZ + j], z[j][y]); + next[i] = {next[i].r + term.r, next[i].i + term.i}; + } + } +#pragma unroll + for (int i = 0; i < NZ; ++i) z[i][y] = next[i]; + } + }; + // A train repeats its interval's length, whatever direction each event + // moves it in: the operator is kept at the length's value and each event + // moves it along its own direction through the slope. + auto relax_forward = [&](int event) { + const T dt0 = event_dt(0, event); + const int row0 = event_row(0, event); + if (uniform && same(dt0, T(0.0f))) return; + if (!uniform) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + ahead[y] = interval_in(mode, geometry, relax_code, event_dt(y, event), in[y], order, + event_row(y, event), atom[y]); + } + ahead_dt = -1.0f; + apply(ahead); + return; + } + if (!same(flat(dt0), ahead_dt) || row0 != ahead_row) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if constexpr (num::is_dual::value) { + const Opened opened = + open_in(mode, geometry, relax_code, flat(dt0), in[y], order, row0, atom[y]); + ahead[y] = opened.value; + ahead_slope[y] = opened.slope; + } else { + ahead[y] = interval_in(mode, geometry, relax_code, dt0, in[y], order, row0, atom[y]); + } + } + ahead_dt = flat(dt0); + ahead_row = row0; + } + if constexpr (num::is_dual::value) { + const float d = direction(dt0); + if (d != 0.0f) { + Interval now[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) now[y] = moved(ahead[y], ahead_slope[y], d); + apply(now); + return; + } + } + apply(ahead); + }; + auto shift = [&](int s) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T up_r = num::shfl_up(plus[s][y].r, 1, width), up_i = num::shfl_up(plus[s][y].i, 1, width); + const T dn_r = num::shfl_down(minus[s][y].r, 1, width), dn_i = num::shfl_down(minus[s][y].i, 1, width); + const bool keep_up = state > 0 && live[y], keep_down = state + 1 < p.state_count && live[y]; + const T pr = keep_up ? up_r : T(0.0f), pi = keep_up ? up_i : T(0.0f); + const T mr = keep_down ? dn_r : T(0.0f), mi = keep_down ? dn_i : T(0.0f); + plus[s][y] = {state == 0 ? mr : pr, state == 0 ? -mi : pi}; + minus[s][y] = {mr, mi}; + } + }; + // The flip, phase and transmit of a pulse for problem y. + auto pulse_terms = [&](int y, int event, T& alpha, T& cphi, T& sphi, T& transmit) { + const int e = uniform ? base + event : (active[y] ? train[y] * p.event_count + event : event); + const T flip = num::load(p.flip, p.d_flip, e); + T ce, se; + event_phase(p, e, ce, se); + T tb1 = b1[y], tc = b1c[y], ts = b1s[y]; + if (p.shimmed) { + const int cell = p.shim_index[event] * p.atom_count + atom[y]; + tb1 = p.transmit ? (active[y] ? num::load(p.b1, p.d_b1, cell) : T(1.0f)) : T(1.0f); + if (p.off_axis) sincos_(active[y] ? num::load(p.b1_phase, p.d_b1_phase, cell) : T(0.0f), ts, tc); + } + alpha = flip * tb1; + transmit = tb1; + cphi = ce * tc - se * ts; + sphi = se * tc + ce * ts; + return flip; + }; + auto pair_of = [&](int y, int event, Cx& a, Cx& b) { + const int row = p.pair_index[(active[y] ? train[y] * p.event_count : 0) + event]; + const long long cell = (static_cast(row) * p.atom_count + atom[y]) * 4; + a = {num::load(p.pairs, p.pair_direction, cell), num::load(p.pairs, p.pair_direction, cell + 1)}; + b = {num::load(p.pairs, p.pair_direction, cell + 2), num::load(p.pairs, p.pair_direction, cell + 3)}; + return cell; + }; + // The rotation a pulse performs on problem y, from its flip and phase. + auto rotation = [&](int y, int event, auto alpha, auto cphi, auto sphi, auto* R) { + using U = typename std::decay::type; + if constexpr (RF == HARD) { + hard_rotation(alpha, cphi, sphi, R); + } else if constexpr (RF == PROFILE) { + const int row = p.profile_index[event] * p.locations + location[y]; + const int last = p.profile_bins - 1; + const U scaled = min_(max_(div_(alpha, U(p.profile_step)), U(0.0f)), U(last + 0.0f)); + const float lower = fminf(floorf(primal(scaled)), last - 1.0f); + U h10, h01, h11; + const U h00 = hermite_weights(scaled, lower, p.profile_step, h10, h01, h11); + const float* knot = p.profile + (row * p.profile_bins + static_cast(lower)) * 8; + U pair[4]; +#pragma unroll + for (int c = 0; c < 4; ++c) { + pair[c] = h00 * __ldg(knot + c) + h10 * __ldg(knot + 4 + c) + h01 * __ldg(knot + 8 + c) + + h11 * __ldg(knot + 12 + c); + } + const Cx turn = {cphi, -sphi}; + spinor_rotation(Cx{pair[0], pair[1]}, cmul(Cx{pair[2], pair[3]}, turn), R); + } + }; + auto absorbed_of = [&](int y, int event, auto alpha, auto offset_b0) { + using U = typename std::decay::type; + const float saturation = p.saturation[event]; + const U shape = lineshape_at(p.lineshape, p.rf_frequency[event] - offset_b0, p.lineshape_bins, p.lineshape_step); + return exp_(saturation * alpha * alpha * shape); + }; + auto rotate = [&](int y, const Cx* R, int s, int zi) { + const Cx fp = plus[s][y], fm = minus[s][y], zz = z[zi][y]; + Cx out[3]; +#pragma unroll + for (int i = 0; i < 3; ++i) { + const Cx a = cmul(R[3 * i], fp), b = cmul(R[3 * i + 1], fm), c = cmul(R[3 * i + 2], zz); + out[i] = {a.r + b.r + c.r, a.i + b.i + c.i}; + } + plus[s][y] = out[0]; + minus[s][y] = out[1]; + z[zi][y] = out[2]; + }; + auto pulse_forward = [&](int event) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + T alpha, cphi, sphi, transmit; + pulse_terms(y, event, alpha, cphi, sphi, transmit); + Cx R[9]; + if constexpr (RF == DYNAMIC) { + Cx a, b; + pair_of(y, event, a, b); + spinor_rotation(a, cmul(b, Cx{cphi, -sphi}), R); + } else { + rotation(y, event, alpha, cphi, sphi, R); + } + rotate(y, R, 0, 0); + if constexpr (POOLS >= 2) rotate(y, R, 1, 1); + if constexpr (POOLS == 1 || POOLS == 3) { + const T absorbed = absorbed_of(y, event, alpha, p.off_axis ? in[y].b0 : T(0.0f)); + Cx& zz = z[POOLS == 1 ? 1 : 2][y]; + zz = {absorbed * zz.r, absorbed * zz.i}; + } + } + }; + auto forward = [&](int event) { + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + relax_forward(event); + if (act & 1) { +#pragma unroll + for (int s = 0; s < NS; ++s) shift(s); + } + if (kind == 1 && (act & 4)) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T eff = inversion[y]; + z[0][y] = {-eff * z[0][y].r, -eff * z[0][y].i}; + if constexpr (POOLS >= 2) z[1][y] = {-eff * z[1][y].r, -eff * z[1][y].i}; + } + } else if (kind == 1) { + pulse_forward(event); + } + if (act & 18) { +#pragma unroll + for (int s = 0; s < NS; ++s) shift(s); + } + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int s = 0; s < NS; ++s) plus[s][y] = minus[s][y] = {T(0.0f), T(0.0f)}; + } + } + }; + // Checkpoints: a complex plane per transverse and longitudinal component. + const long long stride = static_cast(p.event_count) * PLANES * p.state_count; + auto slot = [&](int y, int checkpoint, int plane) { + return (problem[y] - v.problem_base) * stride + + (static_cast(checkpoint) * PLANES + plane) * p.state_count + state; + }; + auto plane_of = [&](int y, int plane) -> Cx& { + if (plane < NS) return plus[plane][y]; + if (plane < 2 * NS) return minus[plane - NS][y]; + return z[plane - 2 * NS][y]; + }; + const int checkpoints = (p.event_count + K - 1) / K; + initial(); +#pragma unroll 1 + for (int event = 0; event < p.event_count; ++event) { + if (event % K == 0) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (live[y]) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + const Cx value = plane_of(y, plane); + put(v.trajectory_r, v.trajectory_tr, slot(y, event / K, plane), value.r); + put(v.trajectory_i, v.trajectory_ti, slot(y, event / K, plane), value.i); + } + } + } + } + forward(event); + } + + // ---- the walk back ---- + // Cotangents of the state, laid out as the state is. + Cx bplus[NS][Y], bminus[NS][Y], bz[NZ][Y]; + // Per-lane gradients of each problem's tissue: the interval's directions, + // then m0, b1, the transmit phase and the inversion. + T grad[KD][Y], g_m0[Y], g_b1[Y], g_b1_phase[Y], g_inversion[Y], g_b0_pulse[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int s = 0; s < NS; ++s) bplus[s][y] = bminus[s][y] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int i = 0; i < NZ; ++i) bz[i][y] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int k = 0; k < KD; ++k) grad[k][y] = 0.0f; + g_m0[y] = g_b1[y] = g_b1_phase[y] = g_inversion[y] = g_b0_pulse[y] = 0.0f; + } + // The interval the walk back is in, its derivative along its own length, + // and what it has met since it began. + Interval back[Y], back_dt[Y], back_curve[Y]; + // What the interval met, weighted by each event's direction along its + // length: the mixed second derivatives contract against it. + Met met_along[Y]; + Met met[Y]; + T back_length[Y]; + int back_row[Y]; + bool back_open = false; + auto clear_met = [&]() { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) { + met[y].t[k] = {T(0.0f), T(0.0f)}; + met_along[y].t[k] = {0.0f, 0.0f}; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) { + met[y].l[k] = {T(0.0f), T(0.0f)}; + met_along[y].l[k] = {0.0f, 0.0f}; + } +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) { + met[y].g[k] = 0.0f; + met_along[y].g[k] = 0.0f; + } + } + }; + auto close_interval = [&]() { + if (!back_open) return; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const Gradients got = contract_in(mode, geometry, relax_code, back_length[y], in[y], + order, back_row[y], atom[y], met[y]); +#pragma unroll + for (int k = 0; k < KD; ++k) grad[k][y] += got.g[k]; + if constexpr (num::is_dual::value) { + // d/dlength of each tissue derivative, against the met the + // events' own length directions weighted. + Met along_met; +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) along_met.t[k] = {T(met_along[y].t[k].r), T(met_along[y].t[k].i)}; +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) along_met.l[k] = {T(met_along[y].l[k].r), T(met_along[y].l[k].i)}; +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) along_met.g[k] = T(met_along[y].g[k]); + const Gradients mixed = contract_in(mode, geometry, relax_code, T(back_length[y].v, 1.0f), flat(in[y]), order, back_row[y], atom[y], along_met); +#pragma unroll + for (int k = 0; k < KD; ++k) grad[k][y].d += mixed.g[k].d; + } + } + clear_met(); + }; + auto open_interval = [&](int event) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T dt = event_dt(y, event); + back_length[y] = uniform ? flat(dt) : dt; + back_row[y] = event_row(y, event); + const Opened opened = + open_in(mode, geometry, relax_code, back_length[y], in[y], order, back_row[y], atom[y]); + back[y] = opened.value; + back_dt[y] = opened.slope; + if constexpr (num::is_dual::value) { + if (uniform) { + // The slope's own derivative along the length. + const Opened curved = open_in(mode, geometry, relax_code, T(back_length[y].v, 1.0f), flat(in[y]), order, back_row[y], atom[y]); + // Held as values, which is what moving the slope reads. +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) { + back_curve[y].t[k] = {T(curved.slope.t[k].r.d), T(curved.slope.t[k].i.d)}; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) { + back_curve[y].l[k] = {T(curved.slope.l[k].r.d), T(curved.slope.l[k].i.d)}; + } +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) back_curve[y].g[k] = T(curved.slope.g[k].d); + } + } + } + back_open = true; + }; + auto event_gradient = [&](float* value, float* tangent, int event, const T* lane_values) { + if (warp_train) { + T sum = lane_values[0]; +#pragma unroll + for (int y = 1; y < Y; ++y) sum += lane_values[y]; + sum = lanes_sum(sum, 1, 32); + if (lane == 0 && active[0]) atomic_add(value, tangent, base + event, sum); + } else { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T sum = lanes_sum(lane_values[y], 1, width); + if (state == 0 && active[y]) atomic_add(value, tangent, train[y] * p.event_count + event, sum); + } + } + }; + auto shift_adjoint = [&](int s) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T carry_r = num::shfl(bplus[s][y].r, 0, width), carry_i = -num::shfl(bplus[s][y].i, 0, width); + const T dn_r = num::shfl_down(bplus[s][y].r, 1, width), dn_i = num::shfl_down(bplus[s][y].i, 1, width); + const T up_r = num::shfl_up(bminus[s][y].r, 1, width), up_i = num::shfl_up(bminus[s][y].i, 1, width); + const bool forward_ok = state + 1 < p.state_count && live[y], back_ok = state > 0 && live[y]; + Cx m = {back_ok ? up_r : T(0.0f), back_ok ? up_i : T(0.0f)}; + if (state == 1 && live[y]) m = {m.r + carry_r, m.i + carry_i}; + bplus[s][y] = {forward_ok ? dn_r : T(0.0f), forward_ok ? dn_i : T(0.0f)}; + bminus[s][y] = m; + } + }; + clear_met(); + const int threads = blockDim.x; + auto smem = [&](int k, int plane, int y, int part) -> T& { + return segment[(((k * PLANES + plane) * Y + y) * 2 + part) * threads + threadIdx.x]; + }; + const long long n_atoms = p.atom_count; + const int past_transmit = 2 * (v.shim_rows - 1); + +#pragma unroll 1 + for (int checkpoint = checkpoints - 1; checkpoint >= 0; --checkpoint) { + const int start = checkpoint * K; + const int stop = min(start + K, p.event_count); +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + plane_of(y, plane) = live[y] ? Cx{get(v.trajectory_r, v.trajectory_tr, slot(y, checkpoint, plane)), + get(v.trajectory_i, v.trajectory_ti, slot(y, checkpoint, plane))} + : Cx{T(0.0f), T(0.0f)}; + } + } +#pragma unroll 1 + for (int event = start; event < stop; ++event) { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + smem(event - start, plane, y, 0) = plane_of(y, plane).r; + smem(event - start, plane, y, 1) = plane_of(y, plane).i; + } + } + if (event + 1 < stop) forward(event); + } +#pragma unroll 1 + for (int event = stop - 1; event >= start; --event) { + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + // The entry state, kept for the interval's contraction. + Cx entry[PLANES][Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + entry[plane][y] = {smem(event - start, plane, y, 0), smem(event - start, plane, y, 1)}; + plane_of(y, plane) = entry[plane][y]; + } + } + // The interval this event relaxes over, and the stage its pulse + // or sample sees. + { + const T dt0 = event_dt(0, event); + const int row0 = event_row(0, event); + bool changed = !back_open || !uniform || row0 != back_row[0]; + if (!changed) changed = !same(flat(dt0), back_length[0]); + if (changed) { + close_interval(); + open_interval(event); + } + } + // This event's operator and slope: the kept ones moved along its + // length's direction. + const float along = uniform ? direction(event_dt(0, event)) : 0.0f; + constexpr bool MOVES = num::is_dual::value; + Interval moved_op[MOVES ? Y : 1], moved_slope[MOVES ? Y : 1]; + if constexpr (MOVES) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + moved_op[y] = moved(back[y], back_dt[y], along); + moved_slope[y] = moved(back_dt[y], back_curve[y], along); + } + } + const Interval* op = MOVES ? moved_op : back; + const Interval* slope = MOVES ? moved_slope : back_dt; + apply(op); + if (act & 1) { +#pragma unroll + for (int s = 0; s < NS; ++s) shift(s); + } + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) { +#pragma unroll + for (int s = 0; s < NS; ++s) bplus[s][y] = bminus[s][y] = {T(0.0f), T(0.0f)}; + } + } else if (act & 18) { +#pragma unroll + for (int s = 0; s < NS; ++s) shift_adjoint(s); + } + if (kind == 1 && (act & 4)) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T eff = inversion[y]; + g_inversion[y] -= rdot(bz[0][y], z[0][y]); + bz[0][y] = {-eff * bz[0][y].r, -eff * bz[0][y].i}; + if constexpr (POOLS >= 2) { + g_inversion[y] -= rdot(bz[1][y], z[1][y]); + bz[1][y] = {-eff * bz[1][y].r, -eff * bz[1][y].i}; + } + } + } else if (kind == 1) { + T flip_gradient[Y], phase_gradient[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + T alpha, cphi, sphi, transmit; + const T flip = pulse_terms(y, event, alpha, cphi, sphi, transmit); + T g_alpha = 0.0f, g_phi = 0.0f; + Cx R[9]; + // Re(conj(cotangent) dR S) along the pulse's own inputs. + auto against = [&](const auto* dR, int s, int zi) { + using U = typename std::decay::type; + const Cx S[3] = {plus[s][y], minus[s][y], z[zi][y]}; + const Cx L[3] = {bplus[s][y], bminus[s][y], bz[zi][y]}; + U sum = 0.0f; +#pragma unroll + for (int i = 0; i < 3; ++i) { +#pragma unroll + for (int j = 0; j < 3; ++j) { + const Cx term = {dR[3 * i + j].r * S[j].r - dR[3 * i + j].i * S[j].i, + dR[3 * i + j].r * S[j].i + dR[3 * i + j].i * S[j].r}; + sum += L[i].r * term.r + L[i].i * term.i; + } + } + return sum; + }; + if constexpr (RF == DYNAMIC) { + using V = Multi<5, T>; + Cx a, b; + const long long cell = pair_of(y, event, a, b); + // The phase turns b by e^{-i phi}: d/dphi of (cos, -sin) is (-sin, -cos). + Cx turn = {V(cphi, -1), V(-sphi, -1)}; + turn.r.d[0] = -sphi; + turn.i.d[0] = -cphi; + Cx mb[9]; + spinor_rotation(Cx{V(a.r, 1), V(a.i, 2)}, cmul(Cx{V(b.r, 3), V(b.i, 4)}, turn), mb); + V got = against(mb, 0, 0); + if constexpr (POOLS >= 2) got += against(mb, 1, 1); + g_phi = got.d[0]; + T pair_sums[4] = {got.d[1], got.d[2], got.d[3], got.d[4]}; +#pragma unroll + for (int c = 0; c < 4; ++c) { + const T sum = lanes_sum(pair_sums[c], 1, width); + if (state == 0 && active[y]) atomic_add(v.grad_pair, v.grad_pair_t, cell + c, sum); + } + Cx RR[9]; + spinor_rotation(a, cmul(b, Cx{cphi, -sphi}), RR); +#pragma unroll + for (int k = 0; k < 9; ++k) R[k] = RR[k]; + } else { + using V = Multi<2, T>; + Cx mr[9]; + V va = V(alpha, 0), vc = V(cphi, 1), vs = V(sphi, 1); + vc.d[1] = -sphi; + vs.d[1] = cphi; + rotation(y, event, va, vc, vs, mr); + V got = against(mr, 0, 0); + if constexpr (POOLS >= 2) got += against(mr, 1, 1); + g_alpha = got.d[0]; + g_phi = got.d[1]; +#pragma unroll + for (int k = 0; k < 9; ++k) R[k] = {mr[k].r.v, mr[k].i.v}; + } + if constexpr (POOLS == 1 || POOLS == 3) { + constexpr int ZS = POOLS == 1 ? 1 : 2; + using V = Multi<2, T>; + const V absorbed = absorbed_of(y, event, V(alpha, 0), p.off_axis ? V(in[y].b0, 1) : V(T(0.0f), -1)); + const T met_z = rdot(bz[ZS][y], z[ZS][y]); + g_alpha += absorbed.d[0] * met_z; + g_b0_pulse[y] += absorbed.d[1] * met_z; + bz[ZS][y] = {absorbed.v * bz[ZS][y].r, absorbed.v * bz[ZS][y].i}; + } + // The cotangent through the rotation: R^H. +#pragma unroll + for (int s = 0; s < NS; ++s) { + const int zi = s; + const Cx L[3] = {bplus[s][y], bminus[s][y], bz[zi][y]}; + Cx back_out[3]; +#pragma unroll + for (int j = 0; j < 3; ++j) { + back_out[j] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int i = 0; i < 3; ++i) { + const Cx term = cmul(cconj(R[3 * i + j]), L[i]); + back_out[j] = {back_out[j].r + term.r, back_out[j].i + term.i}; + } + } + bplus[s][y] = back_out[0]; + bminus[s][y] = back_out[1]; + bz[zi][y] = back_out[2]; + } + flip_gradient[y] = g_alpha * transmit; + phase_gradient[y] = g_phi; + if (p.shimmed) { + const long long cell = static_cast(p.shim_index[event]) * p.atom_count + atom[y]; + const T b1_sum = lanes_sum(g_alpha * flip, 1, width); + const T phase_sum = lanes_sum(g_phi, 1, width); + if (state == 0 && active[y]) { + atomic_add(v.grad_tissue, v.grad_tissue_t, 3 * n_atoms + cell, b1_sum); + atomic_add(v.grad_tissue, v.grad_tissue_t, (3 + v.shim_rows) * n_atoms + cell, phase_sum); + } + } else { + g_b1[y] += g_alpha * flip; + g_b1_phase[y] += g_phi; + } + } + event_gradient(v.grad_flip, v.grad_flip_t, event, flip_gradient); + event_gradient(v.grad_phase, v.grad_phase_t, event, phase_gradient); + } + if ((act & 32) && kind == 2) { + const int out = p.output_index[event]; + if (out >= 0) { + T phase_gradient[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const int e = uniform ? base + event : (active[y] ? train[y] * p.event_count + event : event); + T ac, as; + event_phase(p, e, ac, as); + phase_gradient[y] = 0.0f; + if (state == 0 && active[y]) { + const long long at = static_cast(problem[y]) * p.output_count + out; + const Cx seed = {T(v.grad_output_real[at]), T(v.grad_output_imag[at])}; + Cx read = plus[0][y]; + if constexpr (POOLS >= 2) read = {read.r + plus[1][y].r, read.i + plus[1][y].i}; + const Cx demodulated = cmul(read, Cx{ac, -as}); + g_m0[y] += rdot(seed, demodulated); + phase_gradient[y] = m0[y] * rdot(seed, Cx{demodulated.i, -demodulated.r}); + const Cx back_seed = cmul(Cx{m0[y] * ac, m0[y] * as}, seed); + bplus[0][y] = {bplus[0][y].r + back_seed.r, bplus[0][y].i + back_seed.i}; + if constexpr (POOLS >= 2) { + bplus[1][y] = {bplus[1][y].r + back_seed.r, bplus[1][y].i + back_seed.i}; + } + } + } + event_gradient(v.grad_phase, v.grad_phase_t, event, phase_gradient); + } + } + if (act & 1) { +#pragma unroll + for (int s = 0; s < NS; ++s) shift_adjoint(s); + } + // The interval: what it met, its length's gradient, and the + // cotangent through it. + T duration_gradient[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const Interval& o = op[y]; + Met now; + if constexpr (NS == 1) { + const Cx sp = entry[0][y], sm = entry[1][y]; + now.t[0] = {sp.r * bplus[0][y].r + sp.i * bplus[0][y].i + sm.r * bminus[0][y].r + sm.i * bminus[0][y].i, + sp.i * bplus[0][y].r - sp.r * bplus[0][y].i - sm.i * bminus[0][y].r + sm.r * bminus[0][y].i}; + } else { +#pragma unroll + for (int i = 0; i < 2; ++i) { +#pragma unroll + for (int j = 0; j < 2; ++j) { + const Cx sp = entry[j][y], sm = entry[2 + j][y]; + const Cx lp = bplus[i][y], lm = bminus[i][y]; + now.t[2 * i + j] = {sp.r * lp.r + sp.i * lp.i + sm.r * lm.r + sm.i * lm.i, + sp.i * lp.r - sp.r * lp.i - sm.i * lm.r + sm.r * lm.i}; + } + } + } +#pragma unroll + for (int i = 0; i < NZ; ++i) { +#pragma unroll + for (int j = 0; j < NZ; ++j) { + const Cx s = entry[2 * NS + j][y], l = bz[i][y]; + now.l[i * NZ + j] = {s.r * l.r + s.i * l.i, s.i * l.r - s.r * l.i}; + } + now.g[i] = state == 0 ? bz[i][y].r : T(0.0f); + } + duration_gradient[y] = contract(slope[y], now); + if constexpr (num::is_dual::value) { +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) { + met_along[y].t[k] = {met_along[y].t[k].r + along * now.t[k].r.v, met_along[y].t[k].i + along * now.t[k].i.v}; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) { + met_along[y].l[k] = {met_along[y].l[k].r + along * now.l[k].r.v, met_along[y].l[k].i + along * now.l[k].i.v}; + } +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) met_along[y].g[k] += along * now.g[k].v; + } +#pragma unroll + for (int k = 0; k < Interval::NT; ++k) { + met[y].t[k] = {met[y].t[k].r + now.t[k].r, met[y].t[k].i + now.t[k].i}; + } +#pragma unroll + for (int k = 0; k < Interval::NL; ++k) { + met[y].l[k] = {met[y].l[k].r + now.l[k].r, met[y].l[k].i + now.l[k].i}; + } +#pragma unroll + for (int k = 0; k < Interval::NG; ++k) met[y].g[k] += now.g[k]; + // Cotangent through the operator: its conjugate transpose. + if constexpr (NS == 1) { + bplus[0][y] = cmul(cconj(o.t[0]), bplus[0][y]); + bminus[0][y] = cmul(o.t[0], bminus[0][y]); + } else { + const Cx lp0 = bplus[0][y], lp1 = bplus[1][y], lm0 = bminus[0][y], lm1 = bminus[1][y]; + const Cx a = cmul(cconj(o.t[0]), lp0), b = cmul(cconj(o.t[2]), lp1); + const Cx c = cmul(cconj(o.t[1]), lp0), d = cmul(cconj(o.t[3]), lp1); + bplus[0][y] = {a.r + b.r, a.i + b.i}; + bplus[1][y] = {c.r + d.r, c.i + d.i}; + const Cx e = cmul(o.t[0], lm0), f = cmul(o.t[2], lm1), g = cmul(o.t[1], lm0), h = cmul(o.t[3], lm1); + bminus[0][y] = {e.r + f.r, e.i + f.i}; + bminus[1][y] = {g.r + h.r, g.i + h.i}; + } + Cx next[NZ]; +#pragma unroll + for (int j = 0; j < NZ; ++j) { + next[j] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int i = 0; i < NZ; ++i) { + const Cx term = cmul(cconj(o.l[i * NZ + j]), bz[i][y]); + next[j] = {next[j].r + term.r, next[j].i + term.i}; + } + } +#pragma unroll + for (int j = 0; j < NZ; ++j) bz[j][y] = next[j]; + } + event_gradient(v.grad_duration, v.grad_duration_t, event, duration_gradient); + } + } + close_interval(); + + // The fractions also set where each pool starts. + T g_fb[Y], g_fc[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + g_fb[y] = g_fc[y] = 0.0f; + if (state == 0) { + if constexpr (NZ >= 2) g_fb[y] = bz[1][y].r - bz[0][y].r; + if constexpr (NZ >= 3) g_fc[y] = bz[2][y].r - bz[0][y].r; + } + } + auto store_row = [&](int y, int row, T lane_value) { + const T sum = lanes_sum(lane_value, 1, width); + if (state == 0 && active[y]) atomic_add(v.grad_tissue, v.grad_tissue_t, row * n_atoms + atom[y], sum); + }; + auto from_rate = [&](T gradient, T rate) { return gradient * (-rate * rate * (1.0f / 1000.0f)); }; +#pragma unroll + for (int y = 0; y < Y; ++y) { + store_row(y, 0, from_rate(grad[D_R1][y], in[y].r1)); + store_row(y, 1, from_rate(grad[D_R2][y], in[y].r2)); + store_row(y, 2, g_m0[y]); + if (!p.shimmed) { + store_row(y, 3, g_b1[y]); + store_row(y, 4, g_b1_phase[y]); + } + store_row(y, 5 + past_transmit, grad[D_B0][y] + g_b0_pulse[y]); + store_row(y, 6 + past_transmit, g_inversion[y]); + store_row(y, 7 + past_transmit, grad[D_DIFFUSION][y]); + store_row(y, 8 + past_transmit, grad[D_VELOCITY][y]); + if constexpr (POOLS == 1) { + store_row(y, 9 + past_transmit, grad[D_FB][y] + g_fb[y]); + store_row(y, 10 + past_transmit, grad[D_XB][y]); + store_row(y, 11 + past_transmit, from_rate(grad[D_R1B][y], in[y].r1b)); + } + if constexpr (POOLS >= 2) { + store_row(y, 12 + past_transmit, grad[D_FB][y] + g_fb[y]); + store_row(y, 13 + past_transmit, grad[D_XB][y]); + store_row(y, 14 + past_transmit, from_rate(grad[D_R1B][y], in[y].r1b)); + store_row(y, 15 + past_transmit, from_rate(grad[D_R2B][y], in[y].r2b)); + store_row(y, 16 + past_transmit, grad[D_SHIFT][y]); + } + if constexpr (POOLS == 3) { + store_row(y, 9 + past_transmit, grad[D_FC][y] + g_fc[y]); + store_row(y, 10 + past_transmit, grad[D_XC][y]); + store_row(y, 11 + past_transmit, from_rate(grad[D_R1C][y], in[y].r1c)); + } + } +} + +} // namespace epg_vjp diff --git a/src/blochsim/_layout_numbers.hpp b/src/blochsim/_layout_numbers.hpp index 28076aaf..d49e3263 100644 --- a/src/blochsim/_layout_numbers.hpp +++ b/src/blochsim/_layout_numbers.hpp @@ -6,6 +6,8 @@ #include +#include + namespace num { template @@ -43,6 +45,7 @@ __device__ __forceinline__ F tangent(DualT x) { // Equal as a cache key: the value and, for a dual, the direction too. __device__ __forceinline__ bool same(float a, float b) { return a == b; } +__device__ __forceinline__ bool same(double a, double b) { return a == b; } template __device__ __forceinline__ bool same(DualT a, DualT b) { return a.v == b.v && a.d == b.d; @@ -98,6 +101,11 @@ __device__ __forceinline__ DualT operator/(DualT a, F b) { return {a.v / b, a.d / b}; } template +__device__ __forceinline__ DualT operator/(F a, DualT b) { + const F q = a / b.v; + return {q, -q * b.d / b.v}; +} +template __device__ __forceinline__ DualT& operator+=(DualT& a, DualT b) { return a = a + b; } @@ -180,8 +188,10 @@ __device__ __forceinline__ DualT max_(DualT a, DualT b) { } __device__ __forceinline__ float abs_(float a) { return fabsf(a); } __device__ __forceinline__ double abs_(double a) { return fabs(a); } +// |x| carries no derivative at the origin. template __device__ __forceinline__ DualT abs_(DualT a) { + if (a.v == F(0)) return {a.v, F(0)}; return a.v < F(0) ? -a : a; } @@ -190,6 +200,12 @@ __device__ __forceinline__ Dual64 acos_(Dual64 x) { return {acos(x.v), -x.d / sqrt(1.0 - x.v * x.v)}; } __device__ __forceinline__ double cos_(double x) { return cos(x); } +__device__ __forceinline__ double sin_(double x) { return sin(x); } +__device__ __forceinline__ Dual64 sin_(Dual64 x) { + double s, c; + sincos(x.v, &s, &c); + return {s, c * x.d}; +} __device__ __forceinline__ Dual64 cos_(Dual64 x) { double s, c; sincos(x.v, &s, &c); @@ -250,11 +266,12 @@ __device__ __forceinline__ Dual shfl_from(Dual x, int lane, unsigned mask) { return {__shfl_sync(mask, x.v, lane), __shfl_sync(mask, x.d, lane)}; } -// A value read with its direction where the type carries one. +// A value read with its direction where the type carries one; no direction +// given is a direction of zero. template __device__ __forceinline__ T load(const float* values, const float* tangents, long long at) { if constexpr (is_dual::value) { - return T{__ldg(values + at), __ldg(tangents + at)}; + return T{__ldg(values + at), tangents != nullptr ? __ldg(tangents + at) : 0.0f}; } else { return __ldg(values + at); } @@ -268,4 +285,228 @@ __device__ __forceinline__ Dual rate(Dual ms) { return {r, -1000.0f * ms.d / (ms.v * ms.v)}; } + +// The number a value finally is, through any nesting of duals: what a branch +// decides on. +__device__ __forceinline__ float primal(float x) { return x; } +__device__ __forceinline__ double primal(double x) { return x; } +template +__device__ __forceinline__ auto primal(DualT x) { + return primal(x.v); +} + +// A value with its derivatives along K directions at once: what a gradient +// contracts against when many inputs reach one interval's operator. +template +struct Multi { + S v; + S d[K]; + Multi() = default; + __device__ __forceinline__ Multi(float value) : v(S(value)) { +#pragma unroll + for (int k = 0; k < K; ++k) d[k] = S(0.0f); + } + __device__ __forceinline__ Multi(double value) : v(S(value)) { +#pragma unroll + for (int k = 0; k < K; ++k) d[k] = S(0.0f); + } + __device__ __forceinline__ Multi(S value, int direction) : v(value) { +#pragma unroll + for (int k = 0; k < K; ++k) d[k] = S(k == direction ? 1.0f : 0.0f); + } + template + __device__ __forceinline__ explicit Multi(const Multi& other) : v(S(other.v)) { + static_assert(J == K, "a conversion keeps the directions"); +#pragma unroll + for (int k = 0; k < K; ++k) d[k] = S(other.d[k]); + } +}; +template +struct is_dual> { + static constexpr bool value = true; +}; +template +__device__ __forceinline__ auto primal(const Multi& x) { + return primal(x.v); +} +template +__device__ __forceinline__ S value(const Multi& x) { + return x.v; +} + +#define NUM_MULTI_BINARY(OP, VALUE, TANGENT) \ + template \ + __device__ __forceinline__ Multi operator OP(const Multi& a, const Multi& b) { \ + Multi r; \ + r.v = VALUE; \ + _Pragma("unroll") for (int k = 0; k < K; ++k) r.d[k] = TANGENT; \ + return r; \ + } +NUM_MULTI_BINARY(+, a.v + b.v, a.d[k] + b.d[k]) +NUM_MULTI_BINARY(-, a.v - b.v, a.d[k] - b.d[k]) +NUM_MULTI_BINARY(*, a.v * b.v, a.d[k] * b.v + a.v * b.d[k]) +#undef NUM_MULTI_BINARY +template +__device__ __forceinline__ Multi operator/(const Multi& a, const Multi& b) { + Multi r; + r.v = a.v / b.v; +#pragma unroll + for (int k = 0; k < K; ++k) r.d[k] = (a.d[k] - r.v * b.d[k]) / b.v; + return r; +} +template +__device__ __forceinline__ Multi operator-(const Multi& a) { + Multi r; + r.v = -a.v; +#pragma unroll + for (int k = 0; k < K; ++k) r.d[k] = -a.d[k]; + return r; +} +// With a constant, of the innermost kind or of the direction's own. +template +using primal_of = decltype(primal(S())); +#define NUM_MULTI_SCALAR(OP, LEFT_V, LEFT_D, RIGHT_V, RIGHT_D) \ + template \ + __device__ __forceinline__ Multi operator OP(const Multi& a, primal_of b) { \ + Multi r; \ + r.v = LEFT_V; \ + _Pragma("unroll") for (int k = 0; k < K; ++k) r.d[k] = LEFT_D; \ + return r; \ + } \ + template \ + __device__ __forceinline__ Multi operator OP(primal_of b, const Multi& a) { \ + Multi r; \ + r.v = RIGHT_V; \ + _Pragma("unroll") for (int k = 0; k < K; ++k) r.d[k] = RIGHT_D; \ + return r; \ + } +NUM_MULTI_SCALAR(+, a.v + b, a.d[k], b + a.v, a.d[k]) +NUM_MULTI_SCALAR(-, a.v - b, a.d[k], b - a.v, -a.d[k]) +NUM_MULTI_SCALAR(*, a.v * b, a.d[k] * b, b * a.v, b * a.d[k]) +#undef NUM_MULTI_SCALAR +// With a value of the direction's own kind, where that kind is itself a dual. +#define NUM_MULTI_SAME(OP, LEFT_V, LEFT_D, RIGHT_V, RIGHT_D) \ + template ::value, int>::type = 0> \ + __device__ __forceinline__ Multi operator OP(const Multi& a, const S& b) { \ + Multi r; \ + r.v = LEFT_V; \ + _Pragma("unroll") for (int k = 0; k < K; ++k) r.d[k] = LEFT_D; \ + return r; \ + } \ + template ::value, int>::type = 0> \ + __device__ __forceinline__ Multi operator OP(const S& b, const Multi& a) { \ + Multi r; \ + r.v = RIGHT_V; \ + _Pragma("unroll") for (int k = 0; k < K; ++k) r.d[k] = RIGHT_D; \ + return r; \ + } +NUM_MULTI_SAME(+, a.v + b, a.d[k], b + a.v, a.d[k]) +NUM_MULTI_SAME(-, a.v - b, a.d[k], b - a.v, -a.d[k]) +NUM_MULTI_SAME(*, a.v * b, a.d[k] * b, b * a.v, b * a.d[k]) +#undef NUM_MULTI_SAME +template +__device__ __forceinline__ Multi operator/(const Multi& a, primal_of b) { + Multi r; + r.v = a.v / b; +#pragma unroll + for (int k = 0; k < K; ++k) r.d[k] = a.d[k] / b; + return r; +} +template +__device__ __forceinline__ Multi& operator+=(Multi& a, const Multi& b) { + return a = a + b; +} +template +__device__ __forceinline__ Multi& operator-=(Multi& a, const Multi& b) { + return a = a - b; +} +template +__device__ __forceinline__ Multi& operator*=(Multi& a, const Multi& b) { + return a = a * b; +} +template +__device__ __forceinline__ Multi& operator*=(Multi& a, primal_of b) { + return a = a * b; +} +template +__device__ __forceinline__ bool operator<(const Multi& a, primal_of b) { + return primal(a) < b; +} +template +__device__ __forceinline__ bool operator>(const Multi& a, primal_of b) { + return primal(a) > b; +} +template +__device__ __forceinline__ bool same(const Multi& a, const Multi& b) { + bool equal = same(a.v, b.v); +#pragma unroll + for (int k = 0; k < K; ++k) equal = equal && same(a.d[k], b.d[k]); + return equal; +} + +// A function of one argument along every direction: f(v) and f'(v) d. +template +__device__ __forceinline__ Multi chain(const Multi& x, Value value, Slope slope) { + Multi r; + r.v = value; +#pragma unroll + for (int k = 0; k < K; ++k) r.d[k] = slope * x.d[k]; + return r; +} +template +__device__ __forceinline__ Multi fma_(const Multi& a, const Multi& b, const Multi& c) { + return a * b + c; +} +template +__device__ __forceinline__ Multi exp_(const Multi& x) { + const S e = exp_(x.v); + return chain(x, e, e); +} +template +__device__ __forceinline__ Multi div_(const Multi& a, const Multi& b) { + Multi r; + r.v = div_(a.v, b.v); +#pragma unroll + for (int k = 0; k < K; ++k) r.d[k] = div_(a.d[k] - r.v * b.d[k], b.v); + return r; +} +template +__device__ __forceinline__ Multi sqrt_(const Multi& x) { + const S r = sqrt_(x.v); + return chain(x, r, div_(S(0.5f), r)); +} +template +__device__ __forceinline__ Multi min_(const Multi& a, const Multi& b) { + return primal(b) < primal(a) ? b : a; +} +template +__device__ __forceinline__ Multi max_(const Multi& a, const Multi& b) { + return primal(b) > primal(a) ? b : a; +} +template +__device__ __forceinline__ Multi abs_(const Multi& a) { + if (primal(a) == 0) return Multi(a.v, -1); + return primal(a) < 0 ? -a : a; +} +template +__device__ __forceinline__ void sincos_(const Multi& x, Multi& s, Multi& c) { + S sv, cv; + sincos_(x.v, sv, cv); + s = chain(x, sv, cv); + c = chain(x, cv, -sv); +} +template +__device__ __forceinline__ Multi acos_(const Multi& x) { + return chain(x, acos_(x.v), -1.0 / sqrt_(1.0 - x.v * x.v)); +} +template +__device__ __forceinline__ Multi cos_(const Multi& x) { + return chain(x, cos_(x.v), -sin_(x.v)); +} +template +__device__ __forceinline__ Multi rate(const Multi& ms) { + const S r = rate(ms.v); + return chain(ms, r, -r * r * (1.0f / 1000.0f)); +} + } // namespace num diff --git a/src/blochsim/_layout_real_vjp.cu b/src/blochsim/_layout_real_vjp.cu new file mode 100644 index 00000000..da6aca5c --- /dev/null +++ b/src/blochsim/_layout_real_vjp.cu @@ -0,0 +1,44 @@ +// The real adjoint and its derivative along a direction in their layouts. +#include "_layout.hpp" + +namespace blochsim_layout { +namespace { + +constexpr int SEGMENT = 4; + +template +__global__ void __launch_bounds__(32 * WARPS) real_adjoint_kernel(layout_real_vjp::Params p) { + extern __shared__ unsigned char segment[]; + layout_real_vjp::real_vjp_loop(p, reinterpret_cast(segment)); +} + +template +int launch_real_adjoint(const layout_real_vjp::Params& p, cudaStream_t stream) { + const int groups = 32 / p.width; + const long long problems = p.problem_end - p.problem_base; + const long long per_block = static_cast(WARPS) * groups * Y; + const unsigned grid = static_cast((problems + per_block - 1) / per_block); + const bool maps = p.transmit || p.density || p.inverting; + const size_t shared = static_cast(SEGMENT) * 3 * Y * 32 * WARPS * sizeof(T); +#define BLOCHSIM_GO(MP, ONE) \ + cudaFuncSetAttribute(real_adjoint_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, \ + static_cast(shared)); \ + real_adjoint_kernel<<>>(p) + if (p.single_train) { + if (maps) { BLOCHSIM_GO(true, true); } else { BLOCHSIM_GO(false, true); } + } else { + if (maps) { BLOCHSIM_GO(true, false); } else { BLOCHSIM_GO(false, false); } + } +#undef BLOCHSIM_GO + return static_cast(cudaGetLastError()); +} + +} // namespace + +int real_vjp(const layout_real_vjp::Params& p, cudaStream_t stream) { return launch_real_adjoint(p, stream); } + +int real_vjp_jvp(const layout_real_vjp::Params& p, cudaStream_t stream) { + return launch_real_adjoint(p, stream); +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_real_vjp.hpp b/src/blochsim/_layout_real_vjp.hpp new file mode 100644 index 00000000..a29837f6 --- /dev/null +++ b/src/blochsim/_layout_real_vjp.hpp @@ -0,0 +1,423 @@ +// The real EPG adjoint over a number type: the vector-Jacobian product at +// float, its derivative along a direction (the Hessian-vector product's +// sweep) at num::Dual. +// +// The forward sweep keeps the state every K events; the reverse sweep +// replays each stretch of K events from its checkpoint into shared memory +// and walks it back. One extra forward pass replaces a trajectory of every +// event through global memory. +#pragma once +#include + +#include "_layout_numbers.hpp" + +namespace layout_real_vjp { + +using num::exp_; +using num::same; + +struct Params { + const float *t1, *t2, *m0, *b1, *inversion_efficiency, *diffusion, *duration, *flip; + const float *d_t1, *d_t2, *d_m0, *d_b1, *d_inversion_efficiency, *d_diffusion, *d_duration, *d_flip; + const int *kind, *output_index, *shim_index; + const unsigned char* action; + const float* grad_output_imag; + // Value and, for a dual sweep, direction of each gradient. + float *grad_tissue, *grad_flip, *grad_duration; + float *grad_tissue_t, *grad_flip_t, *grad_duration_t; + float *trajectory, *trajectory_t; + int problem_base, problem_end, atom_count, train_count, event_count, output_count, state_count, width; + int shim_rows; + bool single_train, atom_stride, shimmed, diffusing, transmit, density, inverting; +}; + +template +__device__ __forceinline__ void atomic_add(float* value, float* tangent, long long at, T x) { + if constexpr (num::is_dual::value) { + atomicAdd(value + at, x.v); + atomicAdd(tangent + at, x.d); + } else { + atomicAdd(value + at, x); + } +} + +template +__device__ __forceinline__ T shfl_xor(T x, int mask) { + if constexpr (num::is_dual::value) { + return T{__shfl_xor_sync(0xffffffffu, x.v, mask), __shfl_xor_sync(0xffffffffu, x.d, mask)}; + } else { + return __shfl_xor_sync(0xffffffffu, x, mask); + } +} + +// Sum over the lanes ``from`` apart and wider, up to a warp. +template +__device__ __forceinline__ T lanes_sum(T x, int from, int to) { + for (int offset = from; offset < to; offset <<= 1) x += shfl_xor(x, offset); + return x; +} + +template +__device__ __forceinline__ void store(float* values, float* tangents, long long at, T x) { + if constexpr (num::is_dual::value) { + values[at] = x.v; + tangents[at] = x.d; + } else { + values[at] = x; + } +} + +template +__device__ __forceinline__ T fetch(const float* values, const float* tangents, long long at) { + if constexpr (num::is_dual::value) { + return T{values[at], tangents[at]}; + } else { + return values[at]; + } +} + +template +__device__ __forceinline__ void real_vjp_loop(const Params& p, T* segment) { + constexpr int PLANES = 3; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const int first = p.problem_base + ((blockIdx.x * (blockDim.x >> 5) + warp) * groups + group) * Y; + const float order = static_cast(state); + const float transverse_weight = order * order + order + 0.3333333333333333f; + const float longitudinal_weight = order * order; + + int problem[Y], atom[Y], train[Y]; + bool active[Y], live[Y]; + T t1[Y], t2[Y], r1[Y], r2[Y], m0[MAPS ? Y : 1], b1[MAPS ? Y : 1], inversion[MAPS ? Y : 1], diffusion[Y]; + T plus[Y], minus[Y], z[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + problem[y] = first + y; + active[y] = problem[y] < p.problem_end; + live[y] = active[y] && state < p.state_count; + atom[y] = problem[y] % p.atom_count; + train[y] = problem[y] / p.atom_count; + const int at = p.atom_stride ? atom[y] : 0; + t1[y] = active[y] ? num::load(p.t1, p.d_t1, atom[y]) : T(1.0f); + t2[y] = active[y] ? num::load(p.t2, p.d_t2, atom[y]) : T(1.0f); + r1[y] = num::rate(t1[y]); + r2[y] = num::rate(t2[y]); + if constexpr (MAPS) { + m0[y] = p.density ? (active[y] ? num::load(p.m0, p.d_m0, at) : T(0.0f)) : T(1.0f); + b1[y] = p.transmit ? (active[y] ? num::load(p.b1, p.d_b1, at) : T(1.0f)) : T(1.0f); + inversion[y] = p.inverting ? (active[y] ? num::load(p.inversion_efficiency, p.d_inversion_efficiency, at) + : T(1.0f)) + : T(1.0f); + } + diffusion[y] = p.diffusing && active[y] ? num::load(p.diffusion, p.d_diffusion, at) : T(0.0f); + plus[y] = minus[y] = 0.0f; + z[y] = state == 0 ? 1.0f : 0.0f; + } + const bool uniform = ONE_TRAIN || (active[0] && train[0] == train[Y - 1]); + const int base = ONE_TRAIN ? 0 : train[0] * p.event_count; + // Every problem of the warp in one train: an event's gradient is one + // sum over the warp and one atomic. + const int first_train = __shfl_sync(0xffffffffu, train[0], 0); + const bool warp_train = ONE_TRAIN || __all_sync(0xffffffffu, uniform && train[0] == first_train); + const bool shared_pulse = uniform && !p.transmit; + + // The interval's factors, kept while it repeats: with diffusion's + // damping of this thread's order, and bare. + T e1c[Y], e2c[Y], bare1[Y], bare2[Y], damp_z[Y], damp_t[Y]; + T last_dt = -1.0f; + auto factors = [&](T dt_shared, int event) { + if (uniform && same(dt_shared, last_dt)) return; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T dt = uniform ? dt_shared + : (active[y] ? num::load(p.duration, p.d_duration, train[y] * p.event_count + event) + : T(0.0f)); + bare1[y] = exp_(-r1[y] * dt); + bare2[y] = exp_(-r2[y] * dt); + damp_z[y] = 1.0f; + damp_t[y] = 1.0f; + if (p.diffusing) { + const T b = diffusion[y] * dt; + damp_z[y] = exp_(-b * longitudinal_weight); + damp_t[y] = exp_(-b * transverse_weight); + } + e1c[y] = bare1[y] * damp_z[y]; + e2c[y] = bare2[y] * damp_t[y]; + } + last_dt = uniform ? dt_shared : T(-1.0f); + }; + auto event_dt = [&](int event) { + return uniform ? num::load(p.duration, p.d_duration, base + event) : T(0.0f); + }; + auto shift = [&](T* pl, T* mi) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T up = num::shfl_up(pl[y], 1, width); + const T down = num::shfl_down(mi[y], 1, width); + const T shifted_plus = (state > 0 && live[y]) ? up : T(0.0f); + const T shifted_minus = (state + 1 < p.state_count && live[y]) ? down : T(0.0f); + pl[y] = state == 0 ? -shifted_minus : shifted_plus; + mi[y] = shifted_minus; + } + }; + // Transpose of the shift: the order-zero refill sends plus's adjoint + // back onto minus at order one. + auto shift_adjoint = [&](T* pl, T* mi) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T carry = -num::shfl(pl[y], 0, width); + const T down = num::shfl_down(pl[y], 1, width); + const T up = num::shfl_up(mi[y], 1, width); + T shifted_minus = (state > 0 && live[y]) ? up : T(0.0f); + if (state == 1 && live[y]) shifted_minus += carry; + pl[y] = (state + 1 < p.state_count && live[y]) ? down : T(0.0f); + mi[y] = shifted_minus; + } + }; + auto flip_of = [&](int y, int event) { + return uniform ? num::load(p.flip, p.d_flip, base + event) + : (active[y] ? num::load(p.flip, p.d_flip, train[y] * p.event_count + event) : T(0.0f)); + }; + auto pulse_b1 = [&](int y, int event) { + T value = MAPS ? b1[MAPS ? y : 0] : T(1.0f); + if (p.shimmed && p.transmit) { + value = active[y] ? num::load(p.b1, p.d_b1, p.shim_index[event] * p.atom_count + atom[y]) : T(1.0f); + } + return value; + }; + // One event forward, from its entry state. + auto forward = [&](int event) { + const T dt_shared = event_dt(event); + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + if (!(uniform && same(dt_shared, T(0.0f)))) { + factors(dt_shared, event); +#pragma unroll + for (int y = 0; y < Y; ++y) { + plus[y] *= e2c[y]; + minus[y] *= e2c[y]; + z[y] = z[y] * e1c[y] + (state == 0 ? 1.0f - bare1[y] : T(0.0f)); + } + } + if (act & 1) shift(plus, minus); + if (kind == 1 && (act & 4)) { +#pragma unroll + for (int y = 0; y < Y; ++y) z[y] = -(MAPS ? inversion[MAPS ? y : 0] : T(1.0f)) * z[y]; + } else if (kind == 1) { + T s_u, c_u; + if (shared_pulse) num::sincos_(num::load(p.flip, p.d_flip, base + event), s_u, c_u); +#pragma unroll + for (int y = 0; y < Y; ++y) { + T s = s_u, c = c_u; + if (!shared_pulse) num::sincos_(flip_of(y, event) * pulse_b1(y, event), s, c); + const T chs = 0.5f * (1.0f + c), shs = 0.5f * (1.0f - c), hs = 0.5f * s; + const T pl = chs * plus[y] + shs * minus[y] - s * z[y]; + const T mi = shs * plus[y] + chs * minus[y] + s * z[y]; + z[y] = hs * plus[y] - hs * minus[y] + c * z[y]; + plus[y] = pl; + minus[y] = mi; + } + } + if (act & 18) shift(plus, minus); + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) plus[y] = minus[y] = 0.0f; + } + }; + const long long stride = static_cast(p.event_count) * PLANES * p.state_count; + auto slot = [&](int y, int checkpoint, int plane) { + return (problem[y] - p.problem_base) * stride + (static_cast(checkpoint) * PLANES + plane) * p.state_count + + state; + }; + const int checkpoints = (p.event_count + K - 1) / K; + #pragma unroll 1 + for (int event = 0; event < p.event_count; ++event) { + if (event % K == 0) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (live[y]) { + store(p.trajectory, p.trajectory_t, slot(y, event / K, 0), plus[y]); + store(p.trajectory, p.trajectory_t, slot(y, event / K, 1), minus[y]); + store(p.trajectory, p.trajectory_t, slot(y, event / K, 2), z[y]); + } + } + } + forward(event); + } + + // The walk back. ``plus`` and friends now hold the state, and the bars + // its cotangent. + T plus_bar[Y], minus_bar[Y], z_bar[Y]; + T g_t1[Y], g_t2[Y], g_m0[Y], g_b1[Y], g_inversion[Y], g_diffusion[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + plus_bar[y] = minus_bar[y] = z_bar[y] = 0.0f; + g_t1[y] = g_t2[y] = g_m0[y] = g_b1[y] = g_inversion[y] = g_diffusion[y] = 0.0f; + } + // An event's gradient, summed over the problems that share its row. + auto event_gradient = [&](float* value, float* tangent, int event, const T* lane_values) { + if (warp_train) { + T sum = lane_values[0]; +#pragma unroll + for (int y = 1; y < Y; ++y) sum += lane_values[y]; + sum = lanes_sum(sum, 1, 32); + if (lane == 0 && active[0]) atomic_add(value, tangent, base + event, sum); + } else { +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T sum = lanes_sum(lane_values[y], 1, width); + if (state == 0 && active[y]) atomic_add(value, tangent, train[y] * p.event_count + event, sum); + } + } + }; + const int threads = blockDim.x; + auto smem = [&](int k, int plane, int y) -> T& { return segment[((k * PLANES + plane) * Y + y) * threads + threadIdx.x]; }; + #pragma unroll 1 + for (int checkpoint = checkpoints - 1; checkpoint >= 0; --checkpoint) { + const int start = checkpoint * K; + const int stop = min(start + K, p.event_count); +#pragma unroll + for (int y = 0; y < Y; ++y) { + plus[y] = live[y] ? fetch(p.trajectory, p.trajectory_t, slot(y, checkpoint, 0)) : T(0.0f); + minus[y] = live[y] ? fetch(p.trajectory, p.trajectory_t, slot(y, checkpoint, 1)) : T(0.0f); + z[y] = live[y] ? fetch(p.trajectory, p.trajectory_t, slot(y, checkpoint, 2)) : T(0.0f); + } + #pragma unroll 1 + for (int event = start; event < stop; ++event) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + smem(event - start, 0, y) = plus[y]; + smem(event - start, 1, y) = minus[y]; + smem(event - start, 2, y) = z[y]; + } + if (event + 1 < stop) forward(event); + } + #pragma unroll 1 + for (int event = stop - 1; event >= start; --event) { + T entry_p[Y], entry_m[Y], entry_z[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + entry_p[y] = smem(event - start, 0, y); + entry_m[y] = smem(event - start, 1, y); + entry_z[y] = smem(event - start, 2, y); + } + const T dt_shared = event_dt(event); + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + factors(dt_shared, event); + // The stage the pulse and the sample see. +#pragma unroll + for (int y = 0; y < Y; ++y) { + plus[y] = entry_p[y] * e2c[y]; + minus[y] = entry_m[y] * e2c[y]; + z[y] = entry_z[y] * e1c[y] + (state == 0 ? 1.0f - bare1[y] : T(0.0f)); + } + if (act & 1) shift(plus, minus); + if (act & 8) { +#pragma unroll + for (int y = 0; y < Y; ++y) plus_bar[y] = minus_bar[y] = 0.0f; + } else if (act & 18) { + shift_adjoint(plus_bar, minus_bar); + } + if (kind == 1 && (act & 4)) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + g_inversion[y] -= z_bar[y] * z[y]; + z_bar[y] = -(MAPS ? inversion[MAPS ? y : 0] : T(1.0f)) * z_bar[y]; + } + } else if (kind == 1) { + T flip_gradient[Y]; + // Several shims give each pulse's transmit gradient its shim's row. + const int shim_row = p.shimmed ? p.shim_index[event] * p.atom_count : 0; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T flip = flip_of(y, event), b1y = pulse_b1(y, event); + T s, c; + num::sincos_(flip * b1y, s, c); + const T chs = 0.5f * (1.0f + c), shs = 0.5f * (1.0f - c), hs = 0.5f * s; + // d/dalpha of each output row, against its cotangent. + const T row_p = hs * minus[y] - hs * plus[y] - c * z[y]; + const T row_m = hs * plus[y] - hs * minus[y] + c * z[y]; + const T row_z = 0.5f * c * plus[y] - 0.5f * c * minus[y] - s * z[y]; + const T alpha_bar = plus_bar[y] * row_p + minus_bar[y] * row_m + z_bar[y] * row_z; + const T pb = chs * plus_bar[y] + shs * minus_bar[y] + hs * z_bar[y]; + const T mb = shs * plus_bar[y] + chs * minus_bar[y] - hs * z_bar[y]; + const T zb = -s * plus_bar[y] + s * minus_bar[y] + c * z_bar[y]; + plus_bar[y] = pb; + minus_bar[y] = mb; + z_bar[y] = zb; + flip_gradient[y] = alpha_bar * b1y; + if (p.shimmed) { + const T sum = lanes_sum(alpha_bar * flip, 1, width); + if (state == 0 && active[y]) { + atomic_add(p.grad_tissue, p.grad_tissue_t, 3LL * p.atom_count + shim_row + atom[y], sum); + } + } else { + g_b1[y] += alpha_bar * flip; + } + } + event_gradient(p.grad_flip, p.grad_flip_t, event, flip_gradient); + } + if ((act & 32) && kind == 2) { + const int out = p.output_index[event]; + if (out >= 0) { +#pragma unroll + for (int y = 0; y < Y; ++y) { + if (state == 0 && active[y]) { + const float seed = p.grad_output_imag[static_cast(problem[y]) * p.output_count + out]; + g_m0[y] += seed * plus[y]; + plus_bar[y] += seed * (MAPS ? m0[MAPS ? y : 0] : T(1.0f)); + } + } + } + } + if (act & 1) shift_adjoint(plus_bar, minus_bar); + T duration_gradient[Y]; +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T cot2 = plus_bar[y] * entry_p[y] + minus_bar[y] * entry_m[y]; + const T cot1 = z_bar[y] * entry_z[y]; + const T ge2 = cot2 * damp_t[y]; + const T ge1 = cot1 * damp_z[y] - (state == 0 ? z_bar[y] : T(0.0f)); + const T dt = uniform ? dt_shared + : (active[y] ? num::load(p.duration, p.d_duration, train[y] * p.event_count + event) + : T(0.0f)); + T spread = 0.0f; + if (p.diffusing) { + spread = cot1 * bare1[y] * damp_z[y] * longitudinal_weight + + cot2 * bare2[y] * damp_t[y] * transverse_weight; + g_diffusion[y] -= spread * dt; + } + g_t1[y] += ge1 * bare1[y] * dt * num::div_(T(1000.0f), t1[y] * t1[y]); + g_t2[y] += ge2 * bare2[y] * dt * num::div_(T(1000.0f), t2[y] * t2[y]); + duration_gradient[y] = -ge1 * r1[y] * bare1[y] - ge2 * r2[y] * bare2[y] - spread * diffusion[y]; + plus_bar[y] *= e2c[y]; + minus_bar[y] *= e2c[y]; + z_bar[y] *= e1c[y]; + } + event_gradient(p.grad_duration, p.grad_duration_t, event, duration_gradient); + } + } + const int past_transmit = 2 * (p.shim_rows - 1); +#pragma unroll + for (int y = 0; y < Y; ++y) { + const T t1_sum = lanes_sum(g_t1[y], 1, width), t2_sum = lanes_sum(g_t2[y], 1, width); + const T m0_sum = lanes_sum(g_m0[y], 1, width), b1_sum = lanes_sum(g_b1[y], 1, width); + const T inv_sum = lanes_sum(g_inversion[y], 1, width), diff_sum = lanes_sum(g_diffusion[y], 1, width); + if (state == 0 && active[y]) { + const long long n = p.atom_count; + atomic_add(p.grad_tissue, p.grad_tissue_t, atom[y], t1_sum); + atomic_add(p.grad_tissue, p.grad_tissue_t, n + atom[y], t2_sum); + atomic_add(p.grad_tissue, p.grad_tissue_t, 2 * n + atom[y], m0_sum); + if (!p.shimmed) atomic_add(p.grad_tissue, p.grad_tissue_t, 3 * n + atom[y], b1_sum); + atomic_add(p.grad_tissue, p.grad_tissue_t, (6 + past_transmit) * n + atom[y], inv_sum); + atomic_add(p.grad_tissue, p.grad_tissue_t, (7 + past_transmit) * n + atom[y], diff_sum); + } + } +} + +} // namespace layout_real_vjp From 471e3c4a595d59c16c7ac9140cd47ae04a7cc04e Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 04:09:58 +0200 Subject: [PATCH 10/16] Drop the per-combination specializations; the layouts cover them The layout kernels serve every switch combination from one compile, so the specialized tile kernels, their census and the table they were built from go. A three-pool layout forms its interval operator out of line, which keeps its compile to under a minute. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 50 +- CMakeLists.txt | 44 +- scripts/kernel_census.py | 130 - src/blochsim/_gpu.cu | 120 - src/blochsim/_gpu_launch.py | 30 +- src/blochsim/_gpu_special.cu.in | 29 - src/blochsim/_layout_complex.hpp | 179 +- src/blochsim/_special.hpp | 50 +- src/blochsim/_specializations.json | 2406 ----------------- ...ized_kernels.py => test_layout_kernels.py} | 32 +- 10 files changed, 149 insertions(+), 2921 deletions(-) delete mode 100644 scripts/kernel_census.py delete mode 100644 src/blochsim/_gpu_special.cu.in delete mode 100644 src/blochsim/_specializations.json rename tests/sequence/{test_specialized_kernels.py => test_layout_kernels.py} (76%) diff --git a/CLAUDE.md b/CLAUDE.md index fa6342b5..1681655f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -75,40 +75,26 @@ problems a program carries are another matter: a thread holds `Y_LANES` of them in registers, set per kernel in `src/blochsim/_lanes.hpp`, so that one reading of each event serves all of them. -**A kernel runs compiled for its own switches where one is listed.** Every -EPG kernel takes its feature switches as arguments, so as compiled once it -carries the registers and code of every term it might evaluate. Each entry of -`src/blochsim/_specializations.json` compiles the same kernel again with those -switches as constants (`_special.hpp`): every function is inlined, a constant -switch folds, and the terms it turns off are never generated. The launcher runs -the entry a launch's switches match exactly, and the kernel compiled for all of -them where none does, so the list decides speed and never correctness. -`scripts/kernel_census.py` records the switches launches use and writes the -list; `_gpu_launch.generic_kernels()` runs a block on the general kernels, which -is how `tests/sequence/test_specialized_kernels.py` holds the two to each other. -Each entry is another compile of its kernel, so the list is most of a CUDA -build's time; `--config-settings=cmake.define.BLOCHSIM_SPECIALIZE=OFF` builds -the general kernels alone. A specialized kernel is compiled for rows of at most -32 state orders, so its shifts are shuffles with no test; a wider launch runs -the general one. - **The EPG kernels are written for their layouts** (`_layout.hpp`), and a -launch whose rows fit a warp runs them ahead of any tile kernel. A layout is what is compiled: the pools, how a -pulse is formed, whether the tissue has per-voxel maps, the problems a thread -holds, one train or several; every other switch is read at run time and -steers whole blocks once per event, so one compile serves every combination -of them. The loops are written once over a number type -(`_layout_numbers.hpp`): at `float` they are the forward simulation, at -`num::Dual` -- a value and its derivative along the direction -- the -Jacobian-vector product. The adjoints are written the same way, so their -derivative along a direction is the same source at `num::Dual`: the forward -sweep keeps the state every few events and the reverse sweep replays each -stretch from it, and an interval's gradient is contracted against its +launch whose rows fit a warp runs them ahead of any tile kernel. A layout is +what is compiled: the pools, how a pulse is formed, whether the tissue has +per-voxel maps, the problems a thread holds, one train or several; every other +switch is read at run time and steers whole blocks once per event, so one +compile serves every combination of them. The loops are written once over a +number type (`_layout_numbers.hpp`): at `float` they are the forward +simulation, at `num::Dual` -- a value and its derivative along the direction +-- the Jacobian-vector product. The adjoints are written the same way, so +their derivative along a direction is the same source at `num::Dual`: the +forward sweep keeps the state every few events and the reverse sweep replays +each stretch from it, and an interval's gradient is contracted against its operator's derivatives -- taken along every tissue input at once by -`num::Multi` -- only when the interval changes. `_gpu_launch.generic_kernels()` turns layouts off -with the specializations, and `layout_launches()` counts them. They exist only -on the card: `_gpu_host` compiles the tile kernels, so the host lane holds -those, not these, to the C++ kernels. +`num::Multi` -- only when the interval changes. Code that runs only when an +interval changes is out of line and reads its switches at run time, which is +what keeps a layout's compile to seconds. `_gpu_launch.generic_kernels()` +turns the layouts off and `layout_launches()` counts them; +`tests/sequence/test_layout_kernels.py` holds the two to each other. They +exist only on the card: `_gpu_host` compiles the tile kernels, so the host +lane holds those, not these, to the C++ kernels. **The EPG kernels index in 32 bits** (`bsk::index_t`), as Triton did for every integer argument that fit. An offset that can pass 2^31 is cast to 64 bits diff --git a/CMakeLists.txt b/CMakeLists.txt index f2b78f6f..28712959 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -109,48 +109,6 @@ if(BLOCHSIM_CUDA) list(APPEND _blochsim_gpu_sources "${_source}") endforeach() - # The kernels again, each for the combinations of its feature switches - # listed in _specializations.json: a switch fixed at compile time folds, - # and the terms it turns off are never generated. The launcher runs one - # where a launch's switches match it and the kernel above where none does. - # OFF builds the kernels above alone, which is the quick build for work on - # them. - option(BLOCHSIM_SPECIALIZE "Compile the kernels for the switch combinations listed" ON) - set(_blochsim_special_dir "${CMAKE_CURRENT_BINARY_DIR}/special") - set(_blochsim_special_table "") - if(BLOCHSIM_SPECIALIZE) - set(_list "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim/_specializations.json") - set_property(DIRECTORY APPEND PROPERTY CMAKE_CONFIGURE_DEPENDS "${_list}") - file(READ "${_list}" _json) - string(JSON _count LENGTH "${_json}" specializations) - if(_count GREATER 0) - math(EXPR _last "${_count} - 1") - foreach(INDEX RANGE ${_last}) - string(JSON KERNEL GET "${_json}" specializations ${INDEX} kernel) - string(JSON _fixed GET "${_json}" specializations ${INDEX} fixed) - string(JSON _switches LENGTH "${_fixed}") - set(FIXED "") - set(_text "") - if(_switches GREATER 0) - math(EXPR _last_switch "${_switches} - 1") - foreach(_j RANGE ${_last_switch}) - string(JSON _name MEMBER "${_fixed}" ${_j}) - string(JSON _value GET "${_fixed}" ${_name}) - string(APPEND FIXED ", BLOCHSIM_PARAM(${_name}), ${_value}") - string(APPEND _text "${_name}=${_value},") - endforeach() - endif() - set(_source "${_blochsim_special_dir}/special${INDEX}.cu") - configure_file(src/blochsim/_gpu_special.cu.in "${_source}" @ONLY) - list(APPEND _blochsim_gpu_sources "${_source}") - string(APPEND _blochsim_special_table " X(${INDEX}, ${KERNEL}, \"${_text}\") \\\n") - endforeach() - endif() - endif() - file(CONFIGURE OUTPUT "${_blochsim_special_dir}/_special_table.hpp" - CONTENT "// Written by CMake from _specializations.json: X(index, kernel, \"switch=value,...\").\n#define BLOCHSIM_FOR_EACH_SPECIAL(X) \\\n${_blochsim_special_table}\n" - @ONLY) - # The EPG kernels written for their layouts (_layout.hpp): a unit per # kernel and pool layout, so they compile side by side. list(APPEND _blochsim_gpu_sources src/blochsim/_layout.cu src/blochsim/_layout_real.cu @@ -185,7 +143,7 @@ if(BLOCHSIM_CUDA) WITH_SOABI ${_blochsim_gpu_sources} ) - target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim" "${_blochsim_special_dir}") + target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim") # The runtime is linked in, so a machine needs the driver and nothing else. target_link_libraries(_gpu PRIVATE CUDA::cudart_static) # Warning 221 is a double constant that rounds to zero in float, which the diff --git a/scripts/kernel_census.py b/scripts/kernel_census.py deleted file mode 100644 index bb97bb1b..00000000 --- a/scripts/kernel_census.py +++ /dev/null @@ -1,130 +0,0 @@ -"""Which feature switches the GPU kernels are launched with, and the list to compile. - -As a pytest plugin it records, for every launch on a card, the kernel and the -values of its feature switches, and writes them out when the session ends:: - - PYTHONPATH=scripts pytest -p kernel_census tests/ --census census.json - -Run as a script it turns one or more such records into the entries of -``src/blochsim/_specializations.json``: one per distinct combination of a -kernel's switches, whatever else differed between the launches:: - - python scripts/kernel_census.py census.json > src/blochsim/_specializations.json -""" - -from __future__ import annotations - -import collections -import json -import sys - -#: The switches a kernel is compiled for when it is specialized: every argument -#: that only turns terms on and off. Tile sizes are left out; they shape the -#: launch, not the code. -SWITCHES = ( - "single_train", - "atom_stride", - "shimmed", - "profiled", - "dynamic", - "broadened", - "pools", - "narrow", - "tabulated", - "off_axis", - "moving", - "diffusing", - "transmit", - "density", - "inverting", - "recording", - "directed", -) - -#: The kernels specialized. The pooled kernels' tiles are their pools, and the -#: three-pool tables and PERK have no switches worth a kernel of their own. -SPECIALIZED = ( - "_epg_kernel", - "_epg_jvp_kernel", - "_epg_vjp_kernel", - "_epg_vjp_jvp_kernel", - "_epg_real_kernel", - "_epg_real_jvp_kernel", - "_epg_real_vjp_kernel", - "_epg_real_vjp_jvp_kernel", -) - -#: Switches with which a kernel is left to its general build. Fixing them -#: does not shrink the second-order kernel's compile but multiplies it: one such -#: entry took as long to compile as a dozen without, and they are rare. -COSTLY = { - "_epg_vjp_jvp_kernel": ( - "profiled", - "moving", - "pools", - "dynamic", - "broadened", - "narrow", - "tabulated", - ), -} - -_seen: dict[str, collections.Counter] = collections.defaultdict(collections.Counter) - - -def pytest_addoption(parser) -> None: - parser.addoption( - "--census", default="census.json", help="where to write the census" - ) - - -def pytest_configure(config) -> None: - from blochsim import _gpu_launch - - original = _gpu_launch.Kernel.launch - - def recorded(self, grid, args, kwargs): - names, _ = _gpu_launch._signature(self.name) - values = dict( - zip( - names, - list(args) + [kwargs.get(n) for n in names[len(args) :]], - strict=True, - ) - ) - fixed = {n: int(values[n]) for n in SWITCHES if n in values} - _seen[self.name][json.dumps(fixed, sort_keys=True)] += 1 - return original(self, grid, args, kwargs) - - _gpu_launch.Kernel.launch = recorded - - -def pytest_unconfigure(config) -> None: - census = {name: dict(counts) for name, counts in sorted(_seen.items())} - with open(config.getoption("--census"), "w") as stream: - json.dump(census, stream, indent=1, sort_keys=True) - - -def entries(*censuses: dict) -> list[dict]: - """One entry per distinct combination of a specialized kernel's switches. - - A combination turning on a switch :data:`COSTLY` names for its kernel is - left out. - """ - combinations = { - (name, key) - for census in censuses - for name, counts in census.items() - if name in SPECIALIZED - for key in counts - if not any(json.loads(key).get(switch) for switch in COSTLY.get(name, ())) - } - return [ - {"kernel": name, "fixed": json.loads(key)} for name, key in sorted(combinations) - ] - - -if __name__ == "__main__": - loaded = [json.load(open(path)) for path in sys.argv[1:]] - json.dump({"specializations": entries(*loaded)}, sys.stdout, indent=1) - sys.stdout.write("\n") diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index eca13e5d..a9ac7157 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -11,8 +11,6 @@ #include "_launch.hpp" #include "_layout.hpp" -#include "_special.hpp" -#include "_special_table.hpp" #include @@ -26,11 +24,6 @@ BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_ENTRY) #undef BLOCHSIM_DEVICE_ENTRY -#define BLOCHSIM_SPECIAL_DECLARATION(index, kernel, fixed) \ - __global__ void special##index(bsk::Arguments arguments, int z); -BLOCHSIM_FOR_EACH_SPECIAL(BLOCHSIM_SPECIAL_DECLARATION) -#undef BLOCHSIM_SPECIAL_DECLARATION - namespace { // Each kernel bounded to 256 threads, then to 1024. @@ -40,75 +33,10 @@ namespace { const void* const KERNEL_FUNCTIONS[][2] = {BLOCHSIM_FOR_EACH_KERNEL(BLOCHSIM_DEVICE_POINTER)}; #undef BLOCHSIM_DEVICE_POINTER -// The kernels compiled for one combination of their switches, as CMake listed -// them, and parsed once into the arguments each fixes. -struct SpecialSource { - const char* kernel; - const char* fixed; - const void* function; -}; - -#define BLOCHSIM_SPECIAL_SOURCE(index, kernel, fixed) \ - {#kernel, fixed, reinterpret_cast(&special##index)}, -const SpecialSource SPECIAL_SOURCES[] = { - BLOCHSIM_FOR_EACH_SPECIAL(BLOCHSIM_SPECIAL_SOURCE){nullptr, nullptr, nullptr}}; -#undef BLOCHSIM_SPECIAL_SOURCE - -struct Special { - std::vector> fixed; - const void* function; -}; - -const std::vector>& specials() { - static const std::vector> table = [] { - std::vector> out(sizeof(bsk::KERNELS) / sizeof(bsk::KERNELS[0])); - for (const SpecialSource& source : SPECIAL_SOURCES) { - if (source.kernel == nullptr) { - break; - } - const int kernel = blochsim_launch::find_kernel(source.kernel); - Special special{{}, source.function}; - const std::string text = source.fixed; - std::size_t start = 0; - while (start < text.size()) { - const std::size_t end = text.find(',', start); - const std::string pair = text.substr(start, end - start); - const std::size_t equals = pair.find('='); - special.fixed.emplace_back(bsk::param_index(kernel, pair.substr(0, equals).c_str()), - std::stoll(pair.substr(equals + 1))); - start = end + 1; - } - out[kernel].push_back(std::move(special)); - } - return out; - }(); - return table; -} - -// Whether a launch may run a specialized kernel, and how many have. -bool specializing = true; -unsigned long long specialized_launches = 0; // Whether a launch may run its kernel's layout (_layout.hpp), and how many have. bool laying_out = true; unsigned long long layout_launches = 0; -// The specialized kernel whose fixed switches this launch matches, if any. -const void* matching(const blochsim_launch::Launch& request) { - for (const Special& special : specials()[request.kernel]) { - bool match = true; - for (const auto& [index, value] : special.fixed) { - if (request.arguments.a[index].i != value) { - match = false; - break; - } - } - if (match) { - return special.function; - } - } - return nullptr; -} - PyObject* cuda_error(cudaError_t status, const char* what) { PyErr_Format(PyExc_RuntimeError, "%s: %s", what, cudaGetErrorString(status)); return nullptr; @@ -171,15 +99,7 @@ PyObject* launch(PyObject*, PyObject* args) { sizeof(unsigned long long) * (threads.x * threads.y + bsk::MAX_Z * bsk::MAX_Z); void* parameters[] = {&request.arguments, &request.z}; const int bounded = threads.x * threads.y > 256 ? 1 : 0; - // A specialized kernel is compiled for 256 threads and for rows a warp - // wide; a wider block or row runs the kernel compiled for every combination. const void* function = KERNEL_FUNCTIONS[request.kernel][bounded]; - if (specializing && !bounded && threads.x <= 32) { - if (const void* special = matching(request)) { - function = special; - ++specialized_launches; - } - } status = cudaLaunchKernel(function, blocks, threads, parameters, shared, reinterpret_cast(stream)); if (previous != device) { @@ -191,40 +111,6 @@ PyObject* launch(PyObject*, PyObject* args) { Py_RETURN_NONE; } -PyObject* specializations(PyObject*, PyObject*) { - PyObject* out = PyList_New(0); - if (out == nullptr) { - return nullptr; - } - for (const SpecialSource& source : SPECIAL_SOURCES) { - if (source.kernel == nullptr) { - break; - } - PyObject* entry = Py_BuildValue("(ss)", source.kernel, source.fixed); - if (entry == nullptr || PyList_Append(out, entry) < 0) { - Py_XDECREF(entry); - Py_DECREF(out); - return nullptr; - } - Py_DECREF(entry); - } - return out; -} - -PyObject* use_specializations(PyObject*, PyObject* args) { - int on = 1; - if (!PyArg_ParseTuple(args, "p", &on)) { - return nullptr; - } - const bool previous = specializing; - specializing = on != 0; - return PyBool_FromLong(previous); -} - -PyObject* specialized_launch_count(PyObject*, PyObject*) { - return PyLong_FromUnsignedLongLong(specialized_launches); -} - PyObject* use_layouts(PyObject*, PyObject* args) { int on = 1; if (!PyArg_ParseTuple(args, "p", &on)) { @@ -244,12 +130,6 @@ PyMethodDef METHODS[] = { "Each kernel's parameter names and kinds."}, {"launch", launch, METH_VARARGS, "Queue a kernel over a grid of programs on a device's stream."}, - {"specializations", specializations, METH_NOARGS, - "Each specialized kernel, and the switches it was compiled for."}, - {"use_specializations", use_specializations, METH_VARARGS, - "Whether launches may run specialized kernels; returns the previous setting."}, - {"specialized_launches", specialized_launch_count, METH_NOARGS, - "How many launches have run a specialized kernel."}, {"use_layouts", use_layouts, METH_VARARGS, "Whether launches may run their kernel's layout; returns the previous setting."}, {"layout_launches", layout_launch_count, METH_NOARGS, diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index 19b12a68..d6c0d39c 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -61,32 +61,6 @@ def available() -> bool: return True -def specializations() -> list[tuple[str, dict[str, int]]]: - """Each kernel compiled for one combination of its switches, and the switches. - - Empty where this installation carries no kernels for a card. - """ - if not available(): - return [] - return [ - ( - kernel, - { - name: int(value) - for name, value in ( - pair.split("=") for pair in fixed.split(",") if pair - ) - }, - ) - for kernel, fixed in _module("cuda").specializations() - ] - - -def specialized_launches() -> int: - """How many launches on a card have run a specialized kernel.""" - return _module("cuda").specialized_launches() if available() else 0 - - def layout_launches() -> int: """How many launches on a card have run a kernel written for its layout.""" return _module("cuda").layout_launches() if available() else 0 @@ -94,17 +68,15 @@ def layout_launches() -> int: @contextmanager def generic_kernels() -> Iterator[None]: - """Run every launch inside on the tile kernels compiled for all combinations.""" + """Run every launch inside on the tile kernels rather than on its layout.""" if not available(): yield return module = _module("cuda") - special = module.use_specializations(False) layouts = module.use_layouts(False) try: yield finally: - module.use_specializations(special) module.use_layouts(layouts) diff --git a/src/blochsim/_gpu_special.cu.in b/src/blochsim/_gpu_special.cu.in deleted file mode 100644 index a707a042..00000000 --- a/src/blochsim/_gpu_special.cu.in +++ /dev/null @@ -1,29 +0,0 @@ -// One kernel compiled for one combination of its feature switches. CMake -// writes one of these per entry of _specializations.json: @INDEX@ is the -// entry, @KERNEL@ the kernel, and the switches it fixes follow ``Call``. -#define BLOCHSIM_SIMT 1 -// Launched only on rows a warp wide or narrower; see _tile.hpp. -#define BLOCHSIM_ROWS_IN_A_WARP 1 - -#include "_lanes.hpp" -#define BLOCHSIM_Y_LANES BLOCHSIM_LANES_@KERNEL@ - -#include "_special.hpp" - -namespace { - -constexpr int KERNEL = bsk::kernel_index("@KERNEL@"); -#define BLOCHSIM_PARAM(name) bsk::param_index(KERNEL, #name) - -struct Call { - BSK_HD static void run(const bsk::Arg* a) { bsk::call@KERNEL@(a); } -}; - -using Special = bsk::Fixing; - -} // namespace - -__global__ void __launch_bounds__(256) special@INDEX@(bsk::Arguments arguments, int z) { - bsk::enter(z); - Special::run(arguments.a); -} diff --git a/src/blochsim/_layout_complex.hpp b/src/blochsim/_layout_complex.hpp index 277b1bc0..333036b5 100644 --- a/src/blochsim/_layout_complex.hpp +++ b/src/blochsim/_layout_complex.hpp @@ -585,8 +585,84 @@ struct PoolOperators { T t[Y][NT], e[Y][NE], grow[Y][NG], factor_r[Y], factor_i[Y]; }; -template -__device__ __forceinline__ void pool_operators(const Params& p, const T* dts, const int* rows, +// One problem's operators over one interval at this lane's order, out of +// line: an interval is formed only when it changes, and one copy serves every +// kernel of the layout. The switches are read at run time for the same +// reason. Everything goes in and comes out by value. +template +struct OneOperator { + static constexpr int NT = POOLS >= 2 ? 8 : 2, NE = POOLS == 3 ? 9 : 4, NG = POOLS == 3 ? 3 : 2; + T t[NT], e[NE], grow[NG], factor_r, factor_i, damp_z, wout; +}; +template +struct OneTissue { + T r1, r2, b0, velocity, diffusion, fraction_b, exchange_b, r1_b, free, r2_b, shift_b, fraction_c, exchange_c, r1_c; +}; +struct OperatorGeometry { + float flow_scale, washout_scale; + const float* pool_table; + int atom_count; +}; + +template +__device__ __noinline__ OneOperator one_operator(OperatorGeometry g, int relax_code, T dt, int row, int atom, + float order, OneTissue in, bool three_elsewhere) { + const bool moving = (relax_code & 4) != 0, diffusing = (relax_code & 2) != 0, off_axis = (relax_code & 1) != 0; + OneOperator out; + T wout = 1.0f, turn = 0.0f; + if (moving) { + wout = 1.0f - min_(abs_(in.velocity) * g.washout_scale * dt, T(1.0f)); + turn = in.velocity * g.flow_scale * dt; + } + T damp_z = 1.0f, damp_t = 1.0f; + if (diffusing) { + const T b = in.diffusion * dt; + const float sq = order * order; + damp_z = exp_(-b * sq); + damp_t = exp_(-b * (sq + order + 0.3333333333333333f)); + } + T oc = 1.0f, os = 0.0f; + if (moving || off_axis) sincos_(-TWO_PI * in.b0 * dt - (order + 0.5f) * turn, os, oc); + if constexpr (POOLS == 1) { + const T e2 = exp_(-in.r2 * dt) * wout * damp_t; + out.t[0] = e2 * oc; + out.t[1] = e2 * os; + } else { + T x[8]; + transverse_step(in.r2, in.r2_b, in.exchange_b, in.fraction_b, in.free, in.shift_b, dt, wout, x); + const T rr = damp_t * oc, ri = damp_t * os; +#pragma unroll + for (int k = 0; k < 4; ++k) { + out.t[2 * k] = rr * x[2 * k] - ri * x[2 * k + 1]; + out.t[2 * k + 1] = rr * x[2 * k + 1] + ri * x[2 * k]; + } + } + if constexpr (POOLS == 3) { + if constexpr (MODE == TABLE) { + three_pool_from_table(g.pool_table, row, atom, g.atom_count, dt, wout, in.r1, in.r1_b, in.r1_c, in.exchange_b, + in.exchange_c, in.free, in.fraction_b, in.fraction_c, out.e, out.grow); + } else if (!three_elsewhere) { + three_pool_step(in.r1, in.r1_b, in.r1_c, in.exchange_b, in.exchange_c, in.fraction_b, + in.fraction_c, dt, wout, out.e, out.grow); + } + } else { + two_pool_step(in.r1, in.r1_b, in.exchange_b, in.fraction_b, dt, wout, out.e, out.grow); + } + out.damp_z = damp_z; + out.wout = wout; + out.factor_r = damp_z; + out.factor_i = 0.0f; + if (moving) { + T ts, tc; + sincos_(-order * turn, ts, tc); + out.factor_r = damp_z * tc; + out.factor_i = damp_z * ts; + } + return out; +} + +template +__device__ __forceinline__ void pool_operators(const Params& p, int relax_code, const T* dts, const int* rows, const bool* active, const int* atom, int state, float order, const T* r1, const T* r2, const T* b0, const PoolTissue& tissue, @@ -594,69 +670,46 @@ __device__ __forceinline__ void pool_operators(const Params& p, const T* dts, co // A row of at least Y states spreads the three-pool step over its lanes // where it is formed in double; in float the shuffles cost what they save. const bool SPREAD = POOLS == 3 && MODE == ROOTS && Y > 1 && p.width >= Y; + const bool moving = (relax_code & 4) != 0; + const OperatorGeometry g{p.flow_scale, p.washout_scale, p.pool_table, p.atom_count}; T attenuations[Y], damps[Y]; #pragma unroll for (int y = 0; y < Y; ++y) { - const T dt = dts[y]; const int at = p.atom_stride ? atom[y] : 0; - T wout = 1.0f, turn = 0.0f; - if constexpr (MOVING) { - const T velocity = active[y] ? read(p.velocity, p.d_velocity, at) : T(0.0f); - wout = 1.0f - min_(abs_(velocity) * p.washout_scale * dt, T(1.0f)); - turn = velocity * p.flow_scale * dt; - } - T damp_z = 1.0f, damp_t = 1.0f; - if constexpr (DIFFUSING) { - const T b = (active[y] ? read(p.diffusion, p.d_diffusion, at) : T(0.0f)) * dt; - const float sq = order * order; - damp_z = exp_(-b * sq); - damp_t = exp_(-b * (sq + order + 0.3333333333333333f)); - } - T oc = 1.0f, os = 0.0f; - if constexpr (MOVING || OFF_AXIS) { - T b0y = 0.0f; - if constexpr (OFF_AXIS && MAPS) b0y = b0[y]; - sincos_(-TWO_PI * b0y * dt - (order + 0.5f) * turn, os, oc); - } - if constexpr (POOLS == 1) { - const T e2 = exp_(-r2[y] * dt) * wout * damp_t; - ops.t[y][0] = e2 * oc; - ops.t[y][1] = e2 * os; + OneTissue in; + in.r1 = r1[y]; + in.r2 = r2[y]; + in.b0 = MAPS && p.off_axis ? b0[MAPS ? y : 0] : T(0.0f); + in.velocity = p.moving && active[y] ? read(p.velocity, p.d_velocity, at) : T(0.0f); + in.diffusion = p.diffusing && active[y] ? read(p.diffusion, p.d_diffusion, at) : T(0.0f); + in.fraction_b = tissue.fraction_b[y]; + in.exchange_b = tissue.exchange_b[y]; + in.r1_b = tissue.r1_b[y]; + in.free = tissue.free[y]; + if constexpr (POOLS >= 2) { + in.r2_b = tissue.r2_b[y]; + in.shift_b = tissue.shift_b[y]; } else { - T x[8]; - transverse_step(r2[y], tissue.r2_b[y], tissue.exchange_b[y], tissue.fraction_b[y], tissue.free[y], - tissue.shift_b[y], dt, wout, x); - const T rr = damp_t * oc, ri = damp_t * os; -#pragma unroll - for (int k = 0; k < 4; ++k) { - ops.t[y][2 * k] = rr * x[2 * k] - ri * x[2 * k + 1]; - ops.t[y][2 * k + 1] = rr * x[2 * k + 1] + ri * x[2 * k]; - } + in.r2_b = in.shift_b = 0.0f; } if constexpr (POOLS == 3) { - if constexpr (MODE == TABLE) { - three_pool_from_table(p.pool_table, rows[y], atom[y], p.atom_count, dt, wout, r1[y], - tissue.r1_b[y], tissue.r1_c[y], tissue.exchange_b[y], - tissue.exchange_c[y], tissue.free[y], tissue.fraction_b[y], - tissue.fraction_c[y], ops.e[y], ops.grow[y]); - } else if (!SPREAD) { - three_pool_step(r1[y], tissue.r1_b[y], tissue.r1_c[y], tissue.exchange_b[y], - tissue.exchange_c[y], tissue.fraction_b[y], - tissue.fraction_c[y], dt, wout, ops.e[y], ops.grow[y]); - } else { - attenuations[y] = wout; - } + in.fraction_c = tissue.fraction_c[y]; + in.exchange_c = tissue.exchange_c[y]; + in.r1_c = tissue.r1_c[y]; } else { - two_pool_step(r1[y], tissue.r1_b[y], tissue.exchange_b[y], tissue.fraction_b[y], dt, wout, ops.e[y], - ops.grow[y]); - } - damps[y] = damp_z; - if constexpr (MOVING) { - T ts, tc; - sincos_(-order * turn, ts, tc); - ops.factor_r[y] = damp_z * tc; - ops.factor_i[y] = damp_z * ts; + in.fraction_c = in.exchange_c = in.r1_c = 0.0f; } + const OneOperator one = one_operator(g, relax_code, dts[y], rows[y], atom[y], order, in, SPREAD); +#pragma unroll + for (int k = 0; k < OneOperator::NT; ++k) ops.t[y][k] = one.t[k]; +#pragma unroll + for (int k = 0; k < OneOperator::NE; ++k) ops.e[y][k] = one.e[k]; +#pragma unroll + for (int k = 0; k < OneOperator::NG; ++k) ops.grow[y][k] = one.grow[k]; + ops.factor_r[y] = one.factor_r; + ops.factor_i[y] = one.factor_i; + damps[y] = one.damp_z; + attenuations[y] = one.wout; } if constexpr (POOLS == 3 && MODE == ROOTS && Y > 1) { if (SPREAD) { @@ -687,7 +740,7 @@ __device__ __forceinline__ void pool_operators(const Params& p, const T* dts, co } } } - if constexpr (!MOVING) { + if (!moving) { #pragma unroll for (int y = 0; y < Y; ++y) { #pragma unroll @@ -1065,16 +1118,8 @@ __device__ __forceinline__ void complex_loop(const Params& p) { rows[y] = 0; if constexpr (MODE == TABLE) rows[y] = uniform ? row_shared : p.duration_row[e]; } - switch (relax_code) { - case 0: pool_operators(EPG_POOL_ARGS); break; - case 1: pool_operators(EPG_POOL_ARGS); break; - case 2: pool_operators(EPG_POOL_ARGS); break; - case 3: pool_operators(EPG_POOL_ARGS); break; - case 4: pool_operators(EPG_POOL_ARGS); break; - case 5: pool_operators(EPG_POOL_ARGS); break; - case 6: pool_operators(EPG_POOL_ARGS); break; - default: pool_operators(EPG_POOL_ARGS); break; - } + pool_operators(p, relax_code, dts, rows, active, atom, state, order, r1, r2, + b0, tissue, ops); last_dt = uniform ? dt_shared : T(-1.0f); last_row = row_shared; } diff --git a/src/blochsim/_special.hpp b/src/blochsim/_special.hpp index 854f1437..850c9b02 100644 --- a/src/blochsim/_special.hpp +++ b/src/blochsim/_special.hpp @@ -1,14 +1,6 @@ -// A kernel with some of its arguments fixed at compile time. -// -// The kernels take every feature switch as an argument and branch on it, so -// one kernel serves every combination and is compiled for the worst of them: -// the registers and code of every term it might evaluate. Called with those -// switches as constants, the same kernel is compiled for one combination -// alone -- every function it calls is inlined, so a constant switch folds and -// the terms it turns off are never generated. _specializations.json lists the -// combinations compiled this way; the launcher runs one where a launch's -// switches match it exactly and the kernel as compiled for all of them where -// none does. +// A kernel's place in the table and an argument's place in its parameter +// list, by name, as constant expressions: the layouts (_layout.hpp) read +// their arguments with them. #pragma once #include "_kernels.hpp" @@ -59,40 +51,4 @@ constexpr int param_index(int kernel, const char* param) { return -1; } -// ``Call::run`` with the arguments at ``Fixed``'s even entries replaced by the -// integers that follow them; the copy and the replaced reads fold away. The -// indices are worked out where a constant expression is host code -- an alias -// at namespace scope -- and ``Call`` is a type, so no device function's -// address is taken there. -template -constexpr bool every_index_named() { - constexpr long long pairs[sizeof...(Fixed) + 1] = {Fixed..., 0}; - for (int k = 0; k + 1 < static_cast(sizeof...(Fixed)); k += 2) { - if (pairs[k] < 0) { - return false; - } - } - return true; -} - -template -struct Fixing { - static_assert(sizeof...(Fixed) % 2 == 0, "fixed arguments come as index, value pairs"); - static_assert(every_index_named(), "a fixed argument names no parameter"); - - BSK_HD static void run(const Arg* in) { - constexpr long long pairs[sizeof...(Fixed) + 1] = {Fixed..., 0}; - Arg a[MAX_ARGUMENTS]; -#pragma unroll - for (int i = 0; i < MAX_ARGUMENTS; ++i) { - a[i] = in[i]; - } -#pragma unroll - for (int k = 0; k + 1 < static_cast(sizeof...(Fixed)); k += 2) { - a[pairs[k]].i = pairs[k + 1]; - } - Call::run(a); - } -}; - } // namespace bsk diff --git a/src/blochsim/_specializations.json b/src/blochsim/_specializations.json deleted file mode 100644 index bf14439d..00000000 --- a/src/blochsim/_specializations.json +++ /dev/null @@ -1,2406 +0,0 @@ -{ - "specializations": [ - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 1, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 1, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 1, - "off_axis": 0, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 1, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 1, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 1, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 1, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 1, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 1, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 1, - "off_axis": 0, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 1, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_jvp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_jvp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 0, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 0 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 0, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 0, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 1, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 0 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 1, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 1, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 0 - } - }, - { - "kernel": "_epg_real_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_vjp_kernel", - "fixed": { - "atom_stride": 0, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 0 - } - }, - { - "kernel": "_epg_real_vjp_kernel", - "fixed": { - "atom_stride": 1, - "density": 0, - "diffusing": 0, - "inverting": 0, - "shimmed": 0, - "single_train": 1, - "transmit": 0 - } - }, - { - "kernel": "_epg_real_vjp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_vjp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 0, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_real_vjp_kernel", - "fixed": { - "atom_stride": 1, - "density": 1, - "diffusing": 1, - "inverting": 1, - "shimmed": 1, - "single_train": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "directed": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 0, - "diffusing": 0, - "directed": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_jvp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "directed": 0, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 1, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 0, - "diffusing": 1, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 0, - "off_axis": 0, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 0, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 1, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 1, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 1, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 2, - "profiled": 1, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 0, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 1, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 0, - "density": 1, - "diffusing": 1, - "dynamic": 1, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 0, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 1, - "off_axis": 0, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 0, - "diffusing": 0, - "dynamic": 0, - "inverting": 0, - "moving": 0, - "narrow": 1, - "off_axis": 0, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 0 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 1, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 1, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 1, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 1, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 1, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 0, - "narrow": 1, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 0, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 0, - "transmit": 1 - } - }, - { - "kernel": "_epg_vjp_kernel", - "fixed": { - "atom_stride": 1, - "broadened": 1, - "density": 1, - "diffusing": 1, - "dynamic": 0, - "inverting": 1, - "moving": 1, - "narrow": 0, - "off_axis": 1, - "pools": 3, - "profiled": 0, - "recording": 1, - "shimmed": 0, - "single_train": 1, - "tabulated": 1, - "transmit": 1 - } - } - ] -} diff --git a/tests/sequence/test_specialized_kernels.py b/tests/sequence/test_layout_kernels.py similarity index 76% rename from tests/sequence/test_specialized_kernels.py rename to tests/sequence/test_layout_kernels.py index 70fc2157..ea9b3fd3 100644 --- a/tests/sequence/test_specialized_kernels.py +++ b/tests/sequence/test_layout_kernels.py @@ -1,6 +1,6 @@ -"""Whether a kernel compiled for its switches or its layout computes what the general one does. +"""Whether a kernel written for its layout computes what the tile kernel does. -A kernel that quietly did not run agrees perfectly, so every case also asserts +A layout that quietly did not run agrees perfectly, so every case also asserts that one did. """ @@ -17,8 +17,8 @@ from blochsim.sequence._simulation import TissueProperties pytestmark = pytest.mark.skipif( - not torch.cuda.is_available() or not _gpu_launch.specializations(), - reason="needs a card and the specialized kernels", + not torch.cuda.is_available() or not _gpu_launch.available(), + reason="needs a card and the kernels compiled for it", ) @@ -75,10 +75,6 @@ def _gradient(phase: float) -> torch.Tensor: return tissue["t2_ms"].grad -def _fast_launches() -> int: - return _gpu_launch.specialized_launches() + _gpu_launch.layout_launches() - - CASES = { "real forward": lambda: _forward(torch.pi / 2), "complex forward": lambda: _forward(0.0), @@ -89,21 +85,21 @@ def _fast_launches() -> int: @pytest.mark.parametrize("case", CASES) -def test_a_specialized_kernel_computes_what_the_general_one_does(case) -> None: +def test_a_layout_computes_what_the_tile_kernel_does(case) -> None: with _gpu_launch.generic_kernels(): - general = CASES[case]() - before = _fast_launches() - special = CASES[case]() + tiled = CASES[case]() + before = _gpu_launch.layout_launches() + laid_out = CASES[case]() - assert _fast_launches() > before - error = (special - general).abs().max() - scale = general.abs().max() + assert _gpu_launch.layout_launches() > before + error = (laid_out - tiled).abs().max() + scale = tiled.abs().max() assert float(error / scale) < 1e-5, f"{float(error):.3e} against {float(scale):.3e}" -def test_the_general_kernels_run_where_asked() -> None: - before = _fast_launches() +def test_the_tile_kernels_run_where_asked() -> None: + before = _gpu_launch.layout_launches() with _gpu_launch.generic_kernels(): _forward(torch.pi / 2) - assert _fast_launches() == before + assert _gpu_launch.layout_launches() == before From a65bd83d8306802ee182ff81b6686b67bcd1cc92 Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 04:38:44 +0200 Subject: [PATCH 11/16] Run the many-pool kernels written for their layouts The forward and its JVP hold a problem's pools in a thread's registers, with lanes along the orders. The adjoint does its own forward sweep, keeping every fourth state, and replays each stretch from it on the way back, so a launch it takes needs no recording first and a buffer of checkpoints rather than the whole trajectory. A table row's cotangent is summed per lane while the row repeats and reduced into the row when it changes; the stretch and those sums are a warp's own scratch in global memory, where they stay in its caches and leave the warp's occupancy to its registers. Co-Authored-By: Claude Opus 5.5 --- CMakeLists.txt | 6 + src/blochsim/_gpu.cu | 16 +- src/blochsim/_gpu_launch.py | 16 + src/blochsim/_layout.cu | 149 +++- src/blochsim/_layout.hpp | 24 +- src/blochsim/_layout_pooled.cu.in | 66 ++ src/blochsim/_layout_pooled.hpp | 1014 +++++++++++++++++++++++++++ src/blochsim/sequence/_pools_gpu.py | 45 +- 8 files changed, 1317 insertions(+), 19 deletions(-) create mode 100644 src/blochsim/_layout_pooled.cu.in create mode 100644 src/blochsim/_layout_pooled.hpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 28712959..4151214e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -138,6 +138,12 @@ if(BLOCHSIM_CUDA) endforeach() endforeach() + foreach(POOLS RANGE 2 8) + set(_source "${CMAKE_CURRENT_BINARY_DIR}/layout/pooled_${POOLS}.cu") + configure_file(src/blochsim/_layout_pooled.cu.in "${_source}" @ONLY) + list(APPEND _blochsim_gpu_sources "${_source}") + endforeach() + python_add_library(_gpu MODULE USE_SABI ${BLOCHSIM_ABI3_VERSION} WITH_SOABI diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index a9ac7157..c47527f9 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -78,7 +78,7 @@ PyObject* launch(PyObject*, PyObject* args) { return cuda_error(status, "selecting the device"); } if (laying_out) { - const int laid = blochsim_layout::launch(request.kernel, request.arguments, + const int laid = blochsim_layout::launch(request.kernel, request.arguments, request.grid[0], reinterpret_cast(stream)); if (laid >= 0) { ++layout_launches; @@ -125,6 +125,18 @@ PyObject* layout_launch_count(PyObject*, PyObject*) { return PyLong_FromUnsignedLongLong(layout_launches); } +PyObject* pooled_layout_floats(PyObject*, PyObject* args) { + long long problems = 0; + int event_count = 0, n = 0, m = 0, width = 0, dual = 0; + if (!PyArg_ParseTuple(args, "Liiiip", &problems, &event_count, &n, &m, &width, &dual)) { + return nullptr; + } + if (!laying_out) { + return PyLong_FromLong(-1); + } + return PyLong_FromLongLong(blochsim_layout::pooled_adjoint_floats(problems, event_count, n, m, width, dual != 0)); +} + PyMethodDef METHODS[] = { {"kernels", blochsim_launch::kernel_table, METH_NOARGS, "Each kernel's parameter names and kinds."}, @@ -134,6 +146,8 @@ PyMethodDef METHODS[] = { "Whether launches may run their kernel's layout; returns the previous setting."}, {"layout_launches", layout_launch_count, METH_NOARGS, "How many launches have run a kernel's layout."}, + {"pooled_layout_floats", pooled_layout_floats, METH_VARARGS, + "The floats the many-pool adjoint's layout takes, or -1 where the tile kernels run it."}, {nullptr, nullptr, 0, nullptr}, }; diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index d6c0d39c..6493384e 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -66,6 +66,22 @@ def layout_launches() -> int: return _module("cuda").layout_launches() if available() else 0 +def pooled_layout_floats( + problems: int, event_count: int, n: int, m: int, width: int, dual: bool +) -> int | None: + """The floats the many-pool adjoint's layout takes, or None where it does not run. + + A launch the layout takes records nothing first: the adjoint keeps its own + checkpoints in the buffer it is given. + """ + if not available(): + return None + floats = _module("cuda").pooled_layout_floats( + problems, event_count, n, m, width, dual + ) + return None if floats < 0 else floats + + @contextmanager def generic_kernels() -> Iterator[None]: """Run every launch inside on the tile kernels rather than on its layout.""" diff --git a/src/blochsim/_layout.cu b/src/blochsim/_layout.cu index 891ac2e9..9f918e4a 100644 --- a/src/blochsim/_layout.cu +++ b/src/blochsim/_layout.cu @@ -311,6 +311,141 @@ layout_real_vjp::Params real_adjoint_params(bool dual, const Reader& r) { return p; } +// What the many-pool forward kernel and its adjoint read alike. +epg_pooled::Params pooled_params(const Reader& r, long long programs) { + epg_pooled::Params p{}; + p.m0 = r.floats("m0"); + p.b1 = r.floats("b1"); + p.b1_phase = r.floats("b1_phase"); + p.b0 = r.floats("b0"); + p.efficiency = r.floats("efficiency"); + p.diffusion = r.floats("diffusion"); + p.velocity = r.floats("velocity"); + p.dm0 = r.floats("dm0"); + p.db1 = r.floats("db1"); + p.db1_phase = r.floats("db1_phase"); + p.db0 = r.floats("db0"); + p.defficiency = r.floats("defficiency"); + p.ddiffusion = r.floats("ddiffusion"); + p.dvelocity = r.floats("dvelocity"); + p.duration = r.floats("duration"); + p.flip = r.floats("flip"); + p.phase = r.floats("phase"); + p.saturation = r.floats("saturation"); + p.rf_frequency = r.floats("rf_frequency"); + p.dduration = r.floats("dduration"); + p.dflip = r.floats("dflip"); + p.dphase = r.floats("dphase"); + p.table = r.floats("table"); + p.dtable = r.floats("dtable"); + p.profile = r.floats("profile"); + p.lineshape = r.floats("lineshape"); + p.pairs = r.floats("pairs"); + p.dpairs = r.floats("dpairs"); + p.kind = r.ints("kind"); + p.output_index = r.ints("output_index"); + p.shim_index = r.ints("shim_index"); + p.pool_index = r.ints("pool_index"); + p.profile_index = r.ints("profile_index"); + p.pair_index = r.ints("pair_index"); + p.action = static_cast(r.a[r.at("action")].p); + p.base = r.integer("base"); + p.problems = static_cast(programs); + p.atom_count = static_cast(r.integer("atom_count")); + p.event_count = static_cast(r.integer("event_count")); + p.output_count = static_cast(r.integer("output_count")); + p.state_count = static_cast(r.integer("state_count")); + p.rows = static_cast(r.integer("rows")); + p.width = static_cast(r.integer("S")); + p.m = static_cast(r.integer("m")); + p.blocks = static_cast(r.integer("blocks")); + p.locations = static_cast(r.integer("locations")); + p.profile_bins = static_cast(r.integer("profile_bins")); + p.lineshape_bins = static_cast(r.integer("lineshape_bins")); + p.flow_scale = r.real("flow_scale"); + p.washout_scale = r.real("washout_scale"); + p.profile_step = r.real("profile_step"); + p.lineshape_step = r.real("lineshape_step"); + p.atom_stride = r.flag("atom_stride"); + p.shimmed = r.flag("shimmed"); + p.directed_pairs = r.flag("directed_pairs"); + p.directed_table = r.flag("directed_table"); + p.off_axis = r.flag("off_axis"); + p.moving = r.flag("moving"); + p.diffusing = r.flag("diffusing"); + p.transmit = r.flag("transmit"); + p.density = r.flag("density"); + p.inverting = r.flag("inverting"); + return p; +} + +epg_pooled::Adjoint pooled_adjoint(const Reader& r) { + epg_pooled::Adjoint g{}; + g.grad_real = r.floats("grad_real"); + g.grad_imag = r.floats("grad_imag"); + g.grad_tissue = r.outputs("grad_tissue"); + g.dgrad_tissue = r.outputs("dgrad_tissue"); + g.grad_duration = r.outputs("grad_duration"); + g.dgrad_duration = r.outputs("dgrad_duration"); + g.grad_flip = r.outputs("grad_flip"); + g.dgrad_flip = r.outputs("dgrad_flip"); + g.grad_phase = r.outputs("grad_phase"); + g.dgrad_phase = r.outputs("dgrad_phase"); + g.grad_table = r.outputs("grad_table"); + g.dgrad_table = r.outputs("dgrad_table"); + g.grad_pairs = r.outputs("grad_pairs"); + g.dgrad_pairs = r.outputs("dgrad_pairs"); + g.trajectory = r.outputs("trajectory"); + g.m0_row = static_cast(r.integer("m0_row")); + g.b1_row = static_cast(r.integer("b1_row")); + g.b1_phase_row = static_cast(r.integer("b1_phase_row")); + g.b0_row = static_cast(r.integer("b0_row")); + g.efficiency_row = static_cast(r.integer("efficiency_row")); + g.diffusion_row = static_cast(r.integer("diffusion_row")); + g.velocity_row = static_cast(r.integer("velocity_row")); + return g; +} + +// Whether the many-pool layouts carry this many pools and orders. +bool pooled_fits(long long n, long long width) { return n >= 2 && n <= 8 && width <= 32; } + +int pooled_launch(bool adjoint, const Reader& r, long long programs, cudaStream_t stream) { + const long long n = r.integer("n"); + if (!pooled_fits(n, r.integer("S"))) { + return -1; + } + // A recording for the tile adjoint is the tile kernel's to make. + if (!adjoint && r.flag("keep")) { + return -1; + } + epg_pooled::Params p = pooled_params(r, programs); + if (!adjoint) { + p.output_real = r.outputs("output_real"); + p.output_imag = r.outputs("output_imag"); + } + const int rf = r.flag("dynamic") ? epg::DYNAMIC : (r.flag("profiled") ? epg::PROFILE : epg::HARD); + const bool dual = r.flag("following"); + if (adjoint) { + const epg_pooled::Adjoint g = pooled_adjoint(r); + switch (n) { +#define BLOCHSIM_LAYOUT_POOLED_CASE(pools) \ + case pools: return pooled_adjoint_##pools(p, g, rf, dual, stream); + BLOCHSIM_LAYOUT_POOLED(BLOCHSIM_LAYOUT_POOLED_CASE) +#undef BLOCHSIM_LAYOUT_POOLED_CASE + } + } else { + switch (n) { +#define BLOCHSIM_LAYOUT_POOLED_CASE(pools) \ + case pools: return pooled_forward_##pools(p, rf, dual, stream); + BLOCHSIM_LAYOUT_POOLED(BLOCHSIM_LAYOUT_POOLED_CASE) +#undef BLOCHSIM_LAYOUT_POOLED_CASE + } + } + return -1; +} + +const int POOLED = bsk::kernel_index("_pooled_kernel"); +const int POOLED_ADJOINT = bsk::kernel_index("_pooled_adjoint_kernel"); const int COMPLEX = bsk::kernel_index("_epg_kernel"); const int COMPLEX_VJP = bsk::kernel_index("_epg_vjp_kernel"); const int COMPLEX_VJP_JVP = bsk::kernel_index("_epg_vjp_jvp_kernel"); @@ -322,8 +457,20 @@ const int REAL_JVP = bsk::kernel_index("_epg_real_jvp_kernel"); } // namespace -int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream) { +long long pooled_adjoint_floats(long long problems, int event_count, int n, int m, int width, bool dual) { + if (!pooled_fits(n, width)) { + return -1; + } + const long long programs = (problems + 32 / width - 1) / (32 / width); + return problems * epg_pooled::adjoint_kept_floats(n, width, event_count, POOLED_SEGMENT, dual) + + programs * epg_pooled::adjoint_scratch_floats(n, m, 32, POOLED_SEGMENT, dual); +} + +int launch(int kernel, const bsk::Arguments& arguments, long long programs, cudaStream_t stream) { const Reader r{kernel, arguments.a}; + if (kernel == POOLED || kernel == POOLED_ADJOINT) { + return pooled_launch(kernel == POOLED_ADJOINT, r, programs, stream); + } const bool adjoint = kernel == COMPLEX_VJP || kernel == COMPLEX_VJP_JVP || kernel == REAL_VJP || kernel == REAL_VJP_JVP; if (kernel != COMPLEX && kernel != COMPLEX_JVP && kernel != REAL && kernel != REAL_JVP && !adjoint) { return -1; diff --git a/src/blochsim/_layout.hpp b/src/blochsim/_layout.hpp index 519a9e28..1f660041 100644 --- a/src/blochsim/_layout.hpp +++ b/src/blochsim/_layout.hpp @@ -14,14 +14,21 @@ #include "_kernels.hpp" #include "_layout_complex.hpp" #include "_layout_complex_vjp.hpp" +#include "_layout_pooled.hpp" #include "_layout_real.hpp" #include "_layout_real_vjp.hpp" namespace blochsim_layout { -// Queue ``kernel`` on ``stream`` and return its launch status, or -1 where it -// has no layout for these arguments and the tile kernel is to run instead. -int launch(int kernel, const bsk::Arguments& arguments, cudaStream_t stream); +// Queue ``kernel`` over ``programs`` on ``stream`` and return its launch +// status, or -1 where it has no layout for these arguments and the tile +// kernel is to run instead. +int launch(int kernel, const bsk::Arguments& arguments, long long programs, cudaStream_t stream); + +// The floats of the buffer the many-pool adjoint's layout takes for +// ``problems``, or -1 where the tile kernels run it. A launch the layout +// takes records nothing first: the adjoint keeps its own checkpoints. +long long pooled_adjoint_floats(long long problems, int event_count, int n, int m, int width, bool dual); // One translation unit per kernel and pool layout, so a build compiles them // side by side. ``rf`` and ``mode`` are epg::Rf and epg::Mode. @@ -49,6 +56,17 @@ BLOCHSIM_LAYOUT_ADJOINT(BLOCHSIM_LAYOUT_ADJOINT_DECLARATION) int real_vjp(const layout_real_vjp::Params& p, cudaStream_t stream); int real_vjp_jvp(const layout_real_vjp::Params& p, cudaStream_t stream); +// The many-pool kernels, a unit per pool count. The adjoint keeps every +// POOLED_SEGMENT-th state and replays the stretches between. +constexpr int POOLED_SEGMENT = 4; +#define BLOCHSIM_LAYOUT_POOLED(X) X(2) X(3) X(4) X(5) X(6) X(7) X(8) +#define BLOCHSIM_LAYOUT_POOLED_DECLARATION(pools) \ + int pooled_forward_##pools(const epg_pooled::Params& p, int rf, bool dual, cudaStream_t stream); \ + int pooled_adjoint_##pools(const epg_pooled::Params& p, const epg_pooled::Adjoint& g, int rf, bool dual, \ + cudaStream_t stream); +BLOCHSIM_LAYOUT_POOLED(BLOCHSIM_LAYOUT_POOLED_DECLARATION) +#undef BLOCHSIM_LAYOUT_POOLED_DECLARATION + // Programs of two warps each. constexpr int WARPS = 2; diff --git a/src/blochsim/_layout_pooled.cu.in b/src/blochsim/_layout_pooled.cu.in new file mode 100644 index 00000000..6b936091 --- /dev/null +++ b/src/blochsim/_layout_pooled.cu.in @@ -0,0 +1,66 @@ +// Written by CMake from _layout_pooled.cu.in: the many-pool kernels for +// @POOLS@ pools in their layout. +#include "_layout.hpp" + +namespace blochsim_layout { +namespace { + +constexpr int SEGMENT = POOLED_SEGMENT; +constexpr int N = @POOLS@; + +template +__global__ void __launch_bounds__(32 * WARPS) forward_kernel(epg_pooled::Params p) { + epg_pooled::pooled_loop(p); +} + +// One warp a program: its scratch is a warp's. +template +__global__ void __launch_bounds__(32) adjoint_kernel(epg_pooled::Params p, epg_pooled::Adjoint g) { + const long long scratch = epg_pooled::adjoint_scratch_floats(N, p.m, 32, SEGMENT, num::is_dual::value); + epg_pooled::pooled_adjoint_loop(p, g, g.scratch + blockIdx.x * scratch); +} + +template +int forward_rf(const epg_pooled::Params& p, cudaStream_t stream) { + const int per_block = WARPS * (32 / p.width); + const unsigned grid = static_cast((p.problems + per_block - 1) / per_block); + forward_kernel<<>>(p); + return static_cast(cudaGetLastError()); +} + +template +int adjoint_rf(const epg_pooled::Params& p, epg_pooled::Adjoint g, cudaStream_t stream) { + const int groups = 32 / p.width; + const unsigned grid = static_cast((p.problems + groups - 1) / groups); + g.kept = epg_pooled::adjoint_kept_floats(N, p.width, p.event_count, SEGMENT, num::is_dual::value); + g.scratch = g.trajectory + p.problems * g.kept; + adjoint_kernel<<>>(p, g); + return static_cast(cudaGetLastError()); +} + +template +int forward_t(const epg_pooled::Params& p, int rf, cudaStream_t stream) { + if (rf == epg::DYNAMIC) return forward_rf(p, stream); + if (rf == epg::PROFILE) return forward_rf(p, stream); + return forward_rf(p, stream); +} + +template +int adjoint_t(const epg_pooled::Params& p, const epg_pooled::Adjoint& g, int rf, cudaStream_t stream) { + if (rf == epg::DYNAMIC) return adjoint_rf(p, g, stream); + if (rf == epg::PROFILE) return adjoint_rf(p, g, stream); + return adjoint_rf(p, g, stream); +} + +} // namespace + +int pooled_forward_@POOLS@(const epg_pooled::Params& p, int rf, bool dual, cudaStream_t stream) { + return dual ? forward_t(p, rf, stream) : forward_t(p, rf, stream); +} + +int pooled_adjoint_@POOLS@(const epg_pooled::Params& p, const epg_pooled::Adjoint& g, int rf, bool dual, + cudaStream_t stream) { + return dual ? adjoint_t(p, g, rf, stream) : adjoint_t(p, g, rf, stream); +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout_pooled.hpp b/src/blochsim/_layout_pooled.hpp new file mode 100644 index 00000000..813472f9 --- /dev/null +++ b/src/blochsim/_layout_pooled.hpp @@ -0,0 +1,1014 @@ +// The many-pool EPG kernels over a number type: the forward simulation at +// float, its Jacobian-vector product at num::Dual, and the adjoint whose +// derivative along a direction is the same source at num::Dual. Every interval's exchange +// is read from the tissue's own table: a longitudinal operator over the N +// pools, what each recovers, and a complex transverse operator over the m +// exchanging pools. Lanes along the states, groups of lanes across problems, +// one problem's pools in each thread's registers. +#pragma once +#include "_layout_complex.hpp" +#include "_layout_complex_vjp.hpp" + +namespace epg_pooled { + +using namespace epg; + +struct Params { + const float *m0, *b1, *b1_phase, *b0, *efficiency, *diffusion, *velocity; + const float *dm0, *db1, *db1_phase, *db0, *defficiency, *ddiffusion, *dvelocity; + const float *duration, *flip, *phase, *saturation, *rf_frequency, *dduration, *dflip, *dphase; + const float *table, *dtable, *profile, *lineshape, *pairs, *dpairs; + const int *kind, *output_index, *shim_index, *pool_index, *profile_index, *pair_index; + const unsigned char* action; + float *output_real, *output_imag; + long long base; + int problems, atom_count, event_count, output_count, state_count, rows, width, m, blocks; + int locations, profile_bins, lineshape_bins; + float flow_scale, washout_scale, profile_step, lineshape_step; + bool atom_stride, shimmed, directed_pairs, directed_table, off_axis, moving, diffusing, transmit, density, + inverting; +}; + +template +struct Z { + T r, i; +}; +template +__device__ __forceinline__ Z zmul(const Z& a, const Z& b) { + return {a.r * b.r - a.i * b.i, a.r * b.i + a.i * b.r}; +} +template +__device__ __forceinline__ Z zconj(const Z& a) { + return {a.r, -a.i}; +} + +template +__device__ __forceinline__ void pooled_loop(const Params& p) { + constexpr bool DUAL = num::is_dual::value; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const long long index = (static_cast(blockIdx.x) * (blockDim.x >> 5) + warp) * groups + group; + const bool active = index < p.problems; + const long long problem = p.base + (active ? index : 0); + const bool live = active && state < p.state_count; + const int atom = static_cast(problem % p.atom_count); + const int train = static_cast(problem / p.atom_count); + const int event_base = train * p.event_count; + const int voxel_at = p.atom_stride ? atom : 0; + const int location = RF == PROFILE ? atom % p.locations : 0; + const float order = static_cast(state); + const float squared = order * order, transverse_weight = squared + order + 0.3333333333333333f; + const int m = p.m; + const int row_width = N * N + N + 2 * m * m; + const long long table_width = N + static_cast(p.rows) * row_width * p.blocks; + const float* slot = p.table + atom * table_width; + const float* directions = p.directed_table ? p.dtable + atom * table_width : nullptr; + const long long sloped_at = static_cast(p.rows) * row_width; + + auto read = [&](const float* values, const float* tangents, long long at, bool on, float identity) -> T { + if (!on) return T(identity); + if constexpr (DUAL) { + return T{__ldg(values + at), __ldg(tangents + at)}; + } else { + return __ldg(values + at); + } + }; + // A table entry, moving with the table's direction and, where the table + // carries a slope along the interval's length, along that. + auto entry = [&](long long at, float along, bool sloped) -> T { + const float value = __ldg(slot + at); + if constexpr (DUAL) { + float tangent = 0.0f; + if (directions != nullptr) tangent += __ldg(directions + at); + if (sloped && p.blocks > 1) tangent += __ldg(slot + sloped_at + at) * along; + return T{value, tangent}; + } else { + return value; + } + }; + const T density = read(p.m0, p.dm0, voxel_at, p.density, 1.0f); + const T voxel_b1 = read(p.b1, p.db1, voxel_at, p.transmit, 1.0f); + const T voxel_b1_phase = read(p.b1_phase, p.db1_phase, voxel_at, p.off_axis, 0.0f); + const T voxel_b0 = read(p.b0, p.db0, voxel_at, p.off_axis, 0.0f); + const T inversion = read(p.efficiency, p.defficiency, voxel_at, p.inverting, 1.0f); + const T damping_rate = read(p.diffusion, p.ddiffusion, voxel_at, p.diffusing, 0.0f); + const T moved = read(p.velocity, p.dvelocity, voxel_at, p.moving, 0.0f); + const T flow_rate = p.flow_scale * moved; + const T washout_rate = p.washout_scale * abs_(moved); + + T equilibrium[N]; + Z plus[N], minus[N], z[N]; +#pragma unroll + for (int i = 0; i < N; ++i) { + equilibrium[i] = entry(i, 0.0f, false); + plus[i] = minus[i] = {T(0.0f), T(0.0f)}; + z[i] = {state == 0 ? equilibrium[i] : T(0.0f), T(0.0f)}; + } + auto shift = [&]() { +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i >= m) break; + const T up_r = num::shfl_up(plus[i].r, 1, width), up_i = num::shfl_up(plus[i].i, 1, width); + const T dn_r = num::shfl_down(minus[i].r, 1, width), dn_i = num::shfl_down(minus[i].i, 1, width); + const bool keep_up = state > 0 && live, keep_down = state + 1 < p.state_count && live; + const T pr = keep_up ? up_r : T(0.0f), pi = keep_up ? up_i : T(0.0f); + const T mr = keep_down ? dn_r : T(0.0f), mi = keep_down ? dn_i : T(0.0f); + plus[i] = {state == 0 ? mr : pr, state == 0 ? -mi : pi}; + minus[i] = {mr, mi}; + } + }; + auto event_value = [&](const float* values, const float* tangents, int event) -> T { + if constexpr (DUAL) { + return T{__ldg(values + event_base + event), __ldg(tangents + event_base + event)}; + } else { + return __ldg(values + event_base + event); + } + }; + +#pragma unroll 1 + for (int event = 0; event < p.event_count; ++event) { + const T dt = event_value(p.duration, p.dduration, event); + // What the interval does to every pool alike at this order. + T wout = 1.0f, damp_z = 1.0f, damp_t = 1.0f; + if (p.diffusing) { + const T b = damping_rate * dt; + damp_z = exp_(-squared * b); + damp_t = exp_(-transverse_weight * b); + } + if (p.moving) wout = 1.0f - min_(washout_rate * dt, T(1.0f)); + Z carried = {wout * damp_t, T(0.0f)}, spin = {wout * damp_z, T(0.0f)}; + if (p.off_axis || p.moving) { + const T turn = p.moving ? flow_rate * dt : T(0.0f); + T s, c; + sincos_(-TWO_PI * voxel_b0 * dt - (order + 0.5f) * turn, s, c); + carried = {wout * damp_t * c, wout * damp_t * s}; + if (p.moving) { + sincos_(-order * turn, s, c); + spin = {wout * damp_z * c, wout * damp_z * s}; + } + } + // The interval's exchange, from the table row its length reads. + const int row = p.pool_index[event_base + event]; + const long long row_at = N + static_cast(row) * row_width; + const float along = DUAL ? num::tangent(dt) : 0.0f; + Z mixed_plus[N], mixed_minus[N]; +#pragma unroll + for (int i = 0; i < N; ++i) { + mixed_plus[i] = mixed_minus[i] = {T(0.0f), T(0.0f)}; + if (i >= m) continue; +#pragma unroll + for (int j = 0; j < N; ++j) { + if (j >= m) break; + const long long at = row_at + N * N + N + 2 * (i * m + j); + const Z x = {entry(at, along, true), entry(at + 1, along, true)}; + const Z a = zmul(x, plus[j]), b = zmul(zconj(x), minus[j]); + mixed_plus[i] = {mixed_plus[i].r + a.r, mixed_plus[i].i + a.i}; + mixed_minus[i] = {mixed_minus[i].r + b.r, mixed_minus[i].i + b.i}; + } + } + Z mixed_z[N]; +#pragma unroll + for (int i = 0; i < N; ++i) { + mixed_z[i] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int j = 0; j < N; ++j) { + const T l = entry(row_at + i * N + j, along, true); + mixed_z[i] = {mixed_z[i].r + l * z[j].r, mixed_z[i].i + l * z[j].i}; + } + } +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) { + plus[i] = zmul(carried, mixed_plus[i]); + minus[i] = zmul(zconj(carried), mixed_minus[i]); + } + z[i] = zmul(spin, mixed_z[i]); + if (state == 0) z[i].r += equilibrium[i] - wout * entry(row_at + N * N + i, along, true); + } + + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + if (act & 1) shift(); + if (kind == 1) { + if (act & 4) { + // Every exchanging pool is free water and inverts like it; a + // semisolid one is saturated by the pulse's own term. +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) z[i] = {-inversion * z[i].r, -inversion * z[i].i}; + } + } else { + T pulse_b1 = voxel_b1, pulse_b1_phase = voxel_b1_phase; + if (p.shimmed) { + const long long transmit_at = static_cast(p.shim_index[event]) * p.atom_count + atom; + pulse_b1 = read(p.b1, p.db1, transmit_at, p.transmit, 1.0f); + pulse_b1_phase = read(p.b1_phase, p.db1_phase, transmit_at, true, 0.0f); + } + const T alpha = event_value(p.flip, p.dflip, event) * pulse_b1; + const T phi = event_value(p.phase, p.dphase, event) + pulse_b1_phase; + if (N > m) { + const T shape = lineshape_at(p.lineshape, p.rf_frequency[event] - voxel_b0, p.lineshape_bins, + p.lineshape_step); + const T absorbed = exp_(p.saturation[event] * alpha * alpha * shape); + z[N - 1] = {absorbed * z[N - 1].r, absorbed * z[N - 1].i}; + } + T ts, tc; + sincos_(-phi, ts, tc); + const Z turn = {tc, ts}; + Z a, b; + if constexpr (RF == DYNAMIC) { + const int pair_row = p.pair_index[event_base + event]; + const long long cell = (static_cast(pair_row) * p.atom_count + atom) * 4; + const float* direction = p.directed_pairs ? p.dpairs : nullptr; + a = {num::load(p.pairs, direction, cell), num::load(p.pairs, direction, cell + 1)}; + b = {num::load(p.pairs, direction, cell + 2), num::load(p.pairs, direction, cell + 3)}; + } else if constexpr (RF == PROFILE) { + const int table_row = p.profile_index[event] * p.locations + location; + const int last = p.profile_bins - 1; + const T scaled = min_(max_(div_(alpha, T(p.profile_step)), T(0.0f)), T(last + 0.0f)); + const float lower = fminf(floorf(primal(scaled)), last - 1.0f); + T h10, h01, h11; + const T h00 = hermite_weights(scaled, lower, p.profile_step, h10, h01, h11); + const float* knot = p.profile + (table_row * p.profile_bins + static_cast(lower)) * 8; + T pair[4]; +#pragma unroll + for (int c = 0; c < 4; ++c) { + pair[c] = h00 * __ldg(knot + c) + h10 * __ldg(knot + 4 + c) + h01 * __ldg(knot + 8 + c) + + h11 * __ldg(knot + 12 + c); + } + a = {pair[0], pair[1]}; + b = {pair[2], pair[3]}; + } else { + T hs, hc; + sincos_(0.5f * alpha, hs, hc); + a = {hc, T(0.0f)}; + b = {T(0.0f), -hs}; + } + const Z spun = zmul(b, turn); +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) rotate_spinor(a.r, a.i, spun.r, spun.i, plus[i].r, plus[i].i, minus[i].r, minus[i].i, + z[i].r, z[i].i); + } + } + } + if (kind == 2 && (act & 32)) { + const int out = p.output_index[event]; + T recorded_r = 0.0f, recorded_i = 0.0f; +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) { + recorded_r += plus[i].r; + recorded_i += plus[i].i; + } + } + if (out >= 0 && state == 0 && active) { + T ts, tc; + sincos_(-event_value(p.phase, p.dphase, event), ts, tc); + const Z signal = zmul(Z{density * recorded_r, density * recorded_i}, Z{tc, ts}); + const long long written = problem * p.output_count + out; + if constexpr (DUAL) { + p.output_real[written] = signal.r.d; + p.output_imag[written] = signal.i.d; + } else { + p.output_real[written] = signal.r; + p.output_imag[written] = signal.i; + } + } + } + if (act & 2) shift(); + if (act & 8) { +#pragma unroll + for (int i = 0; i < N; ++i) plus[i] = minus[i] = {T(0.0f), T(0.0f)}; + } else if (act & 16) { + shift(); + } + } +} + + +// The adjoint's own buffers: the output's cotangent, the value and +// derivative planes of every gradient it writes, the checkpoints (``kept`` +// floats a problem) and the programs' scratch after them. +struct Adjoint { + const float *grad_real, *grad_imag; + float *grad_tissue, *dgrad_tissue, *grad_duration, *dgrad_duration, *grad_flip, *dgrad_flip; + float *grad_phase, *dgrad_phase, *grad_table, *dgrad_table, *grad_pairs, *dgrad_pairs; + float *trajectory, *scratch; + long long kept; + int m0_row, b1_row, b1_phase_row, b0_row, efficiency_row, diffusion_row, velocity_row; +}; + +template +__device__ __forceinline__ void add_both(float* value, float* tangent, long long at, T x) { + if constexpr (num::is_dual::value) { + atomicAdd(value + at, x.v); + atomicAdd(tangent + at, x.d); + } else { + atomicAdd(value + at, x); + } +} +// The sum over a group's lanes, in every lane of it. +template +__device__ __forceinline__ T group_sum(T x, int width) { + for (int offset = 1; offset < width; offset <<= 1) { + if constexpr (num::is_dual::value) { + x = x + T{__shfl_xor_sync(0xffffffffu, x.v, offset), __shfl_xor_sync(0xffffffffu, x.d, offset)}; + } else { + x += __shfl_xor_sync(0xffffffffu, x, offset); + } + } + return x; +} + +// What an interval does to every pool alike at an order, as a function of the +// inputs it depends on: its length, B0, the diffusion and the velocity. +template +__device__ __forceinline__ void pool_factors(const Params& p, U dt, U b0, U damping_rate, U moved, float order, + U& wout, Z& carried, Z& spin) { + const float squared = order * order, weight = squared + order + 0.3333333333333333f; + U damp_z = 1.0f, damp_t = 1.0f; + wout = 1.0f; + if (p.diffusing) { + const U b = damping_rate * dt; + damp_z = exp_(-squared * b); + damp_t = exp_(-weight * b); + } + if (p.moving) wout = 1.0f - min_(p.washout_scale * abs_(moved) * dt, U(1.0f)); + carried = {wout * damp_t, U(0.0f)}; + spin = {wout * damp_z, U(0.0f)}; + if (p.off_axis || p.moving) { + const U turn = p.moving ? p.flow_scale * moved * dt : U(0.0f); + U s, c; + sincos_(-TWO_PI * b0 * dt - (order + 0.5f) * turn, s, c); + carried = {wout * damp_t * c, wout * damp_t * s}; + if (p.moving) { + sincos_(-order * turn, s, c); + spin = {wout * damp_z * c, wout * damp_z * s}; + } + } +} + +// The floats a problem keeps its checkpoints in, every K-th state with its +// tangent at num::Dual; and the scratch a program of ``threads`` works in -- +// the states of the stretch it replays, and the per-lane sums of a table +// row's cotangent -- which is a warp's own and stays in its caches. +__host__ __device__ constexpr long long adjoint_kept_floats(int n, int width, int event_count, int k, bool dual) { + return static_cast((event_count + k - 1) / k) * 6 * n * width * (dual ? 2 : 1); +} +__host__ __device__ constexpr long long adjoint_scratch_floats(int n, int m, int threads, int k, bool dual) { + return static_cast(threads) * (k * 6 * n * (dual ? 2 : 1) + (n * n + n + 2 * m * m) * (dual ? 3 : 1)); +} + +template +__device__ __forceinline__ void pooled_adjoint_loop(const Params& p, const Adjoint& g, float* scratch) { + constexpr bool DUAL = num::is_dual::value; + constexpr int PLANES = 6 * N; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int width = p.width; + const int state = lane & (width - 1); + const int group = lane / width; + const int groups = 32 / width; + const long long index = (static_cast(blockIdx.x) * (blockDim.x >> 5) + warp) * groups + group; + const bool active = index < p.problems; + const long long problem = p.base + (active ? index : 0); + const bool live = active && state < p.state_count; + const int atom = static_cast(problem % p.atom_count); + const int train = static_cast(problem / p.atom_count); + const int event_base = train * p.event_count; + const int voxel_at = p.atom_stride ? atom : 0; + const int location = RF == PROFILE ? atom % p.locations : 0; + const float order = static_cast(state); + const int m = p.m; + const int row_width = N * N + N + 2 * m * m; + const long long table_width = N + static_cast(p.rows) * row_width * p.blocks; + const float* slot = p.table + atom * table_width; + const float* directions = p.directed_table ? p.dtable + atom * table_width : nullptr; + const long long sloped_at = static_cast(p.rows) * row_width; + const long long n_atoms = p.atom_count; + const int threads = blockDim.x; + + auto read = [&](const float* values, const float* tangents, long long at, bool on, float identity) -> T { + if (!on) return T(identity); + if constexpr (DUAL) { + return T{__ldg(values + at), __ldg(tangents + at)}; + } else { + return __ldg(values + at); + } + }; + // A table entry, moving with the table's direction and, where the table + // carries the next block, along the interval's length with it. + auto entry = [&](long long at, float along, bool sloped) -> T { + const float value = __ldg(slot + at); + if constexpr (DUAL) { + float tangent = 0.0f; + if (directions != nullptr) tangent += __ldg(directions + at); + if (sloped && p.blocks > 1) tangent += __ldg(slot + sloped_at + at) * along; + return T{value, tangent}; + } else { + return value; + } + }; + auto event_value = [&](const float* values, const float* tangents, int event) -> T { + if constexpr (DUAL) { + return T{__ldg(values + event_base + event), __ldg(tangents + event_base + event)}; + } else { + return __ldg(values + event_base + event); + } + }; + const T density = read(p.m0, p.dm0, voxel_at, p.density, 1.0f); + const T voxel_b1 = read(p.b1, p.db1, voxel_at, p.transmit, 1.0f); + const T voxel_b1_phase = read(p.b1_phase, p.db1_phase, voxel_at, p.off_axis, 0.0f); + const T voxel_b0 = read(p.b0, p.db0, voxel_at, p.off_axis, 0.0f); + const T inversion = read(p.efficiency, p.defficiency, voxel_at, p.inverting, 1.0f); + const T damping_rate = read(p.diffusion, p.ddiffusion, voxel_at, p.diffusing, 0.0f); + const T moved = read(p.velocity, p.dvelocity, voxel_at, p.moving, 0.0f); + + T equilibrium[N]; + Z plus[N], minus[N], z[N]; +#pragma unroll + for (int i = 0; i < N; ++i) { + equilibrium[i] = entry(i, 0.0f, false); + plus[i] = minus[i] = {T(0.0f), T(0.0f)}; + z[i] = {state == 0 ? equilibrium[i] : T(0.0f), T(0.0f)}; + } + auto shift = [&]() { +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i >= m) break; + const T up_r = num::shfl_up(plus[i].r, 1, width), up_i = num::shfl_up(plus[i].i, 1, width); + const T dn_r = num::shfl_down(minus[i].r, 1, width), dn_i = num::shfl_down(minus[i].i, 1, width); + const bool keep_up = state > 0 && live, keep_down = state + 1 < p.state_count && live; + const T pr = keep_up ? up_r : T(0.0f), pi = keep_up ? up_i : T(0.0f); + const T mr = keep_down ? dn_r : T(0.0f), mi = keep_down ? dn_i : T(0.0f); + plus[i] = {state == 0 ? mr : pr, state == 0 ? -mi : pi}; + minus[i] = {mr, mi}; + } + }; + // The interval at the event's row; ``mixed_*`` are what the exchange + // makes of the state before the factors every pool shares. + auto relax = [&](T dt, int row, T& wout, Z& carried, Z& spin, Z* mixed_plus, Z* mixed_minus, + Z* mixed_z) { + pool_factors(p, dt, voxel_b0, damping_rate, moved, order, wout, carried, spin); + const long long row_at = N + static_cast(row) * row_width; + const float along = DUAL ? num::tangent(dt) : 0.0f; +#pragma unroll + for (int i = 0; i < N; ++i) { + mixed_plus[i] = mixed_minus[i] = {T(0.0f), T(0.0f)}; + if (i >= m) continue; +#pragma unroll + for (int j = 0; j < N; ++j) { + if (j >= m) break; + const long long at = row_at + N * N + N + 2 * (i * m + j); + const Z x = {entry(at, along, true), entry(at + 1, along, true)}; + const Z a = zmul(x, plus[j]), b = zmul(zconj(x), minus[j]); + mixed_plus[i] = {mixed_plus[i].r + a.r, mixed_plus[i].i + a.i}; + mixed_minus[i] = {mixed_minus[i].r + b.r, mixed_minus[i].i + b.i}; + } + } +#pragma unroll + for (int i = 0; i < N; ++i) { + mixed_z[i] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int j = 0; j < N; ++j) { + const T l = entry(row_at + i * N + j, along, true); + mixed_z[i] = {mixed_z[i].r + l * z[j].r, mixed_z[i].i + l * z[j].i}; + } + } +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) { + plus[i] = zmul(carried, mixed_plus[i]); + minus[i] = zmul(zconj(carried), mixed_minus[i]); + } + z[i] = zmul(spin, mixed_z[i]); + if (state == 0) z[i].r += equilibrium[i] - wout * entry(row_at + N * N + i, along, true); + } + }; + // The pulse's flip, phase and transmit field. + auto pulse_inputs = [&](int event, T& pulse_b1, T& alpha, T& phi, T& nominal, int& shim) { + pulse_b1 = voxel_b1; + T pulse_b1_phase = voxel_b1_phase; + shim = 0; + if (p.shimmed) { + shim = p.shim_index[event]; + const long long transmit_at = static_cast(shim) * n_atoms + atom; + pulse_b1 = read(p.b1, p.db1, transmit_at, p.transmit, 1.0f); + pulse_b1_phase = read(p.b1_phase, p.db1_phase, transmit_at, true, 0.0f); + } + nominal = event_value(p.flip, p.dflip, event); + alpha = nominal * pulse_b1; + phi = event_value(p.phase, p.dphase, event) + pulse_b1_phase; + }; + // The Cayley-Klein pair a flip forms, turned by the phase. + auto pair_of = [&](int event, auto alpha, auto phi, auto& a, auto& b) { + using U = typename std::decay::type; + if constexpr (RF == PROFILE) { + const int table_row = p.profile_index[event] * p.locations + location; + const int last = p.profile_bins - 1; + const U scaled = min_(max_(div_(alpha, U(p.profile_step)), U(0.0f)), U(last + 0.0f)); + const float lower = fminf(floorf(primal(scaled)), last - 1.0f); + U h10, h01, h11; + const U h00 = hermite_weights(scaled, lower, p.profile_step, h10, h01, h11); + const float* knot = p.profile + (table_row * p.profile_bins + static_cast(lower)) * 8; + U pair[4]; +#pragma unroll + for (int c = 0; c < 4; ++c) { + pair[c] = h00 * __ldg(knot + c) + h10 * __ldg(knot + 4 + c) + h01 * __ldg(knot + 8 + c) + + h11 * __ldg(knot + 12 + c); + } + a = {pair[0], pair[1]}; + b = {pair[2], pair[3]}; + } else { + U hs, hc; + sincos_(0.5f * alpha, hs, hc); + a = {hc, U(0.0f)}; + b = {U(0.0f), -hs}; + } + U ts, tc; + sincos_(-phi, ts, tc); + b = zmul(b, Z{tc, ts}); + }; + auto dynamic_cell = [&](int event) { + return (static_cast(p.pair_index[event_base + event]) * n_atoms + atom) * 4; + }; + auto saturate = [&](int event, T alpha) { + if (N > m) { + const T shape = lineshape_at(p.lineshape, p.rf_frequency[event] - voxel_b0, p.lineshape_bins, + p.lineshape_step); + const T absorbed = exp_(p.saturation[event] * alpha * alpha * shape); + z[N - 1] = {absorbed * z[N - 1].r, absorbed * z[N - 1].i}; + } + }; + auto forward = [&](int event) { + T wout; + Z carried, spin, mixed_plus[N], mixed_minus[N], mixed_z[N]; + relax(event_value(p.duration, p.dduration, event), p.pool_index[event_base + event], wout, carried, spin, + mixed_plus, mixed_minus, mixed_z); + const unsigned char act = p.action[event]; + if (act & 1) shift(); + if (p.kind[event] == 1) { + if (act & 4) { +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) z[i] = {-inversion * z[i].r, -inversion * z[i].i}; + } + } else { + T pulse_b1, alpha, phi, nominal; + int shim; + pulse_inputs(event, pulse_b1, alpha, phi, nominal, shim); + saturate(event, alpha); + Z a, b; + if constexpr (RF == DYNAMIC) { + const long long cell = dynamic_cell(event); + const float* direction = p.directed_pairs ? p.dpairs : nullptr; + a = {num::load(p.pairs, direction, cell), num::load(p.pairs, direction, cell + 1)}; + b = {num::load(p.pairs, direction, cell + 2), num::load(p.pairs, direction, cell + 3)}; + T ts, tc; + sincos_(-phi, ts, tc); + b = zmul(b, Z{tc, ts}); + } else { + pair_of(event, alpha, phi, a, b); + } +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) rotate_spinor(a.r, a.i, b.r, b.i, plus[i].r, plus[i].i, minus[i].r, minus[i].i, z[i].r, + z[i].i); + } + } + } + if (act & 2) shift(); + if (act & 8) { +#pragma unroll + for (int i = 0; i < N; ++i) plus[i] = minus[i] = {T(0.0f), T(0.0f)}; + } else if (act & 16) { + shift(); + } + }; + + // Checkpoints, in the problem's share of the trajectory buffer: the + // values, then the tangents. + const int checkpoints = (p.event_count + K - 1) / K; + float* kept = g.trajectory + (problem - p.base) * g.kept; + const long long tangents_at = static_cast(checkpoints) * PLANES * width; + auto component = [&](int plane) -> T& { + const int i = plane / 6, part = plane % 6; + Z& c = part < 2 ? plus[i] : (part < 4 ? minus[i] : z[i]); + return (part & 1) ? c.i : c.r; + }; + auto kept_at = [&](int checkpoint, int plane) { + return (static_cast(checkpoint) * PLANES + plane) * width + state; + }; + T* segment = reinterpret_cast(scratch); + float* sums = scratch + K * PLANES * threads * (DUAL ? 2 : 1); + auto stretch = [&](int k, int plane) -> T& { return segment[(k * PLANES + plane) * threads + threadIdx.x]; }; + +#pragma unroll 1 + for (int event = 0; event < p.event_count; ++event) { + if (event % K == 0 && live) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + const long long at = kept_at(event / K, plane); + if constexpr (DUAL) { + kept[at] = component(plane).v; + kept[tangents_at + at] = component(plane).d; + } else { + kept[at] = component(plane); + } + } + } + forward(event); + } + + // ---- the walk back ---- + Z bplus[N], bminus[N], bz[N]; +#pragma unroll + for (int i = 0; i < N; ++i) bplus[i] = bminus[i] = bz[i] = {T(0.0f), T(0.0f)}; + T g_m0 = 0.0f, g_b1 = 0.0f, g_b1_phase = 0.0f, g_b0 = 0.0f, g_efficiency = 0.0f, g_damping = 0.0f, + g_velocity = 0.0f; + T g_equilibrium[N]; +#pragma unroll + for (int i = 0; i < N; ++i) g_equilibrium[i] = 0.0f; + auto shift_back = [&]() { +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i >= m) break; + const T dn_r = num::shfl_down(bplus[i].r, 1, width), dn_i = num::shfl_down(bplus[i].i, 1, width); + const T up_r = num::shfl_up(bminus[i].r, 1, width), up_i = num::shfl_up(bminus[i].i, 1, width); + const T head_r = num::shfl_up(bplus[i].r, 1, width), head_i = num::shfl_up(bplus[i].i, 1, width); + const bool keep_down = state + 1 < p.state_count && live, keep_up = state > 0 && live; + bplus[i] = {keep_down ? dn_r : T(0.0f), keep_down ? dn_i : T(0.0f)}; + T mr = keep_up ? up_r : T(0.0f), mi = keep_up ? up_i : T(0.0f); + if (state == 1 && live) { + mr = mr + head_r; + mi = mi - head_i; + } + bminus[i] = {mr, mi}; + } + }; + auto event_gradient = [&](float* value, float* tangent, int event, T lane_value) { + const T sum = group_sum(lane_value, width); + if (state == 0 && active) add_both(value, tangent, event_base + event, sum); + }; + + // A table row's cotangent, summed per lane while the row repeats: the + // value, and at num::Dual its tangent and what the slope block takes. + const int entries = row_width; + auto sum_value = [&](int k) -> float& { return sums[(DUAL ? 3 * k : k) * threads + threadIdx.x]; }; + auto sum_tangent = [&](int k) -> float& { return sums[(3 * k + 1) * threads + threadIdx.x]; }; + auto sum_sloped = [&](int k) -> float& { return sums[(3 * k + 2) * threads + threadIdx.x]; }; +#pragma unroll 1 + for (int k = 0; k < entries; ++k) { + sum_value(k) = 0.0f; + if constexpr (DUAL) sum_tangent(k) = sum_sloped(k) = 0.0f; + } + float* grad_row = g.grad_table + problem * table_width; + float* curve_row = g.dgrad_table + problem * table_width; + int summed_row = -1; + auto flush = [&]() { + if (!__any_sync(0xffffffffu, summed_row >= 0)) return; + const bool mine = summed_row >= 0 && state == 0 && active; + const long long row_at = N + static_cast(summed_row < 0 ? 0 : summed_row) * row_width; +#pragma unroll 1 + for (int k = 0; k < entries; ++k) { + const float v = group_sum(sum_value(k), width); + sum_value(k) = 0.0f; + if (mine) grad_row[row_at + k] += v; + if constexpr (DUAL) { + const float d = group_sum(sum_tangent(k), width); + const float sd = group_sum(sum_sloped(k), width); + sum_tangent(k) = sum_sloped(k) = 0.0f; + if (mine) { + curve_row[row_at + k] += d; + if (p.blocks > 1) curve_row[sloped_at + row_at + k] += sd; + } + } + } + }; + +#pragma unroll 1 + for (int checkpoint = checkpoints - 1; checkpoint >= 0; --checkpoint) { + const int start = checkpoint * K; + const int stop = min(start + K, p.event_count); +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) { + const long long at = kept_at(checkpoint, plane); + if constexpr (DUAL) { + component(plane) = live ? T{kept[at], kept[tangents_at + at]} : T(0.0f); + } else { + component(plane) = live ? kept[at] : 0.0f; + } + } +#pragma unroll 1 + for (int event = start; event < stop; ++event) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) stretch(event - start, plane) = component(plane); + if (event + 1 < stop) forward(event); + } +#pragma unroll 1 + for (int event = stop - 1; event >= start; --event) { +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) component(plane) = stretch(event - start, plane); + const unsigned char act = p.action[event]; + const int kind = p.kind[event]; + const T dt = event_value(p.duration, p.dduration, event); + const int row = p.pool_index[event_base + event]; + T wout; + Z carried, spin, mixed_plus[N], mixed_minus[N], mixed_z[N]; + relax(dt, row, wout, carried, spin, mixed_plus, mixed_minus, mixed_z); + if (act & 1) shift(); + + // The trailing spoil or shift, then the shift before it. + if (act & 8) { +#pragma unroll + for (int i = 0; i < N; ++i) bplus[i] = bminus[i] = {T(0.0f), T(0.0f)}; + } else if (act & 16) { + shift_back(); + } + if (act & 2) shift_back(); + + // A sample records the stage the event reached. + if (kind == 2 && (act & 32) && p.output_index[event] >= 0) { + T phase_gradient = 0.0f; + if (state == 0 && active) { + const long long at = problem * p.output_count + p.output_index[event]; + const Z seed = {T(g.grad_real[at]), T(g.grad_imag[at])}; + T ts, tc; + sincos_(-event_value(p.phase, p.dphase, event), ts, tc); + Z recorded = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) recorded = {recorded.r + plus[i].r, recorded.i + plus[i].i}; + } + const Z demodulated = zmul(recorded, Z{tc, ts}); + g_m0 += seed.r * demodulated.r + seed.i * demodulated.i; + phase_gradient = density * (seed.r * demodulated.i - seed.i * demodulated.r); + const Z weighted = zmul(Z{density * tc, -(density * ts)}, seed); +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) bplus[i] = {bplus[i].r + weighted.r, bplus[i].i + weighted.i}; + } + } + event_gradient(g.grad_phase, g.dgrad_phase, event, phase_gradient); + } + + if (kind == 1 && (act & 4)) { + T taken = 0.0f; +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i < m) { + taken += bz[i].r * z[i].r + bz[i].i * z[i].i; + bz[i] = {-inversion * bz[i].r, -inversion * bz[i].i}; + } + } + g_efficiency -= taken; + } else if (kind == 1) { + T pulse_b1, alpha, phi, nominal; + int shim; + pulse_inputs(event, pulse_b1, alpha, phi, nominal, shim); + // Each product of a state with its cotangent, over the + // exchanging pools, so the rotation's derivatives are taken + // against nine numbers. + Z met[9]; +#pragma unroll + for (int k = 0; k < 9; ++k) met[k] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i >= m) continue; + const Z S[3] = {plus[i], minus[i], z[i]}; + const Z L[3] = {bplus[i], bminus[i], bz[i]}; +#pragma unroll + for (int r = 0; r < 3; ++r) { +#pragma unroll + for (int c = 0; c < 3; ++c) { + const Z term = zmul(S[c], zconj(L[r])); + met[3 * r + c] = {met[3 * r + c].r + term.r, met[3 * r + c].i + term.i}; + } + } + } + auto against = [&](const auto* dR) { + using U = typename std::decay::type; + U sum = 0.0f; +#pragma unroll + for (int k = 0; k < 9; ++k) sum = sum + (dR[k].r * met[k].r - dR[k].i * met[k].i); + return sum; + }; + T g_alpha = 0.0f, g_phi = 0.0f; + Z R[9]; + if constexpr (RF == DYNAMIC) { + using V = num::Multi<5, T>; + const long long cell = dynamic_cell(event); + const float* direction = p.directed_pairs ? p.dpairs : nullptr; + const Z a = {V(num::load(p.pairs, direction, cell), 1), + V(num::load(p.pairs, direction, cell + 1), 2)}; + Z b = {V(num::load(p.pairs, direction, cell + 2), 3), + V(num::load(p.pairs, direction, cell + 3), 4)}; + V ts, tc; + sincos_(-V(phi, 0), ts, tc); + b = zmul(b, Z{tc, ts}); + epg_vjp::Cx dR[9]; + epg_vjp::spinor_rotation(epg_vjp::Cx{a.r, a.i}, epg_vjp::Cx{b.r, b.i}, dR); + const V got = against(dR); + g_phi = got.d[0]; +#pragma unroll + for (int c = 0; c < 4; ++c) { + const T total = group_sum(got.d[1 + c], width); + if (state == 0 && active) add_both(g.grad_pairs, g.dgrad_pairs, cell + c, total); + } +#pragma unroll + for (int k = 0; k < 9; ++k) R[k] = {dR[k].r.v, dR[k].i.v}; + } else { + using V = num::Multi<2, T>; + Z a, b; + pair_of(event, V(alpha, 0), V(phi, 1), a, b); + epg_vjp::Cx dR[9]; + epg_vjp::spinor_rotation(epg_vjp::Cx{a.r, a.i}, epg_vjp::Cx{b.r, b.i}, dR); + const V got = against(dR); + g_alpha = got.d[0]; + g_phi = got.d[1]; +#pragma unroll + for (int k = 0; k < 9; ++k) R[k] = {dR[k].r.v, dR[k].i.v}; + } + // The semisolid pool is saturated before the rotation, which + // leaves it alone. + if (N > m) { + using V = num::Multi<2, T>; + const V shape = lineshape_at(p.lineshape, V(p.rf_frequency[event]) - V(voxel_b0, 1), + p.lineshape_bins, p.lineshape_step); + const V absorbed = exp_(p.saturation[event] * V(alpha, 0) * V(alpha, 0) * shape); + const T taken = bz[N - 1].r * z[N - 1].r + bz[N - 1].i * z[N - 1].i; + g_alpha += absorbed.d[0] * taken; + g_b0 += absorbed.d[1] * taken; + bz[N - 1] = {absorbed.v * bz[N - 1].r, absorbed.v * bz[N - 1].i}; + } +#pragma unroll + for (int i = 0; i < N; ++i) { + if (i >= m) continue; + const Z L[3] = {bplus[i], bminus[i], bz[i]}; + Z back[3]; +#pragma unroll + for (int c = 0; c < 3; ++c) { + back[c] = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int r = 0; r < 3; ++r) { + const Z term = zmul(zconj(R[3 * r + c]), L[r]); + back[c] = {back[c].r + term.r, back[c].i + term.i}; + } + } + bplus[i] = back[0]; + bminus[i] = back[1]; + bz[i] = back[2]; + } + event_gradient(g.grad_flip, g.dgrad_flip, event, g_alpha * pulse_b1); + event_gradient(g.grad_phase, g.dgrad_phase, event, g_phi); + if (p.shimmed) { + const T b1_total = group_sum(g_alpha * nominal, width); + const T phase_total = group_sum(g_phi, width); + if (state == 0 && active) { + add_both(g.grad_tissue, g.dgrad_tissue, (g.b1_row + shim) * n_atoms + atom, b1_total); + add_both(g.grad_tissue, g.dgrad_tissue, (g.b1_phase_row + shim) * n_atoms + atom, + phase_total); + } + } else { + g_b1 += g_alpha * nominal; + g_b1_phase += g_phi; + } + } + if (act & 1) shift_back(); + + // The interval, from the state it met: one pass over the row + // takes the row's cotangent, sends the state's back through it + // and rebuilds what the exchange made, which the shared factors' + // derivatives are taken against. +#pragma unroll + for (int plane = 0; plane < PLANES; ++plane) component(plane) = stretch(event - start, plane); + if (__any_sync(0xffffffffu, row != summed_row)) { + flush(); + summed_row = row; + } + const long long row_at = N + static_cast(row) * row_width; + const float along = DUAL ? num::tangent(dt) : 0.0f; + const bool factored = p.off_axis || p.moving || p.diffusing; + T table_duration = 0.0f; + auto take = [&](int k, T cotangent) { + if constexpr (DUAL) { + sum_value(k) += cotangent.v; + sum_tangent(k) += cotangent.d; + if (p.blocks > 1) { + sum_sloped(k) += cotangent.v * along; + const long long at = sloped_at + row_at + k; + float tangent = __ldg(slot + sloped_at + at) * along; + if (directions != nullptr) tangent += __ldg(directions + at); + table_duration = table_duration + cotangent * T{__ldg(slot + at), tangent}; + } + } else { + sum_value(k) += cotangent; + if (p.blocks > 1) table_duration += cotangent * __ldg(slot + sloped_at + row_at + k); + } + }; + Z back_plus[N], back_minus[N], back_z[N]; +#pragma unroll + for (int j = 0; j < N; ++j) back_plus[j] = back_minus[j] = back_z[j] = {T(0.0f), T(0.0f)}; + Z across_plus = {T(0.0f), T(0.0f)}, across_minus = {T(0.0f), T(0.0f)}, along_z = {T(0.0f), T(0.0f)}; + T restoring = 0.0f; +#pragma unroll + for (int i = 0; i < N; ++i) { + const Z lp = zmul(zconj(carried), bplus[i]), lm = zmul(carried, bminus[i]); + const Z lz = zmul(zconj(spin), bz[i]); + Z mp = {T(0.0f), T(0.0f)}, mm = {T(0.0f), T(0.0f)}, mz = {T(0.0f), T(0.0f)}; +#pragma unroll + for (int j = 0; j < N; ++j) { + if (i < m && j < m) { + const int k = N * N + N + 2 * (i * m + j); + const Z x = {entry(row_at + k, along, true), entry(row_at + k + 1, along, true)}; + if (factored) { + const Z a = zmul(x, plus[j]), b = zmul(zconj(x), minus[j]); + mp = {mp.r + a.r, mp.i + a.i}; + mm = {mm.r + b.r, mm.i + b.i}; + } + const Z a = zmul(zconj(x), lp), b = zmul(x, lm); + back_plus[j] = {back_plus[j].r + a.r, back_plus[j].i + a.i}; + back_minus[j] = {back_minus[j].r + b.r, back_minus[j].i + b.i}; + const Z gp = zmul(lp, zconj(plus[j])), gm = zmul(zconj(lm), minus[j]); + take(k, gp.r + gm.r); + take(k + 1, gp.i + gm.i); + } + const T l = entry(row_at + i * N + j, along, true); + if (factored) mz = {mz.r + l * z[j].r, mz.i + l * z[j].i}; + back_z[j] = {back_z[j].r + l * lz.r, back_z[j].i + l * lz.i}; + take(i * N + j, lz.r * z[j].r + lz.i * z[j].i); + } + if (state == 0) { + g_equilibrium[i] += bz[i].r; + take(N * N + i, -(wout * bz[i].r)); + if (factored) restoring += bz[i].r * entry(row_at + N * N + i, along, true); + } + if (factored) { + if (i < m) { + const Z a = zmul(mp, zconj(bplus[i])), b = zmul(mm, zconj(bminus[i])); + across_plus = {across_plus.r + a.r, across_plus.i + a.i}; + across_minus = {across_minus.r + b.r, across_minus.i + b.i}; + } + const Z c = zmul(mz, zconj(bz[i])); + along_z = {along_z.r + c.r, along_z.i + c.i}; + } + } + T duration_gradient = table_duration; + if (factored) { + using V = num::Multi<4, T>; + V w; + Z c, sp; + pool_factors(p, V(dt, 0), V(voxel_b0, 1), V(damping_rate, 2), V(moved, 3), order, w, c, sp); + const V got = (c.r * across_plus.r - c.i * across_plus.i) + + (c.r * across_minus.r + c.i * across_minus.i) + (sp.r * along_z.r - sp.i * along_z.i) - + w * restoring; + duration_gradient = duration_gradient + got.d[0]; + g_b0 += got.d[1]; + g_damping += got.d[2]; + g_velocity += got.d[3]; + } + event_gradient(g.grad_duration, g.dgrad_duration, event, duration_gradient); +#pragma unroll + for (int i = 0; i < N; ++i) { + bplus[i] = back_plus[i]; + bminus[i] = back_minus[i]; + bz[i] = back_z[i]; + } + } + } + flush(); + // The equilibrium is also where every pool starts. +#pragma unroll + for (int i = 0; i < N; ++i) { + if (state == 0) g_equilibrium[i] += bz[i].r; + const T total = group_sum(g_equilibrium[i], width); + if (state == 0 && active) { + if constexpr (DUAL) { + grad_row[i] += total.v; + curve_row[i] += total.d; + } else { + grad_row[i] += total; + } + } + } + auto store_row = [&](int row, T lane_value) { + const T total = group_sum(lane_value, width); + if (state == 0 && active) add_both(g.grad_tissue, g.dgrad_tissue, row * n_atoms + atom, total); + }; + store_row(g.m0_row, g_m0); + if (!p.shimmed) { + store_row(g.b1_row, g_b1); + store_row(g.b1_phase_row, g_b1_phase); + } + store_row(g.b0_row, g_b0); + store_row(g.efficiency_row, g_efficiency); + store_row(g.diffusion_row, g_damping); + store_row(g.velocity_row, g_velocity); +} + +} // namespace epg_pooled diff --git a/src/blochsim/sequence/_pools_gpu.py b/src/blochsim/sequence/_pools_gpu.py index 434999e6..5668e9b2 100644 --- a/src/blochsim/sequence/_pools_gpu.py +++ b/src/blochsim/sequence/_pools_gpu.py @@ -21,7 +21,7 @@ import torch -from .._gpu_launch import Kernel, next_power_of_2 +from .._gpu_launch import Kernel, next_power_of_2, pooled_layout_floats from ._accelerators import _shim_count, _train_count from ._epg_gpu import _TRAJECTORY_BUDGET_BYTES, _atom_stride, _output_shape from ._parameters import ( @@ -367,9 +367,25 @@ def _adjoint( ) tiles = _tiles(pools.layout, state_count) planes = 12 if following else 6 - held = planes * tiles["P"] * tiles["S"] * max(1, event_count) + + # The layout keeps its own checkpoints in the buffer it is given; the tile + # kernels read a recording the forward kernel makes first. + def laid_floats(problems: int) -> int | None: + if device.type != "cuda": + return None + return pooled_layout_floats( + problems, event_count, tiles["n"], tiles["m"], tiles["S"], following + ) + + sample = min(total, 32) + laid = laid_floats(sample) + if laid is None: + held = planes * tiles["P"] * tiles["S"] * max(1, event_count) + else: + held = -(-laid // sample) wave = max(1, min(total, _TRAJECTORY_BUDGET_BYTES // (4 * held))) - trajectory = torch.empty(wave * held, dtype=torch.float32, device=device) + floats = wave * held if laid is None else laid_floats(wave) + trajectory = torch.empty(floats, dtype=torch.float32, device=device) grad_output = grad_output.resolve_conj() grad_real = grad_output.real.contiguous() @@ -401,17 +417,18 @@ def plane() -> tuple[torch.Tensor, ...]: base, tissue, events, output_count, state_count, pools, geometry, profile, lineshape, ) # fmt: skip - _pooled_kernel[(span,)]( - *inputs, - grad_real, - grad_imag, - trajectory, - *scalars, - planes=planes, - keep=True, - **switches, - **tiles, - ) + if laid is None: + _pooled_kernel[(span,)]( + *inputs, + grad_real, + grad_imag, + trajectory, + *scalars, + planes=planes, + keep=True, + **switches, + **tiles, + ) _pooled_adjoint_kernel[(span,)]( *inputs, grad_real, From 170288c0d1f02d1ee7b09f4d58d576f8fb2e8cd4 Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 05:08:43 +0200 Subject: [PATCH 12/16] Split PERK's features across programs when the voxels alone leave the card idle A launch of a few thousand voxels is a few dozen programs, which leaves most multiprocessors with nothing to run. Such a launch now also splits the features across a second axis of programs, each adding its share into an output that starts at zero, until there are four programs a multiprocessor. Co-Authored-By: Claude Opus 5.5 --- src/blochsim/_kernels.hpp | 10 +++--- src/blochsim/_perk_kernels.hpp | 52 +++++++++++++++++++++++----- src/blochsim/estimators/_perk_gpu.py | 41 ++++++++++++++++++---- 3 files changed, 84 insertions(+), 19 deletions(-) diff --git a/src/blochsim/_kernels.hpp b/src/blochsim/_kernels.hpp index 0a307d30..74e79831 100644 --- a/src/blochsim/_kernels.hpp +++ b/src/blochsim/_kernels.hpp @@ -788,7 +788,8 @@ BSK_HD void call_regress_kernel(const Arg* a) { a[9].i, a[10].i, static_cast(a[11].f), - a[12].i); + a[12].i, + a[13].i); } BSK_HD void call_regress_vjp_kernel(const Arg* a) { @@ -804,7 +805,8 @@ BSK_HD void call_regress_vjp_kernel(const Arg* a) { a[8].i, a[9].i, static_cast(a[10].f), - a[11].i); + a[11].i, + a[12].i); } #endif @@ -822,8 +824,8 @@ inline constexpr KernelInfo KERNELS[] = { {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1, BLOCHSIM_LANES__epg_jvp_kernel}, {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69, BLOCHSIM_LANES__pooled_kernel}, {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87, BLOCHSIM_LANES__pooled_adjoint_kernel}, - {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,threads", "pppppppiiiifi", 12, -1, -1, BLOCHSIM_LANES__regress_kernel}, - {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,threads", "ppppppiiiifi", 11, -1, -1, BLOCHSIM_LANES__regress_vjp_kernel}, + {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,splits,threads", "pppppppiiiifii", 13, -1, -1, BLOCHSIM_LANES__regress_kernel}, + {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,splits,threads", "ppppppiiiifii", 12, -1, -1, BLOCHSIM_LANES__regress_vjp_kernel}, }; #define BLOCHSIM_FOR_EACH_KERNEL(X) \ diff --git a/src/blochsim/_perk_kernels.hpp b/src/blochsim/_perk_kernels.hpp index d03622a6..d5806f96 100644 --- a/src/blochsim/_perk_kernels.hpp +++ b/src/blochsim/_perk_kernels.hpp @@ -11,6 +11,10 @@ // adjoint forms the angles again rather than keeping them. Every array is // contiguous and row-major. // +// A launch too small to fill the card splits the features across a second +// axis of programs: each forms its share of them and adds what it makes into +// an output that starts at zero, and the first adds the parameter mean. +// // Staged blocks are zero past the last feature, parameter or voxel, and a zero // weight is what removes a padded feature from both products, so the inner // loops carry no bounds. On a card a program is THREADS threads; on the host @@ -72,6 +76,31 @@ BSK_HD int clamp_width(std::int64_t left, int block) { return left < block ? static_cast(left) : block; } +// The features this program forms, ``[begin, end)``: whole blocks of +// ``BLOCK``, as evenly as ``splits`` programs share them. +template +BSK_HD void _feature_share(std::int64_t features, std::int64_t splits, std::int64_t& begin, + std::int64_t& end) { + const std::int64_t blocks = (features + BLOCK - 1) / BLOCK; + const std::int64_t share = (blocks + splits - 1) / splits * BLOCK; + begin = bsk::program_id(1) * share; + end = begin + share < features ? begin + share : features; +} + +// One output element: written where the features are not split, added to +// where they are. +BSK_HD void _emit(float* output, float value, std::int64_t splits) { + if (splits == 1) { + *output = value; + } else { +#if defined(BLOCHSIM_SIMT) + atomicAdd(output, value); +#else + *output += value; +#endif + } +} + // Signals of voxels ``first`` onward, contrasts ``c0`` onward, ``width`` of // them, into ``staged[contrast][voxel]``. Consecutive threads copy consecutive // elements of the rows, which are contiguous in memory. @@ -162,8 +191,11 @@ BSK_HD void _regress_kernel(const float* signal, const float* frequency, const f const float* feature_mean, const float* weight, const float* parameter_mean, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, - std::int64_t parameters, float scale, std::int64_t threads) { + std::int64_t parameters, float scale, std::int64_t splits, + std::int64_t threads) { static_cast(threads); + std::int64_t begin = 0, end = 0; + _feature_share(features, splits, begin, end); PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; PERK_SHARED_ROWS float frequencies[CONTRASTS][FEATURES + PAD]; PERK_SHARED_ROWS float weights[FEATURES][PARAMETERS]; @@ -183,8 +215,8 @@ BSK_HD void _regress_kernel(const float* signal, const float* frequency, const f } } } - for (std::int64_t f0 = 0; f0 < features; f0 += FEATURES) { - const int width = clamp_width(features - f0, FEATURES); + for (std::int64_t f0 = begin; f0 < end; f0 += FEATURES) { + const int width = clamp_width(end - f0, FEATURES); PERK_SYNC(); PERK_EACH_THREAD(t) { for (int j = t; j < FEATURES; j += THREADS) { @@ -227,8 +259,8 @@ BSK_HD void _regress_kernel(const float* signal, const float* frequency, const f #pragma unroll for (int k = 0; k < PARAMETERS; ++k) { if (voxel < voxels && k < held) { - output[voxel * parameters + p0 + k] = - total[t][r][k] + parameter_mean[p0 + k]; + const float mean = bsk::program_id(1) == 0 ? parameter_mean[p0 + k] : 0.0f; + _emit(&output[voxel * parameters + p0 + k], total[t][r][k] + mean, splits); } } } @@ -241,9 +273,11 @@ BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, con const float* weight, const float* cotangent, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, std::int64_t parameters, float scale, - std::int64_t threads) { + std::int64_t splits, std::int64_t threads) { static_cast(threads); constexpr int BLOCK = ADJOINT_FEATURES; + std::int64_t begin = 0, end = 0; + _feature_share(features, splits, begin, end); PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; PERK_SHARED_ROWS float frequencies[CONTRASTS][BLOCK + PAD]; PERK_SHARED_ROWS float back[BLOCK][GRADIENT + PAD]; @@ -264,8 +298,8 @@ BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, con } } } - for (std::int64_t f0 = 0; f0 < features; f0 += BLOCK) { - const int width = clamp_width(features - f0, BLOCK); + for (std::int64_t f0 = begin; f0 < end; f0 += BLOCK) { + const int width = clamp_width(end - f0, BLOCK); PERK_SYNC(); PERK_EACH_THREAD(t) { for (int j = t; j < BLOCK; j += THREADS) { @@ -350,7 +384,7 @@ BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, con #pragma unroll for (int c = 0; c < GRADIENT; ++c) { if (voxel < voxels && c < held) { - output[voxel * contrasts + g0 + c] = gradient[t][r][c]; + _emit(&output[voxel * contrasts + g0 + c], gradient[t][r][c], splits); } } } diff --git a/src/blochsim/estimators/_perk_gpu.py b/src/blochsim/estimators/_perk_gpu.py index 909ace13..d4ee06a6 100644 --- a/src/blochsim/estimators/_perk_gpu.py +++ b/src/blochsim/estimators/_perk_gpu.py @@ -18,15 +18,21 @@ __all__ = ["regress", "regress_vjp"] import math +from functools import cache import torch from .._gpu_launch import Kernel, cdiv -#: Threads per program, and the voxels a program holds: ``THREADS`` and -#: ``THREADS * VOXELS`` in ``_perk_kernels.hpp``. +#: Threads per program, the voxels a program holds, and the features the +#: forward pass and the adjoint form at once: ``THREADS``, ``THREADS * VOXELS``, +#: ``FEATURES`` and ``ADJOINT_FEATURES`` in ``_perk_kernels.hpp``. _THREADS = 64 _BLOCK_VOXELS = 128 +_FEATURES = 32 +_ADJOINT_FEATURES = 16 +#: Programs per multiprocessor a launch is split to reach. +_PROGRAMS_PER_SM = 4 _regress_kernel = Kernel("_regress_kernel") _regress_vjp_kernel = Kernel("_regress_vjp_kernel") @@ -37,6 +43,25 @@ def _ready(tensor: torch.Tensor) -> torch.Tensor: return tensor.detach().to(torch.float32).contiguous() +@cache +def _multiprocessors(device: torch.device) -> int: + if device.type != "cuda": + return 1 + return torch.cuda.get_device_properties(device).multi_processor_count + + +def _splits(device: torch.device, voxels: int, features: int, block: int) -> int: + """How many programs share a block of voxels' features. + + One, unless the voxels' programs alone leave the card's multiprocessors + short of work; then as many as fill them, a block of features at least + each. + """ + programs = cdiv(voxels, _BLOCK_VOXELS) + wanted = _PROGRAMS_PER_SM * _multiprocessors(device) + return max(1, min(cdiv(features, block), wanted // programs)) + + def regress( signals: torch.Tensor, frequency: torch.Tensor, @@ -56,11 +81,12 @@ def regress( voxels, contrasts = signals.shape features = frequency.shape[0] parameters = weight.shape[0] - output = torch.empty( + splits = _splits(signals.device, voxels, features, _FEATURES) + output = (torch.zeros if splits > 1 else torch.empty)( (voxels, parameters), dtype=torch.float32, device=signals.device ) if voxels: - _regress_kernel[(cdiv(voxels, _BLOCK_VOXELS),)]( + _regress_kernel[(cdiv(voxels, _BLOCK_VOXELS), splits)]( signals, _ready(frequency), _ready(phase), @@ -73,6 +99,7 @@ def regress( features, parameters, math.sqrt(2.0 / features), + splits, _THREADS, ) return output @@ -96,9 +123,10 @@ def regress_vjp( voxels, contrasts = signals.shape features = frequency.shape[0] parameters = weight.shape[0] - output = torch.empty_like(signals) + splits = _splits(signals.device, voxels, features, _ADJOINT_FEATURES) + output = (torch.zeros_like if splits > 1 else torch.empty_like)(signals) if voxels: - _regress_vjp_kernel[(cdiv(voxels, _BLOCK_VOXELS),)]( + _regress_vjp_kernel[(cdiv(voxels, _BLOCK_VOXELS), splits)]( signals, _ready(frequency), _ready(phase), @@ -110,6 +138,7 @@ def regress_vjp( features, parameters, math.sqrt(2.0 / features), + splits, _THREADS, ) return output From 3c5ebe7c5970db4e914cb495ecc1400230522a7a Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 05:08:43 +0200 Subject: [PATCH 13/16] Ship the card's kernels as blochsim-cuda12 and blochsim-cuda13 The card's module is a package of its own per CUDA major version, installed with blochsim[cu12] or blochsim[cu13] beside a torch of that major version: the same CMake project with BLOCHSIM_CUDA_PACKAGE set, which builds _gpu alone into blochsim_cudaNN/. It links the CUDA runtime torch's nvidia wheel brings, found by rpath, rather than carrying one; carries machine code for 7.5, 8.0 and 9.0 and PTX for 9.0, in compressed fatbinaries; and pins the blochsim it was built with. _gpu_launch loads the build for torch's CUDA major version and refuses one for another major or release, naming the extra to install. The base wheel is the CPU kernels alone. The wheels workflow builds each CUDA package in a manylinux container with the oldest toolkit of its major torch ships, installs it beside that torch on the oldest and newest Python, and publishes each project from an environment of its own. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/wheels.yml | 162 ++++++++++++++++++++++++++++++++++- CHANGELOG.md | 11 ++- CLAUDE.md | 27 ++++-- CMakeLists.txt | 58 ++++++++++--- pyproject.toml | 28 +++--- scripts/check_wheel.py | 103 ++++++++++++++-------- src/blochsim/_gpu.cu | 5 +- src/blochsim/_gpu_launch.py | 73 +++++++++++++++- src/cuda/12/pyproject.toml | 52 +++++++++++ src/cuda/13/pyproject.toml | 52 +++++++++++ src/cuda/__init__.py | 5 ++ 11 files changed, 494 insertions(+), 82 deletions(-) create mode 100644 src/cuda/12/pyproject.toml create mode 100644 src/cuda/13/pyproject.toml create mode 100644 src/cuda/__init__.py diff --git a/.github/workflows/wheels.yml b/.github/workflows/wheels.yml index 50bda729..92d932e8 100644 --- a/.github/workflows/wheels.yml +++ b/.github/workflows/wheels.yml @@ -14,6 +14,7 @@ on: - "src/blochsim/*.hpp" - "src/blochsim/*.cu" - "src/blochsim/*.cu.in" + - "src/cuda/**" - scripts/check_wheel.py - .github/workflows/wheels.yml workflow_dispatch: @@ -54,9 +55,7 @@ jobs: wheels: name: ${{ matrix.label }} runs-on: ${{ matrix.os }} - # The x86-64 manylinux wheel compiles the GPU kernels for every listed - # architecture. - timeout-minutes: 150 + timeout-minutes: 60 strategy: fail-fast: false matrix: @@ -92,9 +91,144 @@ jobs: name: dist-${{ matrix.os }} path: wheelhouse/*.whl + # The GPU kernels are a package of their own per CUDA major version, + # blochsim-cuda12 and blochsim-cuda13 (src/cuda/), which blochsim loads when + # torch is built for the same major version: `pip install blochsim[cu12]`. + # Each carries the card's module alone and links the CUDA runtime that + # torch's CUDA builds bring, and is compiled with the oldest toolkit of its + # major that torch ships, so it runs against every torch of it. + wheels-cuda: + name: Build blochsim-cuda${{ matrix.cuda }} (linux x86_64) + runs-on: ubuntu-latest + timeout-minutes: 150 + container: + image: quay.io/pypa/manylinux_2_28_x86_64 + strategy: + fail-fast: false + matrix: + include: + # CUDA 12.6's nvcc takes GCC 13 at the newest as its host compiler. + - cuda: "12" + toolkit: "12-6" + home: /usr/local/cuda-12.6 + host: gcc-toolset-13-gcc-c++ + hostcxx: /opt/rh/gcc-toolset-13/root/usr/bin/g++ + - cuda: "13" + toolkit: "13-0" + home: /usr/local/cuda-13.0 + host: "" + hostcxx: "" + env: + CUDA_HOME: ${{ matrix.home }} + steps: + - uses: actions/checkout@v7 + with: + # setuptools-scm reads the version from the tag history, and the + # CUDA build pins the blochsim of the same version. + fetch-depth: 0 + - name: Install the CUDA toolkit + run: | + dnf install -y dnf-plugins-core + dnf config-manager --add-repo \ + https://developer.download.nvidia.com/compute/cuda/repos/rhel8/x86_64/cuda-rhel8.repo + dnf install -y --nogpgcheck cuda-nvcc-${{ matrix.toolkit }} cuda-cudart-devel-${{ matrix.toolkit }} ${{ matrix.host }} + echo $CUDA_HOME/bin >> $GITHUB_PATH + echo /opt/python/cp312-cp312/bin >> $GITHUB_PATH + - name: Build the wheel + env: + CUDACXX: ${{ matrix.home }}/bin/nvcc + run: | + if [ -n "${{ matrix.hostcxx }}" ]; then export CUDAHOSTCXX=${{ matrix.hostcxx }}; fi + python -m pip install build + python -m build --wheel --outdir dist/ src/cuda/${{ matrix.cuda }} + # The CUDA runtime is excluded rather than vendored: torch already brings + # it, from the nvidia wheel the module's rpath points into. + - name: Repair to a manylinux tag + run: | + python -m pip install auditwheel + LD_LIBRARY_PATH=$CUDA_HOME/lib64 python -m auditwheel repair dist/*.whl \ + -w wheelhouse/ --plat manylinux_2_28_x86_64 --exclude libcudart.so.${{ matrix.cuda }} + # The card's module and the package around it, and nothing vendored. + # PyPI takes a file of 100 MB at most unless a project is granted more. + - name: Inspect the CUDA wheel + run: | + python -m auditwheel show wheelhouse/*.whl + python -m zipfile -l wheelhouse/*.whl | tee contents.txt + ! grep -Ei 'libcudart|\.libs/' contents.txt + grep -q 'blochsim_cuda${{ matrix.cuda }}/_gpu' contents.txt + ! grep -E '^ *blochsim/' contents.txt + size=$(stat -c %s wheelhouse/*.whl) + echo "wheel: $size bytes" + test "$size" -lt 100000000 + - uses: actions/upload-artifact@v7 + with: + name: cuda-wheel-${{ matrix.cuda }} + path: wheelhouse/*.whl + + # What a user installs: the blochsim wheel with its extra, beside torch's + # build for the same CUDA major version, on the oldest and the newest Python + # the package supports. The runners have no card, so this proves that the + # install resolves -- the CUDA build's pins against torch's -- that blochsim + # loads the CUDA build, and that the runtime it links is the one torch's + # nvidia wheel installed, found from the module's own directory. + test-cuda-wheels: + name: Test blochsim[cu${{ matrix.cuda }}] on Python ${{ matrix.python }} + needs: [wheels, wheels-cuda] + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + cuda: ["12", "13"] + python: ["3.10", "3.14"] + include: + - cuda: "12" + torch: cu126 + - cuda: "13" + torch: cu130 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-python@v7 + with: + python-version: ${{ matrix.python }} + - uses: actions/download-artifact@v8 + with: + name: dist-ubuntu-latest + path: base + - uses: actions/download-artifact@v8 + with: + name: cuda-wheel-${{ matrix.cuda }} + path: cuda + # Both wheels are named as files, so pip takes them whatever their + # version, and the extra's requirement is the CUDA wheel given beside it. + - name: Install blochsim[cu${{ matrix.cuda }}] with torch's ${{ matrix.torch }} build + run: | + python -m pip install --upgrade pip + pip install "$(ls base/*manylinux*x86_64*.whl)[cu${{ matrix.cuda }}]" cuda/*.whl \ + torch --index-url https://download.pytorch.org/whl/${{ matrix.torch }} \ + --extra-index-url https://pypi.org/simple + pip list | grep -Ei '^(torch|blochsim|nvidia-cuda-runtime)' + - name: The card's module loads bare, against torch's runtime + run: python scripts/check_wheel.py + - name: blochsim loads the CUDA build for torch's major version + run: | + python -c " + import torch + from blochsim import _gpu_launch + module = _gpu_launch._module('cuda') + print(torch.__version__, module.__name__) + assert torch.version.cuda.split('.')[0] == '${{ matrix.cuda }}', torch.version.cuda + assert module.__name__ == 'blochsim_cuda${{ matrix.cuda }}._gpu', module.__name__ + assert _gpu_launch.available() + " + + # PyPI hands a trusted publisher a token for one project, the one whose + # publisher matches the job's repository, workflow and environment, so each + # project is published from an environment of its own: pypi for blochsim, + # pypi-cuda12 and pypi-cuda13 for the CUDA builds. The CUDA builds pin the + # blochsim they were built with, so they follow it. publish: name: Publish to PyPI - needs: [sdist, wheels] + needs: [sdist, wheels, wheels-cuda, test-cuda-wheels] if: startsWith(github.ref, 'refs/tags/v') runs-on: ubuntu-latest environment: @@ -111,3 +245,23 @@ jobs: path: dist - uses: pypa/gh-action-pypi-publish@release/v1 + + publish-cuda: + name: Publish blochsim-cuda${{ matrix.cuda }} to PyPI + needs: [publish] + if: startsWith(github.ref, 'refs/tags/v') + runs-on: ubuntu-latest + strategy: + matrix: + cuda: ["12", "13"] + environment: + name: pypi-cuda${{ matrix.cuda }} + url: https://pypi.org/p/blochsim-cuda${{ matrix.cuda }} + permissions: + id-token: write + steps: + - uses: actions/download-artifact@v8 + with: + name: cuda-wheel-${{ matrix.cuda }} + path: dist + - uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/CHANGELOG.md b/CHANGELOG.md index 210d8a77..27db2a3c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,10 +9,13 @@ the same modules and names beneath it. - **The GPU kernels are CUDA compiled ahead of time, and Triton is not used.** The EPG, pooled and PERK kernels are C++ (`_epg_kernels.hpp`, - `_pools_kernels.hpp`, `_perk_kernels.hpp`) compiled by `nvcc` into - `blochsim._gpu`, which links the CUDA runtime statically; the x86-64 - manylinux wheel carries it, and a source build compiles it wherever CMake - finds `nvcc`. No kernel is compiled at the first call. The same kernels are + `_pools_kernels.hpp`, `_perk_kernels.hpp`, and the layouts of `_layout.hpp`) + compiled by `nvcc` into a module of their own per CUDA major version: + `pip install blochsim[cu12]` or `blochsim[cu13]` installs `blochsim-cuda12` + or `blochsim-cuda13` beside a torch of the same major version, which links + the CUDA runtime that torch brings and carries code for 7.5, 8.0 and 9.0 + cards. A source build compiles it beside the package wherever CMake finds + `nvcc`. No kernel is compiled at the first call. The same kernels are compiled for the host as `blochsim._gpu_host`, which the suite holds to the C++ kernels; the `interpreted` marker is gone. A launch is at most 1024 threads, so an EPG run of more than 1024 state orders on a card is refused. diff --git a/CLAUDE.md b/CLAUDE.md index 1681655f..1583c799 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -221,10 +221,25 @@ plain CPython extensions against the **stable ABI** from 3.10 on: they call no PyTorch API and link no PyTorch library, which is why one `cp310-abi3` wheel per platform serves every supported interpreter and why that wheel is a couple of megabytes rather than the size of libtorch. Keep it that way — a `#include -` in any of them ends all of that. The GPU module links the CUDA -runtime statically, so a machine needs the driver and nothing else, and the -x86-64 manylinux wheel is the one built with it. +` in any of them ends all of that. + +The card's module is a package of its own per CUDA major version, +`blochsim-cuda12` and `blochsim-cuda13` (`src/cuda/12`, `src/cuda/13`), which +`pip install blochsim[cu12]` or `blochsim[cu13]` installs beside a torch of +the same major version. Each is this CMake project with `BLOCHSIM_CUDA_PACKAGE` +set, which builds `_gpu` alone into `blochsim_cudaNN/`; it links the CUDA +runtime dynamically -- torch's nvidia wheel, found by rpath -- carries machine +code for 7.5, 8.0 and 9.0 and PTX for 9.0, and pins the blochsim it was built +with. `_gpu_launch` loads the build for torch's CUDA major version and refuses +one for another major or another release, naming the extra to install; with +none installed it takes the `_gpu` a source build leaves beside the package. -Wheels are built by cibuildwheel and published to PyPI by trusted publishing on -a `v*.*.*` tag. `scripts/check_wheel.py` loads each compiled kernel by path, -without importing the package, and is what every built wheel is tested with. +```sh +python -m build --wheel src/cuda/12 # with CUDA 12.6's nvcc as CUDACXX +``` + +Wheels are built by cibuildwheel, the CUDA packages in a manylinux container +of their own, and published to PyPI by trusted publishing on a `v*.*.*` tag, +each project from its own environment (`pypi`, `pypi-cuda12`, `pypi-cuda13`). +`scripts/check_wheel.py` loads each compiled kernel by path, without importing +the package, and is what every built wheel is tested with. diff --git a/CMakeLists.txt b/CMakeLists.txt index 4151214e..d9197c45 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -54,13 +54,22 @@ function(blochsim_add_kernel name source) install(TARGETS ${name} DESTINATION blochsim) endfunction() -blochsim_add_kernel(_epg_cpu src/blochsim/_epg_cpu.cpp) -blochsim_add_kernel(_perk_cpu src/blochsim/_perk_cpu.cpp) +# The GPU kernels are a package of their own per CUDA major version, +# blochsim-cuda12 and blochsim-cuda13 (src/cuda/): the same project with +# BLOCHSIM_CUDA_PACKAGE set, which builds the card's module into +# blochsim_cudaNN/ and nothing else. Empty builds blochsim itself. +set(BLOCHSIM_CUDA_PACKAGE "" CACHE STRING + "CUDA major version of the blochsim-cudaNN wheel being built; empty for blochsim") + +if(NOT BLOCHSIM_CUDA_PACKAGE) + blochsim_add_kernel(_epg_cpu src/blochsim/_epg_cpu.cpp) + blochsim_add_kernel(_perk_cpu src/blochsim/_perk_cpu.cpp) +endif() # The GPU kernels compiled for the host, one program at a time. Nothing in the # package dispatches to them; they are how the suite checks the GPU kernels on # a machine with no card. Linux only, like the card's build. -if(CMAKE_SYSTEM_NAME STREQUAL "Linux") +if(CMAKE_SYSTEM_NAME STREQUAL "Linux" AND NOT BLOCHSIM_CUDA_PACKAGE) set(_blochsim_host_default ON) else() set(_blochsim_host_default OFF) @@ -75,7 +84,7 @@ endif() # compiler is found unless BLOCHSIM_CUDA says otherwise. include(CheckLanguage) check_language(CUDA) -if(CMAKE_CUDA_COMPILER AND CMAKE_SYSTEM_NAME STREQUAL "Linux") +if(BLOCHSIM_CUDA_PACKAGE OR (CMAKE_CUDA_COMPILER AND CMAKE_SYSTEM_NAME STREQUAL "Linux")) set(_blochsim_cuda_default ON) else() set(_blochsim_cuda_default OFF) @@ -83,13 +92,22 @@ endif() option(BLOCHSIM_CUDA "Compile the GPU kernels for the card" ${_blochsim_cuda_default}) if(BLOCHSIM_CUDA) - # Real code for each generation and PTX for the newest, which the driver - # compiles for a card newer than any listed. + # Machine code for 7.5, 8.0 and 9.0, and PTX for 9.0, which the driver + # compiles for a card newer than any listed. Code for 8.0 runs on every + # 8.x card, so 8.6 and 8.9 take it. if(NOT DEFINED CMAKE_CUDA_ARCHITECTURES) - set(CMAKE_CUDA_ARCHITECTURES 75-real 80-real 86-real 89-real 90) + set(CMAKE_CUDA_ARCHITECTURES 75-real 80-real 90) endif() enable_language(CUDA) find_package(CUDAToolkit REQUIRED) + if(BLOCHSIM_CUDA_PACKAGE) + if(NOT CUDAToolkit_VERSION_MAJOR STREQUAL BLOCHSIM_CUDA_PACKAGE) + message(FATAL_ERROR "blochsim-cuda${BLOCHSIM_CUDA_PACKAGE} is being built with CUDA ${CUDAToolkit_VERSION}") + endif() + set(_blochsim_gpu_destination "blochsim_cuda${BLOCHSIM_CUDA_PACKAGE}") + else() + set(_blochsim_gpu_destination "blochsim") + endif() set(CMAKE_CUDA_STANDARD 17) set(CMAKE_CUDA_STANDARD_REQUIRED ON) @@ -150,10 +168,26 @@ if(BLOCHSIM_CUDA) ${_blochsim_gpu_sources} ) target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim") - # The runtime is linked in, so a machine needs the driver and nothing else. - target_link_libraries(_gpu PRIVATE CUDA::cudart_static) + # The runtime is the one torch's CUDA builds bring in the nvidia wheels: + # torch loads it before this module, and the rpath covers a bare import + # from site-packages. CUDA 12's runtime wheel keeps it in + # nvidia/cuda_runtime, CUDA 13's in nvidia/cu13. + set_target_properties(_gpu PROPERTIES CUDA_RUNTIME_LIBRARY Shared) + target_link_libraries(_gpu PRIVATE CUDA::cudart) + set_property(TARGET _gpu APPEND PROPERTY INSTALL_RPATH + "$ORIGIN/../nvidia/cuda_runtime/lib" + "$ORIGIN/../nvidia/cu13/lib") # Warning 221 is a double constant that rounds to zero in float, which the - # kernels write deliberately as the floor of a clamp. - target_compile_options(_gpu PRIVATE $<$:-diag-suppress=221>) - install(TARGETS _gpu DESTINATION blochsim) + # kernels write deliberately as the floor of a clamp. The fatbinaries are + # compressed, which keeps three architectures' machine code in a wheel + # PyPI takes. + target_compile_options(_gpu PRIVATE + $<$:-diag-suppress=221;-Xfatbin=-compress-all>) + install(TARGETS _gpu DESTINATION ${_blochsim_gpu_destination}) + if(BLOCHSIM_CUDA_PACKAGE) + install(FILES src/cuda/__init__.py DESTINATION ${_blochsim_gpu_destination}) + install(FILES LICENSE.txt DESTINATION "${SKBUILD_METADATA_DIR}/licenses") + endif() +elseif(BLOCHSIM_CUDA_PACKAGE) + message(FATAL_ERROR "BLOCHSIM_CUDA_PACKAGE=${BLOCHSIM_CUDA_PACKAGE} needs BLOCHSIM_CUDA=ON") endif() diff --git a/pyproject.toml b/pyproject.toml index 60d8d342..34a260b0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,6 +54,16 @@ dependencies = [ ] [project.optional-dependencies] +# The GPU kernels, compiled for the CUDA major version torch is built for: a +# package per major version (src/cuda/), which blochsim loads when torch's +# matches it. Install the one for the torch beside it -- cu12 with a CUDA 12 +# build of torch, cu13 with a CUDA 13 one. +cu12 = [ +"blochsim-cuda12; sys_platform == 'linux' and platform_machine == 'x86_64'", +] +cu13 = [ +"blochsim-cuda13; sys_platform == 'linux' and platform_machine == 'x86_64'", +] # What it takes to read a Pulseq .seq file. The format's reference # implementation parses it and computes its trajectory; nothing else in the # package imports it. @@ -143,27 +153,15 @@ test-command = "python {project}/scripts/check_wheel.py" # publishes no musllinux wheel for that install to resolve. The workflow loads # the musl kernels in an Alpine container instead, where nothing resolves. test-skip = "*musllinux*" -# The GPU kernels compiled for the host are for the suite, not for a user. -config-settings = {"cmake.define.BLOCHSIM_HOST_KERNELS" = "OFF"} +# The GPU kernels compiled for the host are for the suite, not for a user, and +# those for the card are the blochsim-cudaNN packages' (src/cuda/). +config-settings = {"cmake.define.BLOCHSIM_HOST_KERNELS" = "OFF", "cmake.define.BLOCHSIM_CUDA" = "OFF"} [tool.cibuildwheel.linux] # scikit-build-core installs cmake and ninja as build-time wheels inside the # manylinux container; the image's compiler is what builds the kernels. archs = ["auto64"] -# The GPU kernels, on the one platform PyTorch publishes CUDA builds for that -# a manylinux image can compile: the toolkit's compiler and its static -# runtime, nothing the wheel carries beyond the module itself. -[[tool.cibuildwheel.overrides]] -select = "*-manylinux_x86_64" -before-all = [ - "dnf install -y dnf-plugins-core", - "dnf config-manager --add-repo https://developer.download.nvidia.com/compute/cuda/repos/rhel8/x86_64/cuda-rhel8.repo", - "dnf install -y cuda-nvcc-12-8 cuda-cudart-devel-12-8", -] -environment = {PATH = "/usr/local/cuda-12.8/bin:$PATH", CUDACXX = "/usr/local/cuda-12.8/bin/nvcc"} -config-settings = {"cmake.define.BLOCHSIM_HOST_KERNELS" = "OFF", "cmake.define.BLOCHSIM_CUDA" = "ON"} - [tool.cibuildwheel.macos] # An Apple silicon runner builds its own wheel and cross-compiles the Intel # one beside it; cibuildwheel skips the test for the architecture the runner diff --git a/scripts/check_wheel.py b/scripts/check_wheel.py index 4a7926c1..ef4f8d1e 100644 --- a/scripts/check_wheel.py +++ b/scripts/check_wheel.py @@ -1,10 +1,13 @@ """Load the compiled kernels out of an installed wheel, and report their size. Run by cibuildwheel against every wheel it builds, on the interpreter that -wheel claims to support. It loads each extension by file path rather than by -importing :mod:`blochsim`, so the check needs no PyTorch and says something +wheel claims to support, and against each blochsim-cudaNN wheel installed +beside a CUDA build of torch. It loads each extension by file path rather than +by importing :mod:`blochsim`, so the check needs no PyTorch and says something about the binary alone: that the stable-ABI module initialises on this -interpreter, and that it carries no vendored library. +interpreter, that it carries no vendored library, and -- for the card's module +-- that the CUDA runtime it links is the one in torch's nvidia wheels, found +from the module's own directory. python scripts/check_wheel.py """ @@ -18,49 +21,81 @@ #: A kernel calls no PyTorch API and links no PyTorch library, so anything #: approaching this size means something was bundled that should not have been. LARGEST_REASONABLE_BYTES = 16 * 1024 * 1024 -#: The GPU kernels carry machine code for each architecture they were compiled -#: for, and the CUDA runtime linked in. +#: The card's module carries compressed machine code for each architecture it +#: was compiled for, and links the CUDA runtime rather than carrying it. LARGEST_REASONABLE_GPU_BYTES = 96 * 1024 * 1024 +#: The CUDA major versions a blochsim-cudaNN package is built for. +CUDA_MAJORS = (12, 13) -def kernel_directory() -> pathlib.Path: - """Where the installed package lives, without executing it.""" - spec = importlib.util.find_spec("blochsim") +def package_directory(name: str) -> pathlib.Path | None: + """Where an installed package lives, without executing it.""" + spec = importlib.util.find_spec(name) if spec is None or not spec.submodule_search_locations: - raise SystemExit("blochsim is not installed") + return None return pathlib.Path(next(iter(spec.submodule_search_locations))) -def main() -> int: - """Load every compiled kernel in the installed package.""" - root = kernel_directory() - kernels = sorted(root.glob("_*_cpu.*")) - kernels = [path for path in kernels if path.suffix in {".so", ".pyd", ".dylib"}] - if len(kernels) != 2: - raise SystemExit(f"expected two kernels in {root}, found {kernels}") - # Present where the wheel was built with a CUDA compiler. Loading it needs - # no card: nothing reaches the driver until a launch. - kernels += [ +def extensions(root: pathlib.Path, pattern: str) -> list[pathlib.Path]: + """The compiled extensions in ``root`` whose names match ``pattern``.""" + return [ path - for path in sorted(root.glob("_gpu.*")) + for path in sorted(root.glob(pattern)) if path.suffix in {".so", ".pyd", ".dylib"} ] + +def load(path: pathlib.Path, package: str, largest: int) -> None: + """Initialise the extension at ``path`` and hold it to a size.""" + name = path.name.split(".")[0] + spec = importlib.util.spec_from_file_location(f"{package}.{name}", path) + if spec is None or spec.loader is None: + raise SystemExit(f"no loader for {path}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + size = path.stat().st_size + print(f"{package}/{path.name}: {size} bytes, {len(dir(module))} attributes") + if size > largest: + raise SystemExit(f"{path.name} is {size} bytes; something was bundled") + + +def runtime_loaded(major: int) -> pathlib.Path: + """The CUDA runtime this process mapped, which must be an nvidia wheel's.""" + soname = f"libcudart.so.{major}" + for line in pathlib.Path("/proc/self/maps").read_text().splitlines(): + if line.endswith(soname) or f"/{soname}" in line: + path = pathlib.Path(line.split()[-1]) + if "nvidia" not in path.parts: + raise SystemExit(f"{soname} came from {path}, not torch's nvidia wheel") + return path + raise SystemExit(f"{soname} is not mapped") + + +def main() -> int: + """Load every compiled kernel installed.""" + root = package_directory("blochsim") + if root is None: + raise SystemExit("blochsim is not installed") + kernels = extensions(root, "_*_cpu.*") + if len(kernels) != 2: + raise SystemExit(f"expected two kernels in {root}, found {kernels}") for path in kernels: - name = path.name.split(".")[0] - spec = importlib.util.spec_from_file_location(f"blochsim.{name}", path) - if spec is None or spec.loader is None: - raise SystemExit(f"no loader for {path}") - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - - size = path.stat().st_size - print(f"{path.name}: {size} bytes, {len(dir(module))} attributes") - largest = ( - LARGEST_REASONABLE_GPU_BYTES if name == "_gpu" else LARGEST_REASONABLE_BYTES - ) - if size > largest: - raise SystemExit(f"{path.name} is {size} bytes; something was bundled") + load(path, "blochsim", LARGEST_REASONABLE_BYTES) + # A build from source keeps the card's module beside the package. + for path in extensions(root, "_gpu.*"): + load(path, "blochsim", LARGEST_REASONABLE_GPU_BYTES) + + # Loading the card's module needs no card: nothing reaches the driver + # until a launch. It does need the runtime, which its rpath finds. + for major in CUDA_MAJORS: + build = package_directory(f"blochsim_cuda{major}") + if build is None: + continue + modules = extensions(build, "_gpu.*") + if len(modules) != 1: + raise SystemExit(f"expected the card's module in {build}, found {modules}") + load(modules[0], f"blochsim_cuda{major}", LARGEST_REASONABLE_GPU_BYTES) + print(f"blochsim_cuda{major}: the runtime is {runtime_loaded(major)}") print(f"ok on {sys.implementation.name} {'.'.join(map(str, sys.version_info[:3]))}") return 0 diff --git a/src/blochsim/_gpu.cu b/src/blochsim/_gpu.cu index c47527f9..016359cd 100644 --- a/src/blochsim/_gpu.cu +++ b/src/blochsim/_gpu.cu @@ -2,9 +2,8 @@ // // One CUDA block runs one program: its threads hold a tile an element each, // x along a row and y across the rows. The module links the CUDA runtime -// statically, so a machine needs the driver and nothing else, and it calls no -// PyTorch API: a launch takes the addresses of the tensors' data and the -// stream PyTorch is queueing on. +// torch's CUDA builds bring, and calls no PyTorch API: a launch takes the +// addresses of the tensors' data and the stream PyTorch is queueing on. // // The kernels themselves are compiled one to a file, from _gpu_kernel.cu.in. #define BLOCHSIM_TABLE_ONLY 1 diff --git a/src/blochsim/_gpu_launch.py b/src/blochsim/_gpu_launch.py index 6493384e..31ffeb3c 100644 --- a/src/blochsim/_gpu_launch.py +++ b/src/blochsim/_gpu_launch.py @@ -11,9 +11,12 @@ __all__: list[str] = [] +import importlib +import importlib.util from collections.abc import Iterator from contextlib import contextmanager from functools import cache +from importlib import metadata from typing import Any import torch @@ -29,12 +32,65 @@ def cdiv(numerator: int, denominator: int) -> int: return -(-int(numerator) // int(denominator)) -@cache -def _module(device_type: str) -> Any: - if device_type == "cuda": +# The CUDA major versions a blochsim-cudaNN package is built for (src/cuda/). +CUDA_MAJORS = (12, 13) + + +class CudaBuildMismatch(ImportError): + """A CUDA build is installed, but not the one this torch and this blochsim load.""" + + +def _installed(name: str) -> str | None: + try: + return metadata.version(name) + except metadata.PackageNotFoundError: + return None + + +def _cuda_module() -> Any: + """The card's module: the CUDA build for torch's CUDA major version. + + ``blochsim[cu12]`` and ``blochsim[cu13]`` install it as a package of its + own. Where none is installed, a build from source keeps it beside the + package. A build for another CUDA major version, or from another release, + is refused by name rather than loaded. + """ + builds = [ + major + for major in CUDA_MAJORS + if importlib.util.find_spec(f"blochsim_cuda{major}") is not None + ] + if not builds: from blochsim import _gpu return _gpu + major = ( + max(builds) + if torch.version.cuda is None + else int(torch.version.cuda.split(".")[0]) + ) + if major not in builds: + names = ", ".join(f"blochsim-cuda{m}" for m in builds) + raise CudaBuildMismatch( + f"blochsim: torch is built for CUDA {torch.version.cuda}, and the CUDA " + f"build installed is {names}; install blochsim[cu{major}]." + ) + release, built = ( + _installed(name) for name in ("blochsim", f"blochsim-cuda{major}") + ) + if built != release: + raise CudaBuildMismatch( + f"blochsim: blochsim {release} is installed with blochsim-cuda{major} " + f"{built}, whose kernels are another release's; install " + f"blochsim[cu{major}]=={release}." + ) + return importlib.import_module(f"blochsim_cuda{major}._gpu") + + +@cache +def _module(device_type: str) -> Any: + if device_type == "cuda": + return _cuda_module() from blochsim import _gpu_host return _gpu_host @@ -53,9 +109,18 @@ def _signature(name: str) -> tuple[tuple[str, ...], str]: @cache def available() -> bool: - """Whether this installation carries the kernels compiled for a card.""" + """Whether this installation carries the kernels compiled for a card. + + Raises + ------ + CudaBuildMismatch + A CUDA build is installed for another CUDA major version than torch's, + or from another release than this blochsim. + """ try: _module("cuda") + except CudaBuildMismatch: + raise except ImportError: return False return True diff --git a/src/cuda/12/pyproject.toml b/src/cuda/12/pyproject.toml new file mode 100644 index 00000000..a3ff0e0b --- /dev/null +++ b/src/cuda/12/pyproject.toml @@ -0,0 +1,52 @@ +# blochsim's GPU kernels compiled with CUDA 12, which `blochsim[cu12]` installs. +# The same CMake project as blochsim's own pyproject, configured with +# BLOCHSIM_CUDA_PACKAGE, which builds the card's module into blochsim_cuda12/ +# and nothing else; blochsim loads it from there when torch is a CUDA 12 +# build. Built with CUDA 12.6, the runtime torch's oldest CUDA 12 build ships, +# so it runs against every CUDA 12 torch. +[build-system] +requires = ["scikit-build-core>=0.11", "setuptools-scm>=8"] +build-backend = "scikit_build_core.build" + +[project] +name = "blochsim-cuda12" +dynamic = ["version", "dependencies"] +description = "blochsim's GPU kernels compiled with CUDA 12, for blochsim[cu12]" +readme = { text = "blochsim's GPU kernels compiled with CUDA 12. Install them as `pip install blochsim[cu12]`; see https://github.com/pulserver/blochsim.", content-type = "text/markdown" } +license = "MIT" +requires-python = ">=3.10" + +[project.urls] +Homepage = "https://github.com/pulserver/blochsim" + +[tool.scikit-build] +cmake.source-dir = "../../.." +cmake.version = ">=3.26" +cmake.build-type = "Release" +build-dir = "../../../build/cuda12-{wheel_tag}" +cmake.define.BLOCHSIM_CUDA_PACKAGE = "12" +cmake.define.BLOCHSIM_HOST_KERNELS = "OFF" +wheel.packages = [] +# The module is a CPython extension against the stable ABI from 3.10 on. +wheel.py-api = "cp310" + +[[tool.dynamic-metadata]] +field = "version" +provider = "scikit_build_core.metadata.setuptools_scm" + +# The module is blochsim's launcher ABI, so it is pinned to the blochsim it +# was built with. The runtime is the one torch's CUDA 12 builds bring. +[[tool.dynamic-metadata]] +field = "dependencies" +provider = "scikit_build_core.metadata.template" +result = [ + "blochsim=={project[version]}", + "nvidia-cuda-runtime-cu12>=12.6", +] + +# The version blochsim's own pyproject computes, so the pin above matches it. +[tool.setuptools_scm] +root = "../../.." +version_scheme = "python-simplified-semver" +local_scheme = "no-local-version" +fallback_version = "v99-dev" diff --git a/src/cuda/13/pyproject.toml b/src/cuda/13/pyproject.toml new file mode 100644 index 00000000..e80caef6 --- /dev/null +++ b/src/cuda/13/pyproject.toml @@ -0,0 +1,52 @@ +# blochsim's GPU kernels compiled with CUDA 13, which `blochsim[cu13]` installs. +# The same CMake project as blochsim's own pyproject, configured with +# BLOCHSIM_CUDA_PACKAGE, which builds the card's module into blochsim_cuda13/ +# and nothing else; blochsim loads it from there when torch is a CUDA 13 +# build. Built with CUDA 13.0, the runtime torch's oldest CUDA 13 build ships, +# so it runs against every CUDA 13 torch. +[build-system] +requires = ["scikit-build-core>=0.11", "setuptools-scm>=8"] +build-backend = "scikit_build_core.build" + +[project] +name = "blochsim-cuda13" +dynamic = ["version", "dependencies"] +description = "blochsim's GPU kernels compiled with CUDA 13, for blochsim[cu13]" +readme = { text = "blochsim's GPU kernels compiled with CUDA 13. Install them as `pip install blochsim[cu13]`; see https://github.com/pulserver/blochsim.", content-type = "text/markdown" } +license = "MIT" +requires-python = ">=3.10" + +[project.urls] +Homepage = "https://github.com/pulserver/blochsim" + +[tool.scikit-build] +cmake.source-dir = "../../.." +cmake.version = ">=3.26" +cmake.build-type = "Release" +build-dir = "../../../build/cuda13-{wheel_tag}" +cmake.define.BLOCHSIM_CUDA_PACKAGE = "13" +cmake.define.BLOCHSIM_HOST_KERNELS = "OFF" +wheel.packages = [] +# The module is a CPython extension against the stable ABI from 3.10 on. +wheel.py-api = "cp310" + +[[tool.dynamic-metadata]] +field = "version" +provider = "scikit_build_core.metadata.setuptools_scm" + +# The module is blochsim's launcher ABI, so it is pinned to the blochsim it +# was built with. The runtime is the one torch's CUDA 13 builds bring. +[[tool.dynamic-metadata]] +field = "dependencies" +provider = "scikit_build_core.metadata.template" +result = [ + "blochsim=={project[version]}", + "nvidia-cuda-runtime>=13.0", +] + +# The version blochsim's own pyproject computes, so the pin above matches it. +[tool.setuptools_scm] +root = "../../.." +version_scheme = "python-simplified-semver" +local_scheme = "no-local-version" +fallback_version = "v99-dev" diff --git a/src/cuda/__init__.py b/src/cuda/__init__.py new file mode 100644 index 00000000..6e921c4b --- /dev/null +++ b/src/cuda/__init__.py @@ -0,0 +1,5 @@ +"""blochsim's GPU kernels compiled with CUDA, for the CUDA major version in this package's name. + +blochsim loads ``_gpu`` from this package when torch is built for the same CUDA +major version. ``pip install blochsim[cu12]`` or ``blochsim[cu13]`` installs it. +""" From 90d402b38b8d604ea6daa22392119d47908aa08c Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 06:21:06 +0200 Subject: [PATCH 14/16] Contract a wide interval one direction at a time A three-pool interval at float, or any at num::Dual, is a multi-valued operator whose entries are wide, and it is held on the stack, which the driver backs for every thread the card can hold. One direction per formation of it keeps that stack small: the three-pool adjoint's from 1.9 KB to 0.9 KB, under the default the driver already holds, and the dual three-pool's from 7.0 KB to 3.6 KB. Both are also faster. Co-Authored-By: Claude Opus 5.5 --- src/blochsim/_layout_complex_vjp.hpp | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/src/blochsim/_layout_complex_vjp.hpp b/src/blochsim/_layout_complex_vjp.hpp index b2844adc..ea5e87e9 100644 --- a/src/blochsim/_layout_complex_vjp.hpp +++ b/src/blochsim/_layout_complex_vjp.hpp @@ -310,7 +310,12 @@ struct Gradients { template __device__ __noinline__ Gradients contract_interval(Geometry g, int relax_code, T dt, Inputs in, float order, int row, int atom, Met met) { - constexpr int KD = DIRECTIONS, CHUNK = 4; + // Directions taken per formation of the operator. The multi-valued + // operator is held on the stack, which the driver backs for every thread + // a card can hold: where its entries are wide -- three pools, or a dual + // -- one direction at a time keeps that stack small, and is no slower. + constexpr int KD = DIRECTIONS; + constexpr int CHUNK = num::is_dual::value || POOLS == 3 ? 1 : 4; Gradients out; #pragma unroll for (int k = 0; k < KD; ++k) out.g[k] = 0.0f; From 39ad97522ac97250c51dc6d58d2d60023752c05f Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 06:52:12 +0200 Subject: [PATCH 15/16] Take the package's own structural derivatives in reverse mode PyTorch's forward mode registers decompositions through TorchScript the first time it makes a dual tensor in a process. Resolving a packing and building a transition table each read a directional derivative of the package's own structure, which is now taken by a vector-Jacobian product of the vector-Jacobian product, so a simulation and its gradients compile nothing at run time. Both directions a packing is read along share one walk of the stream. Co-Authored-By: Claude Opus 5.5 --- src/blochsim/_derivative.py | 83 ++++++++++++++++++++++++++++ src/blochsim/model/_binding.py | 3 +- src/blochsim/sequence/_transition.py | 5 +- tests/optim/test_binding.py | 12 ++-- 4 files changed, 95 insertions(+), 8 deletions(-) create mode 100644 src/blochsim/_derivative.py diff --git a/src/blochsim/_derivative.py b/src/blochsim/_derivative.py new file mode 100644 index 00000000..9553b38b --- /dev/null +++ b/src/blochsim/_derivative.py @@ -0,0 +1,83 @@ +"""Directional derivatives taken in reverse mode. + +PyTorch's forward mode loads, the first time it makes a dual tensor in a +process, decompositions it registers through TorchScript. What the package +works out about its own structure -- how a packing moves with its arguments, +how a transition table moves with its flip -- is therefore taken in reverse +mode, so that a simulation and its gradients compile nothing at run time. +""" + +from __future__ import annotations + +__all__ = ["directional_derivatives"] + +from collections.abc import Callable, Sequence +from typing import Any + +import torch + + +def directional_derivatives( + function: Callable[..., Any], + primals: Sequence[torch.Tensor], + directions: Sequence[Sequence[torch.Tensor]], +) -> tuple[Any, tuple[Any, ...]]: + """``function`` at ``primals``, and its derivative along each direction. + + What :func:`torch.func.jvp` returns for each direction, by reverse passes: + the vector-Jacobian product is linear in its cotangent, so its own + vector-Jacobian product, taken with a direction as the cotangent, is the + Jacobian times that direction. One forward pass serves every direction, and + the passes compose with an enclosing :mod:`torch.func` transform. + + Parameters + ---------- + function: + Takes the primals positionally and returns a tensor or a tuple of + tensors, real or complex. + primals: + Floating-point tensors. + directions: + Each a tangent per primal, of the primal's shape. + + Returns + ------- + tuple + The outputs, as ``function`` returns them, and per direction their + derivatives in the same form. + """ + shape: dict[str, Any] = {} + + def real(*given: torch.Tensor) -> tuple[torch.Tensor, ...]: + outputs = function(*given) + shape["single"] = isinstance(outputs, torch.Tensor) + held = (outputs,) if shape["single"] else tuple(outputs) + shape["complex"] = tuple(out.is_complex() for out in held) + return tuple( + torch.view_as_real(out) if out.is_complex() else out for out in held + ) + + values, pull = torch.func.vjp(real, *primals) + _, push = torch.func.vjp( + lambda *cotangents: pull(cotangents), *(torch.zeros_like(v) for v in values) + ) + + def shaped(parts: Sequence[torch.Tensor]) -> Any: + out = tuple( + torch.view_as_complex(part.contiguous()) if is_complex else part + for part, is_complex in zip(parts, shape["complex"], strict=True) + ) + return out[0] if shape["single"] else out + + slopes = tuple( + shaped( + push( + tuple( + tangent.to(primal.dtype) + for tangent, primal in zip(direction, primals, strict=True) + ) + ) + ) + for direction in directions + ) + return shaped(values), slopes diff --git a/src/blochsim/model/_binding.py b/src/blochsim/model/_binding.py index 4925481e..d9db5eb9 100644 --- a/src/blochsim/model/_binding.py +++ b/src/blochsim/model/_binding.py @@ -37,6 +37,7 @@ import torch +from .._derivative import directional_derivatives from ..sequence._accelerators import ( _PackedEvents, pack_description, @@ -231,7 +232,7 @@ def floats(*given: torch.Tensor) -> tuple[torch.Tensor, ...]: seeds = _seeds(primals, position) if seeds is None: continue - along = tuple(torch.func.jvp(floats, primals, seed)[1] for seed in seeds) + _, along = directional_derivatives(floats, primals, seeds) for buffer, scale, ramp in zip(_VALUES, *along, strict=True): term = _term(scale, ramp, primals[position].numel()) if term is None: diff --git a/src/blochsim/sequence/_transition.py b/src/blochsim/sequence/_transition.py index 97cbae92..31422166 100644 --- a/src/blochsim/sequence/_transition.py +++ b/src/blochsim/sequence/_transition.py @@ -46,6 +46,7 @@ import torch +from .._derivative import directional_derivatives from ._description import RfDefinition @@ -616,7 +617,9 @@ def integrate(flip: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: # rather than shaping it. return a * torch.exp(0.5j * carried)[:, None].expand(wide), b - values, slopes = torch.func.jvp(integrate, (theta,), (torch.ones_like(theta),)) + values, (slopes,) = directional_derivatives( + integrate, (theta,), ((torch.ones_like(theta),),) + ) return TransitionTable( a=values[0].to(torch.complex64).contiguous(), b=values[1].to(torch.complex64).contiguous(), diff --git a/tests/optim/test_binding.py b/tests/optim/test_binding.py index ade5a224..8dd14360 100644 --- a/tests/optim/test_binding.py +++ b/tests/optim/test_binding.py @@ -425,13 +425,13 @@ def test_a_schedule_that_moves_its_timing_answers_what_the_plain_one_does() -> N ), ], ) -def test_resolving_walks_the_stream_once_per_direction(design, packings) -> None: +def test_resolving_walks_the_stream_once_per_argument(design, packings) -> None: """Which is the whole cost of resolving, and what it is spent on. - One walk for the structure, two per argument for the derivatives the map is - read from, and one to check the map against a packing it never saw. A walk - is per-event Python under a forward-mode interpreter, so an extra one is - not a rounding error in what resolving costs. + One walk for the structure, one per argument for the derivatives the map is + read from -- both directions it is read along share it -- and one to check + the map against a packing it never saw. A walk is per-event Python, so an + extra one is not a rounding error in what resolving costs. """ simulator = MRFSimulator(TR=10.0, TI=20.0, states=10) played = simulator.played(**design) @@ -442,7 +442,7 @@ def test_resolving_walks_the_stream_once_per_direction(design, packings) -> None ) assert packing is not None - assert len(packings) == 2 + 2 * len(packing.varying) + assert len(packings) == 2 + len(packing.varying) def test_the_sample_times_are_the_clock_the_intervals_make() -> None: From d3d4a7c3bdfaa935ea277600e6dc2f1d1f9f8f9b Mon Sep 17 00:00:00 2001 From: mcencini Date: Thu, 8 Oct 2026 06:53:41 +0200 Subject: [PATCH 16/16] Describe the kernels in their own terms Comments and the agent notes said what the kernels do by reference to the Triton source they replaced; they now say what the code does. Co-Authored-By: Claude Opus 5.5 --- CLAUDE.md | 5 ++--- benchmarks/validate.py | 4 ++-- src/blochsim/_launch.hpp | 4 ++-- src/blochsim/_layout_complex.hpp | 2 +- src/blochsim/_layout_numbers.hpp | 6 +++--- src/blochsim/_tile.hpp | 6 +++--- src/blochsim/sequence/_epg_gpu.py | 12 +++++++++--- tests/sequence/test_both_pools.py | 2 +- 8 files changed, 23 insertions(+), 18 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 1583c799..13ee5c72 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -96,9 +96,8 @@ turns the layouts off and `layout_launches()` counts them; exist only on the card: `_gpu_host` compiles the tile kernels, so the host lane holds those, not these, to the C++ kernels. -**The EPG kernels index in 32 bits** (`bsk::index_t`), as Triton did for every -integer argument that fit. An offset that can pass 2^31 is cast to 64 bits -where it is formed, as the Triton source cast it, and the launcher refuses an +**The EPG kernels index in 32 bits** (`bsk::index_t`). An offset that can +pass 2^31 is cast to 64 bits where it is formed, and the launcher refuses an integer argument that does not fit rather than truncating it. **`--cov` is on by default** through `addopts`, so a bare `pytest` writes diff --git a/benchmarks/validate.py b/benchmarks/validate.py index ba49e81c..b0aa6256 100644 --- a/benchmarks/validate.py +++ b/benchmarks/validate.py @@ -1,6 +1,6 @@ """Whether BlochSim, sycomore and epgpy compute the same fingerprint. -Three independent implementations of the same model -- a fused C++ or Triton +Three independent implementations of the same model -- a fused C++ or CUDA state machine over a packed event stream, a C++ EPG library driven one tissue at a time from Python, and an operator-per-event NumPy library -- should agree to the precision the coarsest of them carries. This says by how much they do, @@ -166,7 +166,7 @@ def main() -> None: if torch.cuda.is_available(): # The two kernels are separate implementations of the same recursion, - # so this is a comparison of the C++ one against the Triton one rather + # so this is a comparison of the C++ one against the CUDA one rather # than a check that a tensor made the trip. on_cpu = blochsim_signal(T1, T2, flip, TR, arguments.states) on_card = blochsim_signal(T1, T2, flip, TR, arguments.states, device="cuda") diff --git a/src/blochsim/_launch.hpp b/src/blochsim/_launch.hpp index 37a60496..26c4e277 100644 --- a/src/blochsim/_launch.hpp +++ b/src/blochsim/_launch.hpp @@ -93,8 +93,8 @@ inline bool read_launch(PyObject* name, PyObject* grid, PyObject* args, Launch& break; default: arg.i = PyLong_AsLongLong(item); - // The EPG kernels index in 32 bits, as Triton did for an - // argument that fit; one that does not is refused, not cut. + // The EPG kernels index in 32 bits; an argument that does not + // fit is refused, not cut. if (!PyErr_Occurred() && (arg.i > 2147483647LL || arg.i < -2147483648LL)) { PyErr_Format(PyExc_ValueError, "%s: argument %zd is %lld, past what a kernel indexes with", diff --git a/src/blochsim/_layout_complex.hpp b/src/blochsim/_layout_complex.hpp index 333036b5..8c3e61e1 100644 --- a/src/blochsim/_layout_complex.hpp +++ b/src/blochsim/_layout_complex.hpp @@ -91,7 +91,7 @@ __device__ __forceinline__ float stored(T x) { } } -// One pool through a hard pulse, in Triton's fused order. +// One pool through a hard pulse, its phase folded into the rotation's entries. template __device__ __forceinline__ void rotate_flip_phase(T cosine, T sine, T cos_phi, T sin_phi, T cos_2phi, T sin_2phi, T& fp_r, T& fp_i, T& fm_r, T& fm_i, diff --git a/src/blochsim/_layout_numbers.hpp b/src/blochsim/_layout_numbers.hpp index d49e3263..a44c7d2b 100644 --- a/src/blochsim/_layout_numbers.hpp +++ b/src/blochsim/_layout_numbers.hpp @@ -152,7 +152,7 @@ __device__ __forceinline__ DualT exp_(DualT x) { return {e, e * x.d}; } -// Division the way Triton's float32 ``/`` divides: approximate. +// Float division, approximate: the fast reciprocal, two ulps at most. __device__ __forceinline__ float div_(float a, float b) { return __fdividef(a, b); } __device__ __forceinline__ double div_(double a, double b) { return a / b; } template @@ -212,8 +212,8 @@ __device__ __forceinline__ Dual64 cos_(Dual64 x) { return {c, -s * x.d}; } -// Triton's _sincos: one Cody-Waite reduction by a quarter turn, the Cephes -// polynomials either side of zero. +// One Cody-Waite reduction by a quarter turn, then the Cephes polynomials +// either side of zero. __device__ __forceinline__ void sincos_cw(float x, float& s, float& c) { const float quarter = rintf(x * 0.6366197723675814f); float r = fmaf(-quarter, 1.5703125f, x); diff --git a/src/blochsim/_tile.hpp b/src/blochsim/_tile.hpp index cf2f5079..b07470ad 100644 --- a/src/blochsim/_tile.hpp +++ b/src/blochsim/_tile.hpp @@ -39,9 +39,9 @@ namespace bsk { -// The integers the EPG kernels index with: 32 bits, as Triton passed every -// integer argument that fit. An offset that can pass 2^31 is cast to 64 bits -// where it is formed, and the launcher refuses an argument that does not fit. +// The integers the EPG kernels index with: 32 bits. An offset that can pass +// 2^31 is cast to 64 bits where it is formed, and the launcher refuses an +// argument that does not fit. using index_t = std::int32_t; // --------------------------------------------------------------------------- diff --git a/src/blochsim/sequence/_epg_gpu.py b/src/blochsim/sequence/_epg_gpu.py index b618d465..eb84c3a1 100644 --- a/src/blochsim/sequence/_epg_gpu.py +++ b/src/blochsim/sequence/_epg_gpu.py @@ -473,7 +473,9 @@ def simulate_into( pools = _pool_flag(lineshape, exchanging) block_states = next_power_of_2(state_count) total = train_count * atom_count - problems = _problems_per_program(block_states, _epg_real_kernel if real_axis == 1 else _epg_kernel) + problems = _problems_per_program( + block_states, _epg_real_kernel if real_axis == 1 else _epg_kernel + ) grid = (cdiv(total, problems),) # A kernel argument has to be a tensor even where the branch reading it is # compiled out, so an unprofiled launch passes one it already has. @@ -724,7 +726,9 @@ def simulate_jvp_into( shims = _shim_count(tissue) block_states = next_power_of_2(state_count) total = train_count * atom_count - problems = _problems_per_program(block_states, _epg_real_jvp_kernel if real_axis == 1 else _epg_jvp_kernel) + problems = _problems_per_program( + block_states, _epg_real_jvp_kernel if real_axis == 1 else _epg_jvp_kernel + ) grid = (cdiv(total, problems),) if real_axis == 1: @@ -1688,7 +1692,9 @@ def simulate_vjp_jvp_into( wave * row_count * 36, dtype=torch.float32, device=t1.device ) - problems = _problems_per_program(block_states, _epg_real_vjp_jvp_kernel if real else _epg_vjp_jvp_kernel) + problems = _problems_per_program( + block_states, _epg_real_vjp_jvp_kernel if real else _epg_vjp_jvp_kernel + ) for base in range(0, total, wave): span = min(wave, total - base) if pool_bars is not None: diff --git a/tests/sequence/test_both_pools.py b/tests/sequence/test_both_pools.py index 55a33fbe..56be77f4 100644 --- a/tests/sequence/test_both_pools.py +++ b/tests/sequence/test_both_pools.py @@ -1679,5 +1679,5 @@ def leaves(value): ) # The roots branch moves by about 1e-5 of the largest value between # correct compilations of itself -- one fused multiply-add more or - # less -- and Triton's differs from the CUDA build's by 2e-5. + # less -- and the card's build differs from the CPU's by 2e-5. assert worst / largest < 5e-5, (name, worst / largest)