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