deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
11.1 kB
"""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)