"""FakeSTormer deepfake video detection service. Wraps the FakeSTormer (ICCV 2025) video deepfake detection model with a FastAPI endpoint. """ import base64 import gc import logging import os import platform import sys import tempfile import time from typing import Any, Dict, List import cv2 import numpy as np import torch import uvicorn from fastapi import FastAPI, HTTPException from PIL import Image from pydantic import BaseModel, ConfigDict, Field # Add model_code to path to allow imports sys.path.insert(0, os.path.join(os.getcwd(), "model_code")) from configs.get_config import load_config from models import MODELS, build_model, load_pretrained from package_utils.transform import final_transform logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) MODEL_PORT = int(os.environ.get("MODEL_PORT", 7001)) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" CONFIG_PATH = os.environ.get( "CONFIG_PATH", "model_code/configs/temporal/FakeSFormer_base_c23.yaml" ) WEIGHTS_PATH = os.environ.get( "WEIGHTS_PATH", "model_code/weights/TopDownDetector_C23_ViTBase224_ST_hm100_tempLOC0.2_4SBI_SAM_mp0.01_temp2_0vidAug_0.35_temp_normlms_normREAL_model_best.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") # Global model references _model = None _cfg = None _transforms = None _device = None def _load_models(): """Load FakeSTormer model.""" global _model, _cfg, _transforms, _device # noqa: F824 if _model is not None: return _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_mem / 1024**3, ) else: logger.warning( "Device: %s (no CUDA available -- check nvidia-container-toolkit)", _device, ) logger.info(f"Loading FakeSTormer model on {_device}...") try: # Load config _cfg = load_config(CONFIG_PATH) # Build model _model = build_model(_cfg.MODEL, MODELS).to(torch.float) # Load weights if not os.path.exists(WEIGHTS_PATH): logger.error(f"Weights file not found at {WEIGHTS_PATH}") # Use placeholder if weights missing? User asked for as true as possible results, # so we should fail if weights are missing, but let's see. raise FileNotFoundError(f"Weights missing: {WEIGHTS_PATH}") logger.info(f"Loading weight ... {WEIGHTS_PATH}") _model = load_pretrained(_model, WEIGHTS_PATH) _model = _model.to(_device) _model.eval() # Setup transforms _transforms = final_transform(_cfg.DATASET) logger.info("FakeSTormer model loaded successfully.") except Exception as e: logger.error(f"Failed to load FakeSTormer model: {e}") raise e def _is_model_loaded(): return _model is not None app = FastAPI( title="FakeSTormer Detection Service", description="Vulnerability-Aware Spatio-Temporal Learning for Generalizable Deepfake Video Detection", version="1.0.0", ) class PredictRequest(BaseModel): video_data: str # Base64 encoded video threshold: float = 0.5 class PredictResponse(BaseModel): model_config = ConfigDict(populate_by_name=True) model: str = "fakestormer" probability: float prediction: int class_name: str = Field(..., alias="class") inference_time: float metadata: Dict[str, Any] @app.on_event("startup") async def startup_event(): if PRELOAD_MODEL: _load_models() @app.get("/") def root(): return { "service": "fakestormer", "port": MODEL_PORT, "model_loaded": _is_model_loaded(), "device": str(_device) if _device else "unknown", } def _gpu_health_info() -> dict: """Return GPU metrics for the health endpoint.""" if torch.cuda.is_available() and _device is not None 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_mem / 1024**2 ), } return {} @app.get("/health") def health(): return { "status": "ok", "model_loaded": _is_model_loaded(), "weights_exist": os.path.exists(WEIGHTS_PATH), "device": str(_device) if _device else "cpu", **_gpu_health_info(), } def extract_frames(video_path: str, num_frames: int = 4) -> List[Image.Image]: """Extract frames from video file uniformly.""" cap = cv2.VideoCapture(video_path) total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) if total_frames <= 0: cap.release() return [] # Get indices for uniform sampling indices = np.linspace(0, total_frames - 1, num_frames, dtype=int) frames = [] for idx in indices: cap.set(cv2.CAP_PROP_POS_FRAMES, idx) ret, frame = cap.read() if ret: # Convert BGR to RGB frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) # Crop 15 pixels from each side as in test.py H, W, _ = frame.shape if H > 30 and W > 30: frame = frame[15 : H - 15, 15 : W - 15] frames.append(Image.fromarray(frame)) cap.release() # Pad if not enough frames while len(frames) < num_frames and len(frames) > 0: frames.append(frames[-1]) return frames @app.post("/predict", response_model=PredictResponse) async def predict(request: PredictRequest): global _model, _cfg, _transforms, _device # noqa: F824 if not _is_model_loaded(): _load_models() start_time = time.time() # Create a temporary file to save the video with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp_video: try: video_bytes = base64.b64decode(request.video_data) tmp_video.write(video_bytes) tmp_video_path = tmp_video.name except Exception as e: raise HTTPException(status_code=400, detail=f"Failed to decode video: {e}") try: num_frames = _cfg.DATASET.DATA.SAMPLES_PER_VIDEO.NUM_FRAMES or 4 frames = extract_frames(tmp_video_path, num_frames) if not frames: raise HTTPException( status_code=500, detail="Failed to extract frames from video." ) # Preprocess frames transformed_imgs = [] image_size = _cfg.DATASET.IMAGE_SIZE # [224, 224] for frame in frames: img_resize = frame.resize((int(image_size[0]), int(image_size[1]))) img_resize = np.array(img_resize) / 255.0 img_tensor = _transforms(img_resize).to(torch.float) transformed_imgs.append(img_tensor.unsqueeze(0)) # Stack and prepare for model [B, T, C, H, W] input_tensor = torch.cat(transformed_imgs, 0) # [T, C, H, W] input_tensor = input_tensor.to(_device) input_tensor = input_tensor.unsqueeze(0) # [1, T, C, H, W] # FakeSTormer model expects [1, C, T, H, W] input_tensor = input_tensor.transpose(1, 2) # [1, C, T, H, W] with torch.no_grad(): outputs = _model(input_tensor) if isinstance(outputs, list): outputs = outputs[0] prob = outputs["cls"].sigmoid().cpu().item() prediction = 1 if prob >= request.threshold else 0 class_name = "fake" if prediction == 1 else "real" return PredictResponse( probability=float(prob), prediction=prediction, class_name=class_name, inference_time=time.time() - start_time, metadata={"frames_extracted": len(frames), "device": str(_device)}, ) except Exception as e: logger.exception("Error during FakeSTormer prediction") raise HTTPException(status_code=500, detail=str(e)) finally: # Cleanup temporary file if os.path.exists(tmp_video_path): os.remove(tmp_video_path) gc.collect() if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=MODEL_PORT)