Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -1,235 +1,214 @@
|
|
|
|
|
|
|
|
| 1 |
import torch
|
| 2 |
import torch.nn as nn
|
| 3 |
import torch.nn.functional as F
|
| 4 |
-
import
|
| 5 |
import torchaudio
|
| 6 |
import gradio as gr
|
|
|
|
|
|
|
| 7 |
from transformers import (
|
|
|
|
|
|
|
|
|
|
| 8 |
XLMRobertaForSequenceClassification,
|
| 9 |
-
XLMRobertaTokenizer
|
| 10 |
-
WavLMModel,
|
| 11 |
-
AutoFeatureExtractor,
|
| 12 |
-
ClapAudioModel,
|
| 13 |
-
ClapProcessor
|
| 14 |
)
|
| 15 |
-
from huggingface_hub import hf_hub_download
|
| 16 |
import warnings
|
| 17 |
|
| 18 |
warnings.filterwarnings('ignore')
|
| 19 |
|
| 20 |
-
# --
|
| 21 |
-
|
|
|
|
|
|
|
| 22 |
REPO_TEXT_MBTI = "anggars/xlm-mbti"
|
| 23 |
-
REPO_TEXT_EMO
|
| 24 |
-
|
|
|
|
|
|
|
|
|
|
| 25 |
|
| 26 |
-
# -- GLOBAL LABELS --
|
| 27 |
MBTI_LABELS = sorted(["INTJ", "INTP", "ENTJ", "ENTP", "INFJ", "INFP", "ENFJ", "ENFP", "ISTJ", "ISFJ", "ESTJ", "ESFJ", "ISTP", "ISFP", "ESTP", "ESFP"])
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
INTENSITY_LABELS = ['low', 'medium', 'high']
|
| 31 |
-
TEMPO_LABELS = ['slow', 'moderate', 'fast']
|
| 32 |
|
| 33 |
-
|
| 34 |
-
CLAP_SR = 48000
|
| 35 |
|
| 36 |
-
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
-
# --
|
| 39 |
-
|
| 40 |
-
|
|
|
|
|
|
|
| 41 |
super().__init__()
|
| 42 |
-
self.
|
| 43 |
-
self.
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
nn.
|
| 51 |
-
nn.
|
| 52 |
-
nn.
|
| 53 |
-
nn.
|
| 54 |
-
nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
|
| 55 |
-
nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)),
|
| 56 |
-
nn.Flatten(),
|
| 57 |
-
nn.Linear(128, 512)
|
| 58 |
)
|
| 59 |
-
|
| 60 |
-
self.wlm_proj = nn.Linear(768, 512)
|
| 61 |
-
self.clp_proj = nn.Linear(768, 512)
|
| 62 |
-
|
| 63 |
-
# Expanded dimension to 1536 to hold CNN features
|
| 64 |
-
self.fusion = nn.Sequential(
|
| 65 |
-
nn.Linear(1536, 512),
|
| 66 |
-
nn.LayerNorm(512),
|
| 67 |
-
nn.Tanh(),
|
| 68 |
-
nn.Dropout(0.3)
|
| 69 |
-
)
|
| 70 |
-
|
| 71 |
-
self.emo_head = nn.Sequential(
|
| 72 |
-
nn.Linear(512, 256), nn.GELU(), nn.Dropout(0.2),
|
| 73 |
-
nn.Linear(256, len(EMO_LABELS))
|
| 74 |
-
)
|
| 75 |
-
self.vibe_head = nn.Linear(512, len(VIBE_LABELS))
|
| 76 |
-
self.int_head = nn.Linear(512, len(INTENSITY_LABELS))
|
| 77 |
-
self.tmp_head = nn.Linear(512, len(TEMPO_LABELS))
|
| 78 |
|
| 79 |
-
def forward(self,
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
|
| 91 |
-
|
| 92 |
-
return self.emo_head(fused), self.vibe_head(fused), self.int_head(fused), self.tmp_head(fused)
|
| 93 |
-
|
| 94 |
-
# -- MODEL INITIALIZATION --
|
| 95 |
-
print("Fetching Model Weights...")
|
| 96 |
-
ckpt_path = hf_hub_download(repo_id=REPO_AUDIO, filename="model.pth")
|
| 97 |
-
ckpt = torch.load(ckpt_path, map_location=DEVICE, weights_only=False)
|
| 98 |
|
| 99 |
-
audio_model = AudioMathRockModel().to(DEVICE)
|
| 100 |
-
audio_model.load_state_dict(ckpt['model_state_dict'], strict=True)
|
| 101 |
-
audio_model.eval()
|
| 102 |
-
|
| 103 |
-
wavlm_extractor = AutoFeatureExtractor.from_pretrained("microsoft/wavlm-base")
|
| 104 |
-
clap_processor = ClapProcessor.from_pretrained("laion/clap-htsat-unfused")
|
| 105 |
-
tokenizer = XLMRobertaTokenizer.from_pretrained(REPO_TEXT_MBTI)
|
| 106 |
-
text_mbti_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_MBTI).to(DEVICE).eval()
|
| 107 |
-
text_emo_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_EMO).to(DEVICE).eval()
|
| 108 |
-
|
| 109 |
-
# -- ANALYSIS ENGINE --
|
| 110 |
def analyze_track(audio_path, lyrics_input):
|
| 111 |
has_audio = audio_path is not None
|
| 112 |
-
|
| 113 |
-
|
|
|
|
| 114 |
if not has_audio and not has_lyrics:
|
| 115 |
-
return {"Error": 1.0}, {"Error": 1.0}
|
| 116 |
|
| 117 |
-
res_mbti
|
|
|
|
| 118 |
|
| 119 |
-
#
|
| 120 |
-
|
| 121 |
-
|
|
|
|
| 122 |
with torch.no_grad():
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
mbti_dict = {MBTI_LABELS[i]: float(
|
| 127 |
-
emo_dict = {EMO_LABELS[i]: float(t_emo_probs[i]) for i in range(len(EMO_LABELS))}
|
| 128 |
-
|
| 129 |
res_mbti = dict(sorted(mbti_dict.items(), key=lambda x: x[1], reverse=True)[:3])
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
# --- AUDIO PROCESSING ---
|
| 134 |
-
try:
|
| 135 |
-
wav_orig, orig_sr = librosa.load(audio_path, sr=None, mono=True)
|
| 136 |
-
waveform = torch.tensor(wav_orig).unsqueeze(0)
|
| 137 |
-
|
| 138 |
-
wlm_wave = torchaudio.functional.resample(waveform, orig_sr, WAVLM_SR).squeeze(0)
|
| 139 |
-
clp_wave = torchaudio.functional.resample(waveform, orig_sr, CLAP_SR).squeeze(0)
|
| 140 |
-
|
| 141 |
-
wavlm_samples = WAVLM_SR * 15
|
| 142 |
-
clap_samples = CLAP_SR * 15
|
| 143 |
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
clp_inputs = clap_processor(audio=[c_c], sampling_rate=CLAP_SR, return_tensors="pt")["input_features"].to(DEVICE)
|
| 157 |
-
|
| 158 |
-
el, vl, il, tl = audio_model(wlm_inputs, clp_inputs)
|
| 159 |
-
el_list.append(el.cpu())
|
| 160 |
-
vl_list.append(vl.cpu())
|
| 161 |
-
il_list.append(il.cpu())
|
| 162 |
-
tl_list.append(tl.cpu())
|
| 163 |
-
|
| 164 |
-
raw_el = torch.cat(el_list, dim=0).mean(dim=0, keepdim=True)
|
| 165 |
-
raw_vl = torch.cat(vl_list, dim=0).mean(dim=0, keepdim=True)
|
| 166 |
-
raw_il = torch.cat(il_list, dim=0).mean(dim=0, keepdim=True)
|
| 167 |
-
raw_tl = torch.cat(tl_list, dim=0).mean(dim=0, keepdim=True)
|
| 168 |
-
|
| 169 |
-
# Enforce highly competitive soft-scaling for audio emotion logits
|
| 170 |
-
scaled_el = raw_el / raw_el.std(dim=-1, keepdim=True).clamp(min=1e-6)
|
| 171 |
-
probs_el_audio = F.softmax(scaled_el / 0.3, dim=-1).squeeze().numpy()
|
| 172 |
-
|
| 173 |
-
probs_vl = F.softmax(raw_vl, dim=-1).squeeze().numpy()
|
| 174 |
-
probs_il = F.softmax(raw_il, dim=-1).squeeze().numpy()
|
| 175 |
-
probs_tl = F.softmax(raw_tl, dim=-1).squeeze().numpy()
|
| 176 |
-
|
| 177 |
-
res_emo_raw = {EMO_LABELS[i]: float(probs_el_audio[i]) for i in range(len(EMO_LABELS))}
|
| 178 |
-
res_vibe = {VIBE_LABELS[i]: float(probs_vl[i]) for i in range(len(VIBE_LABELS))}
|
| 179 |
-
res_int = {INTENSITY_LABELS[i]: float(probs_il[i]) for i in range(len(INTENSITY_LABELS))}
|
| 180 |
-
res_tmp = {TEMPO_LABELS[i]: float(probs_tl[i]) for i in range(len(TEMPO_LABELS))}
|
| 181 |
-
|
| 182 |
-
res_vibe = dict(sorted(res_vibe.items(), key=lambda x: x[1], reverse=True)[:3])
|
| 183 |
-
res_int = dict(sorted(res_int.items(), key=lambda x: x[1], reverse=True)[:3])
|
| 184 |
-
res_tmp = dict(sorted(res_tmp.items(), key=lambda x: x[1], reverse=True)[:3])
|
| 185 |
-
|
| 186 |
-
# --- MULTIMODAL LATE FUSION ENGAGEMENT ---
|
| 187 |
-
if has_lyrics:
|
| 188 |
-
t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
|
| 189 |
with torch.no_grad():
|
| 190 |
-
|
| 191 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 192 |
|
| 193 |
-
|
| 194 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 195 |
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
return {"System Error": 1.0}, {"System Error": 1.0}, {"System Error": 1.0}, {"System Error": 1.0}, {"System Error": 1.0}
|
| 210 |
-
|
| 211 |
-
# -- INTERFACE BLOCK --
|
| 212 |
with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
|
| 213 |
gr.Markdown("# Neural Math Rock Multimodal Analysis")
|
| 214 |
-
gr.Markdown("Identify personality and emotional states from
|
| 215 |
|
| 216 |
with gr.Row():
|
| 217 |
with gr.Column():
|
| 218 |
-
audio_box = gr.Audio(type="filepath", label="Audio Source")
|
| 219 |
-
lyrics_box = gr.Textbox(lines=8, label="Lyrics Source", placeholder="Paste lyrics here
|
| 220 |
-
run_btn = gr.Button("RUN ANALYSIS", variant="primary")
|
| 221 |
|
| 222 |
with gr.Column():
|
| 223 |
-
res_mbti = gr.Label(label="Personality (MBTI)")
|
| 224 |
-
res_emo = gr.Label(label="Emotional State")
|
| 225 |
-
res_vibe = gr.Label(label="Acoustic Vibe")
|
| 226 |
-
res_int = gr.Label(label="Intensity Level")
|
| 227 |
-
res_tmp = gr.Label(label="Tempo Classification")
|
| 228 |
|
| 229 |
run_btn.click(
|
| 230 |
fn=analyze_track,
|
| 231 |
inputs=[audio_box, lyrics_box],
|
| 232 |
-
outputs=[res_mbti, res_emo
|
| 233 |
)
|
| 234 |
|
| 235 |
if __name__ == "__main__":
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import gc
|
| 3 |
import torch
|
| 4 |
import torch.nn as nn
|
| 5 |
import torch.nn.functional as F
|
| 6 |
+
import soundfile as sf
|
| 7 |
import torchaudio
|
| 8 |
import gradio as gr
|
| 9 |
+
import numpy as np
|
| 10 |
+
from huggingface_hub import hf_hub_download
|
| 11 |
from transformers import (
|
| 12 |
+
AutoModel,
|
| 13 |
+
AutoFeatureExtractor,
|
| 14 |
+
AutoTokenizer,
|
| 15 |
XLMRobertaForSequenceClassification,
|
| 16 |
+
XLMRobertaTokenizer
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
)
|
|
|
|
| 18 |
import warnings
|
| 19 |
|
| 20 |
warnings.filterwarnings('ignore')
|
| 21 |
|
| 22 |
+
# ------------------------------------------------------------------------------
|
| 23 |
+
# CONFIGURATION & LABEL MAPPING
|
| 24 |
+
# ------------------------------------------------------------------------------
|
| 25 |
+
REPO_MAIN = "anggars/neural-mathrock"
|
| 26 |
REPO_TEXT_MBTI = "anggars/xlm-mbti"
|
| 27 |
+
REPO_TEXT_EMO = "anggars/xlm-emotion"
|
| 28 |
+
MODEL_FILE = "model.pth"
|
| 29 |
+
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 30 |
+
CROP_SEC = 5
|
| 31 |
+
SR_TARGET = 24000
|
| 32 |
|
|
|
|
| 33 |
MBTI_LABELS = sorted(["INTJ", "INTP", "ENTJ", "ENTP", "INFJ", "INFP", "ENFJ", "ENFP", "ISTJ", "ISFJ", "ESTJ", "ESFJ", "ISTP", "ISFP", "ESTP", "ESFP"])
|
| 34 |
+
EMOTION_CLASSES_AUDIO = ["anger", "disgust", "fear", "joy", "neutral", "sadness", "surprise"]
|
| 35 |
+
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']
|
|
|
|
|
|
|
| 36 |
|
| 37 |
+
text_emo2id = {e: i for i, e in enumerate(EMOTION_CLASSES_TEXT)}
|
|
|
|
| 38 |
|
| 39 |
+
# Calibration Weights derived from SOTA Report 3
|
| 40 |
+
RAW_CE_WEIGHTS = [0.563, 16.326, 4.525, 0.968, 2.704, 0.473, 0.702]
|
| 41 |
+
SOFT_WEIGHTS = torch.sqrt(torch.tensor(RAW_CE_WEIGHTS, dtype=torch.float32)).to(DEVICE)
|
| 42 |
+
TEMPERATURE = 2.0
|
| 43 |
|
| 44 |
+
# ------------------------------------------------------------------------------
|
| 45 |
+
# NEURAL ARCHITECTURE (7-CLASS AUDIO FUSION)
|
| 46 |
+
# ------------------------------------------------------------------------------
|
| 47 |
+
class MultimodalFusionClassifier(nn.Module):
|
| 48 |
+
def __init__(self, num_classes=7):
|
| 49 |
super().__init__()
|
| 50 |
+
self.audio_model = AutoModel.from_pretrained("m-a-p/MERT-v1-330M", trust_remote_code=True)
|
| 51 |
+
self.text_model = AutoModel.from_pretrained("FacebookAI/roberta-base")
|
| 52 |
+
self.fusion_head = nn.Sequential(
|
| 53 |
+
nn.Linear(1024 + 768, 512),
|
| 54 |
+
nn.LayerNorm(512),
|
| 55 |
+
nn.GELU(),
|
| 56 |
+
nn.Dropout(0.5),
|
| 57 |
+
nn.Linear(512, 256),
|
| 58 |
+
nn.LayerNorm(256),
|
| 59 |
+
nn.GELU(),
|
| 60 |
+
nn.Dropout(0.4),
|
| 61 |
+
nn.Linear(256, num_classes)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
|
| 64 |
+
def forward(self, audio_vals, text_ids, text_mask):
|
| 65 |
+
audio_feats = self.audio_model(audio_vals, output_hidden_states=True).last_hidden_state.mean(dim=1)
|
| 66 |
+
text_feats = self.text_model(input_ids=text_ids, attention_mask=text_mask).last_hidden_state.mean(dim=1)
|
| 67 |
+
return self.fusion_head(torch.cat((audio_feats, text_feats), dim=-1))
|
| 68 |
+
|
| 69 |
+
# ------------------------------------------------------------------------------
|
| 70 |
+
# MODEL INITIALIZATION
|
| 71 |
+
# ------------------------------------------------------------------------------
|
| 72 |
+
print("[SYSTEM] Booting Models into Memory...")
|
| 73 |
+
mert_ext = AutoFeatureExtractor.from_pretrained("m-a-p/MERT-v1-330M", trust_remote_code=True)
|
| 74 |
+
roberta_tok = AutoTokenizer.from_pretrained("FacebookAI/roberta-base")
|
| 75 |
+
|
| 76 |
+
fusion_model = MultimodalFusionClassifier(num_classes=len(EMOTION_CLASSES_AUDIO)).to(DEVICE)
|
| 77 |
+
ckpt_path = hf_hub_download(repo_id=REPO_MAIN, filename=MODEL_FILE)
|
| 78 |
+
fusion_model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE)['model_state_dict'])
|
| 79 |
+
fusion_model.eval()
|
| 80 |
+
|
| 81 |
+
# Load Text-Only Models (Isolated)
|
| 82 |
+
mbti_tok = XLMRobertaTokenizer.from_pretrained(REPO_TEXT_MBTI)
|
| 83 |
+
mbti_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_MBTI).to(DEVICE)
|
| 84 |
+
mbti_model.eval()
|
| 85 |
+
|
| 86 |
+
# Load High-Resolution 28-Class Text Emotion Model
|
| 87 |
+
emo28_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_EMO).to(DEVICE)
|
| 88 |
+
emo28_model.eval()
|
| 89 |
+
|
| 90 |
+
# ------------------------------------------------------------------------------
|
| 91 |
+
# INFERENCE ENGINE
|
| 92 |
+
# ------------------------------------------------------------------------------
|
| 93 |
+
def process_audio_chunks(audio_path):
|
| 94 |
+
waveform_np, orig_sr = sf.read(audio_path)
|
| 95 |
+
if len(waveform_np.shape) > 1:
|
| 96 |
+
waveform_np = waveform_np.mean(axis=1)
|
| 97 |
+
|
| 98 |
+
waveform = torch.tensor(waveform_np, dtype=torch.float32).unsqueeze(0)
|
| 99 |
+
if orig_sr != SR_TARGET:
|
| 100 |
+
waveform = torchaudio.functional.resample(waveform, orig_sr, SR_TARGET)
|
| 101 |
+
|
| 102 |
+
waveform = waveform.squeeze(0)
|
| 103 |
+
chunk_samples = SR_TARGET * CROP_SEC
|
| 104 |
+
|
| 105 |
+
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]
|
| 106 |
+
if not chunks:
|
| 107 |
+
chunks = [F.pad(waveform, (0, chunk_samples - waveform.shape[0])).numpy()]
|
| 108 |
|
| 109 |
+
return chunks[:10]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
def analyze_track(audio_path, lyrics_input):
|
| 112 |
has_audio = audio_path is not None
|
| 113 |
+
safe_lyrics = str(lyrics_input).strip() if lyrics_input else ""
|
| 114 |
+
has_lyrics = len(safe_lyrics) > 10
|
| 115 |
+
|
| 116 |
if not has_audio and not has_lyrics:
|
| 117 |
+
return {"Error": 1.0}, {"Error": 1.0}
|
| 118 |
|
| 119 |
+
res_mbti = {}
|
| 120 |
+
res_emo_final = {}
|
| 121 |
|
| 122 |
+
# 1. ISOLATED MBTI & 28-CLASS TEXT INFERENCE
|
| 123 |
+
t_emo28_probs = None
|
| 124 |
+
if has_lyrics:
|
| 125 |
+
t_in = mbti_tok(safe_lyrics, truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
|
| 126 |
with torch.no_grad():
|
| 127 |
+
mbti_probs = F.softmax(mbti_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
|
| 128 |
+
t_emo28_probs = F.softmax(emo28_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
|
| 129 |
+
|
| 130 |
+
mbti_dict = {MBTI_LABELS[i]: float(mbti_probs[i]) for i in range(len(MBTI_LABELS))}
|
|
|
|
|
|
|
| 131 |
res_mbti = dict(sorted(mbti_dict.items(), key=lambda x: x[1], reverse=True)[:3])
|
| 132 |
+
else:
|
| 133 |
+
res_mbti = {"Requires lyrics for MBTI": 1.0}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
|
| 135 |
+
# 2. AUDIO FUSION INFERENCE
|
| 136 |
+
if has_audio:
|
| 137 |
+
try:
|
| 138 |
+
audio_chunks = process_audio_chunks(audio_path)
|
| 139 |
+
|
| 140 |
+
t_inputs = roberta_tok([safe_lyrics] * len(audio_chunks), padding=True, truncation=True, max_length=128, return_tensors="pt")
|
| 141 |
+
t_ids = t_inputs["input_ids"].to(DEVICE)
|
| 142 |
+
t_mask = t_inputs["attention_mask"].to(DEVICE)
|
| 143 |
+
|
| 144 |
+
a_inputs = mert_ext(audio_chunks, sampling_rate=SR_TARGET, return_tensors="pt", padding="max_length", truncation=True, max_length=SR_TARGET * CROP_SEC)
|
| 145 |
+
a_vals = a_inputs["input_values"].to(DEVICE)
|
| 146 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 147 |
with torch.no_grad():
|
| 148 |
+
logits = fusion_model(a_vals, t_ids, t_mask)
|
| 149 |
+
avg_logits = logits.mean(dim=0)
|
| 150 |
+
|
| 151 |
+
scaled_logits = avg_logits / TEMPERATURE
|
| 152 |
+
probs = F.softmax(scaled_logits, dim=-1)
|
| 153 |
+
adjusted_probs = probs * SOFT_WEIGHTS
|
| 154 |
+
audio_7_probs = (adjusted_probs / adjusted_probs.sum()).cpu().numpy()
|
| 155 |
+
|
| 156 |
+
audio_emo_dict = {EMOTION_CLASSES_AUDIO[i]: float(audio_7_probs[i]) for i in range(len(EMOTION_CLASSES_AUDIO))}
|
| 157 |
|
| 158 |
+
# --- TRUE MULTIMODAL LATE FUSION ENGAGEMENT ---
|
| 159 |
+
if has_lyrics:
|
| 160 |
+
# Bikin dictionary kosong untuk 28 emosi
|
| 161 |
+
fusion_28_dict = {label: 0.0 for label in EMOTION_CLASSES_TEXT}
|
| 162 |
+
|
| 163 |
+
# Masukin probabilitas murni dari Teks (Bobot 60%)
|
| 164 |
+
for i, label in enumerate(EMOTION_CLASSES_TEXT):
|
| 165 |
+
fusion_28_dict[label] += float(t_emo28_probs[i]) * 0.60
|
| 166 |
+
|
| 167 |
+
# Suntikin probabilitas dari Audio 7 Kelas ke 28 Kelas Teks (Bobot 40%)
|
| 168 |
+
for label in EMOTION_CLASSES_AUDIO:
|
| 169 |
+
if label in fusion_28_dict:
|
| 170 |
+
fusion_28_dict[label] += audio_emo_dict[label] * 0.40
|
| 171 |
+
|
| 172 |
+
# Normalisasi ulang biar totalnya 1.0
|
| 173 |
+
total_prob = sum(fusion_28_dict.values())
|
| 174 |
+
fusion_28_dict = {k: v / total_prob for k, v in fusion_28_dict.items()}
|
| 175 |
+
|
| 176 |
+
res_emo_final = dict(sorted(fusion_28_dict.items(), key=lambda x: x[1], reverse=True)[:5])
|
| 177 |
+
else:
|
| 178 |
+
# Unimodal Audio (Hanya ngeluarin 7 Kelas)
|
| 179 |
+
res_emo_final = dict(sorted(audio_emo_dict.items(), key=lambda x: x[1], reverse=True)[:4])
|
| 180 |
|
| 181 |
+
except Exception as e:
|
| 182 |
+
res_emo_final = {f"Audio Error: {str(e)}": 1.0}
|
| 183 |
+
|
| 184 |
+
# 3. TEXT-ONLY FALLBACK (Kalau audio gak dimasukin)
|
| 185 |
+
elif has_lyrics:
|
| 186 |
+
emo_dict = {EMOTION_CLASSES_TEXT[i]: float(t_emo28_probs[i]) for i in range(len(EMOTION_CLASSES_TEXT))}
|
| 187 |
+
res_emo_final = dict(sorted(emo_dict.items(), key=lambda x: x[1], reverse=True)[:5])
|
| 188 |
+
|
| 189 |
+
return res_mbti, res_emo_final
|
| 190 |
+
|
| 191 |
+
# ------------------------------------------------------------------------------
|
| 192 |
+
# GRADIO INTERFACE
|
| 193 |
+
# ------------------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
| 194 |
with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
|
| 195 |
gr.Markdown("# Neural Math Rock Multimodal Analysis")
|
| 196 |
+
gr.Markdown("Identify personality (Text) and emotional states (MERT+RoBERTa Multimodal) from Math Rock & Midwest Emo tracks.")
|
| 197 |
|
| 198 |
with gr.Row():
|
| 199 |
with gr.Column():
|
| 200 |
+
audio_box = gr.Audio(type="filepath", label="Audio Source (.wav / .mp3)")
|
| 201 |
+
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.)")
|
| 202 |
+
run_btn = gr.Button("RUN SOTA ANALYSIS", variant="primary")
|
| 203 |
|
| 204 |
with gr.Column():
|
| 205 |
+
res_mbti = gr.Label(label="Personality (MBTI - Text Only)")
|
| 206 |
+
res_emo = gr.Label(label="Emotional State (Multimodal 28-Class Fusion)")
|
|
|
|
|
|
|
|
|
|
| 207 |
|
| 208 |
run_btn.click(
|
| 209 |
fn=analyze_track,
|
| 210 |
inputs=[audio_box, lyrics_box],
|
| 211 |
+
outputs=[res_mbti, res_emo]
|
| 212 |
)
|
| 213 |
|
| 214 |
if __name__ == "__main__":
|