Spaces:
Sleeping
Sleeping
File size: 6,651 Bytes
20d7fde 1b69166 20d7fde 1b69166 20d7fde 1b69166 20d7fde 1b69166 20d7fde 1b69166 20d7fde | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | """Train the Togyzkumalak move classifier on on-the-fly synthetic cells.
Local smoke test (Mac, ~minutes):
python train.py --epochs 1 --samples-per-epoch 4000 --batch-size 64
Full run (Colab GPU or any CUDA machine):
python train.py --epochs 30
Prints synthetic validation accuracy AND accuracy on the real labeled crops
(data/real_crops) every epoch; saves checkpoints/best.pt and last.pt.
"""
import argparse
import json
import time
from pathlib import Path
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from togyz.classes import CLASSES, DIAGRAM_CLASSES
from togyz.dataset import RealCropDataset, SyntheticCellDataset
from togyz.model import auto_device, build_model, save_checkpoint
def evaluate(model, loader, device) -> float:
model.eval()
correct = total = 0
with torch.no_grad():
for images, labels in loader:
images, labels = images.to(device), labels.to(device)
predictions = model(images).argmax(dim=1)
correct += (predictions == labels).sum().item()
total += labels.numel()
return correct / max(total, 1)
def evaluate_real(model, real: RealCropDataset, device) -> tuple[float, float]:
"""Top-1 and top-3 accuracy on the real crops."""
if len(real) == 0:
return float("nan"), float("nan")
images, labels, _ = real.batch()
model.eval()
with torch.no_grad():
logits = model(images.to(device)).cpu()
top3 = logits.topk(3, dim=1).indices
top1_acc = (top3[:, 0] == labels).float().mean().item()
top3_acc = (top3 == labels[:, None]).any(dim=1).float().mean().item()
return top1_acc, top3_acc
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--task", choices=["moves", "diagram"], default="moves",
help="moves: 163-class cell classifier; diagram: unified "
"board-diagram reader (0-81, 'x', '-') for the kazan "
"boxes and pit cells of the summary strips "
"(replaces the old 'kazan' task)")
parser.add_argument("--epochs", type=int, default=20)
parser.add_argument("--samples-per-epoch", type=int, default=50_000)
parser.add_argument("--val-size", type=int, default=4_000)
parser.add_argument("--batch-size", type=int, default=128)
parser.add_argument("--lr", type=float, default=3e-4)
parser.add_argument("--weight-decay", type=float, default=1e-4)
parser.add_argument("--workers", type=int, default=4)
parser.add_argument("--out", default=None,
help="default: checkpoints (moves) / checkpoints/diagram")
parser.add_argument("--resume", default=None, help="path to last.pt to continue")
parser.add_argument("--device", default=None, help="cuda / mps / cpu (default: auto)")
args = parser.parse_args()
device = torch.device(args.device) if args.device else auto_device()
print(f"Device: {device}, task: {args.task}")
classes = DIAGRAM_CLASSES if args.task == "diagram" else CLASSES
out_dir = Path(args.out or ("checkpoints/diagram" if args.task == "diagram" else "checkpoints"))
out_dir.mkdir(parents=True, exist_ok=True)
(out_dir / "classes.json").write_text(json.dumps(classes))
train_set = SyntheticCellDataset(args.samples_per_epoch, seed=None,
classes=classes, task=args.task)
val_set = SyntheticCellDataset(args.val_size, seed=1234,
classes=classes, task=args.task)
# the labeled real crops are move cells; other tasks have no real eval set
real_set = RealCropDataset() if args.task == "moves" else RealCropDataset("/nonexistent")
print(f"Glyph pools:\n{train_set.sampler.describe()}")
print(f"Real eval crops: {len(real_set)}")
loader_kwargs = dict(
batch_size=args.batch_size,
num_workers=args.workers,
pin_memory=(device.type == "cuda"),
persistent_workers=args.workers > 0,
)
train_loader = DataLoader(train_set, shuffle=False, **loader_kwargs)
val_loader = DataLoader(val_set, shuffle=False, **loader_kwargs)
model = build_model(len(classes)).to(device)
criterion = nn.CrossEntropyLoss(label_smoothing=0.05)
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
start_epoch, best_acc = 0, 0.0
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=True)
model.load_state_dict(ckpt["model_state"])
if ckpt.get("optimizer_state"):
optimizer.load_state_dict(ckpt["optimizer_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("val_acc", 0.0)
for _ in range(start_epoch):
scheduler.step()
print(f"Resumed from {args.resume} at epoch {start_epoch}")
for epoch in range(start_epoch, args.epochs):
model.train()
epoch_start = time.time()
running_loss, seen = 0.0, 0
for step, (images, labels) in enumerate(train_loader):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad(set_to_none=True)
loss = criterion(model(images), labels)
loss.backward()
optimizer.step()
running_loss += loss.item() * labels.numel()
seen += labels.numel()
if step % 50 == 0:
print(
f" epoch {epoch + 1} step {step + 1}/{len(train_loader)} "
f"loss {running_loss / seen:.4f}",
flush=True,
)
scheduler.step()
val_acc = evaluate(model, val_loader, device)
real_top1, real_top3 = evaluate_real(model, real_set, device)
elapsed = time.time() - epoch_start
print(
f"Epoch {epoch + 1}/{args.epochs} [{elapsed:.0f}s] "
f"loss {running_loss / max(seen, 1):.4f} | synth val {val_acc:.2%} | "
f"real top-1 {real_top1:.2%} top-3 {real_top3:.2%}",
flush=True,
)
save_checkpoint(out_dir / "last.pt", model, epoch, val_acc, optimizer, classes=classes)
if val_acc >= best_acc:
best_acc = val_acc
save_checkpoint(out_dir / "best.pt", model, epoch, val_acc, classes=classes)
print(f" new best ({val_acc:.2%}) -> {out_dir / 'best.pt'}")
print(f"Done. Best synthetic val accuracy: {best_acc:.2%}")
if __name__ == "__main__":
main()
|