Spaces:
Sleeping
Sleeping
File size: 10,027 Bytes
0df2078 d22cc5f 78bd59e 0df2078 e0c07a1 7aca11d 0df2078 cf6c644 0df2078 cf6c644 0df2078 cf6c644 d22cc5f ae5507d d22cc5f 0df2078 cf6c644 0df2078 8cc551e 78bd59e 0df2078 d22cc5f 0df2078 7bd2706 0df2078 8678f34 0df2078 d22cc5f 0df2078 dd5c062 e0c07a1 0df2078 2b779ce 0df2078 7bd2706 cf6c644 602e3e9 0df2078 602e3e9 0df2078 602e3e9 0df2078 e0c07a1 0df2078 602e3e9 0df2078 602e3e9 0df2078 53b4fdb 0df2078 3bda16f 0df2078 3bda16f 0df2078 3bda16f 0df2078 1e82fc1 0df2078 52119a0 0df2078 1e82fc1 52119a0 0df2078 1e82fc1 cf6c644 0df2078 1e82fc1 ae5507d 92d86ab | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 | 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 = ['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)}
# Calibration Weights derived from SOTA Report 3
RAW_CE_WEIGHTS = [0.563, 16.326, 4.525, 0.968, 2.704, 0.473, 0.702]
SOFT_WEIGHTS = torch.sqrt(torch.tensor(RAW_CE_WEIGHTS, dtype=torch.float32)).to(DEVICE)
TEMPERATURE = 2.0
# ------------------------------------------------------------------------------
# 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 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 = {}
# 1. ISOLATED MBTI & 28-CLASS TEXT INFERENCE
t_emo28_probs = None
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}
# 2. AUDIO FUSION INFERENCE
if has_audio:
try:
audio_chunks = process_audio_chunks(audio_path)
t_inputs = roberta_tok([safe_lyrics] * 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)
avg_logits = logits.mean(dim=0)
scaled_logits = avg_logits / TEMPERATURE
probs = F.softmax(scaled_logits, dim=-1)
adjusted_probs = probs * SOFT_WEIGHTS
audio_7_probs = (adjusted_probs / adjusted_probs.sum()).cpu().numpy()
audio_emo_dict = {EMOTION_CLASSES_AUDIO[i]: float(audio_7_probs[i]) for i in range(len(EMOTION_CLASSES_AUDIO))}
# --- TRUE MULTIMODAL LATE FUSION ENGAGEMENT ---
if has_lyrics:
# Bikin dictionary kosong untuk 28 emosi
fusion_28_dict = {label: 0.0 for label in EMOTION_CLASSES_TEXT}
# Masukin probabilitas murni dari Teks (Bobot 60%)
for i, label in enumerate(EMOTION_CLASSES_TEXT):
fusion_28_dict[label] += float(t_emo28_probs[i]) * 0.60
# Suntikin probabilitas dari Audio 7 Kelas ke 28 Kelas Teks (Bobot 40%)
for label in EMOTION_CLASSES_AUDIO:
if label in fusion_28_dict:
fusion_28_dict[label] += audio_emo_dict[label] * 0.40
# Normalisasi ulang biar totalnya 1.0
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:
# Unimodal Audio (Hanya ngeluarin 7 Kelas)
res_emo_final = dict(sorted(audio_emo_dict.items(), key=lambda x: x[1], reverse=True)[:4])
except Exception as e:
res_emo_final = {f"Audio Error: {str(e)}": 1.0}
# 3. TEXT-ONLY FALLBACK (Kalau audio gak dimasukin)
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 personality (Text) and emotional states (MERT+RoBERTa Multimodal) from Math Rock & Midwest Emo tracks.")
with gr.Row():
with gr.Column():
audio_box = gr.Audio(type="filepath", label="Audio Source (.wav / .mp3)")
lyrics_box = gr.Textbox(lines=8, 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("RUN SOTA ANALYSIS", variant="primary")
with gr.Column():
res_mbti = gr.Label(label="Personality (MBTI - Text Only)")
res_emo = gr.Label(label="Emotional State (Multimodal 28-Class Fusion)")
run_btn.click(
fn=analyze_track,
inputs=[audio_box, lyrics_box],
outputs=[res_mbti, res_emo]
)
if __name__ == "__main__":
demo.launch() |