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