AQIT 0.1.0
Loading...
Searching...
No Matches
aquin.compute.model_loader Namespace Reference

Classes

class  ComputeNotAvailableError

Functions

None clear_sae_cache ()
None evict_sae_cache (str model_id, int|None layer=None)
None _save_active_model (str model_id)
str|None get_active_model_id (*, bool allow_daemon=True)
None reload_vram_for_model (str model_id)
None clear_active_model_file ()
None clear_llm_models ()
str|None get_loaded_llm_id ()
str resolve_model_id (str model_id)
dict get_config (str model_id)
str get_hf_name (str model_id)
int get_sae_layer (str model_id)
int|None _layer_from_sae_filename (Path path)
bool _is_loadable_sae_file (Path path)
str|None corrupt_sae_checkpoint_hint (str model_id, int layer)
Any load_sae_from_disk (str model_id, int layer, *, str|None device=None)
Path|None resolve_sae_checkpoint_path (str model_id, int layer)
list[int] get_available_sae_layers (str model_id)
list[int] get_catalog_sae_layers (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)
int require_sae_layer (str model_id, int|None layer, *, str command="trace", str example_suffix="")
Path resolve_sae_path (str model_id, int|None layer=None)
list[str] get_lora_target_modules (str model_id)
list[str] infer_lora_target_modules (Any model)
str get_sae_source (str model_id)
Any|None get_loaded_model ()
bool _tl_not_in_catalog (BaseException exc)
bool _tl_conversion_failed (BaseException exc)
Any _wrap_hf_causal_lm (str hf_name, dict[str, Any] cfg, Any hf_model, *, Any dtype, str device)
Any _build_hooked_transformer (dict[str, Any] cfg, *, Any dtype, str device)
Any load_model (str model_id)
Any _load_model_unlocked (str model_id)
Any load_sae (Any model, int layer, str model_id, Path|None sae_dir=None)

Variables

dict MODEL_CONFIGS
dict _HF_TO_SHORT
int MAX_LOADED_MODELS = 1
OrderedDict _models = OrderedDict()
 _model_io_lock = threading.RLock()
dict _sae_cache = {}
str _ACTIVE_MODEL_PATH = Path.home() / ".aquin" / "active_model.txt"

Function Documentation

◆ _build_hooked_transformer()

Any _build_hooked_transformer ( dict[str, Any] cfg,
* ,
Any dtype,
str device )
protected
Load HookedTransformer, or HfLlmShim for custom HF-only architectures.

Definition at line 649 of file model_loader.py.

References _tl_conversion_failed(), and _wrap_hf_causal_lm().

Referenced by _load_model_unlocked().

◆ _is_loadable_sae_file()

bool _is_loadable_sae_file ( Path path)
protected

Definition at line 242 of file model_loader.py.

Referenced by corrupt_sae_checkpoint_hint(), and resolve_sae_checkpoint_path().

◆ _layer_from_sae_filename()

int | None _layer_from_sae_filename ( Path path)
protected

Definition at line 232 of file model_loader.py.

Referenced by get_available_sae_layers().

◆ _load_model_unlocked()

Any _load_model_unlocked ( str model_id)
protected

◆ _save_active_model()

None _save_active_model ( str model_id)
protected

Definition at line 138 of file model_loader.py.

Referenced by _load_model_unlocked().

◆ _tl_conversion_failed()

bool _tl_conversion_failed ( BaseException exc)
protected
True when TransformerLens cannot build/wrap this HF architecture (soft-fallback).

Definition at line 564 of file model_loader.py.

References _tl_not_in_catalog().

Referenced by _build_hooked_transformer(), and _wrap_hf_causal_lm().

◆ _tl_not_in_catalog()

bool _tl_not_in_catalog ( BaseException exc)
protected

Definition at line 559 of file model_loader.py.

Referenced by _tl_conversion_failed().

◆ _wrap_hf_causal_lm()

Any _wrap_hf_causal_lm ( str hf_name,
dict[str, Any] cfg,
Any hf_model,
* ,
Any dtype,
str device )
protected
Wrap a loaded HuggingFace causal LM as HookedTransformer or HfLlmShim.

Definition at line 581 of file model_loader.py.

References _tl_conversion_failed().

Referenced by _build_hooked_transformer().

◆ clear_active_model_file()

None clear_active_model_file ( )

Definition at line 173 of file model_loader.py.

◆ clear_llm_models()

None clear_llm_models ( )

Definition at line 178 of file model_loader.py.

Referenced by reload_vram_for_model().

◆ clear_sae_cache()

None clear_sae_cache ( )

Definition at line 124 of file model_loader.py.

◆ corrupt_sae_checkpoint_hint()

str | None corrupt_sae_checkpoint_hint ( str model_id,
int layer )
If a file exists on disk but cannot be loaded, return a re-download hint.

Definition at line 248 of file model_loader.py.

References _is_loadable_sae_file(), resolve_model_id(), and resolve_sae_path().

Referenced by format_sae_layer_choice_message(), and load_sae_from_disk().

◆ evict_sae_cache()

None evict_sae_cache ( str model_id,
int | None layer = None )
Drop a cached SAE so the next load re-reads from disk (after a rebind).

Definition at line 128 of file model_loader.py.

References resolve_model_id().

◆ format_sae_layer_choice_message()

str format_sae_layer_choice_message ( str model_id,
* ,
str command,
str layer_flag = "--layer",
str example_suffix = "",
int | None requested_layer = None )
Human-readable hint listing on-disk SAEs and aquin load sae pull commands.

Definition at line 384 of file model_loader.py.

References corrupt_sae_checkpoint_hint(), get_available_sae_layers(), get_catalog_sae_layers(), resolve_model_id(), resolve_sae_checkpoint_path(), and resolve_sae_path().

Referenced by load_sae_from_disk(), and require_sae_layer().

◆ get_active_model_id()

str | None get_active_model_id ( * ,
bool allow_daemon = True )
Return the model slug written by the last successful `aquin load`.

Definition at line 143 of file model_loader.py.

◆ get_available_sae_layers()

list[int] get_available_sae_layers ( str model_id)
SAE layer indices with loadable checkpoints on disk for this model.

Definition at line 342 of file model_loader.py.

References _layer_from_sae_filename(), get_catalog_sae_layers(), resolve_model_id(), and resolve_sae_checkpoint_path().

Referenced by format_sae_layer_choice_message(), and require_sae_layer().

◆ get_catalog_sae_layers()

list[int] get_catalog_sae_layers ( str model_id)
All SAE layer indices published for this model (may not be downloaded yet).

Definition at line 374 of file model_loader.py.

References get_config(), and resolve_model_id().

Referenced by format_sae_layer_choice_message(), and get_available_sae_layers().

◆ get_config()

◆ get_hf_name()

str get_hf_name ( str model_id)
HuggingFace repo id for transformers.from_pretrained (resolves Aquin short slugs).

Definition at line 223 of file model_loader.py.

References get_config().

◆ get_loaded_llm_id()

str | None get_loaded_llm_id ( )

Definition at line 189 of file model_loader.py.

Referenced by reload_vram_for_model().

◆ get_loaded_model()

Any | None get_loaded_model ( )
Return the most recently used loaded model, or None.

Definition at line 552 of file model_loader.py.

◆ get_lora_target_modules()

list[str] get_lora_target_modules ( str model_id)

Definition at line 521 of file model_loader.py.

References get_config().

◆ get_sae_layer()

int get_sae_layer ( str model_id)

Definition at line 228 of file model_loader.py.

References get_config().

◆ get_sae_source()

str get_sae_source ( str model_id)

Definition at line 546 of file model_loader.py.

References get_config().

◆ infer_lora_target_modules()

list[str] infer_lora_target_modules ( Any model)
Pick LoRA targets from module names when config defaults do not match.

Definition at line 526 of file model_loader.py.

◆ load_model()

Any load_model ( str model_id)
Load a HookedTransformer model by short slug or HF name.
Uses an LRU cache :  at most MAX_LOADED_MODELS kept in VRAM.
Raises ComputeNotAvailableError when no accelerator (unless AQUIN_ALLOW_CPU=1).

Serialized: never interleave two builds (MPS unified RAM doubles fast).

Definition at line 702 of file model_loader.py.

References _load_model_unlocked().

Referenced by reload_vram_for_model().

◆ load_sae()

Any load_sae ( Any model,
int layer,
str model_id,
Path | None sae_dir = None )
Load a SparseAutoencoder for the given layer from ~/.aquin/sae/ or user bindings.
Use `aquin load sae` (Aquin catalog) or `aquin load sae --path` for local files.

Definition at line 828 of file model_loader.py.

References load_sae_from_disk(), and resolve_model_id().

◆ load_sae_from_disk()

Any load_sae_from_disk ( str model_id,
int layer,
* ,
str | None device = None )
Load a validated on-disk SAE or raise with a re-download hint.

Definition at line 265 of file model_loader.py.

References corrupt_sae_checkpoint_hint(), format_sae_layer_choice_message(), resolve_model_id(), and resolve_sae_checkpoint_path().

Referenced by load_sae().

◆ reload_vram_for_model()

None reload_vram_for_model ( str model_id)
Reload the given model into VRAM (used after heavy jobs release the GPU).

Definition at line 165 of file model_loader.py.

References clear_llm_models(), get_loaded_llm_id(), load_model(), and resolve_model_id().

◆ require_sae_layer()

int require_sae_layer ( str model_id,
int | None layer,
* ,
str command = "trace",
str example_suffix = "" )
Resolve an explicit SAE layer for tools that must not silently default.

Raises ValueError with an actionable message when layer is omitted or missing on disk.

Definition at line 447 of file model_loader.py.

References format_sae_layer_choice_message(), get_available_sae_layers(), resolve_model_id(), and resolve_sae_checkpoint_path().

◆ resolve_model_id()

◆ resolve_sae_checkpoint_path()

Path | None resolve_sae_checkpoint_path ( str model_id,
int layer )
Return a loadable SAE checkpoint for model+layer, or None if missing/ambiguous.

Definition at line 312 of file model_loader.py.

References _is_loadable_sae_file(), resolve_model_id(), and resolve_sae_path().

Referenced by format_sae_layer_choice_message(), get_available_sae_layers(), load_sae_from_disk(), and require_sae_layer().

◆ resolve_sae_path()

Path resolve_sae_path ( str model_id,
int | None layer = None )
Canonical on-disk path for a model SAE (existing file, or expected location).

Definition at line 493 of file model_loader.py.

References get_config(), and resolve_model_id().

Referenced by corrupt_sae_checkpoint_hint(), format_sae_layer_choice_message(), and resolve_sae_checkpoint_path().

Variable Documentation

◆ _ACTIVE_MODEL_PATH

str aquin.compute.model_loader._ACTIVE_MODEL_PATH = Path.home() / ".aquin" / "active_model.txt"
protected

Definition at line 121 of file model_loader.py.

◆ _HF_TO_SHORT

dict aquin.compute.model_loader._HF_TO_SHORT
protected
Initial value:
= {
cfg["hf_name"].lower(): key for key, cfg in MODEL_CONFIGS.items()
}

Definition at line 106 of file model_loader.py.

◆ _model_io_lock

aquin.compute.model_loader._model_io_lock = threading.RLock()
protected

Definition at line 115 of file model_loader.py.

◆ _models

OrderedDict aquin.compute.model_loader._models = OrderedDict()
protected

Definition at line 111 of file model_loader.py.

◆ _sae_cache

dict aquin.compute.model_loader._sae_cache = {}
protected

Definition at line 119 of file model_loader.py.

◆ MAX_LOADED_MODELS

int aquin.compute.model_loader.MAX_LOADED_MODELS = 1

Definition at line 110 of file model_loader.py.

◆ MODEL_CONFIGS

dict aquin.compute.model_loader.MODEL_CONFIGS

Definition at line 19 of file model_loader.py.