MJ-Cat / app /main.py
Irell Zane
feat: set up service for MJ-Cat
e6c076c
Raw
History Blame Contribute Delete
4.16 kB
from fastapi import FastAPI, UploadFile, File, HTTPException
from pathlib import Path
import numpy as np
import tensorflow_hub as hub
import tensorflow as tf
import librosa
import io
import joblib
from loguru import logger
from pydub import AudioSegment
from rich.logging import RichHandler
from fastapi.middleware.cors import CORSMiddleware
# -------------------
# Logging
# -------------------
logger.remove()
logger.add(RichHandler(), level="INFO")
# -------------------
# Paths & constants
# -------------------
MODEL_DIR = Path("models")
VGGISH_MODEL_URL = "https://tfhub.dev/google/vggish/1"
SR = 16000
# -------------------
# FastAPI instance
# -------------------
app = FastAPI(title="Audio Embedding & Classification API")
# -------------------
# CORS settings
# -------------------
origins = [
"https://zane-dev16.github.io/MJ-Cat-Frontend/",
]
app.add_middleware(
CORSMiddleware,
allow_origins=origins, # or ["*"] for all origins
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# -------------------
# Preload models
# -------------------
logger.info("Loading VGGish model from TF Hub...")
vggish_model = hub.load(VGGISH_MODEL_URL)
logger.info("VGGish model loaded successfully.")
lgbm_model_path = MODEL_DIR / "lgbm_model.pkl"
if not lgbm_model_path.exists():
raise FileNotFoundError(f"LightGBM model not found at {lgbm_model_path}")
logger.info("Loading LightGBM model...")
lgbm_model = joblib.load(lgbm_model_path)
logger.info("LightGBM model loaded successfully.")
# -------------------
# Helper: Audio preprocessing
# -------------------
def preprocess_audio(file_bytes: bytes) -> tf.Tensor:
waveform, _ = librosa.load(io.BytesIO(file_bytes), sr=SR, mono=True)
return tf.convert_to_tensor(np.array(waveform, dtype=np.float32))
def extract_embedding(tensor: tf.Tensor) -> np.ndarray:
embedding = vggish_model(tensor).numpy()
feature_vector = np.concatenate([
np.mean(embedding, axis=0),
np.std(embedding, axis=0),
np.max(embedding, axis=0)
])
return feature_vector
# -------------------
# API endpoints
# -------------------
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
logger.info(f"Received file: {file.filename}")
file_bytes = await file.read()
if file.filename.endswith(".webm"):
try:
audio = AudioSegment.from_file(io.BytesIO(file_bytes), format="webm")
wav_io = io.BytesIO()
audio.export(wav_io, format="wav")
wav_io.seek(0)
file_bytes = wav_io.read()
logger.info(f"Converted WebM to WAV: {file.filename}")
except Exception as e:
logger.error(f"Failed to convert WebM to WAV: {e}")
raise HTTPException(status_code=500, detail="Failed to convert WebM to WAV")
elif not file.filename.endswith(".wav"):
raise HTTPException(status_code=400, detail="Only WAV or WebM files are supported")
try:
tensor = preprocess_audio(file_bytes)
features = extract_embedding(tensor).reshape(1, -1)
proba = lgbm_model.predict(features)[0]
prediction = int(proba >= 0.5)
logger.info(f"Prediction: {prediction} (prob={proba:.2f})")
return {"prediction": prediction, "probability": float(proba)}
except Exception as e:
logger.error(f"Failed to process file {file.filename}: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/extract_embedding")
async def get_embedding(file: UploadFile = File(...)):
if not file.filename.endswith(".wav"):
raise HTTPException(status_code=400, detail="Only WAV files are supported")
try:
file_bytes = await file.read()
tensor = preprocess_audio(file_bytes)
embedding = extract_embedding(tensor)
return {"embedding": embedding.tolist()}
except Exception as e:
logger.error(f"Failed to process file {file.filename}: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.get("/")
def health_check():
return {"status": "ok", "message": "Service running with VGGish preloaded."}