import os import pickle import tempfile import numpy as np import librosa import tensorflow as tf from fastapi import FastAPI, UploadFile, File, HTTPException from fastapi.middleware.cors import CORSMiddleware from tensorflow.keras import layers import keras # ── Custom SpecAugment layer ─────────────────────────────────────────────────── @keras.saving.register_keras_serializable() class SpecAugment(layers.Layer): def __init__(self, time_mask=20, freq_mask=15, **kwargs): super().__init__(**kwargs) self.time_mask = time_mask self.freq_mask = freq_mask def call(self, x, training=None): if not training: return x shape = tf.shape(x) B = shape[0] T = shape[1] F = shape[2] # ── Time mask ───────────────────────────────────────────────────── t = tf.random.uniform((), minval=0, maxval=self.time_mask, dtype=tf.int32) t0 = tf.random.uniform((), minval=0, maxval=T - t, dtype=tf.int32) time_mask = tf.concat([ tf.ones ([1, t0, F], dtype=x.dtype), tf.zeros([1, t, F], dtype=x.dtype), tf.ones ([1, T - t0 - t, F], dtype=x.dtype), ], axis=1) # ── Freq mask ───────────────────────────────────────────────────── f = tf.random.uniform((), minval=0, maxval=self.freq_mask, dtype=tf.int32) f0 = tf.random.uniform((), minval=0, maxval=F - f, dtype=tf.int32) freq_mask = tf.concat([ tf.ones ([1, T, f0 ], dtype=x.dtype), tf.zeros([1, T, f ], dtype=x.dtype), tf.ones ([1, T, F - f0 - f], dtype=x.dtype), ], axis=2) return x * time_mask * freq_mask def get_config(self): cfg = super().get_config() cfg.update({"time_mask": self.time_mask, "freq_mask": self.freq_mask}) return cfg # ── Constants ────────────────────────────────────────────────────────────────── SR = 22050 DURATION = 3 HOP_LENGTH = 512 TIME_FRAMES = 130 N_BINS = 84 N_MFCC = 40 # 40 × 3 (+ delta + delta2) = 120 features CLASSES = ["Muwa", "east", "loyal", "poems", "rai"] # ── App ──────────────────────────────────────────────────────────────────────── app = FastAPI(title="Audio Classifier API") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"], ) # ── Load artifacts on startup ────────────────────────────────────────────────── @app.on_event("startup") async def load_artifacts(): global model, scaler_cqt, scaler_mfcc, label_encoder model = tf.keras.models.load_model( "model_artifacts/best_cross_attn.keras", custom_objects={"SpecAugment": SpecAugment}, compile=False ) print("Model loaded") with open("model_artifacts/scaler_cqt.pkl", "rb") as f: scaler_cqt = pickle.load(f) with open("model_artifacts/scaler_mfcc.pkl", "rb") as f: scaler_mfcc = pickle.load(f) with open("model_artifacts/label_encoder.pkl", "rb") as f: label_encoder = pickle.load(f) print(" Scalers and label encoder loaded") # ── Feature extraction (mirrors training pipeline exactly) ───────────────────── def extract_features_from_array(y): # Pad or trim to fixed duration expected = SR * DURATION if len(y) < expected: y = np.pad(y, (0, expected - len(y))) else: y = y[:expected] # CQT → amplitude_to_db → (T, 84) cqt = librosa.cqt(y=y, sr=SR, hop_length=HOP_LENGTH, n_bins=N_BINS, bins_per_octave=12) cqt = librosa.amplitude_to_db(np.abs(cqt)).T # (T, 84) # MFCC + delta + delta2 → (T, 120) mfcc = librosa.feature.mfcc(y=y, sr=SR, n_mfcc=N_MFCC, hop_length=HOP_LENGTH) delta = librosa.feature.delta(mfcc) delta2 = librosa.feature.delta(mfcc, order=2) mfcc = np.vstack([mfcc, delta, delta2]).T # (T, 120) # Align & pad/trim to TIME_FRAMES min_t = min(cqt.shape[0], mfcc.shape[0]) cqt, mfcc = cqt[:min_t], mfcc[:min_t] if min_t < TIME_FRAMES: cqt = np.pad(cqt, ((0, TIME_FRAMES - min_t), (0, 0))) mfcc = np.pad(mfcc, ((0, TIME_FRAMES - min_t), (0, 0))) else: cqt = cqt[:TIME_FRAMES] mfcc = mfcc[:TIME_FRAMES] return cqt, mfcc # (130, 84), (130, 120) def extract_features(audio_path: str): y, _ = librosa.load(audio_path, sr=SR, duration=DURATION) return extract_features_from_array(y) # ── Predict endpoint ─────────────────────────────────────────────────────────── @app.post("/predict") async def predict(file: UploadFile = File(...)): if not file.filename.endswith((".wav", ".mp3", ".flac", ".ogg")): raise HTTPException(status_code=400, detail="Unsupported format. Use wav, mp3, flac, or ogg.") with tempfile.NamedTemporaryFile(delete=False, suffix=os.path.splitext(file.filename)[1]) as tmp: tmp.write(await file.read()) tmp_path = tmp.name try: cqt, mfcc = extract_features(tmp_path) # Apply scalers exactly as in training cqt_shape = cqt.shape mfcc_shape = mfcc.shape cqt = scaler_cqt.transform(cqt.reshape(-1, 84)).reshape(cqt_shape) mfcc = scaler_mfcc.transform(mfcc.reshape(-1, 120)).reshape(mfcc_shape) # Add batch dim preds = model.predict([cqt[np.newaxis], mfcc[np.newaxis]])[0] top_idx = int(np.argmax(preds)) try: label = label_encoder.inverse_transform([top_idx])[0] except Exception: label = CLASSES[top_idx] return { "label": label, "confidence": round(float(preds[top_idx]), 4), "all_scores": { CLASSES[i]: round(float(p), 4) for i, p in enumerate(preds) } } except Exception as e: raise HTTPException(status_code=500, detail=f"Inference error: {str(e)}") finally: os.remove(tmp_path) # ── Health endpoints ─────────────────────────────────────────────────────────── @app.get("/") def root(): return {"status": "ok", "message": "Audio Classifier is running"} @app.get("/health") def health(): return {"status": "healthy", "model": "CrossAttnModel", "classes": CLASSES}