| """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 |
|
|