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