Instructions to use deepsafe/deepsafe-services with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use deepsafe/deepsafe-services with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("deepsafe/deepsafe-services", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download image/aide/app.py from deepsafe/deepsafe-services: direct link, hf CLI and curl.
- Browser
- Download file 15.1 kB
-
https://huggingface.co/deepsafe/deepsafe-services/resolve/main/image/aide/app.py
- Command line
-
hf download hf://deepsafe/deepsafe-services/image/aide/app.py
-
curl -L -o app.py https://huggingface.co/deepsafe/deepsafe-services/resolve/main/image/aide/app.py
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=["*"], | |
| ) | |
| 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 {} | |
| 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(), | |
| } | |
| 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."} | |
| 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)) | |
| 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) | |