build_baseline.py 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115
  1. """Step 2: build the per-part baseline table.
  2. Loads the frozen TSPulse checkpoint, encodes every normal cycle sample, and
  3. for each part computes the centroid (mean of its 128-dim fingerprints) and the
  4. radius (max distance from the centroid). Results are saved as a single npz.
  5. Output baseline/baseline.npz:
  6. part_names : (P,) string array of part names
  7. centroids : (P,128) float32 centroid vectors
  8. radii : (P,) float32 max-distance radii
  9. counts : (P,) number of samples used per part
  10. """
  11. import argparse
  12. import csv
  13. from pathlib import Path
  14. import numpy as np
  15. import torch
  16. from torch.utils.data import DataLoader
  17. from tqdm import tqdm
  18. from tspulse import CycleDataset, TSPulse, build_split_datasets, describe, get_device
  19. ROOT = Path(__file__).resolve().parent
  20. CHECKPOINT_PATH = ROOT / "checkpoints" / "tspulse_frozen.pt"
  21. BASELINE_DIR = ROOT / "baseline"
  22. BASELINE_PATH = BASELINE_DIR / "baseline.npz"
  23. def load_model(checkpoint: dict, device: torch.device) -> TSPulse:
  24. config = checkpoint["config"]
  25. model = TSPulse(dim=config["dim"], depth=config["depth"], heads=config["heads"]).to(device)
  26. model.load_state_dict(checkpoint["state_dict"])
  27. model.eval()
  28. return model
  29. def main() -> int:
  30. parser = argparse.ArgumentParser()
  31. parser.add_argument("--batch-size", type=int, default=128)
  32. args = parser.parse_args()
  33. if not CHECKPOINT_PATH.exists():
  34. print(f"未找到模型 {CHECKPOINT_PATH}, 请先运行 tspulse/train.py")
  35. return 1
  36. device = get_device()
  37. print(f"设备: {describe(device)}")
  38. checkpoint = torch.load(CHECKPOINT_PATH, map_location="cpu", weights_only=False)
  39. mean = checkpoint["mean"]
  40. std = checkpoint["std"]
  41. model = load_model(checkpoint, device)
  42. samples, _, _ = build_split_datasets()
  43. parts = sorted({item["part"] for item in samples})
  44. part_index = {part: i for i, part in enumerate(parts)}
  45. centroids = np.zeros((len(parts), model.dim), dtype=np.float32)
  46. radii = np.zeros(len(parts), dtype=np.float32)
  47. counts = np.zeros(len(parts), dtype=np.int64)
  48. all_set = CycleDataset(samples, mean, std)
  49. loader = DataLoader(all_set, batch_size=args.batch_size, shuffle=False)
  50. fingerprints: list[np.ndarray] = []
  51. part_of: list[str] = []
  52. with torch.no_grad():
  53. for batch, _lengths in tqdm(loader, desc="编码所有样本"):
  54. batch = batch.to(device)
  55. vectors, _ = model(batch)
  56. fingerprints.append(vectors.cpu().numpy())
  57. # DataLoader without shuffle preserves order for non-drop_last batches.
  58. ordered = sorted(samples, key=lambda item: item["path"])
  59. part_of = [item["part"] for item in ordered]
  60. all_vectors = np.concatenate(fingerprints, axis=0)
  61. by_part: dict[str, list[np.ndarray]] = {part: [] for part in parts}
  62. for vector, part in zip(all_vectors, part_of):
  63. by_part[part].append(vector)
  64. for part in tqdm(parts, desc="计算质心与半径"):
  65. vectors = np.stack(by_part[part])
  66. centroid = vectors.mean(axis=0)
  67. distances = np.linalg.norm(vectors - centroid, axis=1)
  68. index = part_index[part]
  69. centroids[index] = centroid
  70. radii[index] = float(distances.max())
  71. counts[index] = len(vectors)
  72. BASELINE_DIR.mkdir(parents=True, exist_ok=True)
  73. np.savez(
  74. BASELINE_PATH,
  75. part_names=np.asarray(parts),
  76. centroids=centroids,
  77. radii=radii,
  78. counts=counts,
  79. )
  80. print(f"基准表已保存: {BASELINE_PATH}")
  81. with (BASELINE_DIR / "baseline_summary.csv").open("w", encoding="utf-8", newline="") as handle:
  82. writer = csv.writer(handle)
  83. writer.writerow(["part", "count", "radius"])
  84. for part, count, radius in zip(parts, counts, radii):
  85. writer.writerow([part, int(count), f"{radius:.6f}"])
  86. print("\n部位 / 样本数 / 半径:")
  87. for part, count, radius in zip(parts, counts, radii):
  88. print(f" {part}: {count} 组, 半径 {radius:.4f}")
  89. return 0
  90. if __name__ == "__main__":
  91. raise SystemExit(main())