From faf5be476037d3baf4da3515789906a176f9b040 Mon Sep 17 00:00:00 2001
From: shuishen <1109946754@qq.com>
Date: Tue, 25 Aug 2026 12:23:59 +0800
Subject: [PATCH] feat(pointcloud): auto-select verified GPU runtime

---
 capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py |   49 ++++++++++++++++++++++++++++++++++---------------
 1 files changed, 34 insertions(+), 15 deletions(-)

diff --git a/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py b/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
index 05ef409..21a3b8f 100644
--- a/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
+++ b/capabilities/05-3d-pointcloud/train_pointcloud_semantic_model.py
@@ -7,6 +7,7 @@
 from __future__ import annotations
 
 import argparse
+import copy
 import hashlib
 import json
 import time
@@ -97,6 +98,7 @@
     parser.add_argument("--device", choices={"auto", "cpu", "cuda"}, default="auto")
     parser.add_argument("--epochs", type=int, default=40)
     parser.add_argument("--batch-size", type=int, default=4096)
+    parser.add_argument("--seed", type=int, default=42)
     args = parser.parse_args()
     if args.output.exists() and any(args.output.iterdir()):
         raise SystemExit("Output directory must be new or empty.")
@@ -135,15 +137,32 @@
     if args.device == "cuda" and not torch.cuda.is_available():
         raise SystemExit("CUDA was requested but is unavailable.")
     device = torch.device(device_name)
-    torch.manual_seed(42)
+    torch.manual_seed(args.seed)
+    if device_name == "cuda":
+        torch.cuda.manual_seed_all(args.seed)
+        torch.backends.cudnn.benchmark = False
+        torch.backends.cudnn.deterministic = True
     model = PointWiseNet(len(class_codes)).to(device)
     optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=1e-4)
     weight = torch.tensor([len(train_idx) / max(1, sum(label_by_index[int(i)] == group for i in train_idx)) for group in range(len(class_codes))], dtype=torch.float32, device=device)
     criterion = torch.nn.CrossEntropyLoss(weight=weight)
-    train_labels = np.asarray([label_by_index[int(index)] for index in train_idx], dtype=np.int64)
+    validation_true = np.asarray([target[np.searchsorted(indices, index)] for index in validation_idx], dtype=np.int64)
+    test_true = np.asarray([target[np.searchsorted(indices, index)] for index in test_idx], dtype=np.int64)
+
+    def predict(indices_to_predict: np.ndarray) -> np.ndarray:
+        model.eval(); parts: list[np.ndarray] = []
+        with torch.no_grad():
+            for start in range(0, len(indices_to_predict), args.batch_size):
+                logits = model(torch.from_numpy(features[indices_to_predict[start:start + args.batch_size]]).to(device))
+                parts.append(logits.argmax(dim=1).cpu().numpy())
+        return np.concatenate(parts)
+
     started = time.perf_counter()
-    for _ in range(args.epochs):
-        order = np.random.default_rng(42).permutation(len(train_idx))
+    best_epoch = 0
+    best_validation_macro_f1 = -1.0
+    best_state: dict[str, torch.Tensor] | None = None
+    for epoch in range(1, args.epochs + 1):
+        order = np.random.default_rng(args.seed).permutation(len(train_idx))
         model.train()
         for start in range(0, len(order), args.batch_size):
             subset = train_idx[order[start:start + args.batch_size]]
@@ -152,15 +171,15 @@
             optimizer.zero_grad(set_to_none=True)
             criterion(model(x), y).backward()
             optimizer.step()
-    def predict(indices_to_predict: np.ndarray) -> np.ndarray:
-        model.eval(); parts: list[np.ndarray] = []
-        with torch.no_grad():
-            for start in range(0, len(indices_to_predict), args.batch_size):
-                logits = model(torch.from_numpy(features[indices_to_predict[start:start + args.batch_size]]).to(device))
-                parts.append(logits.argmax(dim=1).cpu().numpy())
-        return np.concatenate(parts)
-    validation_true = np.asarray([target[np.searchsorted(indices, index)] for index in validation_idx], dtype=np.int64)
-    test_true = np.asarray([target[np.searchsorted(indices, index)] for index in test_idx], dtype=np.int64)
+        validation_metrics = metrics(validation_true, np.asarray([class_codes[value] for value in predict(validation_idx)], dtype=np.int64), class_codes)
+        macro_f1 = float(validation_metrics["report"]["macro avg"]["f1-score"])
+        if macro_f1 > best_validation_macro_f1:
+            best_epoch = epoch
+            best_validation_macro_f1 = macro_f1
+            best_state = copy.deepcopy(model.state_dict())
+    if best_state is None:
+        raise RuntimeError("Training did not produce a validation checkpoint.")
+    model.load_state_dict(best_state)
     validation_pred = np.asarray([class_codes[value] for value in predict(validation_idx)], dtype=np.int64)
     test_pred = np.asarray([class_codes[value] for value in predict(test_idx)], dtype=np.int64)
     all_pred = np.asarray([class_codes[value] for value in predict(np.arange(len(features), dtype=np.int64))], dtype=np.uint8)
@@ -177,9 +196,9 @@
         "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),
-        "normalizer": normalizer, "epochs": args.epochs, "batch_size": args.batch_size,
+        "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.", "Spatial blocks reduce leakage but one small source cannot establish field-wide generalization.", "Rule candidate colours were not used as labels or input features."],
+        "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")
     (args.output / "metrics.json").write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")

--
Gitblit v1.9.3