Spaces:
Sleeping
Sleeping
Update app.py
Browse files
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 |
-
#
|
| 177 |
-
|
| 178 |
-
|
| 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(
|
| 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
|
| 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
|
| 209 |
for k in res_emo_raw.keys():
|
| 210 |
idx = emo2id[k]
|
| 211 |
-
res_emo_raw[k] = (res_emo_raw[k] * 0.
|
| 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 = {}
|