File size: 4,704 Bytes
fa3196f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0c534f3
fa3196f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fb2f497
fa3196f
 
 
 
0c534f3
fa3196f
 
 
 
0c534f3
 
fa3196f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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 threading

import audiotools
import gradio as gr
import torch
from pyharp import ModelCard, build_endpoint, save_audio

from audiocraft.data.audio_utils import normalize_audio
from audiocraft.models import MusicGen

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
CHECKPOINT_REPO = "Cyan0731/MusiConGen"

model = None
model_ready = False   # has model been moved onto the GPU yet?
model_loading = True
model_error = None


def load_model():
    """Download + construct on CPU only. ZeroGPU only intercepts CUDA calls 
    made inside @spaces.GPU decorated call."""
    global model, model_loading, model_error
    try:
        model = MusicGen.get_pretrained(CHECKPOINT_REPO, device="cpu")
        print("Model loaded (CPU).")
    except Exception as e:
        model_error = str(e)
        print(f"Load error: {e}")
    finally:
        model_loading = False


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


model_card = ModelCard(
    name="MusiConGen",
    description=(
        "Text-to-music generation with rhythm and chord control. Generates a "
        "music clip from a text description, a chord progression, and a "
        "tempo/time signature."
    ),
    author="Yun-Han Lan, Wen-Yi Hsiao, Hao-Chung Cheng, Yi-Hsuan Yang",
    tags=["music generation", "text-to-music"],
)


@spaces.GPU(duration=120)
@torch.inference_mode()
def process_fn(description, chords, bpm, meter, duration, conditioning_strength):
    """Generate a music clip conditioned on text, chords, and rhythm."""
    global model, model_ready
    if model_loading:
        raise gr.Error("Model is still loading, please wait a moment and try again.")
    if model is None:
        raise gr.Error(f"Model failed to load: {model_error}")
    if not model_ready:
        model = MusicGen.get_pretrained(CHECKPOINT_REPO, device=DEVICE)
        model_ready = True

    model.set_generation_params(
        duration=duration,
        extend_stride=duration // 2,
        cfg_coef=conditioning_strength,
    )
    wav = model.generate_with_chords_and_beats([description], [chords], [bpm], [meter])

    wav = normalize_audio(wav[0].cpu(), strategy="loudness", loudness_compressor=True, sample_rate=model.sample_rate)
    signal = audiotools.AudioSignal(wav, sample_rate=model.sample_rate)
    return save_audio(signal)


with gr.Blocks() as demo:
    input_components = [
        gr.Textbox(
            label="Description",
            value="A laid-back blues shuffle with a relaxed tempo, warm guitar tones, and a comfortable groove. Instruments: electric guitar, bass, drums.",
            info="Text description of the music to generate.",
        ),
        gr.Textbox(
            label="Chord Progression",
            value="C G A:min F",
            info="Space-separated chord symbols, one per bar, repeating to fill the duration (e.g. 'C G A:min F'). Syntax: root note plus optional ':quality', e.g. C, A:min, D:min7.",
        ),
        gr.Number(
            label="Tempo (BPM)", value=120, minimum=40, maximum=240,
            info="Tempo in beats per minute (default: 120, per repo demo script).",
        ),
        gr.Number(
            label="Time Signature (beats per bar)", value=4, minimum=2, maximum=12,
            info="Numerator of the time signature (default: 4, per repo demo script).",
        ),
        gr.Slider(
            minimum=5, maximum=30, step=1, value=30, label="Duration (seconds)",
            info="Length of the generated clip (default/max: 30s, the model's trained segment length, per repo config segment_duration=30).",
        ),
        gr.Slider(
            minimum=0.0, maximum=10.0, step=0.5, value=3.0, label="Conditioning Strength",
            info="How strongly generation follows the description/chords/rhythm vs. sounding more free (default: 3.0, per paper guidance scale γ).",
        ),
    ]
    output_components = [
        gr.Audio(type="filepath", label="Generated Music").set_info("Generated music, 32kHz."),
    ]

    build_endpoint(
        model_card=model_card,
        input_components=input_components,
        output_components=output_components,
        process_fn=process_fn,
    )

demo.queue().launch(pwa=True)