anggars commited on
Commit
8a16cc4
Β·
verified Β·
1 Parent(s): 92d86ab

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +103 -85
app.py CHANGED
@@ -3,7 +3,7 @@ import torch.nn as nn
3
  import librosa
4
  import numpy as np
5
  import gradio as gr
6
- from transformers import AutoTokenizer, AutoModel, WavLMModel, AutoFeatureExtractor
7
  from huggingface_hub import hf_hub_download
8
  import warnings
9
 
@@ -11,131 +11,149 @@ warnings.filterwarnings('ignore')
11
 
12
  # ── CONFIGURATION ──
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 from HF Hub...")
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
- class DummyEncoder:
23
- def __init__(self, classes_list):
24
- self.classes_ = np.array(classes_list)
 
 
 
25
 
26
- # Label Encoders
27
- label_encoders = {
28
- 'mbti': ckpt.get('le_mbti', DummyEncoder(['ENFJ', 'ENFP', 'ENTJ', 'ENTP', 'ESFJ', 'ESFP', 'ESTJ', 'ESTP', 'INFJ', 'INFP', 'INTJ', 'INTP', 'ISFJ', 'ISFP', 'ISTJ', 'ISTP'])),
29
- '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'])),
30
- 'vibe': DummyEncoder(['Aggressive', 'Atmospheric', 'Melancholic', 'Technical']),
31
- 'intensity': DummyEncoder(['High', 'Low', 'Medium']),
32
- 'tempo': DummyEncoder(['Fast', 'Moderate', 'Slow'])
33
- }
34
-
35
- num_classes = {col: len(le.classes_) for col, le in label_encoders.items()}
36
 
37
  # ── ARCHITECTURE ──
38
- class MultimodalMathRock(nn.Module):
39
  def __init__(self):
40
  super().__init__()
41
- self.audio_encoder = WavLMModel.from_pretrained("microsoft/wavlm-base")
42
- self.text_encoder = AutoModel.from_pretrained('xlm-roberta-base')
 
 
 
43
  self.fusion = nn.Sequential(
44
- nn.Dropout(0.3),
45
- nn.Linear(768 + 768, 512),
46
- nn.ReLU()
 
47
  )
48
- self.heads = nn.ModuleDict({
49
- 'mbti': nn.Linear(512, num_classes['mbti']),
50
- 'emotion': nn.Linear(512, num_classes['emotion']),
51
- 'vibe': nn.Linear(512, num_classes['vibe']),
52
- 'intensity': nn.Linear(512, num_classes['intensity']),
53
- 'tempo': nn.Linear(512, num_classes['tempo'])
54
- })
55
-
56
- def forward(self, audio_values, input_ids, attention_mask):
57
- x_a = self.audio_encoder(audio_values).last_hidden_state.mean(dim=1)
58
- x_t = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :]
59
- fused = self.fusion(torch.cat([x_a, x_t], dim=1))
60
- return {col: self.heads[col](fused) for col in TARGET_COLS}
 
 
 
 
 
 
 
 
 
 
61
 
62
  print("Initializing architecture and loading weights...")
63
- model = MultimodalMathRock()
64
- state_dict = ckpt['model_state'] if 'model_state' in ckpt else ckpt
65
- model.load_state_dict(state_dict, strict=False)
66
  model.eval()
67
- tokenizer = AutoTokenizer.from_pretrained('xlm-roberta-base')
68
- audio_processor = AutoFeatureExtractor.from_pretrained("microsoft/wavlm-base")
69
 
70
  # ── HYBRID INFERENCE PIPELINE ──
71
  def predict_multimodal(audio_path, lyrics):
72
  if audio_path is None:
73
- return [None]*5
74
 
75
  try:
76
- # 1. Load Audio & Hitung Real Tempo (BPM) pake Librosa
77
- y_full, _ = librosa.load(audio_path, sr=SR, mono=True)
78
- total_dur = librosa.get_duration(y=y_full, sr=SR)
 
 
 
79
 
80
- # Beat Tracking Logic
81
- onset_env = librosa.onset.onset_strength(y=y_full, sr=SR)
82
- tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=SR)
83
  print(f"Calculated BPM: {tempo_bpm}")
84
 
85
- # 2. Windowing Strategy (8 Windows)
86
- num_windows = min(8, int(total_dur // DURATION))
87
- if num_windows == 0: num_windows = 1
88
- offsets = np.linspace(0, total_dur - DURATION, num_windows)
89
-
90
- if not lyrics or lyrics.strip() == "":
91
- lyrics = "instrumental math rock"
92
-
93
- enc = tokenizer(lyrics, max_length=128, padding='max_length', truncation=True, return_tensors='pt')
 
 
 
94
  all_logits = {col: [] for col in TARGET_COLS}
95
 
96
- # 3. Model Prediction
97
  with torch.no_grad():
98
- for offset in offsets:
99
- start = int(offset * SR)
100
- end = start + int(DURATION * SR)
101
- y_chunk = y_full[start:end]
102
- if len(y_chunk) < int(DURATION * SR):
103
- y_chunk = np.pad(y_chunk, (0, int(DURATION * SR) - len(y_chunk)))
104
-
105
- audio_inputs = audio_processor(y_chunk, sampling_rate=SR, return_tensors="pt")
106
- out = model(audio_inputs.input_values, enc['input_ids'], enc['attention_mask'])
107
  for col in TARGET_COLS:
108
  all_logits[col].append(out[col][0])
109
 
110
- # 4. Result Processing with Tempo Override
111
  final_results = []
 
 
 
 
 
112
  for col in TARGET_COLS:
113
  avg_logits = torch.stack(all_logits[col]).mean(dim=0)
 
114
 
115
- # ── TEMPO LOGIC OVERRIDE ──
116
  if col == 'tempo':
117
- tempo_classes = list(label_encoders['tempo'].classes_)
118
- # Indexing: Fast (0), Moderate (1), Slow (2) berdasarkan dummy encoder lo
119
- if tempo_bpm > 125: # Jika BPM di atas 125, paksa ke arah Fast
120
- avg_logits[0] += 5.0
121
- elif tempo_bpm > 90: # Jika BPM di atas 90, paksa ke arah Moderate
122
- avg_logits[1] += 3.0
123
- # Jika di bawah itu, biarin model mutusin (biasanya Slow)
124
-
125
- probs = torch.nn.functional.softmax(avg_logits, dim=0).numpy()
126
- classes = label_encoders[col].classes_
 
 
 
127
 
128
- pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes)) if i < len(probs)}
129
  sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
130
  final_results.append(sorted_preds)
131
 
132
  return final_results
133
 
134
  except Exception as e:
135
- return [{"Error": str(e)}]*5
136
 
137
- # ── GRADIO UI (Glass Theme) ──
138
- with gr.Blocks(theme=gr.themes.Glass()) as demo:
139
  gr.Markdown("# Neural Math Rock Multimodal Analysis")
140
  gr.Markdown("Analyzing MBTI, Emotion, and Vibes using Hybrid Deep Learning & DSP.")
141
 
@@ -149,7 +167,7 @@ with gr.Blocks(theme=gr.themes.Glass()) as demo:
149
  out_mbti = gr.Label(label="Personality (MBTI)")
150
  out_emotion = gr.Label(label="Emotional State")
151
  out_vibe = gr.Label(label="Acoustic Vibe")
152
- out_intensity = gr.Label(label="Energy Level")
153
  out_tempo = gr.Label(label="Tempo Classification (BPM Adjusted)")
154
 
155
  btn.click(
 
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 warnings
9
 
 
11
 
12
  # ── CONFIGURATION ──
13
  REPO_ID = "anggars/neural-mathrock"
14
+ DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
15
  TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
16
 
17
+ print("Downloading model from Hugging Face Hub...")
18
  model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
19
+ ckpt = torch.load(model_path, map_location=DEVICE, weights_only=False)
20
 
21
+ # ── LOAD LABEL ENCODERS ──
22
+ le_mbti = ckpt['le_mbti']
23
+ le_emotion = ckpt['le_emotion']
24
+ le_vibe = ckpt['le_vibe']
25
+ le_intensity = ckpt['le_intensity']
26
+ le_tempo = ckpt['le_tempo']
27
 
28
+ NUM_MBTI = len(le_mbti.classes_)
29
+ NUM_EMOTION = len(le_emotion.classes_)
30
+ NUM_VIBE = len(le_vibe.classes_)
31
+ NUM_INTENSITY = len(le_intensity.classes_)
32
+ NUM_TEMPO = len(le_tempo.classes_)
 
 
 
 
 
33
 
34
  # ── ARCHITECTURE ──
35
+ class HybridMultimodalModel(nn.Module):
36
  def __init__(self):
37
  super().__init__()
38
+ self.text_model = XLMRobertaModel.from_pretrained('anggars/xlm-mbti')
39
+ self.audio_model = WavLMModel.from_pretrained('microsoft/wavlm-base')
40
+
41
+ self.audio_proj = nn.Linear(768, 256)
42
+
43
  self.fusion = nn.Sequential(
44
+ nn.Linear(1024, 512),
45
+ nn.BatchNorm1d(512),
46
+ nn.ReLU(),
47
+ nn.Dropout(0.4),
48
  )
49
+
50
+ self.head_mbti = nn.Linear(512, NUM_MBTI)
51
+ self.head_emotion = nn.Linear(512, NUM_EMOTION)
52
+ self.head_vibe = nn.Linear(512, NUM_VIBE)
53
+ self.head_intensity = nn.Linear(512, NUM_INTENSITY)
54
+ self.head_tempo = nn.Linear(512, NUM_TEMPO)
55
+
56
+ def forward(self, input_ids, attention_mask, audio_values):
57
+ text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
58
+ text_feat = text_out.pooler_output
59
+
60
+ audio_out = self.audio_model(audio_values).last_hidden_state
61
+ audio_feat = self.audio_proj(audio_out.mean(dim=1))
62
+
63
+ fused = self.fusion(torch.cat([text_feat, audio_feat], dim=-1))
64
+
65
+ return {
66
+ 'mbti': self.head_mbti(fused),
67
+ 'emotion': self.head_emotion(fused),
68
+ 'vibe': self.head_vibe(fused),
69
+ 'intensity': self.head_intensity(fused),
70
+ 'tempo': self.head_tempo(fused),
71
+ }
72
 
73
  print("Initializing architecture and loading weights...")
74
+ model = HybridMultimodalModel().to(DEVICE)
75
+ model.load_state_dict(ckpt['model_state'], strict=False)
 
76
  model.eval()
77
+
78
+ tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')
79
 
80
  # ── HYBRID INFERENCE PIPELINE ──
81
  def predict_multimodal(audio_path, lyrics):
82
  if audio_path is None:
83
+ return [{"Error": "Audio file cannot be empty."}] * 5
84
 
85
  try:
86
+ # 1. Text Processing
87
+ text = str(lyrics).strip() if lyrics else "[INSTRUMENTAL]"
88
+ enc = tokenizer(text, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)
89
+
90
+ # 2. Audio Processing & DSP
91
+ wav, sr = librosa.load(audio_path, sr=16000)
92
 
93
+ onset_env = librosa.onset.onset_strength(y=wav, sr=sr)
94
+ tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=sr)
 
95
  print(f"Calculated BPM: {tempo_bpm}")
96
 
97
+ # 3. Chunking (20-second sequential blocks)
98
+ chunk_size = 16000 * 20
99
+ chunks = []
100
+ for i in range(0, len(wav), chunk_size):
101
+ chunk = wav[i:i + chunk_size]
102
+ if len(chunk) < chunk_size:
103
+ chunk = np.pad(chunk, (0, chunk_size - len(chunk)))
104
+ chunks.append(chunk)
105
+
106
+ if not chunks:
107
+ chunks = [np.zeros(chunk_size)]
108
+
109
  all_logits = {col: [] for col in TARGET_COLS}
110
 
111
+ # 4. Block-by-block Prediction
112
  with torch.no_grad():
113
+ for chunk in chunks:
114
+ audio_tensor = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
115
+ out = model(enc['input_ids'], enc['attention_mask'], audio_tensor)
 
 
 
 
 
 
116
  for col in TARGET_COLS:
117
  all_logits[col].append(out[col][0])
118
 
119
+ # 5. Averaging and Formatting Results
120
  final_results = []
121
+ encoders_map = {
122
+ 'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe,
123
+ 'intensity': le_intensity, 'tempo': le_tempo
124
+ }
125
+
126
  for col in TARGET_COLS:
127
  avg_logits = torch.stack(all_logits[col]).mean(dim=0)
128
+ le = encoders_map[col]
129
 
130
+ # TEMPO LOGIC OVERRIDE
131
  if col == 'tempo':
132
+ tempo_classes = list(le.classes_)
133
+ try:
134
+ fast_idx = tempo_classes.index('Fast')
135
+ mod_idx = tempo_classes.index('Moderate')
136
+ if tempo_bpm > 125:
137
+ avg_logits[fast_idx] += 5.0
138
+ elif tempo_bpm > 90:
139
+ avg_logits[mod_idx] += 3.0
140
+ except ValueError:
141
+ pass
142
+
143
+ probs = torch.nn.functional.softmax(avg_logits, dim=0).cpu().numpy()
144
+ classes = le.classes_
145
 
146
+ pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes))}
147
  sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
148
  final_results.append(sorted_preds)
149
 
150
  return final_results
151
 
152
  except Exception as e:
153
+ return [{"Error": f"Processing failed: {str(e)}"}] * 5
154
 
155
+ # ── GRADIO UI ──
156
+ with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
157
  gr.Markdown("# Neural Math Rock Multimodal Analysis")
158
  gr.Markdown("Analyzing MBTI, Emotion, and Vibes using Hybrid Deep Learning & DSP.")
159
 
 
167
  out_mbti = gr.Label(label="Personality (MBTI)")
168
  out_emotion = gr.Label(label="Emotional State")
169
  out_vibe = gr.Label(label="Acoustic Vibe")
170
+ out_intensity = gr.Label(label="Intensity Level")
171
  out_tempo = gr.Label(label="Tempo Classification (BPM Adjusted)")
172
 
173
  btn.click(