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..fe22f87e 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,15 @@ 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. +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 +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..92d932e8 100644 --- a/.github/workflows/wheels.yml +++ b/.github/workflows/wheels.yml @@ -12,6 +12,9 @@ on: - setup.py - "src/blochsim/*.cpp" - "src/blochsim/*.hpp" + - "src/blochsim/*.cu" + - "src/blochsim/*.cu.in" + - "src/cuda/**" - scripts/check_wheel.py - .github/workflows/wheels.yml workflow_dispatch: @@ -88,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: @@ -107,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/.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..27db2a3c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,18 @@ - **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`, 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. ### Added diff --git a/CLAUDE.md b/CLAUDE.md index 6b6972f0..13ee5c72 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,47 @@ 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. +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. + +**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. + +**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 +operator's derivatives -- taken along every tissue input at once by +`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`). 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 `coverage.xml`. It is ignored, not tracked. @@ -129,7 +137,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,13 +215,30 @@ 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 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 5d80ec2c..d9197c45 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. @@ -54,5 +54,140 @@ 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" AND NOT BLOCHSIM_CUDA_PACKAGE) + 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) +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(BLOCHSIM_CUDA_PACKAGE OR (CMAKE_CUDA_COMPILER AND CMAKE_SYSTEM_NAME STREQUAL "Linux")) + 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) + # 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 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) + + # 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() + + # 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 + src/blochsim/_layout_real_vjp.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() + 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() + + 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 + ${_blochsim_gpu_sources} + ) + target_include_directories(_gpu PRIVATE "${CMAKE_CURRENT_SOURCE_DIR}/src/blochsim") + # 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. 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/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/docs/developer_guide.md b/docs/developer_guide.md index ae089677..7f7b76c0 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.** 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. 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..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. @@ -125,7 +135,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,6 +153,9 @@ 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, 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 @@ -170,10 +183,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 +234,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..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,36 +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 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 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(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 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"}] + """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") - if size > LARGEST_REASONABLE_BYTES: - 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/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/_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/_epg_kernels.hpp b/src/blochsim/_epg_kernels.hpp new file mode 100644 index 00000000..e3f434d9 --- /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 (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)); + 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 (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)); + 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, 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{}; + 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 (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); + 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 (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); + 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 (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); + 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 (bsk::index_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, 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{}; + 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 (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)) { + 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, 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); + 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, 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); + 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, 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{}; + 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 (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); + 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 (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); + 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, 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{}; + 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 (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); + 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, 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{}; + 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 (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); + 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 (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); + 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, 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{}; + 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 (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. + 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, 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{}; + 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 (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); + 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 (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); + 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 (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))); + 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, 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{}; + 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 (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 + // 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..016359cd --- /dev/null +++ b/src/blochsim/_gpu.cu @@ -0,0 +1,160 @@ +// 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 +// 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 + +#include "_launch.hpp" +#include "_layout.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 + +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 + +// 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; + +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; + } + // 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; + } + 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"); + } + if (laying_out) { + const int laid = blochsim_layout::launch(request.kernel, request.arguments, request.grid[0], + 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 + // 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; + const void* function = KERNEL_FUNCTIONS[request.kernel][bounded]; + status = cudaLaunchKernel(function, 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; +} + +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); +} + +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."}, + {"launch", launch, METH_VARARGS, + "Queue a kernel over a grid of programs on a device's stream."}, + {"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."}, + {"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}, +}; + +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..d70864fc --- /dev/null +++ b/src/blochsim/_gpu_kernel.cu.in @@ -0,0 +1,23 @@ +// 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 "_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) { + 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..31ffeb3c --- /dev/null +++ b/src/blochsim/_gpu_launch.py @@ -0,0 +1,218 @@ +"""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] = [] + +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 + + +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)) + + +# 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 + + +@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]: + names, kinds, _ = _entry(name) + return names, kinds + + +@cache +def available() -> bool: + """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 + + +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 + + +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.""" + if not available(): + yield + return + module = _module("cuda") + layouts = module.use_layouts(False) + try: + yield + finally: + module.use_layouts(layouts) + + +class Kernel: + """A compiled kernel, launched as ``kernel[grid](*arguments)``.""" + + 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) + + 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..74e79831 --- /dev/null +++ b/src/blochsim/_kernels.hpp @@ -0,0 +1,847 @@ +// 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 "_lanes.hpp" +#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; + // Rows of y each thread holds on a card. + int lanes; +}; + +#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, + a[13].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, + a[12].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, 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,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) \ + 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/_lanes.hpp b/src/blochsim/_lanes.hpp new file mode 100644 index 00000000..437a0832 --- /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 4 +#define BLOCHSIM_LANES__epg_real_jvp_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 +#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 new file mode 100644 index 00000000..26c4e277 --- /dev/null +++ b/src/blochsim/_launch.hpp @@ -0,0 +1,128 @@ +// 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("(ssi)", info.params, info.kinds, info.lanes); + 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; + // The rows each thread holds on a card. + int lanes = 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); + // 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", + info.name, i, static_cast(arg.i)); + return false; + } + 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); + 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] / launch.lanes); + 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/_layout.cu b/src/blochsim/_layout.cu new file mode 100644 index 00000000..9f918e4a --- /dev/null +++ b/src/blochsim/_layout.cu @@ -0,0 +1,496 @@ +// 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; } +}; + +// 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"); + 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.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"); + 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"); + 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; + } +} + +// 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; +} + +// 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"); +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"); + +} // namespace + +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; + } + // 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); + } + 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); +} + +} // namespace blochsim_layout diff --git a/src/blochsim/_layout.hpp b/src/blochsim/_layout.hpp new file mode 100644 index 00000000..1f660041 --- /dev/null +++ b/src/blochsim/_layout.hpp @@ -0,0 +1,73 @@ +// 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_complex_vjp.hpp" +#include "_layout_pooled.hpp" +#include "_layout_real.hpp" +#include "_layout_real_vjp.hpp" + +namespace blochsim_layout { + +// 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. +#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); + +// 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); + +// 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; + +} // 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..8c3e61e1 --- /dev/null +++ b/src/blochsim/_layout_complex.hpp @@ -0,0 +1,1258 @@ +// 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::primal; +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 if (p.phase_cos != nullptr) { + c = p.phase_cos[at]; + s = p.phase_sin[at]; + } else { + sincos_(__ldg(p.phase + at), s, c); + } +} + +// 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, 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, + 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 = 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) + : 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_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); + 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 (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; + 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(primal(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(primal(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; +}; +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: +// 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(primal(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 * 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, + 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) / (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; + 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(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; + 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]; +}; + +// 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, + 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; + 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 int at = p.atom_stride ? atom[y] : 0; + 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 { + in.r2_b = in.shift_b = 0.0f; + } + if constexpr (POOLS == 3) { + in.fraction_c = tissue.fraction_c[y]; + in.exchange_c = tissue.exchange_c[y]; + in.r1_c = tissue.r1_c[y]; + } else { + 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) { + // 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 (!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(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; +#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]; + } + 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; + } + 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_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..ea5e87e9 --- /dev/null +++ b/src/blochsim/_layout_complex_vjp.hpp @@ -0,0 +1,1298 @@ +// 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) { + // 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; +#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 new file mode 100644 index 00000000..a44c7d2b --- /dev/null +++ b/src/blochsim/_layout_numbers.hpp @@ -0,0 +1,512 @@ +// 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 + +#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; } +__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; +} + +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/(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; +} +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}; +} + +// 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 +__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); } +// |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; +} + +__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__ 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); + return {c, -s * x.d}; +} + +// 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); + 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; 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), tangents != nullptr ? __ldg(tangents + at) : 0.0f}; + } 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)}; +} + + +// 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_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/_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/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 diff --git a/src/blochsim/_perk_kernels.hpp b/src/blochsim/_perk_kernels.hpp new file mode 100644 index 00000000..d5806f96 --- /dev/null +++ b/src/blochsim/_perk_kernels.hpp @@ -0,0 +1,393 @@ +// The PERK feature map and its regression, fused. +// +// y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T +// +// 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. +// +// 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 +// 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 + +struct alignas(16) Four { + float x, y, z, w; +}; + +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; +} + +// 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. +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; + } +} + +// 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 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 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]; + 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 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) { + 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; + } + } + 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; + } + } + } + } + } + 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) { + 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); + } + } + } + } + } +} + +// 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 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]; + 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 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) { + 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; + } + } + 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; + } + } + 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; + } + } + } + } + } + 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) { + _emit(&output[voxel * contrasts + g0 + c], gradient[t][r][c], splits); + } + } + } + } + } +} 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/_special.hpp b/src/blochsim/_special.hpp new file mode 100644 index 00000000..850c9b02 --- /dev/null +++ b/src/blochsim/_special.hpp @@ -0,0 +1,54 @@ +// 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" + +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; +} + +} // namespace bsk diff --git a/src/blochsim/_tile.hpp b/src/blochsim/_tile.hpp new file mode 100644 index 00000000..b07470ad --- /dev/null +++ b/src/blochsim/_tile.hpp @@ -0,0 +1,1298 @@ +// 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 +// 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 + +#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 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; + +// --------------------------------------------------------------------------- +// The launch a program belongs to. +// --------------------------------------------------------------------------- + +#if defined(BLOCHSIM_SIMT) + +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) * Y_LANES; } +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 index_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 index_t program_id(int axis) { return static_cast(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"); + // 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 i = 0; i < lanes; ++i) v[i] = static_cast(scalar); + } + + template = 0> + BSK_HD V(const V& other) { + 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 in its row ``y`` at ``z``. +template +BSK_HD decltype(auto) element(const A& a, int y = 0, int z = 0) { + if constexpr (axes_of == 0) { + return a; + } else { + using Tile = std::decay_t; + return (a.v[((axes_of & 2) ? y : 0) * Tile::depth + ((axes_of & 4) ? z : 0)]); + } +} + +// 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); + if constexpr (AX == 0) { + return f(a...); + } else { + using R = decltype(f(element(a)...)); + V out; + 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 { +#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; + } +} + +#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); + } +} + +// 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(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; +#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; + 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; + } + 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 { +#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 y, int z) { + if (element(mask, y, z)) { + auto p = element(pointer, y, z); + *p = static_cast>(element(value, y, 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 y, int z) { + if (element(mask, y, z)) { + auto p = element(pointer, y, z); + atomicAdd(p, static_cast>(element(value, y, 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. +// +// 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(); + 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 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(); + 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) { + 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; + __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 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; + __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); } +}; + +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) + // 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); + 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) + // 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); + 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) + const int nz = width_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; + } +#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) + 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; + 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) + // 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); + 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) + // 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(); + 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..d4ee06a6 --- /dev/null +++ b/src/blochsim/estimators/_perk_gpu.py @@ -0,0 +1,144 @@ +"""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 +from functools import cache + +import torch + +from .._gpu_launch import Kernel, cdiv + +#: 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") + + +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() + + +@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, + 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] + 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), splits)]( + signals, + _ready(frequency), + _ready(phase), + _ready(feature_mean), + _ready(weight), + _ready(parameter_mean), + output, + voxels, + contrasts, + features, + parameters, + math.sqrt(2.0 / features), + splits, + _THREADS, + ) + 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] + 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), splits)]( + signals, + _ready(frequency), + _ready(phase), + _ready(weight), + _ready(cotangent), + output, + voxels, + contrasts, + features, + parameters, + math.sqrt(2.0 / features), + splits, + _THREADS, + ) + 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/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/_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..eb84c3a1 --- /dev/null +++ b/src/blochsim/sequence/_epg_gpu.py @@ -0,0 +1,1903 @@ +"""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 + + +# 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: + """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, 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. 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 + 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. + """ + rows = max(1, _PROGRAM_THREADS // block_states) + return (1 << (rows.bit_length() - 1)) * kernel.lanes + + +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, _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. + 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, _epg_real_jvp_kernel if real_axis == 1 else _epg_jvp_kernel + ) + 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, _epg_vjp_kernel) + 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, _epg_real_vjp_kernel) + 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, _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 + # 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, _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),)]( + 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, _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: + # 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..5668e9b2 --- /dev/null +++ b/src/blochsim/sequence/_pools_gpu.py @@ -0,0 +1,542 @@ +"""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, pooled_layout_floats +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 + + # 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))) + 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() + 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 + 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, + 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/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/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. +""" diff --git a/tests/estimators/test_perk_kernel.py b/tests/estimators/test_perk_kernel.py index 666a8397..1f160011 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,54 @@ 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) + + +@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, and more voxels than one program holds, + so every edge of the tiling is read. + """ + gpu = pytest.importorskip("blochsim.estimators._perk_gpu") + 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) + 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()) + + def on(*tensors: torch.Tensor) -> list[torch.Tensor]: + return [tensor.float().to(device) for tensor in tensors] + + estimated = gpu.regress( + *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 + ) + torch.testing.assert_close( + gradient.double(), expected_gradient, atol=1e-5, rtol=1e-5 + ) 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: diff --git a/tests/sequence/test_both_pools.py b/tests/sequence/test_both_pools.py index 80a9a77c..56be77f4 100644 --- a/tests/sequence/test_both_pools.py +++ b/tests/sequence/test_both_pools.py @@ -1462,54 +1462,24 @@ 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" + 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 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 +1488,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 +1581,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 +1624,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): @@ -1693,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 the card's build differs from the CPU's by 2e-5. + assert worst / largest < 5e-5, (name, worst / largest) diff --git a/tests/sequence/test_cuda_parity.py b/tests/sequence/test_cuda_parity.py index f2be8b6e..2a680a6a 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. """ @@ -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_triton 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_triton 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]) @@ -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..898ba922 --- /dev/null +++ b/tests/sequence/test_host_kernels.py @@ -0,0 +1,68 @@ +"""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 + +pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + +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_layout_kernels.py b/tests/sequence/test_layout_kernels.py new file mode 100644 index 00000000..ea9b3fd3 --- /dev/null +++ b/tests/sequence/test_layout_kernels.py @@ -0,0 +1,105 @@ +"""Whether a kernel written for its layout computes what the tile kernel does. + +A layout 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.available(), + reason="needs a card and the kernels compiled for it", +) + + +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_layout_computes_what_the_tile_kernel_does(case) -> None: + with _gpu_launch.generic_kernels(): + tiled = CASES[case]() + before = _gpu_launch.layout_launches() + laid_out = CASES[case]() + + 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_tile_kernels_run_where_asked() -> None: + before = _gpu_launch.layout_launches() + with _gpu_launch.generic_kernels(): + _forward(torch.pi / 2) + + assert _gpu_launch.layout_launches() == before diff --git a/tests/sequence/test_many_pools_host.py b/tests/sequence/test_many_pools_host.py new file mode 100644 index 00000000..939bda43 --- /dev/null +++ b/tests/sequence/test_many_pools_host.py @@ -0,0 +1,44 @@ +"""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 + +pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + +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")