Level 1 Weighted batches from a registry
A training job mixes several named datasets. You get a registry object with one method:
registry.get_iterator(name: str, offset: int = 0) -> Iterator
It returns a fresh iterator over dataset name that starts at example number offset (0-based) and stops at the end of the dataset. Opening an iterator at any offset is cheap.
Implement DataBatcher(registry, weights: dict[str, int], batch_size: int) with next_batch() -> list.
- Each batch holds
batch_sizeexamples. For this levelbatch_sizeis a multiple ofS = sum(weights.values()), so withk = batch_size // Sdatasetnamecontributes exactlyweights[name] * kexamples. - Order convention: the batch lists the datasets' contributions one after another, in the order the keys appear in
weights. Within a dataset, examples come in the order the iterator yields them. - Each dataset is read sequentially, continuing across batches where the previous batch stopped.
- When a dataset runs out (even in the middle of its share), start it over from example 0 and keep going.
- A weight may be
0; that dataset is never read. Every dataset with a positive weight has at least one example. - Read lazily: pull only the examples you are about to return.
# datasets: "a" = a0..a4 (5 examples), "b" = b0..b1 (2 examples)
batcher = DataBatcher(registry, {"a": 2, "b": 1}, 6)
batcher.next_batch() # ['a0', 'a1', 'a2', 'a3', 'b0', 'b1']
batcher.next_batch() # ['a4', 'a0', 'a1', 'a2', 'b0', 'b1']