neural-mathrock / app.py
anggars's picture
Update app.py
8a16cc4 verified
Raw History Blame
6.8 kB
import torch
import torch.nn as nn
import librosa
import numpy as np
import gradio as gr
from transformers import XLMRobertaModel, XLMRobertaTokenizer, WavLMModel
from huggingface_hub import hf_hub_download
import warnings
warnings.filterwarnings('ignore')
# ── CONFIGURATION ──
REPO_ID = "anggars/neural-mathrock"
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
print("Downloading model from Hugging Face Hub...")
model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
ckpt = torch.load(model_path, map_location=DEVICE, weights_only=False)
# ── LOAD LABEL ENCODERS ──
le_mbti = ckpt['le_mbti']
le_emotion = ckpt['le_emotion']
le_vibe = ckpt['le_vibe']
le_intensity = ckpt['le_intensity']
le_tempo = ckpt['le_tempo']
NUM_MBTI = len(le_mbti.classes_)
NUM_EMOTION = len(le_emotion.classes_)
NUM_VIBE = len(le_vibe.classes_)
NUM_INTENSITY = len(le_intensity.classes_)
NUM_TEMPO = len(le_tempo.classes_)
# ── ARCHITECTURE ──
class HybridMultimodalModel(nn.Module):
def __init__(self):
super().__init__()
self.text_model = XLMRobertaModel.from_pretrained('anggars/xlm-mbti')
self.audio_model = WavLMModel.from_pretrained('microsoft/wavlm-base')
self.audio_proj = nn.Linear(768, 256)
self.fusion = nn.Sequential(
nn.Linear(1024, 512),
nn.BatchNorm1d(512),
nn.ReLU(),
nn.Dropout(0.4),
)
self.head_mbti = nn.Linear(512, NUM_MBTI)
self.head_emotion = nn.Linear(512, NUM_EMOTION)
self.head_vibe = nn.Linear(512, NUM_VIBE)
self.head_intensity = nn.Linear(512, NUM_INTENSITY)
self.head_tempo = nn.Linear(512, NUM_TEMPO)
def forward(self, input_ids, attention_mask, audio_values):
text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
text_feat = text_out.pooler_output
audio_out = self.audio_model(audio_values).last_hidden_state
audio_feat = self.audio_proj(audio_out.mean(dim=1))
fused = self.fusion(torch.cat([text_feat, audio_feat], dim=-1))
return {
'mbti': self.head_mbti(fused),
'emotion': self.head_emotion(fused),
'vibe': self.head_vibe(fused),
'intensity': self.head_intensity(fused),
'tempo': self.head_tempo(fused),
}
print("Initializing architecture and loading weights...")
model = HybridMultimodalModel().to(DEVICE)
model.load_state_dict(ckpt['model_state'], strict=False)
model.eval()
tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')
# ── HYBRID INFERENCE PIPELINE ──
def predict_multimodal(audio_path, lyrics):
if audio_path is None:
return [{"Error": "Audio file cannot be empty."}] * 5
try:
# 1. Text Processing
text = str(lyrics).strip() if lyrics else "[INSTRUMENTAL]"
enc = tokenizer(text, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)
# 2. Audio Processing & DSP
wav, sr = librosa.load(audio_path, sr=16000)
onset_env = librosa.onset.onset_strength(y=wav, sr=sr)
tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=sr)
print(f"Calculated BPM: {tempo_bpm}")
# 3. Chunking (20-second sequential blocks)
chunk_size = 16000 * 20
chunks = []
for i in range(0, len(wav), chunk_size):
chunk = wav[i:i + chunk_size]
if len(chunk) < chunk_size:
chunk = np.pad(chunk, (0, chunk_size - len(chunk)))
chunks.append(chunk)
if not chunks:
chunks = [np.zeros(chunk_size)]
all_logits = {col: [] for col in TARGET_COLS}
# 4. Block-by-block Prediction
with torch.no_grad():
for chunk in chunks:
audio_tensor = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
out = model(enc['input_ids'], enc['attention_mask'], audio_tensor)
for col in TARGET_COLS:
all_logits[col].append(out[col][0])
# 5. Averaging and Formatting Results
final_results = []
encoders_map = {
'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe,
'intensity': le_intensity, 'tempo': le_tempo
}
for col in TARGET_COLS:
avg_logits = torch.stack(all_logits[col]).mean(dim=0)
le = encoders_map[col]
# TEMPO LOGIC OVERRIDE
if col == 'tempo':
tempo_classes = list(le.classes_)
try:
fast_idx = tempo_classes.index('Fast')
mod_idx = tempo_classes.index('Moderate')
if tempo_bpm > 125:
avg_logits[fast_idx] += 5.0
elif tempo_bpm > 90:
avg_logits[mod_idx] += 3.0
except ValueError:
pass
probs = torch.nn.functional.softmax(avg_logits, dim=0).cpu().numpy()
classes = le.classes_
pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes))}
sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
final_results.append(sorted_preds)
return final_results
except Exception as e:
return [{"Error": f"Processing failed: {str(e)}"}] * 5
# ── GRADIO UI ──
with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
gr.Markdown("# Neural Math Rock Multimodal Analysis")
gr.Markdown("Analyzing MBTI, Emotion, and Vibes using Hybrid Deep Learning & DSP.")
with gr.Row():
with gr.Column():
audio_input = gr.Audio(type="filepath", label="Upload Song (Full Analysis)")
lyrics_input = gr.Textbox(lines=5, label="Lyrics Content", placeholder="Paste lyrics here...")
btn = gr.Button("START HYBRID ANALYSIS", variant="primary")
with gr.Column():
out_mbti = gr.Label(label="Personality (MBTI)")
out_emotion = gr.Label(label="Emotional State")
out_vibe = gr.Label(label="Acoustic Vibe")
out_intensity = gr.Label(label="Intensity Level")
out_tempo = gr.Label(label="Tempo Classification (BPM Adjusted)")
btn.click(
fn=predict_multimodal,
inputs=[audio_input, lyrics_input],
outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
)
if __name__ == "__main__":
demo.launch()