Spaces:
Sleeping
Sleeping
File size: 13,001 Bytes
bda104d cbfe92b bda104d 8852300 bda104d cbfe92b bda104d 6876324 bda104d cbfe92b 8852300 bda104d cbfe92b bda104d cbfe92b bda104d cbfe92b bda104d 3ecf48b cbfe92b 3ecf48b cbfe92b 8852300 cbfe92b 8852300 cbfe92b bda104d cbfe92b bda104d cbfe92b 6876324 8852300 6876324 cbfe92b 8852300 cbfe92b 8852300 cbfe92b 6876324 cbfe92b 6876324 cbfe92b 8852300 6876324 9d8127c 6876324 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 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 | """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
|