shuishen
19 hours ago 385be2eca72eb3833efa4be0a0088b34e764788a
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
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()