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
Pin weights to the requested revision; harden the varlen probe; card: all tiers measured on one GPU with stated batch sizes
Browse files- README.md +51 -30
- modeling_hf_nomic_bert.py +112 -47
README.md
CHANGED
|
@@ -54,13 +54,19 @@ file ships all three paths itself, so no runtime patching or post-load hooks are
|
|
| 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
|
| 56 |
whether the installed build has a kernel for the GPU's architecture, e.g. on ROCm); if the probe
|
| 57 |
-
raises
|
| 58 |
-
|
|
|
|
|
|
|
| 59 |
`RuntimeError` if that tier's precondition doesn't hold — a forced tier never falls back
|
| 60 |
silently. An unrecognized override raises `ValueError`.
|
| 61 |
- **See which tier engaged**: `model[0].auto_model.attention_impl` after a forward pass, or the
|
| 62 |
one-time `NomicBert attention impl=... device=... capability=... torch=... flash_attn=...
|
| 63 |
-
override=...` INFO log line (one line per distinct `(impl, device)`).
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
- **Loads bf16 by default.** `flash_attn` and `torch_varlen` both require half precision and the
|
| 65 |
model runs bf16 in any real serving setup, so the weights are stored bf16 and `config.json`
|
| 66 |
declares `torch_dtype: bfloat16`. The upstream custom `from_pretrained` silently dropped
|
|
@@ -90,34 +96,44 @@ d = model.encode(codes, normalize_embeddings=True)
|
|
| 90 |
## Parity & performance
|
| 91 |
|
| 92 |
The weights are the original CodeRankEmbed weights (bf16-cast), so embeddings match the fp32
|
| 93 |
-
original to within bf16 precision. The table below
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
| --- | --- | --- | --- | --- | --- | --- |
|
| 103 |
-
| RTX 3090 Ti (sm_86) |
|
| 104 |
-
| RTX 3090 Ti (sm_86) | `
|
| 105 |
-
| RTX 3090 Ti (sm_86) | `eager` |
|
| 106 |
-
| RTX
|
| 107 |
-
| RTX
|
| 108 |
-
| RTX
|
| 109 |
-
|
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 121 |
|
| 122 |
## What changed vs the source repo
|
| 123 |
|
|
@@ -144,6 +160,11 @@ at batch size 256.
|
|
| 144 |
after unpadding would hand each packed position the wrong sequence's rotation. Set
|
| 145 |
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager` to force a tier (raises if it can't engage);
|
| 146 |
`model[0].auto_model.attention_impl` and the one-time INFO log line report which tier engaged.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 147 |
|
| 148 |
## License & attribution
|
| 149 |
|
|
|
|
| 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
|
| 56 |
whether the installed build has a kernel for the GPU's architecture, e.g. on ROCm); if the probe
|
| 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, 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
|
|
|
|
| 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 each file; 428 tokens mean,
|
| 103 |
+
638 max), encoded with sentence-transformers at an explicit **batch size of 32 and of 64** (the
|
| 104 |
+
`batch_size=` passed to `encode`; with 64 snippets, 64 is one forward). Each tier encodes once
|
| 105 |
+
untimed, then again for the measurement. Cosine similarity is taken on fp32-renormalized output
|
| 106 |
+
against the **pre-dispatch revision of this repo on the same GPU** (its `flash_attn` path).
|
| 107 |
+
`flash_attn` runs the same computation as before the dispatch, so its cosine is a parity check.
|
| 108 |
+
`torch_varlen` is a different, torch-native kernel, and `eager` a different algorithm, so theirs
|
| 109 |
+
are the real signal. `auto` is the default with no override: the probe keeps `torch_varlen`.
|
| 110 |
+
**Encode peak** is the encode's peak CUDA allocation above the loaded model's own ~266 MiB of
|
| 111 |
+
bf16 weights. Wall time is one encode of all 64 snippets.
|
| 112 |
+
|
| 113 |
+
| GPU | batch size | tier | mean cos | min cos | encode peak | wall |
|
| 114 |
| --- | --- | --- | --- | --- | --- | --- |
|
| 115 |
+
| RTX 3090 Ti (sm_86) | 32 | `torch_varlen` | 0.999941 | 0.999856 | 600 MiB | 0.232 s |
|
| 116 |
+
| RTX 3090 Ti (sm_86) | 32 | `flash_attn` | 1.000000 | 1.000000 | 600 MiB | 0.232 s |
|
| 117 |
+
| RTX 3090 Ti (sm_86) | 32 | `eager` | 0.999912 | 0.999856 | 807 MiB | 0.284 s |
|
| 118 |
+
| RTX 3090 Ti (sm_86) | 64 | `torch_varlen` | 0.999941 | 0.999856 | 1201 MiB | 0.248 s |
|
| 119 |
+
| RTX 3090 Ti (sm_86) | 64 | `flash_attn` | 1.000000 | 1.000000 | 1201 MiB | 0.250 s |
|
| 120 |
+
| RTX 3090 Ti (sm_86) | 64 | `eager` | 0.999910 | 0.999856 | 1614 MiB | 0.316 s |
|
| 121 |
+
| RTX 5090 Laptop GPU (sm_120) | 32 | `torch_varlen` | 0.999944 | 0.999919 | 600 MiB | 0.234 s |
|
| 122 |
+
| RTX 5090 Laptop GPU (sm_120) | 32 | `flash_attn` | 1.000000 | 1.000000 | 600 MiB | 0.236 s |
|
| 123 |
+
| RTX 5090 Laptop GPU (sm_120) | 32 | `eager` | 0.999912 | 0.999870 | 807 MiB | 0.318 s |
|
| 124 |
+
| RTX 5090 Laptop GPU (sm_120) | 64 | `torch_varlen` | 0.999944 | 0.999919 | 1201 MiB | 0.272 s |
|
| 125 |
+
| RTX 5090 Laptop GPU (sm_120) | 64 | `flash_attn` | 1.000000 | 1.000000 | 1201 MiB | 0.271 s |
|
| 126 |
+
| RTX 5090 Laptop GPU (sm_120) | 64 | `eager` | 0.999913 | 0.999834 | 1614 MiB | 0.378 s |
|
| 127 |
+
|
| 128 |
+
Cosines rounded to 6 decimal places, VRAM to the nearest MiB. `torch 2.12.1+cu130`,
|
| 129 |
+
`transformers 5.11.0`, `sentence-transformers 6.1.0`, `flash-attn 2.8.3`.
|
| 130 |
+
|
| 131 |
+
Eager's cost grows with `batch × heads × seq²`: at these short snippets it needs about a third more
|
| 132 |
+
encode memory than the varlen tiers. At longer inputs the gap stops being a percentage. At a
|
| 133 |
+
batch size of **256** inputs of ~300–2048 tokens, `torch_varlen` stays under 20 GB peak VRAM,
|
| 134 |
+
while eager's attention scores alone (`256 × 12 heads × 2048² × 2 bytes`) exceed a 24 GB card.
|
| 135 |
+
Both varlen tiers are also gated downstream at min cosine > 0.997 against the fp32
|
| 136 |
+
`nomic-ai/CodeRankEmbed` reference.
|
| 137 |
|
| 138 |
## What changed vs the source repo
|
| 139 |
|
|
|
|
| 160 |
after unpadding would hand each packed position the wrong sequence's rotation. Set
|
| 161 |
`NOMIC_BERT_ATTN_IMPL=torch_varlen|flash_attn|eager` to force a tier (raises if it can't engage);
|
| 162 |
`model[0].auto_model.attention_impl` and the one-time INFO log line report which tier engaged.
|
| 163 |
+
4. **`from_pretrained` forwards the Hub revision**: the upstream custom `from_pretrained` fetched
|
| 164 |
+
the weight file without the caller's `revision` (and cache/token options), so a pinned load
|
| 165 |
+
still took the weights from the default branch. The copy here passes them through, so a
|
| 166 |
+
pinned load fetches the weights from the pinned commit too. It also accepts transformers'
|
| 167 |
+
`dtype=` alongside the older `torch_dtype=`.
|
| 168 |
|
| 169 |
## License & attribution
|
| 170 |
|
modeling_hf_nomic_bert.py
CHANGED
|
@@ -33,13 +33,13 @@ from .configuration_hf_nomic_bert import NomicBertConfig
|
|
| 33 |
# Three-tier attention dispatch (auto by default, override via NOMIC_BERT_ATTN_IMPL):
|
| 34 |
# 1. torch_varlen — torch.nn.attention.varlen, no third-party kernel
|
| 35 |
# 2. flash_attn — the flash-attn varlen-packed kernel (optional dependency)
|
| 36 |
-
# 3. eager — the original padded
|
| 37 |
# Both tier 1 and tier 2 run the same FA2-family kernel and refuse pre-Ampere
|
| 38 |
# hardware, so a single capability predicate gates both (resolve_attention_impl,
|
| 39 |
# below). In auto mode a tiny kernel probe then confirms the chosen varlen tier
|
| 40 |
# actually runs on the device, stepping down a tier if it raises (_select_cached).
|
| 41 |
-
# When neither is available (CPU-only hosts, or builds without them) tier 3 runs
|
| 42 |
-
#
|
| 43 |
try: # pragma: no cover - import guard, exercised by environment
|
| 44 |
from torch.nn.attention.varlen import varlen_attn as _torch_varlen_attn
|
| 45 |
|
|
@@ -47,13 +47,25 @@ try: # pragma: no cover - import guard, exercised by environment
|
|
| 47 |
except ImportError: # pragma: no cover
|
| 48 |
_torch_varlen_attn = None # type: ignore[assignment]
|
| 49 |
_TORCH_VARLEN_AVAILABLE = False
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
try: # pragma: no cover - import guard, exercised by environment
|
| 51 |
from flash_attn import flash_attn_varlen_qkvpacked_func
|
| 52 |
|
| 53 |
_FLASH_AVAILABLE = True
|
| 54 |
-
except Exception: # pragma: no cover
|
| 55 |
flash_attn_varlen_qkvpacked_func = None # type: ignore[assignment]
|
| 56 |
_FLASH_AVAILABLE = False
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
|
| 58 |
ATTENTION_IMPLS = ("torch_varlen", "flash_attn", "eager")
|
| 59 |
ATTN_IMPL_ENV = "NOMIC_BERT_ATTN_IMPL" # auto (default) | torch_varlen | flash_attn | eager
|
|
@@ -91,7 +103,10 @@ def resolve_attention_impl(
|
|
| 91 |
f"is unavailable (torch {torch.__version__}; needs >= 2.10)"
|
| 92 |
)
|
| 93 |
if override == "flash_attn" and not flash_available:
|
| 94 |
-
raise RuntimeError(
|
|
|
|
|
|
|
|
|
|
| 95 |
if not is_cuda:
|
| 96 |
raise RuntimeError(f"{ATTN_IMPL_ENV}={override} requested but device is cpu (not CUDA)")
|
| 97 |
if capability is None or tuple(capability) < (8, 0):
|
|
@@ -114,7 +129,12 @@ def _probe_varlen_kernel(impl: str, device: torch.device) -> None:
|
|
| 114 |
only a proxy: a torch build can ship ``torch.nn.attention.varlen`` without an
|
| 115 |
FA kernel for the device's arch, and on ROCm ``device.type`` is ``"cuda"`` and
|
| 116 |
the capability is the gfx major.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 117 |
"""
|
|
|
|
| 118 |
qkv = torch.randn(2, 3, 1, 64, device=device, dtype=torch.bfloat16)
|
| 119 |
cu = torch.tensor([0, 2], device=device, dtype=torch.int32)
|
| 120 |
if impl == "torch_varlen":
|
|
@@ -125,23 +145,39 @@ def _probe_varlen_kernel(impl: str, device: torch.device) -> None:
|
|
| 125 |
torch.cuda.synchronize(device)
|
| 126 |
|
| 127 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 128 |
@functools.lru_cache(maxsize=None)
|
| 129 |
def _select_cached(dev_type, dev_index, override):
|
| 130 |
"""Resolve the tier for one device, probing the kernel in auto mode.
|
| 131 |
|
| 132 |
A forced tier is returned exactly as ``resolve_attention_impl`` decides it (or
|
| 133 |
raises), with no probe: it never falls back. In ``auto`` mode a varlen tier is
|
| 134 |
-
accepted only if ``_probe_varlen_kernel`` runs; if the probe raises
|
| 135 |
-
|
|
|
|
|
|
|
| 136 |
unavailable, and the pure selector picks again (``torch_varlen`` ->
|
| 137 |
``flash_attn`` -> ``eager``). The probe is the single authority on whether a
|
| 138 |
varlen kernel runs, on every backend: ROCm gets no separate predicate, it is
|
| 139 |
probed like CUDA. The result, including a fallback, is memoized per
|
| 140 |
-
``(device, override)``, so the probe runs once per device per process.
|
| 141 |
-
|
| 142 |
-
instead, with no WARNING and no tier marked
|
| 143 |
-
memoizes a raised call, so
|
| 144 |
-
to a lower tier
|
|
|
|
|
|
|
|
|
|
| 145 |
"""
|
| 146 |
is_cuda = dev_type == "cuda"
|
| 147 |
device = torch.device(dev_type, dev_index)
|
|
@@ -160,9 +196,9 @@ def _select_cached(dev_type, dev_index, override):
|
|
| 160 |
try:
|
| 161 |
_probe_varlen_kernel(impl, device)
|
| 162 |
return impl
|
| 163 |
-
except
|
| 164 |
-
|
| 165 |
-
|
| 166 |
logger.warning(
|
| 167 |
"NomicBert attention: auto rejected impl=%s on device=%s (probe raised %s: %s); "
|
| 168 |
"trying the next tier",
|
|
@@ -201,12 +237,11 @@ def _pad(x_u: torch.Tensor, indices: torch.Tensor, B: int, S: int) -> torch.Tens
|
|
| 201 |
|
| 202 |
|
| 203 |
def select_attention_impl(device: torch.device) -> str:
|
| 204 |
-
"""
|
| 205 |
|
| 206 |
-
Reads
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
``(device.type, device.index, override)``.
|
| 210 |
"""
|
| 211 |
override = os.environ.get(ATTN_IMPL_ENV) or "auto"
|
| 212 |
return _select_cached(device.type, device.index, override)
|
|
@@ -215,9 +250,10 @@ def select_attention_impl(device: torch.device) -> str:
|
|
| 215 |
def _inline_extended_mask(attention_mask: Optional[torch.Tensor], dtype: torch.dtype) -> Optional[torch.Tensor]:
|
| 216 |
"""Build the additive ``[B, 1, 1, S]`` mask the eager tier consumes.
|
| 217 |
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
helper
|
|
|
|
| 221 |
"""
|
| 222 |
if attention_mask is None:
|
| 223 |
return None
|
|
@@ -247,7 +283,7 @@ def _log_attention_impl_once(impl: str, device: torch.device) -> None:
|
|
| 247 |
except importlib.metadata.PackageNotFoundError:
|
| 248 |
fa = "unknown"
|
| 249 |
else:
|
| 250 |
-
fa = "absent"
|
| 251 |
logger.info(
|
| 252 |
"NomicBert attention impl=%s device=%s capability=%s torch=%s flash_attn=%s override=%s",
|
| 253 |
impl,
|
|
@@ -260,7 +296,7 @@ def _log_attention_impl_once(impl: str, device: torch.device) -> None:
|
|
| 260 |
|
| 261 |
|
| 262 |
# adapted from flash attention, added safe serialization option for hf models
|
| 263 |
-
def state_dict_from_pretrained(model_name, safe_serialization=False, device=None, dtype=None):
|
| 264 |
# If not fp32, then we don't want to load directly to the GPU
|
| 265 |
mapped_device = "cpu" if dtype not in [torch.float32, None] else device
|
| 266 |
is_sharded = False
|
|
@@ -288,10 +324,14 @@ def state_dict_from_pretrained(model_name, safe_serialization=False, device=None
|
|
| 288 |
load_safe = True
|
| 289 |
else: # Try loading from HF hub instead of from local files
|
| 290 |
weight_name = WEIGHTS_NAME if not safe_serialization else SAFE_WEIGHTS_NAME
|
| 291 |
-
resolved_archive_file = cached_file(
|
|
|
|
|
|
|
| 292 |
if resolved_archive_file is None:
|
| 293 |
weight_index = WEIGHTS_INDEX_NAME if not safe_serialization else SAFE_WEIGHTS_INDEX_NAME
|
| 294 |
-
resolved_archive_file = cached_file(
|
|
|
|
|
|
|
| 295 |
if resolved_archive_file is not None:
|
| 296 |
is_sharded = True
|
| 297 |
|
|
@@ -308,7 +348,9 @@ def state_dict_from_pretrained(model_name, safe_serialization=False, device=None
|
|
| 308 |
if is_sharded:
|
| 309 |
# resolved_archive_file becomes a list of files that point to the different
|
| 310 |
# checkpoint shards in this case.
|
| 311 |
-
resolved_archive_file, sharded_metadata = get_checkpoint_shard_files(
|
|
|
|
|
|
|
| 312 |
state_dict = {}
|
| 313 |
for sharded_file in resolved_archive_file:
|
| 314 |
state_dict.update(loader(sharded_file))
|
|
@@ -528,9 +570,26 @@ class NomicBertPreTrainedModel(PreTrainedModel):
|
|
| 528 |
*inputs, **kwargs: additional input for the specific NomicBert class
|
| 529 |
(ex: num_labels for NomicBertForSequenceClassification)
|
| 530 |
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 531 |
# Instantiate model.
|
| 532 |
if config is None:
|
| 533 |
-
config = cls.config_class.from_pretrained(model_name)
|
| 534 |
remove_cls = cls != NomicBertForPreTraining
|
| 535 |
remove_bert_prefix = cls != NomicBertForPreTraining
|
| 536 |
ignore_mismatched_shapes = kwargs.pop("ignore_mismatched_sizes", False)
|
|
@@ -569,7 +628,9 @@ class NomicBertPreTrainedModel(PreTrainedModel):
|
|
| 569 |
load_return = model.load_state_dict(state_dict, strict=False)
|
| 570 |
else:
|
| 571 |
# TODO: can probably check config class and see if we need to remap from a bert model
|
| 572 |
-
state_dict = state_dict_from_pretrained(
|
|
|
|
|
|
|
| 573 |
state_dict = remap_bert_state_dict(
|
| 574 |
state_dict,
|
| 575 |
config,
|
|
@@ -581,21 +642,27 @@ class NomicBertPreTrainedModel(PreTrainedModel):
|
|
| 581 |
state_dict = filter_shapes(state_dict, model)
|
| 582 |
|
| 583 |
load_return = model.load_state_dict(state_dict, strict=True)
|
| 584 |
-
# Honor
|
| 585 |
-
#
|
| 586 |
-
#
|
| 587 |
-
# kwarg
|
|
|
|
| 588 |
import torch as _torch
|
| 589 |
-
_td = kwargs.get("
|
| 590 |
if _td is None:
|
| 591 |
-
_td =
|
|
|
|
|
|
|
| 592 |
if _td == "auto" or _td is None:
|
| 593 |
_td = next(iter(state_dict.values())).dtype
|
| 594 |
if isinstance(_td, str):
|
| 595 |
_td = getattr(_torch, _td)
|
| 596 |
if _td is not None:
|
| 597 |
model = model.to(_td)
|
| 598 |
-
|
|
|
|
|
|
|
|
|
|
| 599 |
return model
|
| 600 |
|
| 601 |
def _set_gradient_checkpointing(self, module, value=False):
|
|
@@ -1028,11 +1095,11 @@ class NomicBertAttention(nn.Module):
|
|
| 1028 |
impl = select_attention_impl(hidden_states.device)
|
| 1029 |
if impl != "eager":
|
| 1030 |
# --- torch_varlen / flash_attn varlen paths ---
|
| 1031 |
-
# Both consume the
|
| 1032 |
-
# from cu_seqlens) — NOT the additive [B, 1, 1, S] mask that
|
| 1033 |
# NomicBertModel.forward builds inline (_inline_extended_mask) for the
|
| 1034 |
# eager tier. NomicBertModel.forward calls the same select_attention_impl
|
| 1035 |
-
# and, on a varlen tier, passes
|
| 1036 |
#
|
| 1037 |
# Correctness keystone: RoPE was already applied above, on the dense
|
| 1038 |
# [B, S, 3, H, D] tensor, BEFORE unpadding — so per-sequence positions
|
|
@@ -1040,12 +1107,10 @@ class NomicBertAttention(nn.Module):
|
|
| 1040 |
# absolute positions of sequence 1's tail and silently drift embeddings.
|
| 1041 |
B, S = qkv.shape[0], qkv.shape[1]
|
| 1042 |
|
| 1043 |
-
# Both varlen kernels require fp16/bf16 inputs.
|
| 1044 |
-
#
|
| 1045 |
-
#
|
| 1046 |
-
#
|
| 1047 |
-
# out_proj. This makes the model work regardless of how it was loaded —
|
| 1048 |
-
# no external cast needed.
|
| 1049 |
orig_dtype = qkv.dtype
|
| 1050 |
if orig_dtype not in (torch.float16, torch.bfloat16):
|
| 1051 |
qkv = qkv.to(torch.bfloat16)
|
|
@@ -1376,7 +1441,7 @@ class NomicBertModel(NomicBertPreTrainedModel):
|
|
| 1376 |
# No varlen tier engaged (CPU, pre-Ampere, forced, or no kernel):
|
| 1377 |
# build the additive [B, 1, 1, S] mask inline for the eager block.
|
| 1378 |
attention_mask = _inline_extended_mask(attention_mask, hidden_states.dtype)
|
| 1379 |
-
# else: torch_varlen / flash_attn consume the
|
| 1380 |
# NomicBertAttention.forward derives lengths from it via cu_seqlens.
|
| 1381 |
sequence_output = self.encoder(
|
| 1382 |
hidden_states, attention_mask=attention_mask, return_dict=return_dict
|
|
|
|
| 33 |
# Three-tier attention dispatch (auto by default, override via NOMIC_BERT_ATTN_IMPL):
|
| 34 |
# 1. torch_varlen — torch.nn.attention.varlen, no third-party kernel
|
| 35 |
# 2. flash_attn — the flash-attn varlen-packed kernel (optional dependency)
|
| 36 |
+
# 3. eager — the original padded algorithm, runs everywhere
|
| 37 |
# Both tier 1 and tier 2 run the same FA2-family kernel and refuse pre-Ampere
|
| 38 |
# hardware, so a single capability predicate gates both (resolve_attention_impl,
|
| 39 |
# below). In auto mode a tiny kernel probe then confirms the chosen varlen tier
|
| 40 |
# actually runs on the device, stepping down a tier if it raises (_select_cached).
|
| 41 |
+
# When neither is available (CPU-only hosts, or builds without them) tier 3 runs,
|
| 42 |
+
# so this model loads and encodes everywhere.
|
| 43 |
try: # pragma: no cover - import guard, exercised by environment
|
| 44 |
from torch.nn.attention.varlen import varlen_attn as _torch_varlen_attn
|
| 45 |
|
|
|
|
| 47 |
except ImportError: # pragma: no cover
|
| 48 |
_torch_varlen_attn = None # type: ignore[assignment]
|
| 49 |
_TORCH_VARLEN_AVAILABLE = False
|
| 50 |
+
# Why flash_attn is unavailable, or None. A missing package and a broken install
|
| 51 |
+
# (e.g. an extension built against another torch) both leave the tier unavailable,
|
| 52 |
+
# but only a broken install is worth a WARNING; the INFO tier line and a forced
|
| 53 |
+
# flash_attn error both name the cause.
|
| 54 |
+
_FLASH_IMPORT_ERROR: Optional[str] = None
|
| 55 |
try: # pragma: no cover - import guard, exercised by environment
|
| 56 |
from flash_attn import flash_attn_varlen_qkvpacked_func
|
| 57 |
|
| 58 |
_FLASH_AVAILABLE = True
|
| 59 |
+
except Exception as _exc: # pragma: no cover
|
| 60 |
flash_attn_varlen_qkvpacked_func = None # type: ignore[assignment]
|
| 61 |
_FLASH_AVAILABLE = False
|
| 62 |
+
_FLASH_IMPORT_ERROR = f"{type(_exc).__name__}: {_exc}"
|
| 63 |
+
if not isinstance(_exc, ModuleNotFoundError) or _exc.name != "flash_attn":
|
| 64 |
+
logging.getLogger(__name__).warning(
|
| 65 |
+
"NomicBert attention: flash_attn is installed but failed to import (%s); "
|
| 66 |
+
"the flash_attn tier is unavailable",
|
| 67 |
+
_FLASH_IMPORT_ERROR,
|
| 68 |
+
)
|
| 69 |
|
| 70 |
ATTENTION_IMPLS = ("torch_varlen", "flash_attn", "eager")
|
| 71 |
ATTN_IMPL_ENV = "NOMIC_BERT_ATTN_IMPL" # auto (default) | torch_varlen | flash_attn | eager
|
|
|
|
| 103 |
f"is unavailable (torch {torch.__version__}; needs >= 2.10)"
|
| 104 |
)
|
| 105 |
if override == "flash_attn" and not flash_available:
|
| 106 |
+
raise RuntimeError(
|
| 107 |
+
f"{ATTN_IMPL_ENV}=flash_attn requested but flash_attn is not importable"
|
| 108 |
+
f" ({_FLASH_IMPORT_ERROR})"
|
| 109 |
+
)
|
| 110 |
if not is_cuda:
|
| 111 |
raise RuntimeError(f"{ATTN_IMPL_ENV}={override} requested but device is cpu (not CUDA)")
|
| 112 |
if capability is None or tuple(capability) < (8, 0):
|
|
|
|
| 129 |
only a proxy: a torch build can ship ``torch.nn.attention.varlen`` without an
|
| 130 |
FA kernel for the device's arch, and on ROCm ``device.type`` is ``"cuda"`` and
|
| 131 |
the capability is the gfx major.
|
| 132 |
+
|
| 133 |
+
The device is synchronized BEFORE the kernel call, outside anything that would
|
| 134 |
+
read a failure as "this tier has no kernel": an asynchronous CUDA error left by
|
| 135 |
+
earlier work surfaces there as itself instead of being blamed on the probe.
|
| 136 |
"""
|
| 137 |
+
torch.cuda.synchronize(device)
|
| 138 |
qkv = torch.randn(2, 3, 1, 64, device=device, dtype=torch.bfloat16)
|
| 139 |
cu = torch.tensor([0, 2], device=device, dtype=torch.int32)
|
| 140 |
if impl == "torch_varlen":
|
|
|
|
| 145 |
torch.cuda.synchronize(device)
|
| 146 |
|
| 147 |
|
| 148 |
+
@functools.lru_cache(maxsize=None)
|
| 149 |
+
def _is_oom(exc: BaseException) -> bool:
|
| 150 |
+
"""True for an out-of-memory failure, allocator-level or driver-level.
|
| 151 |
+
|
| 152 |
+
``torch.cuda.OutOfMemoryError`` is the caching allocator's; a driver-level
|
| 153 |
+
allocation failure inside a kernel surfaces as a plain ``RuntimeError`` (or
|
| 154 |
+
``torch.AcceleratorError``) whose message says ``out of memory``.
|
| 155 |
+
"""
|
| 156 |
+
return isinstance(exc, torch.cuda.OutOfMemoryError) or "out of memory" in str(exc).lower()
|
| 157 |
+
|
| 158 |
+
|
| 159 |
@functools.lru_cache(maxsize=None)
|
| 160 |
def _select_cached(dev_type, dev_index, override):
|
| 161 |
"""Resolve the tier for one device, probing the kernel in auto mode.
|
| 162 |
|
| 163 |
A forced tier is returned exactly as ``resolve_attention_impl`` decides it (or
|
| 164 |
raises), with no probe: it never falls back. In ``auto`` mode a varlen tier is
|
| 165 |
+
accepted only if ``_probe_varlen_kernel`` runs; if the probe raises a
|
| 166 |
+
``RuntimeError`` (which covers ``NotImplementedError`` and
|
| 167 |
+
``torch.AcceleratorError``: the ways a missing or unsupported kernel fails),
|
| 168 |
+
one WARNING names the rejected tier and the exception, that tier is marked
|
| 169 |
unavailable, and the pure selector picks again (``torch_varlen`` ->
|
| 170 |
``flash_attn`` -> ``eager``). The probe is the single authority on whether a
|
| 171 |
varlen kernel runs, on every backend: ROCm gets no separate predicate, it is
|
| 172 |
probed like CUDA. The result, including a fallback, is memoized per
|
| 173 |
+
``(device, override)``, so the probe runs once per device per process.
|
| 174 |
+
|
| 175 |
+
Three failures propagate instead, with no WARNING and no tier marked
|
| 176 |
+
unavailable, and ``lru_cache`` never memoizes a raised call, so none of them
|
| 177 |
+
pins the device to a lower tier: an out-of-memory error (``_is_oom``: transient,
|
| 178 |
+
not a missing kernel), an exception that is not a ``RuntimeError`` (a
|
| 179 |
+
``TypeError`` from a changed kernel signature is a bug to fix, not a tier to
|
| 180 |
+
skip), and an error from the pre-probe synchronize (earlier work's fault).
|
| 181 |
"""
|
| 182 |
is_cuda = dev_type == "cuda"
|
| 183 |
device = torch.device(dev_type, dev_index)
|
|
|
|
| 196 |
try:
|
| 197 |
_probe_varlen_kernel(impl, device)
|
| 198 |
return impl
|
| 199 |
+
except RuntimeError as exc:
|
| 200 |
+
if _is_oom(exc):
|
| 201 |
+
raise
|
| 202 |
logger.warning(
|
| 203 |
"NomicBert attention: auto rejected impl=%s on device=%s (probe raised %s: %s); "
|
| 204 |
"trying the next tier",
|
|
|
|
| 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 |
|
| 242 |
+
Reads the override on every call, so changing the env var between forwards
|
| 243 |
+
takes effect, and delegates to ``_select_cached``, which resolves (and in
|
| 244 |
+
auto mode probes) once per ``(device.type, device.index, override)``.
|
|
|
|
| 245 |
"""
|
| 246 |
override = os.environ.get(ATTN_IMPL_ENV) or "auto"
|
| 247 |
return _select_cached(device.type, device.index, override)
|
|
|
|
| 250 |
def _inline_extended_mask(attention_mask: Optional[torch.Tensor], dtype: torch.dtype) -> Optional[torch.Tensor]:
|
| 251 |
"""Build the additive ``[B, 1, 1, S]`` mask the eager tier consumes.
|
| 252 |
|
| 253 |
+
``attention_mask`` is the ``{0,1}`` ``[B, S]`` tokenizer mask: kept positions
|
| 254 |
+
become ``0.0`` and padding becomes ``torch.finfo(dtype).min``, the arithmetic
|
| 255 |
+
of transformers' extended-attention-mask helper, inline so this file does not
|
| 256 |
+
depend on that helper. ``None`` stays ``None``.
|
| 257 |
"""
|
| 258 |
if attention_mask is None:
|
| 259 |
return None
|
|
|
|
| 283 |
except importlib.metadata.PackageNotFoundError:
|
| 284 |
fa = "unknown"
|
| 285 |
else:
|
| 286 |
+
fa = f"absent({_FLASH_IMPORT_ERROR})"
|
| 287 |
logger.info(
|
| 288 |
"NomicBert attention impl=%s device=%s capability=%s torch=%s flash_attn=%s override=%s",
|
| 289 |
impl,
|
|
|
|
| 296 |
|
| 297 |
|
| 298 |
# adapted from flash attention, added safe serialization option for hf models
|
| 299 |
+
def state_dict_from_pretrained(model_name, safe_serialization=False, device=None, dtype=None, **hub_kwargs):
|
| 300 |
# If not fp32, then we don't want to load directly to the GPU
|
| 301 |
mapped_device = "cpu" if dtype not in [torch.float32, None] else device
|
| 302 |
is_sharded = False
|
|
|
|
| 324 |
load_safe = True
|
| 325 |
else: # Try loading from HF hub instead of from local files
|
| 326 |
weight_name = WEIGHTS_NAME if not safe_serialization else SAFE_WEIGHTS_NAME
|
| 327 |
+
resolved_archive_file = cached_file(
|
| 328 |
+
model_name, weight_name, _raise_exceptions_for_missing_entries=False, **hub_kwargs
|
| 329 |
+
)
|
| 330 |
if resolved_archive_file is None:
|
| 331 |
weight_index = WEIGHTS_INDEX_NAME if not safe_serialization else SAFE_WEIGHTS_INDEX_NAME
|
| 332 |
+
resolved_archive_file = cached_file(
|
| 333 |
+
model_name, weight_index, _raise_exceptions_for_missing_entries=False, **hub_kwargs
|
| 334 |
+
)
|
| 335 |
if resolved_archive_file is not None:
|
| 336 |
is_sharded = True
|
| 337 |
|
|
|
|
| 348 |
if is_sharded:
|
| 349 |
# resolved_archive_file becomes a list of files that point to the different
|
| 350 |
# checkpoint shards in this case.
|
| 351 |
+
resolved_archive_file, sharded_metadata = get_checkpoint_shard_files(
|
| 352 |
+
model_name, resolved_archive_file, **hub_kwargs
|
| 353 |
+
)
|
| 354 |
state_dict = {}
|
| 355 |
for sharded_file in resolved_archive_file:
|
| 356 |
state_dict.update(loader(sharded_file))
|
|
|
|
| 570 |
*inputs, **kwargs: additional input for the specific NomicBert class
|
| 571 |
(ex: num_labels for NomicBertForSequenceClassification)
|
| 572 |
"""
|
| 573 |
+
# The Hub-resolution kwargs transformers' Auto classes forward (revision above
|
| 574 |
+
# all) reach every file this method downloads, so a pinned revision pins the
|
| 575 |
+
# weights, not just the code and config.
|
| 576 |
+
hub_kwargs = {
|
| 577 |
+
k: kwargs[k]
|
| 578 |
+
for k in (
|
| 579 |
+
"revision",
|
| 580 |
+
"cache_dir",
|
| 581 |
+
"token",
|
| 582 |
+
"local_files_only",
|
| 583 |
+
"force_download",
|
| 584 |
+
"proxies",
|
| 585 |
+
"subfolder",
|
| 586 |
+
"_commit_hash",
|
| 587 |
+
)
|
| 588 |
+
if kwargs.get(k) is not None
|
| 589 |
+
}
|
| 590 |
# Instantiate model.
|
| 591 |
if config is None:
|
| 592 |
+
config = cls.config_class.from_pretrained(model_name, **hub_kwargs)
|
| 593 |
remove_cls = cls != NomicBertForPreTraining
|
| 594 |
remove_bert_prefix = cls != NomicBertForPreTraining
|
| 595 |
ignore_mismatched_shapes = kwargs.pop("ignore_mismatched_sizes", False)
|
|
|
|
| 628 |
load_return = model.load_state_dict(state_dict, strict=False)
|
| 629 |
else:
|
| 630 |
# TODO: can probably check config class and see if we need to remap from a bert model
|
| 631 |
+
state_dict = state_dict_from_pretrained(
|
| 632 |
+
model_name, safe_serialization=kwargs.get("safe_serialization", False), **hub_kwargs
|
| 633 |
+
)
|
| 634 |
state_dict = remap_bert_state_dict(
|
| 635 |
state_dict,
|
| 636 |
config,
|
|
|
|
| 642 |
state_dict = filter_shapes(state_dict, model)
|
| 643 |
|
| 644 |
load_return = model.load_state_dict(state_dict, strict=True)
|
| 645 |
+
# Honor the requested dtype like transformers' native from_pretrained does.
|
| 646 |
+
# This custom override instantiates fp32 and load_state_dict upcasts the
|
| 647 |
+
# checkpoint into the fp32 params, so the dtype is applied here. Resolve
|
| 648 |
+
# explicit kwarg (`dtype`, or the older name `torch_dtype`) > config >
|
| 649 |
+
# checkpoint dtype, then cast the model.
|
| 650 |
import torch as _torch
|
| 651 |
+
_td = kwargs.get("dtype")
|
| 652 |
if _td is None:
|
| 653 |
+
_td = kwargs.get("torch_dtype")
|
| 654 |
+
if _td is None:
|
| 655 |
+
_td = getattr(config, "dtype", None) or getattr(config, "torch_dtype", None)
|
| 656 |
if _td == "auto" or _td is None:
|
| 657 |
_td = next(iter(state_dict.values())).dtype
|
| 658 |
if isinstance(_td, str):
|
| 659 |
_td = getattr(_torch, _td)
|
| 660 |
if _td is not None:
|
| 661 |
model = model.to(_td)
|
| 662 |
+
if load_return.missing_keys or load_return.unexpected_keys:
|
| 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):
|
|
|
|
| 1095 |
impl = select_attention_impl(hidden_states.device)
|
| 1096 |
if impl != "eager":
|
| 1097 |
# --- torch_varlen / flash_attn varlen paths ---
|
| 1098 |
+
# Both consume the tokenizer's {0,1} [B, S] mask (per-sequence lengths
|
| 1099 |
+
# derived from cu_seqlens) — NOT the additive [B, 1, 1, S] mask that
|
| 1100 |
# NomicBertModel.forward builds inline (_inline_extended_mask) for the
|
| 1101 |
# eager tier. NomicBertModel.forward calls the same select_attention_impl
|
| 1102 |
+
# and, on a varlen tier, passes the {0,1} mask straight through.
|
| 1103 |
#
|
| 1104 |
# Correctness keystone: RoPE was already applied above, on the dense
|
| 1105 |
# [B, S, 3, H, D] tensor, BEFORE unpadding — so per-sequence positions
|
|
|
|
| 1107 |
# absolute positions of sequence 1's tail and silently drift embeddings.
|
| 1108 |
B, S = qkv.shape[0], qkv.shape[1]
|
| 1109 |
|
| 1110 |
+
# Both varlen kernels require fp16/bf16 inputs. The model loads bf16 by
|
| 1111 |
+
# default, but a caller can ask for fp32 (dtype=torch.float32), so an
|
| 1112 |
+
# fp32 qkv is cast to bf16 for the kernel call and the result is cast
|
| 1113 |
+
# back to the model dtype before out_proj.
|
|
|
|
|
|
|
| 1114 |
orig_dtype = qkv.dtype
|
| 1115 |
if orig_dtype not in (torch.float16, torch.bfloat16):
|
| 1116 |
qkv = qkv.to(torch.bfloat16)
|
|
|
|
| 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.
|
| 1443 |
attention_mask = _inline_extended_mask(attention_mask, hidden_states.dtype)
|
| 1444 |
+
# else: torch_varlen / flash_attn consume the {0,1} [B, S] mask as-is —
|
| 1445 |
# NomicBertAttention.forward derives lengths from it via cu_seqlens.
|
| 1446 |
sequence_output = self.encoder(
|
| 1447 |
hidden_states, attention_mask=attention_mask, return_dict=return_dict
|