Pliploop's picture
Upload folder using huggingface_hub
9d8127c verified
Raw
History Blame Contribute Delete
13 kB
"""Shared engine for the demo Space and website example generation.
Loads a trained BatchTopK SAE + MuQ text tower and a music4all corpus, then steers a
seed track toward free-form concepts and retrieves nearest neighbours. The SAE,
inversion, edit masks, and retrieval corpus stay on CPU; ZeroGPU is used only while
embedding a new text concept for slider instantiation.
Configuration (env):
SSR_CHECKPOINT local .ckpt path, or a HF model repo id (default: local L0=20 run)
SSR_L0 subfolder in the model repo to load, e.g. "L0-20" (Hub repos only)
SSR_CORPUS_REPO HF dataset repo id holding corpus.npz + meta.json (overrides local dir)
SSR_CORPUS_DIR local corpus dir (default: demo/corpus)
SSR_DEVICE "cuda" / "cpu" (auto-detected)
"""
from __future__ import annotations
import json
import os
import time
import numpy as np
import torch
import torch.nn.functional as F
from steerable_retrieval.steer import Slider
from steerable_retrieval.steer.loading import load_steerable_sae
from steerable_retrieval.steer.steering import decode_normalized
HERE = os.path.dirname(__file__)
CORPUS_DIR = os.environ.get("SSR_CORPUS_DIR", os.path.join(HERE, "corpus"))
CORPUS_REPO = os.environ.get("SSR_CORPUS_REPO") # HF dataset repo id, if hosted
CHECKPOINT = os.environ.get("SSR_CHECKPOINT", os.path.join(HERE, os.pardir, "logs/xps/66a0caf9/checkpoints/last.ckpt"))
L0_SUBFOLDER = os.environ.get("SSR_L0") # e.g. "L0-20" when CHECKPOINT is a Hub repo
CONFIG = os.environ.get("SSR_CONFIG")
INVERSION_METHOD = os.environ.get("SSR_INVERSION_METHOD", "fista")
DEBUG = os.environ.get("SSR_DEBUG", "1").lower() not in {"0", "false", "no", "off"}
def _log(msg, *args):
if DEBUG:
print("[ssr-core] " + msg.format(*args), flush=True)
def _stamp(label: str) -> str:
return f"{label}-{int(time.time() * 1000) % 100000}"
def device() -> str:
return os.environ.get("SSR_DEVICE") or ("cuda" if torch.cuda.is_available() else "cpu")
def load_corpus():
"""Return (embeddings[np], track_ids[list], meta[dict]). Fetches from the HF dataset
repo if SSR_CORPUS_REPO is set, else reads the local corpus dir. CPU only."""
if CORPUS_REPO:
from huggingface_hub import hf_hub_download
npz_path = hf_hub_download(CORPUS_REPO, filename="corpus.npz", repo_type="dataset")
meta_path = hf_hub_download(CORPUS_REPO, filename="meta.json", repo_type="dataset")
else:
npz_path = os.path.join(CORPUS_DIR, "corpus.npz")
meta_path = os.path.join(CORPUS_DIR, "meta.json")
npz = np.load(npz_path, allow_pickle=True)
with open(meta_path) as fh:
meta = json.load(fh)
return npz["embeddings"].astype(np.float32), [str(t) for t in npz["track_ids"].tolist()], meta
class DemoEngine:
def __init__(self, corpus=None, dev: str | None = None):
# Keep SAE inversion, sparse edits, and retrieval on CPU. ZeroGPU is only
# needed while embedding new text concepts with MuQ.
self.device = "cpu"
self.text_device = dev or device()
embs, ids, meta = corpus if corpus is not None else load_corpus()
self.embeddings = torch.from_numpy(embs).float()
self.gallery = F.normalize(self.embeddings, dim=-1)
self.track_ids = ids
self.id_to_idx = {t: i for i, t in enumerate(ids)}
self.meta = meta
self.model = load_steerable_sae(CHECKPOINT, device="cpu", subfolder=L0_SUBFOLDER, config_path=CONFIG)
self._slider_cache: dict[str, Slider] = {}
self._text_encoder_device = "cpu"
def _slider_key(self, concept: str) -> str:
return " ".join(concept.strip().lower().split())
def has_slider(self, concept: str) -> bool:
return self._slider_key(concept) in self._slider_cache
def _set_text_encoder_device(self, target: str) -> None:
if self._text_encoder_device == target:
return
text_encoder = self.model.text_encoder
if text_encoder is None:
return
text_encoder.to(target)
if hasattr(text_encoder, "device"):
text_encoder.device = target
self._text_encoder_device = target
def _slider(self, concept: str, *, create: bool = True) -> Slider:
key = self._slider_key(concept)
if key not in self._slider_cache:
if not create:
raise KeyError(f"Slider for concept {concept!r} has not been created yet.")
target = self.text_device if str(self.text_device).startswith("cuda") and torch.cuda.is_available() else "cpu"
_log("slider-create concept={!r} key={!r} text_target={} inversion_method={}", concept, key, target, INVERSION_METHOD)
self._set_text_encoder_device(target)
try:
self._slider_cache[key] = Slider(concept.strip(), model=self.model, method=INVERSION_METHOD, device="cpu")
finally:
# The cached slider stores its text embedding/mask; live edits no
# longer need the text tower or any GPU memory.
self._set_text_encoder_device("cpu")
slider = self._slider_cache[key]
mask = slider.mask.detach().cpu()
_log(
"slider-ready concept={!r} support={} mask_norm={:.6f} mask_max={:.6f} text_cos={:.4f}",
concept,
len(slider),
float(mask.norm().item()),
float(mask.abs().max().item()) if mask.numel() else 0.0,
float(getattr(slider.inversion, "final_text_cosine", float("nan"))),
)
return self._slider_cache[key]
def track_meta(self, track_id: str) -> dict:
m = dict(self.meta.get(track_id, {}))
m["track_id"] = track_id
return m
def seed_embedding(self, track_id: str) -> torch.Tensor:
return self.embeddings[self.id_to_idx[track_id]]
def _retrieve_dense(self, query: torch.Tensor, k: int, *, exclude_idx: int | None = None):
query = F.normalize(query.reshape(-1), dim=0)
sims = self.gallery @ query
if exclude_idx is not None:
sims[exclude_idx] = -1e9
n_avail = self.gallery.shape[0] - (1 if exclude_idx is not None else 0)
vals, idx = torch.topk(sims, k=min(int(k), int(n_avail)), largest=True, sorted=True)
return idx.detach().cpu(), vals.detach().cpu()
def steer_and_retrieve(self, seed_track_id: str, concept: str, alpha: float = 1.0, k: int = 8) -> list[dict]:
slider = self._slider(concept)
z = self.seed_embedding(seed_track_id)
seed_idx = self.id_to_idx[seed_track_id]
edited = slider.steer(z, alpha=alpha)
idx, scores = self._retrieve_dense(edited, k, exclude_idx=seed_idx)
return self._format_results(idx, scores)
def multi_steer_and_retrieve(self, seed_track_id: str, sliders: list[tuple[str, float]], k: int = 8) -> list[dict]:
"""Apply several concept sliders to one seed query, then retrieve once.
Each concept is inverted/cached independently, but the edit itself is additive
in sparse SAE space: encode the seed once, add every active slider mask scaled
by alpha, decode once, and search from the combined query.
"""
call = _stamp("multi")
z = self.seed_embedding(seed_track_id).reshape(1, -1)
seed_idx = self.id_to_idx[seed_track_id]
slider_masks = []
mask_debug = []
for concept, alpha in sliders:
concept = concept.strip()
alpha = float(alpha)
if not concept or abs(alpha) < 1e-6:
continue
mask = self._slider(concept, create=False).mask
slider_masks.append((mask, alpha))
mask_debug.append({
"concept": concept,
"alpha": round(alpha, 4),
"mask_norm": round(float(mask.norm().item()), 6),
"nnz": int(torch.count_nonzero(mask).item()),
})
with torch.inference_mode():
_, sparse, _, _ = self.model.inference(z)
edited_sparse = sparse.clone()
for mask, alpha in slider_masks:
mask = mask.to(device=self.device, dtype=edited_sparse.dtype)
edited_sparse = edited_sparse + alpha * mask.reshape(1, -1)
edited_sparse = edited_sparse.clamp_min(0.0)
edited = decode_normalized(self.model, edited_sparse)
idx, scores = self._retrieve_dense(edited, k, exclude_idx=seed_idx)
sparse_delta = float((edited_sparse - sparse).norm().item())
dense_delta = float((edited.reshape(-1) - F.normalize(z.reshape(-1), dim=0)).norm().item())
dense_cos = float(torch.dot(edited.reshape(-1), F.normalize(z.reshape(-1), dim=0)).item())
top = [
(self.track_ids[int(i)], round(float(s), 4))
for i, s in zip(idx.tolist()[:5], scores.tolist()[:5])
]
_log(
"{} seed={} sliders={} sparse_delta={:.6f} dense_delta={:.6f} dense_seed_cos={:.6f} top={}",
call,
seed_track_id,
mask_debug,
sparse_delta,
dense_delta,
dense_cos,
top,
)
return self._format_results(idx, scores)
def _mask_from_payload(self, payload) -> torch.Tensor:
"""Restore a sparse edit mask serialized through Gradio session state."""
if isinstance(payload, torch.Tensor):
return payload.detach().cpu().float().view(-1)
dict_size = int(self.model.sae_encoder.dict_size)
mask = torch.zeros(dict_size, dtype=torch.float32)
if payload is None:
return mask
if isinstance(payload, dict):
payload = payload.items()
payload = list(payload)
if not payload:
return mask
first = payload[0]
if isinstance(first, (list, tuple)) and len(first) == 2:
for idx, val in payload:
idx = int(idx)
if 0 <= idx < dict_size:
mask[idx] = float(val)
return mask
if len(payload) == dict_size:
return torch.tensor(payload, dtype=torch.float32)
raise ValueError(f"Unsupported slider mask payload with length {len(payload)}")
def multi_mask_steer_and_retrieve(self, seed_track_id: str, masks: list[tuple[str, object, float]], k: int = 8) -> list[dict]:
"""Apply serialized slider masks from Gradio state, then retrieve once."""
call = _stamp("mask")
z = self.seed_embedding(seed_track_id).reshape(1, -1)
seed_idx = self.id_to_idx[seed_track_id]
slider_masks = []
mask_debug = []
for concept, payload, alpha in masks:
concept = str(concept).strip()
alpha = float(alpha)
if not concept or abs(alpha) < 1e-6:
continue
mask = self._mask_from_payload(payload)
slider_masks.append((mask, alpha))
mask_debug.append({
"concept": concept,
"alpha": round(alpha, 4),
"mask_norm": round(float(mask.norm().item()), 6),
"nnz": int(torch.count_nonzero(mask).item()),
})
with torch.inference_mode():
_, sparse, _, _ = self.model.inference(z)
edited_sparse = sparse.clone()
for mask, alpha in slider_masks:
mask = mask.to(device=self.device, dtype=edited_sparse.dtype)
edited_sparse = edited_sparse + alpha * mask.reshape(1, -1)
edited_sparse = edited_sparse.clamp_min(0.0)
edited = decode_normalized(self.model, edited_sparse)
idx, scores = self._retrieve_dense(edited, k, exclude_idx=seed_idx)
sparse_delta = float((edited_sparse - sparse).norm().item())
dense_delta = float((edited.reshape(-1) - F.normalize(z.reshape(-1), dim=0)).norm().item())
dense_cos = float(torch.dot(edited.reshape(-1), F.normalize(z.reshape(-1), dim=0)).item())
top = [
(self.track_ids[int(i)], round(float(s), 4))
for i, s in zip(idx.tolist()[:5], scores.tolist()[:5])
]
_log(
"{} seed={} masks={} sparse_delta={:.6f} dense_delta={:.6f} dense_seed_cos={:.6f} top={}",
call,
seed_track_id,
mask_debug,
sparse_delta,
dense_delta,
dense_cos,
top,
)
return self._format_results(idx, scores)
def _format_results(self, idx, scores) -> list[dict]:
out = []
for i, s in zip(idx.tolist(), scores.tolist()):
m = self.track_meta(self.track_ids[i])
m["affinity"] = float(s)
out.append(m)
return out
_ENGINE: DemoEngine | None = None
def get_engine(corpus=None) -> DemoEngine:
global _ENGINE
if _ENGINE is None:
_ENGINE = DemoEngine(corpus=corpus)
return _ENGINE