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
|