Buckets:
| """Train a figure-skating action classifier on the pipeline's (T=64, 94)-feature tensors. | |
| Architecture ported from training/cnn_lstm_attention.py (Keras): a Conv1D residual backbone, | |
| a stack of multi-head self-attention blocks, temporal pooling, and a deep GELU dense head. | |
| Kept deliberately simple on the training-loop side (Adam + CrossEntropy + val-based best-model | |
| selection). Metrics (precision / recall / F1) come from scikit-learn. | |
| LABEL CONSISTENCY (the thing that must not break): | |
| The pipeline saves integer labels that are *exactly* labels.TAXONOMY indices (0..27). | |
| This model's output layer has num_classes = len(TAXONOMY) units, so output unit i is, by | |
| construction, the score for TAXONOMY[i]. No compaction / remapping is done anywhere, so | |
| argmax(logits) is directly a TAXONOMY index. Assertions below enforce this before training. | |
| Usage: | |
| python model.py # train on /data/processed_500 | |
| python model.py --data-dir DIR # train on another processed dir | |
| python model.py --smoke # self-test the plumbing on synthetic data (no real data needed) | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import pickle | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from sklearn.metrics import ( | |
| accuracy_score, | |
| classification_report, | |
| confusion_matrix, | |
| f1_score, | |
| precision_recall_fscore_support, | |
| ) | |
| from sklearn.utils.class_weight import compute_class_weight | |
| from torch.utils.data import DataLoader, TensorDataset | |
| try: | |
| from . import labels as labels_mod | |
| except ImportError: | |
| import labels as labels_mod | |
| # --------------------------------------------------------------------------- | |
| # Config | |
| # --------------------------------------------------------------------------- | |
| DEFAULT_DATA_DIR = "/data/processed_500" | |
| EPOCHS = 100 | |
| BATCH_SIZE = 64 | |
| LEARNING_RATE = 5e-4 # lowered from 1e-3: the raw velocity features are heavy-tailed (|v|~960) | |
| WEIGHT_DECAY = 1e-4 | |
| GRAD_CLIP = 1.0 # clip grad norm; prevents rare-class / heavy-tail gradient spikes -> NaN | |
| MAX_CLASS_WEIGHT = 5.0 # cap balanced class weights (singletons would otherwise get ~18x) | |
| EARLY_STOP_PATIENCE = 20 | |
| SEED = 42 | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| # --------------------------------------------------------------------------- | |
| # Model (ported from training/cnn_lstm_attention.py) | |
| # --------------------------------------------------------------------------- | |
| class ConvBlock(nn.Module): | |
| """Conv1D + BN + Dropout with a (projected) residual connection. Input/return: (B, C, T).""" | |
| def __init__(self, in_ch: int, out_ch: int, kernel: int, dropout: float = 0.15): | |
| super().__init__() | |
| self.shortcut = nn.Conv1d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity() | |
| self.conv = nn.Conv1d(in_ch, out_ch, kernel, padding="same") | |
| self.bn = nn.BatchNorm1d(out_ch) | |
| self.drop = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| y = F.relu(self.conv(x)) | |
| y = self.drop(self.bn(y)) | |
| return self.shortcut(x) + y | |
| class AttnBlock(nn.Module): | |
| """Multi-head self-attention with residual + LayerNorm + Dropout. Input/return: (B, T, D).""" | |
| def __init__(self, dim: int = 384, heads: int = 8, attn_dropout: float = 0.1, dropout: float = 0.2): | |
| super().__init__() | |
| self.mha = nn.MultiheadAttention(dim, heads, dropout=attn_dropout, batch_first=True) | |
| self.norm = nn.LayerNorm(dim) | |
| self.drop = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| attn, _ = self.mha(x, x, x, need_weights=False) | |
| return self.drop(self.norm(x + attn)) | |
| class DenseBlock(nn.Module): | |
| """Linear + GELU + LayerNorm + BatchNorm + Dropout. Input/return: (B, features).""" | |
| def __init__(self, in_f: int, out_f: int, dropout: float): | |
| super().__init__() | |
| self.fc = nn.Linear(in_f, out_f) | |
| self.ln = nn.LayerNorm(out_f) | |
| self.bn = nn.BatchNorm1d(out_f) | |
| self.drop = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = F.gelu(self.fc(x)) | |
| return self.drop(self.bn(self.ln(x))) | |
| class SkatingActionClassifier(nn.Module): | |
| """(B, T, F) sequence of skeleton features -> (B, num_classes) class logits.""" | |
| def __init__(self, in_features: int, num_classes: int): | |
| super().__init__() | |
| # Conv backbone (expects channels-first (B, F, T)) | |
| self.stem = nn.Conv1d(in_features, 128, 3, padding="same") | |
| self.stem_bn = nn.BatchNorm1d(128) | |
| self.cb1 = ConvBlock(128, 192, 3) | |
| self.cb2 = ConvBlock(192, 256, 3) | |
| self.cb3 = ConvBlock(256, 384, 5) | |
| # Attention stack (operates on (B, T, 384)) | |
| self.attn = nn.ModuleList([AttnBlock(384, heads=8) for _ in range(3)]) | |
| # Deep classification head (temporal mean-pool -> 384). The reference concatenated two | |
| # identical mean-pools (768); a single mean-pool carries the same information. | |
| self.head = nn.Sequential( | |
| DenseBlock(384, 1024, 0.5), | |
| DenseBlock(1024, 512, 0.4), | |
| DenseBlock(512, 256, 0.3), | |
| nn.LayerNorm(256), | |
| ) | |
| self.out = nn.Linear(256, num_classes) # logits; softmax handled by CrossEntropyLoss | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = x.transpose(1, 2) # (B, F, T) | |
| x = F.relu(self.stem(x)) | |
| x = self.stem_bn(x) | |
| x = self.cb3(self.cb2(self.cb1(x))) # (B, 384, T) | |
| x = x.transpose(1, 2) # (B, T, 384) | |
| for blk in self.attn: | |
| x = blk(x) | |
| x = x.mean(dim=1) # temporal average pool -> (B, 384) | |
| return self.out(self.head(x)) | |
| # --------------------------------------------------------------------------- | |
| # Data | |
| # --------------------------------------------------------------------------- | |
| def _load_pickle(path: Path): | |
| with path.open("rb") as f: | |
| return pickle.load(f) | |
| def load_split(data_dir: Path, split: str) -> tuple[np.ndarray, np.ndarray]: | |
| X = np.asarray(_load_pickle(data_dir / f"{split}_features.pkl"), dtype=np.float32) | |
| y = np.asarray(_load_pickle(data_dir / f"{split}_labels.pkl"), dtype=np.int64) | |
| return X, y | |
| def resolve_num_classes(data_dir: Path) -> int: | |
| """num_classes comes from the taxonomy; cross-checked against the pipeline's metadata.json.""" | |
| n = len(labels_mod.TAXONOMY) | |
| meta_path = data_dir / "metadata.json" | |
| if meta_path.exists(): | |
| meta_n = int(json.loads(meta_path.read_text())["num_classes"]) | |
| assert meta_n == n, ( | |
| f"metadata num_classes={meta_n} disagrees with labels.TAXONOMY ({n}). " | |
| "The pipeline and the model are using different taxonomies -- refusing to train." | |
| ) | |
| return n | |
| def assert_label_consistency(num_classes: int, *label_arrays: np.ndarray) -> None: | |
| """Every saved label must be a valid TAXONOMY index; guarantees softmax<->label alignment.""" | |
| for y in label_arrays: | |
| if y.size == 0: | |
| continue | |
| lo, hi = int(y.min()), int(y.max()) | |
| assert 0 <= lo and hi < num_classes, ( | |
| f"label out of range: found [{lo}, {hi}] but valid TAXONOMY indices are " | |
| f"[0, {num_classes - 1}]. Softmax index space would not match the labels." | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Train / evaluate | |
| # --------------------------------------------------------------------------- | |
| def run_epoch(model, loader, criterion, optimizer=None) -> tuple[float, float]: | |
| train = optimizer is not None | |
| model.train(train) | |
| total_loss, correct, n = 0.0, 0, 0 | |
| for xb, yb in loader: | |
| xb, yb = xb.to(DEVICE), yb.to(DEVICE) | |
| with torch.set_grad_enabled(train): | |
| logits = model(xb) | |
| loss = criterion(logits, yb) | |
| if train: | |
| if not torch.isfinite(loss): # defensive: never step on a non-finite loss | |
| continue | |
| optimizer.zero_grad() | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP) | |
| optimizer.step() | |
| total_loss += loss.item() * xb.size(0) | |
| correct += (logits.argmax(1) == yb).sum().item() | |
| n += xb.size(0) | |
| return total_loss / n, correct / n | |
| def predict(model, X: np.ndarray) -> np.ndarray: | |
| model.eval() | |
| out = [] | |
| for i in range(0, len(X), 256): | |
| xb = torch.from_numpy(X[i : i + 256]).to(DEVICE) | |
| out.append(model(xb).argmax(1).cpu().numpy()) | |
| return np.concatenate(out) if out else np.array([], dtype=np.int64) | |
| def report_metrics(y_true: np.ndarray, y_pred: np.ndarray, taxonomy: dict) -> dict: | |
| num_classes = len(taxonomy) | |
| class_ids = list(range(num_classes)) | |
| names = [taxonomy[i] for i in class_ids] | |
| acc = accuracy_score(y_true, y_pred) | |
| p_macro, r_macro, f_macro, _ = precision_recall_fscore_support( | |
| y_true, y_pred, labels=class_ids, average="macro", zero_division=0 | |
| ) | |
| p_w, r_w, f_w, _ = precision_recall_fscore_support( | |
| y_true, y_pred, labels=class_ids, average="weighted", zero_division=0 | |
| ) | |
| print("\n" + "=" * 72) | |
| print("TEST METRICS") | |
| print("=" * 72) | |
| print(f"accuracy : {acc:.4f}") | |
| print(f"precision (macro/weighted): {p_macro:.4f} / {p_w:.4f}") | |
| print(f"recall (macro/weighted): {r_macro:.4f} / {r_w:.4f}") | |
| print(f"f1 (macro/weighted): {f_macro:.4f} / {f_w:.4f}") | |
| present = sorted(set(y_true.tolist()) | set(y_pred.tolist())) | |
| print("\nper-class report (classes present in test true/pred only):") | |
| print( | |
| classification_report( | |
| y_true, y_pred, | |
| labels=present, | |
| target_names=[taxonomy[i] for i in present], | |
| zero_division=0, | |
| digits=3, | |
| ) | |
| ) | |
| return { | |
| "accuracy": acc, | |
| "precision_macro": p_macro, "recall_macro": r_macro, "f1_macro": f_macro, | |
| "precision_weighted": p_w, "recall_weighted": r_w, "f1_weighted": f_w, | |
| } | |
| def train(data_dir: Path, coarse: bool = False) -> dict: | |
| torch.manual_seed(SEED) | |
| np.random.seed(SEED) | |
| fine_n = resolve_num_classes(data_dir) # 28, cross-checked against metadata.json | |
| taxonomy = labels_mod.COARSE_TAXONOMY if coarse else labels_mod.TAXONOMY | |
| num_classes = len(taxonomy) | |
| Xtr, ytr = load_split(data_dir, "train") | |
| Xva, yva = load_split(data_dir, "val") | |
| Xte, yte = load_split(data_dir, "test") | |
| # Saved labels are always fine TAXONOMY indices -> validate in that space first, then coarsen. | |
| # This keeps the whole chain consistent: saved fine idx -> FINE_TO_COARSE_IDX -> softmax unit. | |
| assert_label_consistency(fine_n, ytr, yva, yte) | |
| if coarse: | |
| def coarsen(y): | |
| return np.array([labels_mod.FINE_TO_COARSE_IDX[int(v)] for v in y], dtype=np.int64) | |
| ytr, yva, yte = coarsen(ytr), coarsen(yva), coarsen(yte) | |
| assert_label_consistency(num_classes, ytr, yva, yte) | |
| in_features = Xtr.shape[-1] | |
| print(f"label space: {'COARSE action-level' if coarse else 'FINE'} | {num_classes} classes") | |
| # Per-feature standardization from TRAIN stats only (applied to all splits + stored for | |
| # inference). The velocity block is heavy-tailed (|v| up to ~960 vs coords ~1); feeding it raw | |
| # blows up activations -> NaN. z-scoring brings every feature onto a comparable scale. | |
| mu = Xtr.mean(axis=(0, 1), keepdims=True) | |
| sd = Xtr.std(axis=(0, 1), keepdims=True) + 1e-6 | |
| Xtr = (Xtr - mu) / sd | |
| Xva = (Xva - mu) / sd | |
| Xte = (Xte - mu) / sd | |
| print(f"device={DEVICE} | num_classes={num_classes} (== len(taxonomy)) | in_features={in_features}") | |
| print(f"shapes: train={Xtr.shape} val={Xva.shape} test={Xte.shape}") | |
| print(f"standardized features: train |max|={np.abs(Xtr).max():.1f} (was heavy-tailed pre-scaling)") | |
| print(f"train classes present: {sorted(set(ytr.tolist()))}") | |
| # class-balanced loss weights (train distribution), capped so singleton classes don't get | |
| # ~18x weight and destabilize training. "proper" imbalance handling without the spikes. | |
| present = np.unique(ytr) | |
| cw = np.clip(compute_class_weight(class_weight="balanced", classes=present, y=ytr), | |
| None, MAX_CLASS_WEIGHT) | |
| weight = torch.ones(num_classes) | |
| for c, w in zip(present, cw): | |
| weight[int(c)] = float(w) | |
| criterion = nn.CrossEntropyLoss(weight=weight.to(DEVICE)) | |
| model = SkatingActionClassifier(in_features, num_classes).to(DEVICE) | |
| n_params = sum(p.numel() for p in model.parameters()) | |
| print(f"model params: {n_params/1e6:.2f}M") | |
| optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY) | |
| tr_loader = DataLoader( | |
| TensorDataset(torch.from_numpy(Xtr), torch.from_numpy(ytr)), | |
| batch_size=BATCH_SIZE, shuffle=True, drop_last=False, | |
| ) | |
| va_loader = DataLoader( | |
| TensorDataset(torch.from_numpy(Xva), torch.from_numpy(yva)), | |
| batch_size=BATCH_SIZE, shuffle=False, | |
| ) | |
| best_val_f1, best_state, since_improved = -1.0, None, 0 | |
| for epoch in range(1, EPOCHS + 1): | |
| tr_loss, tr_acc = run_epoch(model, tr_loader, criterion, optimizer) | |
| va_loss, va_acc = run_epoch(model, va_loader, criterion) | |
| va_f1 = f1_score(yva, predict(model, Xva), labels=list(range(num_classes)), | |
| average="macro", zero_division=0) | |
| if va_f1 > best_val_f1: | |
| best_val_f1, since_improved = va_f1, 0 | |
| best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()} | |
| else: | |
| since_improved += 1 | |
| if epoch % 5 == 0 or epoch == 1: | |
| print(f"epoch {epoch:3d} | train loss {tr_loss:.3f} acc {tr_acc:.3f} " | |
| f"| val loss {va_loss:.3f} acc {va_acc:.3f} f1(macro) {va_f1:.3f}") | |
| if since_improved >= EARLY_STOP_PATIENCE: | |
| print(f"early stop at epoch {epoch} (no val-F1 improvement for {EARLY_STOP_PATIENCE})") | |
| break | |
| if best_state is not None: | |
| model.load_state_dict(best_state) | |
| print(f"\nrestored best model (val macro-F1 = {best_val_f1:.3f})") | |
| metrics = report_metrics(yte, predict(model, Xte), taxonomy) | |
| ckpt_name = "model_coarse.pt" if coarse else "model.pt" | |
| torch.save({"state_dict": model.state_dict(), "num_classes": num_classes, | |
| "in_features": in_features, "taxonomy": taxonomy, "coarse": coarse, | |
| "feature_mean": mu, "feature_std": sd}, # standardizer for inference consistency | |
| data_dir / ckpt_name) | |
| print(f"\nsaved model -> {data_dir / ckpt_name}") | |
| return metrics | |
| def smoke() -> None: | |
| """Validate the full plumbing on synthetic data shaped exactly like the pipeline output.""" | |
| print("SMOKE: synthetic (N,64,94) data, labels in TAXONOMY index space") | |
| num_classes = len(labels_mod.TAXONOMY) | |
| rng = np.random.default_rng(0) | |
| X = rng.standard_normal((40, 64, 94)).astype(np.float32) | |
| y = rng.integers(0, num_classes, size=40).astype(np.int64) | |
| assert_label_consistency(num_classes, y) | |
| model = SkatingActionClassifier(94, num_classes).to(DEVICE) | |
| logits = model(torch.from_numpy(X[:4]).to(DEVICE)) | |
| assert logits.shape == (4, num_classes), logits.shape | |
| loss = nn.CrossEntropyLoss()(logits, torch.from_numpy(y[:4]).to(DEVICE)) | |
| loss.backward() | |
| print(f"forward OK: logits {tuple(logits.shape)} | loss {loss.item():.3f} | backward OK") | |
| print(f"num_classes={num_classes} matches softmax units; argmax range in [0,{num_classes-1}]") | |
| print("SMOKE PASSED") | |
| def main() -> int: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--data-dir", type=Path, default=Path(DEFAULT_DATA_DIR)) | |
| parser.add_argument("--coarse", action="store_true", | |
| help="collapse jump rotations into action-level classes (11 instead of 28)") | |
| parser.add_argument("--smoke", action="store_true", help="self-test on synthetic data") | |
| args = parser.parse_args() | |
| if args.smoke: | |
| smoke() | |
| return 0 | |
| if not (args.data_dir / "train_features.pkl").exists(): | |
| raise SystemExit(f"No processed data at {args.data_dir}. Run the pipeline first.") | |
| train(args.data_dir, coarse=args.coarse) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |
Xet Storage Details
- Size:
- 16.6 kB
- Xet hash:
- b8370ce2c2d2789d46696c7bb594a330aa07c5d44feda9cb953b9be8a7d38502
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.