From 48d6f504c2ea968c6719010df56947706156591f Mon Sep 17 00:00:00 2001 From: Ayush Chaurasia Date: Wed, 12 Aug 2026 15:20:27 +0530 Subject: [PATCH] implement global blocks per epoch --- python/python/lancedb/streaming.py | 273 +++++++++++++++--- .../python/tests/test_elastic_dataloader.py | 94 +++++- 2 files changed, 313 insertions(+), 54 deletions(-) diff --git a/python/python/lancedb/streaming.py b/python/python/lancedb/streaming.py index fd63467e6..a94b3c89f 100644 --- a/python/python/lancedb/streaming.py +++ b/python/python/lancedb/streaming.py @@ -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"] diff --git a/python/python/tests/test_elastic_dataloader.py b/python/python/tests/test_elastic_dataloader.py index 348171e2b..a1e336891 100644 --- a/python/python/tests/test_elastic_dataloader.py +++ b/python/python/tests/test_elastic_dataloader.py @@ -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)