deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
8.31 kB
"""SONICS (SpecTTTra) Synthetic Music Detection API.
Detects AI-generated music (Suno, Udio, etc.) using the SpecTTTra
architecture from the SONICS project (ICLR 2025).
The model performs binary classification on raw audio waveforms via
internal MelSpectrogram features. It outputs a single logit; we apply
sigmoid to obtain the fake probability.
Reference: https://github.com/awsaf49/sonics
"""
import base64
import io
import logging
import os
import platform
import time
from typing import Optional
import librosa
import numpy as np
import torch
import uvicorn
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field
# Configure logging
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger("sonics_api")
# Constants
MODEL_NAME = "sonics_detection"
HF_MODEL_ID = "awsaf49/sonics-spectttra-alpha-120s"
SAMPLE_RATE = 16000
MAX_TIME = 120 # seconds (matches alpha-120s config)
MAX_LEN = MAX_TIME * SAMPLE_RATE # 1_920_000 samples
PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "true").lower() == "true"
MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "300"))
def _get_device():
"""Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
if override == "cpu":
return torch.device("cpu")
if override == "cuda" and torch.cuda.is_available():
return torch.device("cuda")
if (
override == "mps"
and hasattr(torch.backends, "mps")
and torch.backends.mps.is_available()
):
return torch.device("mps")
if override:
pass # Invalid override, fall through to auto-detect
if (
platform.system() == "Darwin"
and hasattr(torch.backends, "mps")
and torch.backends.mps.is_available()
):
return torch.device("mps")
if torch.cuda.is_available():
return torch.device("cuda")
return torch.device("cpu")
DEVICE = _get_device()
if DEVICE.type == "cuda":
torch.backends.cudnn.benchmark = True
torch.set_float32_matmul_precision("high")
if DEVICE.type == "cuda":
logger.info(
"Device: cuda (%s, %.1f GB VRAM)",
torch.cuda.get_device_name(0),
torch.cuda.get_device_properties(0).total_memory / 1024**3,
)
else:
logger.warning(
"Device: %s (no CUDA available -- check nvidia-container-toolkit)",
DEVICE,
)
# Global model instance
model = None
class AudioInput(BaseModel):
"""Request schema for audio deepfake detection."""
audio_data: str = Field(
..., description="Base64 encoded audio string (WAV/MP3/FLAC/etc)"
)
threshold: Optional[float] = Field(
0.5, ge=0.0, le=1.0, description="Classification threshold"
)
app = FastAPI(
title="SONICS Synthetic Music Detection API",
description=(
"Service for detecting AI-generated music using the "
"SpecTTTra model from the SONICS project (ICLR 2025)."
),
version="1.0.0",
)
def load_model():
"""Load the SONICS HFAudioClassifier from HuggingFace Hub.
Returns:
The loaded model, or None if loading fails.
"""
global model
if model is not None:
return model
logger.info("Loading SONICS model '%s' onto %s...", HF_MODEL_ID, DEVICE)
try:
from sonics import HFAudioClassifier
model = HFAudioClassifier.from_pretrained(
HF_MODEL_ID,
map_location=str(DEVICE),
)
model.to(DEVICE)
model.eval()
logger.info("SONICS model loaded successfully.")
return model
except Exception:
logger.exception("Failed to load SONICS model")
model = None
return None
@app.on_event("startup")
async def startup_event():
"""Optionally preload model on service startup."""
if PRELOAD_MODEL:
load_model()
@app.get("/")
async def root():
"""Root info endpoint."""
return {
"service": "SONICS Synthetic Music Detection",
"model": MODEL_NAME,
"version": "1.0.0",
}
def _gpu_health_info() -> dict:
"""Return GPU metrics for the health endpoint."""
if torch.cuda.is_available() and DEVICE.type == "cuda":
return {
"gpu_name": torch.cuda.get_device_name(0),
"vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
"vram_total_mb": round(
torch.cuda.get_device_properties(0).total_memory / 1024**2
),
}
return {}
@app.get("/health")
async def health():
"""Health check endpoint."""
return {
"status": "healthy" if model is not None else "degraded",
"model": MODEL_NAME,
"device": str(DEVICE),
**_gpu_health_info(),
}
def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
"""Preprocess audio for SONICS inference.
Loads audio from raw bytes, resamples to 16 kHz mono,
crops or zero-pads to MAX_LEN samples, and normalises by
standard deviation (matching the training pipeline).
Args:
audio_bytes: Raw audio file bytes (WAV, MP3, FLAC, etc.).
Returns:
Audio tensor of shape (1, MAX_LEN) on DEVICE.
Raises:
ValueError: If audio preprocessing fails.
"""
try:
logger.info("Starting audio preprocessing...")
audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True)
logger.info("Audio loaded. Length: %d samples at %dHz", len(audio), sr)
# Crop or pad to fixed length (matching SONICS dataset.py)
if len(audio) > MAX_LEN:
# Crop from 3/4 position (matching eval-mode logic)
idx = int((len(audio) - MAX_LEN) / 4 * 3)
audio = audio[idx : idx + MAX_LEN]
elif len(audio) < MAX_LEN:
audio = np.pad(audio, (0, MAX_LEN - len(audio)), mode="constant")
# Normalise by standard deviation (matching training pipeline)
audio /= np.maximum(np.std(audio), 1e-6)
logger.info("Audio preprocessed to %d samples", len(audio))
audio_tensor = torch.from_numpy(audio).float().unsqueeze(0)
audio_tensor = audio_tensor.to(DEVICE)
return audio_tensor
except Exception as e:
logger.error("Error preprocessing audio: %s", e)
raise ValueError(f"Audio preprocessing failed: {str(e)}")
@app.post("/predict")
async def predict(input_data: AudioInput):
"""Run synthetic music detection on base64-encoded audio.
The model uses BCEWithLogitsLoss with num_classes=1, so it
outputs a single logit. We apply sigmoid to obtain the fake
probability.
"""
if model is None:
if load_model() is None:
raise HTTPException(status_code=503, detail="Model not loaded")
try:
start_time = time.time()
logger.info(
"Prediction request. Data size: %d chars",
len(input_data.audio_data),
)
# Decode base64 audio
audio_bytes = base64.b64decode(input_data.audio_data)
# Preprocess
audio_tensor = preprocess_audio(audio_bytes)
# Inference
logger.info("Starting model inference...")
with torch.no_grad():
logits = model(audio_tensor)
# logits shape: (1, 1) -- single logit for binary classification
prob_fake = torch.sigmoid(logits).squeeze().item()
prediction = 1 if prob_fake >= input_data.threshold else 0
verdict = "fake" if prediction == 1 else "real"
inference_time = time.time() - start_time
logger.info(
"Prediction: %s (prob_fake=%.4f, time=%.3fs)",
verdict,
prob_fake,
inference_time,
)
return {
"model": MODEL_NAME,
"probability": float(prob_fake),
"prediction": int(prediction),
"class": verdict,
"inference_time": float(inference_time),
}
except Exception as e:
logger.exception("Error during prediction: %s", e)
raise HTTPException(status_code=500, detail=str(e))
if __name__ == "__main__":
port = int(os.environ.get("MODEL_PORT", 8003))
uvicorn.run(app, host="0.0.0.0", port=port)