validate_model.py 5.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152
  1. """Step 5: validate the trained TSPulse model capability.
  2. Uses the held-out evaluation split (5%) to verify the model produces a useful
  3. fingerprint without having seen any part label during training:
  4. 1. Training/validation loss convergence (from checkpoint history).
  5. 2. Intra-part vs inter-part fingerprint distances (should be clearly < 1).
  6. 3. Evaluation reconstruction MSE.
  7. 4. Fingerprint health: 128-dim vectors finite, no collapsed dimensions.
  8. Writes a human-readable report to baseline/validation_report.txt.
  9. """
  10. import argparse
  11. from collections import defaultdict
  12. from pathlib import Path
  13. import numpy as np
  14. import torch
  15. from torch.utils.data import DataLoader
  16. from tqdm import tqdm
  17. from tspulse import CycleDataset, TSPulse, build_split_datasets, describe, get_device, make_random_mask
  18. ROOT = Path(__file__).resolve().parent
  19. CHECKPOINT_PATH = ROOT / "checkpoints" / "tspulse_frozen.pt"
  20. BASELINE_DIR = ROOT / "baseline"
  21. REPORT_PATH = BASELINE_DIR / "validation_report.txt"
  22. def masked_mse(recon: torch.Tensor, target: torch.Tensor, mask: torch.BoolTensor, patch_size: int) -> torch.Tensor:
  23. sample_mask = mask.repeat_interleave(patch_size, dim=1)
  24. diff = (recon - target) ** 2
  25. masked_diff = diff[sample_mask]
  26. if masked_diff.numel() == 0:
  27. return torch.tensor(0.0, device=recon.device)
  28. return masked_diff.mean()
  29. def main() -> int:
  30. parser = argparse.ArgumentParser()
  31. parser.add_argument("--batch-size", type=int, default=128)
  32. parser.add_argument("--mask-ratio", type=float, default=0.15)
  33. args = parser.parse_args()
  34. if not CHECKPOINT_PATH.exists():
  35. print(f"未找到模型 {CHECKPOINT_PATH}, 请先运行 tspulse/train.py")
  36. return 1
  37. device = get_device()
  38. print(f"设备: {describe(device)}")
  39. checkpoint = torch.load(CHECKPOINT_PATH, map_location="cpu", weights_only=False)
  40. config = checkpoint["config"]
  41. model = TSPulse(dim=config["dim"], depth=config["depth"], heads=config["heads"]).to(device)
  42. model.load_state_dict(checkpoint["state_dict"])
  43. model.eval()
  44. samples, mean, std = build_split_datasets()
  45. eval_samples = [item for item in samples if item["split"] == "eval"]
  46. if not eval_samples:
  47. print("没有评估集样本")
  48. return 1
  49. lines: list[str] = []
  50. lines.append("TSPulse 模型能力验证报告")
  51. lines.append("=" * 40)
  52. lines.append(f"设备: {describe(device)}")
  53. lines.append(f"评估集样本数: {len(eval_samples)}")
  54. # 1. Loss convergence
  55. history = checkpoint.get("history", [])
  56. lines.append("\n[1] 损失收敛曲线 (epoch: train / val)")
  57. for record in history:
  58. lines.append(
  59. f" epoch {record['epoch']:>3}: train {record['train_loss']:.5f} / val {record['val_loss']:.5f}"
  60. )
  61. # Encode all eval samples
  62. eval_set = CycleDataset(eval_samples, mean, std)
  63. loader = DataLoader(eval_set, batch_size=args.batch_size, shuffle=False)
  64. vectors: list[np.ndarray] = []
  65. recon_losses: list[float] = []
  66. total_rows = 0
  67. with torch.no_grad():
  68. for batch, _lengths in tqdm(loader, desc="编码评估集"):
  69. batch = batch.to(device)
  70. mask = make_random_mask(batch.shape[0], model.num_patches, args.mask_ratio, device)
  71. fingerprint, recon = model(batch, mask)
  72. vectors.append(fingerprint.cpu().numpy())
  73. recon_losses.append(masked_mse(recon, batch, mask, model.patch_size).item() * batch.shape[0])
  74. total_rows += batch.shape[0]
  75. ordered = sorted(eval_samples, key=lambda item: item["path"])
  76. all_vectors = np.concatenate(vectors, axis=0)
  77. parts_of = [item["part"] for item in ordered]
  78. # 3. Reconstruction MSE
  79. recon_mse = sum(recon_losses) / max(total_rows, 1)
  80. lines.append(f"\n[3] 评估集重构 MSE: {recon_mse:.5f}")
  81. # 4. Fingerprint health
  82. finite = bool(np.isfinite(all_vectors).all())
  83. std_per_dim = all_vectors.std(axis=0)
  84. collapsed = int((std_per_dim < 1e-6).sum())
  85. lines.append("\n[4] 指纹健康检查")
  86. lines.append(f" 向量形状: {all_vectors.shape}")
  87. lines.append(f" 全部有限值: {finite}")
  88. lines.append(f" 128 维中退化维数(std<1e-6): {collapsed}")
  89. # 2. Intra vs inter part distance
  90. by_part: dict[str, list[np.ndarray]] = defaultdict(list)
  91. for vector, part in zip(all_vectors, parts_of):
  92. by_part[part].append(vector)
  93. centroids = {}
  94. intra_dists: list[float] = []
  95. for part, vecs in by_part.items():
  96. arr = np.stack(vecs)
  97. centroid = arr.mean(axis=0)
  98. centroids[part] = centroid
  99. if len(vecs) > 1:
  100. dists = np.linalg.norm(arr - centroid, axis=1)
  101. intra_dists.append(float(dists.mean()))
  102. inter_dists: list[float] = []
  103. part_keys = sorted(centroids)
  104. for i in range(len(part_keys)):
  105. for j in range(i + 1, len(part_keys)):
  106. inter_dists.append(float(np.linalg.norm(centroids[part_keys[i]] - centroids[part_keys[j]])))
  107. intra_mean = float(np.mean(intra_dists)) if intra_dists else float("nan")
  108. inter_mean = float(np.mean(inter_dists)) if inter_dists else float("nan")
  109. ratio = intra_mean / inter_mean if inter_mean else float("nan")
  110. lines.append("\n[2] 指纹区分度")
  111. lines.append(f" 平均同部位距离: {intra_mean:.6f}")
  112. lines.append(f" 平均跨部位距离: {inter_mean:.6f}")
  113. lines.append(f" 同部位/跨部位比值: {ratio:.4f} (应明显 < 1)")
  114. ok = finite and collapsed == 0 and ratio < 1.0
  115. lines.append("\n结论: " + ("通过" if ok else "未通过"))
  116. BASELINE_DIR.mkdir(parents=True, exist_ok=True)
  117. REPORT_PATH.write_text("\n".join(lines), encoding="utf-8")
  118. print("\n".join(lines))
  119. print(f"\n报告已保存: {REPORT_PATH}")
  120. return 0 if ok else 1
  121. if __name__ == "__main__":
  122. raise SystemExit(main())