Source code for openworld_radio_twin.plotly_preview

"""VS Code-compatible interactive preview for compiled OWRT scenes."""

from __future__ import annotations

import json
from pathlib import Path
from typing import Any

import numpy as np

from openworld_radio_twin.rt import scene_path

_GROUP_STYLE = {
    "ground": ("Ground", "#D8D0B4", 1.0),
    "vegetation": ("Vegetation", "#43A56F", 1.0),
    "paved": ("Paved", "#929BA4", 1.0),
    "water": ("Water", "#2AA8D2", 0.82),
    "walls": ("Building walls", "#AEB9C5", 1.0),
    "roofs": ("Building roofs", "#697887", 1.0),
}

__all__ = ["interactive_preview"]


def _read_owrt_ply(path: Path) -> tuple[np.ndarray, np.ndarray]:
    """Read the fixed binary triangle PLY format emitted by OWRT."""
    vertex_count = face_count = None
    with path.open("rb") as handle:
        while True:
            raw = handle.readline()
            if not raw:
                raise ValueError(f"Incomplete PLY header: {path}")
            line = raw.decode("ascii").strip()
            if line == "format binary_little_endian 1.0":
                continue
            if line.startswith("format "):
                raise ValueError(f"Unsupported PLY encoding in {path}: {line}")
            if line.startswith("element vertex "):
                vertex_count = int(line.rsplit(" ", 1)[1])
            elif line.startswith("element face "):
                face_count = int(line.rsplit(" ", 1)[1])
            elif line == "end_header":
                break
        if vertex_count is None or face_count is None:
            raise ValueError(f"PLY counts are missing: {path}")
        vertices = np.frombuffer(handle.read(vertex_count * 12), dtype="<f4").reshape(-1, 3)
        face_dtype = np.dtype([("size", "u1"), ("indices", "<i4", (3,))])
        records = np.frombuffer(handle.read(face_count * face_dtype.itemsize), dtype=face_dtype)
    if vertices.shape[0] != vertex_count or records.shape[0] != face_count:
        raise ValueError(f"Truncated PLY payload: {path}")
    if face_count and not np.all(records["size"] == 3):
        raise ValueError(f"Non-triangular face in {path}")
    return vertices.copy(), records["indices"].astype(np.int64, copy=True)


def _mesh_group(path: Path) -> str | None:
    name = path.stem
    if name == "ground":
        return "ground"
    if name == "terrain_vegetation":
        return "vegetation"
    if name == "terrain_paved":
        return "paved"
    if name == "terrain_water":
        return "water"
    if name.startswith("building_") and name.endswith("_wall"):
        return "walls"
    if name.startswith("building_") and name.endswith("_rooftop"):
        return "roofs"
    return None


def _load_mesh_groups(directory: Path) -> dict[str, tuple[np.ndarray, np.ndarray]]:
    grouped: dict[str, list[tuple[np.ndarray, np.ndarray]]] = {}
    for path in sorted((directory / "mesh").glob("*.ply")):
        group = _mesh_group(path)
        if group is None:
            continue
        vertices, faces = _read_owrt_ply(path)
        if not len(vertices) or not len(faces):
            continue
        grouped.setdefault(group, []).append((vertices, faces))

    result = {}
    for group, meshes in grouped.items():
        vertex_parts = []
        face_parts = []
        offset = 0
        for vertices, faces in meshes:
            vertex_parts.append(vertices)
            face_parts.append(faces + offset)
            offset += len(vertices)
        result[group] = (np.concatenate(vertex_parts), np.concatenate(face_parts))
    return result


def _thin_mesh(
    vertices: np.ndarray,
    faces: np.ndarray,
    limit: int,
) -> tuple[np.ndarray, np.ndarray]:
    if len(faces) <= limit:
        return vertices, faces
    chosen = np.linspace(0, len(faces) - 1, limit, dtype=np.int64)
    selected = faces[chosen]
    used, inverse = np.unique(selected, return_inverse=True)
    return vertices[used], inverse.reshape(-1, 3)


def _measurement_record(directory: Path, cell_size: tuple[float, float], height: float) -> dict:
    info = json.loads((directory / "scene_info.json").read_text(encoding="utf-8"))
    records = info.get("shared_assets", {}).get("measurement_surfaces", [])
    for record in reversed(records):
        if np.allclose(record["cell_size_m"], cell_size) and np.isclose(
            record["receiver_height_agl_m"], height
        ):
            return record
    raise ValueError("No matching OWRT measurement surface was found")


def _metric_surface(
    directory: Path,
    radio_map: Any,
    metric: str,
    cell_size: tuple[float, float],
    receiver_height_m: float,
) -> tuple[np.ndarray, np.ndarray, str]:
    record = _measurement_record(directory, cell_size, receiver_height_m)
    with np.load(directory / record["arrays"]) as arrays:
        weights = arrays["triangle_areas_m2"].astype(np.float64)
        centers = arrays["cell_centers_enu_m"].astype(np.float64)
    rows, columns = weights.shape[:2]
    raw = np.asarray(getattr(radio_map, metric).numpy(), dtype=np.float64)
    if raw.ndim == 1:
        raw = raw[None, :]
    triangles = raw.reshape(raw.shape[0], rows, columns, 2)
    linear = np.sum(triangles * weights[None, ...], axis=-1)
    linear /= np.sum(weights, axis=-1)[None, ...]
    strongest = np.max(linear, axis=0)
    offset, unit = {"rss": (30.0, "dBm"), "path_gain": (0.0, "dB"), "sinr": (0.0, "dB")}[metric]
    with np.errstate(divide="ignore", invalid="ignore"):
        values = np.where(strongest > 0, 10.0 * np.log10(strongest) + offset, np.nan)
    return centers, values, unit


def _transmitters(scene: Any) -> tuple[np.ndarray, list[str]]:
    positions = []
    names = []
    for name, transmitter in scene.transmitters.items():
        position = transmitter.position
        if hasattr(position, "numpy"):
            position = position.numpy()
        positions.append(np.asarray(position, dtype=np.float64).reshape(3))
        names.append(name)
    if not positions:
        return np.empty((0, 3)), []
    return np.stack(positions), names


[docs] def interactive_preview( scene: Any, *, radio_map: Any | None = None, metric: str = "rss", cell_size: tuple[float, float] = (10.0, 10.0), receiver_height_m: float = 1.5, title: str = "OpenWorld Radio Twin", max_mesh_faces: int = 600_000, max_surface_side: int = 420, ): """Return a Plotly 3D figure without relying on the ipywidgets renderer.""" try: import plotly.graph_objects as go except ImportError as exc: raise RuntimeError("Plotly preview requires 'openworld-radio-twin[tutorial]'") from exc directory = scene_path(scene) groups = _load_mesh_groups(directory) total_faces = sum(len(faces) for _, faces in groups.values()) traces = [] for group in _GROUP_STYLE: if group not in groups: continue vertices, faces = groups[group] allowance = len(faces) if total_faces > max_mesh_faces: allowance = max(1, round(max_mesh_faces * len(faces) / total_faces)) vertices, faces = _thin_mesh(vertices, faces, allowance) label, color, opacity = _GROUP_STYLE[group] traces.append( go.Mesh3d( x=vertices[:, 0], y=vertices[:, 1], z=vertices[:, 2], i=faces[:, 0], j=faces[:, 1], k=faces[:, 2], name=label, color=color, opacity=opacity, flatshading=True, hoverinfo="skip", showscale=False, lighting=dict(ambient=0.62, diffuse=0.72, roughness=0.88, specular=0.08), ) ) if radio_map is not None: centers, values, unit = _metric_surface( directory, radio_map, metric, cell_size, receiver_height_m ) stride = max(1, int(np.ceil(max(values.shape) / max_surface_side))) centers = centers[::stride, ::stride] values = values[::stride, ::stride] finite = values[np.isfinite(values)] if finite.size: cmin, cmax = np.percentile(finite, [2, 98]) traces.append( go.Surface( x=centers[..., 0], y=centers[..., 1], z=centers[..., 2] + 0.2, surfacecolor=values, colorscale="Viridis", cmin=cmin, cmax=cmax, opacity=0.78, name=metric.upper(), showscale=True, colorbar=dict(title=f"{metric.upper()} ({unit})", len=0.62, thickness=14), hovertemplate=( "East %{x:.1f} m<br>North %{y:.1f} m<br>" + metric.upper() + " %{surfacecolor:.2f} " + unit + "<extra></extra>" ), ) ) tx_positions, tx_names = _transmitters(scene) if len(tx_positions): traces.append( go.Scatter3d( x=tx_positions[:, 0], y=tx_positions[:, 1], z=tx_positions[:, 2], mode="markers+text", text=tx_names, textposition="top center", name="Transmitters", marker=dict(size=7, color="#D43F3A", symbol="diamond"), hovertemplate="%{text}<br>(%{x:.1f}, %{y:.1f}, %{z:.1f}) m<extra></extra>", ) ) figure = go.Figure(traces) figure.update_layout( title=dict(text=title, x=0.02, xanchor="left"), template="plotly_white", height=720, margin=dict(l=0, r=0, t=50, b=0), legend=dict(orientation="h", x=0.01, y=0.99, bgcolor="rgba(255,255,255,0.72)"), scene=dict( aspectmode="data", xaxis_title="East (m)", yaxis_title="North (m)", zaxis_title="Elevation (m)", camera=dict(eye=dict(x=1.35, y=-1.45, z=1.05)), ), ) return figure