"""Stage-A for MUSDB18-HQ: build a SAME-latent restoration cache in the SAME schema as the rebetiko cache, so train.py / data.py consume it UNCHANGED (just point --cache-roots at it). Per 3s excerpt of a MUSDB song: clean_mix = mixture (full-band) -> deg_mix = robustdeg_v1(clean_mix) [time-aligned] clean stems = htdemucs(clean_mix) (TARGET; matches the train convention: target = htdemucs(clean)) deg stems = htdemucs(deg_mix) (INPUT) encode each stem + the deg mix -> latents; write metadata rows (rms/peak dB per stem, for the router). Song-level split (NO excerpt leakage): train songs -> train/val caches; official test songs -> test cache. Uses RESTOFLOW_SEP_CACHE (fast separation cache). Run per split; 2 splits can run on 2 GPUs in parallel: RESTOFLOW_SEP_CACHE=/sep_cache python -m restoflow.build_musdb_cache --split train --device cuda:0 RESTOFLOW_SEP_CACHE=/sep_cache python -m restoflow.build_musdb_cache --split test --device cuda:1 """ from __future__ import annotations import argparse, csv, glob, os, random from pathlib import Path import numpy as np import soundfile as sf BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") MUSDB = BASE / "data" / "musdb18hq" STEMS = ("other", "vocals", "drums", "bass") COLS = ["pair_id", "stem", "class", "clean_present", "deg_present", "clean_rms_db", "deg_rms_db", "clean_peak_db", "deg_peak_db", "clean_latent_path", "deg_latent_path", "latent_dim", "n_frames", "sample_rate", "frame_rate", "duration_s"] def _read(fp, SR): import librosa a, sr = sf.read(str(fp), dtype="float32"); a = a.T if a.ndim == 2 else np.stack([a, a]) if a.ndim == 1: a = np.stack([a, a]) if sr != SR: a = np.stack([librosa.resample(a[c], orig_sr=sr, target_sr=SR) for c in range(2)]) return np.ascontiguousarray(a.astype(np.float32)) def load_song(song_dir: Path, SR: int): """Return (mix[2,N], {stem:[2,N]}) — TRUE clean stems (MUSDB ground truth = the restoration target).""" stems = {s: _read(song_dir / f"{s}.wav", SR) for s in STEMS if (song_dir / f"{s}.wav").exists()} if len(stems) < len(STEMS): return None, None n = min(v.shape[1] for v in stems.values()); stems = {s: v[:, :n] for s, v in stems.items()} mp = song_dir / "mixture.wav" mix = _read(mp, SR)[:, :n] if mp.exists() else np.ascontiguousarray(sum(stems.values())) return np.ascontiguousarray(mix.astype(np.float32)), stems def pick_segments(mix, SR, seg_s, max_seg, rng, rms_floor=1e-3): """Non-overlapping 3s excerpts with enough energy; up to max_seg, shuffled selection for variety.""" L = int(seg_s * SR); n = mix.shape[1] // L cands = [] for i in range(n): s = i * L; seg = mix[:, s:s + L] if float(np.sqrt(np.mean(seg ** 2))) >= rms_floor: cands.append(s) rng.shuffle(cands) return sorted(cands[:max_seg]) def main(): p = argparse.ArgumentParser() p.add_argument("--split", required=True, choices=["train", "val", "test"]) p.add_argument("--device", default="cuda:0") p.add_argument("--seg-s", type=float, default=3.0) p.add_argument("--max-seg", type=int, default=30, help="excerpts per song (train); test/val auto-lower") p.add_argument("--val-songs", type=int, default=10, help="songs held out from train -> val split") p.add_argument("--max-songs", type=int, default=0, help="limit #songs (0=all); for smoke tests") p.add_argument("--seed", type=int, default=0) p.add_argument("--out-root", default=str(BASE / "data" / "musdb_cache")) a = p.parse_args() os.environ["RESTOFLOW_DEVICE"] = a.device os.environ.setdefault("RESTOFLOW_SEP_CACHE", str(BASE / "sep_cache")) import restoflow.app as APP from restoflow.robustdeg import degrade, RobustDegConfig SR = APP.SR; dcfg = RobustDegConfig(); APP.models() sep = APP.get_separator() # --- song-level split (no excerpt leakage) --- src = "train" if a.split in ("train", "val") else "test" songs = sorted([d for d in (MUSDB / src).glob("*") if d.is_dir()]) if a.split in ("train", "val"): rng = random.Random(12345); shuffled = songs[:]; rng.shuffle(shuffled) val_set = set(shuffled[:a.val_songs]) songs = [s for s in songs if (s in val_set) == (a.split == "val")] if a.max_songs: songs = songs[:a.max_songs] max_seg = a.max_seg if a.split == "train" else max(4, a.max_seg // 3) out_root = Path(a.out_root) / a.split; out_root.mkdir(parents=True, exist_ok=True) meta_fp = out_root / "metadata.csv" write_header = not meta_fp.exists() mf = open(meta_fp, "a", newline=""); w = csv.writer(mf) if write_header: w.writerow(COLS) print(f"[musdb-cache] split={a.split} songs={len(songs)} max_seg/song={max_seg} -> {out_root}", flush=True) import torch npairs = 0; L = int(a.seg_s * SR) for si, song in enumerate(songs): mix, stems_full = load_song(song, SR) if mix is None: continue segs = pick_segments(mix, SR, a.seg_s, max_seg, random.Random(a.seed + si)) for ti, s0 in enumerate(segs): pair_id = f"{a.split[0]}{si:03d}_{ti:03d}" pdir = out_root / pair_id / "latents" if (pdir / "degraded_mix.pt").exists(): # resumable: skip done excerpts npairs += 1; continue clean_mix = np.ascontiguousarray(mix[:, s0:s0 + L]) deg, _ = degrade(clean_mix, SR, random.Random(a.seed * 7919 + si * 131 + ti), dcfg) cs = {s: np.ascontiguousarray(stems_full[s][:, s0:s0 + L]) for s in STEMS} # TRUE clean stems (target) ds = APP.separate_chunk(sep, deg) # htdemucs on the degraded mix (input) — 1 separation # batch-encode: 4 clean stems + 4 deg stems + deg mix order = list(STEMS) auds = [cs[s] for s in order] + [ds[s] for s in order] + [deg] lat = APP.encode_batch(auds) clat = {order[i]: lat[i] for i in range(4)} dlat = {order[i]: lat[4 + i] for i in range(4)} mlat = lat[8] (pdir / "clean").mkdir(parents=True, exist_ok=True); (pdir / "degraded").mkdir(parents=True, exist_ok=True) torch.save(mlat, pdir / "degraded_mix.pt") T = int(mlat.shape[1]) for s in order: torch.save(clat[s], pdir / "clean" / f"{s}.pt") torch.save(dlat[s], pdir / "degraded" / f"{s}.pt") crp = APP.rms_peak_db(cs[s].mean(0)); drp = APP.rms_peak_db(ds[s].mean(0)) cpres = int(crp[0] > -40 or crp[1] > -25); dpres = int(drp[0] > -40 or drp[1] > -25) cls = "restoration" if (cpres and dpres) else ("both_silent" if not (cpres or dpres) else "other") w.writerow([pair_id, s, cls, cpres, dpres, f"{crp[0]:.2f}", f"{drp[0]:.2f}", f"{crp[1]:.2f}", f"{drp[1]:.2f}", f"{pair_id}/latents/clean/{s}.pt", f"{pair_id}/latents/degraded/{s}.pt", 256, T, SR, f"{T/a.seg_s:.3f}", a.seg_s]) mf.flush(); npairs += 1 if (si + 1) % 5 == 0: print(f" song {si+1}/{len(songs)} · pairs={npairs}", flush=True) mf.close() print(f"[musdb-cache] DONE split={a.split}: {npairs} excerpt-pairs -> {out_root}") print("MUSDB_CACHE_DONE") if __name__ == "__main__": main()