deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
15.1 kB
import base64
import gc
import io
import logging
import os
import platform
import sys
import threading
import time
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__)
# ── Path setup ──────────────────────────────────────────────────────────────
current_dir = os.path.dirname(os.path.abspath(__file__))
model_code_dir = os.path.join(current_dir, "model_code")
sys.path.insert(0, model_code_dir)
sys.path.insert(0, os.path.join(model_code_dir, "models"))
sys.path.insert(0, os.path.join(model_code_dir, "data"))
# ── Compatibility shim ───────────────────────────────────────────────────────
# AIDE's models/AIDE.py imports `clip` (openai-clip) at module level, but the
# package is not used during inference — only open_clip is. The openai-clip
# package relies on pkg_resources which was removed in Python 3.13. We inject
# a lightweight stub so the import succeeds without installing the full package.
import types as _types
if "clip" not in sys.modules:
_clip_stub = _types.ModuleType("clip")
sys.modules["clip"] = _clip_stub
# ── Config ──────────────────────────────────────────────────────────────────
MODEL_NAME = "aide_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
# The checkpoint is self-contained (includes ConvNeXt weights), so we initialise
# the architecture with no pretrained weights and load everything from the checkpoint.
# Set CONVNEXT_PRETRAINED to a HuggingFace tag only if running without a checkpoint.
CONVNEXT_PRETRAINED = os.environ.get("CONVNEXT_PRETRAINED", None)
# Preferred checkpoint filename (GenImage trains on the most diverse generators)
PREFERRED_CHECKPOINT = os.environ.get("AIDE_CHECKPOINT", "GenImage_train.pth")
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
# ── Weight discovery ─────────────────────────────────────────────────────────
def find_aide_checkpoint() -> Optional[str]:
"""
Return path to the best AIDE checkpoint in WEIGHTS_DIR.
Priority order:
1. PREFERRED_CHECKPOINT filename (GenImage_train.pth by default)
2. Any other .pth file (largest wins)
"""
if not os.path.exists(WEIGHTS_DIR):
logger.warning(f"Weights directory not found: {WEIGHTS_DIR}")
return None
# Try the preferred checkpoint first
preferred = os.path.join(WEIGHTS_DIR, PREFERRED_CHECKPOINT)
if os.path.exists(preferred):
logger.info(f"Using preferred checkpoint: {preferred}")
return preferred
# Fall back to the largest available checkpoint
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(f"Using checkpoint: {best}")
return best
# ── Preprocessing ────────────────────────────────────────────────────────────
def preprocess_image(image_bytes: bytes) -> torch.Tensor:
"""
Preprocess raw image bytes into AIDE's 5-view tensor.
Args:
image_bytes: Raw bytes of a JPEG/PNG/etc. image.
Returns:
Tensor of shape [1, 5, 3, 256, 256] on CPU.
Raises:
Exception: If bytes cannot be decoded or processed.
"""
from data.dct import DCT_base_Rec_Module
from torchvision import transforms
pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
# Ensure minimum 256x256 so DCT unfold has enough patches
w, h = pil_image.size
if w < 256 or h < 256:
pil_image = pil_image.resize((256, 256), Image.BICUBIC)
to_tensor = transforms.ToTensor()
image_tensor = to_tensor(pil_image) # [3, H, W]
# DCT frequency decomposition → 4 patches [3, 32, 32] each
dct_module = DCT_base_Rec_Module()
x_minmin, x_maxmax, x_minmin1, x_maxmax1 = dct_module(image_tensor)
# Resize all views to 256×256 and normalise with ImageNet stats
transform = transforms.Compose(
[
transforms.Resize([256, 256]),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
]
)
x_0 = transform(image_tensor)
x_minmin = transform(x_minmin)
x_maxmax = transform(x_maxmax)
x_minmin1 = transform(x_minmin1)
x_maxmax1 = transform(x_maxmax1)
# Stack → [5, 3, 256, 256], unsqueeze batch → [1, 5, 3, 256, 256]
stacked = torch.stack([x_minmin, x_maxmax, x_minmin1, x_maxmax1, x_0], dim=0)
return stacked.unsqueeze(0).to(DEVICE)
# ── Model loading ─────────────────────────────────────────────────────────────
def load_model_internal():
"""Load AIDE_Model onto CPU with the best available checkpoint."""
global model, last_used_time
with model_lock:
if model is not None:
last_used_time = time.time()
return
logger.info("Loading AIDE model...")
try:
import models.AIDE as AIDE_module
aide_model = AIDE_module.AIDE(
resnet_path=None,
convnext_path=CONVNEXT_PRETRAINED,
)
aide_model.to(DEVICE)
checkpoint_path = find_aide_checkpoint()
if checkpoint_path:
logger.info(f"Loading checkpoint: {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 = aide_model.load_state_dict(cleaned, strict=False)
logger.info(
f"Checkpoint loaded. Missing keys: {len(missing)}, "
f"Unexpected keys: {len(unexpected)}"
)
else:
logger.warning(
"No checkpoint found — model uses pretrained-only weights. "
"Run download_weights.sh to fetch the fine-tuned checkpoint."
)
# Switch to inference mode (no gradient tracking, batch-norm uses running stats)
aide_model.train(mode=False)
model = aide_model
last_used_time = time.time()
logger.info("AIDE model ready.")
except Exception as exc:
logger.exception(f"Failed to load AIDE model: {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 AIDE model to free RAM.")
del model
model = None
gc.collect()
# ── FastAPI app ───────────────────────────────────────────────────────────────
app = FastAPI(
title="AIDE Deepfake Detection Service",
description=(
"AI-generated image detection using AIDE (ICLR 2025) — "
"hybrid DCT frequency analysis + ConvNeXt-xxlarge semantic features."
),
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": "AIDE (ICLR 2025) 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_name": 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 classification 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)
except Exception as exc:
raise HTTPException(status_code=400, detail=f"Invalid image data: {exc}")
with torch.no_grad():
logits = model(input_tensor) # [1, 2]
probs = torch.softmax(logits, dim=-1) # [1, 2]
probability_fake = probs[0, 1].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(
f"Prediction: {class_label} (prob={probability_fake:.4f}, {inference_time:.3f}s)"
)
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(f"Prediction error: {exc}")
raise HTTPException(status_code=500, detail=str(exc))
@app.on_event("startup")
async def startup_event():
"""Startup handler — preloads model if PRELOAD_MODEL=true, else lazy-loads."""
if PRELOAD_MODEL:
logger.info("Preloading AIDE model at startup.")
try:
load_model_internal()
except Exception as exc:
logger.error(f"Preload failed: {exc}")
else:
logger.info("AIDE 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", 5004))
logger.info(f"Starting AIDE service on port {port}")
uvicorn.run(app, host="0.0.0.0", port=port)