""" Music Genre Classification — 5-Model Ensemble Hugging Face Spaces Deployment (Gradio) Public F1 Score: 0.801 """ import os, torch, numpy as np, librosa from scipy.ndimage import zoom as scipyzoom import torch.nn as nn import gradio as gr # ── Constants ────────────────────────────────────────────────────────── NCLASSES = 10 MAXFRAMES = 256 SR = 22050 NMELS = 128 NFFT = 2048 HOP = 512 GENRES = ["blues", "classical", "country", "disco", "hiphop", "jazz", "metal", "pop", "reggae", "rock"] VAL_F1 = [0.8255, 0.8391, 0.7601, 0.7966, 0.7872] _raw = np.exp(np.array(VAL_F1, dtype=np.float32) * 10) ENS_WEIGHTS = torch.FloatTensor(_raw / _raw.sum()).view(-1, 1, 1) device = torch.device("cpu") # ── Building Blocks ─────────────────────────────────────────────────── class ResBlock(nn.Module): def __init__(self, c_in, c_out, stride=1): super().__init__() self.body = nn.Sequential( nn.Conv2d(c_in, c_out, 3, stride=stride, padding=1, bias=False), nn.BatchNorm2d(c_out), nn.ReLU(True), nn.Conv2d(c_out, c_out, 3, padding=1, bias=False), nn.BatchNorm2d(c_out)) self.skip = (nn.Sequential( nn.Conv2d(c_in, c_out, 1, stride=stride, bias=False), nn.BatchNorm2d(c_out)) if stride != 1 or c_in != c_out else nn.Identity()) self.act = nn.ReLU(True) def forward(self, x): return self.act(self.body(x) + self.skip(x)) class SEBlock(nn.Module): def __init__(self, c, ratio=4): super().__init__() h = max(4, c // ratio) self.pool = nn.AdaptiveAvgPool2d(1) self.fc = nn.Sequential( nn.Flatten(), nn.Linear(c, h), nn.ReLU(True), nn.Linear(h, c), nn.Sigmoid()) def forward(self, x): return x * self.fc(self.pool(x)).view(x.size(0), -1, 1, 1) class SEResBlock(nn.Module): def __init__(self, c_in, c_out, stride=1, ratio=4): super().__init__() self.body = nn.Sequential( nn.Conv2d(c_in, c_out, 3, stride=stride, padding=1, bias=False), nn.BatchNorm2d(c_out), nn.ReLU(True), nn.Conv2d(c_out, c_out, 3, padding=1, bias=False), nn.BatchNorm2d(c_out)) self.se = SEBlock(c_out, ratio) self.skip = (nn.Sequential( nn.Conv2d(c_in, c_out, 1, stride=stride, bias=False), nn.BatchNorm2d(c_out)) if stride != 1 or c_in != c_out else nn.Identity()) self.act = nn.ReLU(True) def forward(self, x): return self.act(self.se(self.body(x)) + self.skip(x)) # ── Model Definitions ──────────────────────────────────────────────── class FreqFirstNet(nn.Module): def __init__(self): super().__init__() self.freqcnn = nn.Sequential( nn.Conv2d(3, 32, (8,1), stride=(4,1), padding=(2,0), bias=False), nn.BatchNorm2d(32), nn.ReLU(True), nn.Conv2d(32, 64, (8,1), stride=(4,1), padding=(2,0), bias=False), nn.BatchNorm2d(64), nn.ReLU(True), nn.AdaptiveAvgPool2d((1, MAXFRAMES))) self.tempcnn = nn.Sequential( nn.Conv1d(64, 64, 7, padding=3, bias=False), nn.BatchNorm1d(64), nn.ReLU(True), nn.MaxPool1d(4), nn.Conv1d(64, 128, 5, padding=2, bias=False), nn.BatchNorm1d(128), nn.ReLU(True), nn.AdaptiveAvgPool1d(1), nn.Flatten()) self.fc = nn.Sequential(nn.Dropout(0.5), nn.Linear(128, NCLASSES)) def forward(self, x): return self.fc(self.tempcnn(self.freqcnn(x).squeeze(2))) class FreqFirstNetV2(nn.Module): def __init__(self): super().__init__() self.freqcnn = nn.Sequential( nn.Conv2d(3, 48, (8,1), stride=(4,1), padding=(2,0), bias=False), nn.BatchNorm2d(48), nn.ReLU(True), nn.Conv2d(48, 96, (8,1), stride=(4,1), padding=(2,0), bias=False), nn.BatchNorm2d(96), nn.ReLU(True), nn.AdaptiveAvgPool2d((1, MAXFRAMES))) self.tempcnn = nn.Sequential( nn.Conv1d(96, 96, 5, padding=2, bias=False), nn.BatchNorm1d(96), nn.ReLU(True), nn.MaxPool1d(4), nn.Conv1d(96, 192, 3, padding=1, bias=False), nn.BatchNorm1d(192), nn.ReLU(True), nn.AdaptiveAvgPool1d(1), nn.Flatten()) self.fc = nn.Sequential(nn.Dropout(0.5), nn.Linear(192, NCLASSES)) def forward(self, x): return self.fc(self.tempcnn(self.freqcnn(x).squeeze(2))) class DeepResNet(nn.Module): def __init__(self): super().__init__() self.stem = nn.Sequential( nn.Conv2d(3, 32, 5, stride=2, padding=2, bias=False), nn.BatchNorm2d(32), nn.ReLU(True), nn.MaxPool2d(2, 2)) self.stage1 = nn.Sequential(ResBlock(32, 32), ResBlock(32, 32)) self.stage2 = nn.Sequential(ResBlock(32, 48, stride=2), ResBlock(48, 48)) self.stage3 = nn.Sequential(ResBlock(48, 64, stride=2), ResBlock(64, 64)) self.head = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Dropout(0.4), nn.Linear(64, NCLASSES)) def forward(self, x): return self.head(self.stage3(self.stage2(self.stage1(self.stem(x))))) class TinyResNet(nn.Module): def __init__(self): super().__init__() self.stem = nn.Sequential( nn.Conv2d(3, 32, 5, stride=2, padding=2, bias=False), nn.BatchNorm2d(32), nn.ReLU(True), nn.MaxPool2d(2, 2)) self.stage1 = nn.Sequential(ResBlock(32, 32), ResBlock(32, 32)) self.stage2 = nn.Sequential(ResBlock(32, 64, stride=2), ResBlock(64, 64)) self.head = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Dropout(0.5), nn.Linear(64, NCLASSES)) def forward(self, x): return self.head(self.stage2(self.stage1(self.stem(x)))) class TinySENet(nn.Module): def __init__(self): super().__init__() self.stem = nn.Sequential( nn.Conv2d(3, 32, 5, stride=2, padding=2, bias=False), nn.BatchNorm2d(32), nn.ReLU(True), nn.MaxPool2d(2, 2)) self.stage1 = nn.Sequential(SEResBlock(32, 32), SEResBlock(32, 32)) self.stage2 = nn.Sequential(SEResBlock(32, 64, stride=2), SEResBlock(64, 64)) self.head = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Dropout(0.5), nn.Linear(64, NCLASSES)) def forward(self, x): return self.head(self.stage2(self.stage1(self.stem(x)))) # ── Load Models ────────────────────────────────────────────────────── MODEL_DIR = os.path.dirname(os.path.abspath(__file__)) MODELS_INFO = [ ("FreqFirstNet", FreqFirstNet, "freqfirstnet_best.pth"), ("FreqFirstNetV2", FreqFirstNetV2, "freqfirstnetv2_best.pth"), ("DeepResNet", DeepResNet, "deepresnet_best.pth"), ("TinyResNet", TinyResNet, "tinyresnet_best.pth"), ("TinySENet", TinySENet, "tinysenet_best.pth"), ] models = [] for name, cls, ckpt in MODELS_INFO: model = cls().to(device) ckpt_path = os.path.join(MODEL_DIR, ckpt) state = torch.load(ckpt_path, map_location=device, weights_only=True) model.load_state_dict(state) model.eval() models.append(model) print(f"Loaded {name}") print("All 5 models loaded.") # ── Feature Extraction ────────────────────────────────────────────── def extract3ch(audio, rate=SR): mel = librosa.feature.melspectrogram( y=audio, sr=rate, n_mels=NMELS, n_fft=NFFT, hop_length=HOP) meldb = librosa.power_to_db(mel, ref=np.max) # Compute deltas on the ORIGINAL time resolution (before zoom) d1 = librosa.feature.delta(meldb, order=1) d2 = librosa.feature.delta(meldb, order=2) # Now zoom all 3 channels to MAXFRAMES if meldb.shape[1] != MAXFRAMES: zf = MAXFRAMES / meldb.shape[1] meldb = scipyzoom(meldb, (1, zf), order=1) d1 = scipyzoom(d1, (1, zf), order=1) d2 = scipyzoom(d2, (1, zf), order=1) if meldb.shape[1] > MAXFRAMES: meldb = meldb[:, :MAXFRAMES] d1 = d1[:, :MAXFRAMES] d2 = d2[:, :MAXFRAMES] elif meldb.shape[1] < MAXFRAMES: pad_w = MAXFRAMES - meldb.shape[1] meldb = np.pad(meldb, ((0, 0), (0, pad_w)), constant_values=meldb.min()) d1 = np.pad(d1, ((0, 0), (0, pad_w)), constant_values=0) d2 = np.pad(d2, ((0, 0), (0, pad_w)), constant_values=0) return np.stack([meldb, d1, d2], axis=0).astype(np.float32) # ── Prediction Function ───────────────────────────────────────────── def predict(audio_input): if audio_input is None: return {g: 0.0 for g in GENRES} # Gradio may pass filepath string or (sr, numpy_array) tuple if isinstance(audio_input, tuple): import soundfile as sf sr_in, audio_data = audio_input audio = audio_data.astype(np.float32) if audio.ndim > 1: audio = audio.mean(axis=1) if sr_in != SR: audio = librosa.resample(audio, orig_sr=sr_in, target_sr=SR) else: audio, _ = librosa.load(audio_input, sr=SR) spec = extract3ch(audio) spec_t = torch.FloatTensor(spec) for c in range(spec_t.shape[0]): mu = spec_t[c].mean() sg = spec_t[c].std() + 1e-8 spec_t[c] = (spec_t[c] - mu) / sg spec_t = spec_t.unsqueeze(0) logit_list = [] with torch.no_grad(): for model in models: logit_list.append(model(spec_t)) stacked = torch.stack(logit_list, 0) ensemble = (stacked * ENS_WEIGHTS).sum(0) probs = torch.softmax(ensemble, dim=1)[0] return {genre: float(prob) for genre, prob in zip(GENRES, probs)} # ── Gradio Interface ───────────────────────────────────────────────── demo = gr.Interface( fn=predict, inputs=gr.Audio(type="filepath", label="Upload Audio File (.wav, .mp3, etc.)"), outputs=gr.Label(num_top_classes=10, label="Genre Prediction"), title="Music Genre Classifier", description=( "Upload an audio file to classify its music genre using a " "**5-model ensemble** (FreqFirstNet, FreqFirstNetV2, DeepResNet, " "TinyResNet, TinySENet).\n\n" "**Features:** 3-channel mel spectrogram (mel + delta + delta-squared)\n\n" "**Genres:** blues, classical, country, disco, hiphop, jazz, metal, " "pop, reggae, rock\n\n" "**Public F1 Score:** 0.801\n\n" "Use the **API** tab at the bottom of this page for programmatic access." ), api_name="predict", ) if __name__ == "__main__": demo.launch(show_error=True)