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()