sentence-transformers
Safetensors
English
nomic_bert
flash-attention
code-retrieval
nomic-bert
bf16
custom_code
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
Card: every comparison against nomic-ai/CodeRankEmbed; plain wording; citation fixed
#2
by handwoven8588 - opened
README.md
CHANGED
|
@@ -16,23 +16,23 @@ language:
|
|
| 16 |
# CodeRankEmbed-flash-attn
|
| 17 |
|
| 18 |
A **bf16 quantization of [`nomic-ai/CodeRankEmbed`](https://huggingface.co/nomic-ai/CodeRankEmbed)**
|
| 19 |
-
with a
|
| 20 |
-
|
| 21 |
-
(no further training). Two of the three
|
| 22 |
-
`O(N)` unpadded path; the third keeps the original
|
| 23 |
-
universal fallback.
|
| 24 |
|
| 25 |
## Why
|
| 26 |
|
| 27 |
`nomic-ai/CodeRankEmbed` loads through `trust_remote_code`, and its attention path is **eager
|
| 28 |
-
only** β activation memory grows as `batch Γ heads Γ seqΒ²`, which
|
| 29 |
-
the model is only 137M params. This repo adds two attention paths that compute
|
| 30 |
-
`O(N)` memory by packing unpadded sequences, so the large batches that
|
| 31 |
-
comfortably β with **parity embeddings** (no quality change):
|
| 32 |
|
| 33 |
- `torch_varlen` β `torch.nn.attention.varlen.varlen_attn`, shipped in torch itself (no extra
|
| 34 |
package), available from torch **2.10.0** onward.
|
| 35 |
-
- `flash_attn` β the
|
| 36 |
builds that don't yet have `torch.nn.attention.varlen` but do have `flash_attn` installed. This
|
| 37 |
package is **optional**.
|
| 38 |
|
|
@@ -41,7 +41,7 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
|
|
| 41 |
|
| 42 |
## Behavior
|
| 43 |
|
| 44 |
-
- **Three
|
| 45 |
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager`):
|
| 46 |
1. **`torch_varlen`** β CUDA, compute capability sm_80+ (Ampere or newer), torch **β₯ 2.10.0**. No
|
| 47 |
third-party kernel needed.
|
|
@@ -49,7 +49,7 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
|
|
| 49 |
torch that doesn't yet ship `torch.nn.attention.varlen` (typically an older torch). This
|
| 50 |
dependency is **optional**.
|
| 51 |
3. **`eager`** β everything else: CPU, pre-Ampere GPUs, or neither of the above available. The
|
| 52 |
-
original padded attention algorithm, unchanged, runs on any host.
|
| 53 |
|
| 54 |
`auto` (the default) prefers `torch_varlen`, then `flash_attn`, then `eager`. Before accepting a
|
| 55 |
varlen tier, `auto` runs one tiny kernel probe per device (a capability check alone can't tell
|
|
@@ -57,9 +57,10 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
|
|
| 57 |
raises a `RuntimeError` (how a missing or unsupported kernel fails), it logs one WARNING naming
|
| 58 |
the tier and the error and steps down to the next tier. An out-of-memory error, any other
|
| 59 |
exception type, or a CUDA error left over from earlier work propagates instead of demoting the
|
| 60 |
-
tier, so a transient failure never pins the device to a slower tier. A forced override (e.g.
|
| 61 |
-
`
|
| 62 |
-
silently. An unrecognized override
|
|
|
|
| 63 |
- **See which tier engaged**: `model[0].auto_model.attention_impl` after a forward pass. The
|
| 64 |
modeling file also logs the tier, and why, once at INFO.
|
| 65 |
- **Padding-free on the varlen tiers.** The model unpads the batch once, right after the input
|
|
@@ -71,11 +72,11 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
|
|
| 71 |
batch flattened (`[1, Ξ£L]` token ids with sequence boundaries, no padding at all), and the model
|
| 72 |
returns one `[1, Ξ£L, 768]` output that sentence-transformers pools per sequence. The opt-in needs
|
| 73 |
the `flash_attn` package installed, because transformers checks for it at load and raises
|
| 74 |
-
otherwise, even though the attention itself still runs on whichever tier the
|
| 75 |
the `eager` tier the flattened batch is re-padded and runs the padded path.
|
| 76 |
- **Loads bf16 by default.** `flash_attn` and `torch_varlen` both require half precision and the
|
| 77 |
model runs bf16 in any real serving setup, so the weights are stored bf16 and `config.json`
|
| 78 |
-
declares `torch_dtype: bfloat16`. The
|
| 79 |
`torch_dtype` and always loaded fp32; the copy in this repo honors it, so the model loads bf16
|
| 80 |
natively, like any normal HF model. Pass `torch_dtype=torch.float32` to load fp32 (note: the
|
| 81 |
stored weights are bf16-precision, so this only widens the dtype, not the precision).
|
|
@@ -122,106 +123,92 @@ model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic
|
|
| 122 |
|
| 123 |
## Parity & performance
|
| 124 |
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
|
|
|
|
|
|
| 129 |
|
| 130 |
**Corpus.** The first 16,384 `document` fields (real Python functions) of
|
| 131 |
[`lightonai/cornstack`](https://huggingface.co/datasets/lightonai/cornstack)'s Python split, in
|
| 132 |
shard order (`train-00000-of-00682.parquet`, then `train-00001-β¦`): 2,684,878 tokens, 7 to 4,746
|
| 133 |
per function, 164 mean, 92 median. Each row is one `encode()` call over all of them at the stated
|
| 134 |
`batch_size=`, after one untimed warm-up of the same call. Cosine similarity is per function, on
|
| 135 |
-
fp32-renormalized output, against
|
| 136 |
-
the encode's peak CUDA allocation above the loaded weights
|
| 137 |
-
|
| 138 |
-
`
|
| 139 |
-
|
| 140 |
-
### Agreement with the
|
| 141 |
-
|
| 142 |
-
|
|
| 143 |
-
| --- | --- | --- | --- |
|
| 144 |
-
| `eager` | default |
|
| 145 |
-
| `eager` |
|
| 146 |
-
| `
|
| 147 |
-
| `torch_varlen` |
|
| 148 |
-
| `
|
| 149 |
-
| `flash_attn` |
|
| 150 |
-
| `
|
| 151 |
-
| `
|
| 152 |
-
| `
|
| 153 |
-
| `
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
a padded `eager` batch (the same function alone scores 0.99994).
|
| 163 |
|
| 164 |
### Speed and memory
|
| 165 |
|
| 166 |
-
|
|
| 167 |
-
| --- | --- | --- | --- | --- | --- | --- |
|
| 168 |
-
|
|
| 169 |
-
|
|
| 170 |
-
|
|
| 171 |
-
|
|
| 172 |
-
|
|
| 173 |
-
|
|
| 174 |
-
|
|
| 175 |
-
| flattened | 1,024 | 17.0 GiB |
|
| 176 |
-
|
|
| 177 |
-
|
| 178 |
-
`
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 179 |
`model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True)`.
|
| 180 |
|
| 181 |
-
- **
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
[`HuggingFaceTB/stack-edu`](https://huggingface.co/datasets/HuggingFaceTB/stack-edu) (730,431
|
| 199 |
-
tokens after truncation to the model's 8192 `max_seq_length`), on `torch_varlen`, uncompiled.
|
| 200 |
-
**Sorted** is one `encode()` call over all 768, so sentence-transformers length-sorts them;
|
| 201 |
-
**file** is one call per batch of consecutive files, so each batch mixes short and long files.
|
| 202 |
-
Peak memory here includes the weights; **largest** is the largest batch that fit on either card.
|
| 203 |
-
|
| 204 |
-
| load | order | batch size | peak memory | 3090 Ti | 5090 Laptop |
|
| 205 |
-
| --- | --- | --- | --- | --- | --- |
|
| 206 |
-
| default | sorted | 4 | 1.3 GiB | 5.61 s | 4.75 s |
|
| 207 |
-
| default | sorted | 32 | 6.1 GiB | 4.79 s | 4.72 s |
|
| 208 |
-
| default | sorted | 64 | 9.1 GiB | 4.80 s | 4.84 s |
|
| 209 |
-
| default | sorted | 353 (largest) | 18.2 GiB | 5.77 s | 5.98 s |
|
| 210 |
-
| default | file | 4 | 0.6 GiB | 6.49 s | 5.64 s |
|
| 211 |
-
| default | file | 32 | 1.6 GiB | 6.39 s | 5.70 s |
|
| 212 |
-
| default | file | 64 | 2.8 GiB | 6.44 s | 5.84 s |
|
| 213 |
-
| flattened | sorted | 4 | 0.8 GiB | 5.51 s | 4.56 s |
|
| 214 |
-
| flattened | sorted | 32 | 3.8 GiB | 4.54 s | 4.56 s |
|
| 215 |
-
| flattened | sorted | 64 | 6.1 GiB | 4.47 s | 4.63 s |
|
| 216 |
-
| flattened | sorted | 628 (largest) | 19.2 GiB | 4.60 s | 4.75 s |
|
| 217 |
-
| flattened | file | 4 | 0.6 GiB | 5.83 s | 4.75 s |
|
| 218 |
-
| flattened | file | 32 | 1.6 GiB | 4.73 s | 4.70 s |
|
| 219 |
-
| flattened | file | 64 | 2.7 GiB | 4.53 s | 4.66 s |
|
| 220 |
-
|
| 221 |
-
Flattening also makes a file's embedding independent of the other files in its batch: between
|
| 222 |
-
the file-order and sorted runs, the minimum cosine over the 768 files is 0.9999998 under the
|
| 223 |
-
flattened load on both cards and both varlen tiers, against as low as 0.99987 under the default
|
| 224 |
-
load.
|
| 225 |
|
| 226 |
## What changed vs the source repo
|
| 227 |
|
|
@@ -229,29 +216,30 @@ load.
|
|
| 229 |
model runs bf16 in any real serving configuration, so the weights are stored bf16 and (via the
|
| 230 |
load fix below) arrive bf16 β which is simply how this model is used, and removes the need for a
|
| 231 |
post-load dtype cast. Parity-neutral; the smaller download is incidental, not the reason.
|
| 232 |
-
2. **`from_pretrained` dtype fix**: the
|
| 233 |
fp32 and `load_state_dict`-ed the checkpoint into fp32 params, **ignoring `torch_dtype`**. The
|
| 234 |
copy here adds the standard transformers dtype resolution (explicit arg β `config.torch_dtype` β
|
| 235 |
checkpoint dtype) so the model loads in its declared dtype.
|
| 236 |
-
3. **Three
|
| 237 |
-
|
| 238 |
`torch.nn.attention.varlen.varlen_attn`, no third-party kernel, needs torch β₯ 2.10.0),
|
| 239 |
-
`flash_attn` (the
|
| 240 |
-
builds that have it installed), and `eager` (the original padded attention algorithm,
|
| 241 |
-
unchanged, and the default off CUDA sm_80+; it now converts a raw 2-D `[B, S]` mask
|
| 242 |
-
additive form before adding it, for the case where a `device_map` split hands it a mask
|
| 243 |
-
tier upstream never passes). On both varlen tiers `NomicBertModel.forward` unpads the
|
| 244 |
-
with torch-native `_unpad`/`_pad` helpers (a replacement for
|
| 245 |
-
the hidden states packed `[Ξ£L, 768]` through the embeddings
|
| 246 |
-
the end. Rotary embeddings rotate each packed token by
|
| 247 |
-
the same rotation it gets in the padded batch. On the
|
| 248 |
-
the additive attention mask inline instead of calling
|
| 249 |
-
`get_extended_attention_mask` helper. Set
|
| 250 |
-
to force a tier (raises if it can't
|
| 251 |
-
one-time INFO log line report which tier
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
|
|
|
| 255 |
pinned load fetches the weights from the pinned commit too. It also accepts transformers'
|
| 256 |
`dtype=` alongside the older `torch_dtype=`, and returns the model in eval mode, as
|
| 257 |
transformers' own `from_pretrained` does.
|
|
@@ -269,8 +257,12 @@ derives from Tri Dao's BERT implementation, and `CodeRankEmbed` was trained by t
|
|
| 269 |
|
| 270 |
```bibtex
|
| 271 |
@misc{suresh2025cornstackhighqualitycontrastivedata,
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 275 |
}
|
| 276 |
```
|
|
|
|
| 16 |
# CodeRankEmbed-flash-attn
|
| 17 |
|
| 18 |
A **bf16 quantization of [`nomic-ai/CodeRankEmbed`](https://huggingface.co/nomic-ai/CodeRankEmbed)**
|
| 19 |
+
with a custom `modeling_hf_nomic_bert.py`, shipped in this repo, that **picks one of three
|
| 20 |
+
attention implementations for the hardware it runs on.** It is not a finetune β the weights are the
|
| 21 |
+
original CodeRankEmbed weights cast to bf16 (no further training). Two of the three replace the
|
| 22 |
+
original's padded `O(seqΒ²)` attention with an `O(N)` unpadded path; the third keeps the original's
|
| 23 |
+
algorithm as the correctness reference and universal fallback.
|
| 24 |
|
| 25 |
## Why
|
| 26 |
|
| 27 |
`nomic-ai/CodeRankEmbed` loads through `trust_remote_code`, and its attention path is **eager
|
| 28 |
+
only** β activation memory grows as `batch Γ heads Γ seqΒ²`, which runs out of memory at large
|
| 29 |
+
batches even though the model is only 137M params. This repo adds two attention paths that compute
|
| 30 |
+
the same attention in `O(N)` memory by packing unpadded sequences, so the large batches that the
|
| 31 |
+
eager path cannot fit run comfortably β with **parity embeddings** (no quality change):
|
| 32 |
|
| 33 |
- `torch_varlen` β `torch.nn.attention.varlen.varlen_attn`, shipped in torch itself (no extra
|
| 34 |
package), available from torch **2.10.0** onward.
|
| 35 |
+
- `flash_attn` β the `flash_attn` package's varlen kernel. Kept as a fallback for older torch
|
| 36 |
builds that don't yet have `torch.nn.attention.varlen` but do have `flash_attn` installed. This
|
| 37 |
package is **optional**.
|
| 38 |
|
|
|
|
| 41 |
|
| 42 |
## Behavior
|
| 43 |
|
| 44 |
+
- **Three attention tiers, chosen automatically per device** (override with
|
| 45 |
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager`):
|
| 46 |
1. **`torch_varlen`** β CUDA, compute capability sm_80+ (Ampere or newer), torch **β₯ 2.10.0**. No
|
| 47 |
third-party kernel needed.
|
|
|
|
| 49 |
torch that doesn't yet ship `torch.nn.attention.varlen` (typically an older torch). This
|
| 50 |
dependency is **optional**.
|
| 51 |
3. **`eager`** β everything else: CPU, pre-Ampere GPUs, or neither of the above available. The
|
| 52 |
+
original's padded attention algorithm, unchanged, runs on any host.
|
| 53 |
|
| 54 |
`auto` (the default) prefers `torch_varlen`, then `flash_attn`, then `eager`. Before accepting a
|
| 55 |
varlen tier, `auto` runs one tiny kernel probe per device (a capability check alone can't tell
|
|
|
|
| 57 |
raises a `RuntimeError` (how a missing or unsupported kernel fails), it logs one WARNING naming
|
| 58 |
the tier and the error and steps down to the next tier. An out-of-memory error, any other
|
| 59 |
exception type, or a CUDA error left over from earlier work propagates instead of demoting the
|
| 60 |
+
tier, so a transient failure never pins the device to a slower tier. A forced override (e.g.
|
| 61 |
+
`NOMIC_BERT_ATTN_IMPL=torch_varlen`) is not probed and **raises** `RuntimeError` if that tier's
|
| 62 |
+
precondition doesn't hold β a forced tier never falls back silently. An unrecognized override
|
| 63 |
+
raises `ValueError`.
|
| 64 |
- **See which tier engaged**: `model[0].auto_model.attention_impl` after a forward pass. The
|
| 65 |
modeling file also logs the tier, and why, once at INFO.
|
| 66 |
- **Padding-free on the varlen tiers.** The model unpads the batch once, right after the input
|
|
|
|
| 72 |
batch flattened (`[1, Ξ£L]` token ids with sequence boundaries, no padding at all), and the model
|
| 73 |
returns one `[1, Ξ£L, 768]` output that sentence-transformers pools per sequence. The opt-in needs
|
| 74 |
the `flash_attn` package installed, because transformers checks for it at load and raises
|
| 75 |
+
otherwise, even though the attention itself still runs on whichever tier the model picks. On
|
| 76 |
the `eager` tier the flattened batch is re-padded and runs the padded path.
|
| 77 |
- **Loads bf16 by default.** `flash_attn` and `torch_varlen` both require half precision and the
|
| 78 |
model runs bf16 in any real serving setup, so the weights are stored bf16 and `config.json`
|
| 79 |
+
declares `torch_dtype: bfloat16`. The original's custom `from_pretrained` silently dropped
|
| 80 |
`torch_dtype` and always loaded fp32; the copy in this repo honors it, so the model loads bf16
|
| 81 |
natively, like any normal HF model. Pass `torch_dtype=torch.float32` to load fp32 (note: the
|
| 82 |
stored weights are bf16-precision, so this only widens the dtype, not the precision).
|
|
|
|
| 123 |
|
| 124 |
## Parity & performance
|
| 125 |
|
| 126 |
+
Every number below compares this repo with the original model,
|
| 127 |
+
[`nomic-ai/CodeRankEmbed`](https://huggingface.co/nomic-ai/CodeRankEmbed) (revision `3c4b608`),
|
| 128 |
+
run as published: fp32 weights, its own modeling file, padded eager attention. Both run on the
|
| 129 |
+
same GPU, an RTX 3090 Ti (sm_86). The original's modeling file calls
|
| 130 |
+
`get_extended_attention_mask`, which transformers has deprecated and no longer ships by 5.17, so
|
| 131 |
+
it was measured on transformers 5.11; this repo does not call it and loads on 5.17 too.
|
| 132 |
|
| 133 |
**Corpus.** The first 16,384 `document` fields (real Python functions) of
|
| 134 |
[`lightonai/cornstack`](https://huggingface.co/datasets/lightonai/cornstack)'s Python split, in
|
| 135 |
shard order (`train-00000-of-00682.parquet`, then `train-00001-β¦`): 2,684,878 tokens, 7 to 4,746
|
| 136 |
per function, 164 mean, 92 median. Each row is one `encode()` call over all of them at the stated
|
| 137 |
`batch_size=`, after one untimed warm-up of the same call. Cosine similarity is per function, on
|
| 138 |
+
fp32-renormalized output, against the original's embedding of the same function. Peak memory is
|
| 139 |
+
the encode's peak CUDA allocation above the loaded weights. **Load** is how the model was loaded:
|
| 140 |
+
`default` (as in Usage) or `flattened` (the opt-in above). `torch 2.12.1+cu130`,
|
| 141 |
+
`transformers 5.11.0`, `sentence-transformers 6.1.0`, `flash-attn 2.8.3.post1`.
|
| 142 |
+
|
| 143 |
+
### Agreement with the original model
|
| 144 |
+
|
| 145 |
+
| model | attention | compiled | load | batch size | min cosine | mean cosine |
|
| 146 |
+
| --- | --- | --- | --- | --- | --- | --- |
|
| 147 |
+
| `nomic-ai/CodeRankEmbed` | `eager` | β | default | 4 | 1 (the reference) | 1 (the reference) |
|
| 148 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `eager` | β | default | 4 | 0.99902 | 0.99991 |
|
| 149 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `eager` | β | flattened | 4 | 0.99916 | 0.99991 |
|
| 150 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 256 | 0.99896 | 0.99993 |
|
| 151 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 256 | 0.99896 | 0.99993 |
|
| 152 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `flash_attn` | β | default | 256 | 0.99841 | 0.99993 |
|
| 153 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `flash_attn` | β | flattened | 256 | 0.99841 | 0.99993 |
|
| 154 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 256 | 0.99960 | 0.99995 |
|
| 155 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 256 | 0.99960 | 0.99995 |
|
| 156 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `flash_attn` | β | default | 256 | 0.99946 | 0.99995 |
|
| 157 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `flash_attn` | β | flattened | 256 | 0.99946 | 0.99995 |
|
| 158 |
+
|
| 159 |
+
The original is the reference every other row is measured against, so its cosine is 1 by
|
| 160 |
+
definition, not a measurement. The mean cosine is 0.99991 or higher in every configuration; the
|
| 161 |
+
minimum is the one function, of 16,384, that moves most under bf16 weights and arithmetic. The
|
| 162 |
+
batch size does not change these figures (a varlen tier gives the same five decimals at batch 32
|
| 163 |
+
and 256), so the largest batch that ran in every load is shown; `eager` runs at batch 4 (see
|
| 164 |
+
below). Compiled rows sit closer to the original than uncompiled ones.
|
|
|
|
| 165 |
|
| 166 |
### Speed and memory
|
| 167 |
|
| 168 |
+
| model | attention | compiled | load | batch size | peak memory | time |
|
| 169 |
+
| --- | --- | --- | --- | --- | --- | --- |
|
| 170 |
+
| `nomic-ai/CodeRankEmbed` | `eager` | β | default | 4 | 8.4 GiB | 68.9 s |
|
| 171 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 32 | 2.5 GiB | 16.3 s |
|
| 172 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 256 | 10.9 GiB | 15.1 s |
|
| 173 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 1,024 | out of memory | out of memory |
|
| 174 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 2,048 | out of memory | out of memory |
|
| 175 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 32 | 1.4 GiB | 15.5 s |
|
| 176 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 256 | 6.9 GiB | 13.6 s |
|
| 177 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 1,024 | 17.0 GiB | 13.4 s |
|
| 178 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 2,048 | out of memory | out of memory |
|
| 179 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 32 | 1.6 GiB | 14.0 s |
|
| 180 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 256 | 7.1 GiB | 12.8 s |
|
| 181 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 1,024 | 16.3 GiB | 13.8 s |
|
| 182 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | default | 2,048 | out of memory | out of memory |
|
| 183 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 32 | 0.9 GiB | 13.2 s |
|
| 184 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 256 | 4.5 GiB | 11.4 s |
|
| 185 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 1,024 | 11.1 GiB | 10.9 s |
|
| 186 |
+
| `handwoven8588/CodeRankEmbed-flash-attn` | `torch_varlen` | β | flattened | 2,048 | 16.6 GiB | 10.8 s |
|
| 187 |
+
|
| 188 |
+
This repo's rows are `torch_varlen`; `flash_attn` runs within 1% of it in both time and memory
|
| 189 |
+
at every cell. The original runs at batch 4: its padded fp32 attention needs
|
| 190 |
+
`batch Γ heads Γ seqΒ²` memory, 8.4 GiB at batch 4 for this corpus's longest function. The same
|
| 191 |
+
algorithm in bf16 (this repo's `eager` tier, batch 4) takes 41.2 s in 4.2 GiB; the rest of the
|
| 192 |
+
gain comes from the unpadded path and the batches it allows. Compiled means
|
| 193 |
`model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True)`.
|
| 194 |
|
| 195 |
+
- **Against the original**, the same 16,384 functions encode 5.1Γ faster uncompiled (13.4 s,
|
| 196 |
+
flattened at batch 1,024) and 6.4Γ faster compiled (10.8 s, flattened at batch 2,048), against
|
| 197 |
+
68.9 s. At batch 32 this repo needs 2.5 GiB (default load) or 1.4 GiB (flattened), against the
|
| 198 |
+
original's 8.4 GiB at batch 4.
|
| 199 |
+
- **Memory follows the real tokens in the heaviest batch, not the number of functions.** This
|
| 200 |
+
repo runs the encoder on real tokens only, and every out-of-memory cell failed allocating one
|
| 201 |
+
activation over that batch's tokens: the 3,072-wide MLP hidden state, or, compiled, the
|
| 202 |
+
2,304-wide QKV projection. sentence-transformers length-sorts the corpus, so the default load's
|
| 203 |
+
first batch of 1,024 holds the longest functions, about 870,000 tokens; flattened, it pairs the
|
| 204 |
+
longest with the shortest, so its heaviest batch of 1,024 holds about 594,000.
|
| 205 |
+
- **Flattening** cuts peak memory by 32β42% against the default load at the same batch size, and
|
| 206 |
+
wall time by 5β21% (more at larger batches).
|
| 207 |
+
- **Compiling** is 1.17β1.23Γ faster at the same batch size and cuts peak memory by 35%. It
|
| 208 |
+
compiles one graph, with no graph breaks, once per process (the first encode takes about 5 s
|
| 209 |
+
longer), and does not recompile as batch lengths change; the only further compiles come once a
|
| 210 |
+
batch passes about 699,000 and 932,000 tokens, where inductor switches to 64-bit indexing for
|
| 211 |
+
the 3,072-wide MLP activation and the 2,304-wide QKV projection.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
|
| 213 |
## What changed vs the source repo
|
| 214 |
|
|
|
|
| 216 |
model runs bf16 in any real serving configuration, so the weights are stored bf16 and (via the
|
| 217 |
load fix below) arrive bf16 β which is simply how this model is used, and removes the need for a
|
| 218 |
post-load dtype cast. Parity-neutral; the smaller download is incidental, not the reason.
|
| 219 |
+
2. **`from_pretrained` dtype fix**: the original's custom `from_pretrained` instantiated the model
|
| 220 |
fp32 and `load_state_dict`-ed the checkpoint into fp32 params, **ignoring `torch_dtype`**. The
|
| 221 |
copy here adds the standard transformers dtype resolution (explicit arg β `config.torch_dtype` β
|
| 222 |
checkpoint dtype) so the model loads in its declared dtype.
|
| 223 |
+
3. **Three attention tiers**: `NomicBertAttention.forward` now selects one of three attention
|
| 224 |
+
implementations at call time β `torch_varlen` (torch's own
|
| 225 |
`torch.nn.attention.varlen.varlen_attn`, no third-party kernel, needs torch β₯ 2.10.0),
|
| 226 |
+
`flash_attn` (the `flash_attn` package's varlen kernel, kept as a fallback for older torch
|
| 227 |
+
builds that have it installed), and `eager` (the original's padded attention algorithm,
|
| 228 |
+
numerically unchanged, and the default off CUDA sm_80+; it now converts a raw 2-D `[B, S]` mask
|
| 229 |
+
to the additive form before adding it, for the case where a `device_map` split hands it a mask
|
| 230 |
+
a varlen tier upstream never passes). On both varlen tiers `NomicBertModel.forward` unpads the
|
| 231 |
+
batch once with torch-native `_unpad`/`_pad` helpers (a replacement for
|
| 232 |
+
`flash_attn.bert_padding`), keeps the hidden states packed `[Ξ£L, 768]` through the embeddings
|
| 233 |
+
and every layer, and pads back once at the end. Rotary embeddings rotate each packed token by
|
| 234 |
+
its position within its own sequence, the same rotation it gets in the padded batch. On the
|
| 235 |
+
eager tier `NomicBertModel.forward` builds the additive attention mask inline instead of calling
|
| 236 |
+
the (now-removed-upstream) `get_extended_attention_mask` helper. Set
|
| 237 |
+
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager` to force a tier (raises if it can't
|
| 238 |
+
engage); `model[0].auto_model.attention_impl` and the one-time INFO log line report which tier
|
| 239 |
+
engaged.
|
| 240 |
+
4. **`from_pretrained` forwards the Hub revision**: the original's custom `from_pretrained`
|
| 241 |
+
fetched the weight file without the caller's `revision` (and cache/token options), so a pinned
|
| 242 |
+
load still took the weights from the default branch. The copy here passes them through, so a
|
| 243 |
pinned load fetches the weights from the pinned commit too. It also accepts transformers'
|
| 244 |
`dtype=` alongside the older `torch_dtype=`, and returns the model in eval mode, as
|
| 245 |
transformers' own `from_pretrained` does.
|
|
|
|
| 257 |
|
| 258 |
```bibtex
|
| 259 |
@misc{suresh2025cornstackhighqualitycontrastivedata,
|
| 260 |
+
title={CoRNStack: High-Quality Contrastive Data for Better Code Retrieval and Reranking},
|
| 261 |
+
author={Tarun Suresh and Revanth Gangi Reddy and Yifei Xu and Zach Nussbaum and Andriy Mulyar and Brandon Duderstadt and Heng Ji},
|
| 262 |
+
year={2025},
|
| 263 |
+
eprint={2412.01007},
|
| 264 |
+
archivePrefix={arXiv},
|
| 265 |
+
primaryClass={cs.CL},
|
| 266 |
+
url={https://arxiv.org/abs/2412.01007},
|
| 267 |
}
|
| 268 |
```
|