"""PwTF-DVD (Pixel-wise Temporal Frequency) deepfake video detection service. Wraps the PwTF-DVD (ICCV 2025) face forgery detection model with a FastAPI endpoint. Uses a dual-stream architecture: an I3D backbone for spatial features and a ResNet-based attention network for temporal frequency features, fused via spatial and temporal transformer encoders. The preprocessing pipeline performs RetinaFace detection, SORT-based tracking, face alignment via 68-point landmarks, and temporal FFT computation on median-filtered residuals. Reference: "Pixel-wise Temporal Frequency Domain Video Deepfake Detection", ICCV 2025. """ import base64 import gc import logging import os import platform import sys import tempfile import threading import time from typing import Any, Dict, List, Optional import cv2 import numpy as np import torch import torch.nn.functional as F import uvicorn from fastapi import FastAPI, HTTPException from PIL import Image from pydantic import BaseModel, ConfigDict, Field from torchvision.transforms import Compose, Normalize, ToTensor logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) MODEL_PORT = int(os.environ.get("MODEL_PORT", 7005)) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", 1800)) WEIGHTS_PATH = os.environ.get( "WEIGHTS_PATH", "/app/weights/pwtf_dvd_checkpoint.pth" ) MODEL_CODE_DIR = os.environ.get("MODEL_CODE_DIR", "/app/model_code") # PwTF-DVD uses 224x224 face crops after alignment FACE_CROP_SIZE = 224 # Clip size for temporal analysis (from root_setting.yaml clip_size: 32) CLIP_SIZE = 32 # Maximum frames to extract from the video MAX_FRAMES = 768 def _get_device() -> torch.device: """Select optimal device: CUDA (NVIDIA) > MPS (Apple) > 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 torch.cuda.is_available(): return torch.device("cuda") if ( platform.system() == "Darwin" and hasattr(torch.backends, "mps") and torch.backends.mps.is_available() ): return torch.device("mps") return torch.device("cpu") # ── Global state ──────────────────────────────────────────────────────────── _model: Optional[torch.nn.Module] = None _device: Optional[torch.device] = None _load_lock = threading.Lock() # Lazy-loaded references to model_code modules _detect_all = None _grab_all_frames = None _get_crop_box = None _multiple_tracking = None _find_longest = None _FasterCropAlignXRay = None _crop_align_func = None def _ensure_model_code_on_path() -> None: """Add model_code/inference to sys.path so its internal imports work. The model code uses bare imports like ``from model.framework import get_model`` and ``from config_ftcn import config`` which expect the ``inference/`` directory to be on ``sys.path``. """ inference_dir = os.path.join(MODEL_CODE_DIR, "inference") if inference_dir not in sys.path: sys.path.insert(0, inference_dir) def _load_models() -> None: """Load PwTF-DVD model and face detection tools (thread-safe).""" global _model, _device global _detect_all, _grab_all_frames, _get_crop_box global _multiple_tracking, _find_longest global _FasterCropAlignXRay, _crop_align_func if _model is not None: return with _load_lock: # Double-check after acquiring lock 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 -- check nvidia-container-toolkit)", _device, ) logger.info("Loading PwTF-DVD model on %s ...", _device) # ── Import model_code modules ─────────────────────────────── _ensure_model_code_on_path() from model.framework import get_model from test_tools.common import detect_all, grab_all_frames from test_tools.utils import get_crop_box from test_tools.ct.operations import find_longest, multiple_tracking from test_tools.faster_crop_align_xray import FasterCropAlignXRay _detect_all = detect_all _grab_all_frames = grab_all_frames _get_crop_box = get_crop_box _multiple_tracking = multiple_tracking _find_longest = find_longest _FasterCropAlignXRay = FasterCropAlignXRay # ── Face crop alignment ───────────────────────────────────── _crop_align_func = FasterCropAlignXRay(FACE_CROP_SIZE) # ── PwTF-DVD classifier ───────────────────────────────────── if not os.path.exists(WEIGHTS_PATH): raise FileNotFoundError( f"PwTF-DVD weights not found at {WEIGHTS_PATH}" ) model = get_model() state_dict = torch.load( WEIGHTS_PATH, map_location="cpu", weights_only=False ) model.load_state_dict(state_dict) model = model.to(_device) model.eval() _model = model logger.info("PwTF-DVD model loaded successfully.") def _is_model_loaded() -> bool: """Return True if the model is loaded and ready.""" return _model is not None # ── FastAPI app ───────────────────────────────────────────────────────────── app = FastAPI( title="PwTF-DVD Detection Service", description=( "Pixel-wise Temporal Frequency Domain Video Deepfake Detection " "(ICCV 2025)" ), version="1.0.0", ) class PredictRequest(BaseModel): """Incoming prediction request.""" video_data: str # Base64-encoded video bytes threshold: float = 0.5 class PredictResponse(BaseModel): """Outgoing prediction result.""" model_config = ConfigDict(populate_by_name=True) model: str = "pwtf_dvd_detection" 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(): """Optionally preload model at startup.""" if PRELOAD_MODEL: _load_models() @app.get("/") def root(): """Service info endpoint.""" return { "service": "pwtf_dvd_detection", "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(): """Health check endpoint.""" return { "status": "healthy", "model": "pwtf_dvd_detection", "device": str(_device) if _device else "cpu", "model_loaded": _is_model_loaded(), "weights_exist": os.path.exists(WEIGHTS_PATH), **_gpu_health_info(), } # ── Inference pipeline ────────────────────────────────────────────────────── def _run_inference(video_path: str) -> dict: """Run the full PwTF-DVD inference pipeline on a video file. Follows the same logic as ``model_code/inference/test_on_raw_video.py``: 1. Detect faces in all frames (RetinaFace). 2. Track faces across frames (SORT-based tracker). 3. Generate sliding-window clips of ``CLIP_SIZE`` frames. 4. For each clip: align faces, compute temporal FFT residuals, run dual-stream model. 5. Aggregate per-clip predictions. Args: video_path: Path to the video file on disk. Returns: Dict with ``probability``, ``num_frames``, ``num_tracks``, ``num_clips``, ``frames_with_faces``. """ # ── Step 1: Detect faces in all frames ────────────────────────── detect_res, all_lm68, frames = _detect_all( video_path, return_frames=True, max_size=MAX_FRAMES ) if not frames: return { "probability": 0.5, "num_frames": 0, "num_tracks": 0, "num_clips": 0, "frames_with_faces": 0, } shape = frames[0].shape[:2] # Merge 68-landmark data into detection results all_detect_res = [] for faces, faces_lm68 in zip(detect_res, all_lm68): new_faces = [] for (box, lm5, score), face_lm68 in zip(faces, faces_lm68): new_faces.append((box, lm5, face_lm68, score)) all_detect_res.append(new_faces) detect_res = all_detect_res # ── Step 2: Track faces ───────────────────────────────────────── tracks = _multiple_tracking(detect_res) tuples = [(0, len(detect_res))] * len(tracks) if len(tracks) == 0: tuples, tracks = _find_longest(detect_res) if len(tracks) == 0: return { "probability": 0.5, "num_frames": len(frames), "num_tracks": 0, "num_clips": 0, "frames_with_faces": 0, } # ── Step 3: Extract face crops and landmarks ──────────────────── data_storage = {} frame_boxes = {} super_clips = [] for track_i, ((start, end), track) in enumerate( zip(tuples, tracks) ): super_clips.append(len(track)) for face, frame_idx, j in zip( track, range(start, end), range(len(track)) ): box, lm5, lm68 = face[:3] big_box = _get_crop_box(shape, box, scale=0.5) top_left = big_box[:2][None, :] new_lm5 = lm5 - top_left new_lm68 = lm68 - top_left new_box = (box.reshape(2, 2) - top_left).reshape(-1) info = (new_box, new_lm5, new_lm68, big_box) x1, y1, x2, y2 = big_box cropped = frames[frame_idx][y1:y2, x1:x2] base_key = f"{track_i}_{j}_" data_storage[base_key + "img"] = cropped data_storage[base_key + "ldm"] = info data_storage[base_key + "idx"] = frame_idx frame_boxes[frame_idx] = np.rint(box).astype(np.int32) # ── Step 4: Generate sliding-window clips ─────────────────────── clips_for_video = [] pad_length = CLIP_SIZE - 1 for super_clip_idx, super_clip_size in enumerate(super_clips): inner_index = list(range(super_clip_size)) if super_clip_size < CLIP_SIZE: post_module = inner_index[1:-1][::-1] + inner_index l_post = len(post_module) if l_post == 0: continue post_module = post_module * (pad_length // l_post + 1) post_module = post_module[:pad_length] if len(post_module) != pad_length: continue pre_module = inner_index + inner_index[1:-1][::-1] l_pre = len(pre_module) if l_pre == 0: continue pre_module = pre_module * (pad_length // l_pre + 1) pre_module = pre_module[-pad_length:] if len(pre_module) != pad_length: continue inner_index = pre_module + inner_index + post_module padded_size = len(inner_index) frame_range = [ inner_index[i: i + CLIP_SIZE] for i in range(padded_size) if i + CLIP_SIZE <= padded_size ] for indices in frame_range: clip = [(super_clip_idx, t) for t in indices] clips_for_video.append(clip) if not clips_for_video: return { "probability": 0.5, "num_frames": len(frames), "num_tracks": len(tracks), "num_clips": 0, "frames_with_faces": len(frame_boxes), } # ── Step 5: Run inference on clips ────────────────────────────── preds = [] test_transform = Compose([ ToTensor(), Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], ), ]) for clip in clips_for_video: images = [data_storage[f"{i}_{j}_img"] for i, j in clip] landmarks = [data_storage[f"{i}_{j}_ldm"] for i, j in clip] # Align and crop faces landmarks, images = _crop_align_func(landmarks, images) # Build image tensor and temporal frequency features images_tensor = [] ft_images = [] for image in images: image = np.array(image) img_pil = Image.fromarray(image) img_tensor = test_transform(img_pil) images_tensor.append(img_tensor) # Median filter residual -> grayscale for FFT img_filtered = cv2.medianBlur(image.copy(), 5) residual = cv2.cvtColor( (image - img_filtered), cv2.COLOR_RGB2GRAY ) ft_images.append(residual) # Temporal FFT: take first half of frequencies ft_array = np.array(ft_images) ft_array = np.absolute( np.fft.fft(ft_array, axis=0)[:CLIP_SIZE // 2] * (1.0 / CLIP_SIZE) ) ft_tensor = torch.from_numpy(ft_array).to(_device).unsqueeze(0) # Stack image frames: (1, C, T, H, W) img_stack = torch.stack(images_tensor, dim=1).unsqueeze(0) img_stack = img_stack.to(_device) with torch.no_grad(): output = _model(img_stack, ft_tensor) output = torch.sigmoid(output).squeeze(0) pred = float(output.item()) preds.append(pred) # ── Step 6: Aggregate ─────────────────────────────────────────── probability = float(np.mean(preds)) return { "probability": probability, "num_frames": len(frames), "num_tracks": len(tracks), "num_clips": len(clips_for_video), "frames_with_faces": len(frame_boxes), } # ── Prediction endpoint ───────────────────────────────────────────────────── @app.post("/predict", response_model=PredictResponse) async def predict(request: PredictRequest): """Run PwTF-DVD deepfake detection on a base64-encoded video. Pipeline: 1. Decode video and write to temp file. 2. Run full PwTF-DVD pipeline (face detection, tracking, temporal FFT, dual-stream classification). 3. Aggregate per-clip predictions into a single probability. If no faces are detected the service returns probability=0.5 (undetermined) rather than raising an error. """ if not _is_model_loaded(): _load_models() start_time = time.time() # ── Decode video ──────────────────────────────────────────────── with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp: try: video_bytes = base64.b64decode(request.video_data) tmp.write(video_bytes) tmp_path = tmp.name except Exception as exc: raise HTTPException( status_code=400, detail=f"Failed to decode video: {exc}", ) try: # ── Run inference pipeline ────────────────────────────────── result = _run_inference(tmp_path) probability = result["probability"] prediction = 1 if probability >= request.threshold else 0 class_name = "fake" if prediction == 1 else "real" return PredictResponse( probability=probability, prediction=prediction, class_name=class_name, inference_time=time.time() - start_time, metadata={ "frames_sampled": result["num_frames"], "frames_with_faces": result["frames_with_faces"], "num_tracks": result["num_tracks"], "num_clips": result["num_clips"], "device": str(_device), }, ) except HTTPException: raise except Exception as exc: logger.exception("Error during PwTF-DVD prediction") raise HTTPException(status_code=500, detail=str(exc)) finally: if os.path.exists(tmp_path): os.remove(tmp_path) gc.collect() if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=MODEL_PORT)