Repository navigation
Conversation
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>
Contributor
Author
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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.dotover 128×64 tiles, 8 warps). The port had each thread loop over its own voxel's features, with one global load offrequencyper multiply.This PR tiles the two products the way a matrix product is tiled:
(voxels, features)matrix still never exists.The host build still compiles the same source.
PERK_EACH_THREAD/PERK_SYNCare a thread and__syncthreadson 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.
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_adjointnow 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.pytest tests/estimators/on the card: 154 passed, 1 skipped (it needs two CUDA devices).ruff formatandruff checkare clean on the changed Python.cuobjdump --dump-resource-usagefor the 256-thread variants: forward 118 registers, adjoint 199, 32 bytes of stack each.🤖 Generated with Claude Code