diff --git a/python/python/lancedb/streaming.py b/python/python/lancedb/streaming.py index 525ed3d63..b27e606a4 100644 --- a/python/python/lancedb/streaming.py +++ b/python/python/lancedb/streaming.py @@ -11,6 +11,11 @@ Provides StreamingDataset, a PyTorch IterableDataset that guarantees: - **Resumability**: state_dict / load_state_dict capture per-split consumption counts so training can resume from an exact mid-epoch position even when the distributed topology changes between runs. + +Transform failures on bad rows (e.g. nulls or NaNs from incomplete data) can +be tolerated with ``on_transform_error="skip"``; see the parameter +documentation on StreamingDataset for how this interacts with the guarantees +above. """ import ctypes @@ -22,7 +27,7 @@ import time from collections import deque from concurrent.futures import ThreadPoolExecutor from multiprocessing import RawArray -from typing import Any, Callable, Iterator, Optional +from typing import Any, Callable, Iterator, Optional, Union from torch.utils.data import IterableDataset, get_worker_info @@ -127,6 +132,49 @@ class StreamingDataset(IterableDataset): Maximum number of transforms to run concurrently. Must be greater than zero. When ``None`` (the default), uses ``os.cpu_count()`` or 1 when the CPU count is unavailable. + on_transform_error: + What to do when the transform raises an exception: + + - ``"raise"`` (the default): the exception propagates and iteration + aborts. + - ``"skip"``: the failing rows are dropped and iteration continues. + - ``"warn"``: like ``"skip"``, but a warning is logged for each + failing batch. + - a callable ``handler(exc) -> bool``: called with the exception; + return ``True`` to skip the failing rows or ``False`` to re-raise. + Useful to skip only expected error types (compatible with + ``webdataset.handlers`` style handlers). + + When a batch fails, the transform is re-invoked on each single-row + slice of the batch so that only the rows that actually fail are + dropped. Transforms should therefore be deterministic and accept + batches of any size (including one row). Skipped rows are counted in + ``rows_skipped``. + + Skipping weakens the elastic-determinism guarantee at the end of the + epoch: splits that lose more rows than others run dry earlier, and + each rank's iterator ends at the last cycle where every split *it + owns* still has a row. Because bad rows are not distributed evenly + across splits, this means one rank's iterator can yield noticeably + fewer or more steps than another rank's *in the same run* — there is + no cross-rank coordination that stops every rank at the same global + step. This is generally safe for asynchronous or single-rank use, + but synchronous distributed training (e.g. ranks that call + ``all_reduce`` every step) can hang or deadlock if one rank's + iterator is exhausted while others are still stepping; callers doing + synchronous multi-rank training with ``on_transform_error != "raise"`` + are responsible for their own cross-rank stopping mechanism (e.g. + broadcasting a stop signal on ``StopIteration``). The final few + global steps can also differ across topologies (bounded by the skew + in bad-row counts across splits). The sequence of samples yielded + from each split remains deterministic. Mid-epoch + checkpoints remain exact provided the transform fails + deterministically; in multi-rank training each rank must save its + own ``state_dict`` and the states must be combined with + ``merge_state_dicts`` before resuming on a different topology. + Prefer the ``filter`` parameter when bad rows can be expressed as a + SQL predicate (e.g. ``"col IS NOT NULL"``) — filtering happens before + splits are built, so every guarantee is fully preserved. worker_info_override: If set, used in place of ``torch.utils.data.get_worker_info()`` to determine the DataLoader worker assignment. Intended for unit tests @@ -152,6 +200,7 @@ class StreamingDataset(IterableDataset): filter: Optional[str] = None, transform: Optional[Callable] = None, transform_parallelism: Optional[int] = None, + on_transform_error: Union[str, Callable[[Exception], bool]] = "raise", connection_factory: Optional[Callable[[str], Any]] = None, worker_info_override=None, ): @@ -167,6 +216,13 @@ class StreamingDataset(IterableDataset): ) if transform_parallelism is not None and transform_parallelism <= 0: raise ValueError("transform_parallelism must be greater than 0") + if on_transform_error not in ("raise", "skip", "warn") and not callable( + on_transform_error + ): + raise ValueError( + "on_transform_error must be 'raise', 'skip', 'warn', or a " + f"callable, got {on_transform_error!r}" + ) self._table = table self._num_splits = num_splits @@ -182,6 +238,7 @@ class StreamingDataset(IterableDataset): self._filter = filter self._transform = transform self._transform_parallelism = transform_parallelism + self._on_transform_error = on_transform_error self._connection_factory = connection_factory self._worker_info_override = worker_info_override @@ -199,19 +256,28 @@ class StreamingDataset(IterableDataset): # in the main process. RawArray is picklable via the forkserver # reduction protocol so it survives the dataset pickle round-trip. # Layout: [unscanned_rows, raw_rows, cooked_rows, consumed_rows, - # bytes_loaded, fetch_time_us, transform_time_us] - self._worker_stats: RawArray = RawArray(ctypes.c_int64, 7) + # bytes_loaded, fetch_time_us, transform_time_us, + # rows_skipped] + self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8) # Cumulative bytes of Arrow buffer data fetched across all iterations. self._bytes_loaded: int = 0 # Cumulative seconds spent in LanceDB I/O and in transform functions. self._fetch_time: float = 0.0 self._transform_time: float = 0.0 + # Cumulative rows dropped by on_transform_error across all iterations. + self._rows_skipped: int = 0 # Number of samples each split has already been consumed. At global # step boundaries all splits have consumed this many samples, so a # single scalar captures the topology-independent checkpoint state. self._resume_offset: int = 0 + # Permutation position each split has consumed through, keyed by + # global split index. Equal to _resume_offset for every split unless + # on_transform_error skipped rows, in which case skipped positions + # push the watermark of the affected splits further ahead. Splits + # this instance has never iterated have no entry. + self._resume_positions: dict[int, int] = {} # Build the permutation table once, deterministically. builder = permutation_builder(table) @@ -275,6 +341,7 @@ class StreamingDataset(IterableDataset): # Set identity transform on each Permutation so __getitems__ returns # the raw RecordBatch. Stage 2 applies the real transform. permutations: list[Permutation] = [] + initial_positions: list[int] = [] for split_idx in my_splits: perm = Permutation.from_tables( self._table, self._perm_table, split=split_idx @@ -282,14 +349,20 @@ class StreamingDataset(IterableDataset): if self._columns is not None: perm = perm.select_columns(self._columns) perm = perm.with_transform(lambda batch: batch) - if self._resume_offset > 0: - perm = perm.with_skip(self._resume_offset) + start_pos = self._resume_positions.get(split_idx, self._resume_offset) + if start_pos > 0: + perm = perm.with_skip(start_pos) + initial_positions.append(start_pos) permutations.append(perm) n = len(permutations) split_sizes = [perm.num_rows for perm in permutations] initial_offset = self._resume_offset local_consumed = [0] * n + # Permutation position each split has consumed through (absolute, + # i.e. counted from the start of the unskipped split). Runs ahead of + # initial + local_consumed when rows are skipped. + pos_consumed = list(initial_positions) batch_size = self._read_batch_size max_prefetch = self._prefetch_batches @@ -302,12 +375,14 @@ class StreamingDataset(IterableDataset): self._transform if self._transform is not None else Transforms.arrow2python ) - # Per-split pipeline state. + # Per-split pipeline state. Batches are paired with the absolute + # permutation position of their first row so that skipped rows can be + # accounted for in pos_consumed. fetch_head = [0] * n - io_pending = [deque() for _ in range(n)] # Future[RecordBatch] - raw_batches = [deque() for _ in range(n)] # RecordBatch — fetched, awaiting tx - tx_pending = [deque() for _ in range(n)] # Future[list[Any]] - cooked = [deque() for _ in range(n)] # rows ready to yield + io_pending = [deque() for _ in range(n)] # (abs_start, Future[RecordBatch]) + raw_batches = [deque() for _ in range(n)] # (abs_start, RecordBatch) + tx_pending = [deque() for _ in range(n)] # Future[list[(abs_pos, row)]] + cooked = [deque() for _ in range(n)] # (abs_pos, row) ready to yield # Limit simultaneous transforms to transform_workers across all splits. tx_semaphore = threading.Semaphore(transform_workers) @@ -330,7 +405,8 @@ class StreamingDataset(IterableDataset): fetch_head[i] += fetch perm_i = permutations[i] indices = list(range(start, start + fetch)) - io_pending[i].append(io_pool.submit(_io_call, perm_i, indices)) + abs_start = initial_positions[i] + start + io_pending[i].append((abs_start, io_pool.submit(_io_call, perm_i, indices))) def _fill_io(i: int) -> None: while len(io_pending[i]) < max_prefetch and fetch_head[i] < split_sizes[i]: @@ -338,15 +414,72 @@ class StreamingDataset(IterableDataset): def _drain_io(i: int) -> None: """Move completed I/O futures into raw_batches non-blockingly.""" - while io_pending[i] and io_pending[i][0].done(): - raw_batches[i].append(io_pending[i].popleft().result()) + while io_pending[i] and io_pending[i][0][1].done(): + abs_start, fut = io_pending[i].popleft() + raw_batches[i].append((abs_start, fut.result())) # ── Stage 2 helpers ─────────────────────────────────────────────────── - def _tx_call_guarded(batch): + on_error = self._on_transform_error + + def _should_skip(exc: Exception) -> bool: + if on_error == "raise": + return False + if callable(on_error): + return bool(on_error(exc)) + return True # "skip" or "warn" + + def _check_row_count(rows: list, num_rows: int) -> None: + if len(rows) != num_rows: + raise ValueError( + f"transform returned {len(rows)} rows for a batch of " + f"{num_rows}; transforms must return exactly one output " + "row per input row. To drop bad rows, raise inside the " + "transform and pass on_transform_error='skip'." + ) + + def _transform_isolated(abs_start, batch, batch_exc): + """Re-run the transform on single-row slices, dropping failures.""" + out = [] + skipped = 0 + first_exc = None + for j in range(batch.num_rows): + try: + rows = list(final_transform(batch.slice(j, 1))) + except Exception as exc: + if not _should_skip(exc): + raise + skipped += 1 + if first_exc is None: + first_exc = exc + continue + _check_row_count(rows, 1) + out.append((abs_start + j, rows[0])) + self._rows_skipped += skipped + if skipped and on_error == "warn": + logger.warning( + "Skipped %d of %d rows whose transform failed (first error: %r)", + skipped, + batch.num_rows, + first_exc if first_exc is not None else batch_exc, + ) + return out + + def _transform_batch(abs_start, batch): + """Apply the transform, returning [(abs_pos, row), ...].""" + try: + rows = list(final_transform(batch)) + except Exception as exc: + if not _should_skip(exc): + raise + return _transform_isolated(abs_start, batch, exc) + _check_row_count(rows, batch.num_rows) + return [(abs_start + j, row) for j, row in enumerate(rows)] + + def _tx_call_guarded(abs_start, batch): try: t0 = time.perf_counter() - result = final_transform(batch) + result = _transform_batch(abs_start, batch) self._transform_time += time.perf_counter() - t0 return result finally: @@ -355,8 +488,8 @@ class StreamingDataset(IterableDataset): def _try_submit_tx(i: int) -> None: """Submit transforms for raw_batches[i] up to available capacity.""" while raw_batches[i] and tx_semaphore.acquire(blocking=False): - batch = raw_batches[i].popleft() - tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch)) + abs_start, batch = raw_batches[i].popleft() + tx_pending[i].append(tx_pool.submit(_tx_call_guarded, abs_start, batch)) def _drain_tx(i: int) -> None: """Move completed transform futures into cooked non-blockingly.""" @@ -384,11 +517,14 @@ class StreamingDataset(IterableDataset): # Acquire a transform slot (may block briefly if all # transform_workers are busy with other splits). tx_semaphore.acquire() - batch = raw_batches[i].popleft() - tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch)) + abs_start, batch = raw_batches[i].popleft() + tx_pending[i].append( + tx_pool.submit(_tx_call_guarded, abs_start, batch) + ) elif io_pending[i]: # Block on the oldest in-flight I/O fetch. - raw_batches[i].append(io_pending[i].popleft().result()) + abs_start, fut = io_pending[i].popleft() + raw_batches[i].append((abs_start, fut.result())) _advance(i) else: break # split exhausted @@ -407,15 +543,28 @@ class StreamingDataset(IterableDataset): _fill_io(i) while True: - # Stop when any split is exhausted (all exhaust - # simultaneously: equal split sizes + round-robin). - if any(local_consumed[i] >= split_sizes[i] for i in range(n)): + # A cycle only runs if every split can still produce a + # row. Without skips all splits exhaust simultaneously + # (equal split sizes + round-robin); when + # on_transform_error drops rows a split can run dry + # early, ending the epoch at the last complete cycle. + # This check only sees splits owned by this rank/worker + # (my_splits) — there is no cross-rank coordination, so + # a different rank with fewer skipped rows keeps going; + # see the on_transform_error docstring. + exhausted = False + for i in range(n): + _ensure_cooked(i) + if not cooked[i]: + exhausted = True + break + if exhausted: break for i in range(n): - _ensure_cooked(i) - row = cooked[i].popleft() + pos, row = cooked[i].popleft() local_consumed[i] += 1 + pos_consumed[i] = pos + 1 _advance(i) # After the last split in each cycle: update the @@ -424,21 +573,39 @@ class StreamingDataset(IterableDataset): # even when __iter__ runs in a worker process. if i == n - 1: self._resume_offset = initial_offset + local_consumed[i] + for j, split_idx in enumerate(my_splits): + self._resume_positions[split_idx] = pos_consumed[j] ws = self._worker_stats ws[0] = sum( split_sizes[j] - fetch_head[j] for j in range(n) ) ws[1] = sum( - batch.num_rows for q in raw_batches for batch in q + batch.num_rows + for q in raw_batches + for _, batch in q ) ws[2] = sum(len(q) for q in cooked) ws[3] = sum(local_consumed) ws[4] = self._bytes_loaded ws[5] = int(self._fetch_time * 1_000_000) ws[6] = int(self._transform_time * 1_000_000) + ws[7] = self._rows_skipped yield row finally: + # Final stats flush: the per-cycle write above never runs + # when iteration ends mid-cycle (e.g. a split whose rows + # were all skipped before completing a single cycle), so + # counters like rows_skipped would otherwise be stale. + ws = self._worker_stats + ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n)) + ws[1] = 0 # queue-depth properties document 0 when idle + ws[2] = 0 + ws[3] = sum(local_consumed) + ws[4] = self._bytes_loaded + ws[5] = int(self._fetch_time * 1_000_000) + ws[6] = int(self._transform_time * 1_000_000) + ws[7] = self._rows_skipped self._raw_batches_ref = None self._cooked_ref = None self._fetch_head_ref = None @@ -492,7 +659,7 @@ class StreamingDataset(IterableDataset): batches. Returns 0 when not iterating. """ if self._raw_batches_ref is not None: - return sum(batch.num_rows for q in self._raw_batches_ref for batch in q) + return sum(batch.num_rows for q in self._raw_batches_ref for _, batch in q) return int(self._worker_stats[1]) @property @@ -522,6 +689,19 @@ class StreamingDataset(IterableDataset): ) return int(self._worker_stats[0]) + @property + def rows_skipped(self) -> int: + """Number of rows dropped because their transform raised an exception. + + Only ever non-zero when ``on_transform_error`` is set to ``"skip"``, + ``"warn"``, or a callable that returned ``True``. Accumulates across + multiple iterations of the same dataset instance and is never reset + automatically. + """ + if self._raw_batches_ref is not None: + return self._rows_skipped + return int(self._worker_stats[7]) + @property def consumed_rows(self) -> int: """Number of rows already yielded to the caller across all splits. @@ -587,12 +767,27 @@ class StreamingDataset(IterableDataset): every split has been consumed the same number of times (by the round-robin design), so the per-split count is a single uniform value that is identical across all ranks and DataLoader workers. + + ``positions_consumed_per_split`` records how far into each split's + permutation iteration has advanced. It only differs from + ``samples_consumed_per_split`` when ``on_transform_error`` skipped + rows, in which case entries are exact for the splits this instance + iterated and a lower bound (the sample count) for splits owned by + other ranks or workers. Combine the state dicts from all ranks with + [merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts] + to recover the exact value for every split before resuming on a + different topology. """ + positions = [ + self._resume_positions.get(split, self._resume_offset) + for split in range(self._num_splits) + ] return { "shuffle_seed": self._shuffle_seed, "num_splits": self._num_splits, "epoch": self._epoch, "samples_consumed_per_split": [self._resume_offset] * self._num_splits, + "positions_consumed_per_split": positions, } def load_state_dict(self, state: dict) -> None: @@ -618,3 +813,96 @@ class StreamingDataset(IterableDataset): self._resume_offset = consumed[0] if consumed else 0 else: self._resume_offset = int(consumed) + # Older checkpoints predate positions_consumed_per_split; without + # skipped rows positions equal sample counts, so falling back to + # _resume_offset (the .get default in __iter__) is exact. + positions = state.get("positions_consumed_per_split") + if positions is None: + self._resume_positions = {} + else: + self._resume_positions = { + split: int(pos) for split, pos in enumerate(positions) + } + + @staticmethod + 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. + + Raises ``ValueError`` if the states are empty or were not produced by + the same run (mismatched seed, split count, epoch, or sample counts). + + 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. + + For example, checkpointing 8 ranks and resuming on 4 (the same + pattern applies when growing, e.g. 4 ranks resuming on 8):: + + states = [ds.state_dict() for ds in previous_run_datasets] # 8 + merged = StreamingDataset.merge_state_dicts(states) + for ds in resumed_datasets: # now only 4 ranks + ds.load_state_dict(merged) # same dict on every rank + + The rank count on either side never affects the merge itself, since + ``merge_state_dicts`` only cares about the list of states it is + given. Each split's position is recovered by elementwise maximum; + here rank 0 owned split 0 (and skipped two rows there) while rank 1 + owned split 1 (and skipped one row): + + >>> rank0 = { + ... "shuffle_seed": 0, "num_splits": 2, "epoch": 0, + ... "samples_consumed_per_split": [3, 3], + ... "positions_consumed_per_split": [5, 3], + ... } + >>> rank1 = { + ... "shuffle_seed": 0, "num_splits": 2, "epoch": 0, + ... "samples_consumed_per_split": [3, 3], + ... "positions_consumed_per_split": [3, 4], + ... } + >>> merged = StreamingDataset.merge_state_dicts([rank0, rank1]) + >>> merged["positions_consumed_per_split"] + [5, 4] + """ + if not states: + raise ValueError("merge_state_dicts requires at least one state dict") + first = states[0] + for state in states[1:]: + for key in ("shuffle_seed", "num_splits", "epoch"): + if state[key] != first[key]: + raise ValueError( + f"{key} mismatch across state dicts: " + f"{state[key]} != {first[key]}" + ) + if ( + state["samples_consumed_per_split"] + != first["samples_consumed_per_split"] + ): + raise ValueError( + "samples_consumed_per_split mismatch across state dicts; " + "state_dict() must be called at the same global step " + "boundary on every rank" + ) + merged = dict(first) + all_positions = [ + state.get( + "positions_consumed_per_split", state["samples_consumed_per_split"] + ) + for state in states + ] + merged["positions_consumed_per_split"] = [ + max(per_split) for per_split in zip(*all_positions) + ] + return merged diff --git a/python/python/tests/test_elastic_dataloader.py b/python/python/tests/test_elastic_dataloader.py index 22918082b..734f835c6 100644 --- a/python/python/tests/test_elastic_dataloader.py +++ b/python/python/tests/test_elastic_dataloader.py @@ -1456,6 +1456,408 @@ def test_shuffle_clump_size_yields_all_rows(lance_table): ) +# --------------------------------------------------------------------------- +# on_transform_error tests +# --------------------------------------------------------------------------- + + +class BadRowError(ValueError): + """Raised by the failing transforms below when a batch contains a bad id.""" + + +def _failing_transform(bad_ids: set): + """A transform that raises BadRowError whenever the batch has a bad id. + + Raises on the full batch and on any single-row slice containing a bad id, + so per-row isolation drops exactly the bad rows. + """ + + def transform(batch: pa.RecordBatch) -> list: + ids = batch.column("id").to_pylist() + bad = sorted(set(ids) & bad_ids) + if bad: + raise BadRowError(f"bad ids in batch: {bad}") + return [{"id": i} for i in ids] + + return transform + + +def _sequential_split_members(table) -> list[list[int]]: + """Return each split's ids in yield order for shuffle=False. + + With a single rank and no workers the round-robin yields one row per split + per cycle, so item k of a clean run belongs to split k % NUM_SPLITS. + """ + ds = StreamingDataset(table, num_splits=NUM_SPLITS, shuffle=False) + members: list[list[int]] = [[] for _ in range(NUM_SPLITS)] + for k, row in enumerate(ds): + members[k % NUM_SPLITS].append(row["id"]) + return members + + +def test_on_transform_error_default_raises(lance_table): + """By default a transform exception propagates and aborts iteration.""" + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle_seed=SHUFFLE_SEED, + transform=_failing_transform({7}), + ) + with pytest.raises(BadRowError): + list(ds) + + +def test_on_transform_error_invalid_value(lance_table): + with pytest.raises(ValueError, match="on_transform_error"): + StreamingDataset(lance_table, num_splits=NUM_SPLITS, on_transform_error="bogus") + + +def test_on_transform_error_skip_drops_bad_rows(lance_table): + """With one bad row per split, 'skip' yields every good row exactly once + and counts the dropped rows in rows_skipped.""" + members = _sequential_split_members(lance_table) + bad_ids = {members[i][4] for i in range(NUM_SPLITS)} + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + assert ds.rows_skipped == 0 + + ids = [row["id"] for row in ds] + + assert sorted(ids) == sorted(set(range(NUM_ROWS)) - bad_ids) + assert ds.rows_skipped == NUM_SPLITS + + +def test_on_transform_error_skip_uneven_ends_at_last_complete_cycle(lance_table): + """When one split loses more rows than the others, the epoch ends at the + last cycle where every split still has a row — no crash, no bad rows, and + every step remains one sample per split.""" + members = _sequential_split_members(lance_table) + bad_ids = set(members[0][:3]) # all 3 bad rows in split 0 + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + items = [row["id"] for row in ds] + + rows_per_split = NUM_ROWS // NUM_SPLITS + expected_cycles = rows_per_split - len(bad_ids) + assert len(items) == expected_cycles * NUM_SPLITS + assert len(set(items)) == len(items), "duplicate samples yielded" + assert not set(items) & bad_ids, "a bad row was yielded" + # Split 0 contributed exactly its surviving rows, in order, one per cycle. + survivors = [i for i in members[0] if i not in bad_ids] + assert items[0::NUM_SPLITS] == survivors[:expected_cycles] + + +def test_on_transform_error_warn_logs(lance_table, caplog): + """'warn' skips like 'skip' but logs a warning for the failing batch.""" + members = _sequential_split_members(lance_table) + bad_ids = {members[i][3] for i in range(NUM_SPLITS)} + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="warn", + ) + with caplog.at_level(logging.WARNING, logger="lancedb.streaming"): + items = list(ds) + + assert len(items) == NUM_ROWS - NUM_SPLITS + assert ds.rows_skipped == NUM_SPLITS + assert "Skipped" in caplog.text + assert "BadRowError" in caplog.text + + +def test_on_transform_error_callable_selective(lance_table): + """A callable handler can skip expected errors and re-raise the rest.""" + members = _sequential_split_members(lance_table) + bad_ids = {members[i][0] for i in range(NUM_SPLITS)} + + handled: list[Exception] = [] + + def handler(exc: Exception) -> bool: + handled.append(exc) + return isinstance(exc, BadRowError) + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error=handler, + ) + items = list(ds) + assert len(items) == NUM_ROWS - NUM_SPLITS + assert handled and all(isinstance(exc, BadRowError) for exc in handled) + + def broken_transform(batch: pa.RecordBatch) -> list: + raise TypeError("boom") + + ds2 = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=broken_transform, + on_transform_error=handler, + ) + with pytest.raises(TypeError, match="boom"): + list(ds2) + + +def test_transform_wrong_row_count_raises(lance_table): + """A transform that returns the wrong number of rows is an error even with + on_transform_error='skip' — silent shrinkage would corrupt accounting.""" + + def drops_rows(batch: pa.RecordBatch) -> list: + return batch.column("id").to_pylist()[:-1] + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle_seed=SHUFFLE_SEED, + transform=drops_rows, + on_transform_error="skip", + ) + with pytest.raises(ValueError, match="one output row per input row"): + list(ds) + + +def test_skip_deterministic_across_runs(lance_table): + """With a fixed seed, skipping produces the identical sample sequence on + every run — skips are data-dependent, not run-dependent.""" + bad_ids = {5, 17, 46} + + def run() -> tuple[list[int], int]: + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle_seed=SHUFFLE_SEED, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + return [row["id"] for row in ds], ds.rows_skipped + + ids_a, skipped_a = run() + ids_b, skipped_b = run() + assert ids_a == ids_b + assert skipped_a == skipped_b + assert not set(ids_a) & bad_ids + + +def test_skip_elastic_det_across_world_sizes(lance_table): + """With equal bad-row counts per split, skipping preserves the full + elastic-determinism guarantee: identical global batches at every step for + every compatible world_size.""" + members = _sequential_split_members(lance_table) + bad_ids = {members[i][6] for i in range(NUM_SPLITS)} + + def collect(world_size: int) -> list[frozenset[int]]: + micro = GLOBAL_BATCH_SIZE // world_size + iters = [ + iter( + StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + rank=rank, + world_size=world_size, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + ) + for rank in range(world_size) + ] + _STOP = object() + batches: list[frozenset[int]] = [] + while True: + step_samples: set[int] = set() + exhausted = 0 + for it in iters: + for _ in range(micro): + val = next(it, _STOP) + if val is _STOP: + exhausted += 1 + break + step_samples.add(val["id"]) + if exhausted == len(iters): + break + assert exhausted == 0, ( + "Rank iterators exhausted at different steps despite equal " + "bad-row counts per split" + ) + batches.append(frozenset(step_samples)) + return batches + + reference = collect(1) + assert len(reference) == NUM_ROWS // NUM_SPLITS - 1 + for ws in (2, 3, 4): + assert collect(ws) == reference, f"world_size={ws} diverged" + + +def test_resumability_with_skips_same_topology(lance_table): + """Checkpointing mid-epoch with skipped rows resumes exactly: no sample + repeated, no sample lost, skipped rows stay skipped.""" + members = _sequential_split_members(lance_table) + # Uneven skips: positions diverge across splits (2 bad in split 0, 1 in + # split 5), which only a position-based checkpoint can resume exactly. + bad_ids = {members[0][2], members[0][3], members[5][7]} + kwargs = dict( + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + + reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)] + rows_per_split = NUM_ROWS // NUM_SPLITS + assert len(reference) == (rows_per_split - 2) * NUM_SPLITS + + steps = 3 + ds = StreamingDataset(lance_table, **kwargs) + it = iter(ds) + consumed = [next(it)["id"] for _ in range(steps * NUM_SPLITS)] + checkpoint = ds.state_dict() + it.close() + + # Split 0 skipped positions 2 and 3 within its first 3 yields; split 5's + # bad row is beyond the checkpoint. Everything else is at 3 = the sample + # count. + positions = checkpoint["positions_consumed_per_split"] + assert positions[0] == 5 + assert positions[1:] == [3] * (NUM_SPLITS - 1) + assert checkpoint["samples_consumed_per_split"] == [3] * NUM_SPLITS + + ds2 = StreamingDataset(lance_table, **kwargs) + ds2.load_state_dict(checkpoint) + resumed = [row["id"] for row in ds2] + + assert consumed == reference[: steps * NUM_SPLITS] + assert resumed == reference[steps * NUM_SPLITS :] + + +def test_resumability_with_skips_elastic_merge(lance_table): + """Elastic resume with skips: each rank's checkpoint knows exact positions + only for its own splits; merge_state_dicts recovers the global state, and + a run on a different world_size continues exactly.""" + members = _sequential_split_members(lance_table) + # Bad rows early in split 0 (rank 0) and split 6 (rank 1 of a ws=2 run) so + # both ranks' position vectors diverge before the checkpoint. + bad_ids = {members[0][0], members[0][2], members[6][1]} + kwargs = dict( + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + + reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)] + + steps = 3 + world_size = 2 + micro = GLOBAL_BATCH_SIZE // world_size + datasets = [ + StreamingDataset(lance_table, rank=rank, world_size=world_size, **kwargs) + for rank in range(world_size) + ] + iters = [iter(ds) for ds in datasets] + seen: list[frozenset[int]] = [] + for _ in range(steps): + step_samples = set() + for it in iters: + for _ in range(micro): + step_samples.add(next(it)["id"]) + seen.append(frozenset(step_samples)) + states = [ds.state_dict() for ds in datasets] + for it in iters: + it.close() + + merged = StreamingDataset.merge_state_dicts(states) + expected_positions = [3] * NUM_SPLITS + expected_positions[0] = 5 # skipped positions 0 and 2 + expected_positions[6] = 4 # skipped position 1 + assert merged["positions_consumed_per_split"] == expected_positions + + # The first 3 global batches match the world_size=1 reference. + ref_batches = [ + frozenset(reference[s * NUM_SPLITS : (s + 1) * NUM_SPLITS]) + for s in range(len(reference) // NUM_SPLITS) + ] + assert seen == ref_batches[:steps] + + # Resume on world_size=1 from the merged state. + ds_resume = StreamingDataset(lance_table, **kwargs) + ds_resume.load_state_dict(merged) + resumed = [row["id"] for row in ds_resume] + assert resumed == reference[steps * NUM_SPLITS :] + + +def test_rows_skipped_flushed_when_split_entirely_bad(lance_table): + """A split whose rows all fail never completes a cycle, so the epoch ends + immediately — but rows_skipped must still report the drops after the + iterator exits (the shared-memory counter is flushed on exhaustion).""" + members = _sequential_split_members(lance_table) + bad_ids = set(members[0]) # every row of split 0 is bad + + ds = StreamingDataset( + lance_table, + num_splits=NUM_SPLITS, + shuffle=False, + transform=_failing_transform(bad_ids), + on_transform_error="skip", + ) + assert list(ds) == [] + assert ds.rows_skipped == len(bad_ids) + + +def test_merge_state_dicts_validates_consistency(lance_table): + ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED) + state = ds.state_dict() + other = dict(state, shuffle_seed=SHUFFLE_SEED + 1) + with pytest.raises(ValueError, match="shuffle_seed mismatch"): + StreamingDataset.merge_state_dicts([state, other]) + with pytest.raises(ValueError, match="at least one"): + StreamingDataset.merge_state_dicts([]) + + +def test_load_state_dict_without_positions_key(lance_table): + """Checkpoints from before positions_consumed_per_split existed still + resume exactly (positions equal sample counts when nothing is skipped).""" + reference = [ + row["id"] + for row in StreamingDataset( + lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED + ) + ] + + steps = 4 + ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED) + it = iter(ds) + for _ in range(steps * NUM_SPLITS): + next(it) + checkpoint = ds.state_dict() + it.close() + del checkpoint["positions_consumed_per_split"] + + ds2 = StreamingDataset( + lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED + ) + ds2.load_state_dict(checkpoint) + resumed = [row["id"] for row in ds2] + assert resumed == reference[steps * NUM_SPLITS :] + + def test_num_splits_defaults_to_world_size(lance_table): """Omitting num_splits gives world_size splits (one per rank).""" ds = StreamingDataset(