soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
45.1 kB
"""Gradio demo for the SAME-latent stem RESTORER (v4b) + GENERATOR pipeline.
Unified pipeline: a 3s clip's degraded full-mix -> htdemucs stems -> SAME latents ->
- present stems : v4b deterministic RESTORE
- missing stems : conditional-flow GENERATE from the (restored) context
-> SAME decode -> remixed output. Val samples use the cached latents (fast); uploads run
the full htdemucs+SAME path and are chunked into 3s (last chunk padded with silence then cropped).
Runs on CPU by default (set RESTOFLOW_DEVICE=cuda to use a GPU).
Launch: python -m restoflow.app
"""
from __future__ import annotations
import os, csv, glob, io, hashlib
# Subprocess ffmpeg fix (audio-separator + librosa/audioread shell out to `ffmpeg`):
# a stale /home/maximos/miniconda3/lib on LD_LIBRARY_PATH shadows system libstdc++ and
# breaks the system ffmpeg; drop it and put a known-good conda ffmpeg first on PATH.
os.environ["LD_LIBRARY_PATH"] = ":".join(
p for p in os.environ.get("LD_LIBRARY_PATH", "").split(":") if p and "maximos/miniconda3" not in p)
_FFMPEG_DIR = "/home/ksoil/.conda/envs/ksoil_torch/bin"
if os.path.isdir(_FFMPEG_DIR):
os.environ["PATH"] = _FFMPEG_DIR + ":" + os.environ.get("PATH", "")
from pathlib import Path
import numpy as np
import torch
import soundfile as sf
import librosa
import librosa.display
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import gradio as gr
STEM_EMOJI = {"other": "🎹", "vocals": "🎤", "drums": "🥁", "bass": "🎸"}
from .config import Cfg, STEMS, STEM_ID
from .model import DetRestorer, CondFlow, AttnRestorer, AttnCondFlow, MixAttnCondFlow
from . import eval as E
from . import gen as G
from . import flow as F
# Paths are env-overridable so the same file runs locally and on a packaged HF Space.
# Local default = the dev tree; on HF the deploy bundle sets these (or we fall back to a
# repo-relative layout: <repo>/restoflow_runs, <repo>/demo_samples).
_LOCAL_BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra")
_REPO_ROOT = Path(__file__).resolve().parent.parent # <repo> containing restoflow/
BASE = Path(os.environ.get("RESTOFLOW_BASE", _LOCAL_BASE if _LOCAL_BASE.exists() else _REPO_ROOT))
SRC = Path(os.environ.get(
"RESTOFLOW_SRC",
"/media/maindisk/melkor169/MSRKit/xlance-msr/stereo_a2sb/synthetic_out"))
if not SRC.exists(): # packaged Space: demo wavs live here
SRC = BASE / "demo_samples"
VAL_CACHE = Path(os.environ.get("RESTOFLOW_VAL_CACHE", str(BASE / "demucs_results_val_full")))
if not VAL_CACHE.exists():
VAL_CACHE = BASE / "demo_samples" / "val_cache"
DEVICE = os.environ.get("RESTOFLOW_DEVICE", "cpu")
T, D = 32, 256
SR = 44100
N_VAL = 80
PRESENT_RMS, PRESENT_PEAK = -40.0, -25.0 # presence gate (rms OR peak), matches training
# Generation policy: every song has a low end, so a MISSING bass is always invented.
# A missing drums/other is legitimate (song may have none) -> we DON'T fabricate it; restore-only.
GEN_TARGETS = ("bass",)
_M = {} # model cache
# ---------------- model loading ----------------
def _disable_flash_attn():
"""SAME's transformer uses flash_attn (CUDA-only). On CPU, force the SDPA fallback
by nulling the module globals apply_attn() checks at call time."""
import stable_audio_tools.models.transformer as _tf
for name in ("flash_attn_func", "flash_attn_kvpacked_func", "flash_attn_varlen_func", "index_first_axis"):
if hasattr(_tf, name):
setattr(_tf, name, None)
def _ckpt_path(run):
d = BASE / "restoflow_runs" / run
for f in ("ckpt_best.pt", "ckpt.pt"):
if (d / f).exists():
return d / f
return None
def _default_runs(rests, gens):
"""Preselected restorer/generator. Env override (set by the HF deploy) wins, else a
'prefer newest best' heuristic, else first available."""
dr = os.environ.get("RESTOFLOW_DEFAULT_REST")
dg = os.environ.get("RESTOFLOW_DEFAULT_GEN")
if dr not in rests:
# remix_glue_v1 is the published default: frozen advramp backbone + a residual-flow detail head,
# co-trained with the generator against a multi-scale latent discriminator for a coherent remix.
# restorer_attn_advramp (the deterministic backbone alone) is kept selectable as the BASELINE.
dr = next((r for r in ("remix_glue_v1", "restorer_attn_advramp", "restorer_attn_w003_gan",
"restorer_attn_v1", "v4b_balonly") if r in rests),
rests[0] if rests else None)
if dg not in gens:
# remix_glue_v1_gen: the generator co-trained in the glue stage with remix_glue_v1. gen_advramp_v1
# (trained on advramp's conditioning) is the BASELINE generator. Both conv 10M (fast on CPU).
dg = next((g for g in ("remix_glue_v1_gen", "gen_advramp_v1", "gen_distvar_baseline",
"gen_v2", "gen_v3") if g in gens),
gens[0] if gens else "(none)")
return dr, dg
# Validated runs promoted past the experiment filter below. The glue model (remix_glue_v1: advramp
# backbone + residual-flow detail head, co-trained with remix_glue_v1_gen) is the published default;
# advramp + gen_advramp_v1 stay selectable as the baseline.
PROMOTED_REST = {"remix_glue_v1"}
PROMOTED_GEN = {"remix_glue_v1_gen"}
def list_runs(kind):
"""Available run names (excluding obsolete) for kind in {'restorer','generator'}.
Convention: gen_* are generators; everything else is a restorer (plus explicit PROMOTED_GEN)."""
out = []
for d in sorted(glob.glob(str(BASE / "restoflow_runs" / "*"))):
name = Path(d).name
if name == "obsolete" or not Path(d).is_dir() or _ckpt_path(name) is None:
continue
promoted = name in PROMOTED_REST or name in PROMOTED_GEN
# skip scratch/experiment dirs (scout_*, _smoke_*, remix_* GAN-refine runs) unless explicitly promoted
if not promoted and name.startswith(("scout", "_", "remix", "exp")):
continue
is_gen = name.startswith("gen") or name in PROMOTED_GEN
if (kind == "generator") == is_gen:
out.append(name)
return out
def load_restorer(run):
ck = torch.load(_ckpt_path(run), map_location="cpu"); rc = ck["cfg"]
kind = rc.get("model_kind", "det")
if kind == "flowattn": # GENERATIVE restorer (deg-anchored flow + mix)
m = MixAttnCondFlow(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
stem_emb=rc["stem_emb_dim"], heads=rc.get("heads", 8), use_mix=rc["use_mix"])
elif kind == "attn": # Conformer-lite deterministic restorer
m = AttnRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"], heads=rc.get("heads", 4))
else:
m = DetRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"])
m.load_state_dict(ck["model"]); m.eval().to(DEVICE)
stats = torch.load(BASE / "restoflow_runs" / run / "norm_stats.pt", map_location="cpu")
d = dict(rest=m, rstats=stats, rest_run=run, rest_ep=ck.get("epoch"), rest_kind=kind,
rest_steps=int(rc.get("sample_steps", 40)), rest_sigma=float(rc.get("sigma", 0.0)),
rest_use_mix=bool(rc.get("use_mix", True)),
rest_params=sum(p.numel() for p in m.parameters()) / 1e6)
if "rflow" in ck: # glue-stage residual-flow detail head (det backbone + sampled residual)
rfc = ck["rflow_cfg"]
rf = CondFlow(D, rfc["hidden"], rfc["depth"], n_stems=len(STEMS), stem_emb=rfc["stem_emb"])
rf.load_state_dict(ck["rflow"]); rf.eval().to(DEVICE)
d["rflow"] = rf; d["rflow_ids"] = set(rfc["restore_ids"]); d["rflow_steps"] = int(rfc.get("gen_steps", 2))
d["rest_params"] += sum(p.numel() for p in rf.parameters()) / 1e6
return d
def load_generator(run):
if run in (None, "(none)"):
return dict(gen=None, gstats=None, gargs=None, gen_run="(none)", gen_ep=None, gen_params=0.0)
ck = torch.load(_ckpt_path(run), map_location="cpu"); ga = ck["args"]
if ga.get("arch") == "attn": # Conformer-lite velocity net
m = AttnCondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS),
stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8))
else:
m = CondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim)
m.load_state_dict(ck["model"]); m.eval().to(DEVICE)
return dict(gen=m, gstats=ck["stats"], gargs=ga, gen_run=run, gen_ep=ck.get("epoch"),
gen_params=sum(p.numel() for p in m.parameters()) / 1e6)
def status_md():
if "rest_run" not in _M:
return "⏳ **Models load on first run** (SAME + restorer + generator) — first action takes ~15 s, then it's fast."
rb = f"**Restorer:** `{_M.get('rest_run','-')}` (ep {_M.get('rest_ep','?')}, {_M.get('rest_params',0):.1f}M)"
g = _M.get("gen_run", "(none)")
gt = (_M.get("gargs") or {}).get("target_stems", "—")
gb = f"**Generator:** `{g}`" + (f" (ep {_M.get('gen_ep','?')}, {_M.get('gen_params',0):.1f}M, targets={gt})" if _M.get("gen") is not None else " — *disabled*")
return f"🎛️ Active models · {rb} · {gb} · **Autoencoder:** `SAME-L` · device `{DEVICE}`"
def set_models(rest_run, gen_run):
_M.update(load_restorer(rest_run)); _M.update(load_generator(gen_run))
print(f"[app] active: restorer={rest_run} generator={gen_run}")
return status_md()
def models():
if _M:
return _M
cfg = Cfg(); cfg.device = DEVICE
if DEVICE == "cpu":
_disable_flash_attn()
sa, sr = E.load_same(cfg) # SAME-L autoencoder (frozen)
_M.update(sa=sa, sr=sr, cfg=cfg)
rests = list_runs("restorer"); gens = list_runs("generator")
default_rest, default_gen = _default_runs(rests, gens)
_M.update(load_restorer(default_rest)); _M.update(load_generator(default_gen))
print(f"[app] models loaded on {DEVICE}. restorer={default_rest} generator={default_gen}")
return _M
# ---------------- core latent ops ----------------
def _fitT(x): return G._fitT(x.float(), T)
def rms_peak_db(audio_mono):
rms = float(np.sqrt(np.mean(audio_mono ** 2)) + 1e-12)
peak = float(np.max(np.abs(audio_mono)) + 1e-12)
return 20 * np.log10(rms), 20 * np.log10(peak)
@torch.inference_mode()
def restore(latent, stem, mix):
m = models(); st = m["rstats"]; mu, sd = st[stem]["mu"], st[stem]["sd"]
sd_d = ((latent - mu[:, None]) / sd[:, None]).to(DEVICE)[None]
sd_m = None # no-mix restorers lack __mix__ stats
if m.get("rest_use_mix", True) and "__mix__" in st:
mm = st["__mix__"]
sd_m = ((mix - mm["mu"][:, None]) / mm["sd"][:, None]).to(DEVICE)[None]
sid = torch.tensor([STEM_ID[stem]], device=DEVICE)
if m.get("rest_kind") == "flowattn": # GENERATIVE restorer: integrate the deg-anchored flow
out = F.sample_mix(m["rest"], sd_d, sd_m, sid, m.get("rest_steps", 40), m.get("rest_sigma", 0.0))
else:
out = m["rest"](sd_d, sid, sd_m)
rflow = m.get("rflow") # glue stage: backbone + sampled residual on restore stems
if rflow is not None and STEM_ID[stem] in m.get("rflow_ids", set()):
out = out + G.sample_gen(rflow, out, sid, m.get("rflow_steps", 2), 1.0)
out = out[0].cpu()
return out * sd[:, None] + mu[:, None]
@torch.inference_mode()
def generate(context_latents, stem):
m = models()
if m["gen"] is None:
return None
gs = m["gstats"]; cm, cs = gs["ctx"][stem]; tm, ts = gs["tgt"][stem]
ctx_raw = sum(_fitT(c) for c in context_latents)
ctx = ((ctx_raw - cm[:, None]) / cs[:, None]).to(DEVICE)[None]
sid = torch.tensor([STEM_ID[stem]], device=DEVICE)
w = m["gargs"].get("cfg_w", 2.0); steps = m["gargs"].get("sample_steps", 40)
z = G.sample_gen(m["gen"], ctx, sid, steps, w)[0].cpu()
return z * ts[:, None] + tm[:, None]
# STOPGAP loudness control (proper fix = the planned generator overhaul). Generated bass decodes
# ~2x too loud (off-manifold flow endpoints); neither more data nor attention fixed it. We ground
# its level in the context: target rms = ratio x restored-context rms, applied in AUDIO space
# (SAME decode is nonlinear). Ratio = geometric mean of clean_bass/ctx over GENERATION-case val
# pairs (0.24); the per-song spread is wide (p25 .12, p75 .81), so we deliberately err quiet —
# too-loud bass is far more jarring than too-quiet.
GEN_ENERGY_RATIO = {"bass": 0.24}
def energy_match(audio_by_stem: dict, dec: dict):
"""Scale each GENERATED stem to its fitted energy relative to the restored context."""
ctx = [audio_by_stem[s] for s, v in dec.items() if v == "restored" and s in audio_by_stem]
if not ctx:
return audio_by_stem
ctx_rms = float(np.sqrt(np.mean(np.square(np.sum(ctx, axis=0)))) + 1e-9)
for s, v in dec.items():
if v == "generated" and s in audio_by_stem and s in GEN_ENERGY_RATIO:
a = audio_by_stem[s]
r = float(np.sqrt(np.mean(np.square(a))) + 1e-9)
audio_by_stem[s] = a * (GEN_ENERGY_RATIO[s] * ctx_rms / r)
return audio_by_stem
@torch.inference_mode()
def decode(latent):
return E._decode(models()["sa"], _fitT(latent)) # [2,S] cpu
@torch.inference_mode()
def encode(audio_2S):
return encode_batch([audio_2S])[0]
@torch.inference_mode()
def encode_batch(auds):
"""Encode many [2,S] clips in ONE SAME forward pass (CPU: far less overhead than N calls)."""
m = models(); xs = []
for a in auds:
x = torch.as_tensor(a).float()
if x.shape[0] != 2: x = (x[:2] if x.shape[0] > 2 else x.repeat(2, 1))
xs.append(x)
L = max(x.shape[1] for x in xs)
X = torch.stack([torch.nn.functional.pad(x, (0, L - x.shape[1])) for x in xs]).to(DEVICE)
Z = m["sa"].encode_audio(X, chunked=False)
return [Z[i].float().cpu() for i in range(Z.shape[0])]
@torch.inference_mode()
def decode_batch(lats):
"""Decode many [256,T] latents in ONE SAME forward pass."""
if not lats:
return []
m = models()
Z = torch.stack([_fitT(l) for l in lats]).to(DEVICE).float()
A = m["sa"].decode_audio(Z)
return [A[i].clamp(-1, 1).cpu() for i in range(A.shape[0])]
def process_chunk(deg_latents: dict, mix_latent, present: dict):
"""deg_latents/present keyed by stem. Restore present, generate missing. Returns
(final_latents dict, decisions dict)."""
final, dec = {}, {}
for s in STEMS:
if present.get(s) and s in deg_latents:
final[s] = restore(deg_latents[s], s, mix_latent); dec[s] = "restored"
ctx = [final[o] for o in final] # restored context
for s in STEMS:
if s in final:
continue
if s in GEN_TARGETS: # always invent missing bass
g = generate(ctx, s) if ctx else None
if g is not None: final[s] = g; dec[s] = "generated"
else: dec[s] = "absent (no generator/context)"
else: # drums/other: don't fabricate
dec[s] = "absent (kept out)"
return final, dec
# ---------------- viz / eval ----------------
def spec_fig(audio_2S, title):
y = np.asarray(audio_2S).mean(0) if np.asarray(audio_2S).ndim == 2 else np.asarray(audio_2S)
S = librosa.amplitude_to_db(np.abs(librosa.stft(y, n_fft=1024, hop_length=256)) + 1e-6, ref=np.max)
fig, ax = plt.subplots(figsize=(4.2, 2.4))
librosa.display.specshow(S, sr=SR, hop_length=256, x_axis="time", y_axis="log", ax=ax, cmap="magma")
ax.set_title(title, fontsize=9); ax.tick_params(labelsize=7)
fig.tight_layout()
return fig
def multi_stft_np(a, b):
return E.multi_stft(torch.as_tensor(a).float(), torch.as_tensor(b).float())
def load_audio_any(path):
"""Read wav/flac via soundfile; fall back to librosa+ffmpeg for mp3/m4a/ogg/etc.
Returns ([C, S] float32, sr)."""
try:
a, sr = sf.read(path, dtype="float32")
return (a.T if a.ndim == 2 else a[None]), sr
except Exception: # mp3 & friends -> ffmpeg/audioread
y, sr = librosa.load(path, sr=None, mono=False)
y = np.asarray(y, dtype="float32")
return (y if y.ndim == 2 else y[None]), sr
# ---------------- val index ----------------
def build_val_index():
rows = {}
meta = VAL_CACHE / "metadata.csv"
for r in csv.DictReader(open(meta)):
rows.setdefault(r["pair_id"], {})[r["stem"]] = r
out = []
for pid, st in rows.items():
if not (SRC / "val/degraded" / f"{pid}.wav").exists():
continue
present = [s for s in STEMS if st.get(s, {}).get("deg_present") == "1"]
# generation showcase = a MISSING bass that exists in clean (policy: always invent bass)
missing = [s for s in GEN_TARGETS if s not in present and st.get(s, {}).get("clean_present") == "1"]
if present and missing: # restore present + generate bass
out.append((pid, present, missing))
out.sort(key=lambda x: (-len(x[2]), x[0]))
return out[:N_VAL]
_VAL = None
def val_list():
global _VAL
if _VAL is None: _VAL = build_val_index()
return _VAL
def val_label(item):
pid, pres, miss = item
return f"{pid} | restore: {','.join(pres) or '-'} | generate: {','.join(miss) or '-'}"
# ---------------- pipelines ----------------
def _empty_val(msg):
main = [None, None, None, None, None, None, msg]
stems = []
for s in STEMS:
stems += [f"**{STEM_EMOJI[s]} {s}**", None, None]
return main + stems
def run_val(label):
if not label:
return _empty_val("Pick a sample and press Run.")
pid = label.split()[0]
deg_a, _ = sf.read(SRC / "val/degraded" / f"{pid}.wav", dtype="float32")
deg_a = deg_a.T if deg_a.ndim == 2 else deg_a[None]
clean_a, _ = sf.read(SRC / "val/clean" / f"{pid}.wav", dtype="float32")
clean_a = clean_a.T if clean_a.ndim == 2 else clean_a[None]
latdir = VAL_CACHE / pid / "latents"
deg_latents, present = {}, {}
for s in STEMS:
p = latdir / "degraded" / f"{s}.pt"
if p.exists():
deg_latents[s] = torch.load(p, map_location="cpu").float(); present[s] = True
mix = torch.load(latdir / "degraded_mix.pt", map_location="cpu").float()
final, dec = process_chunk(deg_latents, mix, present)
fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys]) # one batched decode
after = {fkeys[i]: fdec[i].numpy() for i in range(len(fkeys))}
energy_match(after, dec) # level generated bass
bkeys = list(deg_latents); bdec = decode_batch([deg_latents[s] for s in bkeys])
before = {bkeys[i]: bdec[i].numpy() for i in range(len(bkeys))}
out = sum(after.values()) # sum after leveling
n = min(out.shape[1], clean_a.shape[1], deg_a.shape[1])
e_in = multi_stft_np(deg_a[:, :n].mean(0), clean_a[:, :n].mean(0))
e_out = multi_stft_np(out[:, :n].mean(0), clean_a[:, :n].mean(0))
n_rest = sum(v == "restored" for v in dec.values()); n_gen = sum(v == "generated" for v in dec.values())
report = (f"**`{pid}`** · 🔧 restored {n_rest} · ✨ generated {n_gen} | "
f"full-mix vs clean (lower better): degraded **{e_in:.3f}** → output **{e_out:.3f}** "
f"{'✅' if e_out < e_in else '⚠️'}")
main = [(SR, deg_a.T), (SR, out.T), (SR, clean_a.T),
spec_fig(deg_a, "degraded input"), spec_fig(out, "output"), spec_fig(clean_a, "clean reference"),
report]
stems = []
for s in STEMS:
tag = {"restored": "🔧 restored", "generated": "✨ generated"}.get(dec.get(s), "—")
bef = (SR, before[s].T) if s in before else None
aft = (SR, after[s].T) if s in after else None
stems += [f"**{STEM_EMOJI[s]} {s}** — {tag}", bef, aft]
return main + stems
def _empty_upload(msg):
main = [None, None, None, None, msg]
stems = []
for s in STEMS:
stems += [f"**{STEM_EMOJI[s]} {s}**", None]
return main + stems
def _rms(x):
return float(np.sqrt(np.mean(x.astype(np.float64) ** 2)) + 1e-12)
def run_upload(filepath, start_s=0.0, end_s=0.0):
"""Wrapper: surface the REAL error in the UI (and Space logs) instead of a bare Gradio 'Error'."""
try:
return _run_upload(filepath, start_s, end_s)
except Exception as e:
import traceback; traceback.print_exc() # -> Space Logs tab
return _empty_upload(f"⚠️ upload failed — {type(e).__name__}: {str(e)[:400]}")
def _run_upload(filepath, start_s=0.0, end_s=0.0):
if not filepath:
return _empty_upload("Upload a wav/mp3 and press Run.")
a, sr = load_audio_any(filepath) # wav/flac/mp3/m4a...
if a.shape[0] == 1: a = np.repeat(a, 2, 0)
if sr != SR:
a = np.stack([librosa.resample(a[c], orig_sr=sr, target_sr=SR) for c in range(a.shape[0])])
# explicit trim (seconds): actually chop the file before processing
dur = a.shape[1] / SR
s = max(0.0, float(start_s or 0.0))
e = float(end_s) if (end_s and float(end_s) > 0) else dur
e = min(e, dur)
if e <= s:
s, e = 0.0, dur
a = a[:, int(s * SR):int(e * SR)]
trim_note = f"trim {s:.1f}–{e:.1f}s of {dur:.1f}s · "
chunk = SR * 3; fade = SR // 2; hop = chunk - fade; total = a.shape[1] # 3s windows, 0.5s crossfade
try:
sep = get_separator()
except RuntimeError as e: # separation unavailable on this deployment
return _empty_upload(f"⚠️ {e}")
# ---- separate the WHOLE clip ONCE (the CPU bottleneck) — demucs/audio-sep split internally — then
# process OVERLAPPING 3s windows by SLICING the pre-separated stems (no per-window re-separation) ----
full = separate_chunk(sep, a)
windows = []; summ = {}
for i in range(0, max(1, total), hop):
seg = a[:, i:i + chunk]; orig = seg.shape[1]
if orig <= 0:
break
if orig < chunk:
seg = np.pad(seg, ((0, 0), (0, chunk - orig))) # pad silence then crop output back
stems = {}
for s, fs in full.items():
w = fs[:, i:i + orig]
if w.shape[1] < chunk:
w = np.pad(w, ((0, 0), (0, chunk - w.shape[1])))
stems[s] = w
present_stems = [s for s, au in stems.items()
if (lambda rp: rp[0] > PRESENT_RMS or rp[1] > PRESENT_PEAK)(rms_peak_db(au.mean(0)))]
enc = encode_batch([stems[s] for s in present_stems] + [seg]) # one batched encode (stems + mix)
deg_latents = {present_stems[j]: enc[j] for j in range(len(present_stems))}
present = {s: True for s in present_stems}; mix_lat = enc[-1]
final, dec = process_chunk(deg_latents, mix_lat, present)
for s, v in dec.items(): summ[s] = v
fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys]) # one batched decode
chunk_audio = {s: fdec[j].numpy() for j, s in enumerate(fkeys)} # SAME decode = fixed length/window
energy_match(chunk_audio, dec) # intra-chunk: level generated bass
windows.append({"start": i, "orig": orig, "stems": chunk_audio})
if i + chunk >= total:
break
# ---- GLUE 1 (profile): gently match each window's per-stem level to the song-wide MEDIAN, so chunks
# share a consistent profile instead of drifting. alpha<1 preserves real dynamics; clip = no artifacts.
ALPHA, LO, HI = 0.5, 0.7, 1.4
for s in STEMS:
rmss = [_rms(w["stems"][s]) for w in windows if s in w["stems"] and _rms(w["stems"][s]) > 1e-5]
if len(rmss) >= 2:
tgt = float(np.median(rmss))
for w in windows:
if s in w["stems"] and _rms(w["stems"][s]) > 1e-5:
w["stems"][s] = w["stems"][s] * float(np.clip((tgt / _rms(w["stems"][s])) ** ALPHA, LO, HI))
# ---- GLUE 2 (seams): crossfaded overlap-add, done in OUTPUT-sample space (SAME decodes each 3s window
# to a fixed length Ld, not the input `orig`). out_hop/out_fade scale the input hop/fade by Ld/chunk.
# One global weight timeline (incl. windows where a stem is absent) so a stem fades across a boundary
# into a neighbour that lacks it. ----
n = len(windows)
Ld = next((au.shape[1] for w in windows for au in w["stems"].values()), 0)
if Ld == 0:
return _empty_upload(f"{trim_note}no restorable content found.")
out_fade = min(int(round(fade * Ld / chunk)), Ld // 2)
out_hop = max(1, Ld - out_fade)
def _Lout(w): # real output length (last/short window cropped)
return Ld if w["orig"] >= chunk else min(Ld, int(round(w["orig"] * Ld / chunk)))
starts = [k * out_hop for k in range(n)]
total_out = starts[-1] + _Lout(windows[-1])
Wg = np.zeros(total_out); envs = []
for k, w in enumerate(windows):
L = _Lout(w); env = np.ones(L)
if k > 0 and out_fade > 0: env[:min(out_fade, L)] = np.linspace(0, 1, out_fade, endpoint=False)[:min(out_fade, L)]
if k < n - 1 and L >= out_fade and out_fade > 0: env[-out_fade:] = np.linspace(1, 0, out_fade, endpoint=False)
envs.append(env); Wg[starts[k]:starts[k] + L] += env
Wg = np.maximum(Wg, 1e-6)
per_stem_audio = {}
for s in STEMS:
if not any(s in w["stems"] for w in windows):
continue
buf = np.zeros((2, total_out))
for k, w in enumerate(windows):
if s in w["stems"]:
L = _Lout(w); buf[:, starts[k]:starts[k] + L] += w["stems"][s][:, :L] * envs[k]
per_stem_audio[s] = (buf / Wg).astype(np.float32)
out = sum(per_stem_audio.values()) if per_stem_audio else np.zeros((2, total_out), np.float32)
rep = (f"{trim_note}**{n}** overlapping 3 s windows · hop {hop/SR:.1f}s, {fade/SR:.1f}s crossfade · "
"glued to song profile. " + " · ".join(f"{STEM_EMOJI[st]} {summ.get(st, '-')}" for st in STEMS))
main = [(SR, a.T), (SR, out.T), spec_fig(a, "input"), spec_fig(out, "output"), rep]
stems_out = []
for s in STEMS:
tag = {"restored": "🔧 restored", "generated": "✨ generated"}.get(summ.get(s), "—")
au = (SR, per_stem_audio[s].T) if s in per_stem_audio else None
stems_out += [f"**{STEM_EMOJI[s]} {s}** — {tag}", au]
return main + stems_out
# ---------------- separator (uploads only) ----------------
# EXACT-PARITY: the train/val cache was built with htdemucs v4 (audio-separator's htdemucs.yaml,
# shifts=1 overlap=0.25 — see demucs.py). Live upload MUST use those SAME v4 weights or the restorer
# sees a different stem distribution. Priority: (1) audio-separator (exactly how the cache was built —
# present locally); (2) the official `demucs` package (SAME htdemucs v4 weights, far lighter deps so it
# installs on the CPU Space without the numpy/pywt clash that disabled this before). We deliberately do
# NOT fall back to torchaudio's HDemucs — that's v3 (no transformer), i.e. NOT identical stems.
_SEP = {}
def get_separator():
if _SEP: return _SEP
try: # (1) audio-separator — exact, as used to build the cache
from audio_separator.separator import Separator
s = Separator(output_dir="/tmp/app_sep", model_file_dir="/tmp/audio-separator-models")
s.load_model("htdemucs.yaml")
_SEP.update(kind="audio_separator", sep=s)
return _SEP
except Exception:
pass
try: # (2) demucs package — SAME htdemucs v4 weights
from demucs.pretrained import get_model # NB: on the Space there is no root demucs.py to shadow this
from demucs.apply import apply_model
m = get_model("htdemucs").to(DEVICE).eval()
_SEP.update(kind="demucs", model=m, apply=apply_model, sources=list(m.sources))
return _SEP
except Exception as e:
raise RuntimeError(
"Live upload separation is unavailable: htdemucs v4 (audio-separator or the demucs package) "
"isn't installed here. Use the 🎧 Val samples tab — full restore+generate pipeline. "
"(torchaudio's HDemucs is v3 and would NOT match the trained stems.)"
) from e
# ---- separation cache (OPT-IN: set RESTOFLOW_SEP_CACHE=<dir>) ----
# Separation is identical across restorer/generator VARIANTS (same degraded input -> same htdemucs stems),
# so an experiment that scores N variants re-separates the SAME clips N times. This content-addressed cache
# stores each unique input's stems ONCE (deduped by input-hash) and reuses them. Stored as int16 = byte-exact
# to what audio_separator already writes (16-bit PCM) -> NO added error vs the live path, half the disk of
# float32. UNSET env -> behavior is byte-identical to before (the UI/Space never sets it).
_SEP_CACHE_VER = "htdemucs_v4_a" # bump to invalidate all cached stems
def _sep_cache_dir():
d = os.environ.get("RESTOFLOW_SEP_CACHE")
return Path(d) if d else None
def _sep_cache_key(seg_2S, kind):
a = np.ascontiguousarray(np.asarray(seg_2S, dtype=np.float32))
h = hashlib.sha1(a.tobytes())
h.update(f"|{kind}|{a.shape}|{_SEP_CACHE_VER}".encode())
return h.hexdigest()
def _sep_cache_load(cdir, key):
# 4 stems stored as one 8-channel PCM16 FLAC (lossless, ~40% smaller than zlib-int16). Channel order =
# sorted(STEMS); each stem = 2 consecutive channels. PCM16/float read = exactly what the live wav path
# returns (k/32768), so a hit is byte-identical to live separation.
f = cdir / f"{key}.flac"
if not f.exists():
return None
try:
inter, _ = sf.read(f, dtype="float32") # [L, 8] in [-1,1)
order = sorted(STEMS)
return {order[i]: np.ascontiguousarray(inter[:, 2 * i:2 * i + 2].T) for i in range(len(order))}
except Exception:
return None # corrupt -> recompute live
def _sep_cache_save(cdir, key, stems):
try:
cdir.mkdir(parents=True, exist_ok=True)
order = sorted(STEMS)
L = next((np.asarray(v).shape[-1] for v in stems.values()), 0)
chans = []
for s in order: # canonical order; missing -> silence
v = stems.get(s)
a = np.asarray(v, np.float32) if v is not None else np.zeros((2, L), np.float32)
if a.shape[0] != 2:
a = a[:2] if a.shape[0] > 2 else np.repeat(a, 2, 0)
chans.append(np.clip(np.round(a * 32768.0), -32768, 32767).astype(np.int16))
inter = np.concatenate(chans, axis=0).T # [L, 8] int16
tmp = cdir / f".{key}.{os.getpid()}.tmp.flac" # atomic write (no torn files)
sf.write(tmp, inter, SR, format="FLAC", subtype="PCM_16"); os.replace(tmp, cdir / f"{key}.flac")
except Exception:
pass # cache is best-effort, never fatal
def separate_chunk(sep, seg_2S):
"""[2,S] @ SR -> {stem:[2,S]} via htdemucs v4. Cache-aware (see _sep_cache_*): on a hit the stems are
loaded from disk instead of re-separated, so multi-variant experiment evals separate each clip once."""
cdir = _sep_cache_dir()
key = _sep_cache_key(seg_2S, sep["kind"]) if cdir is not None else None
if cdir is not None:
hit = _sep_cache_load(cdir, key)
if hit is not None:
return hit
out = _separate_chunk_impl(sep, seg_2S)
if cdir is not None:
_sep_cache_save(cdir, key, out)
return out
def _separate_chunk_impl(sep, seg_2S):
"""[2,S] @ SR -> {stem:[2,S]} via the same htdemucs v4 weights the cache used. CPU-tuned: shifts=0
(no test-time augmentation — the big needless cost), light overlap, no_grad (demucs does in-place ops
that throw under inference_mode). Called ONCE on the whole clip (demucs splits internally) — NOT per
overlapping window — so live separation stays ~1 forward instead of N."""
if sep["kind"] == "audio_separator":
# per-process input path: distinct output filenames -> no race when another process (e.g. a precache
# pass on another GPU) separates concurrently into the shared /tmp/app_sep dir.
tmp = f"/tmp/app_sep_in_{os.getpid()}.wav"; sf.write(tmp, seg_2S.T, SR)
files = sep["sep"].separate(tmp)
out = {}
for f in files:
fl = f.lower()
for s in STEMS:
if s in fl:
au, _ = sf.read(os.path.join("/tmp/app_sep", f) if not os.path.isabs(f) else f, dtype="float32")
out[s] = (au.T if au.ndim == 2 else np.repeat(au[None], 2, 0))
return out
wav = torch.as_tensor(seg_2S, dtype=torch.float32, device=DEVICE) # [2,S]
ref = wav.mean(0); mu, sd = ref.mean(), ref.std().clamp_min(1e-8)
with torch.no_grad():
src = sep["apply"](sep["model"], ((wav - mu) / sd)[None],
shifts=0, overlap=0.1, split=True, device=DEVICE)[0] # [4,2,S]
src = (src * sd + mu).cpu().numpy()
return {s: src[i] for i, s in enumerate(sep["sources"]) if s in STEMS}
EXPLAIN = """
## How this works (high level)
**The problem.** Old/lo-fi recordings are *degraded* (muffled, missing instruments). We work in
**SAME-L's latent space** — a pretrained **neural audio autoencoder** (a learned codec, à la the
VAEs behind latent diffusion) that compresses 3 s of stereo into a small `256×32` "summary"
(~11 frames/sec). Everything below operates on these latents, then decodes back to audio — so the
models stay tiny and fast.
**The models, in plain terms:**
- **SAME-L** — the compact "latent" space everything runs in (a frozen neural audio codec: encode ↔ decode).
- **htdemucs** — splits a song into its 4 instrument stems.
- **Restorer** — *cleans up* a muffled/damaged stem so it sounds clear again. It fixes what's there; it doesn't add new parts. The default ("glue") restorer is a deterministic Conformer-style backbone **plus a small residual-flow detail head** that samples back fine high-frequency detail; the plain backbone alone is selectable as the **baseline** (*advramp*).
- **Generator** — *invents* a missing instrument (mainly bass) that fits the song, like an AI session musician. (Conditional flow-matching.)
- **Glue stage (training-time)** — the restorer + generator are co-trained so their *summed* mix matches a real clean mix, judged by a **multi-scale latent discriminator** with feature-matching. This is a coherence objective on the recombined mix; in practice it lands close to the deterministic baseline (the two are A/B-selectable so you can compare).
**Step 1 — Separate.** htdemucs splits the input mix into 4 stems: *other, vocals, drums, bass*.
**Step 2 — Route.** For each stem we ask: *is it actually there?* (energy gate).
- **Present** → it just needs cleanup → **RESTORE**.
- **Missing bass** → an absent bass is **invented** → **GENERATE**. (Note: this fills a low end even for songs that may not originally have a bass instrument — a known limitation on bass-free material.)
- **Missing drums/other** → a song may legitimately have none, so we **do not fabricate** them — restore-only.
**Step 3a — Restorer.** For present stems, a small **Conformer-style network** (local convolutions
for texture **+ global self-attention** so each moment sees the whole clip) nudges the degraded
latent toward a clean one, as a *residual* on the input — mostly *deterministic* because the degraded
signal already tells us the answer. The default model adds a **residual-flow head** that *samples* a
small extra detail residual on top (drums/other/vocals). Trained with a **stem-balanced** loss so
quiet stems (bass) aren't drowned out by loud ones.
**Step 3b — Generator (creative, bass).** When the bass is missing there is no answer to copy —
many different bass lines could fit. So this is a **conditional flow-matching** generator: starting
from noise it follows a learned velocity field to *sample* a plausible **bass conditioned on the
other (restored) stems**, so it locks to the song's harmony. **Classifier-free guidance** lets us
dial how strongly it commits to the accompaniment. (Probes showed bass is strongly determined by the
accompaniment and the codec preserves it well; drums are timing-limited by the ~11 Hz latent, so we
don't fabricate drums.)
**Step 4 — Decode & remix.** Each final latent is decoded to audio and summed into the output mix,
with per-window level smoothing and crossfaded overlap-add across the 3 s windows.
> Restoration leans on the *existing* signal (works best when degradation is mild). Generation
> *creates* absent instruments to accompany what's there. Inference only ever sees restored stems —
> never the clean originals. The **baseline** (deterministic *advramp* restorer + *gen_advramp_v1*) is
> selectable for A/B against the default glue models.
"""
THEME = gr.themes.Base(
primary_hue=gr.themes.colors.orange,
secondary_hue=gr.themes.colors.orange,
neutral_hue=gr.themes.colors.zinc,
font=[gr.themes.GoogleFont("Inter"), "system-ui", "sans-serif"],
radius_size=gr.themes.sizes.radius_sm,
).set(
# dark look as the default (set the light variants to dark so it's black/orange regardless of OS)
body_background_fill="#0a0a0b",
body_text_color="#e9e9ec",
body_text_color_subdued="#8a8a93",
background_fill_primary="#151518",
background_fill_secondary="#1c1c20",
block_background_fill="#151518",
block_border_color="#272730",
block_label_background_fill="#151518",
block_label_text_color="#ff8a3d",
block_title_text_color="#f2f2f4",
border_color_primary="#272730",
input_background_fill="#1c1c20",
button_primary_background_fill="#ff7a18",
button_primary_background_fill_hover="#ff9344",
button_primary_text_color="#0a0a0b",
button_secondary_background_fill="#272730",
button_secondary_text_color="#e9e9ec",
color_accent_soft="#2a1a0e",
slider_color="#ff7a18",
)
CSS = """
.gradio-container {max-width: 1080px !important; margin: auto !important; background:#0a0a0b;}
.fullmix audio {width: 100%;}
h1,h2,h3,h4 {color:#f4f4f6 !important; letter-spacing:.2px;}
h1 {font-weight:700; border-left:4px solid #ff7a18; padding-left:12px;}
h4 {color:#ff9a55 !important; text-transform:uppercase; font-size:.8rem; letter-spacing:.6px;}
a {color:#ff8a3d !important;}
footer {display:none !important;}
::-webkit-scrollbar{width:9px;height:9px} ::-webkit-scrollbar-thumb{background:#3a3a44;border-radius:6px}
"""
def _mix_column(title, with_plot=True):
gr.Markdown(f"#### {title}")
a = gr.Audio(label=None, show_label=False, elem_classes="fullmix")
p = gr.Plot(show_label=False) if with_plot else None
return a, p
def build_ui():
# NOTE: do NOT load models here — that delays the port bind ~15 s and makes the
# auto-opened browser tab hit a dead port. UI uses cheap dir listing; models load
# lazily on first action (warmed in a background thread at launch).
rests = list_runs("restorer"); gens = list_runs("generator")
dr, dg = _default_runs(rests, gens)
with gr.Blocks(title="SAME Restore + Generate", theme=THEME, css=CSS) as demo:
gr.Markdown("# 🎛️ Stem Restoration + Generation\n"
"Upload a muffled/lo-fi song. We split it into instruments, **clean up** the ones that are "
"there, **invent** a missing bass that fits, and remix — all in a compact neural-audio space. "
"A GAN 'judge' keeps the result sounding *real*, not dull.")
status = gr.Markdown(status_md())
with gr.Accordion("⚙️ Choose models", open=False):
with gr.Row():
rest_dd = gr.Dropdown(rests, value=dr, label="Restorer", scale=2)
gen_dd = gr.Dropdown(gens + ["(none)"], value=dg, label="Generator (bass)", scale=2)
load_btn = gr.Button("↻ Load", variant="primary", scale=1)
load_btn.click(set_models, [rest_dd, gen_dd], status)
with gr.Tab("🎧 Val samples"):
with gr.Row():
dd = gr.Dropdown([val_label(x) for x in val_list()], label="3 s val sample (each shows both restore + generate)", scale=5)
btn = gr.Button("▶ Run", variant="primary", scale=1)
rep = gr.Markdown()
with gr.Row(equal_height=True):
with gr.Column(): a_in, s_in = _mix_column("🎚️ Degraded input")
with gr.Column(): a_out, s_out = _mix_column("✨ Output")
with gr.Column(): a_ref, s_ref = _mix_column("🎯 Clean reference")
with gr.Accordion("🔬 Per-stem detail (before → after)", open=False):
comps = []
for s in STEMS:
with gr.Row(equal_height=True):
lbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**")
bef = gr.Audio(label="before (degraded)", show_label=True, scale=2)
aft = gr.Audio(label="after (restored/generated)", show_label=True, scale=2)
comps += [lbl, bef, aft]
btn.click(run_val, dd, [a_in, a_out, a_ref, s_in, s_out, s_ref, rep] + comps)
with gr.Tab("⬆️ Upload"):
with gr.Row():
up = gr.Audio(label="Upload wav/mp3 (chunked into 3 s; last chunk padded then cropped)",
type="filepath", sources=["upload"], editable=True,
waveform_options=gr.WaveformOptions(
waveform_color="#5a5a66", waveform_progress_color="#ff7a18",
trim_region_color="#ff7a18"),
scale=5)
ub = gr.Button("▶ Run", variant="primary", scale=1)
with gr.Row():
trim_start = gr.Number(value=0, label="Trim start (s)", scale=1)
trim_end = gr.Number(value=0, label="Trim end (s, 0 = to end)", scale=1)
gr.Markdown("*Set start/end to actually chop the file before processing "
"(leave both 0 to process the whole upload).*")
urep = gr.Markdown()
with gr.Row(equal_height=True):
with gr.Column(): ua_in, us_in = _mix_column("⬆️ Input")
with gr.Column(): ua_out, us_out = _mix_column("✨ Output")
with gr.Accordion("🔬 Per-stem detail (output)", open=False):
ucomps = []
for s in STEMS:
with gr.Row(equal_height=True):
ulbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**")
uaft = gr.Audio(label="output stem", show_label=True, scale=3)
ucomps += [ulbl, uaft]
ub.click(run_upload, [up, trim_start, trim_end], [ua_in, ua_out, us_in, us_out, urep] + ucomps)
with gr.Accordion("ℹ️ How this works (restorer + generator, intuitively)", open=False):
gr.Markdown(EXPLAIN)
return demo
def _free_port(port):
"""Kill any process still holding the port so re-launch binds cleanly."""
try:
out = os.popen(f"fuser {port}/tcp 2>/dev/null").read().split()
for pid in out:
if pid.strip().isdigit() and int(pid) != os.getpid():
os.system(f"kill -9 {pid} 2>/dev/null")
except Exception:
pass
if __name__ == "__main__":
port = int(os.environ.get("PORT", 7860))
_free_port(port) # clear a stale bind from a previous run
demo = build_ui()
import threading
threading.Thread(target=models, daemon=True).start() # warm models in bg; port still binds instantly
try:
demo.launch(server_name="0.0.0.0", server_port=port, share=False,
ssr_mode=False) # ssr_mode off: Gradio-6 SSR hangs without a node runtime
except KeyboardInterrupt:
print("\n[app] shutting down…")
finally: # flush the port on exit (Ctrl+C / close)
try: demo.close()
except Exception: pass
try: gr.close_all()
except Exception: pass
_free_port(port)
print("[app] port released.")