(
bgr_images: list[np.ndarray],
depth_model: DepthAnythingV2,
vae: VQVAE,
*,
depth_input_size: int,
vae_image_size: int,
batch_size: int,
device: torch.device,
)
| 853 | |
| 854 | @torch.no_grad() |
| 855 | def encode_depth_codes( |
| 856 | bgr_images: list[np.ndarray], |
| 857 | depth_model: DepthAnythingV2, |
| 858 | vae: VQVAE, |
| 859 | *, |
| 860 | depth_input_size: int, |
| 861 | vae_image_size: int, |
| 862 | batch_size: int, |
| 863 | device: torch.device, |
| 864 | ) -> tuple[list[np.ndarray], list[np.ndarray]]: |
| 865 | exact = batch_size == 1 |
| 866 | rgb_resized: list[np.ndarray] = [] |
| 867 | codes: list[np.ndarray] = [] |
| 868 | for start in range(0, len(bgr_images), batch_size): |
| 869 | batch_images = bgr_images[start : start + batch_size] |
| 870 | tensors: list[Tensor] = [] |
| 871 | original_sizes: list[tuple[int, int]] = [] |
| 872 | for bgr in batch_images: |
| 873 | tensor, size = depth_image_to_tensor(bgr, depth_input_size, device) |
| 874 | tensors.append(tensor) |
| 875 | original_sizes.append(size) |
| 876 | rgb_resized.append(cv2.resize(bgr, (vae_image_size, vae_image_size), interpolation=cv2.INTER_LINEAR)) |
| 877 | |
| 878 | if exact: |
| 879 | for idx, tensor in enumerate(tensors): |
| 880 | depth = depth_model(tensor[None])[0] |
| 881 | height, width = original_sizes[idx] |
| 882 | resized_depth = nn_func.interpolate( |
| 883 | depth[None, None], |
| 884 | (height, width), |
| 885 | mode="bilinear", |
| 886 | align_corners=True, |
| 887 | )[0, 0] |
| 888 | raw = resized_depth.detach().cpu().numpy() |
| 889 | dmin = raw.min() |
| 890 | dmax = raw.max() |
| 891 | depth_vis = ((raw - dmin) / (dmax - dmin + 1e-8) * 255.0).astype(np.uint8) |
| 892 | depth_vis = cv2.resize( |
| 893 | depth_vis, |
| 894 | (vae_image_size, vae_image_size), |
| 895 | interpolation=cv2.INTER_NEAREST, |
| 896 | ) |
| 897 | depth_tensor = torch.from_numpy(depth_vis).float() / 255.0 |
| 898 | depth_tensor = (depth_tensor - 0.5) / 0.5 |
| 899 | depth_tensor = depth_tensor.unsqueeze(0).unsqueeze(0).to(device) |
| 900 | code_grid = vae(img=depth_tensor, return_indices=True)[0] |
| 901 | codes.append(code_grid.detach().cpu().numpy().astype(np.int16)) |
| 902 | else: |
| 903 | codes_by_index: list[np.ndarray | None] = [None] * len(tensors) |
| 904 | shape_groups: dict[tuple[int, int, int, int], list[int]] = {} |
| 905 | for idx, tensor in enumerate(tensors): |
| 906 | input_height, input_width = tensor.shape[-2:] |
| 907 | original_height, original_width = original_sizes[idx] |
| 908 | shape_groups.setdefault((input_height, input_width, original_height, original_width), []).append(idx) |
| 909 | for indices in shape_groups.values(): |
| 910 | input_height, input_width = tensors[indices[0]].shape[-2:] |
| 911 | original_height, original_width = original_sizes[indices[0]] |
| 912 | max_depth_batch = max(1, (2**31 - 1) // (64 * input_height * input_width)) |
no test coverage detected