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