"""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: /restoflow_runs, /demo_samples). _LOCAL_BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") _REPO_ROOT = Path(__file__).resolve().parent.parent # 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=) ---- # 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.")