"""Unique-group batches for supervised contrastive learning.""" from __future__ import annotations import math import random from collections import defaultdict from collections.abc import Iterator, Sequence from torch.utils.data import Sampler class UniqueGroupBatchSampler(Sampler[list[int]]): def __init__( self, group_ids: Sequence[str], groups_per_batch: int, examples_per_group: int = 2, seed: int = 0, ) -> None: if groups_per_batch < 2 or examples_per_group < 2: raise ValueError("A contrastive batch needs at least two groups and two examples per group") grouped: dict[str, list[int]] = defaultdict(list) for index, group_id in enumerate(group_ids): grouped[str(group_id)].append(index) self.grouped = {key: value for key, value in grouped.items() if len(value) >= examples_per_group} if len(self.grouped) < 2: raise ValueError("At least two groups have enough examples") self.groups_per_batch = groups_per_batch self.examples_per_group = examples_per_group self.seed = seed self.epoch = 0 def set_epoch(self, epoch: int) -> None: self.epoch = int(epoch) def __len__(self) -> int: return math.ceil(len(self.grouped) / self.groups_per_batch) def __iter__(self) -> Iterator[list[int]]: random_state = random.Random(self.seed + self.epoch) groups = sorted(self.grouped) random_state.shuffle(groups) for start in range(0, len(groups), self.groups_per_batch): selected = groups[start : start + self.groups_per_batch] if len(selected) < 2: selected.extend(groups[: 2 - len(selected)]) batch: list[int] = [] for group_id in selected: batch.extend(random_state.sample(self.grouped[group_id], self.examples_per_group)) yield batch