Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -5,6 +5,11 @@ 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 |
|
| 10 |
warnings.filterwarnings('ignore')
|
|
@@ -14,6 +19,9 @@ 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)
|
|
@@ -56,10 +64,13 @@ class HybridMultimodalModel(nn.Module):
|
|
| 56 |
self.head_intensity = nn.Linear(512, NUM_INTENSITY)
|
| 57 |
self.head_tempo = nn.Linear(512, NUM_TEMPO)
|
| 58 |
|
| 59 |
-
def forward(self, input_ids, attention_mask, audio_values):
|
| 60 |
text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
|
| 61 |
text_feat = text_out.pooler_output
|
| 62 |
|
|
|
|
|
|
|
|
|
|
| 63 |
audio_out = self.audio_model(audio_values).last_hidden_state
|
| 64 |
audio_feat = self.audio_proj(audio_out.mean(dim=1))
|
| 65 |
|
|
@@ -83,24 +94,86 @@ model.eval()
|
|
| 83 |
|
| 84 |
tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')
|
| 85 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 86 |
# ββ HYBRID INFERENCE PIPELINE ββ
|
| 87 |
-
def predict_multimodal(audio_path,
|
| 88 |
-
|
| 89 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
|
| 91 |
try:
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
|
| 96 |
-
|
| 97 |
-
wav, sr = librosa.load(audio_path, sr=16000)
|
| 98 |
|
| 99 |
onset_env = librosa.onset.onset_strength(y=wav, sr=sr)
|
| 100 |
tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=sr)
|
| 101 |
print(f"Calculated BPM: {tempo_bpm}")
|
| 102 |
|
| 103 |
-
# 3. Chunking (20-second sequential blocks)
|
| 104 |
chunk_size = 16000 * 20
|
| 105 |
chunks = []
|
| 106 |
for i in range(0, len(wav), chunk_size):
|
|
@@ -114,15 +187,13 @@ def predict_multimodal(audio_path, lyrics):
|
|
| 114 |
|
| 115 |
all_logits = {col: [] for col in TARGET_COLS}
|
| 116 |
|
| 117 |
-
# 4. Block-by-block Prediction
|
| 118 |
with torch.no_grad():
|
| 119 |
for chunk in chunks:
|
| 120 |
audio_tensor = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
|
| 121 |
-
out = model(enc['input_ids'], enc['attention_mask'], audio_tensor)
|
| 122 |
for col in TARGET_COLS:
|
| 123 |
all_logits[col].append(out[col][0])
|
| 124 |
|
| 125 |
-
# 5. Averaging and Formatting Results
|
| 126 |
final_results = []
|
| 127 |
encoders_map = {
|
| 128 |
'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe,
|
|
@@ -133,7 +204,6 @@ def predict_multimodal(audio_path, lyrics):
|
|
| 133 |
avg_logits = torch.stack(all_logits[col]).mean(dim=0)
|
| 134 |
le = encoders_map[col]
|
| 135 |
|
| 136 |
-
# TEMPO LOGIC OVERRIDE
|
| 137 |
if col == 'tempo':
|
| 138 |
tempo_classes = list(le.classes_)
|
| 139 |
try:
|
|
@@ -153,9 +223,16 @@ def predict_multimodal(audio_path, lyrics):
|
|
| 153 |
sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
|
| 154 |
final_results.append(sorted_preds)
|
| 155 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
return final_results
|
| 157 |
|
| 158 |
except Exception as e:
|
|
|
|
|
|
|
|
|
|
| 159 |
return [{"Error": f"Processing failed: {str(e)}"}] * 5
|
| 160 |
|
| 161 |
# ββ GRADIO UI ββ
|
|
@@ -165,8 +242,9 @@ with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
|
|
| 165 |
|
| 166 |
with gr.Row():
|
| 167 |
with gr.Column():
|
| 168 |
-
audio_input = gr.Audio(type="filepath", label="Upload
|
| 169 |
-
|
|
|
|
| 170 |
btn = gr.Button("START HYBRID ANALYSIS", variant="primary")
|
| 171 |
|
| 172 |
with gr.Column():
|
|
@@ -178,7 +256,7 @@ with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
|
|
| 178 |
|
| 179 |
btn.click(
|
| 180 |
fn=predict_multimodal,
|
| 181 |
-
inputs=[audio_input, lyrics_input],
|
| 182 |
outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
|
| 183 |
)
|
| 184 |
|
|
|
|
| 5 |
import gradio as gr
|
| 6 |
from transformers import XLMRobertaModel, XLMRobertaTokenizer, WavLMModel
|
| 7 |
from huggingface_hub import hf_hub_download
|
| 8 |
+
import yt_dlp
|
| 9 |
+
import os
|
| 10 |
+
import re
|
| 11 |
+
import lyricsgenius
|
| 12 |
+
import syncedlyrics
|
| 13 |
import warnings
|
| 14 |
|
| 15 |
warnings.filterwarnings('ignore')
|
|
|
|
| 19 |
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 20 |
TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']
|
| 21 |
|
| 22 |
+
# Set your Genius token here or use HF Secrets/Environment Variables
|
| 23 |
+
GENIUS_TOKEN = os.environ.get("GENIUS_TOKEN", "z2XGBWXalGUtAdC1qxxXBxUnK1ZuoHPkCu5eP9q-fed-DW1uCJ3NSFpHemk3Unmg")
|
| 24 |
+
|
| 25 |
print("Downloading model from Hugging Face Hub...")
|
| 26 |
model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
|
| 27 |
ckpt = torch.load(model_path, map_location=DEVICE, weights_only=False)
|
|
|
|
| 64 |
self.head_intensity = nn.Linear(512, NUM_INTENSITY)
|
| 65 |
self.head_tempo = nn.Linear(512, NUM_TEMPO)
|
| 66 |
|
| 67 |
+
def forward(self, input_ids, attention_mask, audio_values, text_missing=False):
|
| 68 |
text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
|
| 69 |
text_feat = text_out.pooler_output
|
| 70 |
|
| 71 |
+
if text_missing:
|
| 72 |
+
text_feat = torch.zeros_like(text_feat)
|
| 73 |
+
|
| 74 |
audio_out = self.audio_model(audio_values).last_hidden_state
|
| 75 |
audio_feat = self.audio_proj(audio_out.mean(dim=1))
|
| 76 |
|
|
|
|
| 94 |
|
| 95 |
tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')
|
| 96 |
|
| 97 |
+
# ββ UTILITIES ββ
|
| 98 |
+
def get_audio_from_youtube(query):
|
| 99 |
+
temp_filename = "temp_downloaded_audio"
|
| 100 |
+
|
| 101 |
+
ydl_opts = {
|
| 102 |
+
'format': 'bestaudio/best',
|
| 103 |
+
'outtmpl': f'{temp_filename}.%(ext)s',
|
| 104 |
+
'postprocessors': [{
|
| 105 |
+
'key': 'FFmpegExtractAudio',
|
| 106 |
+
'preferredcodec': 'wav',
|
| 107 |
+
'preferredquality': '192',
|
| 108 |
+
}],
|
| 109 |
+
'noplaylist': True,
|
| 110 |
+
'quiet': True,
|
| 111 |
+
'default_search': 'ytsearch1',
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
try:
|
| 115 |
+
with yt_dlp.YoutubeDL(ydl_opts) as ydl:
|
| 116 |
+
ydl.download([query])
|
| 117 |
+
return f"{temp_filename}.wav"
|
| 118 |
+
except Exception as e:
|
| 119 |
+
print(f"Error yt-dlp: {str(e)}")
|
| 120 |
+
return None
|
| 121 |
+
|
| 122 |
+
def get_lyrics_fallback(query):
|
| 123 |
+
try:
|
| 124 |
+
genius = lyricsgenius.Genius(GENIUS_TOKEN, verbose=False)
|
| 125 |
+
song = genius.search_song(query)
|
| 126 |
+
if song and song.lyrics:
|
| 127 |
+
clean_text = re.sub(r'\[.*?\]', '', song.lyrics)
|
| 128 |
+
return clean_text.strip()
|
| 129 |
+
except Exception as e:
|
| 130 |
+
print(f"Genius API Error: {str(e)}")
|
| 131 |
+
|
| 132 |
+
try:
|
| 133 |
+
lrc_lyrics = syncedlyrics.search(query)
|
| 134 |
+
if lrc_lyrics:
|
| 135 |
+
clean_text = re.sub(r'\[\d{2}:\d{2}\.\d{2}\]', '', lrc_lyrics)
|
| 136 |
+
return clean_text.strip()
|
| 137 |
+
except Exception as e:
|
| 138 |
+
print(f"SyncedLyrics Error: {str(e)}")
|
| 139 |
+
|
| 140 |
+
return None
|
| 141 |
+
|
| 142 |
# ββ HYBRID INFERENCE PIPELINE ββ
|
| 143 |
+
def predict_multimodal(audio_path, yt_query, lyrics_input):
|
| 144 |
+
target_audio = audio_path
|
| 145 |
+
|
| 146 |
+
if target_audio is None and yt_query:
|
| 147 |
+
print(f"Searching audio on YouTube for: {yt_query}")
|
| 148 |
+
target_audio = get_audio_from_youtube(yt_query)
|
| 149 |
+
|
| 150 |
+
if target_audio is None:
|
| 151 |
+
return [{"Error": "Audio file or YouTube query is required."}] * 5
|
| 152 |
|
| 153 |
try:
|
| 154 |
+
final_lyrics = lyrics_input
|
| 155 |
+
if not final_lyrics or str(final_lyrics).strip() == "":
|
| 156 |
+
if yt_query:
|
| 157 |
+
print(f"Fetching lyrics for: {yt_query}")
|
| 158 |
+
fetched_lyrics = get_lyrics_fallback(yt_query)
|
| 159 |
+
if fetched_lyrics:
|
| 160 |
+
final_lyrics = fetched_lyrics
|
| 161 |
+
|
| 162 |
+
is_instrumental = False
|
| 163 |
+
if not final_lyrics or str(final_lyrics).strip() == "":
|
| 164 |
+
is_instrumental = True
|
| 165 |
+
text_to_encode = "[INSTRUMENTAL]"
|
| 166 |
+
else:
|
| 167 |
+
text_to_encode = str(final_lyrics).strip()
|
| 168 |
+
|
| 169 |
+
enc = tokenizer(text_to_encode, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)
|
| 170 |
|
| 171 |
+
wav, sr = librosa.load(target_audio, sr=16000)
|
|
|
|
| 172 |
|
| 173 |
onset_env = librosa.onset.onset_strength(y=wav, sr=sr)
|
| 174 |
tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=sr)
|
| 175 |
print(f"Calculated BPM: {tempo_bpm}")
|
| 176 |
|
|
|
|
| 177 |
chunk_size = 16000 * 20
|
| 178 |
chunks = []
|
| 179 |
for i in range(0, len(wav), chunk_size):
|
|
|
|
| 187 |
|
| 188 |
all_logits = {col: [] for col in TARGET_COLS}
|
| 189 |
|
|
|
|
| 190 |
with torch.no_grad():
|
| 191 |
for chunk in chunks:
|
| 192 |
audio_tensor = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
|
| 193 |
+
out = model(enc['input_ids'], enc['attention_mask'], audio_tensor, text_missing=is_instrumental)
|
| 194 |
for col in TARGET_COLS:
|
| 195 |
all_logits[col].append(out[col][0])
|
| 196 |
|
|
|
|
| 197 |
final_results = []
|
| 198 |
encoders_map = {
|
| 199 |
'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe,
|
|
|
|
| 204 |
avg_logits = torch.stack(all_logits[col]).mean(dim=0)
|
| 205 |
le = encoders_map[col]
|
| 206 |
|
|
|
|
| 207 |
if col == 'tempo':
|
| 208 |
tempo_classes = list(le.classes_)
|
| 209 |
try:
|
|
|
|
| 223 |
sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
|
| 224 |
final_results.append(sorted_preds)
|
| 225 |
|
| 226 |
+
if yt_query and target_audio == "temp_downloaded_audio.wav":
|
| 227 |
+
if os.path.exists(target_audio):
|
| 228 |
+
os.remove(target_audio)
|
| 229 |
+
|
| 230 |
return final_results
|
| 231 |
|
| 232 |
except Exception as e:
|
| 233 |
+
if yt_query and target_audio == "temp_downloaded_audio.wav":
|
| 234 |
+
if os.path.exists(target_audio):
|
| 235 |
+
os.remove(target_audio)
|
| 236 |
return [{"Error": f"Processing failed: {str(e)}"}] * 5
|
| 237 |
|
| 238 |
# ββ GRADIO UI ββ
|
|
|
|
| 242 |
|
| 243 |
with gr.Row():
|
| 244 |
with gr.Column():
|
| 245 |
+
audio_input = gr.Audio(type="filepath", label="1. Upload Audio File (Optional, overrides YouTube search)")
|
| 246 |
+
yt_input = gr.Textbox(lines=1, label="2. Search YouTube (Format: Artist - Song Title)", placeholder="e.g. eleventwelfth - front-and-centre")
|
| 247 |
+
lyrics_input = gr.Textbox(lines=5, label="3. Manual Lyrics Input (Optional, auto-fetched if empty)", placeholder="Paste lyrics here or leave empty for instrumental / auto-fetch...")
|
| 248 |
btn = gr.Button("START HYBRID ANALYSIS", variant="primary")
|
| 249 |
|
| 250 |
with gr.Column():
|
|
|
|
| 256 |
|
| 257 |
btn.click(
|
| 258 |
fn=predict_multimodal,
|
| 259 |
+
inputs=[audio_input, yt_input, lyrics_input],
|
| 260 |
outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
|
| 261 |
)
|
| 262 |
|