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)