handwoven8588 commited on
Commit
1954de4
·
verified ·
1 Parent(s): 4c04d49

Pin weights to the requested revision; harden the varlen probe; card: all tiers measured on one GPU with stated batch sizes

Browse files
Files changed (2) hide show
  1. README.md +51 -30
  2. 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, it logs one WARNING naming the tier and the error and steps down to the next tier. A
58
- forced override (e.g. `NOMIC_BERT_ATTN_IMPL=torch_varlen`) is not probed and **raises**
 
 
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 re-measures all three tiers after adding the
94
- dispatch: each tier's output on a 64-code-snippet corpus (fp32-renormalized), compared by cosine
95
- similarity against this same repo's **pre-dispatch output on the same device** (`flash_attn` on
96
- CUDA, `eager` on CPU — `flash_attn` can't run on CPU). `eager` and `flash_attn` each run the same
97
- computation as before the dispatch was added, so their cosines are parity checks (round to
98
- 1.000000 at 6 decimal places; true values are ≥ 0.9999997). `torch_varlen` is a different,
99
- torch-native kernel, and its cosine is the real signal.
100
-
101
- | GPU | tier | device | transformers | mean cos | min cos | peak VRAM |
 
 
 
 
 
 
102
  | --- | --- | --- | --- | --- | --- | --- |
103
- | RTX 3090 Ti (sm_86) | `flash_attn` | cuda | 5.16.1 | 1.000000 | 1.000000 | 1145 MiB |
104
- | RTX 3090 Ti (sm_86) | `torch_varlen` | cuda | 5.16.1 | 0.999941 | 0.999856 | 875 MiB |
105
- | RTX 3090 Ti (sm_86) | `eager` | cpu | 5.16.1 | 1.000000 | 1.000000 | – (CPU) |
106
- | RTX 5090 Laptop GPU (sm_120) | `flash_attn` | cuda | 5.11.0 | 1.000000 | 1.000000 | 1169 MiB |
107
- | RTX 5090 Laptop GPU (sm_120) | `torch_varlen` | cuda | 5.11.0 | 0.999944 | 0.999919 | 899 MiB |
108
- | RTX 5090 Laptop GPU (sm_120) | `eager` | cpu | 5.11.0 | 1.000000 | 1.000000 | – (CPU) |
109
- | CPU only | `eager` | cpu | 5.17.0 | 1.000000 | 1.000000 | – (CPU) |
110
-
111
- Cosines rounded to 6 decimal places; peak VRAM rounded to the nearest MiB. `eager` always runs on
112
- CPU in this protocol (it is the universal fallback tier); the GPU named in the first column is the
113
- host each row's measurement ran on, not the device `eager` used on that row. `torch` was
114
- `2.12.1+cu130` for every row except the standalone CPU-only row (`torch 2.14.0+cpu`), which used a
115
- separate `transformers==5.17.0` install to check the eager path against a newer transformers.
116
-
117
- Separately, both varlen tiers (`torch_varlen`, `flash_attn`) are gated in this repo's downstream
118
- test suite at min cosine > 0.997 against the fp32 `nomic-ai/CodeRankEmbed` reference (not
119
- re-measured here — see the table above for this repo's own numbers), and stay under 20 GB peak VRAM
120
- at batch size 256.
 
 
 
 
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 path, byte-for-byte, 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
- # unchanged, 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,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(f"{ATTN_IMPL_ENV}=flash_attn requested but flash_attn is not importable")
 
 
 
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, one
135
- WARNING names the rejected tier and the exception, that tier is marked
 
 
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. A
141
- ``torch.cuda.OutOfMemoryError`` from the probe is re-raised immediately
142
- instead, with no WARNING and no tier marked unavailable; ``lru_cache`` never
143
- memoizes a raised call, so a transient OOM never permanently pins the device
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 torch.cuda.OutOfMemoryError:
164
- raise
165
- except Exception as exc:
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
- """Bind ``resolve_attention_impl`` to a device, the env override and a memo.
205
 
206
- Reads ``NOMIC_BERT_ATTN_IMPL`` on every call (so a forced-unavailable tier
207
- keeps raising rather than sticking to a memoized exception — ``lru_cache``
208
- never caches a raised call) and memoizes the resolved tier per
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
- Replaces the ``PreTrainedModel`` helper, which newer transformers removed,
219
- with the same arithmetic inline, so this file has no dependency on that
220
- helper. ``None`` stays ``None``.
 
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(model_name, weight_name, _raise_exceptions_for_missing_entries=False)
 
 
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(model_name, weight_index, _raise_exceptions_for_missing_entries=False)
 
 
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(model_name, resolved_archive_file)
 
 
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(model_name, safe_serialization=kwargs.get("safe_serialization", False))
 
 
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 torch_dtype like transformers' native from_pretrained does. Our custom
585
- # override above bypassed it (it instantiates fp32 and load_state_dict upcasts
586
- # the checkpoint into the fp32 params, ignoring torch_dtype). Resolve explicit
587
- # kwarg > config.torch_dtype > checkpoint dtype, then cast the model.
 
588
  import torch as _torch
589
- _td = kwargs.get("torch_dtype")
590
  if _td is None:
591
- _td = getattr(config, "torch_dtype", None)
 
 
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
- logger.warning(load_return)
 
 
 
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 raw bool [B, S] mask (per-sequence lengths derived
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 this raw bool mask straight through.
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. sentence-transformers
1044
- # loads the weights as fp32 by default (it silently drops
1045
- # config.torch_dtype), so if the model is fp32 we cast qkv to bf16 for
1046
- # the kernel call and cast the result back to the model dtype before
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 raw bool [B, S] mask as-is —
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