File size: 13,679 Bytes
7ab05dd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
"""
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()