File size: 9,656 Bytes
83e59db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fffb8ef
 
83e59db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f8c94d1
83e59db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a80fd99
83e59db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c0089ac
 
 
 
83e59db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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 sys
sys.stdout.reconfigure(line_buffering=True)

import contextlib
import gc
import os
import tempfile
import threading

import gradio as gr
import numpy as np
import torch
from huggingface_hub import hf_hub_download
from pyharp import ModelCard, build_endpoint

from beat_this.inference import File2Beats
from src.audiocraft.data.audio import audio_write
from src.audiocraft.models import MusicGen
from src.loopgen import LoopGen

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

MUSICGEN_MODEL = "facebook/musicgen-medium"
MAGNET_MODEL = "facebook/magnet-medium-30secs"
MAX_CFG = 5.0          # classifier-free guidance ceiling (default: 5.0, per paper)
MIN_CFG = 1.0          # classifier-free guidance floor, annealed down to this (default: 1.0, per paper)
RESCORE_WEIGHT = 0.5   # MAGNeT/MusicGen rescoring interpolation weight (default: 0.5, per paper)
HINT_DURATION = 30.0   # seconds of MusicGen hint audio generated for beat detection (default: 30.0, per repo)
MAX_RETRIES = 5        # beat-alignment retries before falling back to a fixed length (default: 5, per repo)

beat_checkpoint_path = None
beat_checkpoint_ready = False
beat_checkpoint_error = None

musicgen = None
loopgen = None
file2beats = None
model_lock = threading.Lock()


def download_beat_checkpoint():
    """Fetch the beat-tracking checkpoint in the background so the server can
    start immediately. Reuses the already-deployed teamup-tech/beat-this
    Space's copy instead of re-uploading the same file here."""
    global beat_checkpoint_path, beat_checkpoint_ready, beat_checkpoint_error
    try:
        beat_checkpoint_path = hf_hub_download(
            repo_id="teamup-tech/beat-this", filename="final0.ckpt", repo_type="space"
        )
        print("Beat-tracking checkpoint downloaded.")
    except Exception as e:
        beat_checkpoint_error = str(e)
        print(f"Beat checkpoint download error: {e}")
    finally:
        beat_checkpoint_ready = True


threading.Thread(target=download_beat_checkpoint, daemon=True).start()


model_card = ModelCard(
    name="LoopGen",
    description="Generates a seamlessly loopable music clip from a text description, training-free, using MusicGen + MAGNeT. "
                "Output audio is the loop played twice back-to-back, so you can hear it repeat.",
    author="Davide Marincione, Giorgio Strano, Donato Crisostomi, Roberto Ribuoli, Emanuele Rodolà",
    tags=["music generation", "loop", "text-to-music"],
)


def _write_temp_wav(wav, sample_rate):
    """Write a generated waveform to a fresh temp file and return its path.
    Adapted from pipelines.py, which reused hardcoded local paths like
    f"musicgen_base_{name}_{seed}.wav" across calls (reused via
    os.path.exists) -- only valid for one person rerunning experiments
    locally, since concurrent Space users could collide on the same
    seed/name and clobber or read each other's files."""
    fd, stem = tempfile.mkstemp()
    os.close(fd)
    os.remove(stem)  # audio_write creates its own file at stem + extension
    path = audio_write(stem, wav, sample_rate, loudness_compressor=True,
                        strategy="loudness", loudness_headroom_db=16)
    return str(path)


def _beat_aligned_length(path, min_length, max_length, unit_size=0.02, length_retries=2):
    """Find a loop length, in 20ms units, aligned to a whole number of bars,
    within [min_length, max_length] seconds. Adapted from pipelines.py's
    get_beat_perfect_hint, unchanged except for taking an already-constructed
    file2beats (construction moved to process_fn, see its docstring)."""
    min_units = min_length // unit_size
    max_units = max_length // unit_size

    beats, downbeats = file2beats(path)

    diffs = []
    prev = None
    beats_per_bar = []
    beats_per_this_bar = 0
    downbeat_idx = 0
    for beat in beats:
        while downbeat_idx < len(downbeats) and beat > downbeats[downbeat_idx]:
            beats_per_bar.append(beats_per_this_bar)
            beats_per_this_bar = 0
            downbeat_idx += 1
        beats_per_this_bar += 1
        if prev is not None:
            diffs.append(beat - prev)
        prev = beat

    if len(beats_per_bar) < 1:
        return None

    diffs = sorted(diffs)
    beats_per_bar = sorted(beats_per_bar)

    beat_length = diffs[len(diffs) // 2] // unit_size + 1  # quantize the median to unit_size
    beats_per_bar = beats_per_bar[len(beats_per_bar) // 2]  # median beats-per-bar
    bar_length = beats_per_bar * beat_length
    length = bar_length * 12
    if length <= 0:
        return None

    loops_done = 0
    while length < min_units or length > max_units:
        if length < min_units:
            length *= 2
        else:
            length /= 2
        loops_done += 1

    if loops_done > length_retries:
        return None
    return length


@spaces.GPU(duration=300)
@torch.inference_mode()
def process_fn(description: str, min_length: float, max_length: float, seed: int) -> str:
    """Generates a seamlessly loopable music clip: a MusicGen hint is used to
    find a beat-aligned loop length (retrying up to MAX_RETRIES times), then
    MAGNeT generates the final loop as a continuation of that hint. Adapted
    from pipelines.py's loopgen_loop.

    Model construction happens here, inside the GPU call, rather than in the
    background-download thread above: LoopGen/MusicGen.get_pretrained()
    bakes device="cuda" in immediately, and ZeroGPU only allows CUDA calls
    made inside an @spaces.GPU-decorated call, not from a background thread."""
    global musicgen, loopgen, file2beats

    if not beat_checkpoint_ready:
        raise gr.Error("Beat-tracking checkpoint is still downloading, please try again shortly.")
    if beat_checkpoint_error is not None:
        raise gr.Error(f"Beat-tracking checkpoint failed to download: {beat_checkpoint_error}")

    with model_lock:
        if musicgen is None:
            musicgen = MusicGen.get_pretrained(MUSICGEN_MODEL, device=DEVICE)
        if loopgen is None:
            loopgen = LoopGen.get_pretrained(MAGNET_MODEL, device=DEVICE)
        if file2beats is None:
            file2beats = File2Beats(checkpoint_path=beat_checkpoint_path, device=DEVICE)

    torch.manual_seed(seed)
    np.random.seed(seed)

    unit_length = None
    hint = None
    sr = None
    tries = 0
    while unit_length is None:
        gc.collect()
        if DEVICE == "cuda":
            torch.cuda.empty_cache()
        if tries >= MAX_RETRIES:
            unit_length = 350 * 3
            break
        del hint
        tries += 1
        musicgen.set_generation_params(duration=HINT_DURATION)
        hint = musicgen.generate([description])[0].to(device="cpu", dtype=torch.float32)
        sr = musicgen.sample_rate
        hint_path = _write_temp_wav(hint, sr)
        unit_length = _beat_aligned_length(hint_path, min_length, max_length)
        os.remove(hint_path)

    gc.collect()
    if DEVICE == "cuda":
        torch.cuda.empty_cache()

    left_hint = [hint[..., :int(0.02 * sr) * int(unit_length / 2)]]
    valid_tokens = [int(unit_length)]

    setup = {
        "span_arrangement": "stride1",
        "use_sampling": True,
        "top_k": 0,
        "top_p": 0.9,
        "temperature": 3.0,
        "max_cfg_coef": MAX_CFG,
        "min_cfg_coef": MIN_CFG,
        "decoding_steps": [100, 50, 10, 10],
        "rescorer": musicgen.lm,
        "rescore_weights": RESCORE_WEIGHT,
        "offset": False,
    }
    loopgen.set_generation_params(**setup)

    autocast_ctx = (torch.autocast(device_type="cuda", dtype=torch.float16)
                     if DEVICE == "cuda" else contextlib.nullcontext())
    with autocast_ctx:
        results = loopgen.generate_continuation(
            text_prompt=description, negative_text_prompt=None,
            left_hint=left_hint, left_hint_sr=sr, valid_tokens=valid_tokens,
        )
    result = results[0].to(device="cpu", dtype=torch.float32)

    max_duration = int(loopgen.duration * loopgen.frame_rate)
    if valid_tokens[0] == max_duration:
        result = torch.cat([result, result], -1)

    return _write_temp_wav(result, loopgen.sample_rate)


with gr.Blocks() as demo:
    input_components = [
        gr.Textbox(label="Description",
                   info="Text description of the music to generate."),
        gr.Slider(minimum=5, maximum=30, step=1, value=15, label="Min Loop Length (s)",
                  info="Shortest acceptable loop length in seconds (default: 15, per repo)."),
        gr.Slider(minimum=5, maximum=30, step=1, value=30, label="Max Loop Length (s)",
                  info="Longest acceptable loop length in seconds (default: 30, per repo)."),
        gr.Number(value=42, precision=0, label="Seed",
                  info="Random seed for generation (default: 42, per repo)."),
    ]
    output_components = [
        gr.Audio(type="filepath", label="Output Audio").set_info("Generated loopable music clip."),
    ]

    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)