MuQ-MuLan / muq_mulan_runtime.py
harp-dev's picture
Deploy MuQ-MuLan music-text similarity endpoint
96558cb
Raw
History Blame Contribute Delete
3.02 kB
from __future__ import annotations
import json
from functools import lru_cache
from pathlib import Path
import soundfile as sf
import torch
import torchaudio.functional as AF
from huggingface_hub import hf_hub_download, snapshot_download
from muq import MuQMuLan
SAMPLE_RATE = 24_000
MODEL_REPO = "OpenMuQ/MuQ-MuLan-large"
MODEL_REVISION = "2e01c796b71dca71b45251384c04cd7b237c9020"
AUDIO_MODEL_REPO = "OpenMuQ/MuQ-large-msd-iter"
AUDIO_MODEL_REVISION = "0562a57814f6f8bbd9fdea0a25921a2fce1a841a"
TEXT_MODEL_REPO = "xlm-roberta-base"
TEXT_MODEL_REVISION = "e73636d4f797dec63c3081bb6ed5c7b0bb3f2089"
def _select_device() -> str:
if torch.cuda.is_available():
return "cuda"
if torch.backends.mps.is_available():
return "mps"
return "cpu"
@lru_cache(maxsize=1)
def _load_model() -> MuQMuLan:
config_path = hf_hub_download(
repo_id=MODEL_REPO,
filename="config.json",
revision=MODEL_REVISION,
)
config = json.loads(Path(config_path).read_text(encoding="utf-8"))
audio_model_path = snapshot_download(
repo_id=AUDIO_MODEL_REPO,
revision=AUDIO_MODEL_REVISION,
allow_patterns=[
"config.json",
"model.safetensors",
],
)
text_model_path = snapshot_download(
repo_id=TEXT_MODEL_REPO,
revision=TEXT_MODEL_REVISION,
allow_patterns=[
"config.json",
"model.safetensors",
"sentencepiece.bpe.model",
"special_tokens_map.json",
"tokenizer.json",
"tokenizer_config.json",
],
)
config["audio_model"]["name"] = audio_model_path
config["text_model"]["name"] = text_model_path
return MuQMuLan.from_pretrained(
MODEL_REPO,
revision=MODEL_REVISION,
config=config,
).eval()
def _load_audio(path: str) -> torch.Tensor:
audio, sample_rate = sf.read(
Path(path),
dtype="float32",
always_2d=True,
)
waveform = torch.from_numpy(audio).mean(dim=1)
if sample_rate != SAMPLE_RATE:
waveform = AF.resample(waveform, sample_rate, SAMPLE_RATE)
return waveform.unsqueeze(0)
@torch.inference_mode()
def rank_descriptions(
audio_path: str,
descriptions: list[str],
) -> list[dict[str, object]]:
device_name = _select_device()
device = torch.device(device_name)
model = _load_model().to(device)
waveform = _load_audio(audio_path).to(device)
audio_embedding = model(wavs=waveform)
text_embeddings = model(texts=descriptions)
scores = model.calc_similarity(
audio_embedding,
text_embeddings,
)[0].detach().cpu().tolist()
ranked = sorted(
zip(descriptions, scores),
key=lambda item: item[1],
reverse=True,
)
return [
{
"rank": index,
"description": description,
"similarity": round(float(score), 6),
}
for index, (description, score) in enumerate(ranked, start=1)
]