File size: 1,947 Bytes
c87881a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 | """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
|