Skip to content

Tile the PERK kernels as matrix products - #32

Closed
mcencini wants to merge 1 commit into
claude/project-thread-9peiwhfrom
perk-tiled
Closed

mcencini wants to merge 1 commit into
claude/project-thread-9peiwhfrom
perk-tiled

Conversation

@mcencini

@mcencini mcencini commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

What this changes

On a card, the PERK kernels in #31 ran 2.5–2.8× slower than the Triton kernels they replace. They were not a translation of Triton's kernel. Triton computed PERK as a tiled matrix product (tl.dot over 128×64 tiles, 8 warps). The port had each thread loop over its own voxel's features, with one global load of frequency per multiply.

This PR tiles the two products the way a matrix product is tiled:

  • A program stages a block of signals and a block of frequencies in shared memory.
  • Each thread forms the angles of 2 voxels × 32 features (16 in the adjoint) in registers, reading weights four at a time.
  • The cosine and the second product consume the block while it is still in registers. The (voxels, features) matrix still never exists.
  • Staged blocks are zero past the last feature, parameter and voxel, so the inner loops carry no bounds.
  • The accumulators are stored through unrolled loops with a guard. A store loop with a runtime bound had moved the adjoint's gradient into local memory for every multiply; unrolling took the adjoint from 13.1 ms to 4.6 ms.

The host build still compiles the same source. PERK_EACH_THREAD / PERK_SYNC are a thread and __syncthreads on the card. On the host they become a loop over the program's threads, ending at each barrier.

Everything stays single precision.

Measurements

RTX 4060 Laptop (sm_89), CUDA 12.8, PyTorch 2.13. 24 contrasts, 256 features, 2 parameters. Kernel device time comes from the Torch profiler; one process per measurement.

Voxels Triton (main) forward #31 forward this PR forward Triton adjoint #31 adjoint this PR adjoint
2,000 0.081 ms 0.455 ms 0.126 ms
5,000 0.160 ms 0.137 ms
20,000 0.547 ms 0.301 ms 0.572 ms 0.548 ms
50,000 1.31 ms 0.68 ms
200,000 5.09 ms 12.93 ms 2.40 ms 5.35 ms 14.04 ms 4.58 ms

End to end at 200k voxels, the estimator's forward call takes 2.05 ms (Triton 4.11, #31 9.95). Forward plus gradient takes 5.57 ms (Triton 8.26, #31 21.40). Peak and total device memory are identical across the three builds.

Known limit: below about 4k voxels, a program's serial chain over the features is longer than Triton's. Its kernel is behind by up to 45 µs, though the whole call is still faster at 2k voxels (0.23 vs 0.26 ms) because launching costs less. Splitting the features across programs, with atomic accumulation, would close that gap; it is not done here.

How it was checked

  • test_the_gpu_kernels_are_the_fused_line_and_its_adjoint now runs on the card as well as on the host build. Its sizes cross every edge of the new tiling: 300 voxels, 37 contrasts, 45 features, 17 parameters.
  • Against float64 Torch, on host and card, the largest error relative to the answer's scale is below 2e-6. The sizes checked were (voxels, contrasts, features, parameters) = (300, 37, 45, 17), (1, 24, 256, 2), (1000, 5, 1000, 3) and (513, 64, 70, 8).
  • pytest tests/estimators/ on the card: 154 passed, 1 skipped (it needs two CUDA devices).
  • ruff format and ruff check are clean on the changed Python.
  • cuobjdump --dump-resource-usage for the 256-thread variants: forward 118 registers, adjoint 199, 32 bytes of stack each.

🤖 Generated with Claude Code

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>
mcencini added a commit that referenced this pull request Oct 8, 2026
…hem as blochsim[cu12]/[cu13] (#33)

The EPG, many-pool and PERK kernels are CUDA compiled ahead of time and written for their layouts: a layout is what is compiled, every other switch steers whole blocks at run time, and each loop is written once over a number type, so the JVPs and the second-order adjoints are the same source at a dual. The adjoints keep checkpoints instead of a recorded trajectory. Every case measured runs at or under Triton's device time; device memory is equal or lower except the three-pool second-order adjoint, which leaves 84 MiB of local memory against Triton's 52.

The card's module ships as blochsim-cuda12 and blochsim-cuda13, installed with blochsim[cu12] or blochsim[cu13] beside a torch of the same CUDA major version: it links torch's CUDA runtime by rpath, carries code for 7.5, 8.0 and 9.0 and PTX for 9.0, and is refused by name for another major or release. The package's own structural derivatives are taken in reverse mode, so a simulation and its gradients make PyTorch compile nothing through TorchScript.

Supersedes #31 and #32.

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

mcencini commented Oct 8, 2026

Copy link
Copy Markdown
Contributor Author

Superseded by #33, merged as 83a8e16: it contains this PR's commits (562a6cb), and the PERK kernels it tiles now also split their features across programs when a launch is too small to fill the card.

@mcencini mcencini closed this Oct 8, 2026
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.

1 participant