Map-Det3D / mapdet3d /vis /image /depth_visualizer.py
RoyYang0714's picture
feat: Add the Gradio demo for Map-Det3D.
0122a25
Raw
History Blame Contribute Delete
6.15 kB
"""Depth visualizer."""
from __future__ import annotations
import os
from dataclasses import dataclass
import numpy as np
from PIL import Image
from mapdet3d.common.array import array_to_numpy
from mapdet3d.common.typing import (
ArgsType,
ArrayLikeFloat,
NDArrayF32,
NDArrayUI8,
)
from mapdet3d.vis.base import Visualizer
from mapdet3d.vis.image.util import preprocess_image
from mapdet3d.vis.util import generate_color_map
from .util import (
colorize,
get_pointcloud_from_rgbd,
save_depth_map,
save_file_ply,
)
@dataclass
class DataSample:
"""Dataclass storing a data sample that can be visualized."""
image: NDArrayUI8
image_name: str
depth: NDArrayF32
depth_gt: NDArrayF32 | None = None
depth_error: NDArrayF32 | None = None
points_rgb: NDArrayF32 | None = None
class DepthVisualizer(Visualizer):
"""Depth visualizer class."""
def __init__(
self,
*args: ArgsType,
max_depth: None | float = None,
plot_error: bool = False,
lift: bool = False,
color_palette: list[tuple[int, int, int]] | None = None,
**kwargs: ArgsType,
) -> None:
"""Creates a new Visualizer for Depth.
Args:
max_depth (None | float): Maximum depth to visualize.
"""
super().__init__(*args, **kwargs)
self.max_depth = max_depth
self._samples: list[DataSample] = []
self._gt_samples = []
self.plot_error = plot_error
self.lift = lift
self.color_palette = (
generate_color_map(50) if color_palette is None else color_palette
)
def __repr__(self):
"""String representation."""
return f"DepthVisualizer(max_depth={self.max_depth}, plot_error={self.plot_error}, lift={self.lift})"
def reset(self) -> None:
"""Reset the visualizer."""
self._samples.clear()
self._gt_samples.clear()
def process(
self,
cur_iter: int,
images: list[ArrayLikeFloat],
image_names: list[str],
depths: list[ArrayLikeFloat],
depth_gts: ArrayLikeFloat | None = None,
intrinsics: ArrayLikeFloat | None = None,
) -> None:
"""Process data of a batch of data."""
if self._run_on_batch(cur_iter):
for i, image in enumerate(images):
image = preprocess_image(image)
self._samples.append(
self.process_single_image(
image,
image_names[i],
array_to_numpy(depths[i]),
(
array_to_numpy(depth_gts[i])
if depth_gts is not None
else None
),
(
array_to_numpy(intrinsics[i])
if intrinsics is not None
else None
),
)
)
def process_single_image(
self,
image: NDArrayUI8,
image_name: str,
depth: NDArrayF32,
depth_gt: NDArrayF32 | None = None,
intrinsic: NDArrayF32 | None = None,
) -> DataSample:
"""Process data of a batch of data."""
if self.max_depth is not None:
mask = depth <= self.max_depth
else:
mask = np.full(depth.shape, True)
if self.plot_error:
assert (
depth_gt is not None
), "Ground truth depth is required for plotting error."
error = np.zeros_like(depth_gt)
error[depth_gt > 0] = (
np.abs(depth_gt - depth)[depth_gt > 0] / depth_gt[depth_gt > 0]
)
else:
error = None
if self.lift:
assert (
intrinsic is not None
), "Intrinsic matrix is required for lifting."
points_rgb = get_pointcloud_from_rgbd(
image, depth, intrinsic, mask
)
else:
points_rgb = None
return DataSample(
image=image,
image_name=image_name,
depth=depth,
depth_gt=depth_gt,
depth_error=error,
points_rgb=points_rgb,
)
def save_to_disk(self, cur_iter: int, output_folder: str) -> None:
"""Saves the visualization to disk.
Args:
cur_iter (int): Current iteration.
output_folder (str): Folder where the output should be written.
"""
if self._run_on_batch(cur_iter):
for sample in self._samples:
Image.fromarray(sample.image).save(
f"{output_folder}/{sample.image_name}.png",
)
if self.plot_error:
error = sample.depth_error
error_image = Image.fromarray(
colorize(
error.clip(0.0, 0.3),
vmin=0.001,
vmax=0.3,
cmap="coolwarm",
)
)
error_image.save(
f"{output_folder}/{sample.image_name}_error.png"
)
save_depth_map(
sample.depth,
f"{output_folder}/{sample.image_name}_pred.png",
)
if sample.depth_gt is not None:
save_depth_map(
sample.depth_gt,
f"{output_folder}/{sample.image_name}_gt.png",
)
if self.lift:
if sample.points_rgb is not None:
save_file_ply(
sample.points_rgb[:, :3],
sample.points_rgb[:, 3:],
os.path.join(
output_folder, f"{sample.image_name}.ply"
),
)