from __future__ import annotations
|
|
import json
|
import sys
|
import tempfile
|
import unittest
|
from pathlib import Path
|
from unittest import mock
|
|
import numpy as np
|
import open3d as o3d
|
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
from train_multiview_point_transformer import ( # noqa: E402
|
MultiviewLocalPointTransformer,
|
load_fusion_features,
|
main,
|
sha256,
|
spatial_split,
|
)
|
|
|
class MultiviewPointTransformerTests(unittest.TestCase):
|
def make_fixture(self, root: Path) -> Path:
|
artifact = root / "multiview-feature-fixture"
|
artifact.mkdir()
|
# Blocks 0/6/3 deterministically map to train/validation/test.
|
xyz, labels = [], []
|
for code, offset in ((5, 0.0), (17, 0.2)):
|
for block in (0, 3, 6):
|
for sample in range(200):
|
xyz.append([block * 100.0 + 10.0 + (sample % 10) * 0.01 + offset, (sample // 10) * 0.01, float(code) / 16.0])
|
labels.append(code)
|
xyz_array = np.asarray(xyz, dtype=np.float64)
|
rgb = np.column_stack((np.asarray(labels) == 5, np.asarray(labels) == 17, np.full(len(labels), 0.3))).astype(np.float32)
|
# This unlabeled extent point fixes the eight XY blocks at 100-unit width.
|
xyz_array = np.vstack((xyz_array, [[810.0, 0.0, 0.0]]))
|
rgb = np.vstack((rgb, [[0.2, 0.2, 0.2]])).astype(np.float32)
|
cloud = o3d.geometry.PointCloud()
|
cloud.points = o3d.utility.Vector3dVector(xyz_array)
|
cloud.colors = o3d.utility.Vector3dVector(rgb)
|
source = artifact / "multiview-annotation-source.ply"
|
self.assertTrue(o3d.io.write_point_cloud(str(source), cloud, write_ascii=False))
|
feature_mean = np.column_stack((rgb, np.zeros((len(xyz_array), 5), dtype=np.float32)))
|
feature_stddev = np.full((len(xyz_array), 8), 0.05, dtype=np.float32)
|
dataset = artifact / "multiview-point-features.npz"
|
np.savez_compressed(dataset, xyz=xyz_array, las_rgb=rgb, photo_feature_mean=feature_mean, photo_feature_stddev=feature_stddev, visible_view_count=np.full(len(xyz_array), 12, dtype=np.uint16), feature_names=np.asarray(("photo_r", "photo_g", "photo_b", "hue_sin", "hue_cos", "saturation", "gradient_magnitude", "local_intensity_stddev")))
|
(artifact / "run_metadata.json").write_text(json.dumps({"annotation_source": {"kind": "multiview_photo_feature_fusion", "point_cloud": source.name, "feature_dataset": dataset.name, "point_count": len(xyz_array), "point_cloud_sha256": sha256(source), "feature_dataset_sha256": sha256(dataset)}}), encoding="utf-8")
|
annotation = root / "annotation.json"
|
annotation.write_text(json.dumps({"schema_version": 1, "source_path": str(source), "source_sha256": sha256(source), "labels": [[index, code] for index, code in enumerate(labels)], "class_schema": {"5": {"code": 5, "key": "vegetation", "label": "植被", "color": [59, 163, 87]}, "17": {"code": 17, "key": "transformer", "label": "变压器", "color": [30, 144, 255]}}}), encoding="utf-8")
|
return annotation
|
|
def test_loads_verified_features_and_local_attention_contract(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
annotation = self.make_fixture(Path(temp_dir))
|
xyz, features, contract = load_fusion_features(annotation)
|
self.assertEqual(xyz.shape, (1201, 3))
|
self.assertEqual(features.shape, (1201, 23))
|
self.assertEqual(contract["point_count"], 1201)
|
train, validation, test = spatial_split(xyz, np.arange(len(xyz)))
|
self.assertTrue(len(train) and len(validation) and len(test))
|
model = MultiviewLocalPointTransformer(features.shape[1], 2)
|
logits = model(__import__("torch").from_numpy(features[:4, None, :].repeat(4, axis=1)), __import__("torch").zeros((4, 4, 3)))
|
self.assertEqual(tuple(logits.shape), (4, 2))
|
|
def test_cli_writes_multiview_model_and_metrics(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
root = Path(temp_dir)
|
annotation = self.make_fixture(root)
|
output = root / "output"
|
with mock.patch.object(sys, "argv", ["trainer", "--annotation", str(annotation), "--output", str(output), "--device", "cpu", "--epochs", "1", "--batch-size", "256", "--neighbours", "4"]):
|
self.assertEqual(main(), 0)
|
metrics = json.loads((output / "metrics.json").read_text(encoding="utf-8"))
|
self.assertTrue((output / "model.pt").is_file())
|
self.assertTrue((output / "predicted-semantic-preview.ply").is_file())
|
self.assertEqual(metrics["model_input_kind"], "multiview_photo_feature_fusion")
|
self.assertEqual(metrics["classes"]["17"]["key"], "transformer")
|
self.assertIn("not official Point Transformer V3", metrics["model"])
|
|
def test_rejects_changed_feature_dataset(self) -> None:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
annotation = self.make_fixture(Path(temp_dir))
|
dataset = annotation.parent / "multiview-feature-fixture" / "multiview-point-features.npz"
|
dataset.write_bytes(b"changed")
|
with self.assertRaisesRegex(ValueError, "checksum changed"):
|
load_fusion_features(annotation)
|
|
|
if __name__ == "__main__":
|
unittest.main()
|