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