File size: 45,098 Bytes
af4583e
 
 
 
 
 
 
 
 
 
 
 
7153194
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
af4583e
 
7153194
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
 
 
 
af4583e
 
 
7153194
 
 
 
af4583e
 
 
 
7153194
 
 
 
 
 
 
af4583e
 
7153194
af4583e
 
 
7153194
 
 
 
 
af4583e
7153194
af4583e
 
 
 
 
 
 
7153194
 
 
 
 
af4583e
 
 
 
 
 
 
7153194
 
 
 
 
 
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
 
 
 
af4583e
7153194
 
 
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4e73c38
 
 
 
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39e8810
 
 
af4583e
 
 
 
 
 
 
39e8810
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af4583e
7153194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39e8810
 
 
 
af4583e
7153194
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
39e8810
af4583e
39e8810
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
af4583e
7153194
af4583e
 
 
 
 
7153194
af4583e
 
7153194
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
7153194
 
af4583e
 
 
7153194
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
"""Gradio demo for the SAME-latent stem RESTORER (v4b) + GENERATOR pipeline.

Unified pipeline: a 3s clip's degraded full-mix -> htdemucs stems -> SAME latents ->
  - present stems  : v4b deterministic RESTORE
  - missing stems  : conditional-flow GENERATE from the (restored) context
-> SAME decode -> remixed output. Val samples use the cached latents (fast); uploads run
the full htdemucs+SAME path and are chunked into 3s (last chunk padded with silence then cropped).

Runs on CPU by default (set RESTOFLOW_DEVICE=cuda to use a GPU).
Launch:  python -m restoflow.app
"""
from __future__ import annotations
import os, csv, glob, io, hashlib
# Subprocess ffmpeg fix (audio-separator + librosa/audioread shell out to `ffmpeg`):
# a stale /home/maximos/miniconda3/lib on LD_LIBRARY_PATH shadows system libstdc++ and
# breaks the system ffmpeg; drop it and put a known-good conda ffmpeg first on PATH.
os.environ["LD_LIBRARY_PATH"] = ":".join(
    p for p in os.environ.get("LD_LIBRARY_PATH", "").split(":") if p and "maximos/miniconda3" not in p)
_FFMPEG_DIR = "/home/ksoil/.conda/envs/ksoil_torch/bin"
if os.path.isdir(_FFMPEG_DIR):
    os.environ["PATH"] = _FFMPEG_DIR + ":" + os.environ.get("PATH", "")
from pathlib import Path
import numpy as np
import torch
import soundfile as sf
import librosa
import librosa.display
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import gradio as gr

STEM_EMOJI = {"other": "🎹", "vocals": "🎀", "drums": "πŸ₯", "bass": "🎸"}

from .config import Cfg, STEMS, STEM_ID
from .model import DetRestorer, CondFlow, AttnRestorer, AttnCondFlow, MixAttnCondFlow
from . import eval as E
from . import gen as G
from . import flow as F

# Paths are env-overridable so the same file runs locally and on a packaged HF Space.
# Local default = the dev tree; on HF the deploy bundle sets these (or we fall back to a
# repo-relative layout: <repo>/restoflow_runs, <repo>/demo_samples).
_LOCAL_BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra")
_REPO_ROOT = Path(__file__).resolve().parent.parent          # <repo> containing restoflow/
BASE = Path(os.environ.get("RESTOFLOW_BASE", _LOCAL_BASE if _LOCAL_BASE.exists() else _REPO_ROOT))
SRC = Path(os.environ.get(
    "RESTOFLOW_SRC",
    "/media/maindisk/melkor169/MSRKit/xlance-msr/stereo_a2sb/synthetic_out"))
if not SRC.exists():                                         # packaged Space: demo wavs live here
    SRC = BASE / "demo_samples"
VAL_CACHE = Path(os.environ.get("RESTOFLOW_VAL_CACHE", str(BASE / "demucs_results_val_full")))
if not VAL_CACHE.exists():
    VAL_CACHE = BASE / "demo_samples" / "val_cache"
DEVICE = os.environ.get("RESTOFLOW_DEVICE", "cpu")
T, D = 32, 256
SR = 44100
N_VAL = 80
PRESENT_RMS, PRESENT_PEAK = -40.0, -25.0   # presence gate (rms OR peak), matches training
# Generation policy: every song has a low end, so a MISSING bass is always invented.
# A missing drums/other is legitimate (song may have none) -> we DON'T fabricate it; restore-only.
GEN_TARGETS = ("bass",)

_M = {}   # model cache


# ---------------- model loading ----------------
def _disable_flash_attn():
    """SAME's transformer uses flash_attn (CUDA-only). On CPU, force the SDPA fallback
    by nulling the module globals apply_attn() checks at call time."""
    import stable_audio_tools.models.transformer as _tf
    for name in ("flash_attn_func", "flash_attn_kvpacked_func", "flash_attn_varlen_func", "index_first_axis"):
        if hasattr(_tf, name):
            setattr(_tf, name, None)


def _ckpt_path(run):
    d = BASE / "restoflow_runs" / run
    for f in ("ckpt_best.pt", "ckpt.pt"):
        if (d / f).exists():
            return d / f
    return None


def _default_runs(rests, gens):
    """Preselected restorer/generator. Env override (set by the HF deploy) wins, else a
    'prefer newest best' heuristic, else first available."""
    dr = os.environ.get("RESTOFLOW_DEFAULT_REST")
    dg = os.environ.get("RESTOFLOW_DEFAULT_GEN")
    if dr not in rests:
        # remix_glue_v1 is the published default: frozen advramp backbone + a residual-flow detail head,
        # co-trained with the generator against a multi-scale latent discriminator for a coherent remix.
        # restorer_attn_advramp (the deterministic backbone alone) is kept selectable as the BASELINE.
        dr = next((r for r in ("remix_glue_v1", "restorer_attn_advramp", "restorer_attn_w003_gan",
                               "restorer_attn_v1", "v4b_balonly") if r in rests),
                  rests[0] if rests else None)
    if dg not in gens:
        # remix_glue_v1_gen: the generator co-trained in the glue stage with remix_glue_v1. gen_advramp_v1
        # (trained on advramp's conditioning) is the BASELINE generator. Both conv 10M (fast on CPU).
        dg = next((g for g in ("remix_glue_v1_gen", "gen_advramp_v1", "gen_distvar_baseline",
                               "gen_v2", "gen_v3") if g in gens),
                  gens[0] if gens else "(none)")
    return dr, dg


# Validated runs promoted past the experiment filter below. The glue model (remix_glue_v1: advramp
# backbone + residual-flow detail head, co-trained with remix_glue_v1_gen) is the published default;
# advramp + gen_advramp_v1 stay selectable as the baseline.
PROMOTED_REST = {"remix_glue_v1"}
PROMOTED_GEN = {"remix_glue_v1_gen"}


def list_runs(kind):
    """Available run names (excluding obsolete) for kind in {'restorer','generator'}.
    Convention: gen_* are generators; everything else is a restorer (plus explicit PROMOTED_GEN)."""
    out = []
    for d in sorted(glob.glob(str(BASE / "restoflow_runs" / "*"))):
        name = Path(d).name
        if name == "obsolete" or not Path(d).is_dir() or _ckpt_path(name) is None:
            continue
        promoted = name in PROMOTED_REST or name in PROMOTED_GEN
        # skip scratch/experiment dirs (scout_*, _smoke_*, remix_* GAN-refine runs) unless explicitly promoted
        if not promoted and name.startswith(("scout", "_", "remix", "exp")):
            continue
        is_gen = name.startswith("gen") or name in PROMOTED_GEN
        if (kind == "generator") == is_gen:
            out.append(name)
    return out


def load_restorer(run):
    ck = torch.load(_ckpt_path(run), map_location="cpu"); rc = ck["cfg"]
    kind = rc.get("model_kind", "det")
    if kind == "flowattn":                                   # GENERATIVE restorer (deg-anchored flow + mix)
        m = MixAttnCondFlow(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
                            stem_emb=rc["stem_emb_dim"], heads=rc.get("heads", 8), use_mix=rc["use_mix"])
    elif kind == "attn":                                     # Conformer-lite deterministic restorer
        m = AttnRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
                         stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"], heads=rc.get("heads", 4))
    else:
        m = DetRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS),
                        stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"])
    m.load_state_dict(ck["model"]); m.eval().to(DEVICE)
    stats = torch.load(BASE / "restoflow_runs" / run / "norm_stats.pt", map_location="cpu")
    d = dict(rest=m, rstats=stats, rest_run=run, rest_ep=ck.get("epoch"), rest_kind=kind,
             rest_steps=int(rc.get("sample_steps", 40)), rest_sigma=float(rc.get("sigma", 0.0)),
             rest_use_mix=bool(rc.get("use_mix", True)),
             rest_params=sum(p.numel() for p in m.parameters()) / 1e6)
    if "rflow" in ck:                                        # glue-stage residual-flow detail head (det backbone + sampled residual)
        rfc = ck["rflow_cfg"]
        rf = CondFlow(D, rfc["hidden"], rfc["depth"], n_stems=len(STEMS), stem_emb=rfc["stem_emb"])
        rf.load_state_dict(ck["rflow"]); rf.eval().to(DEVICE)
        d["rflow"] = rf; d["rflow_ids"] = set(rfc["restore_ids"]); d["rflow_steps"] = int(rfc.get("gen_steps", 2))
        d["rest_params"] += sum(p.numel() for p in rf.parameters()) / 1e6
    return d


def load_generator(run):
    if run in (None, "(none)"):
        return dict(gen=None, gstats=None, gargs=None, gen_run="(none)", gen_ep=None, gen_params=0.0)
    ck = torch.load(_ckpt_path(run), map_location="cpu"); ga = ck["args"]
    if ga.get("arch") == "attn":                             # Conformer-lite velocity net
        m = AttnCondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS),
                         stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8))
    else:
        m = CondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim)
    m.load_state_dict(ck["model"]); m.eval().to(DEVICE)
    return dict(gen=m, gstats=ck["stats"], gargs=ga, gen_run=run, gen_ep=ck.get("epoch"),
                gen_params=sum(p.numel() for p in m.parameters()) / 1e6)


def status_md():
    if "rest_run" not in _M:
        return "⏳ **Models load on first run** (SAME + restorer + generator) β€” first action takes ~15 s, then it's fast."
    rb = f"**Restorer:** `{_M.get('rest_run','-')}` (ep {_M.get('rest_ep','?')}, {_M.get('rest_params',0):.1f}M)"
    g = _M.get("gen_run", "(none)")
    gt = (_M.get("gargs") or {}).get("target_stems", "β€”")
    gb = f"**Generator:** `{g}`" + (f" (ep {_M.get('gen_ep','?')}, {_M.get('gen_params',0):.1f}M, targets={gt})" if _M.get("gen") is not None else " β€” *disabled*")
    return f"πŸŽ›οΈ Active models Β· {rb} Β· {gb} Β· **Autoencoder:** `SAME-L` Β· device `{DEVICE}`"


def set_models(rest_run, gen_run):
    _M.update(load_restorer(rest_run)); _M.update(load_generator(gen_run))
    print(f"[app] active: restorer={rest_run} generator={gen_run}")
    return status_md()


def models():
    if _M:
        return _M
    cfg = Cfg(); cfg.device = DEVICE
    if DEVICE == "cpu":
        _disable_flash_attn()
    sa, sr = E.load_same(cfg)                     # SAME-L autoencoder (frozen)
    _M.update(sa=sa, sr=sr, cfg=cfg)
    rests = list_runs("restorer"); gens = list_runs("generator")
    default_rest, default_gen = _default_runs(rests, gens)
    _M.update(load_restorer(default_rest)); _M.update(load_generator(default_gen))
    print(f"[app] models loaded on {DEVICE}. restorer={default_rest} generator={default_gen}")
    return _M


# ---------------- core latent ops ----------------
def _fitT(x): return G._fitT(x.float(), T)

def rms_peak_db(audio_mono):
    rms = float(np.sqrt(np.mean(audio_mono ** 2)) + 1e-12)
    peak = float(np.max(np.abs(audio_mono)) + 1e-12)
    return 20 * np.log10(rms), 20 * np.log10(peak)

@torch.inference_mode()
def restore(latent, stem, mix):
    m = models(); st = m["rstats"]; mu, sd = st[stem]["mu"], st[stem]["sd"]
    sd_d = ((latent - mu[:, None]) / sd[:, None]).to(DEVICE)[None]
    sd_m = None                                              # no-mix restorers lack __mix__ stats
    if m.get("rest_use_mix", True) and "__mix__" in st:
        mm = st["__mix__"]
        sd_m = ((mix - mm["mu"][:, None]) / mm["sd"][:, None]).to(DEVICE)[None]
    sid = torch.tensor([STEM_ID[stem]], device=DEVICE)
    if m.get("rest_kind") == "flowattn":                     # GENERATIVE restorer: integrate the deg-anchored flow
        out = F.sample_mix(m["rest"], sd_d, sd_m, sid, m.get("rest_steps", 40), m.get("rest_sigma", 0.0))
    else:
        out = m["rest"](sd_d, sid, sd_m)
    rflow = m.get("rflow")                                   # glue stage: backbone + sampled residual on restore stems
    if rflow is not None and STEM_ID[stem] in m.get("rflow_ids", set()):
        out = out + G.sample_gen(rflow, out, sid, m.get("rflow_steps", 2), 1.0)
    out = out[0].cpu()
    return out * sd[:, None] + mu[:, None]

@torch.inference_mode()
def generate(context_latents, stem):
    m = models()
    if m["gen"] is None:
        return None
    gs = m["gstats"]; cm, cs = gs["ctx"][stem]; tm, ts = gs["tgt"][stem]
    ctx_raw = sum(_fitT(c) for c in context_latents)
    ctx = ((ctx_raw - cm[:, None]) / cs[:, None]).to(DEVICE)[None]
    sid = torch.tensor([STEM_ID[stem]], device=DEVICE)
    w = m["gargs"].get("cfg_w", 2.0); steps = m["gargs"].get("sample_steps", 40)
    z = G.sample_gen(m["gen"], ctx, sid, steps, w)[0].cpu()
    return z * ts[:, None] + tm[:, None]

# STOPGAP loudness control (proper fix = the planned generator overhaul). Generated bass decodes
# ~2x too loud (off-manifold flow endpoints); neither more data nor attention fixed it. We ground
# its level in the context: target rms = ratio x restored-context rms, applied in AUDIO space
# (SAME decode is nonlinear). Ratio = geometric mean of clean_bass/ctx over GENERATION-case val
# pairs (0.24); the per-song spread is wide (p25 .12, p75 .81), so we deliberately err quiet β€”
# too-loud bass is far more jarring than too-quiet.
GEN_ENERGY_RATIO = {"bass": 0.24}

def energy_match(audio_by_stem: dict, dec: dict):
    """Scale each GENERATED stem to its fitted energy relative to the restored context."""
    ctx = [audio_by_stem[s] for s, v in dec.items() if v == "restored" and s in audio_by_stem]
    if not ctx:
        return audio_by_stem
    ctx_rms = float(np.sqrt(np.mean(np.square(np.sum(ctx, axis=0)))) + 1e-9)
    for s, v in dec.items():
        if v == "generated" and s in audio_by_stem and s in GEN_ENERGY_RATIO:
            a = audio_by_stem[s]
            r = float(np.sqrt(np.mean(np.square(a))) + 1e-9)
            audio_by_stem[s] = a * (GEN_ENERGY_RATIO[s] * ctx_rms / r)
    return audio_by_stem


@torch.inference_mode()
def decode(latent):
    return E._decode(models()["sa"], _fitT(latent))         # [2,S] cpu

@torch.inference_mode()
def encode(audio_2S):
    return encode_batch([audio_2S])[0]


@torch.inference_mode()
def encode_batch(auds):
    """Encode many [2,S] clips in ONE SAME forward pass (CPU: far less overhead than N calls)."""
    m = models(); xs = []
    for a in auds:
        x = torch.as_tensor(a).float()
        if x.shape[0] != 2: x = (x[:2] if x.shape[0] > 2 else x.repeat(2, 1))
        xs.append(x)
    L = max(x.shape[1] for x in xs)
    X = torch.stack([torch.nn.functional.pad(x, (0, L - x.shape[1])) for x in xs]).to(DEVICE)
    Z = m["sa"].encode_audio(X, chunked=False)
    return [Z[i].float().cpu() for i in range(Z.shape[0])]


@torch.inference_mode()
def decode_batch(lats):
    """Decode many [256,T] latents in ONE SAME forward pass."""
    if not lats:
        return []
    m = models()
    Z = torch.stack([_fitT(l) for l in lats]).to(DEVICE).float()
    A = m["sa"].decode_audio(Z)
    return [A[i].clamp(-1, 1).cpu() for i in range(A.shape[0])]


def process_chunk(deg_latents: dict, mix_latent, present: dict):
    """deg_latents/present keyed by stem. Restore present, generate missing. Returns
    (final_latents dict, decisions dict)."""
    final, dec = {}, {}
    for s in STEMS:
        if present.get(s) and s in deg_latents:
            final[s] = restore(deg_latents[s], s, mix_latent); dec[s] = "restored"
    ctx = [final[o] for o in final]                          # restored context
    for s in STEMS:
        if s in final:
            continue
        if s in GEN_TARGETS:                                 # always invent missing bass
            g = generate(ctx, s) if ctx else None
            if g is not None: final[s] = g; dec[s] = "generated"
            else: dec[s] = "absent (no generator/context)"
        else:                                                # drums/other: don't fabricate
            dec[s] = "absent (kept out)"
    return final, dec


# ---------------- viz / eval ----------------
def spec_fig(audio_2S, title):
    y = np.asarray(audio_2S).mean(0) if np.asarray(audio_2S).ndim == 2 else np.asarray(audio_2S)
    S = librosa.amplitude_to_db(np.abs(librosa.stft(y, n_fft=1024, hop_length=256)) + 1e-6, ref=np.max)
    fig, ax = plt.subplots(figsize=(4.2, 2.4))
    librosa.display.specshow(S, sr=SR, hop_length=256, x_axis="time", y_axis="log", ax=ax, cmap="magma")
    ax.set_title(title, fontsize=9); ax.tick_params(labelsize=7)
    fig.tight_layout()
    return fig

def multi_stft_np(a, b):
    return E.multi_stft(torch.as_tensor(a).float(), torch.as_tensor(b).float())


def load_audio_any(path):
    """Read wav/flac via soundfile; fall back to librosa+ffmpeg for mp3/m4a/ogg/etc.
    Returns ([C, S] float32, sr)."""
    try:
        a, sr = sf.read(path, dtype="float32")
        return (a.T if a.ndim == 2 else a[None]), sr
    except Exception:                                   # mp3 & friends -> ffmpeg/audioread
        y, sr = librosa.load(path, sr=None, mono=False)
        y = np.asarray(y, dtype="float32")
        return (y if y.ndim == 2 else y[None]), sr


# ---------------- val index ----------------
def build_val_index():
    rows = {}
    meta = VAL_CACHE / "metadata.csv"
    for r in csv.DictReader(open(meta)):
        rows.setdefault(r["pair_id"], {})[r["stem"]] = r
    out = []
    for pid, st in rows.items():
        if not (SRC / "val/degraded" / f"{pid}.wav").exists():
            continue
        present = [s for s in STEMS if st.get(s, {}).get("deg_present") == "1"]
        # generation showcase = a MISSING bass that exists in clean (policy: always invent bass)
        missing = [s for s in GEN_TARGETS if s not in present and st.get(s, {}).get("clean_present") == "1"]
        if present and missing:                              # restore present + generate bass
            out.append((pid, present, missing))
    out.sort(key=lambda x: (-len(x[2]), x[0]))
    return out[:N_VAL]

_VAL = None
def val_list():
    global _VAL
    if _VAL is None: _VAL = build_val_index()
    return _VAL

def val_label(item):
    pid, pres, miss = item
    return f"{pid}  | restore: {','.join(pres) or '-'}  | generate: {','.join(miss) or '-'}"


# ---------------- pipelines ----------------
def _empty_val(msg):
    main = [None, None, None, None, None, None, msg]
    stems = []
    for s in STEMS:
        stems += [f"**{STEM_EMOJI[s]} {s}**", None, None]
    return main + stems


def run_val(label):
    if not label:
        return _empty_val("Pick a sample and press Run.")
    pid = label.split()[0]
    deg_a, _ = sf.read(SRC / "val/degraded" / f"{pid}.wav", dtype="float32")
    deg_a = deg_a.T if deg_a.ndim == 2 else deg_a[None]
    clean_a, _ = sf.read(SRC / "val/clean" / f"{pid}.wav", dtype="float32")
    clean_a = clean_a.T if clean_a.ndim == 2 else clean_a[None]
    latdir = VAL_CACHE / pid / "latents"
    deg_latents, present = {}, {}
    for s in STEMS:
        p = latdir / "degraded" / f"{s}.pt"
        if p.exists():
            deg_latents[s] = torch.load(p, map_location="cpu").float(); present[s] = True
    mix = torch.load(latdir / "degraded_mix.pt", map_location="cpu").float()
    final, dec = process_chunk(deg_latents, mix, present)
    fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys])   # one batched decode
    after = {fkeys[i]: fdec[i].numpy() for i in range(len(fkeys))}
    energy_match(after, dec)                                               # level generated bass
    bkeys = list(deg_latents); bdec = decode_batch([deg_latents[s] for s in bkeys])
    before = {bkeys[i]: bdec[i].numpy() for i in range(len(bkeys))}
    out = sum(after.values())                                             # sum after leveling
    n = min(out.shape[1], clean_a.shape[1], deg_a.shape[1])
    e_in = multi_stft_np(deg_a[:, :n].mean(0), clean_a[:, :n].mean(0))
    e_out = multi_stft_np(out[:, :n].mean(0), clean_a[:, :n].mean(0))
    n_rest = sum(v == "restored" for v in dec.values()); n_gen = sum(v == "generated" for v in dec.values())
    report = (f"**`{pid}`**  Β·  πŸ”§ restored {n_rest}  Β·  ✨ generated {n_gen}  |  "
              f"full-mix vs clean (lower better): degraded **{e_in:.3f}** β†’ output **{e_out:.3f}** "
              f"{'βœ…' if e_out < e_in else '⚠️'}")
    main = [(SR, deg_a.T), (SR, out.T), (SR, clean_a.T),
            spec_fig(deg_a, "degraded input"), spec_fig(out, "output"), spec_fig(clean_a, "clean reference"),
            report]
    stems = []
    for s in STEMS:
        tag = {"restored": "πŸ”§ restored", "generated": "✨ generated"}.get(dec.get(s), "β€”")
        bef = (SR, before[s].T) if s in before else None
        aft = (SR, after[s].T) if s in after else None
        stems += [f"**{STEM_EMOJI[s]} {s}** β€” {tag}", bef, aft]
    return main + stems


def _empty_upload(msg):
    main = [None, None, None, None, msg]
    stems = []
    for s in STEMS:
        stems += [f"**{STEM_EMOJI[s]} {s}**", None]
    return main + stems


def _rms(x):
    return float(np.sqrt(np.mean(x.astype(np.float64) ** 2)) + 1e-12)


def run_upload(filepath, start_s=0.0, end_s=0.0):
    """Wrapper: surface the REAL error in the UI (and Space logs) instead of a bare Gradio 'Error'."""
    try:
        return _run_upload(filepath, start_s, end_s)
    except Exception as e:
        import traceback; traceback.print_exc()                       # -> Space Logs tab
        return _empty_upload(f"⚠️ upload failed β€” {type(e).__name__}: {str(e)[:400]}")


def _run_upload(filepath, start_s=0.0, end_s=0.0):
    if not filepath:
        return _empty_upload("Upload a wav/mp3 and press Run.")
    a, sr = load_audio_any(filepath)                     # wav/flac/mp3/m4a...
    if a.shape[0] == 1: a = np.repeat(a, 2, 0)
    if sr != SR:
        a = np.stack([librosa.resample(a[c], orig_sr=sr, target_sr=SR) for c in range(a.shape[0])])
    # explicit trim (seconds): actually chop the file before processing
    dur = a.shape[1] / SR
    s = max(0.0, float(start_s or 0.0))
    e = float(end_s) if (end_s and float(end_s) > 0) else dur
    e = min(e, dur)
    if e <= s:
        s, e = 0.0, dur
    a = a[:, int(s * SR):int(e * SR)]
    trim_note = f"trim {s:.1f}–{e:.1f}s of {dur:.1f}s Β· "
    chunk = SR * 3; fade = SR // 2; hop = chunk - fade; total = a.shape[1]   # 3s windows, 0.5s crossfade
    try:
        sep = get_separator()
    except RuntimeError as e:                      # separation unavailable on this deployment
        return _empty_upload(f"⚠️ {e}")

    # ---- separate the WHOLE clip ONCE (the CPU bottleneck) β€” demucs/audio-sep split internally β€” then
    #      process OVERLAPPING 3s windows by SLICING the pre-separated stems (no per-window re-separation) ----
    full = separate_chunk(sep, a)
    windows = []; summ = {}
    for i in range(0, max(1, total), hop):
        seg = a[:, i:i + chunk]; orig = seg.shape[1]
        if orig <= 0:
            break
        if orig < chunk:
            seg = np.pad(seg, ((0, 0), (0, chunk - orig)))   # pad silence then crop output back
        stems = {}
        for s, fs in full.items():
            w = fs[:, i:i + orig]
            if w.shape[1] < chunk:
                w = np.pad(w, ((0, 0), (0, chunk - w.shape[1])))
            stems[s] = w
        present_stems = [s for s, au in stems.items()
                         if (lambda rp: rp[0] > PRESENT_RMS or rp[1] > PRESENT_PEAK)(rms_peak_db(au.mean(0)))]
        enc = encode_batch([stems[s] for s in present_stems] + [seg])     # one batched encode (stems + mix)
        deg_latents = {present_stems[j]: enc[j] for j in range(len(present_stems))}
        present = {s: True for s in present_stems}; mix_lat = enc[-1]
        final, dec = process_chunk(deg_latents, mix_lat, present)
        for s, v in dec.items(): summ[s] = v
        fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys])  # one batched decode
        chunk_audio = {s: fdec[j].numpy() for j, s in enumerate(fkeys)}       # SAME decode = fixed length/window
        energy_match(chunk_audio, dec)                                       # intra-chunk: level generated bass
        windows.append({"start": i, "orig": orig, "stems": chunk_audio})
        if i + chunk >= total:
            break

    # ---- GLUE 1 (profile): gently match each window's per-stem level to the song-wide MEDIAN, so chunks
    #      share a consistent profile instead of drifting. alpha<1 preserves real dynamics; clip = no artifacts.
    ALPHA, LO, HI = 0.5, 0.7, 1.4
    for s in STEMS:
        rmss = [_rms(w["stems"][s]) for w in windows if s in w["stems"] and _rms(w["stems"][s]) > 1e-5]
        if len(rmss) >= 2:
            tgt = float(np.median(rmss))
            for w in windows:
                if s in w["stems"] and _rms(w["stems"][s]) > 1e-5:
                    w["stems"][s] = w["stems"][s] * float(np.clip((tgt / _rms(w["stems"][s])) ** ALPHA, LO, HI))

    # ---- GLUE 2 (seams): crossfaded overlap-add, done in OUTPUT-sample space (SAME decodes each 3s window
    #      to a fixed length Ld, not the input `orig`). out_hop/out_fade scale the input hop/fade by Ld/chunk.
    #      One global weight timeline (incl. windows where a stem is absent) so a stem fades across a boundary
    #      into a neighbour that lacks it. ----
    n = len(windows)
    Ld = next((au.shape[1] for w in windows for au in w["stems"].values()), 0)
    if Ld == 0:
        return _empty_upload(f"{trim_note}no restorable content found.")
    out_fade = min(int(round(fade * Ld / chunk)), Ld // 2)
    out_hop = max(1, Ld - out_fade)
    def _Lout(w):                                       # real output length (last/short window cropped)
        return Ld if w["orig"] >= chunk else min(Ld, int(round(w["orig"] * Ld / chunk)))
    starts = [k * out_hop for k in range(n)]
    total_out = starts[-1] + _Lout(windows[-1])
    Wg = np.zeros(total_out); envs = []
    for k, w in enumerate(windows):
        L = _Lout(w); env = np.ones(L)
        if k > 0 and out_fade > 0: env[:min(out_fade, L)] = np.linspace(0, 1, out_fade, endpoint=False)[:min(out_fade, L)]
        if k < n - 1 and L >= out_fade and out_fade > 0: env[-out_fade:] = np.linspace(1, 0, out_fade, endpoint=False)
        envs.append(env); Wg[starts[k]:starts[k] + L] += env
    Wg = np.maximum(Wg, 1e-6)
    per_stem_audio = {}
    for s in STEMS:
        if not any(s in w["stems"] for w in windows):
            continue
        buf = np.zeros((2, total_out))
        for k, w in enumerate(windows):
            if s in w["stems"]:
                L = _Lout(w); buf[:, starts[k]:starts[k] + L] += w["stems"][s][:, :L] * envs[k]
        per_stem_audio[s] = (buf / Wg).astype(np.float32)
    out = sum(per_stem_audio.values()) if per_stem_audio else np.zeros((2, total_out), np.float32)

    rep = (f"{trim_note}**{n}** overlapping 3 s windows Β· hop {hop/SR:.1f}s, {fade/SR:.1f}s crossfade Β· "
           "glued to song profile.  " + " Β· ".join(f"{STEM_EMOJI[st]} {summ.get(st, '-')}" for st in STEMS))
    main = [(SR, a.T), (SR, out.T), spec_fig(a, "input"), spec_fig(out, "output"), rep]
    stems_out = []
    for s in STEMS:
        tag = {"restored": "πŸ”§ restored", "generated": "✨ generated"}.get(summ.get(s), "β€”")
        au = (SR, per_stem_audio[s].T) if s in per_stem_audio else None
        stems_out += [f"**{STEM_EMOJI[s]} {s}** β€” {tag}", au]
    return main + stems_out


# ---------------- separator (uploads only) ----------------
# EXACT-PARITY: the train/val cache was built with htdemucs v4 (audio-separator's htdemucs.yaml,
# shifts=1 overlap=0.25 β€” see demucs.py). Live upload MUST use those SAME v4 weights or the restorer
# sees a different stem distribution. Priority: (1) audio-separator (exactly how the cache was built β€”
# present locally); (2) the official `demucs` package (SAME htdemucs v4 weights, far lighter deps so it
# installs on the CPU Space without the numpy/pywt clash that disabled this before). We deliberately do
# NOT fall back to torchaudio's HDemucs β€” that's v3 (no transformer), i.e. NOT identical stems.
_SEP = {}
def get_separator():
    if _SEP: return _SEP
    try:                                          # (1) audio-separator β€” exact, as used to build the cache
        from audio_separator.separator import Separator
        s = Separator(output_dir="/tmp/app_sep", model_file_dir="/tmp/audio-separator-models")
        s.load_model("htdemucs.yaml")
        _SEP.update(kind="audio_separator", sep=s)
        return _SEP
    except Exception:
        pass
    try:                                          # (2) demucs package β€” SAME htdemucs v4 weights
        from demucs.pretrained import get_model    # NB: on the Space there is no root demucs.py to shadow this
        from demucs.apply import apply_model
        m = get_model("htdemucs").to(DEVICE).eval()
        _SEP.update(kind="demucs", model=m, apply=apply_model, sources=list(m.sources))
        return _SEP
    except Exception as e:
        raise RuntimeError(
            "Live upload separation is unavailable: htdemucs v4 (audio-separator or the demucs package) "
            "isn't installed here. Use the 🎧 Val samples tab β€” full restore+generate pipeline. "
            "(torchaudio's HDemucs is v3 and would NOT match the trained stems.)"
        ) from e

# ---- separation cache (OPT-IN: set RESTOFLOW_SEP_CACHE=<dir>) ----
# Separation is identical across restorer/generator VARIANTS (same degraded input -> same htdemucs stems),
# so an experiment that scores N variants re-separates the SAME clips N times. This content-addressed cache
# stores each unique input's stems ONCE (deduped by input-hash) and reuses them. Stored as int16 = byte-exact
# to what audio_separator already writes (16-bit PCM) -> NO added error vs the live path, half the disk of
# float32. UNSET env -> behavior is byte-identical to before (the UI/Space never sets it).
_SEP_CACHE_VER = "htdemucs_v4_a"          # bump to invalidate all cached stems


def _sep_cache_dir():
    d = os.environ.get("RESTOFLOW_SEP_CACHE")
    return Path(d) if d else None


def _sep_cache_key(seg_2S, kind):
    a = np.ascontiguousarray(np.asarray(seg_2S, dtype=np.float32))
    h = hashlib.sha1(a.tobytes())
    h.update(f"|{kind}|{a.shape}|{_SEP_CACHE_VER}".encode())
    return h.hexdigest()


def _sep_cache_load(cdir, key):
    # 4 stems stored as one 8-channel PCM16 FLAC (lossless, ~40% smaller than zlib-int16). Channel order =
    # sorted(STEMS); each stem = 2 consecutive channels. PCM16/float read = exactly what the live wav path
    # returns (k/32768), so a hit is byte-identical to live separation.
    f = cdir / f"{key}.flac"
    if not f.exists():
        return None
    try:
        inter, _ = sf.read(f, dtype="float32")                             # [L, 8] in [-1,1)
        order = sorted(STEMS)
        return {order[i]: np.ascontiguousarray(inter[:, 2 * i:2 * i + 2].T) for i in range(len(order))}
    except Exception:
        return None                                                        # corrupt -> recompute live


def _sep_cache_save(cdir, key, stems):
    try:
        cdir.mkdir(parents=True, exist_ok=True)
        order = sorted(STEMS)
        L = next((np.asarray(v).shape[-1] for v in stems.values()), 0)
        chans = []
        for s in order:                                                    # canonical order; missing -> silence
            v = stems.get(s)
            a = np.asarray(v, np.float32) if v is not None else np.zeros((2, L), np.float32)
            if a.shape[0] != 2:
                a = a[:2] if a.shape[0] > 2 else np.repeat(a, 2, 0)
            chans.append(np.clip(np.round(a * 32768.0), -32768, 32767).astype(np.int16))
        inter = np.concatenate(chans, axis=0).T                            # [L, 8] int16
        tmp = cdir / f".{key}.{os.getpid()}.tmp.flac"                       # atomic write (no torn files)
        sf.write(tmp, inter, SR, format="FLAC", subtype="PCM_16"); os.replace(tmp, cdir / f"{key}.flac")
    except Exception:
        pass                                                               # cache is best-effort, never fatal


def separate_chunk(sep, seg_2S):
    """[2,S] @ SR -> {stem:[2,S]} via htdemucs v4. Cache-aware (see _sep_cache_*): on a hit the stems are
    loaded from disk instead of re-separated, so multi-variant experiment evals separate each clip once."""
    cdir = _sep_cache_dir()
    key = _sep_cache_key(seg_2S, sep["kind"]) if cdir is not None else None
    if cdir is not None:
        hit = _sep_cache_load(cdir, key)
        if hit is not None:
            return hit
    out = _separate_chunk_impl(sep, seg_2S)
    if cdir is not None:
        _sep_cache_save(cdir, key, out)
    return out


def _separate_chunk_impl(sep, seg_2S):
    """[2,S] @ SR -> {stem:[2,S]} via the same htdemucs v4 weights the cache used. CPU-tuned: shifts=0
    (no test-time augmentation β€” the big needless cost), light overlap, no_grad (demucs does in-place ops
    that throw under inference_mode). Called ONCE on the whole clip (demucs splits internally) β€” NOT per
    overlapping window β€” so live separation stays ~1 forward instead of N."""
    if sep["kind"] == "audio_separator":
        # per-process input path: distinct output filenames -> no race when another process (e.g. a precache
        # pass on another GPU) separates concurrently into the shared /tmp/app_sep dir.
        tmp = f"/tmp/app_sep_in_{os.getpid()}.wav"; sf.write(tmp, seg_2S.T, SR)
        files = sep["sep"].separate(tmp)
        out = {}
        for f in files:
            fl = f.lower()
            for s in STEMS:
                if s in fl:
                    au, _ = sf.read(os.path.join("/tmp/app_sep", f) if not os.path.isabs(f) else f, dtype="float32")
                    out[s] = (au.T if au.ndim == 2 else np.repeat(au[None], 2, 0))
        return out
    wav = torch.as_tensor(seg_2S, dtype=torch.float32, device=DEVICE)        # [2,S]
    ref = wav.mean(0); mu, sd = ref.mean(), ref.std().clamp_min(1e-8)
    with torch.no_grad():
        src = sep["apply"](sep["model"], ((wav - mu) / sd)[None],
                           shifts=0, overlap=0.1, split=True, device=DEVICE)[0]   # [4,2,S]
    src = (src * sd + mu).cpu().numpy()
    return {s: src[i] for i, s in enumerate(sep["sources"]) if s in STEMS}


EXPLAIN = """
## How this works (high level)

**The problem.** Old/lo-fi recordings are *degraded* (muffled, missing instruments). We work in
**SAME-L's latent space** β€” a pretrained **neural audio autoencoder** (a learned codec, Γ  la the
VAEs behind latent diffusion) that compresses 3 s of stereo into a small `256Γ—32` "summary"
(~11 frames/sec). Everything below operates on these latents, then decodes back to audio β€” so the
models stay tiny and fast.

**The models, in plain terms:**
- **SAME-L** β€” the compact "latent" space everything runs in (a frozen neural audio codec: encode ↔ decode).
- **htdemucs** β€” splits a song into its 4 instrument stems.
- **Restorer** β€” *cleans up* a muffled/damaged stem so it sounds clear again. It fixes what's there; it doesn't add new parts. The default ("glue") restorer is a deterministic Conformer-style backbone **plus a small residual-flow detail head** that samples back fine high-frequency detail; the plain backbone alone is selectable as the **baseline** (*advramp*).
- **Generator** β€” *invents* a missing instrument (mainly bass) that fits the song, like an AI session musician. (Conditional flow-matching.)
- **Glue stage (training-time)** β€” the restorer + generator are co-trained so their *summed* mix matches a real clean mix, judged by a **multi-scale latent discriminator** with feature-matching. This is a coherence objective on the recombined mix; in practice it lands close to the deterministic baseline (the two are A/B-selectable so you can compare).

**Step 1 β€” Separate.** htdemucs splits the input mix into 4 stems: *other, vocals, drums, bass*.

**Step 2 β€” Route.** For each stem we ask: *is it actually there?* (energy gate).
- **Present** β†’ it just needs cleanup β†’ **RESTORE**.
- **Missing bass** β†’ an absent bass is **invented** β†’ **GENERATE**. (Note: this fills a low end even for songs that may not originally have a bass instrument β€” a known limitation on bass-free material.)
- **Missing drums/other** β†’ a song may legitimately have none, so we **do not fabricate** them β€” restore-only.

**Step 3a β€” Restorer.** For present stems, a small **Conformer-style network** (local convolutions
for texture **+ global self-attention** so each moment sees the whole clip) nudges the degraded
latent toward a clean one, as a *residual* on the input β€” mostly *deterministic* because the degraded
signal already tells us the answer. The default model adds a **residual-flow head** that *samples* a
small extra detail residual on top (drums/other/vocals). Trained with a **stem-balanced** loss so
quiet stems (bass) aren't drowned out by loud ones.

**Step 3b β€” Generator (creative, bass).** When the bass is missing there is no answer to copy β€”
many different bass lines could fit. So this is a **conditional flow-matching** generator: starting
from noise it follows a learned velocity field to *sample* a plausible **bass conditioned on the
other (restored) stems**, so it locks to the song's harmony. **Classifier-free guidance** lets us
dial how strongly it commits to the accompaniment. (Probes showed bass is strongly determined by the
accompaniment and the codec preserves it well; drums are timing-limited by the ~11 Hz latent, so we
don't fabricate drums.)

**Step 4 β€” Decode & remix.** Each final latent is decoded to audio and summed into the output mix,
with per-window level smoothing and crossfaded overlap-add across the 3 s windows.

> Restoration leans on the *existing* signal (works best when degradation is mild). Generation
> *creates* absent instruments to accompany what's there. Inference only ever sees restored stems β€”
> never the clean originals. The **baseline** (deterministic *advramp* restorer + *gen_advramp_v1*) is
> selectable for A/B against the default glue models.
"""


THEME = gr.themes.Base(
    primary_hue=gr.themes.colors.orange,
    secondary_hue=gr.themes.colors.orange,
    neutral_hue=gr.themes.colors.zinc,
    font=[gr.themes.GoogleFont("Inter"), "system-ui", "sans-serif"],
    radius_size=gr.themes.sizes.radius_sm,
).set(
    # dark look as the default (set the light variants to dark so it's black/orange regardless of OS)
    body_background_fill="#0a0a0b",
    body_text_color="#e9e9ec",
    body_text_color_subdued="#8a8a93",
    background_fill_primary="#151518",
    background_fill_secondary="#1c1c20",
    block_background_fill="#151518",
    block_border_color="#272730",
    block_label_background_fill="#151518",
    block_label_text_color="#ff8a3d",
    block_title_text_color="#f2f2f4",
    border_color_primary="#272730",
    input_background_fill="#1c1c20",
    button_primary_background_fill="#ff7a18",
    button_primary_background_fill_hover="#ff9344",
    button_primary_text_color="#0a0a0b",
    button_secondary_background_fill="#272730",
    button_secondary_text_color="#e9e9ec",
    color_accent_soft="#2a1a0e",
    slider_color="#ff7a18",
)

CSS = """
.gradio-container {max-width: 1080px !important; margin: auto !important; background:#0a0a0b;}
.fullmix audio {width: 100%;}
h1,h2,h3,h4 {color:#f4f4f6 !important; letter-spacing:.2px;}
h1 {font-weight:700; border-left:4px solid #ff7a18; padding-left:12px;}
h4 {color:#ff9a55 !important; text-transform:uppercase; font-size:.8rem; letter-spacing:.6px;}
a {color:#ff8a3d !important;}
footer {display:none !important;}
::-webkit-scrollbar{width:9px;height:9px} ::-webkit-scrollbar-thumb{background:#3a3a44;border-radius:6px}
"""


def _mix_column(title, with_plot=True):
    gr.Markdown(f"#### {title}")
    a = gr.Audio(label=None, show_label=False, elem_classes="fullmix")
    p = gr.Plot(show_label=False) if with_plot else None
    return a, p


def build_ui():
    # NOTE: do NOT load models here β€” that delays the port bind ~15 s and makes the
    # auto-opened browser tab hit a dead port. UI uses cheap dir listing; models load
    # lazily on first action (warmed in a background thread at launch).
    rests = list_runs("restorer"); gens = list_runs("generator")
    dr, dg = _default_runs(rests, gens)
    with gr.Blocks(title="SAME Restore + Generate", theme=THEME, css=CSS) as demo:
        gr.Markdown("# πŸŽ›οΈ Stem Restoration + Generation\n"
                    "Upload a muffled/lo-fi song. We split it into instruments, **clean up** the ones that are "
                    "there, **invent** a missing bass that fits, and remix β€” all in a compact neural-audio space. "
                    "A GAN 'judge' keeps the result sounding *real*, not dull.")
        status = gr.Markdown(status_md())
        with gr.Accordion("βš™οΈ Choose models", open=False):
            with gr.Row():
                rest_dd = gr.Dropdown(rests, value=dr, label="Restorer", scale=2)
                gen_dd = gr.Dropdown(gens + ["(none)"], value=dg, label="Generator (bass)", scale=2)
                load_btn = gr.Button("↻ Load", variant="primary", scale=1)
            load_btn.click(set_models, [rest_dd, gen_dd], status)
        with gr.Tab("🎧 Val samples"):
            with gr.Row():
                dd = gr.Dropdown([val_label(x) for x in val_list()], label="3 s val sample (each shows both restore + generate)", scale=5)
                btn = gr.Button("β–Ά Run", variant="primary", scale=1)
            rep = gr.Markdown()
            with gr.Row(equal_height=True):
                with gr.Column(): a_in, s_in = _mix_column("🎚️ Degraded input")
                with gr.Column(): a_out, s_out = _mix_column("✨ Output")
                with gr.Column(): a_ref, s_ref = _mix_column("🎯 Clean reference")
            with gr.Accordion("πŸ”¬ Per-stem detail (before β†’ after)", open=False):
                comps = []
                for s in STEMS:
                    with gr.Row(equal_height=True):
                        lbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**")
                        bef = gr.Audio(label="before (degraded)", show_label=True, scale=2)
                        aft = gr.Audio(label="after (restored/generated)", show_label=True, scale=2)
                    comps += [lbl, bef, aft]
            btn.click(run_val, dd, [a_in, a_out, a_ref, s_in, s_out, s_ref, rep] + comps)
        with gr.Tab("⬆️ Upload"):
            with gr.Row():
                up = gr.Audio(label="Upload wav/mp3 (chunked into 3 s; last chunk padded then cropped)",
                              type="filepath", sources=["upload"], editable=True,
                              waveform_options=gr.WaveformOptions(
                                  waveform_color="#5a5a66", waveform_progress_color="#ff7a18",
                                  trim_region_color="#ff7a18"),
                              scale=5)
                ub = gr.Button("β–Ά Run", variant="primary", scale=1)
            with gr.Row():
                trim_start = gr.Number(value=0, label="Trim start (s)", scale=1)
                trim_end = gr.Number(value=0, label="Trim end (s, 0 = to end)", scale=1)
                gr.Markdown("*Set start/end to actually chop the file before processing "
                            "(leave both 0 to process the whole upload).*")
            urep = gr.Markdown()
            with gr.Row(equal_height=True):
                with gr.Column(): ua_in, us_in = _mix_column("⬆️ Input")
                with gr.Column(): ua_out, us_out = _mix_column("✨ Output")
            with gr.Accordion("πŸ”¬ Per-stem detail (output)", open=False):
                ucomps = []
                for s in STEMS:
                    with gr.Row(equal_height=True):
                        ulbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**")
                        uaft = gr.Audio(label="output stem", show_label=True, scale=3)
                    ucomps += [ulbl, uaft]
            ub.click(run_upload, [up, trim_start, trim_end], [ua_in, ua_out, us_in, us_out, urep] + ucomps)
        with gr.Accordion("ℹ️ How this works (restorer + generator, intuitively)", open=False):
            gr.Markdown(EXPLAIN)
    return demo


def _free_port(port):
    """Kill any process still holding the port so re-launch binds cleanly."""
    try:
        out = os.popen(f"fuser {port}/tcp 2>/dev/null").read().split()
        for pid in out:
            if pid.strip().isdigit() and int(pid) != os.getpid():
                os.system(f"kill -9 {pid} 2>/dev/null")
    except Exception:
        pass


if __name__ == "__main__":
    port = int(os.environ.get("PORT", 7860))
    _free_port(port)                                  # clear a stale bind from a previous run
    demo = build_ui()
    import threading
    threading.Thread(target=models, daemon=True).start()   # warm models in bg; port still binds instantly
    try:
        demo.launch(server_name="0.0.0.0", server_port=port, share=False,
                    ssr_mode=False)   # ssr_mode off: Gradio-6 SSR hangs without a node runtime
    except KeyboardInterrupt:
        print("\n[app] shutting down…")
    finally:                                          # flush the port on exit (Ctrl+C / close)
        try: demo.close()
        except Exception: pass
        try: gr.close_all()
        except Exception: pass
        _free_port(port)
        print("[app] port released.")