Download segment_scan.py from dfki-av/BIMStruct3D-segmentation: direct link, hf CLI and curl.
- Browser
- Download file 13.7 kB
-
https://huggingface.co/dfki-av/BIMStruct3D-segmentation/resolve/main/segment_scan.py
- Command line
-
hf download hf://dfki-av/BIMStruct3D-segmentation/segment_scan.py
-
curl -L -o segment_scan.py https://huggingface.co/dfki-av/BIMStruct3D-segmentation/resolve/main/segment_scan.py
13.7 kB
| """ | |
| 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 | |
| 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() | |