Download source/src/bgc_retrieval/sampling.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 1.95 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/sampling.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/src/bgc_retrieval/sampling.py
-
curl -L -o sampling.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/sampling.py
1.95 kB
| """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 | |