Optimization as a Model for Few-shot Learning
Pytorch implementation of Optimization as a Model for Few-shot Learning in ICLR 2017 (Oral)
Prerequisites
Preparation
- data/
- miniImagenet/
- train/
- n01532829/
- n0153282900000005.jpg
- ...
- n01558993/
- ...
- val/
- n01855672/
- ...
- test/
- ...
- main.py
- ...
It'd be set if you download and extract Mini-Imagenet from the link above
Check out scripts/train_5s_5c.sh, make sure --data-root is properly set
For 5-shot, 5-class training, run
bash scripts/train_5s_5c.sh
Hyper-parameters are referred to the author's repo.
For 5-shot, 5-class evaluation, run (remember to change --resume and --seed arguments)
bash scripts/eval_5s_5c.sh
Notes
Results (This repo is developed following the pytorch reproducibility guideline):
The results I get from directly running the author's repo can be found here, I have slightly better performance (~5%) but neither results match the number in the paper (60%) (Discussion and help are welcome!).
Training with the default settings takes ~2.5 hours on a single Titan Xp while occupying ~2GB GPU memory.
The implementation replicates two learners similar to the author's repo:
learner_w_grad functions as a regular model, get gradients and loss as inputs to meta learner.
learner_wo_grad constructs the graph for meta learner:
All the parameters in learner_wo_grad are replaced by cI output by meta learner.
nn.Parameters in this model are casted to torch.Tensor to connect the graph to meta learner.
Several ways to copy a parameters from meta learner to learner depends on the scenario:
copy_flat_params: we only need the parameter values and keep the original grad_fn.
transfer_params: we want the values as well as the grad_fn (from cI to learner_wo_grad).
.data.copy_ v.s. clone() -> the latter retains all the properties of a tensor including grad_fn.
To maintain the batch statistics, load_state_dict is used (from learner_w_grad to learner_wo_grad).
CloserLookFewShot (Data loader)
pytorch-meta-optimizer (Casting nn.Parameters to torch.Tensor inspired from here)
meta-learning-lstm (Author's repo in Lua Torch)