Codebase for experiments on latent equivariant operators for MNIST classification under geometric transformations.
The model learns latent representations that are either:
no_op: no transformation operator (baseline)fixed_op: fixed cyclic latent operatorlearned_op: learned latent operator with periodicity regularization
Supported transforms:
rotateshift_xshift_yshift_xy
.
├── datasets/
│ ├── mnist_download.py # download + export raw MNIST PNGs
│ ├── mnist_gen.py # generate transformed RGB images
│ ├── make_split_csv.py # build train/val/test CSV split
│ ├── polygon_dataset.py # main training dataset
│ └── compound_dataset.py
├── models/
│ └── linear_classifier.py
├── training/
│ ├── base.py
│ ├── rotate.py
│ ├── xshift.py
│ ├── yshift.py
│ └── xyshift.py
├── mnist_train_cls.py # main training entrypoint
├── knn_xshift.py # x-shift evaluation script
├── knn_yshift.py # y-shift evaluation script
└── README.md
pip install torch torchvision tqdm numpy pandas pillowRun from repository root.
python datasets/mnist_download.pyCreates:
mnist_local/train/<digit>/*.pngmnist_local/test/<digit>/*.png
python datasets/mnist_gen.pyCreates:
mnist_combined/train/<digit>/*.pngmnist_combined/test/<digit>/*.png
python datasets/make_split_csv.pyCreates:
mnist_dataset_split.csv
Expected CSV columns:
file_path(relative tomnist_local, e.g.train/3/123.png)class(integer)split(train,val,test)
Main command:
python mnist_train_cls.py \
--transform {rotate|shift_x|shift_y|shift_xy} \
[--modes no_op fixed_op learned_op] \
[--epochs 20] [--lr 1e-3] [--batch-size 512] \
[--device cuda|cpu] [--num-workers N] [--prefetch-factor 8] \
[--checkpoint-root checkpoints] \
[--csv mnist_dataset_split.csv] \
[--trained-degrees 0,36,72]Examples:
python mnist_train_cls.py --transform rotate
python mnist_train_cls.py --transform shift_x --modes fixed_op
python mnist_train_cls.py --transform shift_y --epochs 40 --lr 5e-4 --batch-size 256Useful args:
--csvpath to split file (default:mnist_dataset_split.csv)--epochs(default:20)--lr(default:1e-3)--batch-size(default:512)--num-workers(default:os.cpu_count())--prefetch-factor(default:8)--no-pin-memory--checkpoint-root(default:checkpoints)--modessubset ofno_op fixed_op learned_op--trained-degreescomma list override (example:0,36,72)--device(default: autocudaif available, elsecpu)
Show full CLI:
python mnist_train_cls.py --help- Checkpoints are saved under
checkpoints/(or--checkpoint-root). - Naming follows
mnist_<transform>_cls_<mode>/best.pthandlatest.pth.
Run after training checkpoints exist:
python knn_xshift.py
python knn_yshift.py- Current
PolygonDatasetfilters to classes< 9(digits0-8). shift_xyuses the transform-specific loop intraining/xyshift.py.
@inproceedings{
dinh2026latent,
title={Latent Equivariant Operators for Robust Object Recognition: Promises and Challenges},
author={Minh T. Dinh and Stephane Deny},
booktitle={ICLR 2026 Workshop on Geometry-grounded Representation Learning and Generative Modeling},
year={2026},
url={https://openreview.net/forum?id=81gVwLVcXQ}
}