FLASH-MAXSIM: IO-Aware Fused Kernels for Late-Interaction Scoring

📅 2026-05-28
📈 Citations: 0
✨ Influential: 0
📄 PDF
🤖 AI Summary
This work addresses the GPU memory bottleneck in traditional late-interaction retrieval caused by explicit construction of large similarity tensors during MaxSim computation, which severely limits batch size and scalability. The authors propose Flash-MaxSim, an I/O-aware fused GPU kernel that streams query and document chunks through on-chip SRAM and computes row-wise max reductions in a single pass, enabling exact MaxSim evaluation without materializing intermediate tensors for the first time. The method supports backpropagation, INT8 quantization, and variable-length padding-free scoring, and introduces key innovations including inverse-grid CSR layout and atomic-free gradient reduction. Evaluated on A100/H100 GPUs, Flash-MaxSim achieves 3.9–4.7× faster inference, reduces inference and training memory consumption by 16× and 28× respectively, substantially expands tractable corpus and batch sizes, and maintains 100% top-20 ranking consistency.
📝 Abstract
Late-interaction retrieval (ColBERT, ColPali) scores a query against a document with the MaxSim operator: for every query token, the maximum similarity over the document tokens, summed over query tokens. The standard implementation materializes the full query-token x document-token similarity tensor in GPU memory; for visual ColPali at 10K documents this tensor alone is 21 GB in FP16, created only to be reduced to one score per document and discarded. It exhausts a 40 GB GPU and bounds the achievable batch size in both inference and training. We present Flash-MaxSim, an IO-aware fused GPU kernel that computes exactly the same scores without ever materializing the tensor, by streaming query and document tiles through on-chip SRAM and folding the row-maximum reduction into the same pass. We extend the IO-aware principle through the training backward pass, an inverse-grid CSR construction that reuses the forward argmax for an atomic-free, destination-owned gradient reduction, and through INT8xINT8 quantization and variable-length (padding-free) scoring. Flash-MaxSim is up to 3.9x faster on an A100 (4.7x on an H100) than naive PyTorch at matched precision, uses up to 16x less inference memory and ~28x less training memory, unlocks corpus and batch sizes that exhaust PyTorch entirely, preserves the exact ranking (100% top-20 agreement with an FP32 reference)
Problem

Research questions and friction points this paper is trying to address.

late-interaction retrieval
MaxSim
GPU memory bottleneck
similarity tensor
batch size limitation
Innovation

Methods, ideas, or system contributions that make the work stand out.

IO-aware
fused kernel
MaxSim
late-interaction retrieval
memory-efficient
🔎 Similar Papers
No similar papers found.
💼 Related Jobs
No related jobs found.
R
Roi Pony
IBM Research Israel
A
Adi Raz Goldfarb
IBM Research Israel
I
Idan Friedman
IBM Research Israel
D
Daniel Ezer
IBM Research Israel
U
Udi Barzelay
IBM Research Israel