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