No description
Find a file
Prashant Pandey 78f0ac6191 h4
2026-08-13 20:40:49 +05:30
src feat: add empty __init__.py to src package 2026-08-13 19:06:24 +05:30
README.md h4 2026-08-13 20:40:49 +05:30
requirements.txt packages 2026-08-13 20:36:51 +05:30

ViT encoder using LeJEPA (Invariance + SIGReg loss).

Architecture

  • ViT-B/16 (92.5M params; ViT-S/Tiny)
  • 3-layer MLP with BatchNorm (768 → 512 → 2048 → 512)
  • 2 global (224², scale 0.3–1.0) + 6 local (96², scale 0.05–0.3) per image
  • Invariance + λ·SIGReg (Epps-Pulley statistic on random projections)

Install

python -m venv .venv && source .venv/bin/activate
pip install -r requirements.txt
pip install "jax[cuda12]" -f https://jax.github.io/jax/releases.html

Data

expected layout: data/data_name/ train/ valid/ train.csv valid.csv

Train

python -m src.main data.batch_size=128 train.epochs=100 

Evaluate

python -m src.main --eval-only --checkpoint checkpoints/step_0000010

Note

All hyperparameters in src/config.py (DataConfig, ModelConfig, TrainConfig): override via CLI: python -m src.main data.batch_size=64 model.embed_dim=384 train.epochs=30