Buckets:

Rishik001's picture
download
raw
11.1 kB
"""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.