Instructions to use handwoven8588/CodeRankEmbed-flash-attn with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use handwoven8588/CodeRankEmbed-flash-attn with sentence-transformers:
from sentence_transformers import SentenceTransformer model = SentenceTransformer("handwoven8588/CodeRankEmbed-flash-attn", trust_remote_code=True) sentences = [ "The weather is lovely today.", "It's so sunny outside!", "He drove to the stadium." ] embeddings = model.encode(sentences) similarities = model.similarity(embeddings, embeddings) print(similarities.shape) # [3, 3] - Notebooks
- Google Colab
- Kaggle
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β theflash_attnpackage's varlen kernel. Kept as a fallback for older torch builds that don't yet havetorch.nn.attention.varlenbut do haveflash_attninstalled. 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):torch_varlenβ CUDA, compute capability sm_80+ (Ampere or newer), torch β₯ 2.10.0. No third-party kernel needed.flash_attnβ CUDA, compute capability sm_80+, theflash_attnpackage importable, on a torch that doesn't yet shiptorch.nn.attention.varlen(typically an older torch). This dependency is optional.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) preferstorch_varlen, thenflash_attn, theneager. Before accepting a varlen tier,autoruns 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 aRuntimeError(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 raisesRuntimeErrorif that tier's precondition doesn't hold β a forced tier never falls back silently. An unrecognized override raisesValueError.See which tier engaged:
model[0].auto_model.attention_implafter 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_statekeeps 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 theflash_attnpackage 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 theeagertier the flattened batch is re-padded and runs the padded path.Loads bf16 by default.
flash_attnandtorch_varlenboth require half precision and the model runs bf16 in any real serving setup, so the weights are stored bf16 andconfig.jsondeclarestorch_dtype: bfloat16. The original's customfrom_pretrainedsilently droppedtorch_dtypeand always loaded fp32; the copy in this repo honors it, so the model loads bf16 natively, like any normal HF model. Passtorch_dtype=torch.float32to load fp32 (note: the stored weights are bf16-precision, so this only widens the dtype, not the precision).eagerruns 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
- Weights: fp32 β bf16.
flash_attnandtorch_varlenonly 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. from_pretraineddtype fix: the original's customfrom_pretrainedinstantiated the model fp32 andload_state_dict-ed the checkpoint into fp32 params, ignoringtorch_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.- Three attention tiers:
NomicBertAttention.forwardnow selects one of three attention implementations at call time βtorch_varlen(torch's owntorch.nn.attention.varlen.varlen_attn, no third-party kernel, needs torch β₯ 2.10.0),flash_attn(theflash_attnpackage's varlen kernel, kept as a fallback for older torch builds that have it installed), andeager(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 adevice_mapsplit hands it a mask a varlen tier upstream never passes). On both varlen tiersNomicBertModel.forwardunpads the batch once with torch-native_unpad/_padhelpers (a replacement forflash_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 tierNomicBertModel.forwardbuilds the additive attention mask inline instead of calling the (now-removed-upstream)get_extended_attention_maskhelper. SetNOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eagerto force a tier (raises if it can't engage);model[0].auto_model.attention_impland the one-time INFO log line report which tier engaged. from_pretrainedforwards the Hub revision: the original's customfrom_pretrainedfetched the weight file without the caller'srevision(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 oldertorch_dtype=, and returns the model in eval mode, as transformers' ownfrom_pretraineddoes.- Flattened input:
NomicBertModel.forwardalso accepts sentence-transformers' flattened batches (input_ids [1, Ξ£L]with per-sequenceposition_idsandcu_seq_lens_q), and the model declares flash-attention support so that sentence-transformers sends them when loaded withattn_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
Model tree for handwoven8588/CodeRankEmbed-flash-attn
Base model
Snowflake/snowflake-arctic-embed-m-long