Skip to content

Run the GPU kernels as AoT CUDA written for their layouts, and ship them as blochsim[cu12]/[cu13] - #33

Merged
mcencini merged 16 commits into
mainfrom
layout-kernels
Oct 8, 2026
Merged

mcencini merged 16 commits into
mainfrom
layout-kernels

Conversation

@mcencini

@mcencini mcencini commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

The GPU kernels run ahead-of-time CUDA, at or under Triton's time on every case measured, and ship to PyPI as blochsim[cu12] / blochsim[cu13]. This supersedes #31 (AoT CUDA in place of Triton) and #32 (tiled PERK), whose commits it contains.

What changed

Kernels written for their layouts (_layout*.hpp). A layout is what is compiled: pools, how a pulse is formed, per-voxel maps, problems per thread, 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: at float they are the forward simulation and the adjoint, at num::Dual the JVP and the second-order adjoint.

  • EPG kernels. Complex and real forward, JVP, VJP and VJP-JVP, for 0–3 pools.
  • Many-pool kernels. Forward and adjoint for 2–8 pools.
  • Checkpointing. The adjoints keep every fourth state and replay each stretch, so they need no recording launch and no full trajectory.
  • Tile kernels remain for what a layout does not take (rows wider than a warp) and for the host build that the suite checks them on.

Other changes.

  • The per-combination specialisations from Compile the GPU kernels ahead of time as CUDA and drop Triton #31 are removed; the layouts cover them.
  • PERK splits its features across a second grid axis when a launch is too small to fill the card.
  • The package's own structural derivatives (packing binding, transition tables) are taken in reverse mode. A simulation and its gradients therefore no longer make PyTorch register its forward-mode decompositions through TorchScript.

Packaging

The base wheel is CPU-only. The card's module is a package per CUDA major version (src/cuda/12, src/cuda/13), the same CMake project with BLOCHSIM_CUDA_PACKAGE set. Each one:

  • links libcudart from torch's nvidia wheel through an rpath, rather than carrying a runtime;
  • carries machine code for 7.5, 8.0 and 9.0 plus PTX for 9.0, in compressed fatbinaries;
  • pins the blochsim it was built with.

_gpu_launch loads the build for torch's CUDA major version. It refuses one for another major or release with CudaBuildMismatch, which names the extra to install.

pip install torch --index-url https://download.pytorch.org/whl/cu126
pip install 'blochsim[cu12]'        # or blochsim[cu13] beside a CUDA 13 torch
Wheel (75;80;90 + PTX 90), as CI builds it Size
blochsim-cuda12, CUDA 12.6 61.7 MB
blochsim-cuda13, CUDA 13.0 48.8 MB

wheels.yml builds both in a manylinux_2_28 container with the oldest toolkit of each major that torch ships. It installs each beside torch cu126/cu130 on Python 3.10 and 3.14 and checks that libcudart resolves into torch's nvidia wheel. It publishes each project from its own environment (pypi, pypi-cuda12, pypi-cuda13).

Speed against Triton

Device time per case, the median of three interleaved runs on an RTX 4060 Laptop (8.9), sm_89 build.

Case n Triton µs CUDA µs Ratio
epg_real 2000 863 932 1.08*
epg_real 100000 13598 9427 0.69
epg_complex 2000 1231 1192 0.97
epg_complex 100000 24645 20083 0.81
epg_real_grad 2000 4105 2598 0.63
epg_real_grad 100000 115635 69236 0.60
epg_complex_grad 20000 61689 42196 0.68
epg_real_hvp 2000 13975 8005 0.57
epg_real_hvp 20000 94107 52168 0.55
epg_complex_hvp 5000 60277 46209 0.77
pooled (3 pools) 2000 2004 1352 0.67
pooled (3 pools) 50000 21183 8279 0.39
pooled_grad 20000 75813 44555 0.59
pooled_hvp 5000 150924 66449 0.44
pooledb_hvp (2 pools) 20000 356946 98822 0.28
pooled5 500 21786 11056 0.51
pooled5 5000 134057 75933 0.57
pooled5_grad 2000 171391 68360 0.40
perk 2000 245 93 0.38
perk 200000 15232 7724 0.51
perk_grad 2000 361 278 0.77
perk_grad 200000 21451 14882 0.69

* The kernel itself is 0.71× (99 µs against 140). The total at this size is dominated by torch's own host-to-device copies, which take 78 µs or 114 µs per process on either build and in either order.

Per kernel, against Triton's on the same captured launches:

Kernels Ratio
Complex forward 0.16–0.98
Real forward 0.53–0.85
JVPs 0.5–1.0
Real VJP 0.52–0.86
Complex VJP 0.05–0.83
VJP-JVP 0.11–0.78
Many-pool forward 0.04–0.22
Many-pool adjoint 0.13–0.30
Three-pool table kernels 0.10–0.16

The first call takes 0.01–1.1 s, against 0.6–2.5 s for Triton.

Memory against Triton

  • Peak torch allocation: identical in every case except pooled5_grad, which drops from 265 to 81 MiB. The many-pool adjoint keeps checkpoints, not Triton's recording.
  • Device memory beyond the allocator: equal or lower in every case but one. The exception is pooled_hvp (three pools, second order), which leaves 84 MiB against Triton's 52.
    • That is the driver's local-memory reservation for the three-pool dual adjoint's 3.6 KB stack. Down from 7.0 KB, since this PR contracts wide intervals one direction at a time, which is also faster.
    • Getting under Triton's 2.5 KB needs that adjoint's contraction redesigned. It is left as a follow-up.

Verified

  • Full suite: 1614 passed, 3 skipped on the card.
  • cu13 wheel: installed beside torch 2.13+cu130, the full suite passes with CUDA_DISABLE_PTX_JIT=1, so an 8.9 card runs the wheel's sm_80 code. scripts/check_wheel.py loads the module bare and finds libcudart.so.13 in torch's nvidia wheel through the rpath.
  • cu12 wheel: loads bare against a CUDA 12.6 runtime placed where its nvidia wheel would put it. Not run against a CUDA 12 torch locally; the CI job does that.
  • Layout against tile paths: they agree to float32 round-off through autograd: values, first and second derivatives, CPU kernels as reference.
  • Not run: on a T4, A40 or H100.

Publishing

blochsim-cuda12 and blochsim-cuda13 have trusted publishers for wheels.yml with environments pypi-cuda12 and pypi-cuda13, so a v*.*.* tag publishes all three projects.

benchmarks/README.md still reports measurements taken on the Triton build and should be re-measured.

🤖 Generated with Claude Code

claude and others added 16 commits October 6, 2026 17:37
The EPG, pooled and PERK kernels are C++ over a small tile runtime
(_tile.hpp), compiled by nvcc into blochsim._gpu, one file per kernel with
256- and 1024-thread bounded variants, linking the CUDA runtime statically.
CMake builds it wherever it finds nvcc (BLOCHSIM_CUDA overrides), and the
x86-64 manylinux wheel carries it. The same source compiled for the host
(blochsim._gpu_host) replaces Triton's interpreter in the suite and is held
to the C++ kernels. Triton, its modules, the interpreted marker and the
cache-pruning script are gone.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ND2A7uRuhWiay6B4rWF1H
The ported PERK kernels formed each voxel's features one thread at a
time, reading every frequency from global memory once per multiply. They
ran 2.5x slower than the Triton kernels they replaced. Each program now
stages a block of signals and frequencies in shared memory, and each
thread forms two voxels by a block of features in registers, reading
weights four at a time. The accumulators are stored through unrolled
loops, so they stay in registers and nothing spills.

On an RTX 4060 (sm_89), 200k voxels, 24 contrasts, 256 features:
forward 2.4 ms against 5.1 ms for Triton and 12.9 ms before; adjoint
4.6 ms against 5.4 ms and 14.0 ms. Below about 4k voxels a program's
serial chain over the features is longer than Triton's, by up to 45 us.

On the host the same source runs a program's threads one after another
between barriers, so the host build still checks the card's code. The
parity test now runs on the card too.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
On a card a tile held one element per thread. A value varying along y now
holds Y_LANES rows per thread in registers, set per kernel in _lanes.hpp;
the launcher divides a program's rows by it to size the block, and Python
sizes problems per program from the kernel's lanes. One-row values compile
to the same straight-line code as before.

The real forward kernel holds 8 lanes, takes the cosine and sine of a
uniform flip once, and indexes in 32 bits. Measured against Triton on one
captured launch (100k atoms, 32 states), it is 1.48x slower at 4 lanes,
against 2.25x for the port: the lanes help, the generated code caps them.

Not for merging as is: the lane counts other than the real forward kernel
are unswept, and the 32-bit indices need a bound on problems x outputs.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Every EPG kernel takes its feature switches as arguments, so as compiled
once it carries the registers and code of every term it might evaluate.
Each entry of _specializations.json compiles the same kernel with those
switches as constants: every function is inlined, a constant switch
folds, and the terms it turns off are never generated. The launcher runs
the entry whose switches a launch matches exactly and the general kernel
where none does, so the list decides speed and never correctness.

The 124 entries are the combinations the test suite and the benchmark
cases launch (scripts/kernel_census.py), less the second-order kernel's
combinations with slice profile, motion or pools, whose compile took as
long as a dozen others together. The real forward kernel holds 4
problems per thread and the complex forward kernel 2. The kernel source
is the port's, unchanged.

On an RTX 4060, device time against Triton (main): complex forward
1.32x, complex adjoint 0.78-0.82x, complex second order 1.01x, real
forward 1.44x, real adjoint 1.27x; the port was 1.4-4.1x on the same
cases. The sm_89 module grows from 7.9 to 15.7 MB.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Profiled against Triton on an RTX 4060, the kernels issued as well as
Triton's and simply executed more instructions, almost all of them
integer bookkeeping. In the tile runtime and the kernels' types:

- A costly function of a thread's rows is taken once where the rows
  share the argument, as a pulse's flip does; Triton sinks a broadcast
  below the function and the tile runtime could not.
- The EPG kernels index in 32 bits, as Triton passed every integer
  argument that fit; the launcher refuses one that does not.
- A row wider than a warp is gathered and reduced in a function of its
  own: inlined, its barriers and shared accesses were predicated and
  issued on every shift.
- A column of one row is not reduced, and one warp's rows are reduced
  by shuffles.
- Specialized kernels are compiled for rows a warp wide and launched
  only on them; a program is two warps.
- The values the adjoint kernels narrow to float before the per-order
  loop are declared float, as Triton rebinds them; declared double, the
  loop ran in fp64.

The series-against-roots test allows 5e-5: one fused multiply-add moves
the roots branch by 1e-5, and Triton's differs from the CUDA build's by
2e-5.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The complex and real forward kernels and their Jacobian-vector products
are written once over a number type: a float gives the simulation and a
dual number its derivative along a direction. Each is compiled for its
layout -- pools, how a pulse is formed, per-voxel maps, problems per
thread, one train or several -- and reads every other switch at run time,
branching around whole blocks once per event. A launch whose rows fit a
warp runs the layout ahead of any tile kernel.

The two-pool transverse step takes its complex root in a form that does
not cancel just above the negative real axis, where its discriminant
usually sits.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…layouts

The complex and real vector-Jacobian products are written once over a
number type, so their derivative along a direction is the same source at a
dual number. The forward sweep keeps the state every four events and the
reverse sweep replays each stretch into shared memory from its checkpoint,
which replaces a trajectory of every event in global memory. An interval's
tissue gradient sums, per lane, the products its operator met while the
interval repeats and contracts them against the operator's derivatives --
along every tissue input at once -- only when the interval changes.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The layout kernels serve every switch combination from one compile, so the
specialized tile kernels, their census and the table they were built from
go. A three-pool layout forms its interval operator out of line, which keeps
its compile to under a minute.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The forward and its JVP hold a problem's pools in a thread's registers, with
lanes along the orders. The adjoint does its own forward sweep, keeping every
fourth state, and replays each stretch from it on the way back, so a launch
it takes needs no recording first and a buffer of checkpoints rather than
the whole trajectory. A table row's cotangent is summed per lane while the
row repeats and reduced into the row when it changes; the stretch and those
sums are a warp's own scratch in global memory, where they stay in its caches
and leave the warp's occupancy to its registers.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
… card idle

A launch of a few thousand voxels is a few dozen programs, which leaves most
multiprocessors with nothing to run. Such a launch now also splits the
features across a second axis of programs, each adding its share into an
output that starts at zero, until there are four programs a multiprocessor.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The card's module is a package of its own per CUDA major version, installed
with blochsim[cu12] or blochsim[cu13] beside a torch of that major version:
the same CMake project with BLOCHSIM_CUDA_PACKAGE set, which builds _gpu alone
into blochsim_cudaNN/. It links the CUDA runtime torch's nvidia wheel brings,
found by rpath, rather than carrying one; carries machine code for 7.5, 8.0
and 9.0 and PTX for 9.0, in compressed fatbinaries; and pins the blochsim it
was built with. _gpu_launch loads the build for torch's CUDA major version and
refuses one for another major or release, naming the extra to install. The
base wheel is the CPU kernels alone. The wheels workflow builds each CUDA
package in a manylinux container with the oldest toolkit of its major torch
ships, installs it beside that torch on the oldest and newest Python, and
publishes each project from an environment of its own.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
A three-pool interval at float, or any at num::Dual, is a multi-valued
operator whose entries are wide, and it is held on the stack, which the
driver backs for every thread the card can hold. One direction per formation
of it keeps that stack small: the three-pool adjoint's from 1.9 KB to 0.9 KB,
under the default the driver already holds, and the dual three-pool's from
7.0 KB to 3.6 KB. Both are also faster.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
PyTorch's forward mode registers decompositions through TorchScript the first
time it makes a dual tensor in a process. Resolving a packing and building a
transition table each read a directional derivative of the package's own
structure, which is now taken by a vector-Jacobian product of the
vector-Jacobian product, so a simulation and its gradients compile nothing at
run time. Both directions a packing is read along share one walk of the
stream.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Comments and the agent notes said what the kernels do by reference to the
Triton source they replaced; they now say what the code does.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@mcencini
mcencini merged commit 83a8e16 into main Oct 8, 2026
35 checks passed
@mcencini
mcencini deleted the layout-kernels branch October 8, 2026 05:15
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants