"""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)