logogram 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (54) hide show
  1. logogram/__init__.py +6 -0
  2. logogram/__main__.py +5 -0
  3. logogram/analysis.py +419 -0
  4. logogram/atp.py +120 -0
  5. logogram/backends/__init__.py +5 -0
  6. logogram/backends/base.py +202 -0
  7. logogram/backends/hub.py +375 -0
  8. logogram/backends/saes.py +277 -0
  9. logogram/backends/transformer_lens.py +872 -0
  10. logogram/cli.py +496 -0
  11. logogram/compare.py +177 -0
  12. logogram/datasets.py +159 -0
  13. logogram/direct.py +193 -0
  14. logogram/engine.py +550 -0
  15. logogram/examples/ioi-gpt2/.gitignore +3 -0
  16. logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
  17. logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
  18. logogram/examples/ioi-gpt2/project.json +6 -0
  19. logogram/exports.py +33 -0
  20. logogram/features.py +368 -0
  21. logogram/fileio.py +63 -0
  22. logogram/ioi.py +220 -0
  23. logogram/paths.py +204 -0
  24. logogram/project.py +444 -0
  25. logogram/prompts.py +204 -0
  26. logogram/research.py +84 -0
  27. logogram/results.py +240 -0
  28. logogram/runner.py +396 -0
  29. logogram/runs.py +98 -0
  30. logogram/sae.py +161 -0
  31. logogram/schema.py +302 -0
  32. logogram/server/__init__.py +1 -0
  33. logogram/server/app.py +1083 -0
  34. logogram/server/models.py +426 -0
  35. logogram/server/security.py +212 -0
  36. logogram/server/state.py +585 -0
  37. logogram/sites.py +249 -0
  38. logogram/spec.py +518 -0
  39. logogram/stats.py +171 -0
  40. logogram/steering.py +258 -0
  41. logogram/system.py +379 -0
  42. logogram/updates.py +194 -0
  43. logogram/verify.py +39 -0
  44. logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
  45. logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
  46. logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
  47. logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
  48. logogram/web_dist/favicon.svg +1 -0
  49. logogram/web_dist/index.html +15 -0
  50. logogram-0.1.0.dist-info/METADATA +550 -0
  51. logogram-0.1.0.dist-info/RECORD +54 -0
  52. logogram-0.1.0.dist-info/WHEEL +4 -0
  53. logogram-0.1.0.dist-info/entry_points.txt +2 -0
  54. logogram-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,202 @@
1
+ """The interface between experiments and a model library.
2
+
3
+ Experiments speak only in abstract sites (``resid_pre``, ``head`` ...) and tensors. A backend
4
+ maps those sites onto its own hooks. TransformerLens is the only backend in v0.1; others (for
5
+ example remote execution) can be added by implementing :class:`ModelBackend`.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import threading
11
+ from abc import ABC, abstractmethod
12
+ from dataclasses import asdict, dataclass, field
13
+ from typing import Any
14
+
15
+ import torch
16
+
17
+ STREAM_KINDS = ("resid_pre", "resid_mid", "resid_post", "attn_out", "mlp_out")
18
+ ALL_KINDS = (*STREAM_KINDS, "head")
19
+
20
+
21
+ @dataclass(frozen=True)
22
+ class ModelInfo:
23
+ id: str
24
+ revision: str | None
25
+ architecture: str
26
+ n_layers: int
27
+ n_heads: int
28
+ d_model: int
29
+ d_head: int
30
+ d_mlp: int | None
31
+ d_vocab: int
32
+ n_ctx: int
33
+ n_params: int | None
34
+ dtype: str
35
+ device: str # "cpu", "cuda" or "mps"
36
+ device_name: str
37
+ process_weights: bool
38
+ site_kinds: tuple[str, ...]
39
+ backend: str
40
+ backend_version: str
41
+ extra: dict[str, Any] = field(default_factory=dict)
42
+
43
+ def to_dict(self) -> dict[str, Any]:
44
+ data = asdict(self)
45
+ data["site_kinds"] = list(self.site_kinds)
46
+ return data
47
+
48
+
49
+ @dataclass
50
+ class Tokenized:
51
+ ids: list[int]
52
+ tokens: list[str]
53
+ offsets: list[tuple[int, int]] # character span of each token; (0, 0) for added BOS
54
+
55
+
56
+ @dataclass
57
+ class Patch:
58
+ """Replace one site's activation, row by row, during a forward pass.
59
+
60
+ ``heads`` gives each row's head for head sites. ``positions`` gives each row's single
61
+ position, or is ``None`` to replace every position. ``values`` is ``[B, pos, d]`` when every
62
+ position is replaced and ``[B, d]`` otherwise (``d`` is ``d_head`` for heads).
63
+ """
64
+
65
+ kind: str
66
+ layer: int
67
+ values: torch.Tensor
68
+ heads: torch.Tensor | None = None
69
+ positions: torch.Tensor | None = None
70
+
71
+
72
+ class BackendError(RuntimeError):
73
+ """A model can't be loaded or run. The message says what went wrong and how to fix it."""
74
+
75
+
76
+ class Cancelled(Exception):
77
+ """The user cancelled the job (a run between batches, or a model load between steps)."""
78
+
79
+
80
+ class ModelBackend(ABC):
81
+ info: ModelInfo
82
+
83
+ def __init__(self) -> None:
84
+ # All model use goes through this lock, one forward pass at a time, so interactive
85
+ # requests can interleave with a running sweep between batches.
86
+ self.lock = threading.RLock()
87
+
88
+ @abstractmethod
89
+ def tokenize(self, text: str, prepend_bos: bool) -> Tokenized: ...
90
+
91
+ @abstractmethod
92
+ def single_token_id(self, text: str) -> int | None:
93
+ """The token id if ``text`` is exactly one token, else None."""
94
+
95
+ @abstractmethod
96
+ def token_str(self, token_id: int) -> str: ...
97
+
98
+ @abstractmethod
99
+ def final_logits(self, tokens: torch.Tensor, patch: Patch | None = None) -> torch.Tensor:
100
+ """Logits at the last position, ``[B, vocab]`` in float32."""
101
+
102
+ @abstractmethod
103
+ def capture(
104
+ self, tokens: torch.Tensor, sites: list[tuple[str, int]]
105
+ ) -> dict[tuple[str, int], torch.Tensor]:
106
+ """Activations for ``(kind, layer)`` sites: ``[B, pos, d]``, or ``[B, pos, H, d_head]``."""
107
+
108
+ @abstractmethod
109
+ def attention_pattern(self, tokens: torch.Tensor, layer: int) -> torch.Tensor:
110
+ """Attention probabilities ``[B, H, query, key]`` in float32."""
111
+
112
+ def edit_logits(
113
+ self,
114
+ tokens: torch.Tensor,
115
+ kind: str,
116
+ layer: int,
117
+ edit: Any,
118
+ ) -> torch.Tensor:
119
+ """Logits at the last position, ``[B, vocab]``, with the activation at ``(kind, layer)``
120
+ replaced by ``edit(activation)``: ``[B, pos, d]`` in, the same shape out. For
121
+ residual-stream sites the edit changes the stream itself, as patching does."""
122
+ raise BackendError("Editing activations isn't supported by this model backend.")
123
+
124
+ def path_patch(
125
+ self,
126
+ tokens: torch.Tensor,
127
+ sender: Patch,
128
+ frozen_heads: dict[int, torch.Tensor],
129
+ frozen_mlps: dict[int, torch.Tensor] | None,
130
+ receivers: list[tuple[str, int, int, str]],
131
+ ) -> torch.Tensor:
132
+ """Logits at the last position, ``[B, vocab]``, after patching the sender's effect into
133
+ the receivers' inputs only.
134
+
135
+ First pass: the sender is patched, every attention head's output ``z`` is held at
136
+ ``frozen_heads`` (the receiver run's own values, ``[B, pos, H, d_head]`` per layer) except
137
+ the sender's, and so are MLP outputs when ``frozen_mlps`` is given; the receivers' inputs
138
+ are recorded. Second pass: only those inputs are patched in. Receivers are
139
+ ``("head", layer, head, "q" | "k" | "v")`` or ``("logits", -1, -1, "")``.
140
+ """
141
+ raise BackendError("Path patching isn't supported by this model backend.")
142
+
143
+ def gradients(
144
+ self,
145
+ tokens: torch.Tensor,
146
+ answers: torch.Tensor,
147
+ distractors: torch.Tensor,
148
+ sites: list[tuple[str, int]],
149
+ ) -> tuple[dict[tuple[str, int], torch.Tensor], dict[tuple[str, int], torch.Tensor]]:
150
+ """Activations at ``(kind, layer)`` sites and the gradient of logit(answer) -
151
+ logit(distractor) at the last position with respect to each, in one forward and backward
152
+ pass. Shapes as in :meth:`capture`. The gradient of a residual-stream site is with respect
153
+ to the residual stream itself, not only to what the next component reads."""
154
+ raise BackendError("Gradients aren't supported by this model backend.")
155
+
156
+ def direct_effects(
157
+ self,
158
+ tokens: torch.Tensor,
159
+ answers: torch.Tensor,
160
+ distractors: torch.Tensor,
161
+ heads: bool,
162
+ ) -> dict[str, torch.Tensor]:
163
+ """Direct contributions to logit(answer) - logit(distractor) at the last position.
164
+
165
+ Returns float64 tensors on the CPU, one row per prompt: ``embed`` ``[B]``, ``attn_out``
166
+ and ``mlp_out`` ``[B, layers]``, ``head`` ``[B, layers, heads]`` when ``heads`` is true,
167
+ ``logit_diff`` ``[B]`` and ``remainder`` ``[B]`` (what biases add: the logit difference
168
+ minus every component's term).
169
+ """
170
+ raise BackendError("Direct logit attribution isn't supported by this model backend.")
171
+
172
+ def layer_logits(self, tokens: torch.Tensor, position: int, row: int) -> torch.Tensor:
173
+ """Final-norm logit lens at resid_post, ``[layer, vocab]`` for one row, on CPU.
174
+
175
+ Preserve the supplied batch composition. Backends must opt in rather than assuming
176
+ that every model has the same normalization and vocabulary projection.
177
+ """
178
+ raise BackendError("Per-layer predictions aren't supported by this model backend.")
179
+
180
+ @property
181
+ def device(self) -> torch.device:
182
+ return torch.device(self.info.device)
183
+
184
+ def memory_in_use(self) -> int | None:
185
+ """Bytes in use: accelerator memory on a GPU, this process's resident memory on a CPU."""
186
+ dev = self.info.device
187
+ if dev == "cuda" and torch.cuda.is_available():
188
+ return int(torch.cuda.memory_allocated())
189
+ if dev == "mps" and hasattr(torch, "mps"):
190
+ try:
191
+ return int(torch.mps.current_allocated_memory())
192
+ except Exception:
193
+ return None
194
+ try:
195
+ import psutil
196
+
197
+ return int(psutil.Process().memory_info().rss)
198
+ except Exception:
199
+ return None
200
+
201
+ def close(self) -> None: # noqa: B027 - optional hook
202
+ """Release memory held by the model."""
@@ -0,0 +1,375 @@
1
+ """Hugging Face Hub access: resolve an exact revision, download with progress, read metadata.
2
+
3
+ This is the only module that talks to the network, and only when the user starts it (loading a
4
+ model or asking for a memory estimate). Weights are loaded only from safetensors files, which
5
+ cannot execute code; pickle-based ``.bin`` checkpoints are refused.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import fnmatch
11
+ import json
12
+ import logging
13
+ import sys
14
+ import threading
15
+ from collections.abc import Callable
16
+ from dataclasses import dataclass
17
+ from pathlib import Path
18
+ from typing import Any, TypeVar
19
+
20
+ from logogram.backends.base import BackendError, Cancelled
21
+
22
+ log = logging.getLogger(__name__)
23
+
24
+ # Top-level files needed to build the model and tokenizer. Weights are added separately.
25
+ _SUPPORT_PATTERNS = ("*.json", "*.txt", "*.model", "*.tiktoken")
26
+ _SKIP_FILES = {"README.md", ".gitattributes"}
27
+
28
+
29
+ @dataclass
30
+ class RepoFiles:
31
+ revision: str
32
+ files: list[tuple[str, int]] # (path, size in bytes) to download
33
+ n_params: int | None
34
+ gated: bool
35
+
36
+
37
+ ProgressFn = Callable[[int, int, str], None] # (bytes done, bytes total, current file)
38
+ T = TypeVar("T")
39
+
40
+
41
+ def _is_offline_error(exc: Exception) -> bool:
42
+ from huggingface_hub.errors import OfflineModeIsEnabled
43
+
44
+ if isinstance(exc, OfflineModeIsEnabled):
45
+ return True
46
+ name = type(exc).__name__
47
+ return (
48
+ any(k in name for k in ("Connect", "Timeout", "NameResolution"))
49
+ or "connect" in str(exc).lower()
50
+ )
51
+
52
+
53
+ def _friendly_hub_error(model_id: str, exc: Exception) -> BackendError:
54
+ from huggingface_hub.errors import (
55
+ GatedRepoError,
56
+ RepositoryNotFoundError,
57
+ RevisionNotFoundError,
58
+ )
59
+
60
+ if isinstance(exc, GatedRepoError):
61
+ return BackendError(
62
+ f"{model_id} is gated. Accept its license on huggingface.co, then log in on this "
63
+ "machine with `hf auth login` and try again."
64
+ )
65
+ if isinstance(exc, RepositoryNotFoundError):
66
+ return BackendError(
67
+ f"No model called {model_id} was found on Hugging Face. Check the spelling (ids look "
68
+ "like owner/name). Private models need `hf auth login` first."
69
+ )
70
+ if isinstance(exc, RevisionNotFoundError):
71
+ return BackendError(
72
+ f"That revision of {model_id} doesn't exist. Check the commit or branch."
73
+ )
74
+ return BackendError(f"Couldn't reach Hugging Face for {model_id}: {exc}")
75
+
76
+
77
+ def _select_files(siblings: list[tuple[str, int]]) -> list[tuple[str, int]]:
78
+ top_level = [(f, s) for f, s in siblings if "/" not in f and f not in _SKIP_FILES]
79
+ weights = [(f, s) for f, s in top_level if f.endswith(".safetensors")]
80
+ if not weights:
81
+ has_pickle = any(f.endswith((".bin", ".pt", ".pth", ".ckpt")) for f, _ in top_level)
82
+ if has_pickle:
83
+ raise BackendError(
84
+ "This model only publishes pickle weights (.bin). Logogram loads safetensors "
85
+ "weights only, because pickle files can run code when opened. Choose a model "
86
+ "with model.safetensors, or convert the weights yourself."
87
+ )
88
+ raise BackendError("This repository has no safetensors weights at its top level.")
89
+ support = [
90
+ (f, s)
91
+ for f, s in top_level
92
+ if any(fnmatch.fnmatch(f, p) for p in _SUPPORT_PATTERNS)
93
+ and not f.startswith(("onnx", "flax", "tf_"))
94
+ ]
95
+ return sorted(support) + sorted(weights)
96
+
97
+
98
+ def resolve(model_id: str, revision: str | None) -> RepoFiles:
99
+ """Ask the Hub for the exact commit and the files to fetch. Falls back to the local cache."""
100
+ from huggingface_hub import HfApi
101
+
102
+ try:
103
+ info = HfApi().model_info(model_id, revision=revision, files_metadata=True)
104
+ except Exception as exc: # noqa: BLE001 - classified below
105
+ if _is_offline_error(exc):
106
+ cached = cached_snapshot(model_id, revision)
107
+ if cached is not None:
108
+ sha, path = cached
109
+ files = [(p.name, p.stat().st_size) for p in path.iterdir() if p.is_file()]
110
+ return RepoFiles(
111
+ revision=sha, files=_select_files(files), n_params=None, gated=False
112
+ )
113
+ raise BackendError(
114
+ f"Couldn't reach Hugging Face, and {model_id} isn't in your local cache. "
115
+ "Connect to the internet to download it once."
116
+ ) from exc
117
+ raise _friendly_hub_error(model_id, exc) from exc
118
+ siblings = [(s.rfilename, int(s.size or 0)) for s in (info.siblings or [])]
119
+ n_params = None
120
+ if getattr(info, "safetensors", None) is not None:
121
+ n_params = int(info.safetensors.total)
122
+ return RepoFiles(
123
+ revision=info.sha,
124
+ files=_select_files(siblings),
125
+ n_params=n_params,
126
+ gated=bool(info.gated),
127
+ )
128
+
129
+
130
+ def cached_snapshot(model_id: str, revision: str | None) -> tuple[str, Path] | None:
131
+ """The (commit, folder) of a cached snapshot containing config.json, if any."""
132
+ from huggingface_hub import try_to_load_from_cache
133
+
134
+ found = try_to_load_from_cache(model_id, "config.json", revision=revision or "main")
135
+ if isinstance(found, str):
136
+ path = Path(found).parent
137
+ return path.name, path
138
+ return None
139
+
140
+
141
+ def validate_local_weights(folder: Path) -> None:
142
+ """Refuse pickle-only caches and unsafe or incomplete safetensors shard indexes."""
143
+ if (folder / "adapter_config.json").exists():
144
+ raise BackendError(
145
+ "Adapter checkpoints are not supported. Choose a complete base or merged model published as safetensors, without a separate adapter configuration."
146
+ )
147
+ single = folder / "model.safetensors"
148
+ index = folder / "model.safetensors.index.json"
149
+ if single.is_file():
150
+ return
151
+ if not index.is_file():
152
+ raise BackendError(
153
+ "No model.safetensors weights are available. Download a model with safetensors weights."
154
+ )
155
+ try:
156
+ data = json.loads(index.read_text(encoding="utf-8"))
157
+ mapping = data["weight_map"]
158
+ if not isinstance(mapping, dict) or not mapping:
159
+ raise ValueError("Empty weight map")
160
+ for name in mapping.values():
161
+ if (
162
+ not isinstance(name, str)
163
+ or "/" in name
164
+ or "\\" in name
165
+ or ":" in name
166
+ or not name.endswith(".safetensors")
167
+ ):
168
+ raise ValueError("Unsafe shard name")
169
+ if not (folder / name).is_file():
170
+ raise ValueError("Missing shard")
171
+ except (OSError, ValueError, KeyError, TypeError) as exc:
172
+ raise BackendError(
173
+ "The safetensors index is invalid or has missing shards. Download the complete safetensors checkpoint again."
174
+ ) from exc
175
+
176
+
177
+ def _make_bar_class(on_value: Callable[[int], None], size: int) -> type:
178
+ """A silent progress bar that reports how much of one file has arrived.
179
+
180
+ Plain HTTP downloads report bytes as they arrive. Xet downloads report network bytes often
181
+ and bytes written to disk rarely; the larger of the two, capped at the file size, gives
182
+ smooth progress that still ends exactly at the file size.
183
+ """
184
+ from tqdm.std import tqdm
185
+
186
+ class _Bar(tqdm): # type: ignore[misc]
187
+ monitor_interval = 0 # no display, so no tqdm monitor thread per file
188
+
189
+ def __init__(self, *args: Any, **kwargs: Any) -> None:
190
+ kwargs["disable"] = True
191
+ super().__init__(*args, **kwargs)
192
+ self._written = 0
193
+ self._received = 0
194
+
195
+ def _report(self) -> None:
196
+ on_value(min(size, max(self._written, self._received)))
197
+
198
+ def update(self, n: float | None = 1) -> bool | None:
199
+ if n:
200
+ self._written += int(n)
201
+ self._report()
202
+ return None
203
+
204
+ def update_transfer(self, n: float | None = 1) -> None:
205
+ if n:
206
+ self._received += int(n)
207
+ self._report()
208
+
209
+ def set_transfer_postfix_str(self, *args: Any, **kwargs: Any) -> None:
210
+ return None
211
+
212
+ def __getattr__(self, name: str) -> Any:
213
+ # Display hooks that newer huggingface_hub versions may call: nothing to display.
214
+ if name.startswith(("set_", "update_")):
215
+ return lambda *args, **kwargs: None
216
+ raise AttributeError(name)
217
+
218
+ return _Bar
219
+
220
+
221
+ def download(
222
+ model_id: str,
223
+ repo: RepoFiles,
224
+ progress: ProgressFn | None = None,
225
+ *,
226
+ cancel: threading.Event | None = None,
227
+ ) -> Path:
228
+ """Download the selected files at the pinned revision and return the snapshot folder.
229
+
230
+ Setting ``cancel`` raises ``Cancelled`` within a fraction of a second and stops the transfer;
231
+ a partial file stays in the cache, where the next attempt resumes it.
232
+ """
233
+ from huggingface_hub import hf_hub_download, try_to_load_from_cache
234
+
235
+ total = sum(size for _, size in repo.files)
236
+ finished = 0 # bytes of files already complete
237
+ folder: Path | None = None
238
+ for filename, size in repo.files:
239
+ if cancel is not None and cancel.is_set():
240
+ raise Cancelled()
241
+ cached = try_to_load_from_cache(model_id, filename, revision=repo.revision)
242
+ if isinstance(cached, str):
243
+ finished += size
244
+ folder = Path(cached).parent
245
+ if progress:
246
+ progress(finished, total, filename)
247
+ continue
248
+
249
+ def on_value(value: int, _name: str = filename, _base: int = finished) -> None:
250
+ if cancel is not None and cancel.is_set():
251
+ # Raising stops a plain HTTP transfer. Xet transfers are aborted from outside
252
+ # (see _abort_xet); an exception raised in their callback would only be printed.
253
+ if not _called_from_xet():
254
+ raise Cancelled()
255
+ return
256
+ if progress:
257
+ progress(_base + value, total, _name)
258
+
259
+ def fetch(_name: str = filename, _size: int = size, _on: Any = on_value) -> str:
260
+ return hf_hub_download(
261
+ model_id,
262
+ _name,
263
+ revision=repo.revision,
264
+ tqdm_class=_make_bar_class(_on, _size),
265
+ )
266
+
267
+ try:
268
+ path = _cancellable(fetch, cancel)
269
+ except Cancelled:
270
+ raise
271
+ except Exception as exc: # noqa: BLE001
272
+ raise _friendly_hub_error(model_id, exc) from exc
273
+ finished += size
274
+ folder = Path(path).parent
275
+ if progress:
276
+ progress(finished, total, filename)
277
+ if folder is None:
278
+ raise BackendError(f"Nothing to download for {model_id}.")
279
+ return folder
280
+
281
+
282
+ def _cancellable(fn: Callable[[], T], cancel: threading.Event | None) -> T:
283
+ """Run ``fn`` on a helper thread, so waiting for it can stop as soon as ``cancel`` is set."""
284
+ if cancel is None:
285
+ return fn()
286
+ box: dict[str, Any] = {}
287
+ done = threading.Event()
288
+
289
+ def target() -> None:
290
+ try:
291
+ box["value"] = fn()
292
+ except BaseException as exc: # noqa: BLE001 - handed to the waiting thread
293
+ box["error"] = exc
294
+ finally:
295
+ done.set()
296
+
297
+ threading.Thread(target=target, name="logogram-download", daemon=True).start()
298
+ while not done.wait(0.1):
299
+ if cancel.is_set():
300
+ _abort_xet()
301
+ raise Cancelled()
302
+ if "error" in box:
303
+ raise box["error"]
304
+ return box["value"]
305
+
306
+
307
+ def _called_from_xet() -> bool:
308
+ frame = sys._getframe(1)
309
+ while frame is not None:
310
+ if "xet" in frame.f_globals.get("__name__", ""):
311
+ return True
312
+ frame = frame.f_back
313
+ return False
314
+
315
+
316
+ def _abort_xet() -> None:
317
+ """Stop Xet transfers in flight, as huggingface_hub does on Ctrl+C. Without this (if the
318
+ private helper moves), a cancelled Xet download finishes in the background, into the cache."""
319
+ try:
320
+ from huggingface_hub.utils._xet import abort_xet_session
321
+
322
+ abort_xet_session()
323
+ except Exception: # noqa: BLE001
324
+ log.debug("couldn't abort Xet transfers", exc_info=True)
325
+
326
+
327
+ @dataclass
328
+ class ArchitectureSummary:
329
+ n_layers: int
330
+ n_heads: int
331
+ d_model: int
332
+ d_mlp: int
333
+ d_vocab: int
334
+ n_ctx: int
335
+ architecture: str
336
+
337
+
338
+ def read_architecture(config: dict[str, Any]) -> ArchitectureSummary:
339
+ """Read the shape of a transformer from a Hugging Face config.json (best effort)."""
340
+
341
+ def first(*keys: str, default: int | None = None) -> int:
342
+ for key in keys:
343
+ value = config.get(key)
344
+ if isinstance(value, int) and value > 0:
345
+ return value
346
+ text = config.get("text_config")
347
+ if isinstance(text, dict):
348
+ for key in keys:
349
+ value = text.get(key)
350
+ if isinstance(value, int) and value > 0:
351
+ return value
352
+ if default is None:
353
+ raise BackendError(f"config.json has none of {', '.join(keys)}.")
354
+ return default
355
+
356
+ d_model = first("hidden_size", "n_embd", "d_model")
357
+ return ArchitectureSummary(
358
+ n_layers=first("num_hidden_layers", "n_layer", "num_layers"),
359
+ n_heads=first("num_attention_heads", "n_head", "num_heads"),
360
+ d_model=d_model,
361
+ d_mlp=first("intermediate_size", "n_inner", "ffn_dim", default=4 * d_model),
362
+ d_vocab=first("vocab_size"),
363
+ n_ctx=first("max_position_embeddings", "n_positions", "n_ctx", default=2048),
364
+ architecture=(config.get("architectures") or ["unknown"])[0],
365
+ )
366
+
367
+
368
+ def fetch_config(model_id: str, revision: str) -> dict[str, Any]:
369
+ from huggingface_hub import hf_hub_download
370
+
371
+ try:
372
+ path = hf_hub_download(model_id, "config.json", revision=revision)
373
+ except Exception as exc: # noqa: BLE001
374
+ raise _friendly_hub_error(model_id, exc) from exc
375
+ return json.loads(Path(path).read_text(encoding="utf-8"))