From 385be2eca72eb3833efa4be0a0088b34e764788a Mon Sep 17 00:00:00 2001
From: shuishen <1109946754@qq.com>
Date: Mon, 31 Aug 2026 09:04:45 +0800
Subject: [PATCH] feat(pointcloud): complete annotation and result lifecycle workflows
---
capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py | 41 +++++++++++++++++++++++++++++++++--------
1 files changed, 33 insertions(+), 8 deletions(-)
diff --git a/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py b/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
index 21a3b8f..ae258e9 100644
--- a/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
+++ b/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
@@ -80,8 +80,32 @@
return train, validation, test
-def metrics(y_true: np.ndarray, y_pred: np.ndarray, classes: list[int]) -> dict[str, Any]:
- names = [CLASS_SCHEMA[code]["key"] for code in classes]
+def annotation_schema(annotation: dict[str, Any]) -> dict[int, dict[str, Any]]:
+ """Use the immutable revision snapshot, while accepting old six-class revisions."""
+ schema = {code: {"code": code, **value} for code, value in CLASS_SCHEMA.items()}
+ values = annotation.get("class_schema")
+ if not isinstance(values, dict):
+ return schema
+ for raw_code, raw_value in values.items():
+ try:
+ code = int(raw_code)
+ except (TypeError, ValueError):
+ raise SystemExit("Annotation class schema has an invalid class code.")
+ if not 1 <= code <= 255:
+ raise SystemExit("Annotation class schema has a non-LAS-compatible class code.")
+ value = raw_value if isinstance(raw_value, dict) else {}
+ inherited = schema.get(code, {})
+ key = value.get("key", inherited.get("key"))
+ label = value.get("label", inherited.get("label"))
+ color = value.get("color", inherited.get("color"))
+ if not isinstance(key, str) or not isinstance(label, str) or not isinstance(color, list) or len(color) != 3 or not all(isinstance(item, int) and 0 <= item <= 255 for item in color):
+ raise SystemExit("Annotation class schema is incomplete for a labelled class.")
+ schema[code] = {"code": code, "key": key, "label": label, "color": color}
+ return schema
+
+
+def metrics(y_true: np.ndarray, y_pred: np.ndarray, classes: list[int], schema: dict[int, dict[str, Any]]) -> dict[str, Any]:
+ names = [schema[code]["key"] for code in classes]
return {
"class_codes": classes,
"class_keys": names,
@@ -105,6 +129,7 @@
annotation = json.loads(args.annotation.read_text(encoding="utf-8"))
if annotation.get("schema_version") != 1:
raise SystemExit("Unsupported annotation schema.")
+ class_schema = annotation_schema(annotation)
source = Path(str(annotation.get("source_path") or ""))
if not source.is_file() or sha256(source) != annotation.get("source_sha256"):
raise SystemExit("Annotation source is unavailable or its checksum changed.")
@@ -119,7 +144,7 @@
if not isinstance(item, list) or len(item) != 2 or not all(isinstance(value, int) for value in item):
raise SystemExit("Annotation contains an invalid point label.")
index, code = item
- if not 0 <= index < len(features) or code not in CLASS_SCHEMA:
+ if not 0 <= index < len(features) or code not in class_schema:
raise SystemExit("Annotation contains an out-of-range point label.")
assigned[index] = code
indices = np.asarray(sorted(assigned), dtype=np.int64)
@@ -171,7 +196,7 @@
optimizer.zero_grad(set_to_none=True)
criterion(model(x), y).backward()
optimizer.step()
- validation_metrics = metrics(validation_true, np.asarray([class_codes[value] for value in predict(validation_idx)], dtype=np.int64), class_codes)
+ validation_metrics = metrics(validation_true, np.asarray([class_codes[value] for value in predict(validation_idx)], dtype=np.int64), class_codes, class_schema)
macro_f1 = float(validation_metrics["report"]["macro avg"]["f1-score"])
if macro_f1 > best_validation_macro_f1:
best_epoch = epoch
@@ -186,21 +211,21 @@
args.output.mkdir(parents=True, exist_ok=True)
predicted = o3d.geometry.PointCloud()
predicted.points = o3d.utility.Vector3dVector(xyz)
- predicted.colors = o3d.utility.Vector3dVector(np.asarray([CLASS_SCHEMA[int(code)]["color"] for code in all_pred], dtype=np.float64) / 255.0)
+ predicted.colors = o3d.utility.Vector3dVector(np.asarray([class_schema[int(code)]["color"] for code in all_pred], dtype=np.float64) / 255.0)
o3d.io.write_point_cloud(str(args.output / "predicted-semantic-preview.ply"), predicted, write_ascii=False)
payload = {
"capability": "05-3d-pointcloud", "classification": "B", "created_at": datetime.now(UTC).isoformat(),
"model": "PointWiseNet shared MLP (human-confirmed point labels)", "device": device_name,
"torch_version": torch.__version__, "annotation": str(args.annotation), "source": str(source),
- "source_sha256": annotation["source_sha256"], "classes": {str(code): CLASS_SCHEMA[code] for code in class_codes},
+ "source_sha256": annotation["source_sha256"], "classes": {str(code): class_schema[code] for code in class_codes},
"label_counts": {str(code): count for code, count in per_class.items()},
"split_counts": {"train": int(len(train_idx)), "validation": int(len(validation_idx)), "test": int(len(test_idx))},
- "validation": metrics(validation_true, validation_pred, class_codes), "test": metrics(test_true, test_pred, class_codes),
+ "validation": metrics(validation_true, validation_pred, class_codes, class_schema), "test": metrics(test_true, test_pred, class_codes, class_schema),
"normalizer": normalizer, "epochs": args.epochs, "best_epoch": best_epoch, "best_validation_macro_f1": best_validation_macro_f1, "batch_size": args.batch_size, "seed": args.seed,
"elapsed_seconds": round(time.perf_counter() - started, 3),
"limitations": ["Metrics cover only human-confirmed points in this annotation revision.", "The retained checkpoint is selected by validation macro F1; the test split remains separate from that selection.", "Spatial blocks reduce leakage but one small source cannot establish field-wide generalization.", "Rule candidate colours were not used as labels or input features."],
}
- torch.save({"state_dict": model.cpu().state_dict(), "class_codes": class_codes, "normalizer": normalizer, "schema_version": 1}, args.output / "model.pt")
+ torch.save({"state_dict": model.cpu().state_dict(), "class_codes": class_codes, "classes": {str(code): class_schema[code] for code in class_codes}, "normalizer": normalizer, "schema_version": 2}, args.output / "model.pt")
(args.output / "metrics.json").write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps(payload, ensure_ascii=False))
return 0
--
Gitblit v1.9.3