mohsinramzan's picture
Fixing a lethal bug
46cb848 verified
Raw History Blame Contribute Delete
7.33 kB
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}