AQIT 0.1.0
Loading...
Searching...
No Matches
model_loader.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6from __future__ import annotations
7
8import os
9import threading
10from collections import OrderedDict
11from pathlib import Path
12from typing import Any
13
14# Model config registry — mirrors inspection-backend/model_config.py
15MODEL_CONFIGS: dict[str, dict[str, Any]] = {
16 "llama-3.2-1b": {
17 "hf_name": "meta-llama/Llama-3.2-1B-Instruct",
18 "d_model": 2048,
19 "n_layers": 16,
20 "sae_layer": 8,
21 "sae_source": "native",
22 "sae_layers": {i: f"sae_layer{i}.pt" for i in range(16)},
23 "lora_target_modules": ["q_proj", "v_proj"],
24 },
25 "pythia-2.8b": {
26 "hf_name": "EleutherAI/pythia-2.8b",
27 "d_model": 2560,
28 "n_layers": 32,
29 "sae_layer": 16,
30 "sae_source": "native",
31 "sae_layers": {},
32 "lora_target_modules": ["query_key_value", "dense"],
33 },
34 "gpt2-small": {
35 "hf_name": "gpt2",
36 "d_model": 768,
37 "n_layers": 12,
38 "sae_layer": 8,
39 "sae_source": "native",
40 "sae_layers": {8: "sae_layer8.pt"},
41 "lora_target_modules": ["c_attn", "c_fc"],
42 },
43 "pythia-70m": {
44 "hf_name": "EleutherAI/pythia-70m-deduped",
45 "d_model": 512,
46 "n_layers": 6,
47 "sae_layer": 3,
48 "sae_source": "native",
49 "sae_layers": {3: "sae_layer3.pt"},
50 "lora_target_modules": ["query_key_value", "dense"],
51 },
52 "sarvam-30b": {
53 "hf_name": "sarvamai/sarvam-30b",
54 "d_model": 4096,
55 "n_layers": 19,
56 "n_heads": 16,
57 "sae_layer": 9,
58 "sae_source": "native",
59 "sae_layers": {9: "sarvam_l9.pt"},
60 "trust_remote_code": True,
61 "hf_only": True,
62 "lora_target_modules": ["query_key_value"],
63 },
64 "lfm2.5-1.2b-instruct": {
65 "hf_name": "LiquidAI/LFM2.5-1.2B-Instruct",
66 "d_model": 2048,
67 "n_layers": 16,
68 "n_heads": 32,
69 "sae_layer": 8,
70 "sae_source": "native",
71 "sae_layers": {8: "sae_layer8.pt"},
72 "norm_layers": {8: "norm_layer8.pt"},
73 "hf_only": True,
74 "lora_target_modules": ["q_proj", "v_proj"],
75 },
76 "lfm2.5-1.2b-thinking": {
77 "hf_name": "LiquidAI/LFM2.5-1.2B-Thinking",
78 "d_model": 2048,
79 "n_layers": 16,
80 "n_heads": 32,
81 "sae_layer": 8,
82 "sae_source": "native",
83 "sae_layers": {8: "sae_layer8.pt"},
84 "norm_layers": {8: "norm_layer8.pt"},
85 "hf_only": True,
86 "lora_target_modules": ["q_proj", "v_proj"],
87 },
88 "lfm2.5-230m": {
89 "hf_name": "LiquidAI/LFM2.5-230M",
90 "d_model": 1024,
91 "n_layers": 14,
92 "n_heads": 16,
93 "sae_layer": 7,
94 "sae_source": "native",
95 "sae_layers": {i: f"sae_layer{i}.pt" for i in range(14)},
96 "norm_layers": {i: f"norm_layer{i}.pt" for i in range(14)},
97 "hf_only": True,
98 "lora_target_modules": ["q_proj", "v_proj"],
99 },
100}
101
102_HF_TO_SHORT: dict[str, str] = {
103 cfg["hf_name"].lower(): key for key, cfg in MODEL_CONFIGS.items()
104}
105
106MAX_LOADED_MODELS = 1
107_models: OrderedDict = OrderedDict()
108
109# One in-flight model build at a time. Concurrent from_pretrained / TL conversion
110# on MPS (unified memory) doubles host RAM and looks like a multi-10GB "leak".
111_model_io_lock = threading.RLock()
112
113# Resident SAEs keyed by (short_model_id, layer). Persists across dispatches in
114# the daemon process so repeated feature/inspect commands do not reload from disk.
115_sae_cache: dict[tuple[str, int], Any] = {}
116
117_ACTIVE_MODEL_PATH = Path.home() / ".aquin" / "active_model.txt"
118
120def clear_sae_cache() -> None:
121 _sae_cache.clear()
122
123
124def evict_sae_cache(model_id: str, layer: int | None = None) -> None:
125 """Drop a cached SAE so the next load re-reads from disk (after a rebind)."""
126 short = resolve_model_id(model_id)
127 if layer is None:
128 for key in [k for k in _sae_cache if k[0] == short]:
129 _sae_cache.pop(key, None)
130 else:
131 _sae_cache.pop((short, int(layer)), None)
132
133
134def _save_active_model(model_id: str) -> None:
135 _ACTIVE_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
136 _ACTIVE_MODEL_PATH.write_text(model_id)
137
139def get_active_model_id(*, allow_daemon: bool = True) -> str | None:
140 """Return the model slug written by the last successful `aquin load`."""
141 if _ACTIVE_MODEL_PATH.exists():
142 txt = _ACTIVE_MODEL_PATH.read_text().strip()
143 if txt:
144 return txt
145 if not allow_daemon:
146 return None
147 # If the user suspended `aquin load --model ...` in their shell before the CLI
148 # wrote active_model.txt, the background daemon may still finish loading. Treat
149 # the daemon's resident model as active so follow-up commands still work.
150 try:
151 from aquin.engine import model_daemon
152
153 resident = model_daemon.loaded_model()
154 if resident:
155 return resident
156 except Exception:
157 pass
158 return None
159
160
161def reload_vram_for_model(model_id: str) -> None:
162 """Reload the given model into VRAM (used after heavy jobs release the GPU)."""
163 short = resolve_model_id(model_id)
164 if get_loaded_llm_id() != short:
166 load_model(short)
167
168
169def clear_active_model_file() -> None:
170 if _ACTIVE_MODEL_PATH.exists():
171 _ACTIVE_MODEL_PATH.unlink(missing_ok=True)
172
174def clear_llm_models() -> None:
175 from aquin.compute.device import empty_device_cache
176
177 with _model_io_lock:
178 while _models:
179 _, evicted = _models.popitem(last=False)
180 del evicted
181 _sae_cache.clear()
182 empty_device_cache()
183
184
185def get_loaded_llm_id() -> str | None:
186 if not _models:
187 return None
188 return next(reversed(_models))
190
191class ComputeNotAvailableError(RuntimeError):
192 pass
193
194
195def resolve_model_id(model_id: str) -> str:
196 if model_id in MODEL_CONFIGS:
197 return model_id
198 short = _HF_TO_SHORT.get(model_id.lower())
199 if short:
200 return short
201 from aquin.compute.model_families import resolve_llm_family
202
203 slug, _ = resolve_llm_family(model_id)
204 return slug
205
206
207def get_config(model_id: str) -> dict:
208 short = resolve_model_id(model_id)
209 if short in MODEL_CONFIGS:
210 return MODEL_CONFIGS[short]
211 from aquin.compute.model_families import get_runtime_llm_config
212
213 runtime = get_runtime_llm_config(short)
214 if runtime:
215 return runtime
216 raise ValueError(f"No config for model '{model_id}'")
217
218
219def get_hf_name(model_id: str) -> str:
220 """HuggingFace repo id for transformers.from_pretrained (resolves Aquin short slugs)."""
221 return get_config(model_id)["hf_name"]
222
224def get_sae_layer(model_id: str) -> int:
225 return get_config(model_id)["sae_layer"]
226
227
228def _layer_from_sae_filename(path: Path) -> int | None:
229 stem = path.stem
230 if stem.startswith("sae_layer"):
231 try:
232 return int(stem.replace("sae_layer", ""))
233 except ValueError:
234 return None
235 return None
236
237
238def _is_loadable_sae_file(path: Path) -> bool:
239 from aquin.compute.torch_io import is_valid_sae_checkpoint_path
240
241 return path.is_file() and is_valid_sae_checkpoint_path(path)
243
244def corrupt_sae_checkpoint_hint(model_id: str, layer: int) -> str | None:
245 """If a file exists on disk but cannot be loaded, return a re-download hint."""
246 short = resolve_model_id(model_id)
247 resolved_layer = int(layer)
248 path = resolve_sae_path(short, resolved_layer)
249 if not path.is_file():
250 return None
251 if _is_loadable_sae_file(path):
252 return None
253 return (
254 f"SAE checkpoint at {path} is corrupted or incomplete. "
255 f"Delete it and re-download:\n"
256 f" rm \"{path}\"\n"
257 f" aquin load sae {short}-l{resolved_layer}"
258 )
259
260
262 model_id: str,
263 layer: int,
264 *,
265 device: str | None = None,
266) -> Any:
267 """Load a validated on-disk SAE or raise with a re-download hint."""
268 import torch
269 from aquin.compute.sae import SparseAutoencoder
270
271 from aquin.compute.device import resolve_compute_device
272
273 short = resolve_model_id(model_id)
274 resolved_layer = int(layer)
275 dev = device or resolve_compute_device()
276
277 hint = corrupt_sae_checkpoint_hint(short, resolved_layer)
278 if hint:
279 raise ValueError(hint)
280
281 path = resolve_sae_checkpoint_path(short, resolved_layer)
282 if path is None:
283 raise FileNotFoundError(
285 short,
286 command="feature logit",
287 requested_layer=resolved_layer,
288 )
289 )
290
291 try:
292 return SparseAutoencoder.load(str(path), device=dev)
293 except Exception as exc:
294 from aquin.compute.torch_io import is_corrupt_pytorch_zip
295
296 cause = exc.__cause__ if exc.__cause__ is not None else exc
297 if is_corrupt_pytorch_zip(exc) or is_corrupt_pytorch_zip(cause):
298 raise ValueError(
299 corrupt_sae_checkpoint_hint(short, resolved_layer)
300 or (
301 f"SAE checkpoint at {path} is corrupted or incomplete. "
302 f"Re-download: aquin load sae {short}-l{resolved_layer}"
303 )
304 ) from exc
305 raise
306
307
308def resolve_sae_checkpoint_path(model_id: str, layer: int) -> Path | None:
309 """Return a loadable SAE checkpoint for model+layer, or None if missing/ambiguous."""
310 short = resolve_model_id(model_id)
311 resolved_layer = int(layer)
313 try:
314 from aquin.compute.user_sae import get_active_user_sae_path, list_user_saes
315
316 active = get_active_user_sae_path(short, resolved_layer)
317 if active is not None and _is_loadable_sae_file(active):
318 return active
319
320 catalog = resolve_sae_path(short, resolved_layer)
321 if _is_loadable_sae_file(catalog):
322 return catalog
323
324 user_matches = [
325 Path(str(r["path"]))
326 for r in list_user_saes(short)
327 if r.get("layer") == resolved_layer and r.get("path")
328 ]
329 existing = [p for p in user_matches if _is_loadable_sae_file(p)]
330 if len(existing) == 1:
331 return existing[0]
332 except Exception:
333 pass
334
335 return None
336
337
338def get_available_sae_layers(model_id: str) -> list[int]:
339 """SAE layer indices with loadable checkpoints on disk for this model."""
340 short = resolve_model_id(model_id)
341 found: set[int] = set()
343 for layer in get_catalog_sae_layers(short):
344 if resolve_sae_checkpoint_path(short, int(layer)) is not None:
345 found.add(int(layer))
346
347 try:
348 from aquin.compute.user_sae import list_user_saes
349
350 for row in list_user_saes(short):
351 layer_raw = row.get("layer")
352 if layer_raw is None:
353 continue
354 layer = int(layer_raw)
355 if resolve_sae_checkpoint_path(short, layer) is not None:
356 found.add(layer)
357 except Exception:
358 pass
359
360 sae_dir = Path.home() / ".aquin" / "sae" / short
361 if sae_dir.is_dir():
362 for path in sae_dir.glob("sae_layer*.pt"):
363 layer = _layer_from_sae_filename(path)
364 if layer is not None and path.is_file():
365 found.add(layer)
366
367 return sorted(found)
368
369
370def get_catalog_sae_layers(model_id: str) -> list[int]:
371 """All SAE layer indices published for this model (may not be downloaded yet)."""
372 short = resolve_model_id(model_id)
373 cfg = get_config(short)
374 layers: dict = cfg.get("sae_layers", {})
375 if layers:
376 return sorted(layers.keys())
377 return [int(cfg["sae_layer"])]
378
379
381 model_id: str,
382 *,
383 command: str,
384 layer_flag: str = "--layer",
385 example_suffix: str = "",
386 requested_layer: int | None = None,
387) -> str:
388 """Human-readable hint listing on-disk SAEs and aquin load sae pull commands."""
389 short = resolve_model_id(model_id)
390 available = get_available_sae_layers(short)
391 catalog = get_catalog_sae_layers(short)
392 not_downloaded = [layer for layer in catalog if layer not in available]
393 sae_dir = Path.home() / ".aquin" / "sae" / short
394
395 lines: list[str] = []
396 if requested_layer is not None:
397 corrupt = corrupt_sae_checkpoint_hint(short, requested_layer)
398 if corrupt:
399 lines.append(corrupt)
400 else:
401 lines.append(f"SAE layer {requested_layer} is not downloaded for {short}.")
402 else:
403 lines.append(f"No {layer_flag} specified for {command}.")
404 lines.append(f"Model: {short}")
405 lines.append("")
406
407 if available:
408 lines.append("Downloaded on this machine:")
409 for layer in available:
410 path = resolve_sae_checkpoint_path(short, layer) or resolve_sae_path(short, layer)
411 lines.append(f" {layer_flag} {layer} ({path})")
412 else:
413 lines.append("No SAE checkpoints found on disk.")
414 lines.append(f" Directory: {sae_dir}/")
415
416 lines.append("")
417 if available:
418 pick = requested_layer if requested_layer in available else available[0]
419 example = f" aquin {command}"
420 if example_suffix:
421 example += f" {example_suffix}"
422 example += f" {layer_flag} {pick}"
423 lines.append("Example:")
424 lines.append(example)
425
426 if not_downloaded:
427 lines.append("")
428 lines.append("Pull more layers:")
429 pull_layers = (
430 [requested_layer]
431 if requested_layer is not None and requested_layer not in available
432 else not_downloaded
433 )
434 for layer in pull_layers:
435 lines.append(f" aquin load sae {short}-l{layer}")
436
437 lines.append("")
438 lines.append(f"Try: aquin load sae {short}-l<n> or aquin load sae --path <weights.pt>")
439 lines.append("Self-hosted: aquin load sae --path <weights.pt> --layer <n> [--model <id>]")
440 return "\n".join(lines)
441
442
444 model_id: str,
445 layer: int | None,
446 *,
447 command: str = "trace",
448 example_suffix: str = "",
449) -> int:
450 """
451 Resolve an explicit SAE layer for tools that must not silently default.
452
453 Raises ValueError with an actionable message when layer is omitted or missing on disk.
454 """
455 short = resolve_model_id(model_id)
456 available = get_available_sae_layers(short)
457
458 if layer is not None:
459 resolved = int(layer)
460 if resolve_sae_checkpoint_path(short, resolved) is not None:
461 return resolved
462 raise ValueError(
464 short,
465 command=command,
466 example_suffix=example_suffix,
467 requested_layer=resolved,
468 )
469 )
470
471 if not available:
472 raise ValueError(
474 short,
475 command=command,
476 example_suffix=example_suffix,
477 )
478 )
479
480 raise ValueError(
482 short,
483 command=command,
484 example_suffix=example_suffix,
485 )
486 )
487
488
489def resolve_sae_path(model_id: str, layer: int | None = None) -> Path:
490 """Canonical on-disk path for a model SAE (existing file, or expected location)."""
491 short = resolve_model_id(model_id)
492 cfg = get_config(short)
493 sae_dir = Path.home() / ".aquin" / "sae" / short
494 resolved_layer = int(layer if layer is not None else cfg["sae_layer"])
495
496 candidates: list[Path] = []
497 sae_layers: dict = cfg.get("sae_layers", {})
498 if rel := sae_layers.get(resolved_layer):
499 candidates.append(sae_dir / Path(rel).name)
500 if layer is None and (fn := cfg.get("sae_filename")):
501 candidates.append(sae_dir / Path(fn).name)
502 candidates.append(sae_dir / f"sae_layer{resolved_layer}.pt")
503
504 seen: set[Path] = set()
505 unique: list[Path] = []
506 for p in candidates:
507 if p not in seen:
508 seen.add(p)
509 unique.append(p)
510
511 for p in unique:
512 if p.exists():
513 return p
514 return unique[0]
515
516
517def get_lora_target_modules(model_id: str) -> list[str]:
518 cfg = get_config(model_id)
519 return list(cfg.get("lora_target_modules", ["q_proj", "v_proj"]))
520
522def infer_lora_target_modules(model: Any) -> list[str]:
523 """Pick LoRA targets from module names when config defaults do not match."""
524 leaf_names = {name.split(".")[-1] for name, _ in model.named_modules()}
525 presets: list[list[str]] = [
526 ["q_proj", "v_proj"],
527 ["query_key_value", "dense"],
528 ["query_key_value"],
529 ["c_attn", "c_fc"],
530 ["Wqkv", "out_proj"],
531 ]
532 for candidates in presets:
533 if all(c in leaf_names for c in candidates):
534 return candidates
535 for candidates in presets:
536 hit = [c for c in candidates if c in leaf_names]
537 if hit:
538 return hit
539 return ["q_proj", "v_proj"]
540
541
542def get_sae_source(model_id: str) -> str:
543 return get_config(model_id).get("sae_source", "native")
544
545
547
548def get_loaded_model() -> Any | None:
549 """Return the most recently used loaded model, or None."""
550 if not _models:
551 return None
552 return next(reversed(_models.values()))
553
554
555def _tl_not_in_catalog(exc: BaseException) -> bool:
556 msg = str(exc).lower()
557 return "not found" in msg or "valid official model names" in msg
558
560def _tl_conversion_failed(exc: BaseException) -> bool:
561 """True when TransformerLens cannot build/wrap this HF architecture (soft-fallback)."""
562 if _tl_not_in_catalog(exc):
563 return True
564 msg = str(exc).lower()
565 # e.g. GPTNeoXForCausalLM has no attribute 'embed_out' on some TL + transformers combos
566 if "embed_out" in msg:
567 return True
568 if "has no attribute" in msg and (
569 "neox" in msg or "gptneo" in msg or "gpt_neox" in msg or "pythia" in msg
570 ):
571 return True
572 if isinstance(exc, AttributeError) and ("embed" in msg or "unembed" in msg):
573 return True
574 return False
575
576
578 hf_name: str,
579 cfg: dict[str, Any],
580 hf_model: Any,
581 *,
582 dtype: Any,
583 device: str,
584) -> Any:
585 """Wrap a loaded HuggingFace causal LM as HookedTransformer or HfLlmShim."""
586 if cfg.get("hf_only"):
587 from aquin.compute.hf_llm_shim import HfLlmShim
588 from transformers import AutoTokenizer
589
590 trust = bool(cfg.get("trust_remote_code", False))
591 tokenizer = AutoTokenizer.from_pretrained(hf_name, trust_remote_code=trust)
592 n_heads = int(cfg.get("n_heads", 32))
593 return HfLlmShim(
594 hf_model,
595 tokenizer,
596 hf_name=hf_name,
597 n_layers=int(cfg["n_layers"]),
598 d_model=int(cfg["d_model"]),
599 n_heads=n_heads,
600 )
601
602 from transformer_lens import HookedTransformer
603
604 trust = bool(cfg.get("trust_remote_code", False))
605 try:
606 tl = HookedTransformer.from_pretrained(
607 hf_name,
608 hf_model=hf_model,
609 dtype=dtype,
610 device=device,
611 trust_remote_code=trust,
612 )
613 except (
614 TypeError,
615 ValueError,
616 KeyError,
617 RuntimeError,
618 NotImplementedError,
619 AttributeError,
620 ) as exc:
621 if not _tl_conversion_failed(exc):
622 raise
623 from aquin.compute.hf_llm_shim import HfLlmShim
624 from transformers import AutoTokenizer
625
626 print(
627 f"[model] TransformerLens cannot wrap {hf_name} ({type(exc).__name__}: {exc}) "
628 "— using HuggingFace shim…",
629 flush=True,
630 )
631 tokenizer = AutoTokenizer.from_pretrained(hf_name, trust_remote_code=trust)
632 n_heads = int(cfg.get("n_heads", 32))
633 return HfLlmShim(
634 hf_model,
635 tokenizer,
636 hf_name=hf_name,
637 n_layers=int(cfg["n_layers"]),
638 d_model=int(cfg["d_model"]),
639 n_heads=n_heads,
640 )
641 tl.eval()
642 return tl
643
644
645def _build_hooked_transformer(cfg: dict[str, Any], *, dtype: Any, device: str) -> Any:
646 """Load HookedTransformer, or HfLlmShim for custom HF-only architectures."""
647 from aquin.compute.hf_auth import require_hf_hub_auth
648
649 hf_name = cfg["hf_name"]
650 require_hf_hub_auth(hf_repo=hf_name)
651
652 if cfg.get("hf_only"):
653 from aquin.compute.hf_llm_shim import HfLlmShim
654
655 print(f"[model] {hf_name} — loading via HuggingFace (no TransformerLens wrapper)…", flush=True)
656 return HfLlmShim.from_pretrained(hf_name, cfg, dtype=dtype, device=device)
657
658 from transformers import AutoModelForCausalLM
659
660 trust = bool(cfg.get("trust_remote_code", False))
661 hf_first = bool(cfg.get("hf_first", False))
662
663 if not hf_first:
664 from transformer_lens import HookedTransformer
665
666 try:
667 return HookedTransformer.from_pretrained(
668 hf_name,
669 dtype=dtype,
670 trust_remote_code=trust,
671 )
672 except (
673 ValueError,
674 KeyError,
675 RuntimeError,
676 NotImplementedError,
677 AttributeError,
678 TypeError,
679 ) as exc:
680 if not _tl_conversion_failed(exc):
681 raise
682 print(
683 f"[model] {hf_name} cannot use TransformerLens directly "
684 f"({type(exc).__name__}: {exc}) — loading via HuggingFace…",
685 flush=True,
686 )
687
688 hf_model = AutoModelForCausalLM.from_pretrained(
689 hf_name,
690 torch_dtype=dtype,
691 device_map=device,
692 trust_remote_code=trust,
693 )
694 hf_model.eval()
695 return _wrap_hf_causal_lm(hf_name, cfg, hf_model, dtype=dtype, device=device)
696
697
698def load_model(model_id: str) -> Any:
699 """
700 Load a HookedTransformer model by short slug or HF name.
701 Uses an LRU cache — at most MAX_LOADED_MODELS kept in VRAM.
702 Raises ComputeNotAvailableError when no accelerator (unless AQUIN_ALLOW_CPU=1).
703
704 Serialized: never interleave two builds (MPS unified RAM doubles fast).
705 """
706 with _model_io_lock:
707 return _load_model_unlocked(model_id)
708
709
710def _load_model_unlocked(model_id: str) -> Any:
711 import torch
712 from aquin.compute.device import (
713 default_dtype_for_device,
714 empty_device_cache,
715 is_oom_error,
716 require_load_device,
717 )
718
719 short = resolve_model_id(model_id)
720
721 if short in _models:
722 _models.move_to_end(short)
723 _save_active_model(short)
724 print(f"[model] cache hit: {short}", flush=True)
725 return _models[short]
726
727 # Avoid VRAM duplication: if a persistent daemon holds a model in another
728 # process, free it before we allocate here. The daemon reloads lazily later.
729 if os.environ.get("AQUIN_DAEMON") != "1":
730 try:
731 from aquin.compute.model_runtime import release_foreign_daemon
732
733 release_foreign_daemon()
734 except Exception:
735 pass
736
737 cfg = get_config(short)
738 try:
739 device = require_load_device(short, cfg)
740 except RuntimeError as exc:
741 raise ComputeNotAvailableError(str(exc)) from exc
742 dtype = default_dtype_for_device(device)
743
744 if not cfg.get("hf_only"):
745 try:
746 from transformer_lens import HookedTransformer # noqa: F401
747 except ImportError:
749 "transformer_lens is not installed. "
750 "Reinstall Aquin: pip install -U aquin"
751 )
752
753 from aquin.compute.hf_auth import gated_repo_hint, require_hf_hub_auth
754
755 try:
756 require_hf_hub_auth(hf_repo=cfg["hf_name"])
757 except RuntimeError as exc:
758 raise ComputeNotAvailableError(str(exc)) from exc
759
760 import time
761
762 print(f"[model] loading {cfg['hf_name']} ({dtype} on {device})...", flush=True)
763 start = time.time()
764 stop_ticker = threading.Event()
765
766 def _tick() -> None:
767 while not stop_ticker.wait(5):
768 print(f"[model] still loading... {int(time.time() - start)}s", flush=True)
769
770 threading.Thread(target=_tick, daemon=True).start()
771 try:
772 m = _build_hooked_transformer(cfg, dtype=dtype, device=device)
773 except Exception as exc:
774 msg = str(exc).lower()
775 if "gated" in msg or "401" in msg or "unauthorized" in msg:
776 raise ComputeNotAvailableError(gated_repo_hint(cfg["hf_name"])) from exc
777 if is_oom_error(exc):
778 hint = ""
779 if short == "sarvam-30b" or int(cfg.get("d_model", 0)) >= 4096:
780 hint = (
781 " Large MoE models like Sarvam 30B store all expert weights "
782 "(~60GB+ in bf16); use a GPU with enough VRAM."
783 )
785 f"Accelerator ran out of memory loading '{short}' on {device}.{hint}"
786 ) from exc
787 raise
788 finally:
789 stop_ticker.set()
790
791 # Drop before install if the UI cancelled / unload raced us.
792 try:
793 from aquin.compute.model_runtime import is_load_cancelled
794
795 if is_load_cancelled():
796 del m
797 empty_device_cache(device)
798 print(f"[model] discarded {short} (load cancelled)", flush=True)
799 raise RuntimeError(f"Model load cancelled ({short})")
800 except ImportError:
801 pass
802
803 m.eval()
804 from aquin.compute.hf_llm_shim import HfLlmShim
805
806 if not isinstance(m, HfLlmShim):
807 m.to(device)
808 print(f"[model] loaded in {int(time.time() - start)}s", flush=True)
809
810 _models[short] = m
811 _save_active_model(short)
812
813 # Evict previous model — only one LLM in VRAM at a time
814 while len(_models) > MAX_LOADED_MODELS:
815 evicted_id, evicted_model = _models.popitem(last=False)
816 del evicted_model
817 empty_device_cache(device)
818 print(f"[model] unloaded {evicted_id} from VRAM", flush=True)
819
820 print(f"[model] {short} ready. (only model in VRAM)", flush=True)
821 return _models[short]
822
823
824def load_sae(model: Any, layer: int, model_id: str, sae_dir: Path | None = None) -> Any:
825 """
826 Load a SparseAutoencoder for the given layer from ~/.aquin/sae/ or user bindings.
827 Use `aquin load sae` (Aquin catalog) or `aquin load sae --path` for local files.
828 """
829 from aquin.compute.device import resolve_compute_device
830
831 short = resolve_model_id(model_id)
832 resolved_layer = int(layer)
833
834 cache_key = (short, resolved_layer)
835 cached = _sae_cache.get(cache_key)
836 if cached is not None:
837 return cached
838
839 device = resolve_compute_device()
840 sae = load_sae_from_disk(short, resolved_layer, device=device)
841 print(f"[sae] loaded layer {resolved_layer} for {short}", flush=True)
842 _sae_cache[cache_key] = sae
843 return sae
bool _is_loadable_sae_file(Path path)
bool _tl_conversion_failed(BaseException exc)
str|None get_active_model_id(*, bool allow_daemon=True)
list[int] get_available_sae_layers(str model_id)
Any _load_model_unlocked(str model_id)
None _save_active_model(str model_id)
str resolve_model_id(str model_id)
str get_sae_source(str model_id)
dict get_config(str model_id)
str get_hf_name(str model_id)
int require_sae_layer(str model_id, int|None layer, *, str command="trace", str example_suffix="")
Any _build_hooked_transformer(dict[str, Any] cfg, *, Any dtype, str device)
None evict_sae_cache(str model_id, int|None layer=None)
Path|None resolve_sae_checkpoint_path(str model_id, int layer)
Any load_model(str model_id)
list[str] get_lora_target_modules(str model_id)
Any _wrap_hf_causal_lm(str hf_name, dict[str, Any] cfg, Any hf_model, *, Any dtype, str device)
int get_sae_layer(str model_id)
int|None _layer_from_sae_filename(Path path)
bool _tl_not_in_catalog(BaseException exc)
Any load_sae(Any model, int layer, str model_id, Path|None sae_dir=None)
list[int] get_catalog_sae_layers(str model_id)
str|None corrupt_sae_checkpoint_hint(str model_id, int layer)
Path resolve_sae_path(str model_id, int|None layer=None)
list[str] infer_lora_target_modules(Any model)
Any load_sae_from_disk(str model_id, int layer, *, str|None device=None)
None reload_vram_for_model(str model_id)
str format_sae_layer_choice_message(str model_id, *, str command, str layer_flag="--layer", str example_suffix="", int|None requested_layer=None)