Padding-free encoder, flattening opt-in, torch.compile recipe; card measured on public corpora

#1
Files changed (3) hide show
  1. .gitattributes +3 -0
  2. README.md +148 -80
  3. modeling_hf_nomic_bert.py +178 -5
.gitattributes CHANGED
@@ -33,3 +33,6 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ assets/encode-time.png filter=lfs diff=lfs merge=lfs -text
37
+ assets/encode-memory.png filter=lfs diff=lfs merge=lfs -text
38
+ assets/parity.png filter=lfs diff=lfs merge=lfs -text
README.md CHANGED
@@ -60,13 +60,19 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
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, or the
64
- one-time `NomicBert attention impl=... device=... capability=... torch=... flash_attn=...
65
- override=...` INFO log line (one line per distinct `(impl, device)`). When `flash_attn` is
66
- unavailable the line names why (`flash_attn=absent(<error>)`); an installed `flash_attn` that
67
- fails to import also logs a WARNING.
68
- - **`revision=` pins everything.** Loading with `revision=<commit>` fetches the code, config,
69
- tokenizer **and weights** from that commit.
 
 
 
 
 
 
70
  - **Loads bf16 by default.** `flash_attn` and `torch_varlen` both require half precision and the
71
  model runs bf16 in any real serving setup, so the weights are stored bf16 and `config.json`
72
  declares `torch_dtype: bfloat16`. The upstream custom `from_pretrained` silently dropped
@@ -93,73 +99,129 @@ q = model.encode(queries, normalize_embeddings=True)
93
  d = model.encode(codes, normalize_embeddings=True)
94
  ```
95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
96
  ## Parity & performance
97
 
98
  The weights are the original CodeRankEmbed weights (bf16-cast), so embeddings match the fp32
99
- original to within bf16 precision. The table below measures every tier **on the same GPU**, so
100
- `eager` is the O(seq²) baseline the varlen tiers are compared against.
101
-
102
- **Protocol.** 64 real Python source snippets (the first 40 lines of the first 64 `.py` files,
103
- by path, of the PenStock inference library's source tree; 428 tokens mean, 638 max), encoded
104
- with sentence-transformers at an explicit **batch size of 32 and of 64** (the
105
- `batch_size=` passed to `encode`; with 64 snippets, 64 is one forward). Each tier encodes once
106
- untimed, then again for the measurement. Cosine similarity is taken on fp32-renormalized output
107
- against the **pre-dispatch revision of this repo on the same GPU** (its `flash_attn` path).
108
- `flash_attn` runs the same computation as before the dispatch, so its cosine is a parity check.
109
- `torch_varlen` is a different, torch-native kernel, and `eager` a different algorithm, so theirs
110
- are the real signal. `auto` is the default with no override: the probe keeps `torch_varlen`.
111
- **Encode peak** is the encode's peak CUDA allocation above the loaded model's own ~266 MiB of
112
- bf16 weights. Wall time is one encode of all 64 snippets.
113
-
114
- | GPU | batch size | tier | mean cos | min cos | encode peak | wall |
115
- | --- | --- | --- | --- | --- | --- | --- |
116
- | RTX 3090 Ti (sm_86) | 32 | `torch_varlen` | 0.999941 | 0.999856 | 600 MiB | 0.197 s |
117
- | RTX 3090 Ti (sm_86) | 32 | `flash_attn` | 1.000000 | 1.000000 | 600 MiB | 0.197 s |
118
- | RTX 3090 Ti (sm_86) | 32 | `eager` | 0.999912 | 0.999856 | 807 MiB | 0.262 s |
119
- | RTX 3090 Ti (sm_86) | 64 | `torch_varlen` | 0.999941 | 0.999856 | 1201 MiB | 0.223 s |
120
- | RTX 3090 Ti (sm_86) | 64 | `flash_attn` | 1.000000 | 1.000000 | 1201 MiB | 0.218 s |
121
- | RTX 3090 Ti (sm_86) | 64 | `eager` | 0.999910 | 0.999856 | 1614 MiB | 0.302 s |
122
- | RTX 5090 Laptop GPU (sm_120) | 32 | `torch_varlen` | 0.999944 | 0.999919 | 600 MiB | 0.177 s |
123
- | RTX 5090 Laptop GPU (sm_120) | 32 | `flash_attn` | 1.000000 | 1.000000 | 600 MiB | 0.175 s |
124
- | RTX 5090 Laptop GPU (sm_120) | 32 | `eager` | 0.999912 | 0.999870 | 807 MiB | 0.249 s |
125
- | RTX 5090 Laptop GPU (sm_120) | 64 | `torch_varlen` | 0.999944 | 0.999919 | 1201 MiB | 0.202 s |
126
- | RTX 5090 Laptop GPU (sm_120) | 64 | `flash_attn` | 1.000000 | 1.000000 | 1201 MiB | 0.203 s |
127
- | RTX 5090 Laptop GPU (sm_120) | 64 | `eager` | 0.999913 | 0.999834 | 1614 MiB | 0.298 s |
128
-
129
- Cosines rounded to 6 decimal places, VRAM to the nearest MiB. `torch 2.12.1+cu130`,
130
- `transformers 5.11.0`, `sentence-transformers 6.1.0`, `flash-attn 2.8.3` (`2.8.3.post1` on the
131
- 3090 Ti), all measured loading this repo at `1954de4`.
132
-
133
- Eager's cost grows with `batch × heads × seq²`: at these short snippets it needs about a third more
134
- encode memory than the varlen tiers. At longer inputs the gap stops being a percentage.
135
-
136
- **Real documents at the model defaults.** 768 Python files sampled with a fixed seed from one shard
137
- of [`HuggingFaceTB/stack-edu`](https://huggingface.co/datasets/HuggingFaceTB/stack-edu) (730,431
138
- tokens after truncation), encoded with sentence-transformers at this model's default
139
- `max_seq_length` of 8192 and one `encode()` call per row, so the files are length-sorted before
140
- batching. Wall time covers the whole corpus; peak VRAM is the peak CUDA allocation, weights
141
- included. The largest batch row is the largest batch size that fit on each card (the 3090 Ti was
142
- sharing about 1 GiB with other processes).
143
-
144
- | GPU | batch size | tier | wall (768 files) | peak VRAM |
145
- | --- | --- | --- | --- | --- |
146
- | RTX 3090 Ti (sm_86) | 4 | `torch_varlen` | 6.57 s | 1,295 MiB |
147
- | RTX 3090 Ti (sm_86) | 4 | `flash_attn` | 6.58 s | 1,296 MiB |
148
- | RTX 3090 Ti (sm_86) | 4 | `eager` | 14.18 s | 12,961 MiB |
149
- | RTX 3090 Ti (sm_86) | 32 | `torch_varlen` | 6.37 s | 7,973 MiB |
150
- | RTX 3090 Ti (sm_86) | 64 | `torch_varlen` | 7.24 s | 15,659 MiB |
151
- | RTX 3090 Ti (sm_86) | 71 | `torch_varlen` | 7.60 s | 17,340 MiB |
152
- | RTX 5090 Laptop GPU (sm_120) | 4 | `torch_varlen` | 5.54 s | 1,319 MiB |
153
- | RTX 5090 Laptop GPU (sm_120) | 4 | `flash_attn` | 5.57 s | 1,319 MiB |
154
- | RTX 5090 Laptop GPU (sm_120) | 4 | `eager` | 16.27 s | 12,983 MiB |
155
- | RTX 5090 Laptop GPU (sm_120) | 32 | `torch_varlen` | 6.36 s | 7,995 MiB |
156
- | RTX 5090 Laptop GPU (sm_120) | 64 | `torch_varlen` | 7.48 s | 15,681 MiB |
157
- | RTX 5090 Laptop GPU (sm_120) | 73 | `torch_varlen` | 7.83 s | 17,843 MiB |
158
-
159
- At batch 4 `eager` takes 2.2× (3090 Ti) to 2.9× (5090 Laptop) as long as the varlen tiers and needs
160
- about ten times the memory. At batch 32 with the same files in random order, `eager` ran out of
161
- memory on 21 of the 24 batches (3090 Ti). Both varlen tiers are also gated downstream at min
162
- cosine > 0.997 against the fp32 `nomic-ai/CodeRankEmbed` reference.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
163
 
164
  ## What changed vs the source repo
165
 
@@ -178,19 +240,25 @@ cosine > 0.997 against the fp32 `nomic-ai/CodeRankEmbed` reference.
178
  builds that have it installed), and `eager` (the original padded attention algorithm, numerically
179
  unchanged, and the default off CUDA sm_80+; it now converts a raw 2-D `[B, S]` mask to the
180
  additive form before adding it, for the case where a `device_map` split hands it a mask a varlen
181
- tier upstream never passes). Both varlen tiers unpad the input with torch-native `_unpad`/`_pad`
182
- helpers (a replacement for `flash_attn.bert_padding`) before the packed kernel call, then repad
183
- the output; `NomicBertModel.forward` builds the additive attention mask inline instead of calling
184
- the (now-removed-upstream) `get_extended_attention_mask` helper. Rotary embeddings are applied to
185
- the dense `[B, S, 3, H, D]` tensor **before** unpadding — the correctness keystone: applying RoPE
186
- after unpadding would hand each packed position the wrong sequence's rotation. Set
187
- `NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager` to force a tier (raises if it can't engage);
188
- `model[0].auto_model.attention_impl` and the one-time INFO log line report which tier engaged.
 
189
  4. **`from_pretrained` forwards the Hub revision**: the upstream custom `from_pretrained` fetched
190
  the weight file without the caller's `revision` (and cache/token options), so a pinned load
191
  still took the weights from the default branch. The copy here passes them through, so a
192
  pinned load fetches the weights from the pinned commit too. It also accepts transformers'
193
- `dtype=` alongside the older `torch_dtype=`.
 
 
 
 
 
194
 
195
  ## License & attribution
196
 
 
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
66
+ ids, runs the embeddings and every encoder layer on the real tokens only, and pads back once at
67
+ the end. `last_hidden_state` keeps its `[B, S, 768]` shape, with zeros at padding positions, so
68
+ pooling is unchanged.
69
+ - **Flattened batches (opt-in).** Loading with
70
+ `model_kwargs={"attn_implementation": "flash_attention_2"}` lets sentence-transformers send each
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
 
99
  d = model.encode(codes, normalize_embeddings=True)
100
  ```
101
 
102
+ With `flash_attn` installed, opt in to flattened batches:
103
+
104
+ ```python
105
+ model = SentenceTransformer(
106
+ "handwoven8588/CodeRankEmbed-flash-attn",
107
+ trust_remote_code=True,
108
+ model_kwargs={"attn_implementation": "flash_attention_2"},
109
+ )
110
+ ```
111
+
112
+ On a varlen tier the encoder layers run on packed tokens only, so they compile cleanly. For about
113
+ 1.2× the throughput and a third less peak memory (see Speed and memory below):
114
+
115
+ ```python
116
+ import torch
117
+
118
+ model[0].auto_model.encoder = torch.compile(model[0].auto_model.encoder, dynamic=True)
119
+ ```
120
+
121
+ `dynamic=True` matters: the number of packed tokens changes with every batch.
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
 
 
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.
258
+ 5. **Flattened input**: `NomicBertModel.forward` also accepts sentence-transformers' flattened
259
+ batches (`input_ids [1, ΣL]` with per-sequence `position_ids` and `cu_seq_lens_q`), and the
260
+ model declares flash-attention support so that sentence-transformers sends them when loaded
261
+ with `attn_implementation="flash_attention_2"` (see Behavior).
262
 
263
  ## License & attribution
264
 
modeling_hf_nomic_bert.py CHANGED
@@ -12,7 +12,7 @@ import os
12
  import re
13
  from collections import OrderedDict
14
  from functools import partial
15
- from typing import List, Optional, Tuple, Union
16
 
17
  import torch
18
  import torch.nn as nn
@@ -236,6 +236,21 @@ def _pad(x_u: torch.Tensor, indices: torch.Tensor, B: int, S: int) -> torch.Tens
236
  return out.view(B, S, *x_u.shape[1:])
237
 
238
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
239
  def select_attention_impl(device: torch.device) -> str:
240
  """The attention tier for ``device`` under the current ``NOMIC_BERT_ATTN_IMPL``.
241
 
@@ -541,6 +556,10 @@ class NomicBertPreTrainedModel(PreTrainedModel):
541
  supports_gradient_checkpointing = True
542
  _no_split_modules = ["Block"]
543
  _skip_keys_device_placement = "past_key_values"
 
 
 
 
544
 
545
  def __init__(self, config, *inputs, **kwargs):
546
  super().__init__(config)
@@ -598,6 +617,9 @@ class NomicBertPreTrainedModel(PreTrainedModel):
598
  if rotary_scaling_factor:
599
  config.rotary_scaling_factor = rotary_scaling_factor
600
 
 
 
 
601
  if config.n_positions <= 0 and config.rotary_emb_fraction > 0:
602
  config.n_positions = 2048
603
  if num_labels:
@@ -663,6 +685,9 @@ class NomicBertPreTrainedModel(PreTrainedModel):
663
  logger.warning(load_return)
664
  else:
665
  logger.debug(load_return)
 
 
 
666
  return model
667
 
668
  def _set_gradient_checkpointing(self, module, value=False):
@@ -929,6 +954,29 @@ class NomicBertRotaryEmbedding(nn.Module):
929
  return torch.stack((q_rot, k_rot, qkv[:, :, 2]), dim=2)
930
 
931
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
932
  class NomicBertDynamicNTKRotaryEmbedding(NomicBertRotaryEmbedding):
933
  def __init__(self, rotary_scaling_factor, max_position_embeddings, **kwargs):
934
  super().__init__(**kwargs)
@@ -1069,8 +1117,12 @@ class NomicBertAttention(nn.Module):
1069
  is_padded_inputs: Optional[bool] = True,
1070
  cu_seqlens: Optional[torch.Tensor] = None,
1071
  max_seq_len: Optional[int] = None,
 
1072
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
1073
 
 
 
 
1074
  has_layer_past = past_key_value is not None
1075
 
1076
  if has_layer_past:
@@ -1172,6 +1224,23 @@ class NomicBertAttention(nn.Module):
1172
 
1173
  return attn_output
1174
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1175
 
1176
  class NomicBertBlock(nn.Module):
1177
  def __init__(
@@ -1225,6 +1294,7 @@ class NomicBertBlock(nn.Module):
1225
  use_cache: Optional[bool] = False,
1226
  cu_seqlens: Optional[torch.Tensor] = None,
1227
  max_seq_len: Optional[int] = None,
 
1228
  ):
1229
  r"""Pass the input through the encoder layer.
1230
  Args:
@@ -1244,6 +1314,7 @@ class NomicBertBlock(nn.Module):
1244
  is_padded_inputs=is_padded_inputs,
1245
  cu_seqlens=cu_seqlens,
1246
  max_seq_len=max_seq_len,
 
1247
  )
1248
 
1249
  dropped = self.dropout2(hidden_states)
@@ -1260,6 +1331,7 @@ class NomicBertBlock(nn.Module):
1260
  is_padded_inputs=is_padded_inputs,
1261
  cu_seqlens=cu_seqlens,
1262
  max_seq_len=max_seq_len,
 
1263
  )
1264
  hidden_states = self.norm1((self.dropout1(attn_outputs) + hidden_states).to(dtype=self.norm1.weight.dtype))
1265
  mlp_out = self.mlp(hidden_states)
@@ -1287,6 +1359,7 @@ class NomicBertEncoder(nn.Module):
1287
  output_hidden_states: Optional[bool] = None,
1288
  return_dict: Optional[bool] = None,
1289
  is_padded_inputs: Optional[bool] = True,
 
1290
  ):
1291
  """If subset_mask is not None, we only want output for the subset of the sequence.
1292
  This means that we only compute the last layer output for these tokens.
@@ -1314,6 +1387,7 @@ class NomicBertEncoder(nn.Module):
1314
  None,
1315
  None,
1316
  is_padded_inputs,
 
1317
  # if you freeze ANY layers, you need `use_reentrant=False`
1318
  # https://github.com/huggingface/transformers/issues/21381
1319
  # https://discuss.pytorch.org/t/checkpoint-with-no-grad-requiring-inputs-problem/19117/7
@@ -1331,6 +1405,7 @@ class NomicBertEncoder(nn.Module):
1331
  is_padded_inputs,
1332
  output_attentions,
1333
  use_cache,
 
1334
  )
1335
  return hidden_states
1336
 
@@ -1427,16 +1502,26 @@ class NomicBertModel(NomicBertPreTrainedModel):
1427
  token_type_ids=None,
1428
  position_ids=None,
1429
  return_dict=None,
 
 
 
1430
  ):
1431
  if token_type_ids is None:
1432
  token_type_ids = torch.zeros_like(input_ids)
 
 
 
 
 
 
 
 
 
 
 
1433
  hidden_states = self.embeddings(input_ids, position_ids=position_ids, token_type_ids=token_type_ids)
1434
  hidden_states = self.emb_ln(hidden_states)
1435
  hidden_states = self.emb_drop(hidden_states)
1436
-
1437
- impl = select_attention_impl(hidden_states.device)
1438
- self.attention_impl = impl
1439
- _log_attention_impl_once(impl, hidden_states.device)
1440
  if impl == "eager":
1441
  # No varlen tier engaged (CPU, pre-Ampere, forced, or no kernel):
1442
  # build the additive [B, 1, 1, S] mask inline for the eager block.
@@ -1454,6 +1539,94 @@ class NomicBertModel(NomicBertPreTrainedModel):
1454
  pooler_output=pooled_output,
1455
  )
1456
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1457
 
1458
  class NomicBertForPreTraining(NomicBertPreTrainedModel):
1459
  _tied_weights_keys = ["predictions.decoder.bias", "cls.predictions.decoder.weight"]
 
12
  import re
13
  from collections import OrderedDict
14
  from functools import partial
15
+ from typing import NamedTuple, List, Optional, Tuple, Union
16
 
17
  import torch
18
  import torch.nn as nn
 
236
  return out.view(B, S, *x_u.shape[1:])
237
 
238
 
239
+ class _Packed(NamedTuple):
240
+ """Metadata for hidden states kept packed ``[N, D]`` across the encoder.
241
+
242
+ ``cu_seqlens`` / ``max_seqlen`` feed the varlen kernel; ``positions`` (int64,
243
+ ``[N]``) is each token's column in the padded batch, the position RoPE uses on
244
+ the dense path; ``rope_seqlen`` is the padded width ``S``, the length the dense
245
+ path sizes the RoPE cache to.
246
+ """
247
+
248
+ cu_seqlens: torch.Tensor
249
+ max_seqlen: int
250
+ positions: torch.Tensor
251
+ rope_seqlen: int
252
+
253
+
254
  def select_attention_impl(device: torch.device) -> str:
255
  """The attention tier for ``device`` under the current ``NOMIC_BERT_ATTN_IMPL``.
256
 
 
556
  supports_gradient_checkpointing = True
557
  _no_split_modules = ["Block"]
558
  _skip_keys_device_placement = "past_key_values"
559
+ # Pre-packed input (cu_seq_lens_q / position_ids through forward's **kwargs) runs on
560
+ # either varlen tier, so the model can accept a flattened batch.
561
+ _supports_attention_backend = True
562
+ _supports_flash_attn = True
563
 
564
  def __init__(self, config, *inputs, **kwargs):
565
  super().__init__(config)
 
617
  if rotary_scaling_factor:
618
  config.rotary_scaling_factor = rotary_scaling_factor
619
 
620
+ if kwargs.get("attn_implementation") == "flash_attention_2":
621
+ # Opt-in: advertises flash attention so sentence-transformers flattens batches.
622
+ config._attn_implementation = "flash_attention_2"
623
  if config.n_positions <= 0 and config.rotary_emb_fraction > 0:
624
  config.n_positions = 2048
625
  if num_labels:
 
685
  logger.warning(load_return)
686
  else:
687
  logger.debug(load_return)
688
+ # transformers' native from_pretrained returns the model in eval mode; this
689
+ # override must too, or a direct forward runs the config's dropout.
690
+ model.eval()
691
  return model
692
 
693
  def _set_gradient_checkpointing(self, module, value=False):
 
954
  return torch.stack((q_rot, k_rot, qkv[:, :, 2]), dim=2)
955
 
956
 
957
+ def _rotary_packed(rotary: "NomicBertRotaryEmbedding", qkv_u: torch.Tensor, packed: _Packed) -> torch.Tensor:
958
+ """RoPE on packed ``[N, 3, H, D]``: each token rotated by its own column position.
959
+
960
+ Same arithmetic as ``NomicBertRotaryEmbedding.forward`` on the dense
961
+ ``[B, S, 3, H, D]`` tensor, so every real token gets bit-identical q/k. The caller
962
+ sizes the cache first (``NomicBertModel._size_rotary_caches``), outside the encoder:
963
+ the update branches on the module's cached length, which ``torch.compile`` guards as
964
+ a static int, so updating it here would recompile a compiled encoder on every growth.
965
+ """
966
+ pattern = "... d -> ... 1 (2 d)" if not rotary.interleaved else "... d -> ... 1 (d 2)"
967
+ cos = repeat(rotary._cos_cached[packed.positions], pattern)
968
+ sin = repeat(rotary._sin_cached[packed.positions], pattern)
969
+ ro_dim = cos.shape[-1]
970
+
971
+ def rot(x):
972
+ return torch.cat(
973
+ [x[..., :ro_dim] * cos + rotate_half(x[..., :ro_dim], rotary.interleaved) * sin, x[..., ro_dim:]],
974
+ dim=-1,
975
+ )
976
+
977
+ return torch.stack((rot(qkv_u[:, 0]), rot(qkv_u[:, 1]), qkv_u[:, 2]), dim=1)
978
+
979
+
980
  class NomicBertDynamicNTKRotaryEmbedding(NomicBertRotaryEmbedding):
981
  def __init__(self, rotary_scaling_factor, max_position_embeddings, **kwargs):
982
  super().__init__(**kwargs)
 
1117
  is_padded_inputs: Optional[bool] = True,
1118
  cu_seqlens: Optional[torch.Tensor] = None,
1119
  max_seq_len: Optional[int] = None,
1120
+ packed: Optional[_Packed] = None,
1121
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
1122
 
1123
+ if packed is not None:
1124
+ return self._forward_packed(hidden_states, packed)
1125
+
1126
  has_layer_past = past_key_value is not None
1127
 
1128
  if has_layer_past:
 
1224
 
1225
  return attn_output
1226
 
1227
+ def _forward_packed(self, hidden_states: torch.Tensor, packed: _Packed) -> torch.Tensor:
1228
+ """Varlen attention on packed ``[N, D]`` hidden states; returns ``[N, D]``."""
1229
+ qkv = self.Wqkv(hidden_states)
1230
+ qkv = rearrange(qkv, "n (three h d) -> n three h d", three=3, d=self.head_dim)
1231
+ if self.rotary_emb_dim > 0:
1232
+ qkv = _rotary_packed(self.rotary_emb, qkv, packed)
1233
+ orig_dtype = qkv.dtype
1234
+ if orig_dtype not in (torch.float16, torch.bfloat16):
1235
+ qkv = qkv.to(torch.bfloat16)
1236
+ cu, max_s = packed.cu_seqlens, packed.max_seqlen
1237
+ if select_attention_impl(hidden_states.device) == "torch_varlen":
1238
+ q, k, v = qkv.unbind(1)
1239
+ out = _torch_varlen_attn(q, k, v, cu, cu, max_s, max_s)
1240
+ else:
1241
+ out = flash_attn_varlen_qkvpacked_func(qkv, cu, max_s, dropout_p=0.0, causal=False)
1242
+ return self.out_proj(rearrange(out.to(orig_dtype), "n h d -> n (h d)"))
1243
+
1244
 
1245
  class NomicBertBlock(nn.Module):
1246
  def __init__(
 
1294
  use_cache: Optional[bool] = False,
1295
  cu_seqlens: Optional[torch.Tensor] = None,
1296
  max_seq_len: Optional[int] = None,
1297
+ packed: Optional[_Packed] = None,
1298
  ):
1299
  r"""Pass the input through the encoder layer.
1300
  Args:
 
1314
  is_padded_inputs=is_padded_inputs,
1315
  cu_seqlens=cu_seqlens,
1316
  max_seq_len=max_seq_len,
1317
+ packed=packed,
1318
  )
1319
 
1320
  dropped = self.dropout2(hidden_states)
 
1331
  is_padded_inputs=is_padded_inputs,
1332
  cu_seqlens=cu_seqlens,
1333
  max_seq_len=max_seq_len,
1334
+ packed=packed,
1335
  )
1336
  hidden_states = self.norm1((self.dropout1(attn_outputs) + hidden_states).to(dtype=self.norm1.weight.dtype))
1337
  mlp_out = self.mlp(hidden_states)
 
1359
  output_hidden_states: Optional[bool] = None,
1360
  return_dict: Optional[bool] = None,
1361
  is_padded_inputs: Optional[bool] = True,
1362
+ packed: Optional[_Packed] = None,
1363
  ):
1364
  """If subset_mask is not None, we only want output for the subset of the sequence.
1365
  This means that we only compute the last layer output for these tokens.
 
1387
  None,
1388
  None,
1389
  is_padded_inputs,
1390
+ packed=packed,
1391
  # if you freeze ANY layers, you need `use_reentrant=False`
1392
  # https://github.com/huggingface/transformers/issues/21381
1393
  # https://discuss.pytorch.org/t/checkpoint-with-no-grad-requiring-inputs-problem/19117/7
 
1405
  is_padded_inputs,
1406
  output_attentions,
1407
  use_cache,
1408
+ packed=packed,
1409
  )
1410
  return hidden_states
1411
 
 
1502
  token_type_ids=None,
1503
  position_ids=None,
1504
  return_dict=None,
1505
+ cu_seq_lens_q=None,
1506
+ max_length_q=None,
1507
+ **kwargs,
1508
  ):
1509
  if token_type_ids is None:
1510
  token_type_ids = torch.zeros_like(input_ids)
1511
+ if cu_seq_lens_q is not None:
1512
+ self.attention_impl = select_attention_impl(input_ids.device)
1513
+ _log_attention_impl_once(self.attention_impl, input_ids.device)
1514
+ return self._forward_prepacked(input_ids, token_type_ids, position_ids, cu_seq_lens_q, max_length_q)
1515
+
1516
+ impl = select_attention_impl(input_ids.device)
1517
+ self.attention_impl = impl
1518
+ _log_attention_impl_once(impl, input_ids.device)
1519
+ if impl != "eager" and (attention_mask is None or attention_mask.ndim == 2):
1520
+ return self._forward_packed(input_ids, attention_mask, token_type_ids, position_ids)
1521
+
1522
  hidden_states = self.embeddings(input_ids, position_ids=position_ids, token_type_ids=token_type_ids)
1523
  hidden_states = self.emb_ln(hidden_states)
1524
  hidden_states = self.emb_drop(hidden_states)
 
 
 
 
1525
  if impl == "eager":
1526
  # No varlen tier engaged (CPU, pre-Ampere, forced, or no kernel):
1527
  # build the additive [B, 1, 1, S] mask inline for the eager block.
 
1539
  pooler_output=pooled_output,
1540
  )
1541
 
1542
+ def _size_rotary_caches(self, seqlen, device):
1543
+ """Grow each layer's RoPE cos/sin cache to ``seqlen`` before the packed encoder runs
1544
+ (see ``_rotary_packed``), in the dtype its q/k projection produces."""
1545
+ for layer in self.encoder.layers:
1546
+ attn = layer.attn
1547
+ if attn.rotary_emb_dim > 0:
1548
+ attn.rotary_emb._update_cos_sin_cache(seqlen, device=device, dtype=attn.Wqkv.weight.dtype)
1549
+
1550
+ def _forward_packed(self, input_ids, attention_mask, token_type_ids, position_ids):
1551
+ """Varlen tiers: unpad once, run embeddings and every layer on real tokens only.
1552
+
1553
+ ``last_hidden_state`` comes back padded ``[B, S, D]`` with zeros at padding
1554
+ positions, so masked pooling downstream is unchanged.
1555
+ """
1556
+ B, S = input_ids.shape
1557
+ mask = (
1558
+ torch.ones(B, S, device=input_ids.device, dtype=torch.bool)
1559
+ if attention_mask is None
1560
+ else attention_mask.to(torch.bool)
1561
+ )
1562
+ seqlens = mask.sum(-1, dtype=torch.int32)
1563
+ max_s = int(seqlens.max())
1564
+ if max_s == 0:
1565
+ raise RuntimeError(
1566
+ "NomicBertModel.forward: every sequence is empty (max_s == 0); "
1567
+ "pre-filter empty inputs before encoding."
1568
+ )
1569
+ indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten()
1570
+ cu = F.pad(torch.cumsum(seqlens, 0, dtype=torch.int32), (1, 0))
1571
+ columns = indices % S
1572
+ pos_u = position_ids.flatten()[indices] if position_ids is not None else columns
1573
+ hidden_states = self.embeddings(
1574
+ input_ids.flatten()[indices][None],
1575
+ position_ids=pos_u[None],
1576
+ token_type_ids=token_type_ids.flatten()[indices][None],
1577
+ )[0]
1578
+ hidden_states = self.emb_drop(self.emb_ln(hidden_states))
1579
+ packed = _Packed(cu_seqlens=cu, max_seqlen=max_s, positions=columns, rope_seqlen=S)
1580
+ self._size_rotary_caches(packed.rope_seqlen, hidden_states.device)
1581
+ hidden_states = self.encoder(hidden_states, packed=packed)
1582
+ sequence_output = _pad(hidden_states, indices, B, S)
1583
+ pooled_output = self.pooler(sequence_output) if self.pooler is not None else None
1584
+ return BaseModelOutputWithPoolingAndCrossAttentions(
1585
+ last_hidden_state=sequence_output,
1586
+ pooler_output=pooled_output,
1587
+ )
1588
+
1589
+ def _forward_prepacked(self, input_ids, token_type_ids, position_ids, cu_seq_lens_q, max_length_q):
1590
+ """Input already packed by the caller: ``input_ids [1, ΣL]`` with per-sequence
1591
+ ``position_ids`` and ``cu_seq_lens_q`` boundaries. Returns ``[1, ΣL, D]``."""
1592
+ if self.attention_impl == "eager":
1593
+ return self._forward_prepacked_eager(input_ids, token_type_ids, cu_seq_lens_q)
1594
+ if position_ids is None:
1595
+ raise ValueError("NomicBertModel.forward: pre-packed input needs per-sequence position_ids")
1596
+ positions = position_ids.reshape(-1).long()
1597
+ max_s = int(max_length_q) if max_length_q is not None else int((cu_seq_lens_q[1:] - cu_seq_lens_q[:-1]).max())
1598
+ hidden_states = self.embeddings(input_ids, position_ids=position_ids, token_type_ids=token_type_ids)[0]
1599
+ hidden_states = self.emb_drop(self.emb_ln(hidden_states))
1600
+ packed = _Packed(cu_seqlens=cu_seq_lens_q.to(torch.int32), max_seqlen=max_s, positions=positions,
1601
+ rope_seqlen=max_s)
1602
+ self._size_rotary_caches(packed.rope_seqlen, hidden_states.device)
1603
+ sequence_output = self.encoder(hidden_states, packed=packed)[None]
1604
+ pooled_output = self.pooler(sequence_output) if self.pooler is not None else None
1605
+ return BaseModelOutputWithPoolingAndCrossAttentions(
1606
+ last_hidden_state=sequence_output,
1607
+ pooler_output=pooled_output,
1608
+ )
1609
+
1610
+ def _forward_prepacked_eager(self, input_ids, token_type_ids, cu_seq_lens_q):
1611
+ """Pre-packed input on the eager tier (forced, or no varlen kernel on this device):
1612
+ re-pad to ``[B, S]`` (right padding, so each token keeps its per-sequence
1613
+ position), run the padded eager forward, and gather the real tokens back to
1614
+ ``[1, ΣL, D]``, the shape the caller's pooling expects."""
1615
+ cu = cu_seq_lens_q.long()
1616
+ seqlens = cu[1:] - cu[:-1]
1617
+ B, S = seqlens.shape[0], int(seqlens.max())
1618
+ mask = torch.arange(S, device=input_ids.device)[None, :] < seqlens[:, None]
1619
+ indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten()
1620
+ ids = _pad(input_ids.reshape(-1), indices, B, S)
1621
+ types = _pad(token_type_ids.reshape(-1), indices, B, S)
1622
+ out = self.forward(ids, attention_mask=mask.long(), token_type_ids=types)
1623
+ sequence_output = out.last_hidden_state.flatten(0, 1)[indices][None]
1624
+ pooled_output = self.pooler(sequence_output) if self.pooler is not None else None
1625
+ return BaseModelOutputWithPoolingAndCrossAttentions(
1626
+ last_hidden_state=sequence_output,
1627
+ pooler_output=pooled_output,
1628
+ )
1629
+
1630
 
1631
  class NomicBertForPreTraining(NomicBertPreTrainedModel):
1632
  _tied_weights_keys = ["predictions.decoder.bias", "cls.predictions.decoder.weight"]