| 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: |
| |
| feature_extractor = ASTFeatureExtractor.from_pretrained(MODEL_ID) |
| model = ASTForAudioClassification.from_pretrained(MODEL_ID) |
| except Exception as e: |
| |
| 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 |
| |
| 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=[], |
| flagging_mode="never" |
| ) |
|
|
| if __name__ == "__main__": |
| demo.launch() |
|
|