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()