| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152 |
- """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())
|