Spaces:
Sleeping
Sleeping
Update app.py
Browse files
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 |
-
#
|
|
|
|
| 167 |
probs_el = torch.sigmoid(raw_el).squeeze().numpy()
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
|
|
|
| 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 =
|
| 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:
|