anggars commited on
Commit
8cc551e
Β·
verified Β·
1 Parent(s): b95dfa3

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +95 -17
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, lyrics):
88
- if audio_path is None:
89
- return [{"Error": "Audio file cannot be empty."}] * 5
 
 
 
 
 
 
90
 
91
  try:
92
- # 1. Text Processing
93
- text = str(lyrics).strip() if lyrics else "[INSTRUMENTAL]"
94
- enc = tokenizer(text, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)
 
 
 
 
 
 
 
 
 
 
 
 
 
95
 
96
- # 2. Audio Processing & DSP
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 Song (Full Analysis)")
169
- lyrics_input = gr.Textbox(lines=5, label="Lyrics Content", placeholder="Paste lyrics here...")
 
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