File size: 3,901 Bytes
59630ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import List
import hashlib
import random
import torch
from einops import rearrange
from PIL import Image
import numpy as np


class VideoTransform:
    """
    Adapted from pixelSplat
    https://github.dev/dcharatan/pixelsplat/blob/main/src/dataset/dataset_re10k.py
    """

    def __init__(self, shape: tuple[int, int]):
        self.shape = shape

    def __call__(
        self, images: torch.Tensor  # (*batch, c, h, w)
    ) -> torch.Tensor:  # (*batch, c, *shape)
        return self._rescale_and_crop(images, self.shape)

    @classmethod
    def _rescale(
        cls,
        image: torch.Tensor,  # (c, h, w),
        shape: tuple[int, int],
    ) -> torch.Tensor:  # (c, *shape)
        h, w = shape
        image_new = (image * 255).clip(min=0, max=255).type(torch.uint8)
        image_new = rearrange(image_new, "c h w -> h w c").detach().cpu().numpy()
        image_new = Image.fromarray(image_new)
        image_new = image_new.resize((w, h), Image.Resampling.LANCZOS)
        image_new = np.array(image_new) / 255
        image_new = torch.tensor(image_new, dtype=image.dtype, device=image.device)
        return rearrange(image_new, "h w c -> c h w")

    @classmethod
    def _center_crop(
        cls,
        images: torch.Tensor,  # (*batch, c, h, w),
        shape: tuple[int, int],
    ) -> torch.Tensor:
        *_, h_in, w_in = images.shape
        h_out, w_out = shape

        # Note that odd input dimensions induce half-pixel misalignments.
        row = (h_in - h_out) // 2
        col = (w_in - w_out) // 2

        # Center-crop the image.
        images = images[..., :, row : row + h_out, col : col + w_out]

        return images

    @classmethod
    def _rescale_and_crop(
        cls,
        images: torch.Tensor,  # (*batch, c, h, w),
        shape: tuple[int, int],
    ):
        """
        Rescale and crop the images to the specified shape.
        Args:
            images (torch.Tensor): images tensor of shape (*batch, c, h, w). Range [0, 1].
            shape (tuple[int, int]): target shape.
        Returns:
            torch.Tensor: rescaled and cropped images tensor of shape (*batch, c, *shape). Range [0, 1].
        """
        *_, h_in, w_in = images.shape
        h_out, w_out = shape
        # assert h_out <= h_in and w_out <= w_in

        scale_factor = max(h_out / h_in, w_out / w_in)
        h_scaled = round(h_in * scale_factor)
        w_scaled = round(w_in * scale_factor)
        assert h_scaled == h_out or w_scaled == w_out

        *batch, c, h, w = images.shape
        images = images.reshape(-1, c, h, w)
        images = torch.stack(
            [cls._rescale(image, (h_scaled, w_scaled)) for image in images]
        )
        images = images.reshape(*batch, c, h_scaled, w_scaled)

        return cls._center_crop(images, shape)


def rescale_and_crop(
    video: torch.Tensor,
    resolution: int,
) -> np.ndarray:
    """
    Rescale and crop the video to the specified resolution. Used for preprocessing.
    Args:
        video (torch.Tensor): video tensor of shape (t, h, w, c). uint8.
        resolution (int): target resolution.
    Returns:
        np.ndarray: rescaled and cropped video tensor of shape (t, resolution, resolution, c). uint8.
    """
    *_, h, w, _ = video.shape
    scale_factor = max(resolution / h, resolution / w)
    h_scaled, w_scaled = round(h * scale_factor), round(w * scale_factor)
    assert h_scaled == resolution or w_scaled == resolution
    row = (h_scaled - resolution) // 2
    col = (w_scaled - resolution) // 2

    def _rescale_and_crop(image: torch.Tensor) -> torch.Tensor:
        image = Image.fromarray(image.numpy())
        image = image.resize((w_scaled, h_scaled), Image.Resampling.LANCZOS)
        return np.array(image)[row : row + resolution, col : col + resolution]

    video = np.stack([_rescale_and_crop(frame) for frame in video])
    return video