anggars commited on
Commit
a90b44a
·
verified ·
1 Parent(s): 38bfdba

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +7 -19
app.py CHANGED
@@ -77,21 +77,17 @@ class AudioMathRockModel(nn.Module):
77
  self.tmp_head = nn.Linear(512, len(TEMPO_LABELS))
78
 
79
  def forward(self, wavlm_values, clap_values):
80
- # Extract features from frozen/unfrozen backbones
81
  wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
82
  clap_feats = self.clap(clap_values).pooler_output
83
 
84
- # Audio Mel representation mapping context
85
  with torch.no_grad():
86
  mel = self.mel_spectrogram(wavlm_values.float())
87
  mel_db = self.amplitude_to_db(mel)
88
 
89
- # Align precision and pass through CNN branch
90
  cnn_feats = F.gelu(self.cnn_extractor(mel_db.unsqueeze(1))).to(wavlm_feats.dtype)
91
  wlm_p = F.gelu(self.wlm_proj(wavlm_feats))
92
  clp_p = F.gelu(self.clp_proj(clap_feats))
93
 
94
- # Late fusion sequence
95
  fused = self.fusion(torch.cat([wlm_p, clp_p, cnn_feats], dim=-1))
96
  return self.emo_head(fused), self.vibe_head(fused), self.int_head(fused), self.tmp_head(fused)
97
 
@@ -139,14 +135,12 @@ def analyze_track(audio_path, lyrics_input):
139
  wav_orig, orig_sr = librosa.load(audio_path, sr=None, mono=True)
140
  waveform = torch.tensor(wav_orig).unsqueeze(0)
141
 
142
- # Resampling constraints based on model requirements
143
  wlm_wave = torchaudio.functional.resample(waveform, orig_sr, WAVLM_SR).squeeze(0)
144
  clp_wave = torchaudio.functional.resample(waveform, orig_sr, CLAP_SR).squeeze(0)
145
 
146
  wavlm_samples = WAVLM_SR * 15
147
  clap_samples = CLAP_SR * 15
148
 
149
- # Chunk generation for robust global representation
150
  wlm_chunks = [wlm_wave[i:i+wavlm_samples].numpy() for i in range(0, len(wlm_wave), wavlm_samples) if len(wlm_wave[i:i+wavlm_samples]) > WAVLM_SR]
151
  clp_chunks = [clp_wave[i:i+clap_samples].numpy() for i in range(0, len(clp_wave), clap_samples) if len(clp_wave[i:i+clap_samples]) > CLAP_SR]
152
 
@@ -167,35 +161,29 @@ def analyze_track(audio_path, lyrics_input):
167
  il_list.append(il.cpu())
168
  tl_list.append(tl.cpu())
169
 
170
- # Extract pure raw logits
171
  raw_el = torch.cat(el_list, dim=0).mean(dim=0, keepdim=True)
172
  raw_vl = torch.cat(vl_list, dim=0).mean(dim=0, keepdim=True)
173
  raw_il = torch.cat(il_list, dim=0).mean(dim=0, keepdim=True)
174
  raw_tl = torch.cat(tl_list, dim=0).mean(dim=0, keepdim=True)
175
 
176
- # Apply Dynamic Scaling Variance to break uniform logit binding
177
- logit_std = raw_el.std(dim=-1, keepdim=True).clamp(min=1e-6)
178
- scaled_el = raw_el / logit_std
179
 
180
- # Softmax over Multi-label head to strictly enforce logit class competition
181
- probs_el = F.softmax(scaled_el / 1.5, dim=-1).squeeze().numpy()
182
-
183
- # Vibe, Intensity, Tempo strictly evaluated as Softmax (Single-label exclusivity)
184
  probs_vl = F.softmax(raw_vl, dim=-1).squeeze().numpy()
185
  probs_il = F.softmax(raw_il, dim=-1).squeeze().numpy()
186
  probs_tl = F.softmax(raw_tl, dim=-1).squeeze().numpy()
187
 
188
- res_emo_raw = {EMO_LABELS[i]: float(probs_el[i]) for i in range(len(EMO_LABELS))}
189
  res_vibe = {VIBE_LABELS[i]: float(probs_vl[i]) for i in range(len(VIBE_LABELS))}
190
  res_int = {INTENSITY_LABELS[i]: float(probs_il[i]) for i in range(len(INTENSITY_LABELS))}
191
  res_tmp = {TEMPO_LABELS[i]: float(probs_tl[i]) for i in range(len(TEMPO_LABELS))}
192
 
193
- # Sort and truncate dicts for UI presentation
194
  res_vibe = dict(sorted(res_vibe.items(), key=lambda x: x[1], reverse=True)[:3])
195
  res_int = dict(sorted(res_int.items(), key=lambda x: x[1], reverse=True)[:3])
196
  res_tmp = dict(sorted(res_tmp.items(), key=lambda x: x[1], reverse=True)[:3])
197
 
198
- # --- LATE FUSION PIPELINE WITH TEXT MODEL ---
199
  if has_lyrics:
200
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
201
  with torch.no_grad():
@@ -205,10 +193,10 @@ def analyze_track(audio_path, lyrics_input):
205
  m_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
206
  res_mbti = dict(sorted(m_dict.items(), key=lambda x: x[1], reverse=True)[:3])
207
 
208
- # Weighted ensemble calculation
209
  for k in res_emo_raw.keys():
210
  idx = emo2id[k]
211
- res_emo_raw[k] = (res_emo_raw[k] * 0.3) + (float(t_emo_probs[idx]) * 0.7)
212
  res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
213
  else:
214
  res_mbti = {}
 
77
  self.tmp_head = nn.Linear(512, len(TEMPO_LABELS))
78
 
79
  def forward(self, wavlm_values, clap_values):
 
80
  wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
81
  clap_feats = self.clap(clap_values).pooler_output
82
 
 
83
  with torch.no_grad():
84
  mel = self.mel_spectrogram(wavlm_values.float())
85
  mel_db = self.amplitude_to_db(mel)
86
 
 
87
  cnn_feats = F.gelu(self.cnn_extractor(mel_db.unsqueeze(1))).to(wavlm_feats.dtype)
88
  wlm_p = F.gelu(self.wlm_proj(wavlm_feats))
89
  clp_p = F.gelu(self.clp_proj(clap_feats))
90
 
 
91
  fused = self.fusion(torch.cat([wlm_p, clp_p, cnn_feats], dim=-1))
92
  return self.emo_head(fused), self.vibe_head(fused), self.int_head(fused), self.tmp_head(fused)
93
 
 
135
  wav_orig, orig_sr = librosa.load(audio_path, sr=None, mono=True)
136
  waveform = torch.tensor(wav_orig).unsqueeze(0)
137
 
 
138
  wlm_wave = torchaudio.functional.resample(waveform, orig_sr, WAVLM_SR).squeeze(0)
139
  clp_wave = torchaudio.functional.resample(waveform, orig_sr, CLAP_SR).squeeze(0)
140
 
141
  wavlm_samples = WAVLM_SR * 15
142
  clap_samples = CLAP_SR * 15
143
 
 
144
  wlm_chunks = [wlm_wave[i:i+wavlm_samples].numpy() for i in range(0, len(wlm_wave), wavlm_samples) if len(wlm_wave[i:i+wavlm_samples]) > WAVLM_SR]
145
  clp_chunks = [clp_wave[i:i+clap_samples].numpy() for i in range(0, len(clp_wave), clap_samples) if len(clp_wave[i:i+clap_samples]) > CLAP_SR]
146
 
 
161
  il_list.append(il.cpu())
162
  tl_list.append(tl.cpu())
163
 
 
164
  raw_el = torch.cat(el_list, dim=0).mean(dim=0, keepdim=True)
165
  raw_vl = torch.cat(vl_list, dim=0).mean(dim=0, keepdim=True)
166
  raw_il = torch.cat(il_list, dim=0).mean(dim=0, keepdim=True)
167
  raw_tl = torch.cat(tl_list, dim=0).mean(dim=0, keepdim=True)
168
 
169
+ # Enforce highly competitive soft-scaling for audio emotion logits
170
+ scaled_el = raw_el / raw_el.std(dim=-1, keepdim=True).clamp(min=1e-6)
171
+ probs_el_audio = F.softmax(scaled_el / 0.3, dim=-1).squeeze().numpy()
172
 
 
 
 
 
173
  probs_vl = F.softmax(raw_vl, dim=-1).squeeze().numpy()
174
  probs_il = F.softmax(raw_il, dim=-1).squeeze().numpy()
175
  probs_tl = F.softmax(raw_tl, dim=-1).squeeze().numpy()
176
 
177
+ res_emo_raw = {EMO_LABELS[i]: float(probs_el_audio[i]) for i in range(len(EMO_LABELS))}
178
  res_vibe = {VIBE_LABELS[i]: float(probs_vl[i]) for i in range(len(VIBE_LABELS))}
179
  res_int = {INTENSITY_LABELS[i]: float(probs_il[i]) for i in range(len(INTENSITY_LABELS))}
180
  res_tmp = {TEMPO_LABELS[i]: float(probs_tl[i]) for i in range(len(TEMPO_LABELS))}
181
 
 
182
  res_vibe = dict(sorted(res_vibe.items(), key=lambda x: x[1], reverse=True)[:3])
183
  res_int = dict(sorted(res_int.items(), key=lambda x: x[1], reverse=True)[:3])
184
  res_tmp = dict(sorted(res_tmp.items(), key=lambda x: x[1], reverse=True)[:3])
185
 
186
+ # --- MULTIMODAL LATE FUSION ENGAGEMENT ---
187
  if has_lyrics:
188
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
189
  with torch.no_grad():
 
193
  m_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
194
  res_mbti = dict(sorted(m_dict.items(), key=lambda x: x[1], reverse=True)[:3])
195
 
196
+ # Weighted late fusion to prioritize text model for emotion taxonomy resolution
197
  for k in res_emo_raw.keys():
198
  idx = emo2id[k]
199
+ res_emo_raw[k] = (res_emo_raw[k] * 0.15) + (float(t_emo_probs[idx]) * 0.85)
200
  res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
201
  else:
202
  res_mbti = {}