File size: 8,168 Bytes
e3cf774
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
537dd1a
 
 
 
 
 
 
 
e3cf774
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13b48d6
e3cf774
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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)