Spaces:
Running on Zero
Running on Zero
| from __future__ import annotations | |
| import importlib | |
| import json | |
| import sys | |
| import threading | |
| from pathlib import Path | |
| import librosa | |
| import mir_eval | |
| import numpy as np | |
| import torch | |
| import torchaudio | |
| from huggingface_hub import snapshot_download | |
| from music21 import note, stream | |
| from torch import nn | |
| from transformers import AutoModel, Wav2Vec2FeatureExtractor | |
| SOURCE_REPO = "amaai-lab/music2emo" | |
| SOURCE_REVISION = "b036e59471583c3d5b30c69e63e8c7323cc36c4a" | |
| MERT_REPO = "m-a-p/MERT-v1-95M" | |
| MERT_REVISION = "12af15fef9d0ac838c3f475bfbbf26d2060dd4f5" | |
| SAMPLE_RATE = 24000 | |
| WINDOW_SECONDS = 30 | |
| MOOD_CLASSES = 56 | |
| _LOCK = threading.Lock() | |
| _RUNTIME = None | |
| class PositionalEncoding(nn.Module): | |
| def __init__(self, width: int, max_length: int = 100): | |
| super().__init__() | |
| encoding = torch.zeros(max_length, width) | |
| position = torch.arange(max_length, dtype=torch.float32).unsqueeze(1) | |
| scale = torch.exp( | |
| torch.arange(0, width, 2).float() * (-np.log(10000.0) / width) | |
| ) | |
| encoding[:, 0::2] = torch.sin(position * scale) | |
| encoding[:, 1::2] = torch.cos(position * scale) | |
| self.register_buffer("encoding", encoding.unsqueeze(0), persistent=False) | |
| def forward(self, values: torch.Tensor) -> torch.Tensor: | |
| return values + self.encoding[:, : values.size(1)] | |
| class EmotionHead(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.root_embedding = nn.Embedding(14, 4) | |
| self.attribute_embedding = nn.Embedding(14, 4) | |
| self.position = PositionalEncoding(8) | |
| layer = nn.TransformerEncoderLayer( | |
| d_model=8, | |
| nhead=8, | |
| dim_feedforward=64, | |
| dropout=0.1, | |
| batch_first=True, | |
| ) | |
| self.chord_transformer = nn.TransformerEncoder(layer, num_layers=2) | |
| self.input_projection = nn.Sequential(nn.Linear(1545, 512), nn.ReLU()) | |
| self.classifier = nn.Sequential( | |
| nn.Linear(512, 256), | |
| nn.ReLU(), | |
| nn.Linear(256, MOOD_CLASSES), | |
| ) | |
| self.regressor = nn.Sequential( | |
| nn.Linear(512, 256), | |
| nn.ReLU(), | |
| nn.Linear(256, 2), | |
| ) | |
| def forward( | |
| self, | |
| mert: torch.Tensor, | |
| chord_roots: torch.Tensor, | |
| chord_attributes: torch.Tensor, | |
| mode: torch.Tensor, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| chord_values = torch.cat( | |
| ( | |
| self.root_embedding(chord_roots), | |
| self.attribute_embedding(chord_attributes), | |
| ), | |
| dim=-1, | |
| ) | |
| chord_values = self.position(chord_values) | |
| cls_token = torch.zeros_like(chord_values[:, :1]) | |
| chord_values = self.chord_transformer( | |
| torch.cat((cls_token, chord_values), dim=1) | |
| )[:, 0] | |
| combined = torch.cat((mert, chord_values, mode.float()), dim=1) | |
| hidden = self.input_projection(combined) | |
| return self.classifier(hidden), self.regressor(hidden) | |
| class Music2EmoRuntime: | |
| def __init__(self): | |
| self.source_dir = Path( | |
| snapshot_download( | |
| repo_id=SOURCE_REPO, | |
| revision=SOURCE_REVISION, | |
| allow_patterns=[ | |
| "inference/data/*", | |
| "saved_models/J_all.ckpt", | |
| "utils/*.py", | |
| ], | |
| ) | |
| ) | |
| sys.path.insert(0, str(self.source_dir)) | |
| self._load_source_modules() | |
| self.mert = AutoModel.from_pretrained( | |
| MERT_REPO, | |
| revision=MERT_REVISION, | |
| trust_remote_code=True, | |
| ) | |
| self.processor = Wav2Vec2FeatureExtractor.from_pretrained( | |
| MERT_REPO, | |
| revision=MERT_REVISION, | |
| trust_remote_code=True, | |
| ) | |
| self.head = EmotionHead() | |
| self._load_emotion_checkpoint() | |
| self.chord_model = self.BTCModel(config=self.config.model) | |
| self._load_chord_checkpoint() | |
| tags = np.load(self.data_dir / "tag_list.npy", allow_pickle=True) | |
| self.mood_labels = [ | |
| str(tag).replace("mood/theme---", "") for tag in tags[-MOOD_CLASSES:] | |
| ] | |
| self.root_map = self._read_json("chord_root.json") | |
| self.attribute_map = self._read_json("chord_attr.json") | |
| def data_dir(self) -> Path: | |
| return self.source_dir / "inference" / "data" | |
| def _load_source_modules(self) -> None: | |
| hparams = importlib.import_module("utils.hparams") | |
| btc_model = importlib.import_module("utils.btc_model") | |
| chords = importlib.import_module("utils.mir_eval_modules") | |
| self.config = hparams.HParams.load(self.data_dir / "run_config.yaml") | |
| self.config.feature["large_voca"] = True | |
| self.config.model["num_chords"] = 170 | |
| self.BTCModel = btc_model.BTC_model | |
| self.chord_vocabulary = chords.idx2voca_chord() | |
| def _read_json(self, name: str) -> dict[str, int]: | |
| return json.loads((self.data_dir / name).read_text(encoding="utf-8")) | |
| def _load_emotion_checkpoint(self) -> None: | |
| checkpoint = torch.load( | |
| self.source_dir / "saved_models" / "J_all.ckpt", | |
| map_location="cpu", | |
| weights_only=False, | |
| ) | |
| state = { | |
| key.removeprefix("model."): value | |
| for key, value in checkpoint["state_dict"].items() | |
| } | |
| rename = { | |
| "chord_root_embedding.": "root_embedding.", | |
| "chord_attr_embedding.": "attribute_embedding.", | |
| "positional_encoding.": "position.", | |
| "input_proj.": "input_projection.", | |
| "classification_branch.": "classifier.", | |
| "regression_branch.": "regressor.", | |
| } | |
| converted = {} | |
| for key, value in state.items(): | |
| for source, target in rename.items(): | |
| if key.startswith(source): | |
| key = target + key[len(source) :] | |
| break | |
| converted[key] = value | |
| expected = self.head.state_dict() | |
| converted = {key: value for key, value in converted.items() if key in expected} | |
| self.head.load_state_dict(converted, strict=True) | |
| self.head.eval() | |
| def _load_chord_checkpoint(self) -> None: | |
| checkpoint = torch.load( | |
| self.data_dir / "btc_model_large_voca.pt", | |
| map_location="cpu", | |
| weights_only=False, | |
| ) | |
| self.chord_mean = checkpoint["mean"] | |
| self.chord_std = checkpoint["std"] | |
| self.chord_model.load_state_dict(checkpoint["model"]) | |
| self.chord_model.eval() | |
| def _audio(path: str) -> tuple[torch.Tensor, int]: | |
| waveform, sample_rate = torchaudio.load(path) | |
| waveform = waveform.mean(dim=0) | |
| if sample_rate != SAMPLE_RATE: | |
| waveform = torchaudio.functional.resample( | |
| waveform, | |
| sample_rate, | |
| SAMPLE_RATE, | |
| ) | |
| return waveform, SAMPLE_RATE | |
| def _mert_embedding( | |
| self, | |
| waveform: torch.Tensor, | |
| device: torch.device, | |
| ) -> torch.Tensor: | |
| window = WINDOW_SECONDS * SAMPLE_RATE | |
| chunks = waveform.split(window) | |
| embeddings = [] | |
| for chunk in chunks: | |
| inputs = self.processor( | |
| chunk, | |
| sampling_rate=SAMPLE_RATE, | |
| return_tensors="pt", | |
| ) | |
| inputs = {key: value.to(device) for key, value in inputs.items()} | |
| outputs = self.mert(**inputs, output_hidden_states=True) | |
| layer_means = torch.stack(outputs.hidden_states[1:]).mean(dim=2) | |
| embeddings.append(torch.cat((layer_means[5], layer_means[6]), dim=1)) | |
| return torch.stack(embeddings).mean(dim=0) | |
| def _chord_intervals( | |
| self, | |
| audio_path: str, | |
| device: torch.device, | |
| ) -> list[tuple[float, float, str]]: | |
| config = self.config | |
| audio, sample_rate = librosa.load( | |
| audio_path, | |
| sr=config.mp3["song_hz"], | |
| mono=True, | |
| ) | |
| feature = librosa.cqt( | |
| audio, | |
| sr=sample_rate, | |
| n_bins=config.feature["n_bins"], | |
| bins_per_octave=config.feature["bins_per_octave"], | |
| hop_length=config.feature["hop_length"], | |
| ) | |
| feature = np.log(np.abs(feature) + 1e-6).T | |
| feature = (feature - self.chord_mean) / self.chord_std | |
| timestep = config.model["timestep"] | |
| pad = timestep - (feature.shape[0] % timestep) | |
| feature = np.pad(feature, ((0, pad), (0, 0))) | |
| blocks = feature.shape[0] // timestep | |
| frame_seconds = config.mp3["inst_len"] / timestep | |
| changes: list[tuple[float, float, str]] = [] | |
| start = 0.0 | |
| previous = None | |
| tensor = torch.tensor(feature, dtype=torch.float32).unsqueeze(0).to(device) | |
| for block in range(blocks): | |
| section = tensor[:, block * timestep : (block + 1) * timestep] | |
| encoded, _ = self.chord_model.self_attn_layers(section) | |
| prediction, _ = self.chord_model.output_layer(encoded) | |
| for offset, chord_index in enumerate(prediction.squeeze().tolist()): | |
| frame = block * timestep + offset | |
| if frame >= feature.shape[0] - pad: | |
| break | |
| if previous is None: | |
| previous = chord_index | |
| elif chord_index != previous: | |
| end = frame * frame_seconds | |
| changes.append((start, end, self.chord_vocabulary[previous])) | |
| start = end | |
| previous = chord_index | |
| duration = len(audio) / sample_rate | |
| if previous is not None and duration > start: | |
| changes.append((start, duration, self.chord_vocabulary[previous])) | |
| return changes | |
| def _key(intervals: list[tuple[float, float, str]]) -> tuple[str, str]: | |
| score = stream.Stream() | |
| note_count = 0 | |
| for start, end, chord in intervals: | |
| root, bitmap, _ = mir_eval.chord.encode(chord) | |
| if root < 0: | |
| continue | |
| chroma = mir_eval.chord.rotate_bitmap_to_root(bitmap, root) | |
| for pitch_class, active in enumerate(chroma): | |
| if active: | |
| value = note.Note(48 + pitch_class) | |
| value.duration.quarterLength = max(end - start, 0.01) | |
| score.insert(start, value) | |
| note_count += 1 | |
| if note_count == 0: | |
| return "C", "major" | |
| key = score.analyze("key") | |
| tonic = str(key.tonic).replace("-", "b") | |
| return tonic, str(key.mode) | |
| def _encode_chords( | |
| self, | |
| intervals: list[tuple[float, float, str]], | |
| tonic: str, | |
| mode: str, | |
| device: torch.device, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| pitch_classes = [ | |
| "C", | |
| "C#", | |
| "D", | |
| "D#", | |
| "E", | |
| "F", | |
| "F#", | |
| "G", | |
| "G#", | |
| "A", | |
| "A#", | |
| "B", | |
| ] | |
| flat_to_sharp = { | |
| "Cb": "B", | |
| "Db": "C#", | |
| "Eb": "D#", | |
| "Fb": "E", | |
| "Gb": "F#", | |
| "Ab": "G#", | |
| "Bb": "A#", | |
| } | |
| tonic = flat_to_sharp.get(tonic, tonic) | |
| reference = "A" if mode == "minor" else "C" | |
| shift = (pitch_classes.index(tonic) - pitch_classes.index(reference)) % 12 | |
| roots = [] | |
| attributes = [] | |
| for _, _, chord in intervals[:100]: | |
| if chord in {"N", "X"}: | |
| root, attribute = chord, 0 | |
| else: | |
| parts = chord.split(":", 1) | |
| source_root = flat_to_sharp.get(parts[0], parts[0]) | |
| root = pitch_classes[ | |
| (pitch_classes.index(source_root) - shift) % 12 | |
| ] | |
| attribute_name = parts[1] if len(parts) == 2 else "maj" | |
| attribute = self.attribute_map.get(attribute_name, 0) | |
| roots.append(self.root_map.get(root, 0)) | |
| attributes.append(attribute) | |
| roots.extend([0] * (100 - len(roots))) | |
| attributes.extend([0] * (100 - len(attributes))) | |
| mode_value = 1 if mode == "minor" else 0 | |
| return ( | |
| torch.tensor(roots, dtype=torch.long, device=device).unsqueeze(0), | |
| torch.tensor(attributes, dtype=torch.long, device=device).unsqueeze(0), | |
| torch.tensor([[mode_value]], dtype=torch.long, device=device), | |
| ) | |
| def predict(self, audio_path: str, threshold: float) -> dict: | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.mert.to(device).eval() | |
| self.head.to(device).eval() | |
| self.chord_model.to(device).eval() | |
| waveform, _ = self._audio(audio_path) | |
| with torch.inference_mode(): | |
| mert = self._mert_embedding(waveform, device) | |
| intervals = self._chord_intervals(audio_path, device) | |
| tonic, mode = self._key(intervals) | |
| roots, attributes, mode_tensor = self._encode_chords( | |
| intervals, | |
| tonic, | |
| mode, | |
| device, | |
| ) | |
| logits, dimensions = self.head( | |
| mert, | |
| roots, | |
| attributes, | |
| mode_tensor, | |
| ) | |
| probabilities = torch.sigmoid(logits).squeeze().cpu().tolist() | |
| valence, arousal = dimensions.squeeze().cpu().tolist() | |
| ranked = sorted( | |
| ( | |
| {"label": label, "probability": round(float(score), 4)} | |
| for label, score in zip(self.mood_labels, probabilities) | |
| if score >= threshold | |
| ), | |
| key=lambda item: item["probability"], | |
| reverse=True, | |
| ) | |
| return { | |
| "model": "Music2Emo", | |
| "moods": ranked, | |
| "valence": round(float(valence), 4), | |
| "arousal": round(float(arousal), 4), | |
| "scale": {"valence": [1, 9], "arousal": [1, 9]}, | |
| "threshold": float(threshold), | |
| "estimated_key": f"{tonic} {mode}", | |
| } | |
| def analyze_music(audio_path: str, threshold: float = 0.5) -> dict: | |
| global _RUNTIME | |
| with _LOCK: | |
| if _RUNTIME is None: | |
| _RUNTIME = Music2EmoRuntime() | |
| return _RUNTIME.predict(audio_path, threshold) | |