Download GeometryForcing/datasets/video/utils/transform.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 3.9 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/datasets/video/utils/transform.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/GeometryForcing/datasets/video/utils/transform.py
-
curl -L -o transform.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/datasets/video/utils/transform.py
3.9 kB
| 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) | |
| 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") | |
| 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 | |
| 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 | |