Spaces:
Sleeping
Sleeping
Download app.py from mohsinramzan/AudioClassification: direct link, hf CLI and curl.
- Browser
- Download file 7.33 kB
-
https://huggingface.co/spaces/mohsinramzan/AudioClassification/resolve/main/app.py
- Command line
-
hf download hf://spaces/mohsinramzan/AudioClassification/app.py
-
curl -L -o app.py https://huggingface.co/spaces/mohsinramzan/AudioClassification/resolve/main/app.py
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 βββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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 ββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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 βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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 βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def root(): | |
| return {"status": "ok", "message": "Audio Classifier is running"} | |
| def health(): | |
| return {"status": "healthy", "model": "CrossAttnModel", "classes": CLASSES} |