"""AudioSeal watermark detection service. Detects Meta AudioSeal watermarks embedded in audio files using the audioseal_detector_16bits model. Returns a probability score indicating likelihood that the audio contains a watermark (indicating AI generation). Unlike detection models that analyze acoustic artifacts, this service inspects only for embedded AudioSeal watermarks. """ import base64 import gc import logging import os import platform import sys import tempfile import threading import time from typing import Any, Dict, Optional import torch import torchaudio import uvicorn from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel # ── Logging ──────────────────────────────────────────────────────────────── logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", handlers=[logging.StreamHandler(sys.stdout)], ) logger = logging.getLogger(__name__) # ── Config ───────────────────────────────────────────────────────────────── MODEL_NAME = "audioseal_detector" MODEL_PORT = int(os.environ.get("MODEL_PORT", "9003")) _PRODUCTION = os.environ.get("PRODUCTION", "false").lower() == "true" DETECTION_THRESHOLD = 0.3 TARGET_SAMPLE_RATE = 16000 def _get_device(): """Select optimal device: CUDA > MPS > 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 torch.cuda.is_available(): return torch.device("cuda") if ( platform.system() == "Darwin" and hasattr(torch.backends, "mps") and torch.backends.mps.is_available() ): return torch.device("mps") 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_mem / 1024**3, ) else: logger.warning( "Device: %s (no CUDA available -- check nvidia-container-toolkit)", DEVICE, ) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "600")) # ── Globals ──────────────────────────────────────────────────────────────── detector = None model_lock = threading.Lock() last_used_time = 0 # ── Pydantic models ─────────────────────────────────────────────────────── class AudioInput(BaseModel): """Request body for /predict.""" audio_data: str threshold: float = 0.5 # ── Confidence mapping ──────────────────────────────────────────────────── def _map_confidence_to_probability(confidence: float) -> float: """Map AudioSeal detection confidence to a probability score. Confidence values below the detection threshold are treated as noise and mapped to 0.5 (neutral). Values at or above the threshold are linearly scaled into [0.5, 1.0]. Args: confidence: Raw detection confidence from AudioSeal in [0, 1]. Returns: Probability in [0.5, 1.0]. 0.5 means neutral (no watermark). """ if confidence < DETECTION_THRESHOLD: return 0.5 return 0.5 + confidence * 0.5 # ── Model loading ───────────────────────────────────────────────────────── def load_model_internal(): """Load AudioSeal detector model onto the selected device.""" global detector, last_used_time with model_lock: if detector is not None: last_used_time = time.time() return logger.info("Loading AudioSeal detector model...") try: from audioseal import AudioSeal loaded = AudioSeal.load_detector("audioseal_detector_16bits") loaded = loaded.to(DEVICE) # Switch to inference mode (no gradient tracking) loaded.train(mode=False) detector = loaded last_used_time = time.time() logger.info("AudioSeal detector ready on %s.", str(DEVICE)) except Exception as exc: logger.exception("Failed to load AudioSeal detector: %s", exc) detector = None raise finally: gc.collect() def ensure_model_loaded(): """Load model on first request (lazy loading).""" global last_used_time if detector is None: load_model_internal() else: last_used_time = time.time() def unload_model_if_idle(): """Evict model from RAM after MODEL_TIMEOUT seconds of inactivity.""" global detector if detector is None or PRELOAD_MODEL: return with model_lock: if detector is not None and (time.time() - last_used_time > MODEL_TIMEOUT): logger.info("Unloading idle AudioSeal detector to free RAM.") del detector detector = None gc.collect() # ── FastAPI app ──────────────────────────────────────────────────────────── app = FastAPI( title="AudioSeal Watermark Detection Service", description=( "Detects Meta AudioSeal watermarks embedded in audio files " "using the audioseal_detector_16bits model." ), version="1.0.0", docs_url=None if _PRODUCTION else "/docs", redoc_url=None if _PRODUCTION else "/redoc", openapi_url=None if _PRODUCTION else "/openapi.json", ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) @app.get("/") async def root(): """Root endpoint with service information.""" if _PRODUCTION: return {"status": "ok"} return { "model_name": MODEL_NAME, "description": ("AudioSeal watermark detector for AI-generated audio"), "device": str(DEVICE), "model_loaded": detector is not None, } 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_mem / 1024**2 ), } return {} @app.get("/health") async def health(): """Health check endpoint.""" if _PRODUCTION: return {"status": "healthy"} return { "status": "healthy", "model_name": MODEL_NAME, "device": str(DEVICE), "model_loaded": detector is not None, **_gpu_health_info(), } @app.post("/predict") async def predict(payload: AudioInput) -> Dict[str, Any]: """Detect AudioSeal watermark in base64-encoded audio. Decodes the audio, resamples to 16 kHz mono, and runs the AudioSeal detector to check for embedded watermarks. Args: payload: Base64-encoded audio data and optional threshold. Returns: Dict with model name, probability, prediction, class, and inference time. """ try: ensure_model_loaded() if detector is None: raise HTTPException(status_code=503, detail="Model not loaded.") start = time.time() # Decode base64 audio try: audio_bytes = base64.b64decode(payload.audio_data) except Exception as exc: raise HTTPException( status_code=400, detail=f"Invalid base64 data: {exc}", ) if not audio_bytes: raise HTTPException( status_code=400, detail="Empty audio data after base64 decode.", ) # Write to temp file and load with torchaudio tmp_path = None try: with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: tmp.write(audio_bytes) tmp_path = tmp.name # Try torchaudio first, fall back to soundfile if torchcodec # or other backends are unavailable. try: waveform, sample_rate = torchaudio.load(tmp_path) except Exception as ta_exc: logger.info( "torchaudio.load failed (%s), falling back to soundfile.", ta_exc, ) import soundfile as sf data, sample_rate = sf.read(tmp_path, dtype="float32") waveform = torch.from_numpy(data).T # [channels, samples] if waveform.dim() == 1: waveform = waveform.unsqueeze(0) except Exception as exc: logger.warning("Failed to load audio: %s", exc) inference_time = time.time() - start return { "model": MODEL_NAME, "probability": 0.5, "prediction": 0, "class": "real", "inference_time": float(inference_time), } finally: if tmp_path and os.path.exists(tmp_path): try: os.unlink(tmp_path) except OSError: pass # Stereo to mono: average channels if waveform.shape[0] > 1: waveform = waveform.mean(dim=0, keepdim=True) # Resample to target sample rate if sample_rate != TARGET_SAMPLE_RATE: resampler = torchaudio.transforms.Resample( orig_freq=sample_rate, new_freq=TARGET_SAMPLE_RATE, ) waveform = resampler(waveform) # Add batch dimension: [1, channels, samples] if waveform.dim() == 2: waveform = waveform.unsqueeze(0) waveform = waveform.to(DEVICE) # Run detection with torch.no_grad(): result, message = detector.detect_watermark( waveform, sample_rate=TARGET_SAMPLE_RATE ) confidence = float(result) probability = _map_confidence_to_probability(confidence) prediction = 1 if probability > payload.threshold else 0 class_label = "fake" if prediction == 1 else "real" inference_time = time.time() - start logger.info( "AudioSeal check: %s (conf=%.4f, prob=%.4f, %.3fs)", class_label, confidence, probability, inference_time, ) if not PRELOAD_MODEL and MODEL_TIMEOUT > 0: threading.Timer(MODEL_TIMEOUT + 5.0, unload_model_if_idle).start() return { "model": MODEL_NAME, "probability": float(probability), "prediction": int(prediction), "class": class_label, "inference_time": float(inference_time), } except HTTPException: raise except Exception as exc: logger.exception("Prediction error: %s", exc) raise HTTPException(status_code=500, detail=str(exc)) @app.on_event("startup") async def startup_event(): """Startup: preload model if configured, else lazy-load on first request.""" if PRELOAD_MODEL: logger.info("Preloading AudioSeal detector at startup.") try: load_model_internal() except Exception as exc: logger.error("Preload failed: %s", exc) else: logger.info("AudioSeal service ready — model loads on first request.") if not PRELOAD_MODEL and MODEL_TIMEOUT > 0: def _periodic_check(): unload_model_if_idle() threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start() threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start() if __name__ == "__main__": logger.info("Starting AudioSeal detector service on port %d", MODEL_PORT) uvicorn.run(app, host="0.0.0.0", port=MODEL_PORT)