Instructions to use ERISLab/TGDA with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use ERISLab/TGDA with timm:
import timm model = timm.create_model("hf_hub:ERISLab/TGDA", pretrained=True) - Notebooks
- Google Colab
- Kaggle
|
Download README.md from ERISLab/TGDA: direct link, hf CLI and curl.
- Browser
- Download file 8.96 kB
-
https://huggingface.co/ERISLab/TGDA/resolve/main/README.md
- Command line
-
hf download hf://ERISLab/TGDA/README.md
-
curl -L -o README.md https://huggingface.co/ERISLab/TGDA/resolve/main/README.md
8.96 kB
| pipeline_tag: image-classification | |
| library_name: pytorch | |
| tags: | |
| - fine-grained-image-recognition | |
| - image-classification | |
| - teacher-guided-data-augmentation | |
| - timm | |
| # TGDA: fine-grained image recognition from scratch with teacher-guided data augmentation | |
| These are the checkpoints behind *Fine-Grained Image Recognition from Scratch with Teacher-Guided Data Augmentation* ([arXiv:2507.12157](https://arxiv.org/abs/2507.12157)). Students trained from scratch learn from a fine-grained teacher (a ResNet-101 trained with counterfactual attention learning, CAL), which also guides data-aware augmentation of the student's inputs. Code: [arkel23/TGDA](https://github.com/arkel23/TGDA). | |
| 249 checkpoints, one per configuration, each the last epoch of one training seed. Each file | |
| is a `torch.save` dict with `config` (the full training configuration), `model` (the state dict), | |
| `accuracy` and `epoch`, with no optimizer state. File names are the runs' experiment-log names, | |
| ending in the serial. Load them with [fgir-zoo](https://github.com/arkel23/fgir-zoo). The | |
| [collection](https://huggingface.co/collections/ERISLab/tgda-fgir-from-scratch-with-teacher-guided-augmentation-6ab32cd151f4fc1063b55f64) groups this repo with the paper. | |
| ## Layout | |
| | Folder | Serials | What | Files | | |
| |---|---|---|---| | |
| | `sota_448` | 301-304 | Main state-of-the-art table: ResNet-18/34/50/101 students, TGDA with TrivialAugment, label smoothing and stochastic depth | 12 | | |
| | `lrnet_128` | 201-204 | Main LRNet table: low-resolution students at 128 px from a 448 px teacher | 24 | | |
| | `vitfs_448` | 100, 101 | Efficient ViTs and the ViTFS variants, teacher and student at 448 px | 54 | | |
| | `vitfs_224_t448` | 102 | Teacher at 448 px, student at 224 px | 15 | | |
| | `vitfs_224` | 103 | Teacher and student at 224 px | 9 | | |
| | `vitfs_128` | 109 | ViTFS tiny/micro/nano students at 128 px from a 448 px teacher | 8 | | |
| | `vit_ce_448` | 107 | The serial-101 models trained with cross-entropy only (no teacher) | 35 | | |
| | `vitfs_transfer` | 104-106 | A TGDA-trained student transferred to two new datasets with cross-entropy | 8 | | |
| | `ablations_128` | 212, 214, 221-223, 241-248, 251, 261-263, 271-278, 281 | Ablations: TGDA without extra regularization, one- vs two-stage training, augmentation type, teacher type, transfer | 73 | | |
| | `scratch_vs_pt_448` | 320, 321, 331 | 200 epochs from scratch with cross-entropy or KD, and ImageNet-pretrained with KD | 9 | | |
| | `sweeps_128` | 0, 1 | KD weight and temperature sweeps on CUB | 2 | | |
| `manifest.csv` lists every file with its dataset, student, teacher, serial, seed, image size, | |
| class count, accuracy, SHA-256 and size. | |
| ## Load a checkpoint and classify an image | |
| ```python | |
| import torch | |
| from PIL import Image | |
| from torchvision import transforms | |
| from fgir_zoo import tgda | |
| model = tgda.create_model('sota_448/cub_resnet18_resnet101_cal_tgda_301') | |
| cfg = model.config | |
| tf = transforms.Compose([ | |
| transforms.Resize((cfg.test_resize_size, cfg.test_resize_size), | |
| interpolation=transforms.InterpolationMode.BICUBIC), | |
| transforms.CenterCrop(cfg.student_image_size or cfg.image_size), | |
| transforms.ToTensor(), | |
| transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), | |
| ]) | |
| # a CUB-200-2011 test image, class index 50 (051.Horned_Grebe) | |
| x = tf(Image.open('Horned_Grebe_0050_34561.jpg').convert('RGB')).unsqueeze(0) | |
| with torch.no_grad(): | |
| logits = model(x) | |
| print(logits.argmax(-1).item(), logits.softmax(-1).max().item()) # 50 0.9855 | |
| ``` | |
| ## Accuracy of the released checkpoints | |
| Top-1 accuracy (%) stored in each file, on the dataset's test split at the end of training. | |
| Each value is one seed; the paper reports its own aggregates, so its tables can differ slightly. | |
| ### `sota_448` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `resnet101 <- resnet101` | 94.81 | 95.19 | 89.25 | | |
| | `resnet18 <- resnet101` | 93.73 | 94.86 | 87.49 | | |
| | `resnet34 <- resnet101` | 94.24 | 95.04 | 88.40 | | |
| | `resnet50 <- resnet101` | 94.78 | 95.16 | 88.95 | | |
| ### `lrnet_128` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `lrnet14 <- resnet101` | 89.02 | 92.74 | 80.48 | | |
| | `lrnet18 <- resnet101` | 90.67 | 93.21 | 82.48 | | |
| | `lrnet26 <- resnet101` | 91.30 | 93.57 | 82.65 | | |
| | `lrnet50 <- resnet101` | 91.81 | 93.74 | 84.12 | | |
| | `resnet101 <- resnet101` | 88.57 | 92.19 | 81.48 | | |
| | `resnet18 <- resnet101` | 86.95 | 91.18 | 77.68 | | |
| | `resnet34 <- resnet101` | 85.45 | 91.03 | 78.81 | | |
| | `resnet50 <- resnet101` | 90.19 | 92.96 | 82.74 | | |
| ### `vitfs_448` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `efficientvit_m2 <- resnet101` | 92.50 | 94.11 | 85.42 | | |
| | `efficientvit_m3 <- resnet101` | 92.80 | 94.50 | 85.48 | | |
| | `lcnet_035 <- resnet101` | 91.42 | 91.12 | 78.43 | | |
| | `levit_128s <- resnet101` | 88.66 | 91.98 | 81.24 | | |
| | `mobilevit_s <- resnet101` | 92.68 | 93.79 | 83.48 | | |
| | `mobilevitv2_100 <- resnet101` | 94.36 | 94.57 | 87.59 | | |
| | `mobilevitv2_125 <- resnet101` | 93.64 | 94.63 | 87.35 | | |
| | `nest_tiny <- resnet101` | 89.89 | 93.25 | 85.48 | | |
| | `pvt_v2_b0 <- resnet101` | 91.84 | 93.79 | 86.93 | | |
| | `pvt_v2_b1 <- resnet101` | 92.53 | 94.29 | 86.85 | | |
| | `swin_tiny_patch4_window7_224 <- resnet101` | 89.62 | 93.35 | 82.40 | | |
| | `tiny_vit_5m_224 <- resnet101` | 93.37 | 94.43 | 88.45 | | |
| | `vit_t16 <- resnet101` | 61.00 | | 58.70 | | |
| | `vitfs_tiny_patch16_clc_gap_reg4_dinov2_bn <- resnet101` | 90.58 | 94.01 | 85.47 | | |
| | `vitfs_tiny_patch16_gap_224 <- resnet101` | 85.30 | | 82.72 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init <- resnet101` | 92.20 | 94.48 | 86.26 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_init <- resnet101` | 89.32 | 93.07 | 83.72 | | |
| | `vitfs_tiny_patch16_kvm_gap_reg4_dinov2_bn <- resnet101` | | 94.25 | 83.91 | | |
| ### `vitfs_224_t448` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `vit_t16 <- resnet101` | 55.06 | 51.29 | 49.48 | | |
| | `vitfs_micro_patch16_gap_reg4_dinov2_bn <- resnet101` | 88.81 | 92.41 | 78.25 | | |
| | `vitfs_nano_patch16_gap_reg4_dinov2_bn <- resnet101` | 85.06 | 90.44 | 70.88 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init <- resnet101` | 90.49 | 93.91 | 81.20 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_init <- resnet101` | 86.35 | 92.71 | 77.94 | | |
| ### `vitfs_224` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `vitfs_tiny_patch16_gap_224 <- resnet101` | 52.63 | 86.16 | 75.61 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init <- resnet101` | 60.13 | 92.64 | 79.89 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_init <- resnet101` | 47.76 | 86.57 | 76.46 | | |
| ### `vitfs_128` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `vitfs_micro_patch16_gap_reg4_dinov2_bn <- resnet101` | 79.21 | 88.45 | 66.40 | | |
| | `vitfs_nano_patch16_gap_reg4_dinov2_bn <- resnet101` | 77.47 | 83.97 | 61.32 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init <- resnet101` | 81.10 | 90.24 | | | |
| ### `vit_ce_448` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `efficientvit_m2` | 87.49 | 83.75 | 73.25 | | |
| | `efficientvit_m3` | 87.46 | 80.19 | 72.82 | | |
| | `levit_128s` | 82.27 | 87.07 | 69.59 | | |
| | `mobilevit_s` | 88.36 | 91.20 | 72.01 | | |
| | `mobilevitv2_100` | 89.11 | 90.57 | 78.60 | | |
| | `mobilevitv2_125` | 87.85 | 90.34 | 81.50 | | |
| | `nest_tiny` | 82.99 | 81.23 | 71.13 | | |
| | `pvt_v2_b0` | 86.41 | 89.76 | 75.09 | | |
| | `pvt_v2_b1` | 85.63 | | | | |
| | `swin_tiny_patch4_window7_224` | | 63.26 | | | |
| | `vit_t16` | 45.60 | 26.38 | 30.81 | | |
| | `vitfs_tiny_patch16_gap_224` | | | 0.78 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init` | | 88.41 | 69.92 | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_init` | 76.84 | 81.35 | 60.30 | | |
| ### `vitfs_transfer` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `vit_t16` | 50.23 | 39.12 | | | |
| | `vitfs_tiny_patch16_gap_reg4_dinov2_bn_init` | 87.19 | 90.70 | 82.52 | | |
| ### `ablations_128` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `lrnet14` | 86.05 | 89.34 | 65.10 | | |
| | `lrnet14 <- resnet101` | 85.30 | 93.06 | 81.79 | | |
| | `resnet101d <- resnet101` | 88.18 | 92.54 | | | |
| | `resnet18` | 73.72 | 82.13 | 56.54 | | |
| | `resnet18 <- resnet101` | 82.03 | 88.22 | 75.34 | | |
| | `resnet34d <- resnet101` | 87.46 | 91.64 | 78.65 | | |
| ### `scratch_vs_pt_448` | |
| | Model | aircraft | cars | cub | | |
| |---|---|---|---| | |
| | `resnet101` | 83.08 | 84.36 | 66.41 | | |
| | `resnet101 <- resnet101` | 93.91 | 94.60 | 87.63 | | |
| ### `sweeps_128` | |
| | Model | cub | | |
| |---|---| | |
| | `resnet18 <- resnet101` | 80.34 | | |
| ## Requirements | |
| - `fgir-zoo` (`pip install git+https://github.com/arkel23/fgir-zoo.git`), which pins | |
| `timm==0.9.12` | |
| - `torch` (checked with 2.5.1) | |
| ## Citation | |
| ```bibtex | |
| @misc{rios_tgda_2025, | |
| title = {Fine-Grained Image Recognition from Scratch with Teacher-Guided Data Augmentation}, | |
| author = {Rios, Edwin Arkel and Mikael, Fernando and Gosal, Oswin and Oyerinde, Femiloye and | |
| Liang, Hao-Chun and Lai, Bo-Cheng and Hu, Min-Chun}, | |
| year = {2025}, | |
| eprint = {2507.12157}, | |
| archivePrefix = {arXiv}, | |
| primaryClass = {cs.CV}, | |
| doi = {10.48550/arXiv.2507.12157}, | |
| url = {https://arxiv.org/abs/2507.12157} | |
| } | |
| ``` | |