"""Step 1: unsupervised TSPulse training (masked reconstruction). Trains a single TSPulse encoder on all normal cycle samples from every part, then freezes the weights into a single checkpoint file used later for inference. Progress is printed via tqdm (epochs and batches with live loss). Usage: python3 tspulse/train.py [--epochs 30] [--batch-size 64] [--lr 1e-3] """ import argparse import json import random 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 from tspulse.dataset import N_CHANNELS, SEQ_LEN ROOT = Path(__file__).resolve().parents[1] CHECKPOINT_DIR = ROOT / "checkpoints" CHECKPOINT_PATH = CHECKPOINT_DIR / "tspulse_frozen.pt" def set_seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) 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) # (B, L) 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("--epochs", type=int, default=30) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--lr", type=float, default=1e-3) parser.add_argument("--mask-ratio", type=float, default=0.15) parser.add_argument("--dim", type=int, default=128) parser.add_argument("--depth", type=int, default=4) parser.add_argument("--heads", type=int, default=4) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() set_seed(args.seed) device = get_device() print(f"设备: {describe(device)}") samples, mean, std = build_split_datasets() train_samples = [item for item in samples if item["split"] == "train"] val_samples = [item for item in samples if item["split"] == "val"] print(f"样本: 训练 {len(train_samples)}, 验证 {len(val_samples)}, 评估 {len(samples) - len(train_samples) - len(val_samples)}") train_set = CycleDataset(train_samples, mean, std) val_set = CycleDataset(val_samples, mean, std) train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=True, drop_last=True) val_loader = DataLoader(val_set, batch_size=args.batch_size, shuffle=False) model = TSPulse(dim=args.dim, depth=args.depth, heads=args.heads).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs) config = { "dim": args.dim, "depth": args.depth, "heads": args.heads, "seq_len": SEQ_LEN, "n_channels": N_CHANNELS, "mask_ratio": args.mask_ratio, } best_val = float("inf") CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True) history: list[dict] = [] for epoch in range(1, args.epochs + 1): model.train() epoch_loss = 0.0 seen = 0 bar = tqdm(train_loader, desc=f"epoch {epoch}/{args.epochs}", leave=False) for batch in bar: x, _lengths = batch x = x.to(device) mask = make_random_mask(x.shape[0], model.num_patches, args.mask_ratio, device) _, recon = model(x, mask) loss = masked_mse(recon, x, mask, model.patch_size) optimizer.zero_grad() loss.backward() optimizer.step() epoch_loss += loss.item() * x.shape[0] seen += x.shape[0] bar.set_postfix(loss=f"{loss.item():.5f}") scheduler.step() train_loss = epoch_loss / max(seen, 1) model.eval() val_loss = 0.0 val_seen = 0 with torch.no_grad(): for batch in val_loader: x, _lengths = batch x = x.to(device) mask = make_random_mask(x.shape[0], model.num_patches, args.mask_ratio, device) _, recon = model(x, mask) val_loss += masked_mse(recon, x, mask, model.patch_size).item() * x.shape[0] val_seen += x.shape[0] val_loss /= max(val_seen, 1) history.append({"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss}) print(f"epoch {epoch:>3}/{args.epochs} 训练 loss {train_loss:.5f} 验证 loss {val_loss:.5f}") if val_loss < best_val: best_val = val_loss torch.save( { "state_dict": model.state_dict(), "config": config, "mean": mean, "std": std, "parts": sorted({item["part"] for item in samples}), "history": history, "device": describe(device), }, CHECKPOINT_PATH, ) print(f" -> 已保存最佳模型 {CHECKPOINT_PATH}") print(f"训练完成: 最佳验证 loss {best_val:.5f}, 模型固化于 {CHECKPOINT_PATH}") return 0 if __name__ == "__main__": raise SystemExit(main())