mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-25 15:38:38 +00:00
fix(python): commit streaming worker checkpoints on consumption (#4023)
## Summary - add `StreamingDataLoader`, which transports worker snapshots with prefetched batches and commits them to the parent dataset only when the trainer receives each batch - preserve exact non-uniform per-split progress and resume lagging splits without replaying already-consumed rows - reject stale parent checkpoints after a standard multi-process `DataLoader` has started, with guidance to use the consumer-aware loader - document the new public loader and merge non-uniform state across ranks ## Root cause PyTorch runs `StreamingDataset.__iter__` in private worker-process copies, while callers invoke `state_dict()` on the parent dataset. Sharing producer counters would still be incorrect because DataLoader prefetch can advance workers beyond batches returned to the trainer. ## Validation - `uv run --extra tests pytest python/tests/test_elastic_dataloader.py -q` (154 passed) - focused non-uniform merge regression (1 passed) - `uv run --project python --extra tests --extra dev ruff format .` - `uv run --project python --extra tests --extra dev ruff check .` - `cd docs && PYTHONPATH=. ../python/.venv/bin/mkdocs build` Fixes #3967 <!-- lance-gatekeeper-fix:v1 agent=572be272619660b97e87fd5c85188341 generation=1 --> --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo <github@xuanwo.io>
This commit is contained in:
committed by
GitHub
parent
94d484f539
commit
105fd73bc6
@@ -261,6 +261,8 @@ instead of being materialized with the rest of the row.
|
||||
|
||||
::: lancedb.streaming.StreamingDataset
|
||||
|
||||
::: lancedb.streaming.StreamingDataLoader
|
||||
|
||||
::: lancedb.permutation.permutation_builder
|
||||
|
||||
::: lancedb.permutation.PermutationBuilder
|
||||
|
||||
@@ -19,6 +19,7 @@ above.
|
||||
"""
|
||||
|
||||
import ctypes
|
||||
import heapq
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
@@ -29,12 +30,12 @@ from collections import deque
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from copy import deepcopy
|
||||
from multiprocessing import RawArray
|
||||
from typing import Any, Callable, cast, Iterator, Literal, Optional, Union
|
||||
from typing import Any, Callable, cast, Iterator, Literal, NamedTuple, Optional, Union
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.compute as pc
|
||||
import torch
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
|
||||
|
||||
from .permutation import (
|
||||
Permutation,
|
||||
@@ -55,6 +56,155 @@ DEFAULT_READ_BATCH_SIZE = 64
|
||||
DEFAULT_PREFETCH_BATCHES = 4
|
||||
|
||||
|
||||
class _WorkerSample(NamedTuple):
|
||||
data: Any
|
||||
dataset: "StreamingDataset"
|
||||
|
||||
|
||||
class _WorkerBatch(NamedTuple):
|
||||
data: Any
|
||||
state: dict
|
||||
|
||||
|
||||
class _ConsumerIteratorLease(NamedTuple):
|
||||
owner_token: int
|
||||
owner_thread: int
|
||||
|
||||
|
||||
class _CheckpointCollate:
|
||||
"""Attach the worker's post-fetch state to a collated batch."""
|
||||
|
||||
def __init__(self, collate_fn: Callable):
|
||||
self._collate_fn = collate_fn
|
||||
|
||||
def __call__(self, samples):
|
||||
try:
|
||||
if isinstance(samples, list):
|
||||
if not samples:
|
||||
return _WorkerBatch(self._collate_fn(samples), {})
|
||||
worker_samples = samples
|
||||
data = self._collate_fn([sample.data for sample in worker_samples])
|
||||
dataset = worker_samples[-1].dataset
|
||||
else:
|
||||
data = self._collate_fn(samples.data)
|
||||
dataset = samples.dataset
|
||||
except StopIteration as exc:
|
||||
raise RuntimeError(
|
||||
"collate_fn raised StopIteration before returning a batch"
|
||||
) from exc
|
||||
return _WorkerBatch(data, dataset._checkpoint_snapshot())
|
||||
|
||||
|
||||
class _StreamingDatasetAdapter(IterableDataset):
|
||||
"""Yield private sample wrappers for :class:`StreamingDataLoader`."""
|
||||
|
||||
def __init__(self, dataset: "StreamingDataset"):
|
||||
super().__init__()
|
||||
self.dataset = dataset
|
||||
|
||||
def __iter__(self):
|
||||
for sample in self.dataset._iter(consumer_checkpoint_transport=True):
|
||||
yield _WorkerSample(sample, self.dataset)
|
||||
|
||||
def __getattr__(self, name):
|
||||
dataset = self.__dict__.get("dataset")
|
||||
if dataset is None:
|
||||
raise AttributeError(name)
|
||||
return getattr(dataset, name)
|
||||
|
||||
|
||||
class _ConsumerCommitIterator:
|
||||
def __init__(
|
||||
self,
|
||||
iterator,
|
||||
dataset: "StreamingDataset",
|
||||
*,
|
||||
owner_token: int,
|
||||
require_uniform: bool,
|
||||
):
|
||||
self._iterator = iterator
|
||||
self._dataset = dataset
|
||||
self._owner_token = owner_token
|
||||
self._require_uniform = require_uniform
|
||||
self._released = False
|
||||
self._terminal = False
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
if self._terminal:
|
||||
raise StopIteration
|
||||
try:
|
||||
batch = next(self._iterator)
|
||||
except StopIteration:
|
||||
self._terminal = True
|
||||
self._release()
|
||||
raise
|
||||
except BaseException as exc:
|
||||
self._dataset._invalidate_checkpoint(
|
||||
f"a DataLoader batch failed before it was returned: {exc}"
|
||||
)
|
||||
raise
|
||||
try:
|
||||
if not isinstance(batch, _WorkerBatch):
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader did not receive worker checkpoint metadata"
|
||||
)
|
||||
self._dataset._commit_worker_state(
|
||||
batch.state, require_uniform=self._require_uniform
|
||||
)
|
||||
return batch.data
|
||||
except BaseException as exc:
|
||||
self._dataset._invalidate_checkpoint(
|
||||
f"a DataLoader batch failed before it was returned: {exc}"
|
||||
)
|
||||
raise
|
||||
|
||||
def _release(self) -> None:
|
||||
if self.__dict__.get("_released", True):
|
||||
return
|
||||
self._released = True
|
||||
dataset = self.__dict__.get("_dataset")
|
||||
if dataset is not None:
|
||||
dataset._release_consumer_iterator(self._owner_token)
|
||||
|
||||
def _shutdown_workers(self):
|
||||
if self.__dict__.get("_released", True):
|
||||
return None
|
||||
self._terminal = True
|
||||
iterator = self.__dict__.get("_iterator")
|
||||
shutdown = getattr(iterator, "_shutdown_workers", None)
|
||||
try:
|
||||
if shutdown is not None:
|
||||
shutdown()
|
||||
else:
|
||||
fetcher = getattr(iterator, "_dataset_fetcher", None)
|
||||
dataset_iterator = getattr(fetcher, "dataset_iter", None)
|
||||
close = getattr(dataset_iterator, "close", None)
|
||||
if close is None:
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader could not close its inner iterator"
|
||||
)
|
||||
close()
|
||||
except BaseException as exc:
|
||||
self._dataset._invalidate_checkpoint(
|
||||
f"a DataLoader iterator could not be shut down safely: {exc}"
|
||||
)
|
||||
raise
|
||||
else:
|
||||
self._release()
|
||||
|
||||
def __del__(self):
|
||||
try:
|
||||
self._shutdown_workers()
|
||||
except BaseException:
|
||||
pass
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._iterator, name)
|
||||
|
||||
|
||||
class StreamingDataset(IterableDataset):
|
||||
"""An elastic, resumable PyTorch IterableDataset backed by a LanceDB table.
|
||||
|
||||
@@ -384,6 +534,22 @@ class StreamingDataset(IterableDataset):
|
||||
# rows_skipped]
|
||||
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8)
|
||||
|
||||
# A standard multi-process DataLoader cannot report which prefetched
|
||||
# batches were actually returned to its consumer. Workers set this
|
||||
# shared flag so state_dict() can reject a stale parent checkpoint
|
||||
# unless StreamingDataLoader installed the consumer-commit transport.
|
||||
self._untracked_worker_iteration: RawArray = RawArray(ctypes.c_int64, 1)
|
||||
|
||||
# Parent-side checkpoint lifecycle. A failed DataLoader task creates
|
||||
# a permanent hole in that iterator's delivery stream, while a
|
||||
# multi-worker checkpoint is safe to restore only after all splits
|
||||
# reach the same logical step boundary.
|
||||
self._checkpoint_invalid_reason: Optional[str] = None
|
||||
self._consumer_checkpoint_requires_uniform = False
|
||||
self._consumer_iterator_lock = threading.Lock()
|
||||
self._consumer_iterator_generation = 0
|
||||
self._consumer_iterator_lease: Optional[_ConsumerIteratorLease] = None
|
||||
|
||||
# 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.
|
||||
@@ -396,6 +562,10 @@ class StreamingDataset(IterableDataset):
|
||||
# step boundaries all splits have consumed this many samples, so a
|
||||
# single scalar captures the topology-independent checkpoint state.
|
||||
self._resume_offset: int = 0
|
||||
# Exact yielded-sample counts for splits this process has advanced.
|
||||
# Missing entries use _resume_offset, which remains the lower-bound
|
||||
# checkpoint inherited from an earlier uniform/global state.
|
||||
self._resume_samples: dict[int, int] = {}
|
||||
# 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
|
||||
@@ -521,11 +691,45 @@ class StreamingDataset(IterableDataset):
|
||||
return self._rank_splits[start : start + splits_per_worker]
|
||||
|
||||
def __iter__(self) -> Iterator[dict[str, Any]]:
|
||||
return self._iter()
|
||||
|
||||
def _iter(
|
||||
self, *, consumer_checkpoint_transport: bool = False
|
||||
) -> Iterator[dict[str, Any]]:
|
||||
owner_token = None
|
||||
previous_lease = self._consumer_iterator_lease
|
||||
if consumer_checkpoint_transport:
|
||||
if not self._consumer_iterator_active:
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader worker transport requires an active "
|
||||
"parent iterator reservation"
|
||||
)
|
||||
else:
|
||||
try:
|
||||
owner_token = self._acquire_consumer_iterator()
|
||||
except BaseException:
|
||||
self._release_consumer_iterator_after_failed_acquire(previous_lease)
|
||||
raise
|
||||
try:
|
||||
yield from self._iter_owned(
|
||||
consumer_checkpoint_transport=consumer_checkpoint_transport
|
||||
)
|
||||
finally:
|
||||
if owner_token is not None:
|
||||
self._release_consumer_iterator(owner_token)
|
||||
|
||||
def _iter_owned(
|
||||
self, *, consumer_checkpoint_transport: bool
|
||||
) -> Iterator[dict[str, Any]]:
|
||||
if self._raw_batches_ref is not None:
|
||||
raise RuntimeError(
|
||||
"StreamingDataset does not support concurrent iteration. "
|
||||
"Only one active iterator per dataset instance is allowed."
|
||||
)
|
||||
real_worker = get_worker_info() is not None
|
||||
if real_worker and not consumer_checkpoint_transport:
|
||||
self._untracked_worker_iteration[0] = 1
|
||||
|
||||
my_splits = self._resolve_my_splits()
|
||||
if not my_splits:
|
||||
return
|
||||
@@ -533,6 +737,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_samples: list[int] = []
|
||||
initial_positions: list[int] = []
|
||||
for split_idx in my_splits:
|
||||
perm = Permutation.from_tables(
|
||||
@@ -541,21 +746,22 @@ class StreamingDataset(IterableDataset):
|
||||
if self._columns is not None:
|
||||
perm = perm.select_columns(self._columns)
|
||||
perm = perm.with_transform(Transforms.arrow2arrow)
|
||||
sample_count = self._resume_samples.get(split_idx, self._resume_offset)
|
||||
# Both modes resume from absolute permutation positions. Packing
|
||||
# stores them separately because it also checkpoints partial blocks.
|
||||
start_pos = (
|
||||
self._pack_consumed[split_idx]
|
||||
if self._pack_sequences is not None
|
||||
else self._resume_positions.get(split_idx, self._resume_offset)
|
||||
else self._resume_positions.get(split_idx, sample_count)
|
||||
)
|
||||
if start_pos > 0:
|
||||
perm = perm.with_skip(start_pos)
|
||||
initial_samples.append(sample_count)
|
||||
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
|
||||
@@ -853,6 +1059,27 @@ class StreamingDataset(IterableDataset):
|
||||
for i in range(n):
|
||||
_fill_io(i)
|
||||
|
||||
def _yield_row(i: int):
|
||||
pos, row = cooked[i].popleft()
|
||||
# Surface any completed prefetched failure before the
|
||||
# current row becomes durable checkpoint progress.
|
||||
_advance(i)
|
||||
local_consumed[i] += 1
|
||||
pos_consumed[i] = pos + 1
|
||||
split_idx = my_splits[i]
|
||||
self._resume_samples[split_idx] = (
|
||||
initial_samples[i] + local_consumed[i]
|
||||
)
|
||||
self._resume_positions[split_idx] = pos_consumed[i]
|
||||
return row
|
||||
|
||||
def _update_progress_stats() -> None:
|
||||
if not real_worker:
|
||||
self._resume_offset = min(
|
||||
initial_samples[j] + local_consumed[j] for j in range(n)
|
||||
)
|
||||
_update_stats()
|
||||
|
||||
if self._pack_sequences is not None:
|
||||
first_count = pack_blocks_emitted[my_splits[0]]
|
||||
if any(
|
||||
@@ -878,12 +1105,38 @@ class StreamingDataset(IterableDataset):
|
||||
tokens.extend([pad_id] * (pack_len - len(tokens)))
|
||||
block = _emit_block(i)
|
||||
pack_blocks_emitted[my_splits[i]] += 1
|
||||
# Checkpoint state must advance before yielding so
|
||||
# StreamingDataLoader can attach the exact state to
|
||||
# the batch it transports to the parent process.
|
||||
_commit_pack_state()
|
||||
if i == n - 1:
|
||||
_commit_pack_state()
|
||||
_update_stats()
|
||||
yield block
|
||||
return
|
||||
|
||||
# A checkpoint taken between round-robin split turns has
|
||||
# non-uniform counts. Resume lagging splits first so the
|
||||
# exact canonical sequence continues without replaying
|
||||
# already-consumed rows.
|
||||
if len(set(initial_samples)) > 1:
|
||||
catch_up_to = max(initial_samples)
|
||||
pending = [
|
||||
(initial_samples[i], my_splits[i], i)
|
||||
for i in range(n)
|
||||
if initial_samples[i] < catch_up_to
|
||||
]
|
||||
heapq.heapify(pending)
|
||||
while pending:
|
||||
consumed, _, i = heapq.heappop(pending)
|
||||
_ensure_cooked(i)
|
||||
if not cooked[i]:
|
||||
return
|
||||
row = _yield_row(i)
|
||||
if consumed + 1 < catch_up_to:
|
||||
heapq.heappush(pending, (consumed + 1, my_splits[i], i))
|
||||
_update_progress_stats()
|
||||
yield row
|
||||
|
||||
while True:
|
||||
# A cycle only runs if every split can still produce a
|
||||
# row. Without skips all splits exhaust simultaneously
|
||||
@@ -904,20 +1157,14 @@ class StreamingDataset(IterableDataset):
|
||||
break
|
||||
|
||||
for i in range(n):
|
||||
pos, row = cooked[i].popleft()
|
||||
local_consumed[i] += 1
|
||||
pos_consumed[i] = pos + 1
|
||||
_advance(i)
|
||||
row = _yield_row(i)
|
||||
|
||||
# After the last split in each cycle: update the
|
||||
# global offset and refresh the shared-memory stats
|
||||
# so the main process can observe pipeline depth
|
||||
# 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]
|
||||
_update_stats()
|
||||
_update_progress_stats()
|
||||
|
||||
yield row
|
||||
finally:
|
||||
@@ -1064,6 +1311,7 @@ class StreamingDataset(IterableDataset):
|
||||
"_local_consumed_ref",
|
||||
):
|
||||
state[key] = None
|
||||
state["_consumer_iterator_lock"] = None
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
@@ -1074,6 +1322,7 @@ class StreamingDataset(IterableDataset):
|
||||
table_state = state.pop("_table")
|
||||
perm_name, perm_data = state.pop("_perm_table")
|
||||
self.__dict__.update(state)
|
||||
self._consumer_iterator_lock = threading.Lock()
|
||||
if self._connection_factory is not None:
|
||||
self._table = self._connection_factory(table_name)
|
||||
else:
|
||||
@@ -1083,10 +1332,18 @@ class StreamingDataset(IterableDataset):
|
||||
def state_dict(self) -> dict:
|
||||
"""Snapshot the dataset's consumption state.
|
||||
|
||||
When using DataLoader workers, construct a
|
||||
[StreamingDataLoader][lancedb.streaming.StreamingDataLoader]. It
|
||||
commits worker state only when a prefetched batch is returned to the
|
||||
trainer. A standard multi-process ``DataLoader`` cannot expose that
|
||||
boundary, so calling this method after one has started raises
|
||||
``RuntimeError`` instead of returning stale producer state.
|
||||
|
||||
In row mode, the returned dict is topology-independent at global step
|
||||
boundaries. ``positions_consumed_per_split`` records how far each
|
||||
split's permutation has advanced, which can differ from the sample
|
||||
count when ``on_transform_error`` skips rows. Combine state dicts from
|
||||
count when ``on_transform_error`` skips rows. ``StreamingDataLoader``
|
||||
combines worker state in its parent process. Combine state dicts from
|
||||
every rank with
|
||||
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
|
||||
before resuming on a different topology.
|
||||
@@ -1095,6 +1352,43 @@ class StreamingDataset(IterableDataset):
|
||||
for every logical split. When packing is sharded, merge every rank
|
||||
state with ``merge_state_dicts`` before loading it.
|
||||
"""
|
||||
if self._untracked_worker_iteration[0] and get_worker_info() is None:
|
||||
raise RuntimeError(
|
||||
"StreamingDataset cannot checkpoint a standard DataLoader with "
|
||||
"num_workers > 0 because prefetched worker progress is not "
|
||||
"consumer-committed. Use StreamingDataLoader instead."
|
||||
)
|
||||
if self._checkpoint_invalid_reason is not None:
|
||||
raise RuntimeError(
|
||||
"StreamingDataset checkpointing is invalid because "
|
||||
f"{self._checkpoint_invalid_reason}. Load the last valid "
|
||||
"checkpoint into a fresh dataset before continuing."
|
||||
)
|
||||
state = self._checkpoint_snapshot()
|
||||
if self._pack_sequences is not None:
|
||||
rank_blocks = [
|
||||
state["blocks_emitted_per_split"][split] for split in self._rank_splits
|
||||
]
|
||||
if len(set(rank_blocks)) > 1:
|
||||
raise RuntimeError(
|
||||
"Packed StreamingDataset checkpointing is only safe at a "
|
||||
"complete logical step boundary, when every split assigned "
|
||||
"to this rank has emitted the same block count. Consume more "
|
||||
"batches before calling state_dict()."
|
||||
)
|
||||
elif self._consumer_checkpoint_requires_uniform:
|
||||
samples = state["samples_consumed_per_split"]
|
||||
rank_samples = [samples[split] for split in self._rank_splits]
|
||||
if len(set(rank_samples)) > 1:
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader checkpointing with multiple workers is "
|
||||
"only safe at a complete logical step boundary, when every "
|
||||
"split assigned to this rank has the same consumed-sample "
|
||||
"count. Consume more batches before calling state_dict()."
|
||||
)
|
||||
return state
|
||||
|
||||
def _checkpoint_snapshot(self) -> dict:
|
||||
if self._pack_sequences is not None:
|
||||
return {
|
||||
"shuffle_seed": self._shuffle_seed,
|
||||
@@ -1108,18 +1402,141 @@ class StreamingDataset(IterableDataset):
|
||||
"blocks_emitted_per_split": list(self._pack_blocks_emitted),
|
||||
"pack_buffers": deepcopy(self._pack_buffers),
|
||||
}
|
||||
samples = [
|
||||
self._resume_samples.get(split, self._resume_offset)
|
||||
for split in range(self._num_splits)
|
||||
]
|
||||
positions = [
|
||||
self._resume_positions.get(split, self._resume_offset)
|
||||
self._resume_positions.get(split, samples[split])
|
||||
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,
|
||||
"samples_consumed_per_split": samples,
|
||||
"positions_consumed_per_split": positions,
|
||||
}
|
||||
|
||||
def _invalidate_checkpoint(self, reason: str) -> None:
|
||||
if self._checkpoint_invalid_reason is None:
|
||||
self._checkpoint_invalid_reason = reason
|
||||
|
||||
@property
|
||||
def _consumer_iterator_active(self) -> bool:
|
||||
return self._consumer_iterator_lease is not None
|
||||
|
||||
@property
|
||||
def _consumer_iterator_owner(self) -> Optional[int]:
|
||||
lease = self._consumer_iterator_lease
|
||||
return lease.owner_token if lease is not None else None
|
||||
|
||||
@property
|
||||
def _consumer_iterator_owner_thread(self) -> Optional[int]:
|
||||
lease = self._consumer_iterator_lease
|
||||
return lease.owner_thread if lease is not None else None
|
||||
|
||||
def _acquire_consumer_iterator(self) -> int:
|
||||
"""Reserve this parent dataset for one checkpoint-aware iterator."""
|
||||
with self._consumer_iterator_lock:
|
||||
if self._consumer_iterator_active or self._raw_batches_ref is not None:
|
||||
raise RuntimeError(
|
||||
"StreamingDataset does not support concurrent iteration. "
|
||||
"Only one active iterator per dataset instance is allowed."
|
||||
)
|
||||
owner_thread = threading.get_ident()
|
||||
owner_token = self._consumer_iterator_generation + 1
|
||||
lease = _ConsumerIteratorLease(owner_token, owner_thread)
|
||||
self._consumer_iterator_generation = owner_token
|
||||
self._consumer_iterator_lease = lease
|
||||
return owner_token
|
||||
|
||||
def _release_consumer_iterator(self, owner_token: int) -> None:
|
||||
with self._consumer_iterator_lock:
|
||||
lease = self._consumer_iterator_lease
|
||||
if lease is not None and lease.owner_token == owner_token:
|
||||
self._consumer_iterator_lease = None
|
||||
|
||||
def _release_consumer_iterator_after_failed_acquire(
|
||||
self, previous_lease: Optional[_ConsumerIteratorLease]
|
||||
) -> None:
|
||||
"""Clean up when an interrupted acquire set a lease but did not return it."""
|
||||
owner_thread = threading.current_thread().ident
|
||||
with self._consumer_iterator_lock:
|
||||
lease = self._consumer_iterator_lease
|
||||
if (
|
||||
lease is not None
|
||||
and lease is not previous_lease
|
||||
and lease.owner_thread == owner_thread
|
||||
):
|
||||
self._consumer_iterator_lease = None
|
||||
|
||||
def _commit_worker_state(self, state: dict, *, require_uniform: bool) -> None:
|
||||
"""Merge one trainer-consumed worker batch into parent state."""
|
||||
for key, expected in (
|
||||
("shuffle_seed", self._shuffle_seed),
|
||||
("num_splits", self._num_splits),
|
||||
("epoch", self._epoch),
|
||||
):
|
||||
if state.get(key) != expected:
|
||||
raise ValueError(
|
||||
f"{key} mismatch in worker checkpoint: "
|
||||
f"{state.get(key)} != {expected}"
|
||||
)
|
||||
packed = "pack_buffers" in state
|
||||
if packed != (self._pack_sequences is not None):
|
||||
raise ValueError("worker checkpoint mode does not match the dataset")
|
||||
if packed:
|
||||
for key in ("pack_sequences", "eos_id", "pad_id", "blocks_per_epoch"):
|
||||
expected = getattr(self, f"_{key}")
|
||||
if state.get(key) != expected:
|
||||
raise ValueError(
|
||||
f"{key} mismatch in worker checkpoint: "
|
||||
f"{state.get(key)} != {expected}"
|
||||
)
|
||||
samples = state["samples_consumed_per_split"]
|
||||
emitted = state["blocks_emitted_per_split"]
|
||||
if len(samples) != self._num_splits or len(emitted) != self._num_splits:
|
||||
raise ValueError(
|
||||
"packed worker checkpoint must contain one entry per split"
|
||||
)
|
||||
buffers = state["pack_buffers"]
|
||||
for split, (count, blocks) in enumerate(zip(samples, emitted)):
|
||||
incoming = (int(blocks), int(count))
|
||||
current = (
|
||||
self._pack_blocks_emitted[split],
|
||||
self._pack_consumed[split],
|
||||
)
|
||||
if incoming > current:
|
||||
self._pack_blocks_emitted[split] = incoming[0]
|
||||
self._pack_consumed[split] = incoming[1]
|
||||
buffer = buffers.get(split, buffers.get(str(split)))
|
||||
if buffer is None:
|
||||
self._pack_buffers.pop(split, None)
|
||||
else:
|
||||
self._pack_buffers[split] = {
|
||||
"tokens": list(buffer["tokens"]),
|
||||
"starts": list(buffer["starts"]),
|
||||
}
|
||||
self._consumer_checkpoint_requires_uniform |= require_uniform
|
||||
return
|
||||
|
||||
samples = state["samples_consumed_per_split"]
|
||||
positions = state.get("positions_consumed_per_split", samples)
|
||||
for split, count in enumerate(samples):
|
||||
current = self._resume_samples.get(split, self._resume_offset)
|
||||
self._resume_samples[split] = max(current, int(count))
|
||||
for split, position in enumerate(positions):
|
||||
current = self._resume_positions.get(
|
||||
split, self._resume_samples.get(split, self._resume_offset)
|
||||
)
|
||||
self._resume_positions[split] = max(current, int(position))
|
||||
self._resume_offset = min(
|
||||
self._resume_samples.get(split, self._resume_offset)
|
||||
for split in range(self._num_splits)
|
||||
)
|
||||
self._consumer_checkpoint_requires_uniform |= require_uniform
|
||||
|
||||
def load_state_dict(self, state: dict) -> None:
|
||||
"""Resume from a previously snapshotted state.
|
||||
|
||||
@@ -1139,6 +1556,7 @@ class StreamingDataset(IterableDataset):
|
||||
f"shuffle_seed mismatch: checkpoint has {state['shuffle_seed']}, "
|
||||
f"current dataset has {self._shuffle_seed}"
|
||||
)
|
||||
self._consumer_checkpoint_requires_uniform = False
|
||||
|
||||
if "pack_buffers" in state or self._pack_sequences is not None:
|
||||
for key in (
|
||||
@@ -1165,14 +1583,17 @@ class StreamingDataset(IterableDataset):
|
||||
return
|
||||
|
||||
consumed = state["samples_consumed_per_split"]
|
||||
# All entries are equal at step boundaries; use the first.
|
||||
if isinstance(consumed, list):
|
||||
self._resume_offset = consumed[0] if consumed else 0
|
||||
self._resume_offset = min(consumed) if consumed else 0
|
||||
self._resume_samples = {
|
||||
split: int(count) for split, count in enumerate(consumed)
|
||||
}
|
||||
else:
|
||||
self._resume_offset = int(consumed)
|
||||
self._resume_samples = {}
|
||||
# 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.
|
||||
# the per-split sample count (the .get default in __iter__) is exact.
|
||||
positions = state.get("positions_consumed_per_split")
|
||||
if positions is None:
|
||||
self._resume_positions = {}
|
||||
@@ -1185,10 +1606,11 @@ class StreamingDataset(IterableDataset):
|
||||
def merge_state_dicts(states: list[dict]) -> dict:
|
||||
"""Merge state dicts saved by different ranks into one exact state.
|
||||
|
||||
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 permutation position and partial token buffer. Packed
|
||||
In row mode, each rank records exact consumer-committed progress for
|
||||
its own splits and lower bounds for the rest, so elementwise maxima
|
||||
recover both sample counts and permutation positions. In packed mode,
|
||||
the state that emitted the most blocks for each logical split supplies
|
||||
that split's permutation position and partial token buffer. Packed
|
||||
states must cover every rank at the same global step.
|
||||
|
||||
Raises ``ValueError`` if the states are empty, were not produced by
|
||||
@@ -1299,17 +1721,13 @@ class StreamingDataset(IterableDataset):
|
||||
merged["pack_buffers"] = merged_buffers
|
||||
return merged
|
||||
|
||||
for state in states[1:]:
|
||||
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)
|
||||
merged["samples_consumed_per_split"] = [
|
||||
max(per_split)
|
||||
for per_split in zip(
|
||||
*(state["samples_consumed_per_split"] for state in states)
|
||||
)
|
||||
]
|
||||
all_positions = [
|
||||
state.get(
|
||||
"positions_consumed_per_split", state["samples_consumed_per_split"]
|
||||
@@ -1320,3 +1738,113 @@ class StreamingDataset(IterableDataset):
|
||||
max(per_split) for per_split in zip(*all_positions)
|
||||
]
|
||||
return merged
|
||||
|
||||
|
||||
class StreamingDataLoader(DataLoader):
|
||||
"""A PyTorch DataLoader with consumer-committed dataset checkpoints.
|
||||
|
||||
PyTorch workers prefetch batches ahead of the trainer, so worker-local
|
||||
producer progress is not a safe checkpoint. This loader carries a state
|
||||
snapshot alongside every internal batch and applies it to the parent
|
||||
[StreamingDataset][lancedb.streaming.StreamingDataset] only when that batch
|
||||
is returned by ``next()``.
|
||||
The trainer receives the same collated batch it would receive from a
|
||||
standard ``torch.utils.data.DataLoader``.
|
||||
|
||||
With more than one worker, row-mode ``state_dict()`` is available only at
|
||||
complete logical step boundaries, when every split assigned to the rank has
|
||||
the same consumed-sample count. Packed checkpoints require equal emitted-block
|
||||
counts across the rank's splits for any worker count. ``persistent_workers=True``
|
||||
is not supported because prefetched worker copies cannot be restored from
|
||||
parent-committed state. If batch collation raises, checkpointing remains
|
||||
invalid for that dataset instance; restore the last valid checkpoint into a
|
||||
fresh dataset before continuing.
|
||||
Only one active iterator may own a dataset at a time, including when worker
|
||||
processes are used. Exhausting or explicitly shutting down the iterator
|
||||
releases that ownership. ``drop_last=True`` is not supported because worker
|
||||
replicas discard incomplete tails independently, which cannot produce a
|
||||
topology-independent checkpoint.
|
||||
|
||||
Parameters are the same as ``torch.utils.data.DataLoader`` except that
|
||||
``dataset`` must be a
|
||||
[StreamingDataset][lancedb.streaming.StreamingDataset].
|
||||
Subclasses that override ``StreamingDataset.__iter__`` are not supported
|
||||
because the custom iterator cannot provide the exact per-yield checkpoint
|
||||
snapshots required by this loader.
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> # dataset = StreamingDataset(table, num_splits=2)
|
||||
>>> # loader = StreamingDataLoader(dataset, batch_size=8, num_workers=2)
|
||||
>>> # batch = next(iter(loader))
|
||||
>>> # checkpoint = dataset.state_dict()
|
||||
"""
|
||||
|
||||
def __init__(self, dataset: StreamingDataset, *args, **kwargs):
|
||||
if not isinstance(dataset, StreamingDataset):
|
||||
raise TypeError("StreamingDataLoader requires a StreamingDataset")
|
||||
if type(dataset).__iter__ is not StreamingDataset.__iter__:
|
||||
raise TypeError(
|
||||
"StreamingDataLoader does not support StreamingDataset subclasses "
|
||||
"that override __iter__ because they cannot provide exact "
|
||||
"per-yield checkpoint state"
|
||||
)
|
||||
if kwargs.get("in_order", True) is False:
|
||||
raise ValueError(
|
||||
"StreamingDataLoader requires in_order=True for deterministic "
|
||||
"consumer checkpoints"
|
||||
)
|
||||
if kwargs.get("persistent_workers", False):
|
||||
raise ValueError(
|
||||
"StreamingDataLoader does not support persistent_workers=True "
|
||||
"because worker prefetch state cannot be reset from a checkpoint"
|
||||
)
|
||||
self._streaming_dataset = dataset
|
||||
super().__init__(_StreamingDatasetAdapter(dataset), *args, **kwargs)
|
||||
if self.drop_last:
|
||||
raise ValueError(
|
||||
"StreamingDataLoader does not support drop_last=True because "
|
||||
"discarded worker tails cannot be checkpointed "
|
||||
"topology-independently"
|
||||
)
|
||||
self.collate_fn = _CheckpointCollate(self.collate_fn)
|
||||
|
||||
def __iter__(self):
|
||||
dataset = self._streaming_dataset
|
||||
previous_lease = dataset._consumer_iterator_lease
|
||||
owner_token = None
|
||||
try:
|
||||
owner_token = dataset._acquire_consumer_iterator()
|
||||
state = dataset._checkpoint_snapshot()
|
||||
packed = dataset._pack_sequences is not None
|
||||
if packed:
|
||||
blocks = state["blocks_emitted_per_split"]
|
||||
rank_blocks = [blocks[split] for split in dataset._rank_splits]
|
||||
if len(set(rank_blocks)) > 1:
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader cannot start from a partial packed "
|
||||
"logical step; resume from a checkpoint whose splits "
|
||||
"assigned to this rank have equal emitted-block counts"
|
||||
)
|
||||
elif self.num_workers > 1:
|
||||
samples = state["samples_consumed_per_split"]
|
||||
rank_samples = [samples[split] for split in dataset._rank_splits]
|
||||
if len(set(rank_samples)) > 1:
|
||||
raise RuntimeError(
|
||||
"StreamingDataLoader cannot start multiple workers from a "
|
||||
"partial logical step; resume from a checkpoint whose "
|
||||
"splits assigned to this rank have equal consumed-sample "
|
||||
"counts"
|
||||
)
|
||||
return _ConsumerCommitIterator(
|
||||
super().__iter__(),
|
||||
dataset,
|
||||
owner_token=owner_token,
|
||||
require_uniform=self.num_workers > 1 or packed,
|
||||
)
|
||||
except BaseException:
|
||||
if owner_token is not None:
|
||||
dataset._release_consumer_iterator(owner_token)
|
||||
else:
|
||||
dataset._release_consumer_iterator_after_failed_acquire(previous_lease)
|
||||
raise
|
||||
|
||||
@@ -32,6 +32,7 @@ Parameters used throughout:
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
import threading
|
||||
from unittest.mock import patch
|
||||
|
||||
import lancedb
|
||||
@@ -46,6 +47,7 @@ from utils import (
|
||||
torch = pytest.importorskip("torch")
|
||||
streaming = pytest.importorskip("lancedb.streaming")
|
||||
StreamingDataset = streaming.StreamingDataset
|
||||
StreamingDataLoader = streaming.StreamingDataLoader
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dataset parameters
|
||||
@@ -92,6 +94,27 @@ class FakeWorkerInfo:
|
||||
num_workers: int
|
||||
|
||||
|
||||
def _collate_with_first_batch_error(samples):
|
||||
ids = [sample["id"] for sample in samples]
|
||||
if ids == [0, 1]:
|
||||
raise ValueError("first batch fails")
|
||||
return ids
|
||||
|
||||
|
||||
def _collate_with_first_batch_stop(samples):
|
||||
ids = [sample["id"] for sample in samples]
|
||||
if ids == [0, 1]:
|
||||
raise StopIteration("first batch stopped")
|
||||
return ids
|
||||
|
||||
|
||||
def _collate_with_first_batch_interrupt(samples):
|
||||
ids = [sample["id"] for sample in samples]
|
||||
if ids == [0, 1]:
|
||||
raise KeyboardInterrupt("first batch interrupted")
|
||||
return ids
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1008,6 +1031,565 @@ def test_multi_worker_elastic_det_across_worker_counts(lance_table):
|
||||
# ── Resumability with num_workers ─────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_streaming_dataloader_commits_only_consumed_worker_batches(tmp_path):
|
||||
"""Prefetched worker state is committed only as the trainer receives it."""
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"worker_commit", pa.table({"id": [1, 2, 3, 4, 10, 20, 30, 40]})
|
||||
)
|
||||
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
prefetch_factor=4,
|
||||
)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
first = next(iterator)["id"].tolist()
|
||||
|
||||
assert first == [1, 2]
|
||||
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2, 0]
|
||||
with pytest.raises(RuntimeError, match="complete logical step boundary"):
|
||||
dataset.state_dict()
|
||||
|
||||
second = next(iterator)["id"].tolist()
|
||||
assert second == [10, 20]
|
||||
checkpoint = dataset.state_dict()
|
||||
assert checkpoint["samples_consumed_per_split"] == [2, 2]
|
||||
uninterrupted = [batch["id"].tolist() for batch in iterator]
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
resumed = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
resumed.load_state_dict(checkpoint)
|
||||
resumed_loader = StreamingDataLoader(
|
||||
resumed,
|
||||
batch_size=2,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
prefetch_factor=4,
|
||||
)
|
||||
resumed_iterator = iter(resumed_loader)
|
||||
try:
|
||||
remaining = [batch["id"].tolist() for batch in resumed_iterator]
|
||||
finally:
|
||||
resumed_iterator._shutdown_workers()
|
||||
assert remaining == uninterrupted == [[3, 4], [30, 40]]
|
||||
|
||||
|
||||
def test_distributed_checkpoint_uses_rank_local_worker_boundary(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("rank_boundary", pa.table({"id": list(range(8))}))
|
||||
dataset = StreamingDataset(
|
||||
table,
|
||||
num_splits=4,
|
||||
shuffle=False,
|
||||
rank=0,
|
||||
world_size=2,
|
||||
)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
assert next(iterator)["id"].tolist() == [0]
|
||||
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [
|
||||
1,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
]
|
||||
with pytest.raises(RuntimeError, match="complete logical step boundary"):
|
||||
dataset.state_dict()
|
||||
|
||||
assert next(iterator)["id"].tolist() == [2]
|
||||
checkpoint = dataset.state_dict()
|
||||
remaining = [batch["id"].tolist() for batch in iterator]
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
assert checkpoint["samples_consumed_per_split"] == [1, 1, 0, 0]
|
||||
assert remaining == [[1], [3]]
|
||||
|
||||
|
||||
def test_standard_dataloader_rejects_stale_parent_checkpoint(tmp_path):
|
||||
"""A standard DataLoader must not expose prefetched producer progress."""
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("untracked_workers", pa.table({"id": [1, 2, 10, 20]}))
|
||||
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
# Merely constructing the checkpoint-aware loader must not authorize a
|
||||
# later plain DataLoader's worker progress.
|
||||
StreamingDataLoader(dataset, batch_size=2, num_workers=0)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
assert next(iterator)["id"].tolist() == [1, 2]
|
||||
with pytest.raises(RuntimeError, match="Use StreamingDataLoader"):
|
||||
dataset.state_dict()
|
||||
list(iterator)
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
|
||||
def test_streaming_dataloader_rejects_persistent_workers(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("persistent_workers", pa.table({"id": [1, 2]}))
|
||||
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
|
||||
with pytest.raises(ValueError, match="persistent_workers=True"):
|
||||
StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
num_workers=2,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
def test_collate_failure_invalidates_consumer_checkpoint(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"collate_failure", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]})
|
||||
)
|
||||
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
collate_fn=_collate_with_first_batch_error,
|
||||
prefetch_factor=2,
|
||||
)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="first batch fails"):
|
||||
next(iterator)
|
||||
assert next(iterator) == [100, 101]
|
||||
assert next(iterator) == [2, 3]
|
||||
with pytest.raises(RuntimeError, match="failed before it was returned"):
|
||||
dataset.state_dict()
|
||||
list(iterator)
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
|
||||
def test_collate_stop_iteration_invalidates_consumer_checkpoint(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("collate_stop", pa.table({"id": list(range(6))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=0,
|
||||
collate_fn=_collate_with_first_batch_stop,
|
||||
)
|
||||
iterator = iter(loader)
|
||||
|
||||
with pytest.raises(RuntimeError, match="collate_fn raised StopIteration"):
|
||||
next(iterator)
|
||||
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
|
||||
with pytest.raises(RuntimeError, match="failed before it was returned"):
|
||||
dataset.state_dict()
|
||||
assert list(iterator) == [[2, 3], [4, 5]]
|
||||
|
||||
|
||||
def test_batch_base_exception_invalidates_consumer_checkpoint(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("collate_interrupt", pa.table({"id": list(range(6))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=0,
|
||||
collate_fn=_collate_with_first_batch_interrupt,
|
||||
)
|
||||
iterator = iter(loader)
|
||||
|
||||
with pytest.raises(KeyboardInterrupt, match="first batch interrupted"):
|
||||
next(iterator)
|
||||
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
|
||||
with pytest.raises(RuntimeError, match="failed before it was returned"):
|
||||
dataset.state_dict()
|
||||
assert list(iterator) == [[2, 3], [4, 5]]
|
||||
|
||||
|
||||
def test_parent_commit_base_exception_invalidates_consumer_checkpoint(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("commit_interrupt", pa.table({"id": list(range(4))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
|
||||
iterator = iter(loader)
|
||||
real_commit = dataset._commit_worker_state
|
||||
|
||||
def interrupt_after_commit(state, *, require_uniform):
|
||||
real_commit(state, require_uniform=require_uniform)
|
||||
raise KeyboardInterrupt("after parent commit")
|
||||
|
||||
with patch.object(
|
||||
dataset, "_commit_worker_state", side_effect=interrupt_after_commit
|
||||
):
|
||||
with pytest.raises(KeyboardInterrupt, match="after parent commit"):
|
||||
next(iterator)
|
||||
|
||||
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
|
||||
with pytest.raises(RuntimeError, match="failed before it was returned"):
|
||||
dataset.state_dict()
|
||||
|
||||
|
||||
def test_direct_iteration_surfaces_prefetch_failure_before_committing_row(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("prefetch_failure", pa.table({"id": list(range(4))}))
|
||||
release = threading.Event()
|
||||
failed = threading.Event()
|
||||
real_getitems = streaming.Permutation.__getitems__
|
||||
|
||||
def controlled_getitems(permutation, indices):
|
||||
if indices and indices[0] >= 2:
|
||||
assert release.wait(timeout=5)
|
||||
failed.set()
|
||||
raise RuntimeError("later prefetched I/O failed")
|
||||
return real_getitems(permutation, indices)
|
||||
|
||||
class SignalDict(dict):
|
||||
def __setitem__(self, key, value):
|
||||
super().__setitem__(key, value)
|
||||
release.set()
|
||||
assert failed.wait(timeout=5)
|
||||
|
||||
monkeypatch.setattr(streaming.Permutation, "__getitems__", controlled_getitems)
|
||||
dataset = StreamingDataset(
|
||||
table,
|
||||
num_splits=1,
|
||||
shuffle=False,
|
||||
read_batch_size=2,
|
||||
io_queue_depth=2,
|
||||
)
|
||||
dataset._resume_positions = SignalDict()
|
||||
iterator = iter(dataset)
|
||||
|
||||
assert next(iterator)["id"] == 0
|
||||
with pytest.raises(RuntimeError, match="later prefetched I/O failed"):
|
||||
next(iterator)
|
||||
|
||||
checkpoint = dataset.state_dict()
|
||||
assert checkpoint["samples_consumed_per_split"] == [1]
|
||||
assert checkpoint["positions_consumed_per_split"] == [1]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("workers", [0, 1, 2])
|
||||
def test_streaming_dataloader_rejects_drop_last(tmp_path, workers):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("drop_last", pa.table({"id": [0, 1, 2]}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
worker_options = {"multiprocessing_context": "spawn"} if workers else {}
|
||||
|
||||
with pytest.raises(ValueError, match="drop_last=True"):
|
||||
StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=workers,
|
||||
drop_last=True,
|
||||
**worker_options,
|
||||
)
|
||||
|
||||
|
||||
def test_streaming_dataloader_owns_one_iterator_until_teardown(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("iterator_owner", pa.table({"id": list(range(4))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=1,
|
||||
multiprocessing_context="spawn",
|
||||
)
|
||||
|
||||
first = iter(loader)
|
||||
try:
|
||||
assert next(first)["id"].tolist() == [0, 1]
|
||||
with pytest.raises(RuntimeError, match="concurrent iteration"):
|
||||
iter(loader)
|
||||
finally:
|
||||
first._shutdown_workers()
|
||||
|
||||
second = iter(loader)
|
||||
try:
|
||||
assert [batch["id"].tolist() for batch in second] == [[2, 3]]
|
||||
except BaseException:
|
||||
second._shutdown_workers()
|
||||
raise
|
||||
|
||||
# Natural exhaustion releases ownership too.
|
||||
third = iter(loader)
|
||||
try:
|
||||
assert list(third) == []
|
||||
finally:
|
||||
third._shutdown_workers()
|
||||
|
||||
|
||||
def test_zero_worker_shutdown_closes_inner_iterator_before_release(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("zero_worker_shutdown", pa.table({"id": list(range(6))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
|
||||
|
||||
first = iter(loader)
|
||||
assert next(first)["id"].tolist() == [0, 1]
|
||||
first._shutdown_workers()
|
||||
|
||||
assert dataset._consumer_iterator_active is False
|
||||
assert dataset._raw_batches_ref is None
|
||||
second = iter(loader)
|
||||
try:
|
||||
with pytest.raises(StopIteration):
|
||||
next(first)
|
||||
assert next(second)["id"].tolist() == [2, 3]
|
||||
finally:
|
||||
second._shutdown_workers()
|
||||
|
||||
|
||||
def test_direct_and_loader_admission_share_one_atomic_lease(tmp_path, monkeypatch):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("direct_loader_lease", pa.table({"id": list(range(4))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
|
||||
entered = threading.Event()
|
||||
release = threading.Event()
|
||||
direct_result = []
|
||||
direct_error = []
|
||||
contender = []
|
||||
real_resolve = dataset._resolve_my_splits
|
||||
|
||||
def controlled_resolve():
|
||||
if threading.current_thread().name == "direct-start":
|
||||
entered.set()
|
||||
assert release.wait(timeout=5)
|
||||
return real_resolve()
|
||||
|
||||
def advance_direct(iterator):
|
||||
try:
|
||||
direct_result.append(next(iterator)["id"])
|
||||
except BaseException as exc:
|
||||
direct_error.append(exc)
|
||||
|
||||
monkeypatch.setattr(dataset, "_resolve_my_splits", controlled_resolve)
|
||||
direct = iter(dataset)
|
||||
thread = threading.Thread(
|
||||
target=advance_direct, args=(direct,), name="direct-start"
|
||||
)
|
||||
thread.start()
|
||||
assert entered.wait(timeout=5)
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="concurrent iteration"):
|
||||
contender.append(iter(loader))
|
||||
finally:
|
||||
release.set()
|
||||
thread.join(timeout=5)
|
||||
if contender:
|
||||
contender[0]._shutdown_workers()
|
||||
direct.close()
|
||||
|
||||
assert not thread.is_alive()
|
||||
assert direct_error == []
|
||||
assert direct_result == [0]
|
||||
|
||||
|
||||
def test_loader_acquires_before_snapshot_and_cleans_interrupted_acquire(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("lease_snapshot", pa.table({"id": list(range(4))}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
|
||||
first = iter(loader)
|
||||
assert next(first)["id"].tolist() == [0, 1]
|
||||
|
||||
entered = threading.Event()
|
||||
release = threading.Event()
|
||||
pending = []
|
||||
pending_errors = []
|
||||
observed_snapshots = []
|
||||
real_acquire = dataset._acquire_consumer_iterator
|
||||
real_snapshot = dataset._checkpoint_snapshot
|
||||
|
||||
def controlled_acquire():
|
||||
if threading.current_thread().name == "stale-start":
|
||||
entered.set()
|
||||
assert release.wait(timeout=5)
|
||||
return real_acquire()
|
||||
|
||||
def recording_snapshot():
|
||||
state = real_snapshot()
|
||||
if threading.current_thread().name == "stale-start":
|
||||
observed_snapshots.append(state["samples_consumed_per_split"])
|
||||
return state
|
||||
|
||||
def create_pending_iterator():
|
||||
try:
|
||||
pending.append(iter(loader))
|
||||
except BaseException as exc:
|
||||
pending_errors.append(exc)
|
||||
|
||||
monkeypatch.setattr(dataset, "_acquire_consumer_iterator", controlled_acquire)
|
||||
monkeypatch.setattr(dataset, "_checkpoint_snapshot", recording_snapshot)
|
||||
thread = threading.Thread(target=create_pending_iterator, name="stale-start")
|
||||
thread.start()
|
||||
assert entered.wait(timeout=5)
|
||||
assert next(first)["id"].tolist() == [2, 3]
|
||||
with pytest.raises(StopIteration):
|
||||
next(first)
|
||||
release.set()
|
||||
thread.join(timeout=5)
|
||||
|
||||
assert not thread.is_alive()
|
||||
assert pending_errors == []
|
||||
assert observed_snapshots == [[4]]
|
||||
assert len(pending) == 1
|
||||
assert list(pending[0]) == []
|
||||
assert dataset.state_dict()["samples_consumed_per_split"] == [4]
|
||||
|
||||
def interrupted_acquire():
|
||||
real_acquire()
|
||||
raise KeyboardInterrupt("after acquire")
|
||||
|
||||
monkeypatch.setattr(dataset, "_acquire_consumer_iterator", interrupted_acquire)
|
||||
with pytest.raises(KeyboardInterrupt, match="after acquire"):
|
||||
iter(loader)
|
||||
assert dataset._consumer_iterator_active is False
|
||||
|
||||
|
||||
def test_consumer_iterator_lease_publication_is_atomic(tmp_path, monkeypatch):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("atomic_lease", pa.table({"id": [0, 1]}))
|
||||
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
|
||||
loader = StreamingDataLoader(dataset, batch_size=1, num_workers=0)
|
||||
real_get_ident = streaming.threading.get_ident
|
||||
calls = 0
|
||||
|
||||
def interrupt_during_publication():
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls == 1:
|
||||
raise KeyboardInterrupt("during lease mutation")
|
||||
return real_get_ident()
|
||||
|
||||
monkeypatch.setattr(streaming.threading, "get_ident", interrupt_during_publication)
|
||||
with pytest.raises(KeyboardInterrupt, match="during lease mutation"):
|
||||
iter(loader)
|
||||
monkeypatch.setattr(streaming.threading, "get_ident", real_get_ident)
|
||||
|
||||
assert dataset._consumer_iterator_active is False
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
assert next(iterator)["id"].tolist() == [0]
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
|
||||
def test_streaming_dataloader_rejects_dataset_iter_override(tmp_path):
|
||||
class CustomizedDataset(StreamingDataset):
|
||||
def __iter__(self):
|
||||
return iter([1000, 1001])
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("custom_iteration", pa.table({"id": [0, 1, 2]}))
|
||||
dataset = CustomizedDataset(table, num_splits=1, shuffle=False)
|
||||
|
||||
assert list(dataset) == [1000, 1001]
|
||||
with pytest.raises(TypeError, match="override __iter__"):
|
||||
StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
num_workers=0,
|
||||
collate_fn=list,
|
||||
)
|
||||
|
||||
|
||||
def test_interleaved_adapters_do_not_authorize_plain_iteration(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table_a = db.create_table("adapter_a", pa.table({"id": [0, 1]}))
|
||||
table_b = db.create_table("adapter_b", pa.table({"id": [10, 11]}))
|
||||
dataset_a = StreamingDataset(table_a, num_splits=1, shuffle=False)
|
||||
dataset_b = StreamingDataset(table_b, num_splits=1, shuffle=False)
|
||||
initial_state = dataset_a.state_dict()
|
||||
|
||||
owner_a = dataset_a._acquire_consumer_iterator()
|
||||
owner_b = dataset_b._acquire_consumer_iterator()
|
||||
try:
|
||||
iterator_a = iter(streaming._StreamingDatasetAdapter(dataset_a))
|
||||
iterator_b = iter(streaming._StreamingDatasetAdapter(dataset_b))
|
||||
assert next(iterator_a).data["id"] == 0
|
||||
assert next(iterator_b).data["id"] == 10
|
||||
assert [sample.data["id"] for sample in iterator_a] == [1]
|
||||
assert [sample.data["id"] for sample in iterator_b] == [11]
|
||||
finally:
|
||||
dataset_a._release_consumer_iterator(owner_a)
|
||||
dataset_b._release_consumer_iterator(owner_b)
|
||||
|
||||
dataset_a.load_state_dict(initial_state)
|
||||
with patch(
|
||||
"lancedb.streaming.get_worker_info",
|
||||
return_value=FakeWorkerInfo(id=0, num_workers=1),
|
||||
):
|
||||
plain_iterator = iter(dataset_a)
|
||||
assert next(plain_iterator)["id"] == 0
|
||||
plain_iterator.close()
|
||||
|
||||
assert dataset_a._untracked_worker_iteration[0] == 1
|
||||
with pytest.raises(RuntimeError, match="Use StreamingDataLoader"):
|
||||
dataset_a.state_dict()
|
||||
|
||||
|
||||
def test_resume_from_partial_split_cycle_preserves_remaining_order(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("partial_cycle", pa.table({"id": [1, 2, 10, 20]}))
|
||||
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
iterator = iter(dataset)
|
||||
|
||||
assert next(iterator)["id"] == 1
|
||||
checkpoint = dataset.state_dict()
|
||||
iterator.close()
|
||||
assert checkpoint["samples_consumed_per_split"] == [1, 0]
|
||||
|
||||
resumed = StreamingDataset(table, num_splits=2, shuffle=False)
|
||||
resumed.load_state_dict(checkpoint)
|
||||
assert [row["id"] for row in resumed] == [10, 2, 20]
|
||||
|
||||
|
||||
def test_partial_cycle_resume_preserves_skip_truncation(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"partial_skip", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]})
|
||||
)
|
||||
kwargs = dict(
|
||||
num_splits=2,
|
||||
shuffle=False,
|
||||
transform=_failing_transform({1, 2, 3}),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
dataset = StreamingDataset(table, **kwargs)
|
||||
iterator = iter(dataset)
|
||||
|
||||
assert next(iterator)["id"] == 0
|
||||
checkpoint = dataset.state_dict()
|
||||
uninterrupted = [row["id"] for row in iterator]
|
||||
|
||||
resumed = StreamingDataset(table, **kwargs)
|
||||
resumed.load_state_dict(checkpoint)
|
||||
assert [row["id"] for row in resumed] == uninterrupted == [100]
|
||||
|
||||
|
||||
def test_multi_worker_resumability_same_topology(lance_table):
|
||||
"""Checkpoint with num_workers=2, resume with num_workers=2: exact continuation."""
|
||||
world_size = 1
|
||||
@@ -2018,6 +2600,23 @@ def test_merge_state_dicts_validates_consistency(lance_table):
|
||||
StreamingDataset.merge_state_dicts([])
|
||||
|
||||
|
||||
def test_merge_state_dicts_combines_nonuniform_consumer_progress(lance_table):
|
||||
dataset = StreamingDataset(
|
||||
lance_table, num_splits=2, shuffle=False, shuffle_seed=SHUFFLE_SEED
|
||||
)
|
||||
rank0 = dataset.state_dict()
|
||||
rank0["samples_consumed_per_split"] = [2, 0]
|
||||
rank0["positions_consumed_per_split"] = [2, 0]
|
||||
rank1 = dataset.state_dict()
|
||||
rank1["samples_consumed_per_split"] = [0, 2]
|
||||
rank1["positions_consumed_per_split"] = [0, 2]
|
||||
|
||||
merged = StreamingDataset.merge_state_dicts([rank0, rank1])
|
||||
|
||||
assert merged["samples_consumed_per_split"] == [2, 2]
|
||||
assert merged["positions_consumed_per_split"] == [2, 2]
|
||||
|
||||
|
||||
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)."""
|
||||
@@ -2254,6 +2853,65 @@ def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path):
|
||||
]
|
||||
|
||||
|
||||
def test_packed_checkpoint_requires_complete_split_cycle(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1], [2], [10], [20]])
|
||||
dataset = _packed_dataset(table, pack_sequences=3, blocks_per_epoch=4, num_splits=2)
|
||||
iterator = iter(dataset)
|
||||
|
||||
next(iterator)
|
||||
with pytest.raises(RuntimeError, match="complete logical step boundary"):
|
||||
dataset.state_dict()
|
||||
|
||||
next(iterator)
|
||||
assert dataset.state_dict()["blocks_emitted_per_split"] == [1, 1]
|
||||
iterator.close()
|
||||
|
||||
|
||||
def test_streaming_dataloader_commits_consumed_packed_batches(tmp_path):
|
||||
table = _create_token_table(
|
||||
tmp_path,
|
||||
[[1], [2], [3], [4], [10], [20], [30], [40]],
|
||||
)
|
||||
kwargs = dict(pack_sequences=4, blocks_per_epoch=4, num_splits=2)
|
||||
dataset = _packed_dataset(table, **kwargs)
|
||||
loader = StreamingDataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
prefetch_factor=2,
|
||||
)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
next(iterator)
|
||||
with pytest.raises(RuntimeError, match="complete logical step boundary"):
|
||||
dataset.state_dict()
|
||||
|
||||
next(iterator)
|
||||
checkpoint = dataset.state_dict()
|
||||
uninterrupted = [batch["input_ids"].tolist() for batch in iterator]
|
||||
finally:
|
||||
iterator._shutdown_workers()
|
||||
|
||||
resumed = _packed_dataset(table, **kwargs)
|
||||
resumed.load_state_dict(checkpoint)
|
||||
resumed_loader = StreamingDataLoader(
|
||||
resumed,
|
||||
batch_size=1,
|
||||
num_workers=2,
|
||||
multiprocessing_context="spawn",
|
||||
prefetch_factor=2,
|
||||
)
|
||||
resumed_iterator = iter(resumed_loader)
|
||||
try:
|
||||
remaining = [batch["input_ids"].tolist() for batch in resumed_iterator]
|
||||
finally:
|
||||
resumed_iterator._shutdown_workers()
|
||||
|
||||
assert checkpoint["blocks_emitted_per_split"] == [1, 1]
|
||||
assert remaining == uninterrupted
|
||||
|
||||
|
||||
def test_pack_sequences_validates_configuration_and_tokens(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1, 2]])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user