Spaces:
Sleeping
Sleeping
Fixing a bug
Browse files
app.py
CHANGED
|
@@ -6,7 +6,52 @@ import librosa
|
|
| 6 |
import tensorflow as tf
|
| 7 |
from fastapi import FastAPI, UploadFile, File, HTTPException
|
| 8 |
from fastapi.middleware.cors import CORSMiddleware
|
| 9 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
# ββ Constants (from model_config.json) ββββββββββββββββββββββββββββββββββββββββ
|
| 11 |
SR = 22050
|
| 12 |
DURATION = 3
|
|
@@ -31,8 +76,8 @@ app.add_middleware(
|
|
| 31 |
async def load_artifacts():
|
| 32 |
global model, scaler_cqt, scaler_mfcc, label_encoder
|
| 33 |
|
| 34 |
-
model = tf.keras.models.load_model("model_artifacts/best_cross_attn.keras", compile=False)
|
| 35 |
-
print("
|
| 36 |
|
| 37 |
with open("model_artifacts/scaler_cqt.pkl", "rb") as f:
|
| 38 |
scaler_cqt = pickle.load(f)
|
|
@@ -43,7 +88,7 @@ async def load_artifacts():
|
|
| 43 |
with open("model_artifacts/label_encoder.pkl", "rb") as f:
|
| 44 |
label_encoder = pickle.load(f)
|
| 45 |
|
| 46 |
-
print("
|
| 47 |
|
| 48 |
# ββ Feature extraction βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 49 |
def extract_features(audio_path: str):
|
|
|
|
| 6 |
import tensorflow as tf
|
| 7 |
from fastapi import FastAPI, UploadFile, File, HTTPException
|
| 8 |
from fastapi.middleware.cors import CORSMiddleware
|
| 9 |
+
from tensorflow.keras import layers
|
| 10 |
+
import keras
|
| 11 |
+
|
| 12 |
+
# ββ Custom SpecAugment layer (must match your training definition exactly) βββββ
|
| 13 |
+
@keras.saving.register_keras_serializable()
|
| 14 |
+
class SpecAugment(layers.Layer):
|
| 15 |
+
def __init__(self, time_mask=20, freq_mask=10, **kwargs):
|
| 16 |
+
super().__init__(**kwargs)
|
| 17 |
+
self.time_mask = time_mask
|
| 18 |
+
self.freq_mask = freq_mask
|
| 19 |
+
|
| 20 |
+
def call(self, x, training=None):
|
| 21 |
+
if not training:
|
| 22 |
+
return x
|
| 23 |
+
|
| 24 |
+
shape = tf.shape(x)
|
| 25 |
+
time_steps = shape[1]
|
| 26 |
+
freq_steps = shape[2]
|
| 27 |
+
|
| 28 |
+
# Time masking
|
| 29 |
+
t = tf.random.uniform((), 0, self.time_mask, dtype=tf.int32)
|
| 30 |
+
t0 = tf.random.uniform((), 0, tf.maximum(1, time_steps - t), dtype=tf.int32)
|
| 31 |
+
time_mask_tensor = tf.concat([
|
| 32 |
+
tf.ones([shape[0], t0, freq_steps]),
|
| 33 |
+
tf.zeros([shape[0], t, freq_steps]),
|
| 34 |
+
tf.ones([shape[0], time_steps - t0 - t, freq_steps])
|
| 35 |
+
], axis=1)
|
| 36 |
+
x = x * time_mask_tensor
|
| 37 |
+
|
| 38 |
+
# Frequency masking
|
| 39 |
+
f = tf.random.uniform((), 0, self.freq_mask, dtype=tf.int32)
|
| 40 |
+
f0 = tf.random.uniform((), 0, tf.maximum(1, freq_steps - f), dtype=tf.int32)
|
| 41 |
+
freq_mask_tensor = tf.concat([
|
| 42 |
+
tf.ones([shape[0], time_steps, f0]),
|
| 43 |
+
tf.zeros([shape[0], time_steps, f]),
|
| 44 |
+
tf.ones([shape[0], time_steps, freq_steps - f0 - f])
|
| 45 |
+
], axis=2)
|
| 46 |
+
x = x * freq_mask_tensor
|
| 47 |
+
|
| 48 |
+
return x
|
| 49 |
+
|
| 50 |
+
def get_config(self):
|
| 51 |
+
config = super().get_config()
|
| 52 |
+
config.update({"time_mask": self.time_mask, "freq_mask": self.freq_mask})
|
| 53 |
+
return config
|
| 54 |
+
|
| 55 |
# ββ Constants (from model_config.json) ββββββββββββββββββββββββββββββββββββββββ
|
| 56 |
SR = 22050
|
| 57 |
DURATION = 3
|
|
|
|
| 76 |
async def load_artifacts():
|
| 77 |
global model, scaler_cqt, scaler_mfcc, label_encoder
|
| 78 |
|
| 79 |
+
model = tf.keras.models.load_model("model_artifacts/best_cross_attn.keras", custom_objects={"SpecAugment": SpecAugment}, compile=False)
|
| 80 |
+
print(" Model loaded")
|
| 81 |
|
| 82 |
with open("model_artifacts/scaler_cqt.pkl", "rb") as f:
|
| 83 |
scaler_cqt = pickle.load(f)
|
|
|
|
| 88 |
with open("model_artifacts/label_encoder.pkl", "rb") as f:
|
| 89 |
label_encoder = pickle.load(f)
|
| 90 |
|
| 91 |
+
print(" Scalers and label encoder loaded")
|
| 92 |
|
| 93 |
# ββ Feature extraction βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 94 |
def extract_features(audio_path: str):
|