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