""" Segment a single raw point-cloud scan (.las/.laz/.ply) with the bundled PT-v3m1 model, using the pointcept/ package shipped alongside this script. Large scans are split into overlapping tiles (points don't fit the model / GPU memory in one shot), each tile is run through the real Pointcept inference pipeline (pointcept.datasets.DefaultDataset -> GridSample fragments -> TTA -> model forward, same machinery pointcept/engines/test.py uses), and per-point predictions from overlapping tiles are merged back by majority vote. Usage: python segment_scan.py --input scan.las --output scan_segmented.las python segment_scan.py --input scan.ply --output scan_segmented.ply See README.md for setup / class list. """ import argparse import shutil import tempfile from collections import OrderedDict from pathlib import Path import laspy import numpy as np import open3d as o3d import torch import torch.nn.functional as F from addict import Dict as AttrDict from pointcept.datasets.defaults import DefaultDataset from pointcept.datasets.utils import collate_fn from pointcept.models import build_model from pointcept.utils.config import Config # --------------------------------------------------------------------------- # Preprocessing recipe the model was evaluated with (extracted from the # original training config's `data.test` block; kept here rather than in # configs/model_config.py since it's inference plumbing, not architecture). # --------------------------------------------------------------------------- PREPROCESS_TRANSFORM = [ dict(type="CenterShift", apply_z=True), dict(type="NormalizeColor"), ] TEST_CFG = dict( voxelize=dict( type="GridSample", grid_size=0.04, hash_type="fnv", mode="test", return_grid_coord=True, ), crop=None, post_transform=[ dict(type="CenterShift", apply_z=False), dict(type="ToTensor"), dict( type="Collect", keys=("coord", "grid_coord", "index"), feat_keys=("color", "normal"), ), ], # test-time augmentation: 5 scales x {no-flip, flip} = 10 forward passes # per tile, softmax-averaged. Matches how the checkpoint was evaluated. aug_transform=[ [dict(type="RandomScale", scale=[s, s])] for s in (0.9, 0.95, 1.0, 1.05, 1.1) ] + [ [dict(type="RandomScale", scale=[s, s]), dict(type="RandomFlip", p=1)] for s in (0.9, 0.95, 1.0, 1.05, 1.1) ], ) # --------------------------------------------------------------------------- # Point-cloud IO # --------------------------------------------------------------------------- def read_scan(path: Path): """Read a .las/.laz/.ply file, return (coord[float64,N,3], color[float32,N,3] in 0-255).""" if path.suffix.lower() in (".las", ".laz"): return _read_las(path) elif path.suffix.lower() == ".ply": return _read_ply(path) raise ValueError(f"Unsupported input format {path.suffix!r} (expected .las/.laz/.ply)") def _read_las(path: Path): las = laspy.read(str(path)) coord = np.vstack((las.x, las.y, las.z)).transpose().astype(np.float64) rgb = np.vstack((las.red, las.green, las.blue)).transpose().astype(np.float64) # LAS commonly stores color as 16-bit even when source precision is 8-bit; # NormalizeColor (below) expects 0-255 range, matching Pointcept convention. if rgb.max() > 255: rgb = rgb / 257.0 # 65535 / 255 ~= 257 color = rgb.astype(np.float32) return coord, color def _read_ply(path: Path): pcd = o3d.io.read_point_cloud(str(path)) coord = np.asarray(pcd.points, dtype=np.float64) if pcd.has_colors(): # open3d normalizes PLY colors to [0, 1] on read regardless of the # source's bit depth; NormalizeColor (below) expects 0-255. color = (np.asarray(pcd.colors, dtype=np.float64) * 255.0).astype(np.float32) else: color = np.zeros_like(coord, dtype=np.float32) return coord, color def estimate_normals(coord: np.ndarray) -> np.ndarray: pcd = o3d.geometry.PointCloud() pcd.points = o3d.utility.Vector3dVector(coord) pcd.estimate_normals(fast_normal_computation=False) return np.asarray(pcd.normals, dtype=np.float32) def write_scan(path: Path, coord: np.ndarray, color: np.ndarray, labels: np.ndarray): if path.suffix.lower() in (".las", ".laz"): _write_las(path, coord, color, labels) elif path.suffix.lower() == ".ply": _write_ply(path, coord, color, labels) else: raise ValueError(f"Unsupported output format {path.suffix!r} (expected .las/.laz/.ply)") np.save(str(path.with_name(path.stem + "_labels.npy")), labels) def _write_las(path: Path, coord: np.ndarray, color: np.ndarray, labels: np.ndarray): header = laspy.LasHeader(point_format=3, version="1.2") header.offsets = np.min(coord, axis=0) header.scales = np.array([0.001, 0.001, 0.001]) las = laspy.LasData(header) las.x, las.y, las.z = coord[:, 0], coord[:, 1], coord[:, 2] las.red = np.clip(color[:, 0], 0, 255).astype("uint16") las.green = np.clip(color[:, 1], 0, 255).astype("uint16") las.blue = np.clip(color[:, 2], 0, 255).astype("uint16") las.classification = labels.astype("uint8") las.write(str(path)) def _write_ply(path: Path, coord: np.ndarray, color: np.ndarray, labels: np.ndarray): # Legacy open3d.geometry.PointCloud has no slot for a custom per-point # scalar field, so the label is written via the tensor-based t.geometry # API instead, as a "classification" field (readable in CloudCompare / # MeshLab alongside "positions"/"colors"). pcd = o3d.t.geometry.PointCloud() pcd.point.positions = o3d.core.Tensor(coord.astype(np.float32)) pcd.point.colors = o3d.core.Tensor(np.clip(color, 0, 255).astype(np.uint8)) pcd.point.classification = o3d.core.Tensor(labels.astype(np.uint8).reshape(-1, 1)) o3d.t.io.write_point_cloud(str(path), pcd) # --------------------------------------------------------------------------- # Tiling: split one large scan into overlapping axis-aligned tiles so each # tile fits in memory/GPU. Ports the tiling primitives used elsewhere in # this codebase for large-point-cloud inference, generalized to operate # directly on a point array instead of a named dataset's preprocessing step. # --------------------------------------------------------------------------- def make_tile_boxes(coord: np.ndarray, tile_size_m: float, overlap_pct: float): mins = coord.min(axis=0) maxs = coord.max(axis=0) x_range, y_range = maxs[0] - mins[0], maxs[1] - mins[1] n = max(1, int(max(x_range, y_range) / tile_size_m) + 1) box_width = x_range / (n - overlap_pct / 100 * (n - 1)) if n > 1 else x_range box_height = y_range / (n - overlap_pct / 100 * (n - 1)) if n > 1 else y_range overlap_w = box_width * overlap_pct / 100 overlap_h = box_height * overlap_pct / 100 boxes = [] for i in range(n): for j in range(n): x0 = mins[0] + i * (box_width - overlap_w) y0 = mins[1] + j * (box_height - overlap_h) boxes.append((x0, x0 + box_width, y0, y0 + box_height)) return boxes def points_in_box(coord: np.ndarray, box) -> np.ndarray: x0, x1, y0, y1 = box mask = ( (coord[:, 0] >= x0) & (coord[:, 0] <= x1) & (coord[:, 1] >= y0) & (coord[:, 1] <= y1) ) return np.where(mask)[0] # --------------------------------------------------------------------------- # Model # --------------------------------------------------------------------------- def load_model(config_path: Path, weight_path: Path, device: str): cfg = Config.fromfile(str(config_path)) model = build_model(cfg.model).to(device) # weights_only=False: this checkpoint predates PyTorch 2.6's stricter # default and needs the full unpickler. Only do this for checkpoints # from a source you trust (as this one is -- it's our own training run). checkpoint = torch.load(str(weight_path), map_location=device, weights_only=False) weight = OrderedDict() for key, value in checkpoint["state_dict"].items(): weight[key[7:] if key.startswith("module.") else key] = value model.load_state_dict(weight, strict=True) model.eval() return model, cfg @torch.no_grad() def infer_tile(model, dataset: DefaultDataset, idx: int, num_classes: int, device: str): """Run one tiled DefaultDataset item (with TTA fragments) through the model. Mirrors pointcept.engines.test.SemSegTester's per-fragment loop. """ item = dataset[idx] fragment_list = item["fragment_list"] n_points = item["segment"].shape[0] pred = torch.zeros((n_points, num_classes), device=device) for fragment in fragment_list: batch = collate_fn([fragment]) for key in batch: if isinstance(batch[key], torch.Tensor): batch[key] = batch[key].to(device, non_blocking=True) logits = model(batch)["seg_logits"] probs = F.softmax(logits, dim=-1) pred[batch["index"]] += probs return pred.argmax(dim=1).cpu().numpy() # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input", required=True, type=Path, help="Input .las/.laz/.ply scan") parser.add_argument( "--output", required=True, type=Path, help="Output labeled .las/.laz/.ply" ) parser.add_argument( "--config", type=Path, default=Path(__file__).parent / "configs/model_config.py" ) parser.add_argument( "--weight", type=Path, default=Path(__file__).parent / "weights/model_best.pth" ) parser.add_argument("--tile-size", type=float, default=20.0, help="Tile size in meters") parser.add_argument("--overlap", type=float, default=25.0, help="Tile overlap, percent") parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") parser.add_argument( "--work-dir", type=Path, default=None, help="Temp dir for tiles (default: auto, deleted after run)" ) parser.add_argument( "--keep-work-dir", action="store_true", help="Don't delete the tile working directory" ) args = parser.parse_args() print(f"[1/5] Loading model from {args.weight} (device={args.device}) ...") model, cfg = load_model(args.config, args.weight, args.device) num_classes = cfg.data.num_classes class_names = cfg.data.names print(f"[2/5] Reading {args.input} ...") coord, color = read_scan(args.input) n_points = coord.shape[0] print(f" {n_points:,} points") print(" estimating normals ...") normal = estimate_normals(coord) print(f"[3/5] Tiling ({args.tile_size} m tiles, {args.overlap}% overlap) ...") boxes = make_tile_boxes(coord, args.tile_size, args.overlap) work_dir = args.work_dir or Path(tempfile.mkdtemp(prefix="segment_scan_")) tiles_dir = work_dir / "tiles" tiles_dir.mkdir(parents=True, exist_ok=True) ids_by_tile_name = {} n_tiles = 0 for box in boxes: idx = points_in_box(coord, box) if idx.size == 0: continue tile_name = f"tile_{n_tiles:05d}" tile_dir = tiles_dir / tile_name tile_dir.mkdir(parents=True, exist_ok=True) np.save(tile_dir / "coord.npy", coord[idx].astype(np.float32)) np.save(tile_dir / "color.npy", color[idx]) np.save(tile_dir / "normal.npy", normal[idx]) ids_by_tile_name[tile_name] = idx n_tiles += 1 print(f" {n_tiles} non-empty tiles (out of {len(boxes)} grid cells)") print("[4/5] Running inference per tile ...") dataset = DefaultDataset( split="", data_root=str(tiles_dir), transform=PREPROCESS_TRANSFORM, test_mode=True, test_cfg=AttrDict(TEST_CFG), ) votes = torch.zeros((n_points, num_classes), dtype=torch.int16) for i in range(len(dataset)): print(f" tile {i + 1}/{len(dataset)}", end="\r") # DefaultDataset.get_data_list() uses glob, whose order isn't # guaranteed to match tile-creation order -- look ids up by name. tile_name = Path(dataset.data_list[i]).name tile_pred = infer_tile(model, dataset, i, num_classes, args.device) one_hot = F.one_hot(torch.from_numpy(tile_pred), num_classes=num_classes) votes[ids_by_tile_name[tile_name]] += one_hot.to(torch.int16) print() if not args.keep_work_dir and args.work_dir is None: shutil.rmtree(work_dir, ignore_errors=True) print("[5/5] Merging and writing output ...") covered = votes.sum(dim=1) > 0 labels_full = votes.argmax(dim=1).numpy() labels_full[~covered.numpy()] = 255 # sentinel for "never covered by a tile" out_coord = coord[covered.numpy()] out_color = color[covered.numpy()] out_labels = labels_full[covered.numpy()] args.output.parent.mkdir(parents=True, exist_ok=True) write_scan(args.output, out_coord, out_color, out_labels) dropped = n_points - out_coord.shape[0] print(f" wrote {out_coord.shape[0]:,} labeled points to {args.output}") if dropped: print(f" ({dropped:,} points were not covered by any tile and were dropped)") counts = np.bincount(out_labels, minlength=num_classes) print("\nLabel histogram:") for name, count in zip(class_names, counts): print(f" {name:>10s}: {count:,}") if __name__ == "__main__": main()