Buckets:
| """xLSTM classifier for figure-skating action classification. | |
| Uses the official NX-AI/xlstm implementation (Beck et al., 2024). Same Conv1D residual | |
| backbone as model.py, with an xLSTM block stack (mix of mLSTM and sLSTM blocks) as the | |
| temporal-modeling stage, followed by the same dense classification head used across this | |
| project. | |
| Usage: | |
| python -m temporal_scripts.model_xlstm --data-dir /path/to/processed | |
| python -m temporal_scripts.model_xlstm --data-dir /path/to/processed --smoke | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| 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 f1_score | |
| from sklearn.utils.class_weight import compute_class_weight | |
| from torch.utils.data import DataLoader, TensorDataset | |
| from xlstm import ( | |
| FeedForwardConfig, | |
| mLSTMBlockConfig, | |
| mLSTMLayerConfig, | |
| sLSTMBlockConfig, | |
| sLSTMLayerConfig, | |
| xLSTMBlockStack, | |
| xLSTMBlockStackConfig, | |
| ) | |
| import sys | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) | |
| import labels as labels_mod | |
| from model import ( | |
| ConvBlock, | |
| DenseBlock, | |
| load_split, | |
| resolve_taxonomy, | |
| assert_label_consistency, | |
| report_metrics, | |
| predict, | |
| ) | |
| EPOCHS = 100 | |
| BATCH_SIZE = 64 | |
| LEARNING_RATE = 5e-4 | |
| WEIGHT_DECAY = 1e-4 | |
| GRAD_CLIP = 1.0 | |
| MAX_CLASS_WEIGHT = 5.0 | |
| EARLY_STOP_PATIENCE = 20 | |
| SEED = 42 | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| D_MODEL = 384 | |
| XLSTM_NUM_BLOCKS = 4 | |
| XLSTM_NUM_HEADS = 4 | |
| XLSTM_CONTEXT_LENGTH = 128 | |
| class SkatingxLSTMClassifier(nn.Module): | |
| """(B, T, F) sequence of skeleton features -> (B, num_classes) class logits. | |
| Conv backbone -> xLSTM block stack -> temporal mean pool -> dense head. | |
| The xLSTM stack uses a mix of mLSTM (matrix-valued memory, parallelizable) and | |
| sLSTM (scalar memory with new gating) blocks. | |
| """ | |
| def __init__(self, in_features: int, num_classes: int, | |
| num_blocks: int = XLSTM_NUM_BLOCKS, | |
| context_length: int = XLSTM_CONTEXT_LENGTH): | |
| super().__init__() | |
| 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, D_MODEL, 5) | |
| cfg = xLSTMBlockStackConfig( | |
| mlstm_block=mLSTMBlockConfig( | |
| mlstm=mLSTMLayerConfig( | |
| conv1d_kernel_size=4, | |
| num_heads=XLSTM_NUM_HEADS, | |
| ), | |
| ), | |
| slstm_block=sLSTMBlockConfig( | |
| slstm=sLSTMLayerConfig( | |
| num_heads=XLSTM_NUM_HEADS, | |
| conv1d_kernel_size=4, | |
| backend="vanilla", | |
| ), | |
| feedforward=FeedForwardConfig(proj_factor=1.3, act_fn="gelu"), | |
| ), | |
| context_length=context_length, | |
| num_blocks=num_blocks, | |
| embedding_dim=D_MODEL, | |
| slstm_at=[1], | |
| dropout=0.1, | |
| ) | |
| self.xlstm = xLSTMBlockStack(cfg) | |
| self.head = nn.Sequential( | |
| DenseBlock(D_MODEL, 1024, 0.5), | |
| DenseBlock(1024, 512, 0.4), | |
| DenseBlock(512, 256, 0.3), | |
| nn.LayerNorm(256), | |
| ) | |
| self.out = nn.Linear(256, num_classes) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = x.transpose(1, 2) | |
| x = F.relu(self.stem(x)) | |
| x = self.stem_bn(x) | |
| x = self.cb3(self.cb2(self.cb1(x))) | |
| x = x.transpose(1, 2) | |
| x = self.xlstm(x) | |
| x = x.mean(dim=1) | |
| return self.out(self.head(x)) | |
| def forward_per_step(self, x: torch.Tensor) -> torch.Tensor: | |
| """Return per-timestep logits (B, T, num_classes) for per-frame scoring.""" | |
| x = x.transpose(1, 2) | |
| x = F.relu(self.stem(x)) | |
| x = self.stem_bn(x) | |
| x = self.cb3(self.cb2(self.cb1(x))) | |
| x = x.transpose(1, 2) | |
| x = self.xlstm(x) | |
| B, T, D = x.shape | |
| flat = x.reshape(B * T, D) | |
| logits = self.out(self.head(flat)) | |
| return logits.reshape(B, T, -1) | |
| 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): | |
| 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 train(data_dir: Path, coarse: bool = False, num_blocks: int = XLSTM_NUM_BLOCKS) -> dict: | |
| torch.manual_seed(SEED) | |
| np.random.seed(SEED) | |
| saved_n, taxonomy = resolve_taxonomy(data_dir, coarse) | |
| 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") | |
| assert_label_consistency(saved_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] | |
| if coarse: | |
| label_space_name = "COARSE action-level" | |
| elif taxonomy is labels_mod.FS_JUMP3D_TAXONOMY: | |
| label_space_name = "FS_JUMP3D" | |
| elif taxonomy is labels_mod.FS_JUMP3D_SINGLES_TAXONOMY: | |
| label_space_name = "FS_JUMP3D_SINGLES (Comb excluded)" | |
| else: | |
| label_space_name = "FINE" | |
| print(f"[xLSTM] label space: {label_space_name} | {num_classes} classes") | |
| mu = Xtr.mean(axis=(0, 1), keepdims=True) | |
| sd = Xtr.std(axis=(0, 1), keepdims=True) + 1e-6 | |
| Xtr, Xva, Xte = (Xtr - mu) / sd, (Xva - mu) / sd, (Xte - mu) / sd | |
| context_length = Xtr.shape[1] | |
| print(f"[xLSTM] device={DEVICE} | num_classes={num_classes} | in_features={in_features}") | |
| print(f"[xLSTM] shapes: train={Xtr.shape} val={Xva.shape} test={Xte.shape}") | |
| print(f"[xLSTM] context_length={context_length} | num_blocks={num_blocks}") | |
| print(f"[xLSTM] train classes present: {sorted(set(ytr.tolist()))}") | |
| 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 = SkatingxLSTMClassifier( | |
| in_features, num_classes, | |
| num_blocks=num_blocks, | |
| context_length=context_length, | |
| ).to(DEVICE) | |
| n_params = sum(p.numel() for p in model.parameters()) | |
| print(f"[xLSTM] 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"[xLSTM] 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"[xLSTM] 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"\n[xLSTM] restored best model (val macro-F1 = {best_val_f1:.3f})") | |
| metrics = report_metrics(yte, predict(model, Xte), taxonomy) | |
| ckpt_name = f"model_xlstm{'_coarse' if coarse else ''}.pt" | |
| torch.save({"state_dict": model.state_dict(), "num_classes": num_classes, | |
| "in_features": in_features, "taxonomy": taxonomy, "coarse": coarse, | |
| "num_blocks": num_blocks, "context_length": context_length, | |
| "feature_mean": mu, "feature_std": sd}, | |
| data_dir / ckpt_name) | |
| print(f"\n[xLSTM] saved model -> {data_dir / ckpt_name}") | |
| return metrics | |
| def smoke() -> None: | |
| print("[xLSTM] SMOKE: synthetic (N,128,94) data, labels in TAXONOMY index space") | |
| num_classes = len(labels_mod.TAXONOMY) | |
| rng = np.random.default_rng(0) | |
| X = rng.standard_normal((40, 128, 94)).astype(np.float32) | |
| y = rng.integers(0, num_classes, size=40).astype(np.int64) | |
| assert_label_consistency(num_classes, y) | |
| model = SkatingxLSTMClassifier(94, num_classes, context_length=128).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() | |
| n_params = sum(p.numel() for p in model.parameters()) | |
| print(f"[xLSTM] forward OK: logits {tuple(logits.shape)} | loss {loss.item():.3f} | " | |
| f"backward OK | params {n_params/1e6:.2f}M") | |
| per_step = model.forward_per_step(torch.from_numpy(X[:2]).to(DEVICE)) | |
| assert per_step.shape == (2, 128, num_classes), per_step.shape | |
| print(f"[xLSTM] forward_per_step OK: {tuple(per_step.shape)}") | |
| print("[xLSTM] SMOKE PASSED") | |
| def main() -> int: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--data-dir", type=Path) | |
| parser.add_argument("--coarse", action="store_true") | |
| parser.add_argument("--num-blocks", type=int, default=XLSTM_NUM_BLOCKS) | |
| parser.add_argument("--smoke", action="store_true") | |
| args = parser.parse_args() | |
| if args.smoke: | |
| smoke() | |
| return 0 | |
| if args.data_dir is None or 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, num_blocks=args.num_blocks) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |
Xet Storage Details
- Size:
- 11.1 kB
- Xet hash:
- 8e037ff6728aaa259646a575c2136ef499ac3547d6ba3ae8658c479e9521d900
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.