implement global blocks per epoch

This commit is contained in:
Ayush Chaurasia
2026-08-12 15:20:27 +05:30
parent 51bc3f824b
commit 48d6f504c2
2 changed files with 313 additions and 54 deletions
+233 -40
View File
@@ -29,9 +29,10 @@ from collections import deque
from concurrent.futures import ThreadPoolExecutor
from copy import deepcopy
from multiprocessing import RawArray
from typing import Any, Callable, cast, Iterator, Optional, Union
from typing import Any, Callable, cast, Iterator, Literal, Optional, Union
import pyarrow as pa
import pyarrow.compute as pc
import torch
from torch.utils.data import IterableDataset, get_worker_info
@@ -150,8 +151,6 @@ class StreamingDataset(IterableDataset):
one block per owned split.
* ``eos_id``, ``pad_id``, and ``columns`` naming a single
list-typed column are required; incompatible with ``transform``.
Note: When packing is sharded across ranks or DataLoader workers, a warning
is emitted because padding and checkpoint state are local to each iterator.
eos_id:
Separator token id between packed documents. Required with
``pack_sequences``, ignored otherwise.
@@ -159,6 +158,19 @@ class StreamingDataset(IterableDataset):
Padding token id used to complete blocks when a split runs out of
real tokens mid-cycle or at epoch end. Required with
``pack_sequences``, ignored otherwise.
blocks_per_epoch:
Total number of packed blocks emitted globally per epoch. Required with
``pack_sequences``. An integer must be divisible by ``num_splits``.
Every logical split emits exactly ``blocks_per_epoch / num_splits``
blocks: exhausted splits emit padding, while tokens beyond the budget
are left out of the epoch. This fixed per-split budget keeps packed
iteration and checkpoints independent of rank and worker topology.
Pass ``"auto"`` to estimate the budget by reading only the token column
for approximately 1% of the rows, with a target cap of 100,000 rows and
at least one row from every logical split. Sampling is batched and evenly
distributed across logical splits. This is a convenience for smaller
workloads. For large-scale training, materialize a ``token_count`` column,
calculate an explicit block budget once, and pass that integer instead.
on_transform_error:
What to do when the transform raises an exception:
@@ -230,6 +242,7 @@ class StreamingDataset(IterableDataset):
pack_sequences: Optional[int] = None,
eos_id: Optional[int] = None,
pad_id: Optional[int] = None,
blocks_per_epoch: Optional[Union[int, Literal["auto"]]] = None,
on_transform_error: Union[str, Callable[[Exception], bool]] = "raise",
connection_factory: Optional[Callable[[str], Any]] = None,
worker_info_override=None,
@@ -253,6 +266,24 @@ class StreamingDataset(IterableDataset):
raise ValueError("eos_id is required when pack_sequences is set")
if pad_id is None:
raise ValueError("pad_id is required when pack_sequences is set")
if blocks_per_epoch is None:
raise ValueError(
"blocks_per_epoch is required when pack_sequences is set"
)
if blocks_per_epoch != "auto":
if not isinstance(blocks_per_epoch, int) or isinstance(
blocks_per_epoch, bool
):
raise ValueError(
"blocks_per_epoch must be a positive integer or 'auto'"
)
if blocks_per_epoch <= 0:
raise ValueError("blocks_per_epoch must be greater than 0")
if blocks_per_epoch % num_splits != 0:
raise ValueError(
f"blocks_per_epoch ({blocks_per_epoch}) must be divisible by "
f"num_splits ({num_splits})"
)
if transform is not None:
raise ValueError("transform cannot be combined with pack_sequences")
if columns is None or len(columns) != 1:
@@ -270,6 +301,8 @@ class StreamingDataset(IterableDataset):
f"pack_sequences requires a list-typed token column; "
f"{columns[0]} has type {field.type}"
)
elif blocks_per_epoch is not None:
raise ValueError("blocks_per_epoch requires pack_sequences")
if on_transform_error not in ("raise", "skip", "warn") and not callable(
on_transform_error
):
@@ -295,6 +328,7 @@ class StreamingDataset(IterableDataset):
self._pack_sequences = pack_sequences
self._eos_id = eos_id
self._pad_id = pad_id
self._blocks_per_epoch = blocks_per_epoch
self._on_transform_error = on_transform_error
self._connection_factory = connection_factory
self._worker_info_override = worker_info_override
@@ -302,6 +336,7 @@ class StreamingDataset(IterableDataset):
# Packing resume state: documents consumed and partial-block token buffers.
self._pack_consumed: list[int] = [0] * num_splits
self._pack_buffers: dict[int, dict[str, list[int]]] = {}
self._pack_blocks_emitted: list[int] = [0] * num_splits
# Live references to pipeline state, set only while __iter__ is running
# in the same process. Used by the observability properties when the
@@ -352,6 +387,9 @@ class StreamingDataset(IterableDataset):
else:
self._perm_table = builder.split_sequential(fixed=num_splits).execute()
if self._blocks_per_epoch == "auto":
self._blocks_per_epoch = self._estimate_blocks_per_epoch()
# Contiguous block of global split indices assigned to this rank.
splits_per_rank = num_splits // world_size
rank_start = rank * splits_per_rank
@@ -359,6 +397,76 @@ class StreamingDataset(IterableDataset):
range(rank_start, rank_start + splits_per_rank)
)
def _estimate_blocks_per_epoch(self) -> int:
"""Estimate a fixed packed-block budget from a bounded token sample."""
if self._pack_sequences is None or not self._columns:
raise RuntimeError(
"packing must be configured before estimating its budget"
)
pack_len = self._pack_sequences
token_column = self._columns[0]
sample_cap_per_split = max(1, 100_000 // self._num_splits)
sampled_tokens = 0
total_sampled = 0
total_rows = 0
warnings.warn(
"blocks_per_epoch='auto' is a convenience estimate that samples full "
"token lists and repeats for every process constructing the dataset. It "
"is recommended to materialize a token_count column, calculate an "
"explicit blocks_per_epoch once, and pass that integer instead.",
# TODO: Link to docs example that shows how to do this
)
for split in range(self._num_splits):
permutation = Permutation.from_tables(
self._table, self._perm_table, split=split
)
permutation = permutation.select_columns([token_column])
permutation = permutation.with_transform(Transforms.arrow2arrow)
split_rows = permutation.num_rows
if split_rows == 0:
raise ValueError(
"blocks_per_epoch='auto' cannot estimate an empty dataset"
)
# Roughly 1% per logical split, with at least one row from each
# split and a global cap of approximately 100,000 rows.
sample_rows = min(
split_rows,
max(1, min((split_rows + 99) // 100, sample_cap_per_split)),
)
sample_batch_size = max(1, self._read_batch_size)
for start in range(0, sample_rows, sample_batch_size):
stop = min(start + sample_batch_size, sample_rows)
# approx mid of each bucket
batch_offsets = [
((2 * index + 1) * split_rows) // (2 * sample_rows)
for index in range(start, stop)
]
batch = permutation.__getitems__(batch_offsets)
lengths = pc.list_value_length(batch.column(0))
if lengths.null_count:
raise ValueError("pack_sequences does not support null token lists")
sampled_tokens += int(pc.sum(lengths).as_py())
total_sampled += sample_rows
total_rows += split_rows
# Pool all samples into one global average. Each document contributes
# one EOS token. Round down so estimation
# favors leaving a short tail unused over emitting all-padding blocks.
estimated_tokens = (
(sampled_tokens + total_sampled) * total_rows // total_sampled
)
estimated_blocks = estimated_tokens // pack_len
blocks_per_epoch = max(
self._num_splits,
estimated_blocks - estimated_blocks % self._num_splits,
)
return blocks_per_epoch
def _resolve_my_splits(self) -> list[int]:
"""Return the split indices this instance should read in __iter__."""
torch_worker_info = get_worker_info()
@@ -398,13 +506,6 @@ class StreamingDataset(IterableDataset):
my_splits = self._resolve_my_splits()
if not my_splits:
return
if self._pack_sequences is not None and len(my_splits) < self._num_splits:
warnings.warn(
"Sequence-packing padding is local to each rank or DataLoader "
"worker. Sharded iterators can yield different numbers of packed "
"blocks. Packed checkpoints cannot safely resume with a different "
"world_size or number of DataLoader workers."
)
# Set identity transform on each Permutation so __getitems__ returns
# the raw RecordBatch. Stage 2 applies the real transform.
@@ -416,7 +517,7 @@ class StreamingDataset(IterableDataset):
)
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(lambda batch: batch)
perm = perm.with_transform(Transforms.arrow2arrow)
# Packing tracks documents consumed per split. Row mode tracks
# absolute permutation positions so transform failures can skip
# rows without making resume repeat them.
@@ -618,8 +719,14 @@ class StreamingDataset(IterableDataset):
pack_len = cast(int, self._pack_sequences)
eos_id = cast(int, self._eos_id)
pad_id = cast(int, self._pad_id)
blocks_per_split = (
cast(int, self._blocks_per_epoch) // self._num_splits
if self._pack_sequences is not None
else 0
)
pack_consumed = list(self._pack_consumed)
pack_buffers = deepcopy(self._pack_buffers)
pack_blocks_emitted = list(self._pack_blocks_emitted)
def _pack_buf(i: int) -> dict[str, list[int]]:
return pack_buffers.setdefault(my_splits[i], {"tokens": [], "starts": []})
@@ -661,6 +768,7 @@ class StreamingDataset(IterableDataset):
def _commit_pack_state() -> None:
self._pack_consumed = list(pack_consumed)
self._pack_buffers = deepcopy(pack_buffers)
self._pack_blocks_emitted = list(pack_blocks_emitted)
# ── Main loop ─────────────────────────────────────────────────────────
@@ -691,24 +799,31 @@ class StreamingDataset(IterableDataset):
_fill_io(i)
if self._pack_sequences is not None:
while True:
# Fill every owned split before deciding whether the
# siblings still have enough data for one or more blocks.
first_count = pack_blocks_emitted[my_splits[0]]
if any(
pack_blocks_emitted[split] != first_count
for split in my_splits[1:]
):
raise ValueError(
"Packed checkpoint is not aligned across the splits "
"owned by this iterator; merge every rank or worker "
"state with merge_state_dicts before resuming on a "
"different topology"
)
while pack_blocks_emitted[my_splits[0]] < blocks_per_split:
# Each logical split gets one block per cycle. Exhausted
# splits are padded through the fixed global budget.
for i in range(n):
_fill_block(i)
# No real tokens remain in any buffer, so all owned
# splits are exhausted and the epoch is complete.
# otherwise pad the buffers and emit a block
if not any(_pack_buf(i)["tokens"] for i in range(n)):
break
for i in range(n):
tokens = _pack_buf(i)["tokens"]
tokens.extend([pad_id] * (pack_len - len(tokens)))
for i in range(n):
block = _emit_block(i)
pack_blocks_emitted[my_splits[i]] += 1
if i == n - 1:
_commit_pack_state()
_update_stats()
@@ -930,8 +1045,9 @@ class StreamingDataset(IterableDataset):
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
before resuming on a different topology.
Packed state also includes partial token buffers and is not topology-
independent because packing and padding happen within each iterator.
Packed state includes partial token buffers and emitted block counts
for every logical split. When packing is sharded, merge every rank or
worker state with ``merge_state_dicts`` before loading it.
"""
if self._pack_sequences is not None:
return {
@@ -941,7 +1057,9 @@ class StreamingDataset(IterableDataset):
"pack_sequences": self._pack_sequences,
"eos_id": self._eos_id,
"pad_id": self._pad_id,
"blocks_per_epoch": self._blocks_per_epoch,
"samples_consumed_per_split": list(self._pack_consumed),
"blocks_emitted_per_split": list(self._pack_blocks_emitted),
"pack_buffers": deepcopy(self._pack_buffers),
}
positions = [
@@ -962,7 +1080,8 @@ class StreamingDataset(IterableDataset):
Raises ``ValueError`` if ``num_splits`` or ``shuffle_seed`` differ
from the checkpoint, since a different split structure or shuffle order
makes mid-epoch resumption meaningless. Packed checkpoints
pin ``pack_sequences``, ``eos_id``, ``pad_id``, and ``epoch``.
pin ``pack_sequences``, ``eos_id``, ``pad_id``,
``blocks_per_epoch``, and ``epoch``.
"""
if state["num_splits"] != self._num_splits:
raise ValueError(
@@ -976,7 +1095,13 @@ class StreamingDataset(IterableDataset):
)
if "pack_buffers" in state or self._pack_sequences is not None:
for key in ("pack_sequences", "eos_id", "pad_id", "epoch"):
for key in (
"pack_sequences",
"eos_id",
"pad_id",
"blocks_per_epoch",
"epoch",
):
ours = getattr(self, f"_{key}")
if state.get(key) != ours:
raise ValueError(
@@ -984,6 +1109,9 @@ class StreamingDataset(IterableDataset):
f"current dataset has {ours}"
)
self._pack_consumed = [int(c) for c in state["samples_consumed_per_split"]]
self._pack_blocks_emitted = [
int(c) for c in state["blocks_emitted_per_split"]
]
self._pack_buffers = {
int(g): {"tokens": list(b["tokens"]), "starts": list(b["starts"])}
for g, b in state["pack_buffers"].items()
@@ -1011,25 +1139,22 @@ class StreamingDataset(IterableDataset):
def merge_state_dicts(states: list[dict]) -> dict:
"""Merge state dicts saved by different ranks into one exact state.
Only needed when ``on_transform_error`` skips rows in multi-rank
training: each rank then knows the exact permutation position only for
its own splits, and records a lower bound for the rest. Because
exactly one rank owns each split, the elementwise maximum across all
ranks' ``positions_consumed_per_split`` recovers the exact position of
every split. Without skipped rows every rank's state is already
identical and merging is a no-op.
For row mode, the elementwise maximum of permutation positions recovers
splits advanced by different ranks after transform failures. For packed
mode, the state that emitted the most blocks for each logical split
supplies that split's document count and partial token buffer. Packed
states must cover every rank or worker at the same global step.
Raises ``ValueError`` if the states are empty or were not produced by
the same run (mismatched seed, split count, epoch, or sample counts).
Raises ``ValueError`` if the states are empty, were not produced by
the same run, or do not represent the same global step.
The merge is always all-to-all and topology-agnostic: collect the
``state_dict()`` from every rank of the *previous* run into one list,
merge that whole list, and hand the identical merged result to every
rank of the *next* run — regardless of whether the rank count grew,
shrank, or stayed the same. There is no pairwise or subset merging
step, because each split's exact position is only known to whichever
rank owned that split, and the elementwise maximum needs every rank's
contribution to be correct.
``state_dict()`` from every rank or worker of the *previous* run into
one list, merge that whole list, and hand the identical merged result
to every rank or worker of the *next* run — regardless of whether the
topology grew, shrank, or stayed the same. There is no pairwise or
subset merging step, because each split's exact state is only known to
whichever iterator owned that split.
For example, checkpointing 8 ranks and resuming on 4 (the same
pattern applies when growing, e.g. 4 ranks resuming on 8)::
@@ -1062,6 +1187,7 @@ class StreamingDataset(IterableDataset):
if not states:
raise ValueError("merge_state_dicts requires at least one state dict")
first = states[0]
packed = "pack_buffers" in first
for state in states[1:]:
for key in ("shuffle_seed", "num_splits", "epoch"):
if state[key] != first[key]:
@@ -1069,6 +1195,73 @@ class StreamingDataset(IterableDataset):
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
if ("pack_buffers" in state) != packed:
raise ValueError("cannot merge packed and unpacked state dicts")
if packed:
config_keys = (
"pack_sequences",
"eos_id",
"pad_id",
"blocks_per_epoch",
)
for state in states[1:]:
for key in config_keys:
if state[key] != first[key]:
raise ValueError(
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
num_splits = first["num_splits"]
for state in states:
for key in (
"samples_consumed_per_split",
"blocks_emitted_per_split",
):
if len(state[key]) != num_splits:
raise ValueError(
f"{key} must contain one entry per logical split"
)
merged_consumed = []
merged_emitted = []
merged_buffers = {}
for split in range(num_splits):
owner = states[0]
owner_progress = (
owner["blocks_emitted_per_split"][split],
owner["samples_consumed_per_split"][split],
)
for state in states[1:]:
progress = (
state["blocks_emitted_per_split"][split],
state["samples_consumed_per_split"][split],
)
if progress > owner_progress:
owner = state
owner_progress = progress
merged_consumed.append(owner["samples_consumed_per_split"][split])
merged_emitted.append(owner["blocks_emitted_per_split"][split])
buffer = owner["pack_buffers"].get(
split, owner["pack_buffers"].get(str(split))
)
if buffer is not None:
merged_buffers[split] = deepcopy(buffer)
if len(set(merged_emitted)) > 1:
raise ValueError(
"packed state dicts were not captured at the same global "
"step or do not cover every rank and worker"
)
merged = dict(first)
merged["samples_consumed_per_split"] = merged_consumed
merged["blocks_emitted_per_split"] = merged_emitted
merged["pack_buffers"] = merged_buffers
return merged
for state in states[1:]:
if (
state["samples_consumed_per_split"]
!= first["samples_consumed_per_split"]
+80 -14
View File
@@ -1935,7 +1935,7 @@ def _create_token_table(tmp_path, documents):
return db.create_table("tokens", pa.table({"tokens": tokens}))
def _packed_dataset(table, pack_sequences, *, pad_id=0, **kwargs):
def _packed_dataset(table, pack_sequences, *, blocks_per_epoch, pad_id=0, **kwargs):
return StreamingDataset(
table,
shuffle=False,
@@ -1943,13 +1943,14 @@ def _packed_dataset(table, pack_sequences, *, pad_id=0, **kwargs):
pack_sequences=pack_sequences,
eos_id=9,
pad_id=pad_id,
blocks_per_epoch=blocks_per_epoch,
**kwargs,
)
def test_pack_sequences_emits_blocks_and_pads_final_tail(tmp_path):
table = _create_token_table(tmp_path, [[1, 2], [3, 4], [5]])
dataset = _packed_dataset(table, 6)
dataset = _packed_dataset(table, 6, blocks_per_epoch=2)
blocks = list(dataset)
@@ -1967,7 +1968,7 @@ def test_pack_sequences_pads_lagging_splits(tmp_path):
tmp_path,
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
)
dataset = _packed_dataset(table, 5, num_splits=2)
dataset = _packed_dataset(table, 5, blocks_per_epoch=6, num_splits=2)
input_ids = [block["input_ids"].tolist() for block in dataset]
# Split 0 has four real tokens including EOS markers, while split 1 has
# eleven. Packing must emit three complete two-split cycles.
@@ -1980,19 +1981,67 @@ def test_pack_sequences_pads_lagging_splits(tmp_path):
[9, 0, 0, 0, 0],
]
per_rank = []
for rank in range(2):
rank_dataset = _packed_dataset(
table,
5,
blocks_per_epoch=6,
num_splits=2,
world_size=2,
rank=rank,
)
per_rank.append([block["input_ids"].tolist() for block in rank_dataset])
def test_pack_sequences_checkpoint_resumes_partial_buffers(tmp_path):
assert [len(blocks) for blocks in per_rank] == [3, 3]
sharded = [block for cycle in zip(*per_rank) for block in cycle]
assert sharded == input_ids
def test_pack_sequences_auto_estimates_in_bounded_batches(tmp_path):
table = _create_token_table(tmp_path, [[1, 2, 3, 4]] * 1_000)
batch_sizes = []
getitems = streaming.Permutation.__getitems__
def record_batch_size(permutation, indices):
batch_sizes.append(len(indices))
return getitems(permutation, indices)
with patch.object(streaming.Permutation, "__getitems__", record_batch_size):
with pytest.warns(UserWarning, match="materialize a token_count column"):
dataset = _packed_dataset(
table,
5,
blocks_per_epoch="auto",
num_splits=2,
read_batch_size=3,
)
assert dataset.state_dict()["blocks_per_epoch"] == 1_000
assert sum(batch_sizes) == 10
assert max(batch_sizes) <= 3
def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path):
table = _create_token_table(
tmp_path,
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
)
dataset = _packed_dataset(table, 5, num_splits=2)
iterator = iter(dataset)
first_cycle = [next(iterator), next(iterator)]
checkpoint = dataset.state_dict()
expected_remaining = list(iterator)
kwargs = dict(pack_sequences=5, blocks_per_epoch=6, num_splits=2)
reference = list(_packed_dataset(table, **kwargs))
resumed = _packed_dataset(table, 5, num_splits=2)
datasets = [
_packed_dataset(table, world_size=2, rank=rank, **kwargs) for rank in range(2)
]
iterators = [iter(dataset) for dataset in datasets]
first_cycle = [next(iterator) for iterator in iterators]
checkpoint = StreamingDataset.merge_state_dicts(
[dataset.state_dict() for dataset in datasets]
)
for iterator in iterators:
iterator.close()
resumed = _packed_dataset(table, **kwargs)
resumed.load_state_dict(checkpoint)
actual_remaining = list(resumed)
@@ -2000,11 +2049,12 @@ def test_pack_sequences_checkpoint_resumes_partial_buffers(tmp_path):
[1, 9, 2, 9, 0],
[10, 11, 12, 13, 14],
]
assert checkpoint["blocks_emitted_per_split"] == [1, 1]
assert [block["input_ids"].tolist() for block in actual_remaining] == [
block["input_ids"].tolist() for block in expected_remaining
block["input_ids"].tolist() for block in reference[2:]
]
assert [block["doc_ids"].tolist() for block in actual_remaining] == [
block["doc_ids"].tolist() for block in expected_remaining
block["doc_ids"].tolist() for block in reference[2:]
]
@@ -2020,8 +2070,24 @@ def test_pack_sequences_validates_padding_id(tmp_path):
eos_id=9,
)
checkpoint = _packed_dataset(table, 4).state_dict()
resumed = _packed_dataset(table, 4, pad_id=8)
with pytest.raises(ValueError, match="blocks_per_epoch is required"):
StreamingDataset(
table,
shuffle=False,
columns=["tokens"],
pack_sequences=4,
eos_id=9,
pad_id=0,
)
with pytest.raises(ValueError, match="must be divisible"):
_packed_dataset(table, 4, blocks_per_epoch=3, num_splits=2)
with pytest.raises(ValueError, match="positive integer or 'auto'"):
_packed_dataset(table, 4, blocks_per_epoch="estimate")
checkpoint = _packed_dataset(table, 4, blocks_per_epoch=1).state_dict()
resumed = _packed_dataset(table, 4, blocks_per_epoch=1, pad_id=8)
with pytest.raises(ValueError, match="pad_id mismatch"):
resumed.load_state_dict(checkpoint)