anggars commited on
Commit
993e6ad
·
verified ·
1 Parent(s): 92fb8cc

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +37 -57
app.py CHANGED
@@ -9,24 +9,19 @@ import warnings
9
 
10
  warnings.filterwarnings('ignore')
11
 
12
- # ── CONFIGURATION ──
13
  REPO_ID = "anggars/neural-mathrock"
14
  SR = 16000
15
  DURATION = 10
16
- # Pakai urutan yang sama dengan training
17
  TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
18
 
19
  print("Downloading model from HF Hub...")
20
  model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
21
- # Load checkpoint
22
  ckpt = torch.load(model_path, map_location=torch.device('cpu'), weights_only=False)
23
 
24
- # ── LABEL ENCODER PATCH (Sesuai output training lo) ──
25
  class DummyEncoder:
26
  def __init__(self, classes_list):
27
  self.classes_ = np.array(classes_list)
28
 
29
- # Ambil encoder dari checkpoint, kalo gak ada pake manual sesuai kategori di dataset
30
  label_encoders = {
31
  'mbti': ckpt.get('le_mbti', DummyEncoder(['ENFJ', 'ENFP', 'ENTJ', 'ENTP', 'ESFJ', 'ESFP', 'ESTJ', 'ESTP', 'INFJ', 'INFP', 'INTJ', 'INTP', 'ISFJ', 'ISFP', 'ISTJ', 'ISTP'])),
32
  'emotion': ckpt.get('le_emotion', DummyEncoder(['amusement', 'anger', 'annoyance', 'approval', 'caring', 'confusion', 'curiosity', 'desire', 'disappointment', 'disapproval', 'disgust', 'embarrassment', 'excitement', 'fear', 'gratitude', 'grief', 'love', 'nervousness', 'neutral', 'pride', 'realization', 'relief', 'remorse', 'sadness'])),
@@ -37,25 +32,16 @@ label_encoders = {
37
 
38
  num_classes = {col: len(le.classes_) for col, le in label_encoders.items()}
39
 
40
- # ============================================================
41
- # ── ARCHITECTURE (MATCHING EPOCH 10 WEIGHTS) ──
42
- # ============================================================
43
  class MultimodalMathRock(nn.Module):
44
  def __init__(self):
45
  super().__init__()
46
- # Encoder Audio & Text
47
  self.audio_encoder = WavLMModel.from_pretrained("microsoft/wavlm-base")
48
  self.text_encoder = AutoModel.from_pretrained('xlm-roberta-base')
49
-
50
- # Fusion Layer (Linear 1536 -> 512)
51
  self.fusion = nn.Sequential(
52
  nn.Dropout(0.3),
53
  nn.Linear(768 + 768, 512),
54
  nn.ReLU()
55
  )
56
-
57
- # Classification Heads
58
- # Gunakan ModuleDict agar key-nya cocok dengan state_dict "heads.mbti.weight" dst.
59
  self.heads = nn.ModuleDict({
60
  'mbti': nn.Linear(512, num_classes['mbti']),
61
  'emotion': nn.Linear(512, num_classes['emotion']),
@@ -65,23 +51,13 @@ class MultimodalMathRock(nn.Module):
65
  })
66
 
67
  def forward(self, audio_values, input_ids, attention_mask):
68
- # Audio Feature Extraction
69
  x_a = self.audio_encoder(audio_values).last_hidden_state.mean(dim=1)
70
- # Text Feature Extraction (CLS Token)
71
  x_t = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :]
72
-
73
- # Late Fusion
74
- combined = torch.cat([x_a, x_t], dim=1)
75
- fused = self.fusion(combined)
76
-
77
  return {col: self.heads[col](fused) for col in TARGET_COLS}
78
 
79
- # ============================================================
80
-
81
  print("Initializing architecture and loading weights...")
82
  model = MultimodalMathRock()
83
-
84
- # Load state dict dengan filter prefix jika perlu
85
  state_dict = ckpt['model_state'] if 'model_state' in ckpt else ckpt
86
  model.load_state_dict(state_dict, strict=False)
87
  model.eval()
@@ -89,35 +65,43 @@ model.eval()
89
  tokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')
90
  audio_processor = AutoFeatureExtractor.from_pretrained("microsoft/wavlm-base")
91
 
92
- # ── INFERENCE PIPELINE ──
93
  def predict_multimodal(audio_path, lyrics):
94
  if audio_path is None:
95
  return [None]*5
96
 
97
  try:
98
- # Load audio 10 detik awal
99
- y, _ = librosa.load(audio_path, sr=SR, duration=DURATION, mono=True)
100
- inputs = audio_processor(y, sampling_rate=SR, return_tensors="pt")
 
 
 
101
 
102
- # Handle lyrics kosong
103
  if not lyrics or lyrics.strip() == "":
104
- lyrics = "instrumental math rock music"
105
 
106
- enc = tokenizer(
107
- lyrics, max_length=128, padding='max_length',
108
- truncation=True, return_tensors='pt'
109
- )
110
 
111
  with torch.no_grad():
112
- out = model(inputs.input_values, enc['input_ids'], enc['attention_mask'])
113
-
 
 
 
 
 
 
 
 
 
 
114
  final_results = []
115
  for col in TARGET_COLS:
116
- logits = out[col][0]
117
- probs = torch.nn.functional.softmax(logits, dim=0).numpy()
118
  classes = label_encoders[col].classes_
119
 
120
- # Buat dict untuk Gradio Label (Top 3)
121
  pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes)) if i < len(probs)}
122
  sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
123
  final_results.append(sorted_preds)
@@ -127,23 +111,23 @@ def predict_multimodal(audio_path, lyrics):
127
  except Exception as e:
128
  return [{"Error": str(e)}]*5
129
 
130
- # ── GRADIO UI ──
131
- with gr.Blocks(theme=gr.themes.Soft()) as demo:
132
  gr.Markdown("# Neural Math Rock & Midwest Emo Analysis")
133
- gr.Markdown("Analyze audio and lyrics to predict MBTI, Emotion, Vibe, Intensity, and Tempo using Multimodal Transformers.")
134
 
135
  with gr.Row():
136
  with gr.Column():
137
- audio_input = gr.Audio(type="filepath", label="Upload Song (10s will be analyzed)")
138
- lyrics_input = gr.Textbox(lines=5, label="Lyrics", placeholder="Paste lyrics here...")
139
- btn = gr.Button("Analyze Personality & Emotion", variant="primary")
140
 
141
  with gr.Column():
142
- out_mbti = gr.Label(label="Predicted MBTI")
143
- out_emotion = gr.Label(label="Predicted Emotion")
144
- out_vibe = gr.Label(label="Music Vibe")
145
- out_intensity = gr.Label(label="Energy Intensity")
146
- out_tempo = gr.Label(label="Estimated Tempo")
147
 
148
  btn.click(
149
  fn=predict_multimodal,
@@ -151,10 +135,6 @@ with gr.Blocks(theme=gr.themes.Soft()) as demo:
151
  outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
152
  )
153
 
154
- gr.Examples(
155
- examples=[["example.mp3", "I'm not sad, I'm just tired of being alone in this basement."]],
156
- inputs=[audio_input, lyrics_input]
157
- )
158
-
159
  if __name__ == "__main__":
160
- demo.launch()
 
 
9
 
10
  warnings.filterwarnings('ignore')
11
 
 
12
  REPO_ID = "anggars/neural-mathrock"
13
  SR = 16000
14
  DURATION = 10
 
15
  TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
16
 
17
  print("Downloading model from HF Hub...")
18
  model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
 
19
  ckpt = torch.load(model_path, map_location=torch.device('cpu'), weights_only=False)
20
 
 
21
  class DummyEncoder:
22
  def __init__(self, classes_list):
23
  self.classes_ = np.array(classes_list)
24
 
 
25
  label_encoders = {
26
  'mbti': ckpt.get('le_mbti', DummyEncoder(['ENFJ', 'ENFP', 'ENTJ', 'ENTP', 'ESFJ', 'ESFP', 'ESTJ', 'ESTP', 'INFJ', 'INFP', 'INTJ', 'INTP', 'ISFJ', 'ISFP', 'ISTJ', 'ISTP'])),
27
  'emotion': ckpt.get('le_emotion', DummyEncoder(['amusement', 'anger', 'annoyance', 'approval', 'caring', 'confusion', 'curiosity', 'desire', 'disappointment', 'disapproval', 'disgust', 'embarrassment', 'excitement', 'fear', 'gratitude', 'grief', 'love', 'nervousness', 'neutral', 'pride', 'realization', 'relief', 'remorse', 'sadness'])),
 
32
 
33
  num_classes = {col: len(le.classes_) for col, le in label_encoders.items()}
34
 
 
 
 
35
  class MultimodalMathRock(nn.Module):
36
  def __init__(self):
37
  super().__init__()
 
38
  self.audio_encoder = WavLMModel.from_pretrained("microsoft/wavlm-base")
39
  self.text_encoder = AutoModel.from_pretrained('xlm-roberta-base')
 
 
40
  self.fusion = nn.Sequential(
41
  nn.Dropout(0.3),
42
  nn.Linear(768 + 768, 512),
43
  nn.ReLU()
44
  )
 
 
 
45
  self.heads = nn.ModuleDict({
46
  'mbti': nn.Linear(512, num_classes['mbti']),
47
  'emotion': nn.Linear(512, num_classes['emotion']),
 
51
  })
52
 
53
  def forward(self, audio_values, input_ids, attention_mask):
 
54
  x_a = self.audio_encoder(audio_values).last_hidden_state.mean(dim=1)
 
55
  x_t = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :]
56
+ fused = self.fusion(torch.cat([x_a, x_t], dim=1))
 
 
 
 
57
  return {col: self.heads[col](fused) for col in TARGET_COLS}
58
 
 
 
59
  print("Initializing architecture and loading weights...")
60
  model = MultimodalMathRock()
 
 
61
  state_dict = ckpt['model_state'] if 'model_state' in ckpt else ckpt
62
  model.load_state_dict(state_dict, strict=False)
63
  model.eval()
 
65
  tokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')
66
  audio_processor = AutoFeatureExtractor.from_pretrained("microsoft/wavlm-base")
67
 
 
68
  def predict_multimodal(audio_path, lyrics):
69
  if audio_path is None:
70
  return [None]*5
71
 
72
  try:
73
+ y_full, _ = librosa.load(audio_path, sr=SR, mono=True)
74
+ total_dur = librosa.get_duration(y=y_full, sr=SR)
75
+
76
+ num_windows = min(8, int(total_dur // DURATION))
77
+ if num_windows == 0: num_windows = 1
78
+ offsets = np.linspace(0, total_dur - DURATION, num_windows)
79
 
 
80
  if not lyrics or lyrics.strip() == "":
81
+ lyrics = "instrumental math rock"
82
 
83
+ enc = tokenizer(lyrics, max_length=128, padding='max_length', truncation=True, return_tensors='pt')
84
+ all_logits = {col: [] for col in TARGET_COLS}
 
 
85
 
86
  with torch.no_grad():
87
+ for offset in offsets:
88
+ start = int(offset * SR)
89
+ end = start + int(DURATION * SR)
90
+ y_chunk = y_full[start:end]
91
+ if len(y_chunk) < int(DURATION * SR):
92
+ y_chunk = np.pad(y_chunk, (0, int(DURATION * SR) - len(y_chunk)))
93
+
94
+ audio_inputs = audio_processor(y_chunk, sampling_rate=SR, return_tensors="pt")
95
+ out = model(audio_inputs.input_values, enc['input_ids'], enc['attention_mask'])
96
+ for col in TARGET_COLS:
97
+ all_logits[col].append(out[col][0])
98
+
99
  final_results = []
100
  for col in TARGET_COLS:
101
+ avg_logits = torch.stack(all_logits[col]).mean(dim=0)
102
+ probs = torch.nn.functional.softmax(avg_logits, dim=0).numpy()
103
  classes = label_encoders[col].classes_
104
 
 
105
  pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes)) if i < len(probs)}
106
  sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
107
  final_results.append(sorted_preds)
 
111
  except Exception as e:
112
  return [{"Error": str(e)}]*5
113
 
114
+ # ── GRADIO UI (Glass Theme) ──
115
+ with gr.Blocks(theme=gr.themes.Glass()) as demo:
116
  gr.Markdown("# Neural Math Rock & Midwest Emo Analysis")
117
+ gr.Markdown("Multimodal Affective Computing System for Math Rock & Midwest Emo Genres.")
118
 
119
  with gr.Row():
120
  with gr.Column():
121
+ audio_input = gr.Audio(type="filepath", label="Upload Song (Full Analysis)")
122
+ lyrics_input = gr.Textbox(lines=5, label="Lyrics Content", placeholder="Paste lyrics here...")
123
+ btn = gr.Button("START MULTIMODAL ANALYSIS", variant="primary")
124
 
125
  with gr.Column():
126
+ out_mbti = gr.Label(label="Personality Type (MBTI)")
127
+ out_emotion = gr.Label(label="Emotional State")
128
+ out_vibe = gr.Label(label="Acoustic Vibe")
129
+ out_intensity = gr.Label(label="Energy Level")
130
+ out_tempo = gr.Label(label="Tempo Classification")
131
 
132
  btn.click(
133
  fn=predict_multimodal,
 
135
  outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
136
  )
137
 
 
 
 
 
 
138
  if __name__ == "__main__":
139
+ demo.launch()
140
+