Integrating flash-kmeans into the ColBERT repo (initial setup)

ColBERT
I’m dusting off my ColBERT maintenance hat and getting back into the rhythm of things…
Author

Vishal Bakshi

Published

August 14, 2026

I’m dusting off my ColBERT maintenance hat and getting back into the rhythm of things.

The first thing I’m working on for the next release is exploring the most likely integration of flash-kmeans (arxiv, github, pypi) as a replacement for faiss-gpu.

flash-kmeans (2026 Yang, et al) is a ridiculously fast “IO-aware batched K-Means clustering implemented with Triton GPU kernels.”

I’ll be digging more into its internals later this fall when I read the paper and do a deeper dive into their repo.

In this blog post I’m going to recap my experience debugging some transformers and CUDA errors during a basic flash-kmeans integration into the ColBERT repo.

Replacing faiss-gpu with flash-kmeans

This was pretty trivial to do (a near identical copy/paste of the existing FAISS implementation, thank you Yang and team!):

def compute_flash_kmeans(dim, num_partitions, kmeans_niters, shared_lists, return_value_queue=None):
    from flash_kmeans import batch_kmeans_Euclid # ta-da!

    sample = shared_lists[0][0]  # Extract sample (same as FAISS path)
    sample_cuda = sample.cuda() 
    sample_batched = sample_cuda.unsqueeze(0)
    torch.save(sample_batched, f"{ROOT}/flash_input_sample.pt")

    cluster_ids, centers, _ = batch_kmeans_Euclid(
        sample_batched,
        n_clusters=num_partitions,
        max_iters=kmeans_niters,  
        verbose=True
    )
    torch.save(centers, f"{ROOT}/flash_raw_centers.pt")
    centroids = centers.squeeze(0)  # (1, K, D) → (K, D)
    centroids = centroids.float().cpu()  # Match FAISS output: float32 on CPU

    print_memory_stats(f'RANK:0*')

    if return_value_queue is not None:
        return_value_queue.put(centroids)

    return centroids
...

if KMEANS_ALGO == "flash": # for debugging
    print("USING FLASH KMEANS ===========================") # for debugging
    centroids = compute_flash_kmeans(*args_)
else:
    centroids = compute_faiss_kmeans(*args_)

transformers==4.57.0 is yanked

Don’t use it.

ColBERT’s JIT-compiled CUDA extension decompress_residuals_cpp fails to build

I feel like I get this error every year and then forget how I resolved it, so hopefully this will be the last time I forget it.

First: I was erroneously installing both cuda -c nvidia/label/11.7.1 and torch==2.10.0 which is not compatible with it.
Second: Once I installed it was clashing with gcc 14, giving me a ninja: build stopped: subcommand failed. error so I had to pin it <14:

```
micromamba create -n colbert python=3.11 cuda-toolkit “gxx_linux-64<14” -c nvidia/label/cuda-12.6.0 -c conda-forge
```

Third: I was getting a ninja build error which needed the following as advised by contributor Robin Narsingh Ranabhat and modified a bit by Opus 4.6 to fit my Dockerfile

```
ENV CONDA_DEFAULT_ENV=colbert
ENV PATH=/opt/conda/envs/colbert/bin:\(PATH ENV CONDA\_PREFIX=/opt/conda/envs/colbert ENV CC=\)CONDA_PREFIX/bin/x86_64-conda-linux-gnu-gcc
ENV CXX=\(CONDA\_PREFIX/bin/x86\_64-conda-linux-gnu-g++ ENV CUDA\_HOST\_COMPILER=\)CONDA_PREFIX/bin/x86_64-conda-linux-gnu-g++
```

AttributeError: ‘HF_ColBERT’ object has no attribute ‘all_tied_weights_keys’

transformers defines all_tied_weights_keys in the PretrainedModel.post_init() method. jkaniewski-tii’s comment on this very informative and helpful issue thread advised the reader to call self.post_init() in the model’s init. I added that to hf_colbert.py’s HFColBERT init and it resolved this issue!

Ghost indentation error

After adding that the first couple of times, I rebuilt the image in the modal, and I got an indentation error in the file. The third time I deployed the app, I didn’t, so, shrug.

Next: analyze flash-kmeans artifacts during indexing!

My absolute favorite part of ColBERT Maintenance, and really any project, is following the data through the pipeline, in this case, the indexing pipeline. Every time any data goes through a transformation, I save the artifact so I can inspect it later. This process will probably take at least a week, and I’ll be posting updates on Twitter (@vishal_learner) as it happens, along with another blog post. Happy maintaining!