Spaces:
Running on Zero
Running on Zero
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)
|