wx9Songs's picture
Update app.py
ba4c0e4 verified
Raw
History Blame Contribute Delete
6.83 kB
import os
import time
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
# ZeroGPU must be imported before torch or any module that imports torch.
import spaces
import gradio as gr
import numpy as np
import soundfile as sf
import torch
from scipy.signal import resample_poly
from src.modeling_moss_music import MossMusicModel
from src.processing_moss_music import MossMusicProcessor
MODEL_ID = "OpenMOSS-Team/MOSS-Music-8B-Instruct"
TITLE = "MOSS-Music-8B-Instruct"
DEFAULT_PROMPT = (
"请从风格与速度、调性与和声、乐器编配、结构安排以及整体情绪几个方面描述这段音乐。"
)
# ZeroGPU emulates CUDA while the app starts, packs the model weights, and
# materializes them on the allocated GPU when analyze_audio is called.
model = MossMusicModel.from_pretrained(
MODEL_ID,
trust_remote_code=True,
torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True,
).to("cuda")
model.eval()
processor = MossMusicProcessor.from_pretrained(
MODEL_ID,
trust_remote_code=True,
enable_time_marker=True,
)
def load_audio(path: str, sample_rate: int) -> np.ndarray:
"""Load an audio file as mono float32 and resample it when necessary."""
waveform, original_sample_rate = sf.read(path, dtype="float32", always_2d=False)
if waveform.ndim > 1:
waveform = waveform.mean(axis=1)
if original_sample_rate != sample_rate:
waveform = resample_poly(waveform, sample_rate, original_sample_rate)
return np.asarray(waveform, dtype=np.float32)
@spaces.GPU(size="large", duration=300)
def analyze_audio(
audio_path: str | None,
prompt: str,
max_new_tokens: int,
temperature: float,
top_p: float,
top_k: int,
) -> tuple[str, str]:
question = (prompt or "").strip() or DEFAULT_PROMPT
started_at = time.perf_counter()
if audio_path:
raw_audio = load_audio(audio_path, sample_rate=processor.config.mel_sr)
inputs = processor(text=question, audios=[raw_audio], return_tensors="pt")
else:
inputs = processor(text=question, return_tensors="pt")
inputs = inputs.to(model.device)
if inputs.get("audio_data") is not None:
inputs["audio_data"] = inputs["audio_data"].to(model.dtype)
inputs["audio_input_mask"] = inputs["input_ids"] == processor.audio_token_id
generation_kwargs = {
"max_new_tokens": int(max_new_tokens),
"num_beams": 1,
"use_cache": True,
}
if temperature > 0:
generation_kwargs.update(
do_sample=True,
temperature=float(temperature),
top_p=float(top_p),
top_k=int(top_k),
)
else:
generation_kwargs["do_sample"] = False
with torch.inference_mode():
generated_ids = model.generate(**inputs, **generation_kwargs)
input_length = inputs["input_ids"].shape[1]
answer = processor.decode(
generated_ids[0, input_length:],
skip_special_tokens=True,
).strip()
elapsed = time.perf_counter() - started_at
generated_tokens = generated_ids.shape[1] - input_length
status = (
f"Model: `{MODEL_ID}` \n"
f"Elapsed: `{elapsed:.2f}s` \n"
f"Generated tokens: `{generated_tokens}`"
)
return answer, status
CSS = """
#app-container { max-width: 1120px; margin: 0 auto; }
"""
with gr.Blocks(title=TITLE, css=CSS) as demo:
with gr.Column(elem_id="app-container"):
gr.Markdown(f"# {TITLE}")
gr.Markdown(
"上传整曲音频音乐后,可以进行音乐描述、歌词识别、和弦/调性/速度分析、"
"段落结构分析和开放式音乐问答。"
)
with gr.Row():
with gr.Column(scale=5):
audio_input = gr.Audio(
label="音乐文件",
sources=["upload", "microphone"],
type="filepath",
)
prompt_input = gr.Textbox(
label="问题 / Prompt",
value=DEFAULT_PROMPT,
lines=4,
)
with gr.Accordion("高级参数", open=False):
max_new_tokens_input = gr.Slider(
minimum=64,
maximum=2048,
value=512,
step=32,
label="Max new tokens",
)
temperature_input = gr.Slider(
minimum=0,
maximum=1.5,
value=1.0,
step=0.1,
label="Temperature(0 为贪心解码)",
)
top_p_input = gr.Slider(
minimum=0.1,
maximum=1.0,
value=0.8,
step=0.05,
label="Top-p",
)
top_k_input = gr.Slider(
minimum=1,
maximum=100,
value=50,
step=1,
label="Top-k",
)
with gr.Row():
submit_button = gr.Button("开始分析", variant="primary")
gr.ClearButton(
[audio_input, prompt_input],
value="清空",
)
with gr.Column(scale=5):
output_text = gr.Textbox(
label="分析结果",
lines=20,
show_copy_button=True,
)
status_text = gr.Markdown("等待输入。")
gr.Examples(
examples=[
[None, "请详细描述这段音乐的风格、情绪、速度与主要乐器。"],
[None, "请转录这首歌的歌词,并尽量给出时间戳。"],
[None, "请分析这首歌的调性、速度和和弦进行。"],
[None, "请将这首歌划分为 intro、verse、chorus、bridge 和 outro。"],
],
inputs=[audio_input, prompt_input],
label="示例问题",
)
submit_button.click(
fn=analyze_audio,
inputs=[
audio_input,
prompt_input,
max_new_tokens_input,
temperature_input,
top_p_input,
top_k_input,
],
outputs=[output_text, status_text],
api_name="analyze",
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1, max_size=8).launch(
server_name="0.0.0.0",
server_port=7860,
ssr_mode=False,
)