jananiramaseshan's picture
Update app.py
7b06d08 verified
Raw
History Blame Contribute Delete
2.79 kB
import gradio as gr
import torch
import torchaudio
import numpy as np
from transformers import ASTFeatureExtractor, ASTForAudioClassification
import os
MODEL_ID = "jananiramaseshan/ast-music-genre-classifier"
def load_assets():
try:
# Try loading from Hub
feature_extractor = ASTFeatureExtractor.from_pretrained(MODEL_ID)
model = ASTForAudioClassification.from_pretrained(MODEL_ID)
except Exception as e:
# Fallback to local directory if on a Space with model files
print(f"Loading from Hub failed: {e}. Trying local...")
feature_extractor = ASTFeatureExtractor.from_pretrained("./")
model = ASTForAudioClassification.from_pretrained("./")
return feature_extractor, model
extractor, model = load_assets()
id2label = model.config.id2label
def predict_genre(audio):
if audio is None:
return None
sr, data = audio
if data.dtype == np.int16:
data = data.astype(np.float32) / 32768.0
elif data.dtype == np.int32:
data = data.astype(np.float32) / 2147483648.0
waveform = torch.from_numpy(data.astype(np.float32))
if waveform.ndim > 1:
waveform = waveform.mean(dim=1, keepdim=True).T
else:
waveform = waveform.unsqueeze(0)
if sr != 16000:
resampler = torchaudio.transforms.Resample(sr, 16000)
waveform = resampler(waveform)
full_audio = waveform.squeeze(0).numpy()
target_samples = 16000 * 10 # 10 seconds segment
audio_len = len(full_audio)
if audio_len > target_samples:
offsets = [0, (audio_len - target_samples)//2, audio_len - target_samples]
else:
offsets = [0]
all_logits = []
for offset in offsets:
window = full_audio[offset:offset+target_samples]
if len(window) < target_samples:
window = np.pad(window, (0, target_samples - len(window)))
inputs = extractor(window, sampling_rate=16000, return_tensors="pt")
with torch.no_grad():
outputs = model(**inputs)
all_logits.append(outputs.logits)
avg_logits = torch.stack(all_logits).mean(dim=0)
probs = torch.softmax(avg_logits, dim=-1)[0]
results = {id2label[i]: float(probs[i]) for i in range(len(probs))}
return results
demo = gr.Interface(
fn=predict_genre,
inputs=gr.Audio(type="numpy", label="Upload Song"),
outputs=gr.Label(num_top_classes=5, label="Predicted Genre"),
title="🎵 Music Genre Classifier",
description="Upload a song file (WAV, MP3, etc.) to classify its genre using the Audio Spectrogram Transformer (AST).",
examples=[], # You can add example audio files here
flagging_mode="never"
)
if __name__ == "__main__":
demo.launch()