anggars commited on
Commit
a613a36
Β·
verified Β·
1 Parent(s): e861fa1

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +63 -38
app.py CHANGED
@@ -13,60 +13,77 @@ warnings.filterwarnings('ignore')
13
  REPO_ID = "anggars/neural-mathrock"
14
  SR = 16000
15
  DURATION = 10
 
16
  TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
17
 
18
- print("Downloading model and extracting label encoders...")
19
  model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
 
20
  ckpt = torch.load(model_path, map_location=torch.device('cpu'), weights_only=False)
21
 
22
- # ── LABEL ENCODER PATCH ──
23
  class DummyEncoder:
24
  def __init__(self, classes_list):
25
  self.classes_ = np.array(classes_list)
26
 
 
27
  label_encoders = {
28
- 'mbti': ckpt['le_mbti'],
29
- 'emotion': ckpt['le_emotion'],
30
  'vibe': DummyEncoder(['Aggressive', 'Atmospheric', 'Melancholic', 'Technical']),
31
  'intensity': DummyEncoder(['High', 'Low', 'Medium']),
32
  'tempo': DummyEncoder(['Fast', 'Moderate', 'Slow'])
33
  }
 
34
  num_classes = {col: len(le.classes_) for col, le in label_encoders.items()}
35
 
36
  # ============================================================
37
- # ── ARCHITECTURE (FIXED EXACT MAPPING) ──
38
  # ============================================================
39
  class MultimodalMathRock(nn.Module):
40
  def __init__(self):
41
  super().__init__()
42
- # Langsung panggil model HF, jangan dibungkus class lagi biar path statenya cocok
43
  self.audio_encoder = WavLMModel.from_pretrained("microsoft/wavlm-base")
44
  self.text_encoder = AutoModel.from_pretrained('xlm-roberta-base')
45
 
46
- # Matikan gradien
47
- for p in self.audio_encoder.parameters(): p.requires_grad = False
48
- for p in self.text_encoder.parameters(): p.requires_grad = False
49
-
50
  self.fusion = nn.Sequential(
51
- nn.Dropout(0.3), # index 0
52
- nn.Linear(768 + 768, 512), # index 1 (cocok sama log "fusion.1.weight")
53
- nn.ReLU() # index 2
54
  )
 
 
 
55
  self.heads = nn.ModuleDict({
56
- col: nn.Linear(512, num_classes[col]) for col in TARGET_COLS
 
 
 
 
57
  })
58
 
59
  def forward(self, audio_values, input_ids, attention_mask):
 
60
  x_a = self.audio_encoder(audio_values).last_hidden_state.mean(dim=1)
 
61
  x_t = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :]
62
- fused = self.fusion(torch.cat([x_a, x_t], dim=1))
 
 
 
 
63
  return {col: self.heads[col](fused) for col in TARGET_COLS}
 
64
  # ============================================================
65
 
66
  print("Initializing architecture and loading weights...")
67
  model = MultimodalMathRock()
68
- # STRICT=TRUE: Kalo ada yang beda dikit, mending error di awal daripada ngaco
69
- model.load_state_dict(ckpt['model_state'])
 
 
70
  model.eval()
71
 
72
  tokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')
@@ -78,12 +95,13 @@ def predict_multimodal(audio_path, lyrics):
78
  return [None]*5
79
 
80
  try:
 
81
  y, _ = librosa.load(audio_path, sr=SR, duration=DURATION, mono=True)
82
  inputs = audio_processor(y, sampling_rate=SR, return_tensors="pt")
83
- audio_tensor = inputs.input_values
84
-
85
  if not lyrics or lyrics.strip() == "":
86
- lyrics = "instrumental"
87
 
88
  enc = tokenizer(
89
  lyrics, max_length=128, padding='max_length',
@@ -91,39 +109,41 @@ def predict_multimodal(audio_path, lyrics):
91
  )
92
 
93
  with torch.no_grad():
94
- out = model(audio_tensor, enc['input_ids'], enc['attention_mask'])
95
 
96
  final_results = []
97
  for col in TARGET_COLS:
98
- probs = torch.nn.functional.softmax(out[col][0], dim=0).numpy()
99
- classes = label_encoders[col].classes_.astype(str)
 
100
 
101
- pred_dict = {classes[i]: float(probs[i]) for i in range(len(classes))}
102
- top_dict = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
103
- final_results.append(top_dict)
 
104
 
105
  return final_results
106
 
107
  except Exception as e:
108
  return [{"Error": str(e)}]*5
109
 
110
- # ── UI ──
111
- with gr.Blocks() as demo:
112
- gr.Markdown("# Neural Math Rock & Midwest Emo Classifier")
113
- gr.Markdown("Multimodal WavLM & XLM-RoBERTa based MBTI and emotion classification system.")
114
 
115
  with gr.Row():
116
  with gr.Column():
117
- audio_input = gr.Audio(type="filepath", label="Input Audio")
118
- lyrics_input = gr.Textbox(lines=5, label="Input Lyrics", placeholder="Paste lyrics here (Optional/Instrumental if empty)")
119
- btn = gr.Button("Analyze", variant="primary")
120
 
121
  with gr.Column():
122
- out_mbti = gr.Label(label="MBTI")
123
- out_emotion = gr.Label(label="Emotion")
124
- out_vibe = gr.Label(label="Vibe")
125
- out_intensity = gr.Label(label="Intensity")
126
- out_tempo = gr.Label(label="Tempo")
127
 
128
  btn.click(
129
  fn=predict_multimodal,
@@ -131,5 +151,10 @@ with gr.Blocks() as demo:
131
  outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
132
  )
133
 
 
 
 
 
 
134
  if __name__ == "__main__":
135
  demo.launch()
 
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'])),
33
  'vibe': DummyEncoder(['Aggressive', 'Atmospheric', 'Melancholic', 'Technical']),
34
  'intensity': DummyEncoder(['High', 'Low', 'Medium']),
35
  'tempo': DummyEncoder(['Fast', 'Moderate', 'Slow'])
36
  }
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']),
62
+ 'vibe': nn.Linear(512, num_classes['vibe']),
63
+ 'intensity': nn.Linear(512, num_classes['intensity']),
64
+ 'tempo': nn.Linear(512, 70 if 'tempo' not in num_classes else num_classes['tempo'])
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()
88
 
89
  tokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')
 
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',
 
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)
124
 
125
  return final_results
126
 
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
  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()