"""MINTIME video deepfake detection service. Wraps the MINTIME model (IEEE T-IFS 2024, Multi-Identity-size-iNvariant TIMEsformer) with a FastAPI endpoint. Uses Xception as a feature extractor and a Size-Invariant TimeSformer for temporal classification with identity-aware attention masks. The pipeline is: 1. Extract frames from the video. 2. Detect faces per frame with MTCNN. 3. Crop faces, cluster them by identity (InceptionResnetV1 embeddings). 4. Build identity-ordered sequences with size embeddings and masks. 5. Extract Xception features, run the TimeSformer, return sigmoid score. Reference: Coccomini et al., "MINTIME: Multi-Identity Size-Invariant Video Deepfake Detection", IEEE T-IFS 2024. """ import base64 import gc import logging import os import platform import sys import tempfile import threading import time from statistics import mean from typing import Any, Dict, List, Optional, Tuple import cv2 import networkx as nx import numpy as np import torch import torch.nn.functional as F import uvicorn from einops import rearrange from facenet_pytorch import InceptionResnetV1, fixed_image_standardization from fastapi import FastAPI, HTTPException from PIL import Image from pydantic import BaseModel, ConfigDict, Field from torch import nn # Prepend model_code to sys.path so vendored modules resolve correctly. _MODEL_CODE_DIR = os.path.join(os.path.dirname(__file__), "model_code") sys.path.insert(0, _MODEL_CODE_DIR) # The vendored transforms/albu.py imports crop from an old albumentations # path that was removed in v2.x. The function is never actually called # (IsotropicResize only uses cv2.resize), so inject a harmless stub. import types as _types _compat = _types.ModuleType("albumentations.augmentations.functional") _compat.crop = None sys.modules.setdefault("albumentations.augmentations.functional", _compat) from models.size_invariant_timesformer import SizeInvariantTimeSformer # noqa: E402 from models.xception import xception # noqa: E402 from transforms.albu import IsotropicResize # noqa: E402 logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) MODEL_PORT = int(os.environ.get("MODEL_PORT", 7008)) PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true" MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", 1800)) # Paths -- extractor and SizeInvariantTimeSformer checkpoints EXTRACTOR_WEIGHTS_PATH = os.environ.get( "EXTRACTOR_WEIGHTS_PATH", "/app/weights/MINTIME/MINTIME_XC_Extractor_checkpoint30", ) MODEL_WEIGHTS_PATH = os.environ.get( "MODEL_WEIGHTS_PATH", "/app/weights/MINTIME/MINTIME_XC_Model_checkpoint30", ) # Model hyperparameters from size_invariant_timesformer.yaml NUM_FRAMES = 16 IMAGE_SIZE = 224 NUM_PATCHES = 49 # 7 x 7 spatial from Xception feature map MAX_IDENTITIES = 2 RANGE_SIZE = 5 SIZE_EMB_DICT = [ (1 + i * RANGE_SIZE, (i + 1) * RANGE_SIZE) if i != 0 else (0, RANGE_SIZE) for i in range(20) ] # Model config dict matching size_invariant_timesformer.yaml MODEL_CONFIG = { "model": { "image-size": IMAGE_SIZE, "patch-size": 1, "num-classes": 1, "num-patches": NUM_PATCHES, "num-frames": NUM_FRAMES, "max-identities": MAX_IDENTITIES, "dim": 512, "depth": 9, "dim-head": 64, "channels": 2048, "heads": 8, "attn-dropout": 0.0, "ff-dropout": 0.0, "shift-tokens": False, "enable-size-emb": True, "enable-pos-emb": True, "enable-identity-attention": True, } } # ── Helpers ────────────────────────────────────────────────────────────── def _generate_connected_components( similarities: np.ndarray, similarity_threshold: float = 0.80, ) -> List[List[int]]: """Build a similarity graph and return connected components.""" graph = nx.Graph() n = len(similarities) for i in range(n): for j in range(i + 1, n): if similarities[i, j] > similarity_threshold: graph.add_edge(i, j) components = [sorted(c) for c in nx.connected_components(graph)] # Include isolated nodes (faces not similar to any other) all_in_components = set() for c in components: all_in_components.update(c) for i in range(n): if i not in all_in_components: components.append([i]) return components def _preprocess_face_for_clustering(img: Image.Image) -> np.ndarray: """Resize a PIL face crop to 128x128 for embedding extraction.""" from torchvision.transforms import Resize return np.asarray(Resize([128, 128])(img)) # ── Device selection ───────────────────────────────────────────────────── def _get_device() -> torch.device: """Select optimal device: CUDA > MPS > 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 ───────────────────────────────────────────────────────── _extractor: Optional[nn.Module] = None _model: Optional[nn.Module] = None _embedding_model: Optional[nn.Module] = None _device: Optional[torch.device] = None _load_lock = threading.Lock() def _strip_module_prefix(state_dict: dict) -> dict: """Remove 'module.' prefix from DataParallel state dicts.""" new_sd = {} for k, v in state_dict.items(): new_key = k.replace("module.", "", 1) if k.startswith("module.") else k new_sd[new_key] = v return new_sd def _load_models() -> None: """Load Xception extractor, SizeInvariantTimeSformer, and InceptionResnetV1 for identity clustering (thread-safe).""" global _extractor, _model, _embedding_model, _device if _model is not None: return with _load_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") 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 -- inference may be slow)", _device, ) logger.info("Loading MINTIME models on %s ...", _device) # ── Xception feature extractor ────────────────────────────── if not os.path.exists(EXTRACTOR_WEIGHTS_PATH): raise FileNotFoundError( f"Extractor weights not found: {EXTRACTOR_WEIGHTS_PATH}" ) feat_ext = xception(num_classes=1, pretrain_path=None) ext_sd = torch.load( EXTRACTOR_WEIGHTS_PATH, map_location="cpu", weights_only=False ) feat_ext.load_state_dict(_strip_module_prefix(ext_sd)) feat_ext = feat_ext.to(_device) feat_ext.train(False) _extractor = feat_ext # ── SizeInvariantTimeSformer ──────────────────────────────── if not os.path.exists(MODEL_WEIGHTS_PATH): raise FileNotFoundError(f"Model weights not found: {MODEL_WEIGHTS_PATH}") sit = SizeInvariantTimeSformer(config=MODEL_CONFIG, require_attention=False) model_sd = torch.load( MODEL_WEIGHTS_PATH, map_location="cpu", weights_only=False ) sit.load_state_dict(_strip_module_prefix(model_sd)) sit = sit.to(_device) sit.train(False) _model = sit # ── InceptionResnetV1 for identity clustering ─────────────── emb = InceptionResnetV1(pretrained="vggface2").to(_device) emb.train(False) _embedding_model = emb logger.info("MINTIME models loaded successfully.") def _is_model_loaded() -> bool: """Return True if all three models are loaded.""" return ( _model is not None and _extractor is not None and _embedding_model is not None ) # ── FastAPI app ────────────────────────────────────────────────────────── app = FastAPI( title="MINTIME Detection Service", description=( "Multi-Identity Size-Invariant Video Deepfake Detection " "(Xception + TimeSformer, IEEE T-IFS 2024)" ), 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 = "mintime_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": "mintime_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_memory / 1024**2 ), } return {} @app.get("/health") def health(): """Health check endpoint.""" return { "status": "healthy", "model": "mintime_detection", "device": str(_device) if _device else "cpu", "model_loaded": _is_model_loaded(), "extractor_weights_exist": os.path.exists(EXTRACTOR_WEIGHTS_PATH), "model_weights_exist": os.path.exists(MODEL_WEIGHTS_PATH), **_gpu_health_info(), } # ── Video processing utilities ─────────────────────────────────────────── def _extract_frames( video_path: str, ) -> Tuple[List[np.ndarray], int, int, int]: """Extract all frames from a video file. Returns: Tuple of (all_frames, fps, width, height). """ cap = cv2.VideoCapture(video_path) fps = max(int(cap.get(cv2.CAP_PROP_FPS)), 1) width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) frames: List[np.ndarray] = [] while True: ret, frame = cap.read() if not ret: break frames.append(frame) cap.release() return frames, fps, width, height def _detect_faces_mtcnn( frames: List[np.ndarray], fps: int, ) -> Dict[str, Any]: """Detect faces in sampled frames using MTCNN. Samples one frame per second (every fps frames) and runs MTCNN on half-resolution PIL images, matching the original preprocessing. Returns: Dict mapping str(frame_index) to list of bboxes, or None. """ from facenet_pytorch import MTCNN mtcnn = MTCNN( device=_device, thresholds=[0.85, 0.95, 0.95], margin=0, ) bboxes_dict: Dict[str, Any] = {} # Sample frames at ~1 per second indices = list(range(0, len(frames), max(fps, 1))) if not indices: indices = [0] if frames else [] for idx in indices: frame = frames[idx] rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) pil_img = Image.fromarray(rgb) # Half resolution, matching original VideoDataset pil_img = pil_img.resize([s // 2 for s in pil_img.size]) boxes, _ = mtcnn.detect(pil_img) if boxes is not None: bboxes_dict[str(idx)] = boxes.tolist() else: bboxes_dict[str(idx)] = None return bboxes_dict def _extract_crops( frames: List[np.ndarray], bboxes_dict: Dict[str, Any], fps: int, ) -> List[Tuple[int, Image.Image, list]]: """Extract face crops from video frames using detected bboxes. Follows the original extract_crops logic: iterate per-second windows, find the nearest frame with bboxes, crop with padding, make square. Returns: List of (frame_index, pil_crop, bbox). """ frames_num = len(frames) crops = [] for i in range(0, frames_num, fps): # Find nearest frame with valid bboxes in this window idx = i limit = min(i + fps - 1, frames_num - 1) # Walk forward to find a frame with bboxes while str(idx) not in bboxes_dict or bboxes_dict.get(str(idx)) is None: if idx >= limit: break idx += 1 if str(idx) not in bboxes_dict or bboxes_dict.get(str(idx)) is None: continue bboxes = bboxes_dict[str(idx)] frame = frames[i] if i < frames_num else frames[-1] for bbox in bboxes: xmin, ymin, xmax, ymax = [int(b * 2) for b in bbox] w = xmax - xmin h = ymax - ymin if w <= 0 or h <= 0: continue # Padding p_h = h // 3 p_w = w // 3 crop_h = (ymax + p_h) - max(ymin - p_h, 0) crop_w = (xmax + p_w) - max(xmin - p_w, 0) # Make square if crop_h > crop_w: p_h -= int((crop_h - crop_w) / 2) else: p_w -= int((crop_w - crop_h) / 2) crop = frame[ max(ymin - p_h, 0) : ymax + p_h, max(xmin - p_w, 0) : xmax + p_w, ] h_c, w_c = crop.shape[:2] if h_c <= 0 or w_c <= 0: continue # Final square trim if h_c > w_c: diff = (h_c - w_c) // 2 if diff > 0: crop = crop[diff:-diff, :] else: crop = crop[1:, :] elif h_c < w_c: diff = (w_c - h_c) // 2 if diff > 0: crop = crop[:, diff:-diff] else: crop = crop[:, :-1] if crop.size == 0: continue rgb_crop = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) crops.append((i, Image.fromarray(rgb_crop), bbox)) return crops def _cluster_faces( crops: List[Tuple[int, Image.Image, list]], similarity_threshold: float = 0.45, ) -> Dict[int, List]: """Cluster face crops by identity using InceptionResnetV1 embeddings. Returns: Dict mapping identity_index to list of (frame_idx, pil_img, bbox). """ if not crops: return {} crops_images = [row[1] for row in crops] # Prepare face tensors for embedding extraction faces = [_preprocess_face_for_clustering(face) for face in crops_images] faces = np.stack([np.uint8(f) for f in faces]) faces_tensor = torch.as_tensor(faces).permute(0, 3, 1, 2).float() faces_tensor = fixed_image_standardization(faces_tensor) faces_tensor = faces_tensor.to(_device) with torch.no_grad(): embeddings = _embedding_model(faces_tensor).cpu().numpy() # Cosine similarity matrix similarities = np.dot(embeddings, embeddings.T) components = _generate_connected_components( similarities, similarity_threshold=similarity_threshold ) clustered_faces: Dict[int, List] = {} for identity_index, component in enumerate(components): clustered_faces[identity_index] = [crops[fi] for fi in component] return clustered_faces def _get_sorted_identities( identities: Dict[int, List], num_frames: int = NUM_FRAMES, max_identities: int = MAX_IDENTITIES, ) -> List[list]: """Sort identities by face size and allocate frame slots. Returns list of [identity_id, mean_side, num_faces, faces_list]. """ sorted_ids = [] for identity in identities: faces = identities[identity] mean_side = mean([row[1].size[0] for row in faces]) sorted_ids.append([identity, mean_side, len(faces), faces]) # Sort by face size descending (largest first) sorted_ids.sort(key=lambda x: x[1], reverse=True) if len(sorted_ids) > max_identities: sorted_ids = sorted_ids[:max_identities] identities_number = len(sorted_ids) available_additional = [] if identities_number > 1: max_faces_map = { 1: [num_frames], 2: [num_frames // 2, num_frames // 2], 3: [num_frames // 3, num_frames // 3, num_frames // 4], 4: [num_frames // 3, num_frames // 3, num_frames // 8, num_frames // 8], } alloc = max_faces_map[identities_number] for i in range(identities_number): if sorted_ids[i][2] < alloc[i] and i < identities_number - 1: sorted_ids[i + 1][2] += alloc[i] - sorted_ids[i][2] available_additional.append(0) elif sorted_ids[i][2] > alloc[i]: available_additional.append(sorted_ids[i][2] - alloc[i]) sorted_ids[i][2] = alloc[i] else: available_additional.append(0) else: sorted_ids[0][2] = num_frames available_additional.append(0) # Fill remaining slots if needed input_len = sum(r[2] for r in sorted_ids) if input_len < num_frames: for i in range(identities_number): needed = num_frames - input_len if available_additional[i] > 0: added = min(available_additional[i], needed) sorted_ids[i][2] += added input_len += added if input_len == num_frames: break if input_len < num_frames: sorted_ids[-1][2] += num_frames - input_len return sorted_ids def _create_val_transform(size: int, additional_targets: dict): """Create the validation-time albumentations transform.""" from albumentations import Compose, PadIfNeeded, Resize return Compose( [ IsotropicResize( max_side=size, interpolation_down=cv2.INTER_AREA, interpolation_up=cv2.INTER_CUBIC, ), PadIfNeeded( min_height=size, min_width=size, border_mode=cv2.BORDER_CONSTANT, ), Resize(height=size, width=size), ], additional_targets=additional_targets, ) def _generate_masks( video_width: int, video_height: int, identities: List[list], num_frames: int = NUM_FRAMES, image_size: int = IMAGE_SIZE, num_patches: int = NUM_PATCHES, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, list]: """Build input tensors, masks, size embeddings, and positions. Mirrors generate_masks from predict.py. """ mask = [] sequence = [] size_embeddings = [] images_frames = [] video_area = video_width * video_height / 2 for identity in identities: max_faces = identity[2] identity_images = identity[3] # Uniform sampling if too many faces if len(identity_images) > max_faces: idx = np.round(np.linspace(0, len(identity_images) - 2, max_faces)).astype( int ) identity_images = list(np.asarray(identity_images, dtype=object)[idx]) images_frames.extend(img[0] for img in identity_images) pil_images = [img[1] for img in identity_images] # Size embeddings based on face-frame area ratio identity_size_embs = [] for img in pil_images: face_area = img.size[0] * img.size[1] ratio = int(face_area * 100 / max(video_area, 1)) side_ranges = list( map( lambda a_: ratio in range(a_[0], a_[1] + 1), SIZE_EMB_DICT, ) ) matches = np.where(side_ranges)[0] identity_size_embs.append(int(matches[0] + 1) if len(matches) > 0 else 1) # Pad with empty frames if needed if len(pil_images) < max_faces: diff = max_faces - len(identity_size_embs) identity_size_embs = list(identity_size_embs) + [0] * diff pil_images.extend( [ np.zeros((image_size, image_size, 3), dtype=np.uint8) for _ in range(diff) ] ) mask.extend([1 if i < max_faces - diff else 0 for i in range(max_faces)]) images_frames.extend([max(images_frames)] * diff) else: mask.extend([1] * max_faces) size_embeddings.extend(identity_size_embs) sequence.extend(pil_images) # Convert PIL images to numpy arrays sequence = [np.asarray(img) for img in sequence] # Apply albumentations transform to all frames together additional_targets_keys = [f"image{i}" for i in range(num_frames)] additional_targets_values = ["image"] * num_frames additional_targets = dict(zip(additional_targets_keys, additional_targets_values)) transform = _create_val_transform(image_size, additional_targets) # Build transform kwargs transform_kwargs = {"image": sequence[0]} for i in range(1, len(sequence)): transform_kwargs[f"image{i}"] = sequence[i] transformed = transform(**transform_kwargs) sequence = [transformed[k] for k in transformed] # Build identities_mask identities_mask = [] last_range_end = 0 for identity in identities: n_faces = identity[2] identity_mask = [ last_range_end <= i < last_range_end + n_faces for i in range(num_frames) ] for _ in range(n_faces): identities_mask.append(identity_mask) last_range_end += n_faces # Coherent temporal-positional embedding images_frames_positions = { k: v + 1 for v, k in enumerate(sorted(set(images_frames))) } frame_positions = [images_frames_positions[f] for f in images_frames] if num_patches is not None: positions = [] for fp in frame_positions: positions.extend( [i + 1 for i in range((fp - 1) * num_patches, num_patches * fp)] ) positions.insert(0, 0) # CLS token position else: positions = [] tokens_per_identity = [] for i, ident in enumerate(identities): if i > 0: tokens_per_identity.append( (ident[0], ident[2] * num_patches + identities[i - 1][2] * num_patches) ) else: tokens_per_identity.append((ident[0], ident[2] * num_patches)) return ( torch.tensor([sequence]).float(), torch.tensor([size_embeddings]).int(), torch.tensor([mask]).bool(), torch.tensor([identities_mask]).bool(), torch.tensor([positions]), tokens_per_identity, ) # ── Prediction endpoint ───────────────────────────────────────────────── @app.post("/predict", response_model=PredictResponse) async def predict(request: PredictRequest): """Run MINTIME deepfake detection on a base64-encoded video. Pipeline: 1. Decode video, extract all frames. 2. Detect faces per-second with MTCNN. 3. Crop faces and cluster by identity. 4. Build identity-ordered input with masks and size embeddings. 5. Extract Xception features, run TimeSformer. 6. Return sigmoid probability. Returns probability=0.5 if no faces are detected. """ 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 e: raise HTTPException(status_code=400, detail=f"Failed to decode video: {e}") try: # ── Extract frames ────────────────────────────────────────── frames, fps, vid_w, vid_h = _extract_frames(tmp_path) if not frames: raise HTTPException( status_code=400, detail="Could not extract frames from video.", ) # ── Detect faces ──────────────────────────────────────────── bboxes_dict = _detect_faces_mtcnn(frames, fps) # Check if any faces found has_faces = any(v is not None and len(v) > 0 for v in bboxes_dict.values()) if not has_faces: return PredictResponse( probability=0.5, prediction=0, class_name="real", inference_time=time.time() - start_time, metadata={ "frames_extracted": len(frames), "faces_detected": 0, "identities": 0, "note": "no faces detected", "device": str(_device), }, ) # ── Extract crops ─────────────────────────────────────────── crops = _extract_crops(frames, bboxes_dict, fps) if not crops: return PredictResponse( probability=0.5, prediction=0, class_name="real", inference_time=time.time() - start_time, metadata={ "frames_extracted": len(frames), "faces_detected": 0, "identities": 0, "note": "no valid face crops", "device": str(_device), }, ) # ── Cluster by identity ───────────────────────────────────── clustered = _cluster_faces(crops) num_identities = len(clustered) # ── Build identity sequence ───────────────────────────────── sorted_identities = _get_sorted_identities(clustered) ( videos_tensor, size_embeddings, mask_tensor, identities_mask, positions, tokens_per_identity, ) = _generate_masks( vid_w, vid_h, sorted_identities, ) # ── Run inference ─────────────────────────────────────────── b, f, h, w, c = videos_tensor.shape videos_tensor = videos_tensor.to(_device) identities_mask = identities_mask.to(_device) mask_tensor = mask_tensor.to(_device) positions = positions.to(_device) with torch.no_grad(): # Feature extraction: (B*F, C, H, W) video_input = rearrange(videos_tensor, "b f h w c -> (b f) c h w") features = _extractor(video_input) # (B*F, 2048, 7, 7) features = rearrange(features, "(b f) c h w -> b f c h w", b=b, f=f) # TimeSformer classification pred = _model( features, mask=mask_tensor, size_embedding=size_embeddings, identities_mask=identities_mask, positions=positions, ) probability = float(torch.sigmoid(pred[0]).item()) 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_extracted": len(frames), "faces_detected": len(crops), "identities": num_identities, "device": str(_device), }, ) except HTTPException: raise except Exception as e: logger.exception("Error during MINTIME prediction") raise HTTPException(status_code=500, detail=str(e)) 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)