WorldString_live_demo / robot_backends /gradio_plot_utils.py
robokq's picture
Upload 119 files
2433307 verified
Raw
History Blame Contribute Delete
3.83 kB
from __future__ import annotations
from typing import Any, List, Tuple
import numpy as np
import plotly.graph_objects as go
MAX_RENDER_POINTS = 30000
GT_COLORSCALE = [
[0.0, "rgb(43, 111, 214)"],
[0.55, "rgb(91, 151, 236)"],
[1.0, "rgb(176, 210, 255)"],
]
INF_COLORSCALE = [
[0.0, "rgb(204, 61, 46)"],
[0.55, "rgb(244, 103, 82)"],
[1.0, "rgb(255, 182, 155)"],
]
TWO_COLUMN_SLIDER_CSS = """
.slider-two-column-row {
display: grid !important;
grid-template-columns: minmax(0, 1fr) minmax(0, 1fr) !important;
gap: 16px !important;
align-items: start !important;
}
.slider-two-column-row > div {
min-width: 0 !important;
width: 100% !important;
}
@media (max-width: 900px) {
.slider-two-column-row {
grid-template-columns: 1fr !important;
}
}
"""
def prepare_points_for_plot(points: np.ndarray, max_points: int = MAX_RENDER_POINTS) -> np.ndarray:
if points.shape[0] <= max_points:
return points
indices = np.linspace(0, points.shape[0] - 1, max_points, dtype=np.int64)
return points[indices]
def add_point_cloud_trace(
fig: go.Figure,
points: np.ndarray,
name: str,
colorscale: List[List[Any]],
) -> None:
if points.size == 0:
return
render_points = prepare_points_for_plot(points)
depth_color = render_points[:, 2]
fig.add_trace(
go.Scatter3d(
x=render_points[:, 0],
y=render_points[:, 1],
z=render_points[:, 2],
mode="markers",
marker=dict(
size=2.2,
color=depth_color,
colorscale=colorscale,
opacity=0.92,
line=dict(width=0),
),
name=name,
hoverinfo="skip",
)
)
def make_point_cloud_figures(
points_gt: np.ndarray,
points_inferred: np.ndarray,
gt_name: str,
inf_name: str,
) -> Tuple[go.Figure, go.Figure]:
fig_gt = go.Figure()
fig_inf = go.Figure()
add_point_cloud_trace(fig_gt, points_gt, gt_name, GT_COLORSCALE)
add_point_cloud_trace(fig_inf, points_inferred, inf_name, INF_COLORSCALE)
all_points = (
np.vstack([points_gt, points_inferred])
if points_inferred.size > 0 and points_gt.size > 0
else (points_gt if points_gt.size > 0 else points_inferred)
)
if all_points.size > 0:
min_b, max_b = all_points.min(axis=0), all_points.max(axis=0)
center = (max_b + min_b) / 2.0
max_dim = (max_b - min_b).max()
view_dim = max(max_dim * 1.35, 0.24)
half_dim = view_dim / 2.0
axis_style = dict(visible=False, showgrid=False, zeroline=False, showbackground=False)
scene_config = dict(
xaxis=dict(axis_style, range=[center[0] - half_dim, center[0] + half_dim]),
yaxis=dict(axis_style, range=[center[1] - half_dim, center[1] + half_dim]),
zaxis=dict(axis_style, range=[center[2] - half_dim, center[2] + half_dim]),
aspectmode="cube",
)
else:
axis_style = dict(visible=False, showgrid=False, zeroline=False, showbackground=False)
scene_config = dict(
xaxis=dict(axis_style, range=[-0.2, 0.2]),
yaxis=dict(axis_style, range=[-0.2, 0.2]),
zaxis=dict(axis_style, range=[-0.2, 0.2]),
aspectmode="cube",
)
common_layout = dict(
height=520,
margin=dict(l=0, r=0, b=0, t=0),
paper_bgcolor="white",
plot_bgcolor="white",
showlegend=False,
scene_camera=dict(
eye=dict(x=1.25, y=-1.75, z=0.9),
projection=dict(type="orthographic"),
),
scene=scene_config,
)
fig_gt.update_layout(**common_layout)
fig_inf.update_layout(**common_layout)
return fig_gt, fig_inf