File size: 3,379 Bytes
d72cff1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 | """
preprocessing.py
Lightweight video preprocessing for deepfake detection.
Input:
MP4/AVI/MOV/etc.
Output:
Tensor with shape [T, C, H, W], where T is the sampled frame count.
The default pipeline samples 16 frames uniformly from the video, resizes
them to 224x224, and applies ImageNet normalization.
"""
from pathlib import Path
from typing import List
import cv2
import numpy as np
import torch
from PIL import Image
from torchvision import transforms
DEFAULT_NUM_FRAMES = 16
DEFAULT_IMAGE_SIZE = 224
def build_transform(image_size: int = DEFAULT_IMAGE_SIZE):
"""Transform compatible with ImageNet-pretrained CNN backbones."""
return transforms.Compose([
transforms.Resize((image_size, image_size)),
transforms.ToTensor(),
transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225],
),
])
def sample_frame_indices(total_frames: int, num_frames: int = DEFAULT_NUM_FRAMES) -> np.ndarray:
"""Return uniformly distributed frame indices."""
if total_frames <= 0:
raise ValueError("Video contains no readable frames.")
if total_frames <= num_frames:
return np.linspace(0, total_frames - 1, num_frames).astype(int)
return np.linspace(0, total_frames - 1, num_frames).astype(int)
def read_video_frames(
video_path: str,
num_frames: int = DEFAULT_NUM_FRAMES,
) -> List[Image.Image]:
"""Read uniformly sampled RGB frames from a video."""
cap = cv2.VideoCapture(str(video_path))
if not cap.isOpened():
raise ValueError(f"Could not open video: {video_path}")
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
indices = sample_frame_indices(total_frames, num_frames)
wanted = set(indices.tolist())
frames = {}
frame_idx = 0
while True:
ok, frame = cap.read()
if not ok:
break
if frame_idx in wanted:
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames[frame_idx] = Image.fromarray(frame)
frame_idx += 1
if frame_idx > indices[-1]:
break
cap.release()
# Keep the requested order and duplicate the last available frame if
# a video has fewer readable frames than expected.
output = [frames[i] for i in indices if i in frames]
if not output:
raise ValueError(f"No readable frames found in: {video_path}")
while len(output) < num_frames:
output.append(output[-1].copy())
return output[:num_frames]
def preprocess_video(
video_path: str,
num_frames: int = DEFAULT_NUM_FRAMES,
image_size: int = DEFAULT_IMAGE_SIZE,
) -> torch.Tensor:
"""
Convert one video into a normalized tensor.
Returns:
Tensor of shape [T, C, H, W].
"""
transform = build_transform(image_size)
frames = read_video_frames(video_path, num_frames)
tensor = torch.stack([transform(frame) for frame in frames])
return tensor
def save_preprocessed_video(
video_path: str,
output_path: str,
num_frames: int = DEFAULT_NUM_FRAMES,
image_size: int = DEFAULT_IMAGE_SIZE,
) -> None:
"""Optional utility to cache a processed video as a .pt file."""
tensor = preprocess_video(video_path, num_frames, image_size)
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
torch.save(tensor, output_path)
|