StemFX / app.py
Vansh Chugh
cleanup: docstring polish
537dd1a
Raw
History Blame Contribute Delete
8.17 kB
import sys
sys.stdout.reconfigure(line_buffering=True)
try:
import spaces
except ImportError:
# keep @spaces.GPU usable as a no-op; ZeroGPU requires this exact name.
class spaces:
class GPU:
def __init__(self, func=None, duration=60):
self.func = func
def __call__(self, *args, **kwargs):
if self.func is not None:
return self.func(*args, **kwargs)
func = args[0]
return func
import os
import tempfile
import numpy as np
import pyloudnorm as pyln
import soundfile as sf
import torch
import torchaudio
import gradio as gr
from pyharp import ModelCard, build_endpoint
import stemfx
from stemfx.separator import SCNetSeparator
from multiafx import FXChain
SAMPLE_RATE = 44100
SEGMENT_SAMPLES = SAMPLE_RATE * 10 # StemFX's encoder is trained on fixed 10s clips (stemfx.api.SEGMENT_SECONDS)
STEM_NAMES = ("vocals", "bass", "drums", "other")
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model_card = ModelCard(
name="StemFX",
description=(
"Predicts a per-stem effects chain that makes one mix sound like a "
"reference mix, then applies it to the full track. The chain itself "
"is chosen by listening to only the first 10 seconds of each input "
"(a limit of the underlying model, trained on 10-second clips) and "
"then applied uniformly across the whole song -- it won't adapt if "
"the song's character changes partway through."
),
author="Yuan-Chiao Cheng, Jui-Te Wu, Brian Chen, Yen-Tung Yeh, Yu-Hua Chen, Yi-Hsuan Yang",
tags=["audio-effects", "mixing", "style-transfer"],
)
_model = None
_separator = None
def _get_model():
"""Load StemFX on first use, so the CUDA touch (if any) happens inside
the GPU-attached call, not at import time."""
global _model
if _model is None:
_model = stemfx.load(device=DEVICE)
return _model
def _get_separator():
"""Load the SCNet stem separator on first use -- same GPU-safety reasoning as _get_model."""
global _separator
if _separator is None:
_separator = SCNetSeparator(device=DEVICE)
return _separator
def _load_wav(path: str) -> torch.Tensor:
"""Load a wav as a (2, T) float32 tensor at 44.1kHz.
Adapted from stemfx.api._load_wav: uses soundfile rather than
torchaudio.load(), which would pull in torchcodec.
"""
data, sr = sf.read(path, dtype="float32", always_2d=True)
audio = torch.from_numpy(data.T.copy())
if sr != SAMPLE_RATE:
audio = torchaudio.functional.resample(audio, sr, SAMPLE_RATE)
if audio.shape[0] == 1:
audio = audio.repeat(2, 1)
elif audio.shape[0] > 2:
audio = audio[:2]
return audio.float()
def _loudness_normalize(audio: torch.Tensor, target_lufs: float) -> torch.Tensor:
"""Normalize integrated loudness to a target LUFS -- same approach stemfx.api uses internally."""
meter = pyln.Meter(SAMPLE_RATE)
audio_np = audio.cpu().numpy().astype(np.float32)
integrated = meter.integrated_loudness(audio_np.T)
if not np.isfinite(integrated) or integrated < -70:
return audio
out = pyln.normalize.loudness(audio_np.T, integrated, target_lufs).T
return torch.from_numpy(out.astype(np.float32))
def _pretty_chain(chain: dict) -> str:
"""Render a predicted FX chain as one readable line per stem.
Same format as stemfx.api.TransferResult.pretty(), reimplemented here
because we call model.transfer() directly (chain only, no audio) rather
than transfer_audio() -- see process_fn's docstring for why.
"""
def fmt(v):
return f"{v:.3g}" if isinstance(v, float) else str(v)
lines = []
for stem in STEM_NAMES:
steps = chain.get(stem, [])
if not steps:
lines.append(f" {stem}: (no FX)")
continue
chunks = []
for step in steps:
eff = step["effect"]
params = ", ".join(f"{k}={fmt(v)}" for k, v in step.get("params", {}).items())
chunks.append(f"{eff}({params})" if params else eff)
lines.append(f" {stem}: " + " -> ".join(chunks))
return "\n".join(lines)
@spaces.GPU
@torch.inference_mode()
def process_fn(
original_path: str,
reference_path: str,
normalize_loudness: bool,
target_lufs: float,
) -> tuple[str, str]:
"""Restyle the full original mix to sound like the reference mix.
stemfx's own transfer_audio() separates the full track (same cost as
here) but then crops the rendered output down to 10s too, since it
reuses the same tensor for embedding and rendering. Only the embedding
needs the crop -- StemFX's encoder is trained on fixed 10s clips, but
the predicted FX chain is just static params, applicable to any length.
So here we separate once and pass the full-length stems straight to
embed() (which crops its own copy internally) and to FXChain rendering
-- same separation cost as transfer_audio, full-length output.
"""
model = _get_model()
separator = _get_separator()
orig_audio = _load_wav(original_path)
orig_stems = separator.separate(orig_audio) # full length; embed() below crops its own copy internally
ref_audio = _load_wav(reference_path)[:, :SEGMENT_SAMPLES] # only the first 10s of the reference is ever used
ref_stems = separator.separate(ref_audio)
emb_orig = model.embed(orig_stems)
emb_target = model.embed(ref_stems)
chain = model.transfer(emb_orig, emb_target)
processed = {}
for stem in STEM_NAMES:
audio_np = orig_stems[stem].cpu().numpy().astype(np.float32)
steps = chain.get(stem, [])
if steps:
audio_np = FXChain(steps)(audio_np, SAMPLE_RATE)
processed[stem] = torch.from_numpy(audio_np)
if normalize_loudness:
processed = {k: _loudness_normalize(v, target_lufs) for k, v in processed.items()}
mix = sum(processed.values())
peak = mix.abs().max()
if peak > 0.95:
mix = mix * (0.95 / peak)
if normalize_loudness:
mix = _loudness_normalize(mix, target_lufs)
audio_path = tempfile.NamedTemporaryFile(suffix=".wav", delete=False).name
sf.write(audio_path, np.ascontiguousarray(mix.cpu().numpy().T), SAMPLE_RATE)
chain_path = tempfile.NamedTemporaryFile(suffix=".txt", delete=False).name
with open(chain_path, "w") as f:
f.write(
f"{os.path.basename(original_path)}\n\n"
"Predicted FX Chain\n"
f"{_pretty_chain(chain)}\n"
)
return audio_path, chain_path
with gr.Blocks() as demo:
input_components = [
gr.Audio(type="filepath", label="Original Mix")
.harp_required(True)
.set_info("The mix to restyle. Effects are chosen using its first 10 seconds, then applied to the whole track."),
gr.Audio(type="filepath", label="Reference Mix")
.harp_required(True)
.set_info("The mix whose sound/style to copy. Only its first 10 seconds are used."),
gr.Checkbox(
value=True,
label="Normalize Output Loudness",
info="Normalize output to a target loudness (default: True, per repo config)",
),
gr.Slider(
minimum=-36,
maximum=-9,
step=0.5,
value=-23.0,
label="Target Loudness (LUFS)",
info="Loudness target used when normalization is enabled (default: -23.0, per repo config)",
),
]
output_components = [
gr.Audio(type="filepath", label="Processed Mix").set_info(
"Full original mix, re-rendered with the predicted FX chain in the reference's style."
),
gr.File(type="filepath", file_types=[".txt"], label="FX Chain").set_info(
"Human-readable per-stem effects chain predicted by the model."
),
]
build_endpoint(
model_card=model_card,
input_components=input_components,
output_components=output_components,
process_fn=process_fn,
)
if __name__ == "__main__":
demo.queue().launch(pwa=True)