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
Padding-free encoder, flattening opt-in, torch.compile recipe; card measured on public corpora
#1
by handwoven8588 - opened
- .gitattributes +3 -0
- README.md +148 -80
- 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
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
|
| 116 |
-
|
|
| 117 |
-
|
|
| 118 |
-
|
|
| 119 |
-
|
|
| 120 |
-
|
|
| 121 |
-
|
|
| 122 |
-
|
|
| 123 |
-
|
|
| 124 |
-
|
|
| 125 |
-
|
|
| 126 |
-
|
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
|
| 145 |
-
|
|
| 146 |
-
|
|
| 147 |
-
|
|
| 148 |
-
|
|
| 149 |
-
|
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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).
|
| 182 |
-
helpers (a replacement for `flash_attn.bert_padding`)
|
| 183 |
-
the
|
| 184 |
-
the
|
| 185 |
-
the
|
| 186 |
-
|
| 187 |
-
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager`
|
| 188 |
-
`model[0].auto_model.attention_impl` and the
|
|
|
|
| 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"]
|