| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115 |
- """Step 2: build the per-part baseline table.
- Loads the frozen TSPulse checkpoint, encodes every normal cycle sample, and
- for each part computes the centroid (mean of its 128-dim fingerprints) and the
- radius (max distance from the centroid). Results are saved as a single npz.
- Output baseline/baseline.npz:
- part_names : (P,) string array of part names
- centroids : (P,128) float32 centroid vectors
- radii : (P,) float32 max-distance radii
- counts : (P,) number of samples used per part
- """
- import argparse
- import csv
- 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
- ROOT = Path(__file__).resolve().parent
- CHECKPOINT_PATH = ROOT / "checkpoints" / "tspulse_frozen.pt"
- BASELINE_DIR = ROOT / "baseline"
- BASELINE_PATH = BASELINE_DIR / "baseline.npz"
- def load_model(checkpoint: dict, device: torch.device) -> TSPulse:
- 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()
- return model
- def main() -> int:
- parser = argparse.ArgumentParser()
- parser.add_argument("--batch-size", type=int, default=128)
- 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)
- mean = checkpoint["mean"]
- std = checkpoint["std"]
- model = load_model(checkpoint, device)
- samples, _, _ = build_split_datasets()
- parts = sorted({item["part"] for item in samples})
- part_index = {part: i for i, part in enumerate(parts)}
- centroids = np.zeros((len(parts), model.dim), dtype=np.float32)
- radii = np.zeros(len(parts), dtype=np.float32)
- counts = np.zeros(len(parts), dtype=np.int64)
- all_set = CycleDataset(samples, mean, std)
- loader = DataLoader(all_set, batch_size=args.batch_size, shuffle=False)
- fingerprints: list[np.ndarray] = []
- part_of: list[str] = []
- with torch.no_grad():
- for batch, _lengths in tqdm(loader, desc="编码所有样本"):
- batch = batch.to(device)
- vectors, _ = model(batch)
- fingerprints.append(vectors.cpu().numpy())
- # DataLoader without shuffle preserves order for non-drop_last batches.
- ordered = sorted(samples, key=lambda item: item["path"])
- part_of = [item["part"] for item in ordered]
- all_vectors = np.concatenate(fingerprints, axis=0)
- by_part: dict[str, list[np.ndarray]] = {part: [] for part in parts}
- for vector, part in zip(all_vectors, part_of):
- by_part[part].append(vector)
- for part in tqdm(parts, desc="计算质心与半径"):
- vectors = np.stack(by_part[part])
- centroid = vectors.mean(axis=0)
- distances = np.linalg.norm(vectors - centroid, axis=1)
- index = part_index[part]
- centroids[index] = centroid
- radii[index] = float(distances.max())
- counts[index] = len(vectors)
- BASELINE_DIR.mkdir(parents=True, exist_ok=True)
- np.savez(
- BASELINE_PATH,
- part_names=np.asarray(parts),
- centroids=centroids,
- radii=radii,
- counts=counts,
- )
- print(f"基准表已保存: {BASELINE_PATH}")
- with (BASELINE_DIR / "baseline_summary.csv").open("w", encoding="utf-8", newline="") as handle:
- writer = csv.writer(handle)
- writer.writerow(["part", "count", "radius"])
- for part, count, radius in zip(parts, counts, radii):
- writer.writerow([part, int(count), f"{radius:.6f}"])
- print("\n部位 / 样本数 / 半径:")
- for part, count, radius in zip(parts, counts, radii):
- print(f" {part}: {count} 组, 半径 {radius:.4f}")
- return 0
- if __name__ == "__main__":
- raise SystemExit(main())
|