neural-mathrock / app.py
anggars's picture
Update app.py
5de0d17 verified
Raw History Blame Contribute Delete
9.79 kB
import os
import gc
import torch
import torch.nn as nn
import torch.nn.functional as F
import soundfile as sf
import torchaudio
import gradio as gr
import numpy as np
from huggingface_hub import hf_hub_download
from transformers import (
AutoModel,
AutoFeatureExtractor,
AutoTokenizer,
XLMRobertaForSequenceClassification,
XLMRobertaTokenizer
)
import warnings
warnings.filterwarnings('ignore')
# ------------------------------------------------------------------------------
# CONFIGURATION & LABEL MAPPING
# ------------------------------------------------------------------------------
REPO_MAIN = "anggars/neural-mathrock"
REPO_TEXT_MBTI = "anggars/xlm-mbti"
REPO_TEXT_EMO = "anggars/xlm-emotion"
MODEL_FILE = "model.pth"
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CROP_SEC = 5
SR_TARGET = 24000
MBTI_LABELS = sorted([
"INTJ", "INTP", "ENTJ", "ENTP", "INFJ", "INFP", "ENFJ", "ENFP",
"ISTJ", "ISFJ", "ESTJ", "ESFJ", "ISTP", "ISFP", "ESTP", "ESFP"
])
EMOTION_CLASSES_AUDIO = ["anger", "disgust", "fear", "joy", "neutral", "sadness", "surprise"]
EMOTION_CLASSES_TEXT = sorted([
'admiration', 'amusement', 'anger', 'annoyance', 'approval', 'caring', 'confusion',
'curiosity', 'desire', 'disappointment', 'disapproval', 'disgust', 'embarrassment',
'excitement', 'fear', 'gratitude', 'grief', 'joy', 'love', 'nervousness', 'optimism',
'pride', 'realization', 'relief', 'remorse', 'sadness', 'surprise', 'neutral'
])
text_emo2id = {e: i for i, e in enumerate(EMOTION_CLASSES_TEXT)}
# ------------------------------------------------------------------------------
# NEURAL ARCHITECTURE (7-CLASS AUDIO FUSION)
# ------------------------------------------------------------------------------
class MultimodalFusionClassifier(nn.Module):
def __init__(self, num_classes=7):
super().__init__()
self.audio_model = AutoModel.from_pretrained("m-a-p/MERT-v1-330M", trust_remote_code=True)
self.text_model = AutoModel.from_pretrained("FacebookAI/roberta-base")
self.fusion_head = nn.Sequential(
nn.Linear(1024 + 768, 512),
nn.LayerNorm(512),
nn.GELU(),
nn.Dropout(0.5),
nn.Linear(512, 256),
nn.LayerNorm(256),
nn.GELU(),
nn.Dropout(0.4),
nn.Linear(256, num_classes)
)
def forward(self, audio_vals, text_ids, text_mask):
audio_feats = self.audio_model(audio_vals, output_hidden_states=True).last_hidden_state.mean(dim=1)
text_feats = self.text_model(input_ids=text_ids, attention_mask=text_mask).last_hidden_state.mean(dim=1)
return self.fusion_head(torch.cat((audio_feats, text_feats), dim=-1))
# ------------------------------------------------------------------------------
# MODEL INITIALIZATION
# ------------------------------------------------------------------------------
print("[SYSTEM] Booting Models into Memory...")
mert_ext = AutoFeatureExtractor.from_pretrained("m-a-p/MERT-v1-330M", trust_remote_code=True)
roberta_tok = AutoTokenizer.from_pretrained("FacebookAI/roberta-base")
fusion_model = MultimodalFusionClassifier(num_classes=len(EMOTION_CLASSES_AUDIO)).to(DEVICE)
ckpt_path = hf_hub_download(repo_id=REPO_MAIN, filename=MODEL_FILE)
fusion_model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE)['model_state_dict'])
fusion_model.eval()
# Load Text-Only Models (Isolated)
mbti_tok = XLMRobertaTokenizer.from_pretrained(REPO_TEXT_MBTI)
mbti_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_MBTI).to(DEVICE)
mbti_model.eval()
# Load High-Resolution 28-Class Text Emotion Model
emo28_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_EMO).to(DEVICE)
emo28_model.eval()
# ------------------------------------------------------------------------------
# INFERENCE ENGINE
# ------------------------------------------------------------------------------
def process_audio_chunks(audio_path):
waveform_np, orig_sr = sf.read(audio_path)
if len(waveform_np.shape) > 1:
waveform_np = waveform_np.mean(axis=1)
waveform = torch.tensor(waveform_np, dtype=torch.float32).unsqueeze(0)
if orig_sr != SR_TARGET:
waveform = torchaudio.functional.resample(waveform, orig_sr, SR_TARGET)
waveform = waveform.squeeze(0)
chunk_samples = SR_TARGET * CROP_SEC
chunks = [waveform[i:i+chunk_samples].numpy() for i in range(0, len(waveform), chunk_samples) if len(waveform[i:i+chunk_samples]) > SR_TARGET * 2]
if not chunks:
chunks = [F.pad(waveform, (0, chunk_samples - waveform.shape[0])).numpy()]
return chunks[:10]
def show_track_name(audio_path):
if not audio_path:
return ""
return os.path.basename(audio_path)
def analyze_track(audio_path, lyrics_input):
has_audio = audio_path is not None
safe_lyrics = str(lyrics_input).strip() if lyrics_input else ""
has_lyrics = len(safe_lyrics) > 10
if not has_audio and not has_lyrics:
return {"Error": 1.0}, {"Error": 1.0}
res_mbti = {}
res_emo_final = {}
t_emo28_probs = None
# --- EKSEKUSI TEKS (Jika lirik ada) ---
if has_lyrics:
t_in = mbti_tok(safe_lyrics, truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
with torch.no_grad():
mbti_probs = F.softmax(mbti_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
t_emo28_probs = F.softmax(emo28_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
mbti_dict = {MBTI_LABELS[i]: float(mbti_probs[i]) for i in range(len(MBTI_LABELS))}
res_mbti = dict(sorted(mbti_dict.items(), key=lambda x: x[1], reverse=True)[:3])
else:
res_mbti = {"Requires lyrics for MBTI": 1.0}
# --- EKSEKUSI AUDIO (Jika audio ada) ---
if has_audio:
try:
audio_chunks = process_audio_chunks(audio_path)
# PyTorch butuh text tensor buat fusion model. Pakai dummy text kalo lirik kosong
text_for_fusion = safe_lyrics if has_lyrics else "instrumental"
t_inputs = roberta_tok([text_for_fusion] * len(audio_chunks), padding=True, truncation=True, max_length=128, return_tensors="pt")
t_ids = t_inputs["input_ids"].to(DEVICE)
t_mask = t_inputs["attention_mask"].to(DEVICE)
a_inputs = mert_ext(audio_chunks, sampling_rate=SR_TARGET, return_tensors="pt", padding="max_length", truncation=True, max_length=SR_TARGET * CROP_SEC)
a_vals = a_inputs["input_values"].to(DEVICE)
with torch.no_grad():
logits = fusion_model(a_vals, t_ids, t_mask)
chunk_adjusted_probs = F.softmax(logits, dim=-1)
final_avg_probs = chunk_adjusted_probs.mean(dim=0)
audio_7_probs = (final_avg_probs / final_avg_probs.sum()).cpu().numpy()
audio_emo_dict = {EMOTION_CLASSES_AUDIO[i]: float(audio_7_probs[i]) for i in range(len(EMOTION_CLASSES_AUDIO))}
# LOGIC PENENTUAN HASIL AKHIR
if has_lyrics:
# 1. MULTIMODAL (Audio + Lirik)
fusion_28_dict = {label: 0.0 for label in EMOTION_CLASSES_TEXT}
for i, label in enumerate(EMOTION_CLASSES_TEXT):
fusion_28_dict[label] += float(t_emo28_probs[i]) * 0.60
for label in EMOTION_CLASSES_AUDIO:
if label in fusion_28_dict:
fusion_28_dict[label] += audio_emo_dict[label] * 0.40
total_prob = sum(fusion_28_dict.values())
fusion_28_dict = {k: v / total_prob for k, v in fusion_28_dict.items()}
res_emo_final = dict(sorted(fusion_28_dict.items(), key=lambda x: x[1], reverse=True)[:5])
else:
# 2. AUDIO ONLY
res_emo_final = dict(sorted(audio_emo_dict.items(), key=lambda x: x[1], reverse=True)[:5])
except Exception as e:
res_emo_final = {f"Audio Error: {str(e)}": 1.0}
# 3. TEXT ONLY (Kalo audio kosong tapi lirik ada)
elif has_lyrics:
emo_dict = {EMOTION_CLASSES_TEXT[i]: float(t_emo28_probs[i]) for i in range(len(EMOTION_CLASSES_TEXT))}
res_emo_final = dict(sorted(emo_dict.items(), key=lambda x: x[1], reverse=True)[:5])
return res_mbti, res_emo_final
# ------------------------------------------------------------------------------
# GRADIO INTERFACE
# ------------------------------------------------------------------------------
with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
gr.Markdown("# Neural Math Rock Multimodal Analysis")
gr.Markdown("Identify MBTI personality via **XLM-RoBERTa** and emotional states via **Multimodal Late-Fusion (MERT + RoBERTa)** from Math Rock & Midwest Emo tracks.")
with gr.Row():
with gr.Column():
audio_box = gr.Audio(type="filepath", label="Audio Source (.wav / .mp3)")
track_name_box = gr.Textbox(label="Detected Track", interactive=False, placeholder="Track name appears here after upload")
lyrics_box = gr.Textbox(lines=5, label="Lyrics Source", placeholder="Paste lyrics here...\n(If blank: Analyzes 7 Audio Emotions. If filled: Analyzes 28 High-Resolution Emotions.)")
run_btn = gr.Button("Analyze", variant="primary")
with gr.Column():
res_mbti = gr.Label(label="Personality (MBTI - Text Only)")
res_emo = gr.Label(label="Emotional State (Multimodal/Unimodal)")
audio_box.change(fn=show_track_name, inputs=audio_box, outputs=track_name_box)
run_btn.click(
fn=analyze_track,
inputs=[audio_box, lyrics_box],
outputs=[res_mbti, res_emo]
)
if __name__ == "__main__":
demo.launch()