deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
30.2 kB
"""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)