"""Dataset + normalization stats for SAME-latent restoration. - Reads metadata.csv from one or more cache roots. - Keeps only RESTORATION (pair,stem) under hardened presence (router.is_restoration). - Deterministic by-pair train/val split (or explicit val_cache_roots). - Per-stem per-dim z-score normalization from TRAIN clean latents. Item: dict(deg[256,T], clean[256,T] standardized, stem_id, stem, pair_id, root). """ from __future__ import annotations import csv from pathlib import Path import torch from torch.utils.data import Dataset from .config import Cfg, STEM_ID from . import router def _index(cfg: Cfg, roots, split_mode): """Collect item dicts. split_mode: 'train'/'val'/'all'. If cfg.val_cache_roots set, roots are used as-is for the requested split; else by-pair hash split.""" items = [] explicit_val = bool(cfg.val_cache_roots) for root in roots: root = Path(root) meta = root / "metadata.csv" if not meta.exists(): continue for row in csv.DictReader(open(meta)): if row.get("stem") not in cfg.stems: continue if not router.is_restoration(row, cfg.rms_thr, cfg.peak_thr, cfg.max_drop_db): continue if row.get("clean_latent_path") in ("", None) or row.get("deg_latent_path") in ("", None): continue if not explicit_val and split_mode in ("train", "val"): if router.pair_split(row["pair_id"], cfg.val_frac, cfg.split_seed) != split_mode: continue mix = root / row["pair_id"] / "latents" / "degraded_mix.pt" if cfg.use_mix and not mix.exists(): continue items.append({ "pair_id": row["pair_id"], "stem": row["stem"], "root": str(root), "clean": str(root / row["clean_latent_path"]), "deg": str(root / row["deg_latent_path"]), "mix": str(mix), }) return items def _fit_T(x: torch.Tensor, T: int) -> torch.Tensor: # x [256, t] -> [256, T] t = x.shape[1] if t == T: return x if t > T: return x[:, :T] return torch.cat([x, x[:, -1:].repeat(1, T - t)], dim=1) # edge-repeat pad class LatentRestore(Dataset): def __init__(self, cfg: Cfg, split: str, stats: dict): self.cfg = cfg roots = cfg.val_cache_roots if (split == "val" and cfg.val_cache_roots) else cfg.cache_roots self.items = _index(cfg, roots, split) self.stats = stats # stem -> {'mu':[256], 'sd':[256]} def __len__(self): return len(self.items) def _norm(self, x, stem): mu, sd = self.stats[stem]["mu"], self.stats[stem]["sd"] return (x - mu[:, None]) / sd[:, None] def __getitem__(self, i): it = self.items[i] clean = _fit_T(torch.load(it["clean"], map_location="cpu").float(), self.cfg.T) deg = _fit_T(torch.load(it["deg"], map_location="cpu").float(), self.cfg.T) out = { "deg": self._norm(deg, it["stem"]), "clean": self._norm(clean, it["stem"]), "stem_id": torch.tensor(STEM_ID[it["stem"]], dtype=torch.long), "stem": it["stem"], "pair_id": it["pair_id"], } if self.cfg.use_mix: mix = _fit_T(torch.load(it["mix"], map_location="cpu").float(), self.cfg.T) m = self.stats["__mix__"] out["mix"] = (mix - m["mu"][:, None]) / m["sd"][:, None] return out def build_norm_stats(cfg: Cfg) -> dict: """Per-stem per-dim mean/std from TRAIN clean latents. Saved to cfg.stats_file().""" items = _index(cfg, cfg.cache_roots, "train") acc = {s: {"n": 0, "sum": torch.zeros(cfg.latent_dim, dtype=torch.float64), "sq": torch.zeros(cfg.latent_dim, dtype=torch.float64)} for s in cfg.stems} for it in items: x = torch.load(it["clean"], map_location="cpu").double() # [256,t] a = acc[it["stem"]] a["n"] += x.shape[1]; a["sum"] += x.sum(1); a["sq"] += (x * x).sum(1) stats = {} for s, a in acc.items(): if a["n"] == 0: stats[s] = {"mu": torch.zeros(cfg.latent_dim), "sd": torch.ones(cfg.latent_dim), "n": 0} continue mu = a["sum"] / a["n"] var = (a["sq"] / a["n"] - mu * mu).clamp_min(1e-8) stats[s] = {"mu": mu.float(), "sd": var.sqrt().float(), "n": a["n"]} if cfg.use_mix: # stem-independent mix stats (one degraded_mix per pair) n = 0; ssum = torch.zeros(cfg.latent_dim, dtype=torch.float64); ssq = torch.zeros(cfg.latent_dim, dtype=torch.float64) seen = set() for it in items: if it["pair_id"] in seen or not Path(it["mix"]).exists(): continue seen.add(it["pair_id"]); x = torch.load(it["mix"], map_location="cpu").double() n += x.shape[1]; ssum += x.sum(1); ssq += (x * x).sum(1) if n: mu = ssum / n; var = (ssq / n - mu * mu).clamp_min(1e-8) stats["__mix__"] = {"mu": mu.float(), "sd": var.sqrt().float(), "n": n} out = cfg.stats_file(); out.parent.mkdir(parents=True, exist_ok=True) torch.save(stats, out) print(f"[stats] frames/stem:", {s: stats[s]["n"] for s in cfg.stems}, "-> saved", out) return stats def load_or_build_stats(cfg: Cfg) -> dict: f = cfg.stats_file() return torch.load(f) if f.exists() else build_norm_stats(cfg) if __name__ == "__main__": c = Cfg() build_norm_stats(c) for sp in ("train", "val"): ds = LatentRestore(c, sp, load_or_build_stats(c)) from collections import Counter print(sp, "items:", len(ds), dict(Counter(it["stem"] for it in ds.items)))