Music2Emo / music2emo_runtime.py
harp-dev's picture
Deploy Music2Emo emotion analysis endpoint
33f35c7
Raw
History Blame Contribute Delete
14.6 kB
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")
@property
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()
@staticmethod
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
@staticmethod
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)