Card: every comparison against nomic-ai/CodeRankEmbed; plain wording; citation fixed

#2
Files changed (1) hide show
  1. README.md +120 -128
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 **three-tier attention dispatch built into a custom `modeling_hf_nomic_bert.py` shipped in
20
- this repo.** It is not a finetune β€” the weights are the original CodeRankEmbed weights cast to bf16
21
- (no further training). Two of the three tiers replace the original eager `O(seqΒ²)` attention with an
22
- `O(N)` unpadded path; the third keeps the original eager algorithm as the correctness reference and
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 OOMs at large batches even though
29
- the model is only 137M params. This repo adds two attention paths that compute the same attention in
30
- `O(N)` memory by packing unpadded sequences, so the large batches that OOM the eager path run
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 original `flash_attn` varlen-packed 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,7 +41,7 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
41
 
42
  ## Behavior
43
 
44
- - **Three-tier attention, 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,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. `NOMIC_BERT_ATTN_IMPL=torch_varlen`) is not probed and **raises**
61
- `RuntimeError` if that tier's precondition doesn't hold β€” a forced tier never falls back
62
- silently. An unrecognized override raises `ValueError`.
 
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 dispatch picks. On
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 upstream custom `from_pretrained` silently dropped
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
- The weights are the original CodeRankEmbed weights (bf16-cast), so embeddings match the fp32
126
- original to within bf16 precision. Every number below is measured on two cards, an RTX 3090 Ti
127
- (sm_86) and an RTX 5090 Laptop GPU (sm_120), against the **pre-dispatch revision of this repo
128
- (`e361c6f`) on the same GPU**, which runs the original `flash_attn` path on padded batches.
 
 
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 one batch-32 encode by the pre-dispatch revision. Peak memory is
136
- the encode's peak CUDA allocation above the loaded weights (~266 MiB), the same on both cards to
137
- within 32 MiB. `torch 2.12.1+cu130`, `transformers 5.11.0`, `sentence-transformers 6.1.0`,
138
- `flash-attn 2.8.3` (`2.8.3.post1` on the 3090 Ti).
139
-
140
- ### Agreement with the pre-dispatch revision
141
-
142
- | tier | load | 3090 Ti min / mean | 5090 Laptop min / mean |
143
- | --- | --- | --- | --- |
144
- | `eager` | default | 0.99940 / 0.99992 | 0.99972 / 0.99992 |
145
- | `eager` | flattened | 0.99940 / 0.99992 | 0.99816 / 0.99992 |
146
- | `torch_varlen` | default | 0.99940 / 0.99995 | 0.99975 / 0.99995 |
147
- | `torch_varlen` | flattened | 0.99940 / 0.99995 | 0.99975 / 0.99995 |
148
- | `flash_attn` | default | 0.99978 / 0.99998 | 0.99981 / 0.99999 |
149
- | `flash_attn` | flattened | 0.99978 / 0.99998 | 0.99981 / 0.99999 |
150
- | `torch_varlen`, compiled | default | 0.99816 / 0.99994 | 0.99866 / 0.99994 |
151
- | `torch_varlen`, compiled | flattened | 0.99816 / 0.99994 | 0.99925 / 0.99994 |
152
- | `flash_attn`, compiled | default | 0.99818 / 0.99994 | 0.99876 / 0.99994 |
153
- | `flash_attn`, compiled | flattened | 0.99818 / 0.99994 | 0.99876 / 0.99994 |
154
-
155
- `eager` at batch 4, the others at batch 256. `torch_varlen` is a different, torch-native kernel and
156
- `eager` a different algorithm; `flash_attn` is the same kernel, and where it sits below 1 the rest
157
- of the network is running on packed tokens, so its matrix shapes and bf16 rounding differ.
158
- Compiling reassociates bf16 arithmetic, which leaves the mean unchanged and moves a few functions
159
- further. `auto`, the default with no override, kept `torch_varlen` on both cards. The flattened
160
- `eager` minimum on the 5090 Laptop is one 2,166-token function padded next to a 2,195-token one:
161
- when it flattens, sentence-transformers interleaves long and short texts, which changes who shares
162
- a padded `eager` batch (the same function alone scores 0.99994).
163
 
164
  ### Speed and memory
165
 
166
- | load | batch size | peak memory, eager | peak memory, compiled | 3090 Ti eager | 5090 Laptop eager | 3090 Ti compiled | 5090 Laptop compiled |
167
- | --- | --- | --- | --- | --- | --- | --- | --- |
168
- | pre-dispatch revision | 32 | 4.4 GiB | – | 23.8 s | 21.4 s | – | – |
169
- | default | 32 | 2.5 GiB | 1.6 GiB | 16.4 s | 14.3 s | 14.0 s | 12.2 s |
170
- | default | 256 | 10.9 GiB | 7.1 GiB | 15.1 s | 15.1 s | 12.7 s | 12.3 s |
171
- | default | 1,024 | out of memory | 16.3 GiB | out of memory | out of memory | 13.8 s | 13.4 s |
172
- | default | 2,048 | out of memory | out of memory | out of memory | out of memory | out of memory | out of memory |
173
- | flattened | 32 | 1.4 GiB | 0.9 GiB | 15.6 s | 13.7 s | 13.2 s | 11.7 s |
174
- | flattened | 256 | 6.9 GiB | 4.5 GiB | 13.7 s | 14.4 s | 11.4 s | 11.7 s |
175
- | flattened | 1,024 | 17.0 GiB | 11.1 GiB | 13.4 s | 14.7 s | 10.9 s | 11.8 s |
176
- | flattened | 2,048 | out of memory | 16.6 GiB | out of memory | out of memory | 10.7 s | 11.4 s |
177
-
178
- `torch_varlen`; `flash_attn` runs within 1% of it. Compiled means
 
 
 
 
 
 
 
 
 
 
 
 
179
  `model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True)`.
180
 
181
- - **Memory follows the real tokens in the heaviest batch, not the number of functions.** Both
182
- loads run the encoder padding-free, and every out-of-memory cell failed allocating one 3,072-wide
183
- MLP activation over that batch's tokens. sentence-transformers length-sorts the corpus, so the
184
- default load's first batch of 1,024 holds the longest functions, 869,948 tokens. Flattened, it
185
- interleaves longest with shortest, so its heaviest batch of 1,024 holds 593,716 tokens.
186
- - **Flattening** cuts peak memory by 36–42% against the default load and wall time by up to 9%.
187
- - **Compiling** is 1.17–1.24Γ— faster at the same batch size and cuts peak memory by about a third,
188
- because inductor fuses the element-wise work around each attention call. It compiles one graph,
189
- with no graph breaks, once per process (the first encode takes 4–10 s longer), and does not
190
- recompile as batch lengths change; the one further compile comes once a batch passes about
191
- 699,000 tokens, where inductor switches to 64-bit indexing.
192
- - **Together**, flattened and compiled at batch 2,048 encode the corpus in 10.7 s (3090 Ti) and
193
- 11.4 s (5090 Laptop), 2.2Γ— and 1.9Γ— the pre-dispatch revision. `eager` at batch 4 takes 34–52 s.
194
-
195
- ### Real documents
196
-
197
- 768 Python files sampled with a fixed seed from one shard of
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 upstream custom `from_pretrained` instantiated the model
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-tier attention dispatch**: `NomicBertAttention.forward` now selects one of three
237
- attention implementations at call time β€” `torch_varlen` (torch's own
238
  `torch.nn.attention.varlen.varlen_attn`, no third-party kernel, needs torch β‰₯ 2.10.0),
239
- `flash_attn` (the original flash-attn varlen-packed kernel, kept as a fallback for older torch
240
- builds that have it installed), and `eager` (the original padded attention algorithm, numerically
241
- unchanged, and the default off CUDA sm_80+; it now converts a raw 2-D `[B, S]` mask to the
242
- additive form before adding it, for the case where a `device_map` split hands it a mask a varlen
243
- tier upstream never passes). On both varlen tiers `NomicBertModel.forward` unpads the batch once
244
- with torch-native `_unpad`/`_pad` helpers (a replacement for `flash_attn.bert_padding`), keeps
245
- the hidden states packed `[Ξ£L, 768]` through the embeddings and every layer, and pads back once at
246
- the end. Rotary embeddings rotate each packed token by its position within its own sequence,
247
- the same rotation it gets in the padded batch. On the eager tier `NomicBertModel.forward` builds
248
- the additive attention mask inline instead of calling the (now-removed-upstream)
249
- `get_extended_attention_mask` helper. Set `NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager`
250
- to force a tier (raises if it can't engage); `model[0].auto_model.attention_impl` and the
251
- one-time INFO log line report which tier engaged.
252
- 4. **`from_pretrained` forwards the Hub revision**: the upstream custom `from_pretrained` fetched
253
- the weight file without the caller's `revision` (and cache/token options), so a pinned load
254
- still took the weights from the default branch. The copy here passes them through, so a
 
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
- title = {CoRNStack: High-Quality Contrastive Data for Text and Code Retrieval},
273
- author = {Suresh, K N Q and Wang, Xiang and Khan, Saqib and others},
274
- year = {2025},
 
 
 
 
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
  ```