b11405
Pre-releasecuda: tile the lightning indexer over keys and tokens for 4 heads (#29901)
- cuda: tile the lightning indexer over keys and tokens for 4 heads
With too few heads for a wmma tile, a block scores 64 keys against 8
tokens: the keys are staged once in half precision, the queries one
head at a time, and each thread owns one key for two tokens, so no dot
product needs a cross thread reduction. Batches smaller than a token
tile keep the vector kernel. test-backend-ops measures 4 heads.
- cuda: multiply the lightning indexer tile in float
Address review from am17an: the half2 products overflow once a single
q * k exceeds the f16 range. The queries stay in float in shared memory
and each half2 of keys is widened once for both tokens, so every
product and sum is computed in float.
- cuda: widen each lightning indexer key once for all heads
The tile kernel stages the queries and weights of every head at once,
so each key element is widened from half once and feeds all heads,
with a single barrier. F16 keys are copied into the tile without a
float round trip. Keeping the keys in float in shared memory measures
slower, the occupancy drops.
- cuda: stop the lightning indexer tile from spilling registers on ROCm
Each thread of the tile kernel now scores two keys for a single token,
so a warp shares its token and the query reads are broadcasts: six
shared reads per element pair instead of nine for the same products.
The inner loop is unrolled by 8, which keeps gfx908 at 63 VGPRs with no
spill where the fully unrolled loop needed over a thousand, and makes
the kernel 36x faster on an R9700 and slightly faster on CUDA.
Website:
Attestations:
macOS/iOS:
- macOS Apple Silicon (arm64)
- macOS Apple Silicon (arm64, KleidiAI enabled) DISABLED
- macOS Intel (x64)
- iOS XCFramework
Linux:
- Ubuntu x64 (CPU)
- Ubuntu arm64 (CPU)
- Ubuntu s390x (CPU)
- Ubuntu x64 (Vulkan)
- Ubuntu arm64 (Vulkan)
- Ubuntu x64 (CUDA 12) - CUDA 12.8 libraries
- Ubuntu x64 (CUDA 13) - CUDA 13.4 libraries
- Ubuntu arm64 (CUDA 13) - CUDA 13.4 libraries
- Ubuntu x64 (ROCm 10.0)
- Ubuntu x64 (OpenVINO)
- Ubuntu x64 (SYCL FP32)
- Ubuntu x64 (SYCL FP16)
- Linux arm64 (Snapdragon: CPU, Adreno GPU, Hexagon NPU) - setup guide
Android:
Windows:
- Windows x64 (CPU)
- Windows arm64 (CPU)
- Windows arm64 (OpenCL Adreno)
- Windows x64 (CUDA 12) - CUDA 12.4 DLLs
- Windows x64 (CUDA 13) - CUDA 13.4 DLLs
- Windows arm64 (CUDA 13) - CUDA 13.4 DLLs
- Windows x64 (Vulkan)
- Windows arm64 (Vulkan)
- Windows x64 (OpenVINO)
- Windows x64 (SYCL)
- Windows x64 (ROCm 10.0)
openEuler:
- DISABLED
- openEuler x86 (310p)
- openEuler x86 (910b, ACL Graph)
- openEuler aarch64 (310p)
- openEuler aarch64 (910b, ACL Graph)
UI: