Rishik001's picture
Add checkpoints and model card
fb47c95 verified
|
Raw
History Blame Contribute Delete
6.27 kB
---
license: other
license_name: proprietary-internal
license_link: LICENSE
library_name: pytorch
tags:
- figure-skating
- action-recognition
- skeleton-based
- time-series-classification
- transformer
- pytorch
---
# Figure Skating Transformer Action Classifier
A Conv1D backbone + Transformer-encoder baseline for classifying figure-skating actions (jumps,
spins, sequences) from pose-skeleton time series. Sinusoidal positional encoding + a real
3-layer `nn.TransformerEncoder` (self-attn + FFN per layer) on top of the same Conv1D residual
backbone used across this model family.
This repo holds **two checkpoints** trained on the same architecture and data split, differing
only in label granularity:
| file | label space | classes |
|---|---|---|
| `model_transformer_fine.pt` | fine-grained | 28 (jump rotation-count preserved, e.g. `3Lutz`, `2Axel`) |
| `model_transformer_coarse.pt` | coarse | 11 (jump rotation-count collapsed, e.g. `Lutz`, `Axel`) |
Not `transformers`-compatible in the HF sense — this is a plain PyTorch `nn.Module` +
`state_dict` checkpoint (custom architecture, not a Hub `AutoModel`). Model source code
(`SkatingActionClassifier` in the project's `model_transformers.py`, which depends on
`model.py`) is **not included in this repo** — you'll need the project source to reconstruct
the class before loading the state dict.
## How to load
```python
from huggingface_hub import hf_hub_download
import torch
repo_id = "Mercity/figure-skating-transformer"
ckpt_path = hf_hub_download(repo_id, "model_transformer_fine.pt")
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
# reconstruct with the project's SkatingActionClassifier (model_transformers.py):
# model = SkatingActionClassifier(ckpt["in_features"], ckpt["num_classes"])
# model.load_state_dict(ckpt["state_dict"])
# model.eval()
# ckpt also carries the standardization stats used at train time and the taxonomy:
mu, sd = ckpt["feature_mean"], ckpt["feature_std"] # normalize inputs with (x - mu) / sd
taxonomy = ckpt["taxonomy"] # {0: "3Toeloop", 1: "3Loop", ...}
```
Swap `model_transformer_fine.pt` for `model_transformer_coarse.pt` to load the 11-class variant.
## Config
| | |
|---|---|
| Backbone | Conv1D stem(128) → ConvBlock(192,k3) → ConvBlock(256,k3) → ConvBlock(384,k5) |
| Temporal block | Sinusoidal positional encoding + `nn.TransformerEncoder`, 3 layers, d_model=384, nhead=8, dim_feedforward=1536 (4×d_model), dropout=0.1, GELU, post-norm |
| Head | Dense(384→1024,p0.5) → Dense(1024→512,p0.4) → Dense(512→256,p0.3) → LayerNorm(256) → Linear(256→num_classes) |
| Optimizer | Adam, lr=5e-4, weight_decay=1e-4 |
| Batch size | 64 |
| Grad clip | 1.0 |
| Class weight cap | 5.0 |
| Early stop | patience 20 on val macro-F1, 100-epoch cap |
| Seed | 42 |
| Params | 7.31M (both) |
## Data
Trained on skeleton-extracted figure-skating clips (YOLO11n-pose, COCO-17 joints, imgsz=384)
from an internal Mercity dataset. Stratified train/val/test split, seed 42.
**Not included in this repo** (private internal dataset).
## Results (test split)
### Fine-grained (28-class)
Early-stopped epoch 42 (best val macro-F1 0.572)
**Overall:** accuracy 0.8348 · precision (macro/weighted) 0.4511 / 0.8364 · recall 0.5129 / 0.8348 · **F1 0.4642 / 0.8287**
| class | precision | recall | f1 | support |
|---|---|---|---|---|
| 3Toeloop | 0.514 | 0.500 | 0.507 | 38 |
| 3Loop | 0.607 | 0.810 | 0.694 | 21 |
| 2Axel | 0.919 | 0.791 | 0.850 | 43 |
| CamelSpin | 0.991 | 0.912 | 0.950 | 125 |
| SitSpin | 0.969 | 0.955 | 0.962 | 132 |
| UprightSpin | 0.934 | 0.917 | 0.925 | 108 |
| 2Salchow | 0.286 | 0.400 | 0.333 | 5 |
| 2Toeloop | 0.658 | 0.781 | 0.714 | 32 |
| 3Salchow | 0.381 | 0.381 | 0.381 | 21 |
| 3Axel | 0.867 | 0.867 | 0.867 | 15 |
| 3Flip | 0.538 | 0.389 | 0.452 | 36 |
| 3Lutz | 0.667 | 0.244 | 0.357 | 41 |
| NoBasic | 0.757 | 0.903 | 0.824 | 31 |
| 2Lutz | 0.214 | 0.429 | 0.286 | 7 |
| 4Salchow | 0.200 | 0.333 | 0.250 | 3 |
| 4Flip | 0.000 | 0.000 | 0.000 | 1 |
| 4Toeloop | 0.200 | 0.167 | 0.182 | 6 |
| 4Lutz | 0.000 | 0.000 | 0.000 | 1 |
| 4Loop | 0.000 | 0.000 | 0.000 | 1 |
| 2Flip | 0.333 | 0.500 | 0.400 | 8 |
| 2Loop | 0.500 | 0.750 | 0.600 | 8 |
| 1Axel | 0.333 | 0.333 | 0.333 | 3 |
| 1Loop | 0.500 | 1.000 | 0.667 | 1 |
| 1Salchow | 0.000 | 0.000 | 0.000 | 1 |
| 1Toeloop | 0.333 | 1.000 | 0.500 | 1 |
| 1Flip | 0.000 | 0.000 | 0.000 | 2 |
| 1Lutz | 0.000 | 0.000 | 0.000 | 1 |
| Sequence | 0.930 | 1.000 | 0.964 | 307 |
| **macro avg** | **0.451** | **0.513** | **0.464** | 999 |
| **weighted avg** | **0.836** | **0.835** | **0.829** | 999 |
### Coarse (11-class)
Early-stopped epoch 61 (best val macro-F1 0.860)
**Overall:** accuracy 0.8989 · precision (macro/weighted) 0.8390 / 0.8989 · recall 0.8158 / 0.8989 · **F1 0.8233 / 0.8970**
| class | precision | recall | f1 | support |
|---|---|---|---|---|
| Axel | 0.879 | 0.951 | 0.913 | 61 |
| Toeloop | 0.728 | 0.870 | 0.793 | 77 |
| Salchow | 0.737 | 0.467 | 0.571 | 30 |
| Loop | 0.793 | 0.742 | 0.767 | 31 |
| Flip | 0.714 | 0.638 | 0.674 | 47 |
| Lutz | 0.750 | 0.720 | 0.735 | 50 |
| CamelSpin | 0.976 | 0.976 | 0.976 | 125 |
| SitSpin | 0.976 | 0.917 | 0.945 | 132 |
| UprightSpin | 0.871 | 0.935 | 0.902 | 108 |
| Sequence | 0.977 | 0.984 | 0.981 | 307 |
| NoBasic | 0.828 | 0.774 | 0.800 | 31 |
| **macro avg** | **0.839** | **0.816** | **0.823** | 999 |
| **weighted avg** | **0.899** | **0.899** | **0.897** | 999 |
## Notable failure modes (from confusion matrix analysis, coarse label space)
- Best `Axel` (95.1%) and `Toeloop` (87.0%) recall of all architectures tested, but by far the
**worst `Salchow` recall (46.7%)** — heavily biased toward calling Salchow "Toeloop"
(11/30 test samples, the single worst confusion pair measured across every model tried).
- `Flip ↔ Lutz` confusion present but less severe than in other architectures.
- Fine-grained macro-F1 is dragged down almost entirely by 1-2-sample tail classes
(`4Flip`, `4Lutz`, `4Loop`, `1Salchow`, `1Toeloop`, `1Flip`, `1Lutz`) — a data scarcity issue,
not an architecture weakness.
## License
Internal/proprietary — not licensed for external use. Contact the Mercity team for access terms.