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