anggars commited on
Commit
0df2078
·
verified ·
1 Parent(s): a90b44a

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +166 -187
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 librosa
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
- # -- CONFIGURATION --
21
- REPO_AUDIO = "anggars/neural-mathrock"
 
 
22
  REPO_TEXT_MBTI = "anggars/xlm-mbti"
23
- REPO_TEXT_EMO = "anggars/xlm-emotion"
24
- DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
 
 
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
- EMO_LABELS = ['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']
29
- VIBE_LABELS = ['aggressive', 'atmospheric', 'melancholic', 'technical']
30
- INTENSITY_LABELS = ['low', 'medium', 'high']
31
- TEMPO_LABELS = ['slow', 'moderate', 'fast']
32
 
33
- WAVLM_SR = 16000
34
- CLAP_SR = 48000
35
 
36
- emo2id = {e: i for i, e in enumerate(EMO_LABELS)}
 
 
 
37
 
38
- # -- HYBRID ARCHITECTURE ALIGNED WITH TRAINING EXECUTION --
39
- class AudioMathRockModel(nn.Module):
40
- def __init__(self):
 
 
41
  super().__init__()
42
- self.wavlm = WavLMModel.from_pretrained("microsoft/wavlm-base")
43
- self.clap = ClapAudioModel.from_pretrained("laion/clap-htsat-unfused")
44
-
45
- # 2D-CNN Mel Spectrogram Extractor Block
46
- self.mel_spectrogram = torchaudio.transforms.MelSpectrogram(sample_rate=WAVLM_SR, n_fft=1024, hop_length=512, n_mels=128)
47
- self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB(stype='power', top_db=80.0)
48
-
49
- self.cnn_extractor = nn.Sequential(
50
- nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
51
- nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2, 2),
52
- nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
53
- nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2, 2),
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, wavlm_values, clap_values):
80
- wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
81
- clap_feats = self.clap(clap_values).pooler_output
82
-
83
- with torch.no_grad():
84
- mel = self.mel_spectrogram(wavlm_values.float())
85
- mel_db = self.amplitude_to_db(mel)
86
-
87
- cnn_feats = F.gelu(self.cnn_extractor(mel_db.unsqueeze(1))).to(wavlm_feats.dtype)
88
- wlm_p = F.gelu(self.wlm_proj(wavlm_feats))
89
- clp_p = F.gelu(self.clp_proj(clap_feats))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
90
 
91
- fused = self.fusion(torch.cat([wlm_p, clp_p, cnn_feats], dim=-1))
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
- has_lyrics = lyrics_input is not None and len(str(lyrics_input).strip()) > 15
113
-
 
114
  if not has_audio and not has_lyrics:
115
- return {"Error": 1.0}, {"Error": 1.0}, {"Error": 1.0}, {"Error": 1.0}, {"Error": 1.0}
116
 
117
- res_mbti, res_emo, res_vibe, res_int, res_tmp = {}, {}, {}, {}, {}
 
118
 
119
- # --- TEXT ONLY PROCESSING ---
120
- if has_lyrics and not has_audio:
121
- t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
 
122
  with torch.no_grad():
123
- t_mbti_probs = F.softmax(text_mbti_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
124
- t_emo_probs = F.softmax(text_emo_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
125
-
126
- mbti_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
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
- res_emo = dict(sorted(emo_dict.items(), key=lambda x: x[1], reverse=True)[:3])
131
- return res_mbti, res_emo, {}, {}, {}
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
- wlm_chunks = [wlm_wave[i:i+wavlm_samples].numpy() for i in range(0, len(wlm_wave), wavlm_samples) if len(wlm_wave[i:i+wavlm_samples]) > WAVLM_SR]
145
- clp_chunks = [clp_wave[i:i+clap_samples].numpy() for i in range(0, len(clp_wave), clap_samples) if len(clp_wave[i:i+clap_samples]) > CLAP_SR]
146
-
147
- if not wlm_chunks:
148
- wlm_chunks = [F.pad(wlm_wave, (0, wavlm_samples - wlm_wave.shape[0])).numpy()]
149
- clp_chunks = [F.pad(clp_wave, (0, clap_samples - clp_wave.shape[0])).numpy()]
150
-
151
- el_list, vl_list, il_list, tl_list = [], [], [], []
152
-
153
- with torch.no_grad():
154
- for w_c, c_c in zip(wlm_chunks, clp_chunks):
155
- wlm_inputs = wavlm_extractor([w_c], sampling_rate=WAVLM_SR, return_tensors="pt")["input_values"].to(DEVICE)
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
- t_mbti_probs = F.softmax(text_mbti_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
191
- t_emo_probs = F.softmax(text_emo_model(**t_in).logits, dim=1).cpu().squeeze().numpy()
 
 
 
 
 
 
 
192
 
193
- m_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
194
- res_mbti = dict(sorted(m_dict.items(), key=lambda x: x[1], reverse=True)[:3])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
195
 
196
- # Weighted late fusion to prioritize text model for emotion taxonomy resolution
197
- for k in res_emo_raw.keys():
198
- idx = emo2id[k]
199
- res_emo_raw[k] = (res_emo_raw[k] * 0.15) + (float(t_emo_probs[idx]) * 0.85)
200
- res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
201
- else:
202
- res_mbti = {}
203
- res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
204
-
205
- return res_mbti, res_emo, res_vibe, res_int, res_tmp
206
-
207
- except Exception as e:
208
- print(f"PIPELINE ERROR: {str(e)}")
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 pure music audio and lyrics inferences.")
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 for hybrid analysis...")
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, res_vibe, res_int, res_tmp]
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__":