Newer
Older
oak-d_proto / pointcloud_preview.py
@T.Nakaguchi T.Nakaguchi 21 days ago 13 KB 最初のコミット
import sys
import threading

import depthai as dai
import numpy as np
import pyvista as pv

from pyvistaqt import BackgroundPlotter
from PySide6.QtWidgets import QApplication

from vtkmodules.util.numpy_support import vtk_to_numpy


# ============================================================
# Settings
# ============================================================

FRAME_SIZE = (640, 400)
FPS = 30

MIN_DEPTH_M = 0.3
MAX_DEPTH_M = 5.0

# 640x400 = 256000 points
#
# 2 -> 128000 points
# 4 ->  64000 points
DISPLAY_STRIDE = 2

# GUI update rate
VIEWER_INTERVAL_MS = 33


# ============================================================
# Convert DepthAI point cloud
# ============================================================

def convert_pointcloud(pcl_data, unit_scale):
    """
    Convert DepthAI point cloud to fixed-size arrays.

    Return:
        xyz  : Nx3 float32 [m]
        rgba : Nx4 uint8
        valid_count : int

    Viewer coordinates:
        X : right
        Y : forward
        Z : up
    """

    raw = np.asarray(
        pcl_data.getPoints(),
        dtype=np.float32
    )

    # --------------------------------------------------------
    # Fixed decimation
    #
    # Do NOT remove invalid points because the VTK point count
    # must remain constant between frames.
    # --------------------------------------------------------

    raw = raw[::DISPLAY_STRIDE]

    points_m = raw * unit_scale

    z = points_m[:, 2]

    valid = np.isfinite(points_m).all(axis=1)

    valid &= z >= MIN_DEPTH_M
    valid &= z <= MAX_DEPTH_M

    # --------------------------------------------------------
    # DepthAI:
    #
    # X = right
    # Y = down
    # Z = forward
    #
    # Viewer:
    #
    # X = right
    # Y = forward
    # Z = up
    # --------------------------------------------------------

    xyz = np.empty_like(
        points_m,
        dtype=np.float32
    )

    xyz[:, 0] = points_m[:, 0]
    xyz[:, 1] = points_m[:, 2]
    xyz[:, 2] = -points_m[:, 1]

    # --------------------------------------------------------
    # RGBA color
    # --------------------------------------------------------

    rgba = np.zeros(
        (len(xyz), 4),
        dtype=np.uint8
    )

    if np.any(valid):

        # Depth normalized to 0 ... 1
        t = (
            z[valid] - MIN_DEPTH_M
        ) / (
            MAX_DEPTH_M - MIN_DEPTH_M
        )

        t = np.clip(
            t,
            0.0,
            1.0
        )

        # Simple near/far color
        rgba[valid, 0] = (
            255 * (1.0 - t)
        ).astype(np.uint8)

        rgba[valid, 1] = 180

        rgba[valid, 2] = (
            255 * t
        ).astype(np.uint8)

        rgba[valid, 3] = 255

        # ----------------------------------------------------
        # Invalid points:
        #
        # Place them at an existing valid point and alpha=0.
        #
        # This avoids NaN and avoids changing the point count.
        # ----------------------------------------------------

        first_valid = np.flatnonzero(valid)[0]

        xyz[~valid] = xyz[first_valid]

    else:

        xyz[:] = 0.0

    return (
        xyz,
        rgba,
        int(np.count_nonzero(valid))
    )


# ============================================================
# DepthAI
# ============================================================

with dai.Pipeline() as pipeline:

    # --------------------------------------------------------
    # Cameras
    # --------------------------------------------------------

    color = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_A,
        sensorFps=FPS
    )

    left = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_B,
        sensorFps=FPS
    )

    right = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_C,
        sensorFps=FPS
    )

    # --------------------------------------------------------
    # StereoDepth
    # --------------------------------------------------------

    stereo = pipeline.create(
        dai.node.StereoDepth
    )

    stereo.setDefaultProfilePreset(
        dai.node.StereoDepth.PresetMode.DEFAULT
    )

    stereo.setRectifyEdgeFillColor(0)

    stereo.enableDistortionCorrection(
        True
    )

    left.requestOutput(
        FRAME_SIZE
    ).link(
        stereo.left
    )

    right.requestOutput(
        FRAME_SIZE
    ).link(
        stereo.right
    )

    # --------------------------------------------------------
    # RGBD
    # --------------------------------------------------------

    rgbd = pipeline.create(
        dai.node.RGBD
    ).build(
        color,
        stereo,
        FRAME_SIZE,
        FPS
    )

    pcl_queue = (
        rgbd.pcl.createOutputQueue(
            maxSize=2,
            blocking=False
        )
    )

    pipeline.start()

    print("DepthAI started")
    print("Waiting for first point cloud...")

    # ========================================================
    # First frame
    # ========================================================

    first_data = pcl_queue.get()

    raw = np.asarray(
        first_data.getPoints(),
        dtype=np.float32
    )

    raw_z = raw[:, 2]

    positive = (
        np.isfinite(raw_z)
        & (raw_z > 0)
    )

    median_z = np.median(
        raw_z[positive]
    )

    # --------------------------------------------------------
    # Detect XYZ unit once
    # --------------------------------------------------------

    if median_z > 100.0:

        unit_scale = 0.001
        unit_name = "millimeter"

    else:

        unit_scale = 1.0
        unit_name = "meter"

    print(
        f"Raw Z range: "
        f"{raw_z[positive].min():.3f} - "
        f"{raw_z[positive].max():.3f}"
    )

    print(
        f"Median Z: {median_z:.3f}"
    )

    print(
        f"Detected unit: {unit_name}"
    )

    first_xyz, first_rgba, valid_count = (
        convert_pointcloud(
            first_data,
            unit_scale
        )
    )

    print(
        f"Viewer points: {len(first_xyz)}"
    )

    print(
        f"Valid points: {valid_count}"
    )

    # ========================================================
    # Shared latest-frame buffer
    #
    # Camera thread -> Qt GUI thread
    # ========================================================

    latest_lock = threading.Lock()

    latest_xyz = first_xyz.copy()
    latest_rgba = first_rgba.copy()

    latest_frame_id = 0

    stop_event = threading.Event()

    # ========================================================
    # Capture thread
    #
    # IMPORTANT:
    # DepthAI acquisition is completely separated from Qt/VTK.
    # ========================================================

    def capture_loop():

        nonlocal_vars = {
            "frame_id": 0
        }

        while not stop_event.is_set():

            try:

                pcl_data = pcl_queue.get()

            except Exception:

                break

            xyz, rgba, count = (
                convert_pointcloud(
                    pcl_data,
                    unit_scale
                )
            )

            with latest_lock:

                # In-place replacement of shared references
                globals_placeholder[0] = xyz
                globals_placeholder[1] = rgba

                nonlocal_vars["frame_id"] += 1

                globals_placeholder[2] = (
                    nonlocal_vars["frame_id"]
                )

    # --------------------------------------------------------
    # Python doesn't have writable nonlocal variables at module
    # scope, therefore store current frame data in this list.
    # --------------------------------------------------------

    globals_placeholder = [
        latest_xyz,
        latest_rgba,
        latest_frame_id
    ]

    capture_thread = threading.Thread(
        target=capture_loop,
        daemon=True
    )

    capture_thread.start()

    # ========================================================
    # Qt
    # ========================================================

    app = QApplication.instance()

    if app is None:

        app = QApplication(sys.argv)

    app.setQuitOnLastWindowClosed(True)

    # ========================================================
    # PyVistaQt BackgroundPlotter
    # ========================================================

    plotter = BackgroundPlotter(
        app=app,
        show=True,
        window_size=(1280, 720),
        title="OAK-D Point Cloud"
    )

    # --------------------------------------------------------
    # Create ONE PolyData
    # --------------------------------------------------------

    cloud = pv.PolyData(
        first_xyz.copy()
    )

    cloud.point_data["rgba"] = (
        first_rgba.copy()
    )

    actor = plotter.add_points(
        cloud,
        scalars="rgba",
        rgba=True,
        style="points",
        point_size=3,
        lighting=False,
        name="oak_pointcloud",
        reset_camera=True
    )

    plotter.show_axes()
    plotter.show_grid()

    plotter.add_text(
        "OAK-D Real-time Point Cloud",
        position="upper_left",
        font_size=10
    )

    plotter.reset_camera()

    # ========================================================
    # IMPORTANT:
    #
    # Get direct NumPy views into VTK memory.
    #
    # We will update THESE arrays, rather than assigning new
    # PyVista arrays every frame.
    # ========================================================

    vtk_points_data = (
        cloud.GetPoints().GetData()
    )

    vtk_rgba_data = (
        cloud.GetPointData().GetArray(
            "rgba"
        )
    )

    vtk_points_np = vtk_to_numpy(
        vtk_points_data
    )

    vtk_rgba_np = vtk_to_numpy(
        vtk_rgba_data
    )

    print()
    print("VTK buffers")
    print(
        "points:",
        vtk_points_np.shape,
        vtk_points_np.dtype
    )

    print(
        "rgba:",
        vtk_rgba_np.shape,
        vtk_rgba_np.dtype
    )

    # ========================================================
    # Qt GUI update callback
    # ========================================================

    displayed_frame_id = [-1]

    update_counter = [0]

    def update_view():

        # Window already closed
        if plotter._closed:
            return

        # --------------------------------------------
        # Get newest frame from capture thread
        # --------------------------------------------

        with latest_lock:

            frame_id = globals_placeholder[2]

            if frame_id == displayed_frame_id[0]:
                return

            xyz = globals_placeholder[0]
            rgba = globals_placeholder[1]

            # Copy because capture thread may replace
            # references while VTK is drawing.
            xyz = xyz.copy()
            rgba = rgba.copy()

        # --------------------------------------------
        # Sanity check
        # --------------------------------------------

        if xyz.shape != vtk_points_np.shape:

            print(
                "Point array shape changed:",
                xyz.shape,
                vtk_points_np.shape
            )

            return

        # ====================================================
        # KEY POINT:
        #
        # Write DIRECTLY into the existing VTK arrays.
        # ====================================================

        vtk_points_np[:] = xyz
        vtk_rgba_np[:] = rgba

        # ----------------------------------------------------
        # Explicitly tell VTK which buffers changed
        # ----------------------------------------------------

        vtk_points_data.Modified()
        vtk_rgba_data.Modified()

        cloud.GetPoints().Modified()
        cloud.GetPointData().Modified()
        cloud.Modified()

        # Mapper update
        actor.mapper.Update()

        # Render
        plotter.render()

        displayed_frame_id[0] = frame_id

        update_counter[0] += 1

        if update_counter[0] % 60 == 0:

            print(
                f"Viewer updates: "
                f"{update_counter[0]}"
            )

    # ========================================================
    # Qt QTimer
    #
    # BackgroundPlotter.add_callback() uses QTimer.
    # ========================================================

    plotter.add_callback(
        update_view,
        interval=VIEWER_INTERVAL_MS
    )

    # ========================================================
    # Closing
    # ========================================================

    def window_closed():

        print("Viewer closing...")

        stop_event.set()

        app.quit()

    plotter.app_window.signal_close.connect(
        window_closed
    )

    print()
    print("Viewer started")
    print("-----------------------------")
    print("Left drag   : rotate")
    print("Mouse wheel : zoom")
    print("Middle drag : pan")
    print("Q / X       : close")
    print()

    # ========================================================
    # Qt event loop
    #
    # This is the ONLY GUI event loop.
    # ========================================================

    try:

        app.exec()

    except KeyboardInterrupt:

        print("Interrupted")

    finally:

        # Stop capture thread
        stop_event.set()

        # Stop DepthAI first so blocking get() is released
        pipeline.stop()

        capture_thread.join(
            timeout=1.0
        )

        print("DepthAI stopped")
        print("Finished")