"""Step 5: validate the trained TSPulse model capability. Uses the held-out evaluation split (5%) to verify the model produces a useful fingerprint without having seen any part label during training: 1. Training/validation loss convergence (from checkpoint history). 2. Intra-part vs inter-part fingerprint distances (should be clearly < 1). 3. Evaluation reconstruction MSE. 4. Fingerprint health: 128-dim vectors finite, no collapsed dimensions. Writes a human-readable report to baseline/validation_report.txt. """ import argparse from collections import defaultdict from pathlib import Path import numpy as np import torch from torch.utils.data import DataLoader from tqdm import tqdm from tspulse import CycleDataset, TSPulse, build_split_datasets, describe, get_device, make_random_mask ROOT = Path(__file__).resolve().parent CHECKPOINT_PATH = ROOT / "checkpoints" / "tspulse_frozen.pt" BASELINE_DIR = ROOT / "baseline" REPORT_PATH = BASELINE_DIR / "validation_report.txt" def masked_mse(recon: torch.Tensor, target: torch.Tensor, mask: torch.BoolTensor, patch_size: int) -> torch.Tensor: sample_mask = mask.repeat_interleave(patch_size, dim=1) diff = (recon - target) ** 2 masked_diff = diff[sample_mask] if masked_diff.numel() == 0: return torch.tensor(0.0, device=recon.device) return masked_diff.mean() def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--mask-ratio", type=float, default=0.15) args = parser.parse_args() if not CHECKPOINT_PATH.exists(): print(f"未找到模型 {CHECKPOINT_PATH}, 请先运行 tspulse/train.py") return 1 device = get_device() print(f"设备: {describe(device)}") checkpoint = torch.load(CHECKPOINT_PATH, map_location="cpu", weights_only=False) config = checkpoint["config"] model = TSPulse(dim=config["dim"], depth=config["depth"], heads=config["heads"]).to(device) model.load_state_dict(checkpoint["state_dict"]) model.eval() samples, mean, std = build_split_datasets() eval_samples = [item for item in samples if item["split"] == "eval"] if not eval_samples: print("没有评估集样本") return 1 lines: list[str] = [] lines.append("TSPulse 模型能力验证报告") lines.append("=" * 40) lines.append(f"设备: {describe(device)}") lines.append(f"评估集样本数: {len(eval_samples)}") # 1. Loss convergence history = checkpoint.get("history", []) lines.append("\n[1] 损失收敛曲线 (epoch: train / val)") for record in history: lines.append( f" epoch {record['epoch']:>3}: train {record['train_loss']:.5f} / val {record['val_loss']:.5f}" ) # Encode all eval samples eval_set = CycleDataset(eval_samples, mean, std) loader = DataLoader(eval_set, batch_size=args.batch_size, shuffle=False) vectors: list[np.ndarray] = [] recon_losses: list[float] = [] total_rows = 0 with torch.no_grad(): for batch, _lengths in tqdm(loader, desc="编码评估集"): batch = batch.to(device) mask = make_random_mask(batch.shape[0], model.num_patches, args.mask_ratio, device) fingerprint, recon = model(batch, mask) vectors.append(fingerprint.cpu().numpy()) recon_losses.append(masked_mse(recon, batch, mask, model.patch_size).item() * batch.shape[0]) total_rows += batch.shape[0] ordered = sorted(eval_samples, key=lambda item: item["path"]) all_vectors = np.concatenate(vectors, axis=0) parts_of = [item["part"] for item in ordered] # 3. Reconstruction MSE recon_mse = sum(recon_losses) / max(total_rows, 1) lines.append(f"\n[3] 评估集重构 MSE: {recon_mse:.5f}") # 4. Fingerprint health finite = bool(np.isfinite(all_vectors).all()) std_per_dim = all_vectors.std(axis=0) collapsed = int((std_per_dim < 1e-6).sum()) lines.append("\n[4] 指纹健康检查") lines.append(f" 向量形状: {all_vectors.shape}") lines.append(f" 全部有限值: {finite}") lines.append(f" 128 维中退化维数(std<1e-6): {collapsed}") # 2. Intra vs inter part distance by_part: dict[str, list[np.ndarray]] = defaultdict(list) for vector, part in zip(all_vectors, parts_of): by_part[part].append(vector) centroids = {} intra_dists: list[float] = [] for part, vecs in by_part.items(): arr = np.stack(vecs) centroid = arr.mean(axis=0) centroids[part] = centroid if len(vecs) > 1: dists = np.linalg.norm(arr - centroid, axis=1) intra_dists.append(float(dists.mean())) inter_dists: list[float] = [] part_keys = sorted(centroids) for i in range(len(part_keys)): for j in range(i + 1, len(part_keys)): inter_dists.append(float(np.linalg.norm(centroids[part_keys[i]] - centroids[part_keys[j]]))) intra_mean = float(np.mean(intra_dists)) if intra_dists else float("nan") inter_mean = float(np.mean(inter_dists)) if inter_dists else float("nan") ratio = intra_mean / inter_mean if inter_mean else float("nan") lines.append("\n[2] 指纹区分度") lines.append(f" 平均同部位距离: {intra_mean:.6f}") lines.append(f" 平均跨部位距离: {inter_mean:.6f}") lines.append(f" 同部位/跨部位比值: {ratio:.4f} (应明显 < 1)") ok = finite and collapsed == 0 and ratio < 1.0 lines.append("\n结论: " + ("通过" if ok else "未通过")) BASELINE_DIR.mkdir(parents=True, exist_ok=True) REPORT_PATH.write_text("\n".join(lines), encoding="utf-8") print("\n".join(lines)) print(f"\n报告已保存: {REPORT_PATH}") return 0 if ok else 1 if __name__ == "__main__": raise SystemExit(main())