Spaces:
Sleeping
Sleeping
Download restoflow/app.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 45.1 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/app.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/app.py
-
curl -L -o app.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/app.py
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) | |
| 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] | |
| 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 | |
| def decode(latent): | |
| return E._decode(models()["sa"], _fitT(latent)) # [2,S] cpu | |
| def encode(audio_2S): | |
| return encode_batch([audio_2S])[0] | |
| 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])] | |
| 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.") | |