| from __future__ import annotations |
|
|
| import json |
| import tempfile |
| import uuid |
| from pathlib import Path |
|
|
| import gradio as gr |
| import soundfile as sf |
|
|
| try: |
| import spaces |
| except ImportError: |
| 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) |
| return args[0] |
|
|
| from pyharp import ModelCard, build_endpoint |
|
|
| from muq_mulan_runtime import rank_descriptions |
|
|
| MIN_AUDIO_SECONDS = 10 |
| MAX_AUDIO_SECONDS = 60 |
| MAX_DESCRIPTIONS = 8 |
| MAX_DESCRIPTION_LENGTH = 300 |
| OUTPUT_ROOT = Path(tempfile.gettempdir()) / "muq_mulan_outputs" |
|
|
| model_card = ModelCard( |
| name="MuQ-MuLan", |
| description=( |
| "Rank English music descriptions by their similarity to an " |
| "uploaded music clip." |
| ), |
| author="Tencent AI Lab", |
| tags=[ |
| "music-information-retrieval", |
| "music-text-retrieval", |
| "music-tagging", |
| "audio-analysis", |
| ], |
| ) |
|
|
|
|
| def _validate_audio(path: str | None) -> str: |
| if not path: |
| raise gr.Error("Please upload a music clip.") |
|
|
| try: |
| duration = sf.info(path).duration |
| except Exception as exc: |
| raise gr.Error(f"Could not read the audio file: {exc}") from exc |
|
|
| if duration < MIN_AUDIO_SECONDS: |
| raise gr.Error( |
| f"Audio must be at least {MIN_AUDIO_SECONDS} seconds long. " |
| f"Received {duration:.1f} seconds." |
| ) |
| if duration > MAX_AUDIO_SECONDS: |
| raise gr.Error( |
| f"Audio must be no longer than {MAX_AUDIO_SECONDS} seconds. " |
| f"Received {duration:.1f} seconds." |
| ) |
| return path |
|
|
|
|
| def _parse_descriptions(value: str | None) -> list[str]: |
| descriptions = [ |
| line.strip() |
| for line in (value or "").splitlines() |
| if line.strip() |
| ] |
| if not descriptions: |
| raise gr.Error("Enter at least one music description.") |
| if len(descriptions) > MAX_DESCRIPTIONS: |
| raise gr.Error( |
| f"Enter no more than {MAX_DESCRIPTIONS} descriptions." |
| ) |
| if any(len(description) > MAX_DESCRIPTION_LENGTH for description in descriptions): |
| raise gr.Error( |
| "Each description must be no more than " |
| f"{MAX_DESCRIPTION_LENGTH} characters." |
| ) |
| return descriptions |
|
|
|
|
| @spaces.GPU(duration=240) |
| def process_fn( |
| input_audio: str | None, |
| candidate_descriptions: str | None, |
| ) -> str: |
| input_audio = _validate_audio(input_audio) |
| descriptions = _parse_descriptions(candidate_descriptions) |
|
|
| try: |
| results = rank_descriptions(input_audio, descriptions) |
| except Exception as exc: |
| raise gr.Error(f"MuQ-MuLan inference failed: {exc}") from exc |
|
|
| output_dir = OUTPUT_ROOT / uuid.uuid4().hex |
| output_dir.mkdir(parents=True, exist_ok=True) |
| output_path = output_dir / "muq_mulan_similarity.json" |
| output_path.write_text( |
| json.dumps( |
| { |
| "model": "OpenMuQ/MuQ-MuLan-large", |
| "score_type": "cosine_similarity", |
| "score_range": [-1.0, 1.0], |
| "results": results, |
| }, |
| ensure_ascii=False, |
| indent=2, |
| ) |
| + "\n", |
| encoding="utf-8", |
| ) |
| return str(output_path) |
|
|
|
|
| with gr.Blocks(title="MuQ-MuLan Music-Text Similarity") as demo: |
| input_components = [ |
| gr.Audio( |
| type="filepath", |
| label="Music Audio", |
| ) |
| .harp_required(True) |
| .set_info("Music clip between 10 and 60 seconds long."), |
| gr.Textbox( |
| lines=5, |
| label="Candidate Descriptions", |
| placeholder=( |
| "upbeat electronic dance music\n" |
| "slow acoustic ballad\n" |
| "bright piano melody" |
| ), |
| ) |
| .harp_required(True) |
| .set_info("Enter one English description per line."), |
| ] |
| output_components = [ |
| gr.File( |
| type="filepath", |
| file_types=[".json"], |
| label="Similarity Ranking", |
| ).set_info("Descriptions ranked by cosine similarity."), |
| ] |
| build_endpoint( |
| model_card=model_card, |
| input_components=input_components, |
| output_components=output_components, |
| process_fn=process_fn, |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| demo.queue(default_concurrency_limit=1).launch( |
| show_error=True, |
| pwa=True, |
| ) |
|
|