To generate customized few-shot tasks (e.g., N-way K-shot) from a dataset, follow this pipeline:
- Pre-process input data: Use standard transforms (like
torchvision.transforms) for resizing or tensor conversion. - Load the dataset: Initialize your dataset with the pre-processing transforms.
- Wrap with
MetaDataset: Wrap your dataset in l2l.data.MetaDataset to enable fast indexing of samples. - Define task transforms: Create a list of
learn2learn task transforms to define the task structure (e.g., number of ways, shots, label remapping). - Create a
Taskset: Initialize l2l.data.Taskset with the MetaDataset and your list of transforms. - Sample tasks: Use
.sample() to get a single task or iterate over the Taskset to sample multiple tasks.
Common task transforms include:
NWays(dataset, n): Selects $N$ random classes per task.KShots(dataset, k): Selects $K$ samples per class from the selected $N$ classes.LoadData(dataset): Loads the actual data samples.RemapLabels(dataset): Remaps labels to start from zero.ConsecutiveLabels(dataset): Re-orders samples so they are sorted in consecutive order.RandomClassRotation(dataset, degrees): Randomly rotates vision samples (e.g., [0, 90, 180, 270]).
import learn2learn as l2l
import torchvision as tv
from PIL.Image import LANCZOS
from learn2learn.data.transforms import NWays, KShots, LoadData, RemapLabels, ConsecutiveLabels
from learn2learn.vision.transforms import RandomClassRotation
# 1. Apply transforms on input data
data_transform = tv.transforms.Compose([tv.transforms.Resize((28, 28), interpolation=LANCZOS), tv.transforms.ToTensor()])
# 2. Load the dataset
dataset = l2l.vision.datasets.FullOmniglot(root='~\data', transform=data_transform, download=True)
# 3. Wrap the dataset using MetaDataset for fast indexing
omniglot = l2l.data.MetaDataset(dataset)
# 4. Specify transforms to be used for generating tasks
transforms = [
NWays(omniglot, 5), # N = 5
KShots(omniglot, 1), # K = 1
LoadData(omniglot),
RemapLabels(omniglot),
ConsecutiveLabels(omniglot),
RandomClassRotation(omniglot, [0, 90, 180, 270])
]
# 5. Generate set of tasks
taskset = l2l.data.Taskset(dataset=omniglot, task_transforms=transforms, num_tasks=10)
# Sample a task
X, y = taskset.sample()
print(X.shape)