"""Sparse Steerable Retrieval — live demo (Hugging Face ZeroGPU Space).
Search for a seed track or draw one at random, make one or more concept sliders,
and retrieve live as the sliders move. No audio is hosted — playback is via
Spotify embeds.
ZeroGPU: the corpus + MuQ weights are fetched at startup (CPU); the text tower is
used on GPU only while a new slider is made. Live slider edits reuse cached CPU
masks.
"""
import html
import os
import random
import re
import time
import gradio as gr
import numpy as np
# `spaces` exists only on HF ZeroGPU; shim to a no-op decorator elsewhere.
try:
import spaces
except Exception: # pragma: no cover
class _Spaces:
def GPU(self, *a, **k):
def deco(fn):
return fn
return deco if not (a and callable(a[0])) else a[0]
spaces = _Spaces()
from core import DemoEngine, load_corpus
K = 8
MAX_SLIDERS = 5
DEBUG = os.environ.get("SSR_DEBUG", "1").lower() not in {"0", "false", "no", "off"}
EXAMPLE_CONCEPTS = [
"piano",
"electric guitar",
"female vocals",
"synthwave",
"acoustic",
"dreamy and ethereal",
"warm and intimate",
]
# --- startup (CPU): pre-cache weights so the first GPU call stays within budget -------- #
for _repo in ("OpenMuQ/MuQ-MuLan-large", "OpenMuQ/MuQ-large-msd-iter"):
try:
from huggingface_hub import snapshot_download
snapshot_download(_repo)
except Exception:
pass
CORPUS = load_corpus()
EMB_NORM = CORPUS[0] / (np.linalg.norm(CORPUS[0], axis=1, keepdims=True) + 1e-8)
print(f"corpus loaded: {len(CORPUS[1])} tracks")
_ENGINE = None
_CALL_SEQ = 0
def _log(msg, *args):
if DEBUG:
print("[ssr-app] " + msg.format(*args), flush=True)
def _next_call(label):
global _CALL_SEQ
_CALL_SEQ += 1
return f"{label}-{_CALL_SEQ}-{int(time.time() * 1000) % 100000}"
def ensure_engine():
global _ENGINE
if _ENGINE is None:
_ENGINE = DemoEngine(corpus=CORPUS, dev="cuda")
return _ENGINE
# --- data helpers --------------------------------------------------------------------- #
def _safe(x):
return html.escape(str(x or ""))
def _spotify(sid, height=80):
if not sid:
return '
Spotify preview unavailable
'
sid = _safe(sid)
return (
f''
)
def _meta(track_id):
m = CORPUS[2].get(track_id, {})
return m.get("title", ""), m.get("artist", ""), m.get("genre", ""), m.get("spotify", "")
def _label_for(track_id):
title, artist, genre, _ = _meta(track_id)
bits = [b for b in (title, artist, genre) if b]
return f"{' — '.join(bits[:2])}{f' · {genre}' if genre and len(bits) >= 2 else ''} [{track_id}]"
def _track_from_label(label):
if not label:
return None
match = re.search(r"\[([^\]]+)\]\s*$", str(label))
return match.group(1) if match else str(label)
def _search_text(track_id):
title, artist, genre, _ = _meta(track_id)
return f"{title} {artist} {genre}".lower()
SEARCH_INDEX = [(tid, _search_text(tid)) for tid in CORPUS[1]]
def search_tracks(query):
query = (query or "").strip().lower()
if not query:
choices = [_label_for(tid) for tid in random.sample(CORPUS[1], min(12, len(CORPUS[1])))]
else:
tokens = query.split()
ranked = []
for tid, text in SEARCH_INDEX:
if all(tok in text for tok in tokens):
title, artist, genre, _ = _meta(tid)
starts = int(title.lower().startswith(query)) + int(artist.lower().startswith(query))
ranked.append((starts, tid))
ranked.sort(reverse=True)
choices = [_label_for(tid) for _, tid in ranked[:12]]
if not choices:
return gr.update(choices=[], value=None), 'No matching tracks. Try artist, title, or genre.'
return gr.update(choices=choices, value=choices[0]), 'Select a seed from the matches.'
def seed_card(track_id):
title, artist, genre, sid = _meta(track_id)
return (
''
'
Seed track
'
f'
{_safe(title)}
'
f'
{_safe(artist)} · {_safe(genre)}
'
f'{_spotify(sid, 152)}'
'
'
)
def _baseline_results(seed_track_id):
seed_idx = CORPUS[1].index(seed_track_id)
sims = EMB_NORM @ EMB_NORM[seed_idx]
sims[seed_idx] = -1e9
idx = np.argsort(-sims)[:K]
out = []
for i in idx:
title, artist, genre, sid = _meta(CORPUS[1][int(i)])
out.append({
"title": title,
"artist": artist,
"genre": genre,
"spotify": sid,
"affinity": float(sims[int(i)]),
})
return out
def results_html(results):
cards = []
for i, r in enumerate(results, 1):
cards.append(
''
'
'
f'{i}'
''
f'{_safe(r.get("title"))}'
f'{_safe(r.get("artist"))} · {_safe(r.get("genre"))}'
''
f'{float(r.get("affinity", 0.0)):.3f}'
'
'
f'{_spotify(r.get("spotify"))}'
'
'
)
return '' + "".join(cards) + "
"
def _result_sig(results, n=3):
return [
(
str(r.get("track_id", "")),
str(r.get("title", ""))[:32],
round(float(r.get("affinity", 0.0)), 4),
)
for r in list(results or [])[:n]
]
def _active_pairs(concepts, alphas):
concepts = concepts or []
return [
(str(concept), float(alpha or 0.0))
for concept, alpha in zip(concepts, alphas)
if str(concept).strip()
]
def _alpha_values(values=None):
out = [float(v or 0.0) for v in list(values or [])[:MAX_SLIDERS]]
while len(out) < MAX_SLIDERS:
out.append(0.0)
return out
def _mask_values(values=None):
out = list(values or [])[:MAX_SLIDERS]
while len(out) < MAX_SLIDERS:
out.append(None)
return out
def _mask_payload(slider):
mask = slider.mask.detach().cpu().float().view(-1)
idx = (mask.abs() > 0).nonzero(as_tuple=False).flatten().tolist()
return [[int(i), float(mask[int(i)].item())] for i in idx]
def _status_html(concepts, alphas, *, prefix=None):
concepts = concepts or []
chips = []
for concept, alpha in _active_pairs(concepts, alphas):
tone = "pos" if alpha > 0 else "neg" if alpha < 0 else "zero"
chips.append(f'{_safe(concept)} {alpha:+.1f}')
head = f'{_safe(prefix)}' if prefix else 'Live retrieval query'
return '' + head + '
' + "".join(chips) + '
'
def _slider_components(concepts, values=None):
values = list(values or [])
updates = []
for i in range(MAX_SLIDERS):
if i < len(concepts):
concept = concepts[i]
value = float(values[i]) if i < len(values) else 0.0
label = (
''
f'{_safe(concept)}'
f'α {value:+.1f}'
'
'
)
updates.extend([
gr.update(value=label, visible=True),
gr.update(value=value, visible=True, label=concept),
])
else:
updates.extend([
gr.update(value="", visible=False),
gr.update(value=0.0, visible=False, label=f"Concept {i + 1}"),
])
return updates
def _slider_label_components(concepts, values=None):
values = list(values or [])
updates = []
for i in range(MAX_SLIDERS):
if i < len(concepts):
concept = concepts[i]
value = float(values[i]) if i < len(values) else 0.0
label = (
''
f'{_safe(concept)}'
f'α {value:+.1f}'
'
'
)
updates.append(gr.update(value=label, visible=True))
else:
updates.append(gr.update(value="", visible=False))
return updates
def random_seed():
tid = random.choice(CORPUS[1])
return tid, seed_card(tid), _status_html([], [], prefix="Showing nearest neighbours for the seed."), results_html(_baseline_results(tid))
def select_seed(choice):
tid = _track_from_label(choice)
if not tid:
return None, "", 'Search for a track or choose random.', ""
return tid, seed_card(tid), _status_html([], [], prefix="Showing nearest neighbours for the seed."), results_html(_baseline_results(tid))
@spaces.GPU(duration=120)
def make_slider(seed_track_id, concept, concepts, alpha_values, mask_values):
call = _next_call("make")
values = _alpha_values(alpha_values)
masks = _mask_values(mask_values)
_log(
"{} start seed={} concept={!r} concepts={} alpha_state={} mask_nnz={}",
call,
seed_track_id,
concept,
concepts,
values,
[len(m or []) for m in masks],
)
if not seed_track_id:
_log("{} no-seed", call)
return (
concepts or [],
values,
masks,
'Pick a seed track first.',
"",
*_slider_components(concepts or [], values),
)
concept = (concept or "").strip()
concepts = list(concepts or [])
if not concept:
_log("{} empty-concept", call)
return concepts, values, masks, 'Type a concept, then make a slider.', results_html(_baseline_results(seed_track_id)), *_slider_components(concepts, values)
if concept.lower() not in {c.lower() for c in concepts}:
if len(concepts) >= MAX_SLIDERS:
msg = f"Maximum of {MAX_SLIDERS} sliders reached."
_log("{} max-sliders concepts={}", call, concepts)
return concepts, values, masks, f'{msg}', results_html(_baseline_results(seed_track_id)), *_slider_components(concepts, values)
concepts.append(concept)
values[len(concepts) - 1] = 0.0
eng = ensure_engine()
slider = eng._slider(concept)
support = len(slider)
slot = next((i for i, c in enumerate(concepts) if c.lower() == concept.lower()), len(concepts) - 1)
masks[slot] = _mask_payload(slider)
mask_inputs = [(c, masks[i], values[i]) for i, c in enumerate(concepts)]
results = eng.multi_mask_steer_and_retrieve(seed_track_id, mask_inputs, k=K)
_log(
"{} ready concepts={} alpha_state={} mask_nnz={} support={} top={}",
call,
concepts,
values,
[len(m or []) for m in masks],
support,
_result_sig(results),
)
note = _status_html(concepts, values, prefix=f"Slider ready: {concept} uses {support} sparse features.")
return concepts, values, masks, note, results_html(results), *_slider_components(concepts, values)
def live_retrieve(seed_track_id, concepts, alpha_values, mask_values):
call = _next_call("live")
concepts = list(concepts or [])
values = _alpha_values(alpha_values)
masks = _mask_values(mask_values)
_log(
"{} start seed={} concepts={} alpha_state={} mask_nnz={}",
call,
seed_track_id,
concepts,
values,
[len(m or []) for m in masks],
)
if not seed_track_id:
_log("{} no-seed", call)
return values, masks, 'Pick a seed track first.', "", *_slider_label_components(concepts, values)
if not concepts:
baseline = _baseline_results(seed_track_id)
_log("{} no-concepts baseline_top={}", call, _result_sig(baseline))
return (
values,
masks,
_status_html([], [], prefix="Showing nearest neighbours for the seed."),
results_html(baseline),
*_slider_label_components(concepts, values),
)
pairs = _active_pairs(concepts, values)
active_pairs = [(c, a) for c, a in pairs if abs(float(a)) >= 1e-6]
_log("{} pairs={} active={}", call, pairs, active_pairs)
if not active_pairs:
baseline = _baseline_results(seed_track_id)
_log("{} zero-active baseline_top={}", call, _result_sig(baseline))
return (
values,
masks,
_status_html(concepts, values, prefix="Showing nearest neighbours for the seed."),
results_html(baseline),
*_slider_label_components(concepts, values),
)
eng = ensure_engine()
missing = [concept for concept, _ in active_pairs if masks[concepts.index(concept)] is None]
if missing:
name = missing[0]
baseline = _baseline_results(seed_track_id)
_log("{} missing={} baseline_top={}", call, missing, _result_sig(baseline))
note = _status_html(concepts, values, prefix=f"Preparing slider for {name}; click Make slider if it is not ready.")
return values, masks, note, results_html(baseline), *_slider_label_components(concepts, values)
mask_inputs = [(concept, masks[i], values[i]) for i, concept in enumerate(concepts)]
results = eng.multi_mask_steer_and_retrieve(seed_track_id, mask_inputs, k=K)
_log("{} edited_top={}", call, _result_sig(results))
return values, masks, _status_html(concepts, values), results_html(results), *_slider_label_components(concepts, values)
def _released_slider(slot):
def fn(released_alpha, seed_track_id, concepts, alpha_values, mask_values):
call = _next_call("release")
values = _alpha_values(alpha_values)
masks = _mask_values(mask_values)
old_values = list(values)
if 0 <= slot < len(values):
values[slot] = float(released_alpha or 0.0)
_log(
"{} slot={} released_alpha={} old_state={} new_state={} seed={} concepts={}",
call,
slot,
released_alpha,
old_values,
values,
seed_track_id,
concepts,
)
return live_retrieve(seed_track_id, concepts, values, masks)
return fn
CSS = """
:root {
--ssr-purple: #7b3ff2;
--ssr-green: #27995f;
--ssr-orange: #e88a2a;
--ssr-pink: #df6b96;
--ssr-ink: #171717;
--ssr-muted: #737373;
--ssr-line: rgba(23, 23, 23, 0.12);
--ssr-glass: rgba(255, 255, 255, 0.72);
}
body, .gradio-container {
background: radial-gradient(70% 45% at 50% 0%, rgba(123, 63, 242, 0.11), transparent 72%), #ffffff !important;
color: var(--ssr-ink) !important;
font-family: Inter, ui-sans-serif, system-ui, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif !important;
}
.gradio-container {
max-width: none !important;
width: 100% !important;
padding-left: clamp(16px, 3vw, 44px) !important;
padding-right: clamp(16px, 3vw, 44px) !important;
}
footer, .api-docs { display: none !important; }
.ssr-hero {
padding: 20px 4px 16px;
}
.ssr-hero h1 {
margin: 0;
font-size: clamp(34px, 5vw, 62px);
letter-spacing: 0;
line-height: 1.02;
font-weight: 760;
}
.ssr-hero p {
max-width: 760px;
margin: 16px 0 0;
color: #555;
font-size: 17px;
line-height: 1.62;
}
.ssr-kicker {
color: var(--ssr-purple);
font-size: 12px;
font-weight: 760;
letter-spacing: .14em;
text-transform: uppercase;
}
.ssr-panel, .ssr-seed, .ssr-card {
border: 1px solid var(--ssr-line);
background: var(--ssr-glass);
box-shadow: 0 18px 55px rgba(23, 23, 23, 0.07);
backdrop-filter: blur(18px);
}
.ssr-panel {
border-radius: 22px !important;
padding: 18px !important;
gap: 14px !important;
}
.ssr-panel,
.ssr-panel > *,
.ssr-panel .form,
.ssr-panel .block,
.ssr-panel .wrap,
.ssr-panel .gradio-row,
.ssr-panel .gradio-column {
background-color: transparent !important;
}
.ssr-panel .form,
.ssr-panel .block {
border: 0 !important;
box-shadow: none !important;
}
.ssr-panel h3 {
margin: 0 0 2px !important;
font-size: 15px !important;
font-weight: 760 !important;
letter-spacing: 0 !important;
}
.ssr-panel label span {
color: var(--ssr-purple) !important;
font-size: 12px !important;
font-weight: 680 !important;
}
.ssr-seed {
padding: 16px;
border-radius: 20px;
}
.ssr-title {
display: block;
color: var(--ssr-ink);
font-weight: 720;
line-height: 1.15;
}
.ssr-seed .ssr-title { margin-top: 8px; font-size: 22px; }
.ssr-sub {
display: block;
margin-top: 3px;
color: var(--ssr-muted);
font-size: 13px;
}
.ssr-embed {
margin-top: 12px;
border: 0;
border-radius: 14px;
background: #f5f5f5;
}
.ssr-empty {
min-height: 76px;
display: grid;
place-items: center;
color: #9b9b9b;
font-size: 12px;
}
.ssr-results {
display: flex;
flex-direction: column;
gap: 12px;
}
.ssr-card {
padding: 12px;
border-radius: 18px;
}
.ssr-row {
display: flex;
align-items: center;
gap: 10px;
margin-bottom: 6px;
}
.ssr-track { min-width: 0; flex: 1; }
.ssr-card .ssr-title {
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
font-size: 14px;
}
.ssr-rank {
display: inline-flex;
align-items: center;
justify-content: center;
width: 22px;
height: 22px;
flex: none;
border-radius: 999px;
background: var(--ssr-purple);
color: #fff;
font-size: 11px;
font-weight: 760;
}
.ssr-aff {
color: #a3a3a3;
font-size: 12px;
font-variant-numeric: tabular-nums;
}
.ssr-note {
display: flex;
flex-wrap: wrap;
align-items: center;
gap: 10px;
color: #525252;
font-size: 13px;
}
.ssr-chiprow {
display: flex;
flex-wrap: wrap;
gap: 6px;
}
.ssr-chip {
display: inline-flex;
align-items: center;
gap: 6px;
padding: 5px 9px;
border-radius: 999px;
border: 1px solid rgba(123, 63, 242, 0.18);
background: rgba(123, 63, 242, 0.07);
color: var(--ssr-purple);
}
.ssr-chip.pos { color: var(--ssr-green); border-color: rgba(39, 153, 95, .22); background: rgba(39, 153, 95, .08); }
.ssr-chip.neg { color: var(--ssr-pink); border-color: rgba(223, 107, 150, .22); background: rgba(223, 107, 150, .08); }
.ssr-chip.zero { color: var(--ssr-purple); }
.ssr-muted { color: var(--ssr-muted); font-size: 13px; }
.ssr-slider-label {
display: flex;
justify-content: space-between;
align-items: baseline;
margin: 3px 0 -4px;
color: var(--ssr-purple);
font-weight: 720;
}
.ssr-alpha {
color: #a78bfa;
font-size: 12px;
font-variant-numeric: tabular-nums;
}
.gradio-container .form,
.gradio-container .block {
border-color: rgba(23, 23, 23, 0.08) !important;
border-radius: 16px !important;
}
.gradio-container input,
.gradio-container textarea,
.gradio-container select {
border-radius: 12px !important;
border-color: rgba(23, 23, 23, 0.10) !important;
box-shadow: none !important;
}
.gradio-container button {
border-radius: 999px !important;
font-weight: 650 !important;
min-height: 42px !important;
height: 42px !important;
box-shadow: 0 10px 24px rgba(23, 23, 23, 0.08) !important;
}
.ssr-panel button {
align-self: end !important;
padding-left: 18px !important;
padding-right: 18px !important;
}
.gradio-container button.primary {
background: var(--ssr-ink) !important;
border-color: var(--ssr-ink) !important;
}
input[type='range'] {
accent-color: var(--ssr-purple);
}
"""
THEME = gr.themes.Soft(
primary_hue="purple",
secondary_hue="neutral",
neutral_hue="neutral",
).set(
body_background_fill="#ffffff",
button_primary_background_fill="#171717",
button_primary_background_fill_hover="#000000",
button_primary_text_color="#ffffff",
block_radius="18px",
input_radius="12px",
)
with gr.Blocks(title="Sparse Steerable Retrieval", css=CSS, theme=THEME) as demo:
gr.HTML(
''
'
Sparse Steerable Retrieval
'
'
Steer music retrieval with concept sliders.
'
'
Search for a seed track or draw one at random. Make sliders from free-form concepts, '
'then move several at once: retrieval updates live from the combined sparse edit.
'
'
'
)
seed_state = gr.State()
slider_state = gr.State([])
alpha_state = gr.State([0.0] * MAX_SLIDERS)
mask_state = gr.State([None] * MAX_SLIDERS)
with gr.Row(equal_height=False):
with gr.Column(scale=5):
with gr.Column(elem_classes=["ssr-panel"]):
gr.Markdown("### 1. Choose a seed")
with gr.Row():
search = gr.Textbox(
label="Search music4all",
placeholder="artist, title, or genre",
scale=4,
)
search_btn = gr.Button("Search", variant="secondary", scale=1)
random_btn = gr.Button("Random seed", variant="secondary", scale=1)
matches = gr.Dropdown(label="Search results", choices=[], interactive=True)
pick_btn = gr.Button("Use selected seed", variant="primary")
search_note = gr.HTML('Search for a track, or start from a random seed.')
seed_html = gr.HTML()
with gr.Column(elem_classes=["ssr-panel"]):
gr.Markdown("### 2. Make concept sliders")
with gr.Row():
concept = gr.Textbox(
label="Concept",
placeholder="piano, electric guitar, female vocals, acoustic",
scale=4,
)
make_btn = gr.Button("Make slider", variant="primary", scale=1)
gr.Examples(EXAMPLE_CONCEPTS, inputs=concept, label="Try concepts")
slider_labels = []
slider_controls = []
for i in range(MAX_SLIDERS):
label = gr.HTML(visible=False)
control = gr.Slider(
minimum=-3.0,
maximum=3.0,
value=0.0,
step=0.1,
label=f"Concept {i + 1}",
visible=False,
interactive=True,
)
slider_labels.append(label)
slider_controls.append(control)
note = gr.HTML()
with gr.Column(scale=5):
with gr.Column(elem_classes=["ssr-panel"]):
gr.Markdown("### 3. Live retrieval")
results_out = gr.HTML()
search_btn.click(search_tracks, inputs=[search], outputs=[matches, search_note])
search.submit(search_tracks, inputs=[search], outputs=[matches, search_note])
random_btn.click(random_seed, outputs=[seed_state, seed_html, note, results_out])
pick_btn.click(select_seed, inputs=[matches], outputs=[seed_state, seed_html, note, results_out])
matches.change(select_seed, inputs=[matches], outputs=[seed_state, seed_html, note, results_out])
slider_outputs = [slider_state, alpha_state, mask_state, note, results_out]
for label, control in zip(slider_labels, slider_controls):
slider_outputs.extend([label, control])
make_btn.click(
make_slider,
inputs=[seed_state, concept, slider_state, alpha_state, mask_state],
outputs=slider_outputs,
)
live_outputs = [alpha_state, mask_state, note, results_out]
live_outputs.extend(slider_labels)
for i, control in enumerate(slider_controls):
control.release(
_released_slider(i),
inputs=[control, seed_state, slider_state, alpha_state, mask_state],
outputs=live_outputs,
trigger_mode="always_last",
concurrency_limit=1,
concurrency_id="live-retrieval",
)
demo.load(random_seed, outputs=[seed_state, seed_html, note, results_out])
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=int(os.getenv("PORT", "7860")))