File size: 7,329 Bytes
d3e1c6c
 
 
 
 
 
 
 
5fc4dae
 
d3e1c6c
46cb848
5fc4dae
 
e007778
5fc4dae
 
 
 
 
 
 
 
e007778
 
 
 
 
 
 
 
 
 
5fc4dae
e007778
 
 
 
 
 
 
5fc4dae
e007778
5fc4dae
 
e007778
 
 
46cb848
 
d3e1c6c
 
 
 
 
46cb848
d3e1c6c
 
46cb848
d3e1c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
46cb848
 
 
 
 
 
d3e1c6c
 
 
 
 
 
 
 
 
 
5fc4dae
d3e1c6c
46cb848
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d3e1c6c
 
46cb848
 
 
d3e1c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
46cb848
 
 
 
 
 
 
 
 
 
d3e1c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
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}