anggars commited on
Commit
78bd59e
·
verified ·
1 Parent(s): beaad7b

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +96 -100
app.py CHANGED
@@ -1,9 +1,10 @@
1
  import torch
2
  import torch.nn as nn
 
3
  import librosa
4
  import numpy as np
5
  import gradio as gr
6
- from transformers import XLMRobertaModel, XLMRobertaTokenizer, WavLMModel
7
  from huggingface_hub import hf_hub_download
8
  import yt_dlp
9
  import os
@@ -16,62 +17,46 @@ import warnings
16
  warnings.filterwarnings('ignore')
17
 
18
  # -- CONFIGURATION --
19
- REPO_ID = "anggars/neural-mathrock"
 
 
20
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
21
- TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
22
  GENIUS_TOKEN = os.environ.get("GENIUS_TOKEN", "z2XGBWXalGUtAdC1qxxXBxUnK1ZuoHPkCu5eP9q-fed-DW1uCJ3NSFpHemk3Unmg")
23
 
24
- print("Fetching model weights...")
25
- model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
26
- ckpt = torch.load(model_path, map_location=DEVICE, weights_only=False)
 
 
27
 
28
- le_mbti, le_emotion, le_vibe, le_intensity, le_tempo = ckpt['le_mbti'], ckpt['le_emotion'], ckpt['le_vibe'], ckpt['le_intensity'], ckpt['le_tempo']
29
-
30
- # -- ARCHITECTURE --
31
- class HybridMultimodalModel(nn.Module):
32
  def __init__(self):
33
  super().__init__()
34
- self.text_model = XLMRobertaModel.from_pretrained('anggars/xlm-mbti')
35
- self.audio_model = WavLMModel.from_pretrained('microsoft/wavlm-base')
36
- self.audio_proj = nn.Linear(768, 256)
37
- self.text_gate = nn.Sequential(nn.Linear(768, 768), nn.Sigmoid())
38
- self.audio_gate = nn.Sequential(nn.Linear(256, 256), nn.Sigmoid())
39
- self.fusion = nn.Sequential(
40
- nn.Linear(1024, 512),
41
- nn.BatchNorm1d(512),
42
- nn.ReLU(),
43
- nn.Dropout(0.4),
44
- )
45
- self.head_mbti = nn.Linear(512, len(le_mbti.classes_))
46
- self.head_emotion = nn.Linear(512, len(le_emotion.classes_))
47
- self.head_vibe = nn.Linear(512, len(le_vibe.classes_))
48
- self.head_intensity = nn.Linear(512, len(le_intensity.classes_))
49
- self.head_tempo = nn.Linear(512, len(le_tempo.classes_))
50
-
51
- def forward(self, input_ids, attention_mask, audio_values, text_missing=False):
52
- text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
53
- text_feat = text_out.pooler_output
54
-
55
- # Jika instrumental, matikan representasi teks
56
- if text_missing:
57
- text_feat = torch.zeros_like(text_feat)
58
-
59
- audio_out = self.audio_model(audio_values).last_hidden_state
60
- audio_feat = self.audio_proj(audio_out.mean(dim=1))
61
-
62
- # Gating mechanism
63
- gated_text = text_feat * self.text_gate(text_feat)
64
- gated_audio = audio_feat * self.audio_gate(audio_feat)
65
-
66
- # Fusion
67
- fused = self.fusion(torch.cat([gated_text, gated_audio], dim=-1))
68
-
69
- return {col: getattr(self, f'head_{col}')(fused) for col in TARGET_COLS}
70
 
71
- model = HybridMultimodalModel().to(DEVICE)
72
- model.load_state_dict(ckpt['model_state'], strict=False)
73
- model.eval()
74
- tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')
 
 
 
 
 
 
 
 
 
 
 
 
75
 
76
  # -- UTILITIES --
77
  def search_and_fetch(query):
@@ -80,7 +65,6 @@ def search_and_fetch(query):
80
  search = VideosSearch(query, limit=1)
81
  res = search.result()
82
  if not res['result']: return None, "No results."
83
-
84
  video_url = res['result'][0]['link']
85
  temp_fn = "temp_audio_file"
86
  ydl_opts = {
@@ -107,69 +91,81 @@ def search_and_fetch(query):
107
  def analyze_track(audio_path, lyrics_input):
108
  if not audio_path: return [{"Error": "No audio"}] * 5
109
  try:
110
- is_inst = not lyrics_input or str(lyrics_input).strip() == ""
111
- text = str(lyrics_input).strip() if not is_inst else "[INSTRUMENTAL]"
112
- enc = tokenizer(text, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)
113
-
114
  wav, sr = librosa.load(audio_path, sr=16000)
115
- tempo_bpm, _ = librosa.beat.beat_track(y=wav, sr=sr)
116
-
117
  chunk_len = 16000 * 15
118
  chunks = [wav[i:i + chunk_len] for i in range(0, len(wav), chunk_len) if len(wav[i:i+chunk_len]) >= 16000]
119
- all_logits = {col: [] for col in TARGET_COLS}
 
120
 
121
  with torch.no_grad():
122
  for chunk in chunks:
123
  if len(chunk) < chunk_len: chunk = np.pad(chunk, (0, chunk_len - len(chunk)))
124
- audio_t = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
125
- out = model(enc['input_ids'], enc['attention_mask'], audio_t, text_missing=is_inst)
126
- for col in TARGET_COLS: all_logits[col].append(out[col][0])
127
-
128
- final_results = []
129
- encoders = {'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe, 'intensity': le_intensity, 'tempo': le_tempo}
130
-
131
- for col in TARGET_COLS:
132
- logits_stack = torch.stack(all_logits[col])
133
- # Pake mean tanpa weight manual biar model audio kerja murni
134
- final_logits = logits_stack.mean(dim=0) / 0.7
135
-
136
- if col == 'tempo':
137
- t_cls = list(le_tempo.classes_)
138
- try:
139
- if tempo_bpm > 125: final_logits[t_cls.index('Fast')] += 3.0
140
- except: pass
141
-
142
- probs = torch.nn.functional.softmax(final_logits, dim=0).cpu().numpy()
143
- classes = encoders[col].classes_
144
- res_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes))}
145
- final_results.append(dict(sorted(res_dict.items(), key=lambda x: x[1], reverse=True)[:3]))
146
 
147
- return final_results
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
148
  except Exception as e: return [{"Error": str(e)}] * 5
149
 
150
- # -- INTERFACE --
151
- with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
152
- gr.Markdown("# Neural Math Rock Multimodal Analysis")
153
- gr.Markdown("Identify personality and emotional states from music audio and lyrics.")
154
 
155
  with gr.Row():
156
  with gr.Column():
157
- search_box = gr.Textbox(label="YouTube Search (Artist - Song Title)", placeholder="Enter song name...")
158
- fetch_btn = gr.Button("FETCH AUDIO AND LYRICS", variant="secondary")
159
- gr.HTML("<hr>")
160
- audio_box = gr.Audio(type="filepath", label="Audio Source")
161
- lyrics_box = gr.Textbox(lines=6, label="Lyrics Source", placeholder="Lyrics...")
162
- run_btn = gr.Button("RUN ANALYSIS", variant="primary")
163
-
164
  with gr.Column():
165
- res_mbti = gr.Label(label="Personality (MBTI)")
166
- res_emo = gr.Label(label="Emotional State")
167
- res_vibe = gr.Label(label="Acoustic Vibe")
168
- res_int = gr.Label(label="Intensity Level")
169
- res_tmp = gr.Label(label="Tempo Classification")
170
-
171
- fetch_btn.click(fn=search_and_fetch, inputs=[search_box], outputs=[audio_box, lyrics_box])
172
- run_btn.click(fn=analyze_track, inputs=[audio_box, lyrics_box], outputs=[res_mbti, res_emo, res_vibe, res_int, res_tmp])
173
 
174
  if __name__ == "__main__":
175
  demo.launch()
 
1
  import torch
2
  import torch.nn as nn
3
+ import torch.nn.functional as F
4
  import librosa
5
  import numpy as np
6
  import gradio as gr
7
+ from transformers import XLMRobertaForSequenceClassification, XLMRobertaTokenizer, WavLMModel
8
  from huggingface_hub import hf_hub_download
9
  import yt_dlp
10
  import os
 
17
  warnings.filterwarnings('ignore')
18
 
19
  # -- CONFIGURATION --
20
+ REPO_AUDIO = "anggars/neural-mathrock"
21
+ REPO_TEXT_MBTI = "anggars/xlm-mbti"
22
+ REPO_TEXT_EMO = "anggars/xlm-emotion"
23
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
24
  GENIUS_TOKEN = os.environ.get("GENIUS_TOKEN", "z2XGBWXalGUtAdC1qxxXBxUnK1ZuoHPkCu5eP9q-fed-DW1uCJ3NSFpHemk3Unmg")
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 = 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'])
29
+ VIBE_LABELS = sorted(["Aggressive", "Atmospheric", "Melancholic", "Technical"])
30
+ INTENSITY_LABELS = sorted(["High", "Low", "Medium"])
31
 
32
+ # -- AUDIO ARCHITECTURE (neural-mathrock) --
33
+ class NeuralMathRockAudio(nn.Module):
 
 
34
  def __init__(self):
35
  super().__init__()
36
+ self.wavlm = WavLMModel.from_pretrained("microsoft/wavlm-base-plus")
37
+ self.proj = nn.Sequential(nn.Linear(768, 256), nn.BatchNorm1d(256), nn.GELU(), nn.Dropout(0.4))
38
+ self.mbti_head = nn.Linear(256, len(MBTI_LABELS))
39
+ self.emo_head = nn.Linear(256, len(EMO_LABELS))
40
+ self.vibe_head = nn.Linear(256, len(VIBE_LABELS))
41
+ self.intensity_head = nn.Linear(256, len(INTENSITY_LABELS))
42
+ self.tempo_head = nn.Linear(256, 1)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
43
 
44
+ def forward(self, iv, aam):
45
+ h = self.wavlm(iv, attention_mask=aam).last_hidden_state.mean(dim=1)
46
+ f = self.proj(h)
47
+ return self.mbti_head(f), self.emo_head(f), self.vibe_head(f), self.intensity_head(f), self.tempo_head(f).squeeze(-1)
48
+
49
+ # -- MODEL INITIALIZATION --
50
+ print("Initializing Hybrid Ensemble Models...")
51
+ ckpt_path = hf_hub_download(repo_id=REPO_AUDIO, filename="model.pt")
52
+ ckpt = torch.load(ckpt_path, map_location=DEVICE, weights_only=False)
53
+ audio_model = NeuralMathRockAudio().to(DEVICE)
54
+ audio_model.load_state_dict(ckpt['model'], strict=False)
55
+ audio_model.eval()
56
+
57
+ tokenizer = XLMRobertaTokenizer.from_pretrained(REPO_TEXT_MBTI)
58
+ text_mbti_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_MBTI).to(DEVICE).eval()
59
+ text_emo_model = XLMRobertaForSequenceClassification.from_pretrained(REPO_TEXT_EMO).to(DEVICE).eval()
60
 
61
  # -- UTILITIES --
62
  def search_and_fetch(query):
 
65
  search = VideosSearch(query, limit=1)
66
  res = search.result()
67
  if not res['result']: return None, "No results."
 
68
  video_url = res['result'][0]['link']
69
  temp_fn = "temp_audio_file"
70
  ydl_opts = {
 
91
  def analyze_track(audio_path, lyrics_input):
92
  if not audio_path: return [{"Error": "No audio"}] * 5
93
  try:
 
 
 
 
94
  wav, sr = librosa.load(audio_path, sr=16000)
 
 
95
  chunk_len = 16000 * 15
96
  chunks = [wav[i:i + chunk_len] for i in range(0, len(wav), chunk_len) if len(wav[i:i+chunk_len]) >= 16000]
97
+
98
+ a_mbti, a_emo, a_vibe, a_int, a_tmp = [], [], [], [], []
99
 
100
  with torch.no_grad():
101
  for chunk in chunks:
102
  if len(chunk) < chunk_len: chunk = np.pad(chunk, (0, chunk_len - len(chunk)))
103
+ iv = torch.tensor(chunk).unsqueeze(0).to(DEVICE)
104
+ aam = torch.ones_like(iv).to(DEVICE)
105
+
106
+ m, e, v, it, t = audio_model(iv, aam)
107
+ a_mbti.append(F.softmax(m, dim=1))
108
+ a_emo.append(F.softmax(e, dim=1))
109
+ a_vibe.append(F.softmax(v, dim=1))
110
+ a_int.append(F.softmax(it, dim=1))
111
+ a_tmp.append(t * 200.0) # Denormalize tempo
112
+
113
+ avg_a_mbti = torch.stack(a_mbti).mean(dim=0)
114
+ avg_a_emo = torch.stack(a_emo).mean(dim=0)
115
+ avg_a_vibe = torch.stack(a_vibe).mean(dim=0)
116
+ avg_a_int = torch.stack(a_int).mean(dim=0)
117
+ avg_tempo_bpm = torch.stack(a_tmp).mean().item()
118
+
119
+ has_lyrics = lyrics_input and len(str(lyrics_input).strip()) > 15
120
+ if has_lyrics:
121
+ t_inputs = 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_inputs).logits, dim=1)
124
+ t_emo_probs = F.softmax(text_emo_model(**t_inputs).logits, dim=1)
125
 
126
+ final_mbti_probs = (avg_a_mbti * 0.6) + (t_mbti_probs * 0.4)
127
+ final_emo_probs = (avg_a_emo * 0.6) + (t_emo_probs * 0.4)
128
+ else:
129
+ final_mbti_probs = avg_a_mbti
130
+ final_emo_probs = avg_a_emo
131
+
132
+ def process_probs(probs, labels):
133
+ p = probs.cpu().squeeze().numpy()
134
+ res = {labels[i]: float(p[i]) for i in range(len(labels))}
135
+ return dict(sorted(res.items(), key=lambda x: x[1], reverse=True)[:3])
136
+
137
+ return [
138
+ process_probs(final_mbti_probs, MBTI_LABELS),
139
+ process_probs(final_emo_probs, EMO_LABELS),
140
+ process_probs(avg_a_vibe, VIBE_LABELS),
141
+ process_probs(avg_a_int, INTENSITY_LABELS),
142
+ f"{avg_tempo_bpm:.2f} BPM"
143
+ ]
144
+
145
  except Exception as e: return [{"Error": str(e)}] * 5
146
 
147
+ # -- GRADIO INTERFACE --
148
+ with gr.Blocks(theme=gr.themes.Soft()) as demo:
149
+ gr.Markdown("# 🎸 Neural Math Rock - Hybrid Analysis")
150
+ gr.Markdown("Identify personality and emotional states using Ensemble Audio-Text Models.")
151
 
152
  with gr.Row():
153
  with gr.Column():
154
+ search_input = gr.Textbox(label="YouTube Search", placeholder="Artist - Song Title")
155
+ btn_fetch = gr.Button("FETCH ASSETS", variant="secondary")
156
+ audio_input = gr.Audio(type="filepath", label="Audio Source")
157
+ lyrics_input = gr.Textbox(lines=6, label="Lyrics Source", placeholder="Paste lyrics for better accuracy...")
158
+ btn_run = gr.Button("RUN HYBRID ANALYSIS", variant="primary")
159
+
 
160
  with gr.Column():
161
+ out_mbti = gr.Label(label="Personality (MBTI)")
162
+ out_emo = gr.Label(label="Emotional State")
163
+ out_vibe = gr.Label(label="Acoustic Vibe")
164
+ out_int = gr.Label(label="Intensity Level")
165
+ out_tmp = gr.Textbox(label="Estimated Tempo (BPM)")
166
+
167
+ btn_fetch.click(fn=search_and_fetch, inputs=[search_input], outputs=[audio_input, lyrics_input])
168
+ btn_run.click(fn=analyze_track, inputs=[audio_input, lyrics_input], outputs=[out_mbti, out_emo, out_vibe, out_int, out_tmp])
169
 
170
  if __name__ == "__main__":
171
  demo.launch()