train.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144
  1. """Step 1: unsupervised TSPulse training (masked reconstruction).
  2. Trains a single TSPulse encoder on all normal cycle samples from every part,
  3. then freezes the weights into a single checkpoint file used later for
  4. inference. Progress is printed via tqdm (epochs and batches with live loss).
  5. Usage:
  6. python3 tspulse/train.py [--epochs 30] [--batch-size 64] [--lr 1e-3]
  7. """
  8. import argparse
  9. import json
  10. import random
  11. from pathlib import Path
  12. import numpy as np
  13. import torch
  14. from torch.utils.data import DataLoader
  15. from tqdm import tqdm
  16. from tspulse import CycleDataset, TSPulse, build_split_datasets, describe, get_device, make_random_mask
  17. from tspulse.dataset import N_CHANNELS, SEQ_LEN
  18. ROOT = Path(__file__).resolve().parents[1]
  19. CHECKPOINT_DIR = ROOT / "checkpoints"
  20. CHECKPOINT_PATH = CHECKPOINT_DIR / "tspulse_frozen.pt"
  21. def set_seed(seed: int) -> None:
  22. random.seed(seed)
  23. np.random.seed(seed)
  24. torch.manual_seed(seed)
  25. def masked_mse(recon: torch.Tensor, target: torch.Tensor, mask: torch.BoolTensor, patch_size: int) -> torch.Tensor:
  26. sample_mask = mask.repeat_interleave(patch_size, dim=1) # (B, L)
  27. diff = (recon - target) ** 2
  28. masked_diff = diff[sample_mask]
  29. if masked_diff.numel() == 0:
  30. return torch.tensor(0.0, device=recon.device)
  31. return masked_diff.mean()
  32. def main() -> int:
  33. parser = argparse.ArgumentParser()
  34. parser.add_argument("--epochs", type=int, default=30)
  35. parser.add_argument("--batch-size", type=int, default=64)
  36. parser.add_argument("--lr", type=float, default=1e-3)
  37. parser.add_argument("--mask-ratio", type=float, default=0.15)
  38. parser.add_argument("--dim", type=int, default=128)
  39. parser.add_argument("--depth", type=int, default=4)
  40. parser.add_argument("--heads", type=int, default=4)
  41. parser.add_argument("--seed", type=int, default=42)
  42. args = parser.parse_args()
  43. set_seed(args.seed)
  44. device = get_device()
  45. print(f"设备: {describe(device)}")
  46. samples, mean, std = build_split_datasets()
  47. train_samples = [item for item in samples if item["split"] == "train"]
  48. val_samples = [item for item in samples if item["split"] == "val"]
  49. print(f"样本: 训练 {len(train_samples)}, 验证 {len(val_samples)}, 评估 {len(samples) - len(train_samples) - len(val_samples)}")
  50. train_set = CycleDataset(train_samples, mean, std)
  51. val_set = CycleDataset(val_samples, mean, std)
  52. train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=True, drop_last=True)
  53. val_loader = DataLoader(val_set, batch_size=args.batch_size, shuffle=False)
  54. model = TSPulse(dim=args.dim, depth=args.depth, heads=args.heads).to(device)
  55. optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05)
  56. scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
  57. config = {
  58. "dim": args.dim,
  59. "depth": args.depth,
  60. "heads": args.heads,
  61. "seq_len": SEQ_LEN,
  62. "n_channels": N_CHANNELS,
  63. "mask_ratio": args.mask_ratio,
  64. }
  65. best_val = float("inf")
  66. CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)
  67. history: list[dict] = []
  68. for epoch in range(1, args.epochs + 1):
  69. model.train()
  70. epoch_loss = 0.0
  71. seen = 0
  72. bar = tqdm(train_loader, desc=f"epoch {epoch}/{args.epochs}", leave=False)
  73. for batch in bar:
  74. x, _lengths = batch
  75. x = x.to(device)
  76. mask = make_random_mask(x.shape[0], model.num_patches, args.mask_ratio, device)
  77. _, recon = model(x, mask)
  78. loss = masked_mse(recon, x, mask, model.patch_size)
  79. optimizer.zero_grad()
  80. loss.backward()
  81. optimizer.step()
  82. epoch_loss += loss.item() * x.shape[0]
  83. seen += x.shape[0]
  84. bar.set_postfix(loss=f"{loss.item():.5f}")
  85. scheduler.step()
  86. train_loss = epoch_loss / max(seen, 1)
  87. model.eval()
  88. val_loss = 0.0
  89. val_seen = 0
  90. with torch.no_grad():
  91. for batch in val_loader:
  92. x, _lengths = batch
  93. x = x.to(device)
  94. mask = make_random_mask(x.shape[0], model.num_patches, args.mask_ratio, device)
  95. _, recon = model(x, mask)
  96. val_loss += masked_mse(recon, x, mask, model.patch_size).item() * x.shape[0]
  97. val_seen += x.shape[0]
  98. val_loss /= max(val_seen, 1)
  99. history.append({"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss})
  100. print(f"epoch {epoch:>3}/{args.epochs} 训练 loss {train_loss:.5f} 验证 loss {val_loss:.5f}")
  101. if val_loss < best_val:
  102. best_val = val_loss
  103. torch.save(
  104. {
  105. "state_dict": model.state_dict(),
  106. "config": config,
  107. "mean": mean,
  108. "std": std,
  109. "parts": sorted({item["part"] for item in samples}),
  110. "history": history,
  111. "device": describe(device),
  112. },
  113. CHECKPOINT_PATH,
  114. )
  115. print(f" -> 已保存最佳模型 {CHECKPOINT_PATH}")
  116. print(f"训练完成: 最佳验证 loss {best_val:.5f}, 模型固化于 {CHECKPOINT_PATH}")
  117. return 0
  118. if __name__ == "__main__":
  119. raise SystemExit(main())