"""FSD (Forensic Self-Descriptions) deepfake detection service. Wraps FSDDetector (CVPR 2025) with a FastAPI endpoint that matches the contract of all other DeepSafe image services. Key design decisions: - FSDDetector imported lazily inside load_detector_internal() so that the service can be imported in tests without the fsd package installed. - z_score_to_probability() converts FSD's z-score to [0,1] using a sigmoid centred at THRESHOLD_Z=-2.0 (the paper's default boundary). - SCALE=1.5 is an empirical choice — validated during evaluation. """ import base64 import gc import io import logging import math import os import platform import sys import threading import time from contextlib import asynccontextmanager from typing import Any, Dict, Optional import torch import uvicorn from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from PIL import Image, ImageFile from pydantic import BaseModel ImageFile.LOAD_TRUNCATED_IMAGES = True logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", handlers=[logging.StreamHandler(sys.stdout)], ) logger = logging.getLogger(__name__) # ── Configuration ───────────────────────────────────────────────────────────── MODEL_NAME = "fsd_detection" WEIGHTS_DIR = os.environ.get( "WEIGHTS_DIR", os.path.join(os.path.dirname(__file__), "weights") ) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "600")) def _get_device(): """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU.""" override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower() if override == "cpu": return "cpu" if override == "cuda" and torch.cuda.is_available(): return "cuda" if ( override == "mps" and hasattr(torch.backends, "mps") and torch.backends.mps.is_available() ): return "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 "mps" if torch.cuda.is_available(): return "cuda" return "cpu" DEVICE = _get_device() if DEVICE == "cuda": torch.backends.cudnn.benchmark = True torch.set_float32_matmul_precision("high") if DEVICE == "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, ) # Z-score conversion constants (see design spec §4) THRESHOLD_Z: float = -2.0 # FSD paper's default fake/real boundary SCALE: float = 1.5 # Empirical sigmoid sharpness — validated during evaluation # ── Global state ────────────────────────────────────────────────────────────── detector = None detector_lock = threading.Lock() last_used_time: float = 0.0 class ImageInput(BaseModel): """Request body for /predict.""" image_data: str threshold: Optional[float] = 0.5 # ── Conversion ──────────────────────────────────────────────────────────────── def z_score_to_probability(z_score: float) -> float: """Map an FSD z-score to a [0, 1] fake probability via sigmoid. Args: z_score: FSD detector output. More negative = more likely AI-generated. Returns: Float in [0, 1]. At z=THRESHOLD_Z the output is 0.5. """ exponent = -(THRESHOLD_Z - z_score) * SCALE # Clamp to prevent math.exp overflow on extreme inputs exponent = max(-500.0, min(500.0, exponent)) return 1.0 / (1.0 + math.exp(exponent)) # ── Model loading ───────────────────────────────────────────────────────────── def load_detector_internal(): """Load FSDDetector onto CPU from the local weights directory.""" global detector, last_used_time with detector_lock: if detector is not None: last_used_time = time.time() return logger.info("Loading FSD detector...") try: from fsd import FSDDetector # deferred import — keeps tests import-safe if not os.path.exists(WEIGHTS_DIR): raise RuntimeError( f"Weights directory not found: {WEIGHTS_DIR}. " "Run download_weights.sh first." ) loaded = FSDDetector.load( weights_dir=WEIGHTS_DIR, device=DEVICE, threshold=THRESHOLD_Z, ) detector = loaded last_used_time = time.time() logger.info("FSD detector ready.") except Exception as exc: logger.exception(f"Failed to load FSD detector: {exc}") detector = None raise finally: gc.collect() def ensure_detector_loaded(): """Load detector on first request (lazy loading).""" global last_used_time if detector is None: load_detector_internal() else: last_used_time = time.time() def unload_detector_if_idle(): """Evict detector from RAM after MODEL_TIMEOUT seconds of inactivity.""" global detector if detector is None or PRELOAD_MODEL: return with detector_lock: if detector is not None and (time.time() - last_used_time > MODEL_TIMEOUT): logger.info("Unloading idle FSD detector to free RAM.") del detector detector = None gc.collect() # ── FastAPI lifespan ────────────────────────────────────────────────────────── @asynccontextmanager async def lifespan(app: FastAPI): """Handle startup and shutdown logic for the FSD service.""" if PRELOAD_MODEL: logger.info("Preloading FSD detector at startup.") try: load_detector_internal() except Exception as exc: logger.error(f"Preload failed: {exc}") else: logger.info("FSD service ready — detector loads on first request.") if not PRELOAD_MODEL and MODEL_TIMEOUT > 0: def _periodic_check(): unload_detector_if_idle() threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start() threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start() yield logger.info("Shutting down FSD service.") # ── FastAPI app ─────────────────────────────────────────────────────────────── app = FastAPI( title="FSD Deepfake Detection Service", description=( "Zero-shot AI-generated image detection using Forensic Self-Descriptions " "(CVPR 2025). Trained on real images only; generalises to any generator." ), version="1.0.0", lifespan=lifespan, ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) @app.get("/") async def root(): """Service info.""" return { "model_name": MODEL_NAME, "description": "FSD (CVPR 2025) zero-shot AI-generated image detector", "device": 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 == "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", "model_name": MODEL_NAME, "device": DEVICE, "model_loaded": detector is not None, **_gpu_health_info(), } @app.post("/unload") async def unload_model_endpoint(): """Manually evict detector from RAM.""" global detector with detector_lock: if detector is None: return {"status": "not_loaded"} del detector detector = None gc.collect() return {"status": "success", "message": "Detector unloaded."} @app.post("/predict") async def predict(image_input: ImageInput) -> Dict[str, Any]: """Predict whether the submitted image is AI-generated. Args: image_input: Base64-encoded image and optional probability threshold. Returns: Dict with model name, fake probability, binary prediction, class label, and inference time. """ try: ensure_detector_loaded() if detector is None: raise HTTPException(status_code=503, detail="Detector not loaded.") start = time.time() try: image_bytes = base64.b64decode(image_input.image_data) image_pil = Image.open(io.BytesIO(image_bytes)).convert("RGB") except Exception as exc: raise HTTPException(status_code=400, detail=f"Invalid image data: {exc}") with torch.no_grad(): result = detector.score(image_pil) probability = z_score_to_probability(result.z_score) prediction = 1 if probability >= image_input.threshold else 0 class_label = "fake" if prediction == 1 else "real" inference_time = time.time() - start logger.info( f"z={result.z_score:.4f} prob={probability:.4f} " f"→ {class_label} ({inference_time:.3f}s)" ) if not PRELOAD_MODEL and MODEL_TIMEOUT > 0: threading.Timer(MODEL_TIMEOUT + 5.0, unload_detector_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(f"Prediction error: {exc}") raise HTTPException(status_code=500, detail=str(exc)) if __name__ == "__main__": port = int(os.environ.get("MODEL_PORT", 5005)) logger.info(f"Starting FSD service on port {port}") uvicorn.run(app, host="0.0.0.0", port=port)