Spaces:
Sleeping
Sleeping
File size: 5,826 Bytes
bda104d b0af84d bda104d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | """Loading trained SAEs for the steering API.
A :class:`SteerableSAE` is the minimal model surface the :class:`~steerable_retrieval.steer.slider.Slider`
needs: a trained ``sae_encoder`` + ``sae_decoder`` and a ``text_encoder`` mapping concept
strings to embeddings in the same joint space. This module resolves such a model from a
checkpoint (local path or HuggingFace repo).
NOTE: the released pretrained checkpoint (BatchTopK SAE on MuQ / music4all) is not published
yet. Until then, construct a model in-memory and pass it via ``Slider(..., model=model)``.
"""
from __future__ import annotations
from typing import Optional
import torch
import torch.nn as nn
class SteerableSAE(nn.Module):
"""Minimal SAE wrapper for inference-time steering (no training deps).
Args:
sae_encoder: a trained SAE encoder (``steerable_retrieval.models.sae.encoders``).
sae_decoder: a trained SAE decoder (``steerable_retrieval.models.sae.decoders``).
text_encoder: callable mapping ``list[str] -> [B, d]`` embeddings (e.g. the MuQ
text tower), living in the same joint space as the SAE's audio inputs.
"""
def __init__(self, sae_encoder, sae_decoder, text_encoder):
super().__init__()
self.sae_encoder = sae_encoder
self.sae_decoder = sae_decoder
self.text_encoder = text_encoder
if not hasattr(self.sae_decoder, "b_dec"):
self.sae_decoder.b_dec = self.sae_encoder.b_dec
@torch.no_grad()
def inference(self, x):
"""Encode dense features to a sparse code and back. Returns ``(x, z, xhat, pre)``."""
pre, z = self.sae_encoder(x)
xhat = self.sae_decoder(z)
return x, z, xhat, pre
def _resolve_local_run(checkpoint: str, subfolder: Optional[str] = None):
"""Return (ckpt_path, config_path) for a checkpoint, downloading from the Hub if
``checkpoint`` is a ``org/repo`` id rather than a local path.
A Lightning run stores its Hydra config at ``<run>/.hydra/config.yaml`` and its
checkpoints under ``<run>/checkpoints/``. We use that config to rebuild the SAE
modules before loading the (encoder/decoder-only) weights. For a Hub repo carrying
several models, ``subfolder`` (e.g. ``"L0-20"``) selects which one.
"""
import os
if os.path.exists(checkpoint):
ckpt_path = os.path.abspath(checkpoint)
run_dir = os.path.dirname(os.path.dirname(ckpt_path)) # .../checkpoints/x.ckpt -> run
cfg_path = os.path.join(run_dir, ".hydra", "config.yaml")
if not os.path.exists(cfg_path):
# allow a config.yaml sitting next to the checkpoint
alt = os.path.join(os.path.dirname(ckpt_path), "config.yaml")
cfg_path = alt if os.path.exists(alt) else cfg_path
return ckpt_path, cfg_path
# Otherwise treat it as a HuggingFace repo id: expects last.ckpt + config.yaml
# (optionally under `subfolder`).
from huggingface_hub import hf_hub_download
pre = f"{subfolder}/" if subfolder else ""
ckpt_path = hf_hub_download(checkpoint, filename=f"{pre}last.ckpt")
try:
cfg_path = hf_hub_download(checkpoint, filename=f"{pre}config.yaml")
except Exception:
cfg_path = hf_hub_download(checkpoint, filename=f"{pre}.hydra/config.yaml")
return ckpt_path, cfg_path
def load_steerable_sae(
checkpoint: Optional[str],
*,
model_class: Optional[str] = None,
device: str = "cpu",
text_encoder=None,
build_text_encoder: bool = True,
config_path: Optional[str] = None,
subfolder: Optional[str] = None,
) -> SteerableSAE:
"""Resolve a :class:`SteerableSAE` from a trained Lightning checkpoint.
Args:
checkpoint: local path to a ``.ckpt`` (its run's ``.hydra/config.yaml`` is used
to rebuild the SAE modules), or a HuggingFace ``org/repo`` id carrying
``last.ckpt`` + ``config.yaml``.
device: where to place the model.
text_encoder: a ready callable ``list[str] -> [B, d]``. If ``None`` and
``build_text_encoder`` is True, the text tower from the run config (e.g.
MuQ-MuLan) is instantiated; if False, ``text_encoder`` stays ``None`` (useful
for steering/retrieval that never embeds new text).
config_path: override the run config path explicitly.
"""
from hydra.utils import instantiate
from omegaconf import OmegaConf
ckpt_path, resolved_cfg = _resolve_local_run(checkpoint, subfolder=subfolder)
cfg = OmegaConf.load(config_path or resolved_cfg)
enc = instantiate(cfg.model.sae_encoder, device=device)
dec = instantiate(cfg.model.sae_decoder, device=device)
if not hasattr(dec, "b_dec"):
dec.b_dec = enc.b_dec
# Lightning checkpoints from our training runs can contain OmegaConf metadata
# alongside tensors. PyTorch 2.6 defaults torch.load(weights_only=True), which
# rejects that metadata; this loader is for trusted project checkpoints.
state = torch.load(ckpt_path, map_location=device, weights_only=False)
sd = state.get("state_dict", state)
enc_sd = {k[len("sae_encoder."):]: v for k, v in sd.items() if k.startswith("sae_encoder.")}
dec_sd = {k[len("sae_decoder."):]: v for k, v in sd.items() if k.startswith("sae_decoder.")}
if not enc_sd or not dec_sd:
raise ValueError(
f"No sae_encoder/sae_decoder weights found in {ckpt_path}. "
f"Available prefixes: {sorted({k.split('.')[0] for k in sd})}"
)
enc.load_state_dict(enc_sd, strict=False)
dec.load_state_dict(dec_sd, strict=False)
if text_encoder is None and build_text_encoder:
text_encoder = instantiate(cfg.model.text_encoder, device=device)
model = SteerableSAE(enc, dec, text_encoder)
model.to(device).eval()
return model
|