FLASH-MAXSIM: IO-Aware Fused Kernels for Late-Interaction Retrieval
Late-interaction retrieval (ColBERT, ColPali) scores a query against a document via the MaxSim operator. The standard PyTorch implementation materialises the full query-token $\times$ document-token similarity tensor only to reduce it away. At ColPali scale this is the single largest tensor in the pipeline (e.g. 21 GB in FP16 for 10K documents) and limits both candidate set size at inference and batch size during contrastive training. We present FLASH-MAXSIM (FM), an IO-aware fused GPU kernel that computes the same MaxSim scores without ever materialising the tensor, and extends the same principle to the training backward. At ColPali scale on A100, FM reduces inference peak memory by 1.4-2.6$\times$ relative to the deployed chunked baseline (4.9-8.9$\times$ relative to unchunked eager execution) and reduces MaxSim-operator training memory by two orders of magnitude, enabling exact reranking over larger resident candidate pools and contrastive batch sizes that vanilla autograd cannot fit on a single GPU. The kernel is a drop-in replacement, exact up to floating-point evaluation order under its stated FP32-accumulation protocol: nDCG@10 differs from the FP32 reference by at most $5\times10^{-4}$ on BEIR and REAL-MM-RAG. A separate INT8 path trades exactness for halved index storage at high fidelity. Code, benchmark scripts, and raw results: https://github.com/roipony/flash-maxsim