soilkon commited on
Commit
39e8810
·
verified ·
1 Parent(s): 4e73c38

sync app + active models

Browse files
Files changed (1) hide show
  1. restoflow/app.py +15 -6
restoflow/app.py CHANGED
@@ -433,7 +433,9 @@ def _run_upload(filepath, start_s=0.0, end_s=0.0):
433
  except RuntimeError as e: # separation unavailable on this deployment
434
  return _empty_upload(f"⚠️ {e}")
435
 
436
- # ---- pass 1: process OVERLAPPING 3s windows (hop 2.5s) independently ----
 
 
437
  windows = []; summ = {}
438
  for i in range(0, max(1, total), hop):
439
  seg = a[:, i:i + chunk]; orig = seg.shape[1]
@@ -441,7 +443,12 @@ def _run_upload(filepath, start_s=0.0, end_s=0.0):
441
  break
442
  if orig < chunk:
443
  seg = np.pad(seg, ((0, 0), (0, chunk - orig))) # pad silence then crop output back
444
- stems = separate_chunk(sep, seg)
 
 
 
 
 
445
  present_stems = [s for s, au in stems.items()
446
  if (lambda rp: rp[0] > PRESENT_RMS or rp[1] > PRESENT_PEAK)(rms_peak_db(au.mean(0)))]
447
  enc = encode_batch([stems[s] for s in present_stems] + [seg]) # one batched encode (stems + mix)
@@ -542,7 +549,10 @@ def get_separator():
542
  ) from e
543
 
544
  def separate_chunk(sep, seg_2S):
545
- """seg_2S [2,S] @ SR -> {stem: [2,S] float32}, via the same htdemucs v4 weights the cache used."""
 
 
 
546
  if sep["kind"] == "audio_separator":
547
  tmp = "/tmp/app_sep_in.wav"; sf.write(tmp, seg_2S.T, SR)
548
  files = sep["sep"].separate(tmp)
@@ -554,12 +564,11 @@ def separate_chunk(sep, seg_2S):
554
  au, _ = sf.read(os.path.join("/tmp/app_sep", f) if not os.path.isabs(f) else f, dtype="float32")
555
  out[s] = (au.T if au.ndim == 2 else np.repeat(au[None], 2, 0))
556
  return out
557
- # demucs package: per-mix normalize (the standard htdemucs recipe), shifts=1 overlap=0.25 to match.
558
  wav = torch.as_tensor(seg_2S, dtype=torch.float32, device=DEVICE) # [2,S]
559
  ref = wav.mean(0); mu, sd = ref.mean(), ref.std().clamp_min(1e-8)
560
- with torch.inference_mode():
561
  src = sep["apply"](sep["model"], ((wav - mu) / sd)[None],
562
- shifts=1, overlap=0.25, split=True, device=DEVICE)[0] # [4,2,S]
563
  src = (src * sd + mu).cpu().numpy()
564
  return {s: src[i] for i, s in enumerate(sep["sources"]) if s in STEMS}
565
 
 
433
  except RuntimeError as e: # separation unavailable on this deployment
434
  return _empty_upload(f"⚠️ {e}")
435
 
436
+ # ---- separate the WHOLE clip ONCE (the CPU bottleneck) — demucs/audio-sep split internally — then
437
+ # process OVERLAPPING 3s windows by SLICING the pre-separated stems (no per-window re-separation) ----
438
+ full = separate_chunk(sep, a)
439
  windows = []; summ = {}
440
  for i in range(0, max(1, total), hop):
441
  seg = a[:, i:i + chunk]; orig = seg.shape[1]
 
443
  break
444
  if orig < chunk:
445
  seg = np.pad(seg, ((0, 0), (0, chunk - orig))) # pad silence then crop output back
446
+ stems = {}
447
+ for s, fs in full.items():
448
+ w = fs[:, i:i + orig]
449
+ if w.shape[1] < chunk:
450
+ w = np.pad(w, ((0, 0), (0, chunk - w.shape[1])))
451
+ stems[s] = w
452
  present_stems = [s for s, au in stems.items()
453
  if (lambda rp: rp[0] > PRESENT_RMS or rp[1] > PRESENT_PEAK)(rms_peak_db(au.mean(0)))]
454
  enc = encode_batch([stems[s] for s in present_stems] + [seg]) # one batched encode (stems + mix)
 
549
  ) from e
550
 
551
  def separate_chunk(sep, seg_2S):
552
+ """[2,S] @ SR -> {stem:[2,S]} via the same htdemucs v4 weights the cache used. CPU-tuned: shifts=0
553
+ (no test-time augmentation — the big needless cost), light overlap, no_grad (demucs does in-place ops
554
+ that throw under inference_mode). Called ONCE on the whole clip (demucs splits internally) — NOT per
555
+ overlapping window — so live separation stays ~1 forward instead of N."""
556
  if sep["kind"] == "audio_separator":
557
  tmp = "/tmp/app_sep_in.wav"; sf.write(tmp, seg_2S.T, SR)
558
  files = sep["sep"].separate(tmp)
 
564
  au, _ = sf.read(os.path.join("/tmp/app_sep", f) if not os.path.isabs(f) else f, dtype="float32")
565
  out[s] = (au.T if au.ndim == 2 else np.repeat(au[None], 2, 0))
566
  return out
 
567
  wav = torch.as_tensor(seg_2S, dtype=torch.float32, device=DEVICE) # [2,S]
568
  ref = wav.mean(0); mu, sd = ref.mean(), ref.std().clamp_min(1e-8)
569
+ with torch.no_grad():
570
  src = sep["apply"](sep["model"], ((wav - mu) / sd)[None],
571
+ shifts=0, overlap=0.1, split=True, device=DEVICE)[0] # [4,2,S]
572
  src = (src * sd + mu).cpu().numpy()
573
  return {s: src[i] for i, s in enumerate(sep["sources"]) if s in STEMS}
574