"""Effort (ICML 2025) AI-generated image detection service. Wraps the Effort detector — SVD orthogonal subspace decomposition on CLIP ViT-L/14 — with a FastAPI endpoint matching the DeepSafe contract. Key design decisions: - Model architecture code (EffortDetector, SVDResidualLinear, apply_svd_residual_to_self_attn) is loaded from the cloned repo via sys.path at startup. The API layer in this file is fully self-contained. - Preprocessing mirrors demo.py:preprocess_face() — resize to 224x224, CLIP normalisation ([0.4815, 0.4578, 0.4082], [0.2686, 0.2613, 0.2758]). - No face-alignment (dlib) dependency: the image is resized directly, matching the repo's non-landmark path (landmark_model=False). - Weights are volume-mounted at /app/weights/ — not baked into the image. """ import base64 import gc import io import logging import math import os import platform import sys import threading import time from typing import Any, Dict, Optional import cv2 import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import uvicorn from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from PIL import Image, ImageFile from pydantic import BaseModel from torchvision import transforms 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__) # ── Path setup ────────────────────────────────────────────────────────────── current_dir = os.path.dirname(os.path.abspath(__file__)) model_code_dir = os.path.join(current_dir, "model_code") # Add model code paths so the repo's detector / network / loss modules resolve for _p in [ model_code_dir, os.path.join(model_code_dir, "training"), os.path.join(model_code_dir, "training", "detectors"), os.path.join(model_code_dir, "training", "networks"), os.path.join(model_code_dir, "training", "loss"), os.path.join(model_code_dir, "training", "utils"), os.path.join(model_code_dir, "training", "metrics"), ]: if _p not in sys.path: sys.path.insert(0, _p) # ── Config ────────────────────────────────────────────────────────────────── MODEL_NAME = "effort_detection" WEIGHTS_DIR = os.environ.get("WEIGHTS_DIR", os.path.join(current_dir, "weights")) HF_HOME = os.environ.get("HF_HOME", os.path.join(current_dir, "hf_cache")) os.environ["HF_HOME"] = HF_HOME os.environ["TRANSFORMERS_CACHE"] = os.path.join(HF_HOME, "hub") # Preferred checkpoint filename PREFERRED_CHECKPOINT = os.environ.get("EFFORT_CHECKPOINT", "genimage_effort.pth") # CLIP normalization constants (from the Effort config / OpenAI CLIP) CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073] CLIP_STD = [0.26862954, 0.26130258, 0.27577711] INPUT_RESOLUTION = 224 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, ) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "600")) # ── Globals ───────────────────────────────────────────────────────────────── model = None model_lock = threading.Lock() last_used_time = 0 class ImageInput(BaseModel): """Request body for /predict.""" image_data: str threshold: Optional[float] = 0.5 # ── Monkey-patch CLIP loading ────────────────────────────────────────────── # The Effort repo's effort_detector.py hardcodes a relative path # CLIPModel.from_pretrained("../models--openai--clip-vit-large-patch14") # which is an invalid HuggingFace repo ID. We intercept from_pretrained and # redirect to the canonical hub identifier so transformers downloads/caches # the model correctly. _CLIP_PATH_FIXUPS = { "../models--openai--clip-vit-large-patch14": "openai/clip-vit-large-patch14", "models--openai--clip-vit-large-patch14": "openai/clip-vit-large-patch14", } def _patch_clip_from_pretrained(): """Wrap CLIPModel.from_pretrained to fix hardcoded relative paths.""" from transformers import CLIPModel _original_from_pretrained = CLIPModel.from_pretrained.__func__ @classmethod # type: ignore[misc] def _patched_from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs): fixed = _CLIP_PATH_FIXUPS.get( pretrained_model_name_or_path, pretrained_model_name_or_path ) if fixed != pretrained_model_name_or_path: logger.info( "Intercepted CLIP path '%s' -> '%s'", pretrained_model_name_or_path, fixed, ) return _original_from_pretrained(cls, fixed, *args, **kwargs) CLIPModel.from_pretrained = _patched_from_pretrained _patch_clip_from_pretrained() # ── Weight discovery ──────────────────────────────────────────────────────── def find_effort_checkpoint() -> Optional[str]: """Return path to the best Effort checkpoint in WEIGHTS_DIR. Priority order: 1. PREFERRED_CHECKPOINT filename 2. Any other .pth / .pt file (largest wins) """ if not os.path.exists(WEIGHTS_DIR): logger.warning("Weights directory not found: %s", WEIGHTS_DIR) return None preferred = os.path.join(WEIGHTS_DIR, PREFERRED_CHECKPOINT) if os.path.exists(preferred): logger.info("Using preferred checkpoint: %s", preferred) return preferred candidates = [ os.path.join(WEIGHTS_DIR, f) for f in os.listdir(WEIGHTS_DIR) if f.endswith(".pth") or f.endswith(".pt") ] if not candidates: logger.warning("No .pth checkpoint found in weights directory.") return None best = max(candidates, key=os.path.getsize) logger.info("Using checkpoint: %s", best) return best # ── Preprocessing ─────────────────────────────────────────────────────────── def preprocess_image(image_bytes: bytes) -> torch.Tensor: """Preprocess raw image bytes into an Effort-compatible tensor. Mirrors demo.py:preprocess_face() — resize to 224x224, convert to PIL, apply CLIP normalisation. Args: image_bytes: Raw bytes of a JPEG/PNG/etc. image. Returns: Tensor of shape [1, 3, 224, 224]. Raises: Exception: If bytes cannot be decoded or processed. """ pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB") pil_image = pil_image.resize((INPUT_RESOLUTION, INPUT_RESOLUTION), Image.BICUBIC) transform = transforms.Compose( [ transforms.ToTensor(), transforms.Normalize(mean=CLIP_MEAN, std=CLIP_STD), ] ) tensor = transform(pil_image) # [3, 224, 224] return tensor.unsqueeze(0) # [1, 3, 224, 224] # ── Model loading ─────────────────────────────────────────────────────────── def _build_effort_config() -> dict: """Build the minimal config dict required by EffortDetector.__init__.""" return { "model_name": "effort", "backbone_name": "vit", "pretrained": None, "backbone_config": { "mode": "original", "num_classes": 2, "inc": 3, "dropout": False, }, "resolution": INPUT_RESOLUTION, "mean": CLIP_MEAN, "std": CLIP_STD, "loss_func": "cross_entropy", } def load_model_internal(): """Load EffortDetector onto the selected device.""" global model, last_used_time with model_lock: if model is not None: last_used_time = time.time() return logger.info("Loading Effort model...") try: # Import detector from the cloned repo code from detectors.effort_detector import ( EffortDetector, ) cfg = _build_effort_config() effort_model = EffortDetector(config=cfg) effort_model.to(DEVICE) checkpoint_path = find_effort_checkpoint() if checkpoint_path: logger.info("Loading checkpoint: %s", checkpoint_path) ckpt = torch.load(checkpoint_path, map_location=DEVICE) if isinstance(ckpt, dict): state_dict = ckpt.get("model") or ckpt.get("state_dict") or ckpt else: state_dict = ckpt # Strip DataParallel "module." prefix if present cleaned = { k[7:] if k.startswith("module.") else k: v for k, v in state_dict.items() } missing, unexpected = effort_model.load_state_dict( cleaned, strict=False ) logger.info( "Checkpoint loaded. Missing keys: %d, " "Unexpected keys: %d", len(missing), len(unexpected), ) else: logger.warning( "No checkpoint found -- model uses pretrained-only " "weights. Download the fine-tuned checkpoint for " "accurate predictions." ) effort_model.train(mode=False) model = effort_model last_used_time = time.time() logger.info("Effort model ready on %s.", DEVICE) except Exception as exc: logger.exception("Failed to load Effort model: %s", exc) model = None raise finally: gc.collect() def ensure_model_loaded(): """Load model on first request (lazy loading).""" global last_used_time if model 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 model if model is None or PRELOAD_MODEL: return with model_lock: if model is not None and (time.time() - last_used_time > MODEL_TIMEOUT): logger.info("Unloading idle Effort model to free RAM.") del model model = None gc.collect() # ── FastAPI app ───────────────────────────────────────────────────────────── app = FastAPI( title="Effort Deepfake Detection Service", description=( "AI-generated image detection using Effort (ICML 2025 Oral) -- " "SVD orthogonal subspace decomposition on CLIP ViT-L/14." ), version="1.0.0", ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) @app.get("/") async def root(): """Root endpoint with service information.""" return { "model_name": MODEL_NAME, "description": ("Effort (ICML 2025 Oral) AI-generated image detector"), "device": str(DEVICE), "model_loaded": model 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_memory / 1024**2 ), } return {} @app.get("/health") async def health(): """Health check endpoint.""" return { "status": "healthy", "model": MODEL_NAME, "device": str(DEVICE), "model_loaded": model is not None, **_gpu_health_info(), } @app.post("/unload") async def unload_model_endpoint(): """Manually unload the model to free RAM.""" global model if model is None: return {"status": "not_loaded"} del model model = None gc.collect() return {"status": "success", "message": "Model 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 threshold. Returns: Dict with model name, fake probability, binary prediction, class label, and inference time. """ try: ensure_model_loaded() if model is None: raise HTTPException(status_code=503, detail="Model not loaded.") start = time.time() try: image_bytes = base64.b64decode(image_input.image_data) input_tensor = preprocess_image(image_bytes) input_tensor = input_tensor.to(DEVICE) except Exception as exc: raise HTTPException(status_code=400, detail=f"Invalid image data: {exc}") with torch.no_grad(): data_dict = { "image": input_tensor, "label": torch.tensor([0]).to(DEVICE), } preds = model(data_dict, inference=True) probability_fake = preds["prob"].squeeze().cpu().item() prediction = 1 if probability_fake >= image_input.threshold else 0 class_label = "fake" if prediction == 1 else "real" inference_time = time.time() - start logger.info( "Prediction: %s (prob=%.4f, %.3fs)", class_label, probability_fake, 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_fake), "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(): """Preload model if PRELOAD_MODEL=true, else lazy-load on first request.""" if PRELOAD_MODEL: logger.info("Preloading Effort model at startup.") try: load_model_internal() except Exception as exc: logger.error("Preload failed: %s", exc) else: logger.info("Effort 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__": port = int(os.environ.get("MODEL_PORT", 5006)) logger.info("Starting Effort service on port %d", port) uvicorn.run(app, host="0.0.0.0", port=port)