anggars commited on
Commit
8678f34
·
verified ·
1 Parent(s): dd98f8b

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +20 -5
app.py CHANGED
@@ -33,6 +33,8 @@ TEMPO_LABELS = ['slow', 'moderate', 'fast']
33
  WAVLM_SR = 16000
34
  CLAP_SR = 48000
35
 
 
 
36
  # -- HYBRID ARCHITECTURE ALIGNED WITH TRAINING EXECUTION --
37
  class AudioMathRockModel(nn.Module):
38
  def __init__(self):
@@ -75,6 +77,7 @@ class AudioMathRockModel(nn.Module):
75
  self.tmp_head = nn.Linear(512, len(TEMPO_LABELS))
76
 
77
  def forward(self, wavlm_values, clap_values):
 
78
  wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
79
  clap_feats = self.clap(clap_values).pooler_output
80
 
@@ -83,10 +86,12 @@ class AudioMathRockModel(nn.Module):
83
  mel = self.mel_spectrogram(wavlm_values.float())
84
  mel_db = self.amplitude_to_db(mel)
85
 
 
86
  cnn_feats = F.gelu(self.cnn_extractor(mel_db.unsqueeze(1))).to(wavlm_feats.dtype)
87
  wlm_p = F.gelu(self.wlm_proj(wavlm_feats))
88
  clp_p = F.gelu(self.clp_proj(clap_feats))
89
 
 
90
  fused = self.fusion(torch.cat([wlm_p, clp_p, cnn_feats], dim=-1))
91
  return self.emo_head(fused), self.vibe_head(fused), self.int_head(fused), self.tmp_head(fused)
92
 
@@ -115,6 +120,7 @@ def analyze_track(audio_path, lyrics_input):
115
 
116
  res_mbti, res_emo, res_vibe, res_int, res_tmp = {}, {}, {}, {}, {}
117
 
 
118
  if has_lyrics and not has_audio:
119
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
120
  with torch.no_grad():
@@ -128,16 +134,19 @@ def analyze_track(audio_path, lyrics_input):
128
  res_emo = dict(sorted(emo_dict.items(), key=lambda x: x[1], reverse=True)[:3])
129
  return res_mbti, res_emo, {}, {}, {}
130
 
 
131
  try:
132
  wav_orig, orig_sr = librosa.load(audio_path, sr=None, mono=True)
133
  waveform = torch.tensor(wav_orig).unsqueeze(0)
134
 
 
135
  wlm_wave = torchaudio.functional.resample(waveform, orig_sr, WAVLM_SR).squeeze(0)
136
  clp_wave = torchaudio.functional.resample(waveform, orig_sr, CLAP_SR).squeeze(0)
137
 
138
  wavlm_samples = WAVLM_SR * 15
139
  clap_samples = CLAP_SR * 15
140
 
 
141
  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]
142
  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]
143
 
@@ -158,26 +167,31 @@ def analyze_track(audio_path, lyrics_input):
158
  il_list.append(il.cpu())
159
  tl_list.append(tl.cpu())
160
 
 
161
  raw_el = torch.cat(el_list, dim=0).mean(dim=0, keepdim=True)
162
  raw_vl = torch.cat(vl_list, dim=0).mean(dim=0, keepdim=True)
163
  raw_il = torch.cat(il_list, dim=0).mean(dim=0, keepdim=True)
164
  raw_tl = torch.cat(tl_list, dim=0).mean(dim=0, keepdim=True)
165
 
166
- # FIXED LOGIC UNIFIED TO SIGMOID INFERENCE AS SPECIFIED IN MULTI-LABEL TRAINING
 
167
  probs_el = torch.sigmoid(raw_el).squeeze().numpy()
168
- probs_vl = torch.sigmoid(raw_vl).squeeze().numpy()
169
- probs_il = torch.sigmoid(raw_il).squeeze().numpy()
170
- probs_tl = torch.sigmoid(raw_tl).squeeze().numpy()
 
171
 
172
  res_emo_raw = {EMO_LABELS[i]: float(probs_el[i]) for i in range(len(EMO_LABELS))}
173
  res_vibe = {VIBE_LABELS[i]: float(probs_vl[i]) for i in range(len(VIBE_LABELS))}
174
  res_int = {INTENSITY_LABELS[i]: float(probs_il[i]) for i in range(len(INTENSITY_LABELS))}
175
  res_tmp = {TEMPO_LABELS[i]: float(probs_tl[i]) for i in range(len(TEMPO_LABELS))}
176
 
 
177
  res_vibe = dict(sorted(res_vibe.items(), key=lambda x: x[1], reverse=True)[:3])
178
  res_int = dict(sorted(res_int.items(), key=lambda x: x[1], reverse=True)[:3])
179
  res_tmp = dict(sorted(res_tmp.items(), key=lambda x: x[1], reverse=True)[:3])
180
 
 
181
  if has_lyrics:
182
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
183
  with torch.no_grad():
@@ -187,8 +201,9 @@ def analyze_track(audio_path, lyrics_input):
187
  m_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
188
  res_mbti = dict(sorted(m_dict.items(), key=lambda x: x[1], reverse=True)[:3])
189
 
 
190
  for k in res_emo_raw.keys():
191
- idx = EMO_LABELS.index(k)
192
  res_emo_raw[k] = (res_emo_raw[k] * 0.3) + (float(t_emo_probs[idx]) * 0.7)
193
  res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
194
  else:
 
33
  WAVLM_SR = 16000
34
  CLAP_SR = 48000
35
 
36
+ emo2id = {e: i for i, e in enumerate(EMO_LABELS)}
37
+
38
  # -- HYBRID ARCHITECTURE ALIGNED WITH TRAINING EXECUTION --
39
  class AudioMathRockModel(nn.Module):
40
  def __init__(self):
 
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
 
 
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
 
 
120
 
121
  res_mbti, res_emo, res_vibe, res_int, res_tmp = {}, {}, {}, {}, {}
122
 
123
+ # --- TEXT ONLY PROCESSING ---
124
  if has_lyrics and not has_audio:
125
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
126
  with torch.no_grad():
 
134
  res_emo = dict(sorted(emo_dict.items(), key=lambda x: x[1], reverse=True)[:3])
135
  return res_mbti, res_emo, {}, {}, {}
136
 
137
+ # --- AUDIO PROCESSING ---
138
  try:
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
  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 correct activation logic (Pure Model Inference)
177
+ # Emotion remains Sigmoid (Multi-label distribution)
178
  probs_el = torch.sigmoid(raw_el).squeeze().numpy()
179
+ # Vibe, Intensity, Tempo strictly evaluated as Softmax (Single-label exclusivity)
180
+ probs_vl = F.softmax(raw_vl, dim=-1).squeeze().numpy()
181
+ probs_il = F.softmax(raw_il, dim=-1).squeeze().numpy()
182
+ probs_tl = F.softmax(raw_tl, dim=-1).squeeze().numpy()
183
 
184
  res_emo_raw = {EMO_LABELS[i]: float(probs_el[i]) for i in range(len(EMO_LABELS))}
185
  res_vibe = {VIBE_LABELS[i]: float(probs_vl[i]) for i in range(len(VIBE_LABELS))}
186
  res_int = {INTENSITY_LABELS[i]: float(probs_il[i]) for i in range(len(INTENSITY_LABELS))}
187
  res_tmp = {TEMPO_LABELS[i]: float(probs_tl[i]) for i in range(len(TEMPO_LABELS))}
188
 
189
+ # Sort and truncate dicts for UI presentation
190
  res_vibe = dict(sorted(res_vibe.items(), key=lambda x: x[1], reverse=True)[:3])
191
  res_int = dict(sorted(res_int.items(), key=lambda x: x[1], reverse=True)[:3])
192
  res_tmp = dict(sorted(res_tmp.items(), key=lambda x: x[1], reverse=True)[:3])
193
 
194
+ # --- LATE FUSION PIPELINE WITH TEXT MODEL ---
195
  if has_lyrics:
196
  t_in = tokenizer(str(lyrics_input), truncation=True, padding=True, max_length=256, return_tensors="pt").to(DEVICE)
197
  with torch.no_grad():
 
201
  m_dict = {MBTI_LABELS[i]: float(t_mbti_probs[i]) for i in range(len(MBTI_LABELS))}
202
  res_mbti = dict(sorted(m_dict.items(), key=lambda x: x[1], reverse=True)[:3])
203
 
204
+ # Weighted ensemble calculation
205
  for k in res_emo_raw.keys():
206
+ idx = emo2id[k]
207
  res_emo_raw[k] = (res_emo_raw[k] * 0.3) + (float(t_emo_probs[idx]) * 0.7)
208
  res_emo = dict(sorted(res_emo_raw.items(), key=lambda x: x[1], reverse=True)[:3])
209
  else: