CodeRankEmbed-flash-attn

A bf16 quantization of nomic-ai/CodeRankEmbed with a custom modeling_hf_nomic_bert.py, shipped in this repo, that picks one of three attention implementations for the hardware it runs on. It is not a finetune β€” the weights are the original CodeRankEmbed weights cast to bf16 (no further training). Two of the three replace the original's padded O(seqΒ²) attention with an O(N) unpadded path; the third keeps the original's algorithm as the correctness reference and universal fallback.

Why

nomic-ai/CodeRankEmbed loads through trust_remote_code, and its attention path is eager only β€” activation memory grows as batch Γ— heads Γ— seqΒ², which runs out of memory at large batches even though the model is only 137M params. This repo adds two attention paths that compute the same attention in O(N) memory by packing unpadded sequences, so the large batches that the eager path cannot fit run comfortably β€” with parity embeddings (no quality change):

  • torch_varlen β€” torch.nn.attention.varlen.varlen_attn, shipped in torch itself (no extra package), available from torch 2.10.0 onward.
  • flash_attn β€” the flash_attn package's varlen kernel. Kept as a fallback for older torch builds that don't yet have torch.nn.attention.varlen but do have flash_attn installed. This package is optional.

Both run the same FA2-family kernel and are gated by the same GPU-capability check. The modeling file ships all three paths itself, so no runtime patching or post-load hooks are needed.

Behavior

  • Three attention tiers, chosen automatically per device (override with NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager):

    1. torch_varlen β€” CUDA, compute capability sm_80+ (Ampere or newer), torch β‰₯ 2.10.0. No third-party kernel needed.
    2. flash_attn β€” CUDA, compute capability sm_80+, the flash_attn package importable, on a torch that doesn't yet ship torch.nn.attention.varlen (typically an older torch). This dependency is optional.
    3. eager β€” everything else: CPU, pre-Ampere GPUs, or neither of the above available. The original's padded attention algorithm, unchanged, runs on any host.

    auto (the default) prefers torch_varlen, then flash_attn, then eager. Before accepting a varlen tier, auto runs one tiny kernel probe per device (a capability check alone can't tell whether the installed build has a kernel for the GPU's architecture, e.g. on ROCm); if the probe raises a RuntimeError (how a missing or unsupported kernel fails), it logs one WARNING naming the tier and the error and steps down to the next tier. An out-of-memory error, any other exception type, or a CUDA error left over from earlier work propagates instead of demoting the tier, so a transient failure never pins the device to a slower tier. A forced override (e.g. NOMIC_BERT_ATTN_IMPL=torch_varlen) is not probed and raises RuntimeError if that tier's precondition doesn't hold β€” a forced tier never falls back silently. An unrecognized override raises ValueError.

  • See which tier engaged: model[0].auto_model.attention_impl after a forward pass. The modeling file also logs the tier, and why, once at INFO.

  • Padding-free on the varlen tiers. The model unpads the batch once, right after the input ids, runs the embeddings and every encoder layer on the real tokens only, and pads back once at the end. last_hidden_state keeps its [B, S, 768] shape, with zeros at padding positions, so pooling is unchanged.

  • Flattened batches (opt-in). Loading with model_kwargs={"attn_implementation": "flash_attention_2"} lets sentence-transformers send each batch flattened ([1, Ξ£L] token ids with sequence boundaries, no padding at all), and the model returns one [1, Ξ£L, 768] output that sentence-transformers pools per sequence. The opt-in needs the flash_attn package installed, because transformers checks for it at load and raises otherwise, even though the attention itself still runs on whichever tier the model picks. On the eager tier the flattened batch is re-padded and runs the padded path.

  • Loads bf16 by default. flash_attn and torch_varlen both require half precision and the model runs bf16 in any real serving setup, so the weights are stored bf16 and config.json declares torch_dtype: bfloat16. The original's custom from_pretrained silently dropped torch_dtype and always loaded fp32; the copy in this repo honors it, so the model loads bf16 natively, like any normal HF model. Pass torch_dtype=torch.float32 to load fp32 (note: the stored weights are bf16-precision, so this only widens the dtype, not the precision).

  • eager runs in bf16 too (because the stored weights are bf16), numerically equivalent to the varlen tiers, just without their memory and throughput wins. The model loads and encodes on any host regardless of which tier engages.

Usage

Identical to the original. The query prompt must include the task-instruction prefix "Represent this query for searching relevant code: "; documents need no prefix.

from sentence_transformers import SentenceTransformer

model = SentenceTransformer("handwoven8588/CodeRankEmbed-flash-attn", trust_remote_code=True)
queries = ["Represent this query for searching relevant code: Calculate the n-th factorial"]
codes   = ["def fact(n):\n    if n < 0:\n        raise ValueError\n    return 1 if n == 0 else n * fact(n - 1)"]

q = model.encode(queries, normalize_embeddings=True)
d = model.encode(codes,   normalize_embeddings=True)

With flash_attn installed, opt in to flattened batches:

model = SentenceTransformer(
    "handwoven8588/CodeRankEmbed-flash-attn",
    trust_remote_code=True,
    model_kwargs={"attn_implementation": "flash_attention_2"},
)

On a varlen tier the encoder layers run on packed tokens only, so they compile cleanly. For about 1.2Γ— the throughput and a third less peak memory (see Speed and memory below):

import torch

model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True)

dynamic=True matters: the number of packed tokens changes with every batch.

Parity & performance

Every number below compares this repo with the original model, nomic-ai/CodeRankEmbed (revision 3c4b608), run as published: fp32 weights, its own modeling file, padded eager attention. Both run on the same GPU, an RTX 3090 Ti (sm_86). The original's modeling file calls get_extended_attention_mask, which transformers has deprecated and no longer ships by 5.17, so it was measured on transformers 5.11; this repo does not call it and loads on 5.17 too.

Corpus. The first 16,384 document fields (real Python functions) of lightonai/cornstack's Python split, in shard order (train-00000-of-00682.parquet, then train-00001-…): 2,684,878 tokens, 7 to 4,746 per function, 164 mean, 92 median. Each row is one encode() call over all of them at the stated batch_size=, after one untimed warm-up of the same call. Cosine similarity is per function, on fp32-renormalized output, against the original's embedding of the same function. Peak memory is the encode's peak CUDA allocation above the loaded weights. Load is how the model was loaded: default (as in Usage) or flattened (the opt-in above). torch 2.12.1+cu130, transformers 5.11.0, sentence-transformers 6.1.0, flash-attn 2.8.3.post1.

Agreement with the original model

model attention compiled load batch size min cosine mean cosine
nomic-ai/CodeRankEmbed eager – default 4 1 (the reference) 1 (the reference)
handwoven8588/CodeRankEmbed-flash-attn eager – default 4 0.99902 0.99991
handwoven8588/CodeRankEmbed-flash-attn eager – flattened 4 0.99916 0.99991
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – default 256 0.99896 0.99993
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – flattened 256 0.99896 0.99993
handwoven8588/CodeRankEmbed-flash-attn flash_attn – default 256 0.99841 0.99993
handwoven8588/CodeRankEmbed-flash-attn flash_attn – flattened 256 0.99841 0.99993
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ default 256 0.99960 0.99995
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ flattened 256 0.99960 0.99995
handwoven8588/CodeRankEmbed-flash-attn flash_attn βœ“ default 256 0.99946 0.99995
handwoven8588/CodeRankEmbed-flash-attn flash_attn βœ“ flattened 256 0.99946 0.99995

The original is the reference every other row is measured against, so its cosine is 1 by definition, not a measurement. The mean cosine is 0.99991 or higher in every configuration; the minimum is the one function, of 16,384, that moves most under bf16 weights and arithmetic. The batch size does not change these figures (a varlen tier gives the same five decimals at batch 32 and 256), so the largest batch that ran in every load is shown; eager runs at batch 4 (see below). Compiled rows sit closer to the original than uncompiled ones.

Speed and memory

model attention compiled load batch size peak memory time
nomic-ai/CodeRankEmbed eager – default 4 8.4 GiB 68.9 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – default 32 2.5 GiB 16.3 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – default 256 10.9 GiB 15.1 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – default 1,024 out of memory out of memory
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – default 2,048 out of memory out of memory
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – flattened 32 1.4 GiB 15.5 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – flattened 256 6.9 GiB 13.6 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – flattened 1,024 17.0 GiB 13.4 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen – flattened 2,048 out of memory out of memory
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ default 32 1.6 GiB 14.0 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ default 256 7.1 GiB 12.8 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ default 1,024 16.3 GiB 13.8 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ default 2,048 out of memory out of memory
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ flattened 32 0.9 GiB 13.2 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ flattened 256 4.5 GiB 11.4 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ flattened 1,024 11.1 GiB 10.9 s
handwoven8588/CodeRankEmbed-flash-attn torch_varlen βœ“ flattened 2,048 16.6 GiB 10.8 s

This repo's rows are torch_varlen; flash_attn runs within 1% of it in both time and memory at every cell. The original runs at batch 4: its padded fp32 attention needs batch Γ— heads Γ— seqΒ² memory, 8.4 GiB at batch 4 for this corpus's longest function. The same algorithm in bf16 (this repo's eager tier, batch 4) takes 41.2 s in 4.2 GiB; the rest of the gain comes from the unpadded path and the batches it allows. Compiled means model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True).

  • Against the original, the same 16,384 functions encode 5.1Γ— faster uncompiled (13.4 s, flattened at batch 1,024) and 6.4Γ— faster compiled (10.8 s, flattened at batch 2,048), against 68.9 s. At batch 32 this repo needs 2.5 GiB (default load) or 1.4 GiB (flattened), against the original's 8.4 GiB at batch 4.
  • Memory follows the real tokens in the heaviest batch, not the number of functions. This repo runs the encoder on real tokens only, and every out-of-memory cell failed allocating one activation over that batch's tokens: the 3,072-wide MLP hidden state, or, compiled, the 2,304-wide QKV projection. sentence-transformers length-sorts the corpus, so the default load's first batch of 1,024 holds the longest functions, about 870,000 tokens; flattened, it pairs the longest with the shortest, so its heaviest batch of 1,024 holds about 594,000.
  • Flattening cuts peak memory by 32–42% against the default load at the same batch size, and wall time by 5–21% (more at larger batches).
  • Compiling is 1.17–1.23Γ— faster at the same batch size and cuts peak memory by 35%. It compiles one graph, with no graph breaks, once per process (the first encode takes about 5 s longer), and does not recompile as batch lengths change; the only further compiles come once a batch passes about 699,000 and 932,000 tokens, where inductor switches to 64-bit indexing for the 3,072-wide MLP activation and the 2,304-wide QKV projection.

What changed vs the source repo

  1. Weights: fp32 β†’ bf16. flash_attn and torch_varlen only accept half precision and the model runs bf16 in any real serving configuration, so the weights are stored bf16 and (via the load fix below) arrive bf16 β€” which is simply how this model is used, and removes the need for a post-load dtype cast. Parity-neutral; the smaller download is incidental, not the reason.
  2. from_pretrained dtype fix: the original's custom from_pretrained instantiated the model fp32 and load_state_dict-ed the checkpoint into fp32 params, ignoring torch_dtype. The copy here adds the standard transformers dtype resolution (explicit arg β†’ config.torch_dtype β†’ checkpoint dtype) so the model loads in its declared dtype.
  3. Three attention tiers: NomicBertAttention.forward now selects one of three attention implementations at call time β€” torch_varlen (torch's own torch.nn.attention.varlen.varlen_attn, no third-party kernel, needs torch β‰₯ 2.10.0), flash_attn (the flash_attn package's varlen kernel, kept as a fallback for older torch builds that have it installed), and eager (the original's padded attention algorithm, numerically unchanged, and the default off CUDA sm_80+; it now converts a raw 2-D [B, S] mask to the additive form before adding it, for the case where a device_map split hands it a mask a varlen tier upstream never passes). On both varlen tiers NomicBertModel.forward unpads the batch once with torch-native _unpad/_pad helpers (a replacement for flash_attn.bert_padding), keeps the hidden states packed [Ξ£L, 768] through the embeddings and every layer, and pads back once at the end. Rotary embeddings rotate each packed token by its position within its own sequence, the same rotation it gets in the padded batch. On the eager tier NomicBertModel.forward builds the additive attention mask inline instead of calling the (now-removed-upstream) get_extended_attention_mask helper. Set NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager to force a tier (raises if it can't engage); model[0].auto_model.attention_impl and the one-time INFO log line report which tier engaged.
  4. from_pretrained forwards the Hub revision: the original's custom from_pretrained fetched the weight file without the caller's revision (and cache/token options), so a pinned load still took the weights from the default branch. The copy here passes them through, so a pinned load fetches the weights from the pinned commit too. It also accepts transformers' dtype= alongside the older torch_dtype=, and returns the model in eval mode, as transformers' own from_pretrained does.
  5. Flattened input: NomicBertModel.forward also accepts sentence-transformers' flattened batches (input_ids [1, Ξ£L] with per-sequence position_ids and cu_seq_lens_q), and the model declares flash-attention support so that sentence-transformers sends them when loaded with attn_implementation="flash_attention_2" (see Behavior).

License & attribution

MIT β€” same license as nomic-ai/CodeRankEmbed (see NOTICE). The weights, tokenizer, and the bulk of the modeling file are a verbatim derivative of nomic-ai/CodeRankEmbed; the modeling file derives from Tri Dao's BERT implementation, and CodeRankEmbed was trained by the CoRNStack team (Suresh et al., 2025). Cite their work:

@misc{suresh2025cornstackhighqualitycontrastivedata,
      title={CoRNStack: High-Quality Contrastive Data for Better Code Retrieval and Reranking},
      author={Tarun Suresh and Revanth Gangi Reddy and Yifei Xu and Zach Nussbaum and Andriy Mulyar and Brandon Duderstadt and Heng Ji},
      year={2025},
      eprint={2412.01007},
      archivePrefix={arXiv},
      primaryClass={cs.CL},
      url={https://arxiv.org/abs/2412.01007},
}
Downloads last month
12,715
Safetensors
Model size
0.1B params
Tensor type
BF16
Β·
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Model tree for handwoven8588/CodeRankEmbed-flash-attn

Quantized
(15)
this model

Paper for handwoven8588/CodeRankEmbed-flash-attn