mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
inital pass
This commit is contained in:
@@ -19,11 +19,15 @@ import os
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
import warnings
|
||||
from collections import deque
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from copy import deepcopy
|
||||
from multiprocessing import RawArray
|
||||
from typing import Any, Callable, Iterator, Optional
|
||||
from typing import Any, Callable, cast, Iterator, Optional
|
||||
|
||||
import pyarrow as pa
|
||||
import torch
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
|
||||
from .permutation import (
|
||||
@@ -127,6 +131,29 @@ 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.
|
||||
pack_sequences:
|
||||
Sequence-packing mode: token lists from consecutive documents are
|
||||
joined with ``eos_id`` and sliced into blocks of this many tokens.
|
||||
Each item is then a dict of two ``(pack_sequences,)`` LongTensors —
|
||||
``input_ids`` and ``doc_ids`` (per-position document index within
|
||||
the block, for block-diagonal masks or position-id resets).
|
||||
* Packing happens independently per owned split and preserves per-split
|
||||
resume state.
|
||||
* When a split cannot fill a real block for a cycle but an owned sibling
|
||||
still can, or when only a short tail remains at epoch end, the short buffer
|
||||
is padded to ``pack_sequences`` with ``pad_id`` so every local cycle emits
|
||||
one block per owned split.
|
||||
* ``eos_id``, ``pad_id``, and ``columns`` naming a single
|
||||
list-typed column are required; incompatible with ``transform``.
|
||||
Note: When packing is sharded across ranks or DataLoader workers, a warning
|
||||
is emitted because padding and checkpoint state are local to each iterator.
|
||||
eos_id:
|
||||
Separator token id between packed documents. Required with
|
||||
``pack_sequences``, ignored otherwise.
|
||||
pad_id:
|
||||
Padding token id used to complete blocks when a split runs out of
|
||||
real tokens mid-cycle or at epoch end. Required with
|
||||
``pack_sequences``, ignored otherwise.
|
||||
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 +179,9 @@ class StreamingDataset(IterableDataset):
|
||||
filter: Optional[str] = None,
|
||||
transform: Optional[Callable] = None,
|
||||
transform_parallelism: Optional[int] = None,
|
||||
pack_sequences: Optional[int] = None,
|
||||
eos_id: Optional[int] = None,
|
||||
pad_id: Optional[int] = None,
|
||||
connection_factory: Optional[Callable[[str], Any]] = None,
|
||||
worker_info_override=None,
|
||||
):
|
||||
@@ -167,6 +197,30 @@ class StreamingDataset(IterableDataset):
|
||||
)
|
||||
if transform_parallelism is not None and transform_parallelism <= 0:
|
||||
raise ValueError("transform_parallelism must be greater than 0")
|
||||
if pack_sequences is not None:
|
||||
if pack_sequences <= 0:
|
||||
raise ValueError("pack_sequences must be greater than 0")
|
||||
if eos_id is None:
|
||||
raise ValueError("eos_id is required when pack_sequences is set")
|
||||
if pad_id is None:
|
||||
raise ValueError("pad_id is required when pack_sequences is set")
|
||||
if transform is not None:
|
||||
raise ValueError("transform cannot be combined with pack_sequences")
|
||||
if columns is None or len(columns) != 1:
|
||||
raise ValueError(
|
||||
"pack_sequences requires columns to name exactly one "
|
||||
"list-typed column of token ids"
|
||||
)
|
||||
field = table.schema.field(columns[0])
|
||||
if not (
|
||||
pa.types.is_list(field.type)
|
||||
or pa.types.is_large_list(field.type)
|
||||
or pa.types.is_fixed_size_list(field.type)
|
||||
):
|
||||
raise ValueError(
|
||||
f"pack_sequences requires a list-typed token column; "
|
||||
f"{columns[0]!r} has type {field.type}"
|
||||
)
|
||||
|
||||
self._table = table
|
||||
self._num_splits = num_splits
|
||||
@@ -182,9 +236,16 @@ class StreamingDataset(IterableDataset):
|
||||
self._filter = filter
|
||||
self._transform = transform
|
||||
self._transform_parallelism = transform_parallelism
|
||||
self._pack_sequences = pack_sequences
|
||||
self._eos_id = eos_id
|
||||
self._pad_id = pad_id
|
||||
self._connection_factory = connection_factory
|
||||
self._worker_info_override = worker_info_override
|
||||
|
||||
# Packing resume state: documents consumed and partial-block token buffers.
|
||||
self._pack_consumed: list[int] = [0] * num_splits
|
||||
self._pack_buffers: dict[int, dict[str, list[int]]] = {}
|
||||
|
||||
# Live references to pipeline state, set only while __iter__ is running
|
||||
# in the same process. Used by the observability properties when the
|
||||
# DataLoader runs with num_workers=0.
|
||||
@@ -271,6 +332,13 @@ class StreamingDataset(IterableDataset):
|
||||
my_splits = self._resolve_my_splits()
|
||||
if not my_splits:
|
||||
return
|
||||
if self._pack_sequences is not None and len(my_splits) < self._num_splits:
|
||||
warnings.warn(
|
||||
"Sequence-packing padding is local to each rank or DataLoader "
|
||||
"worker. Sharded iterators can yield different numbers of packed "
|
||||
"blocks. Packed checkpoints cannot safely resume with a different "
|
||||
"world_size or number of DataLoader workers."
|
||||
)
|
||||
|
||||
# Set identity transform on each Permutation so __getitems__ returns
|
||||
# the raw RecordBatch. Stage 2 applies the real transform.
|
||||
@@ -282,8 +350,15 @@ 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)
|
||||
# Packing consumes documents at different rates per split, so it
|
||||
# tracks per-split offsets; row mode uses the uniform offset.
|
||||
skip = (
|
||||
self._pack_consumed[split_idx]
|
||||
if self._pack_sequences is not None
|
||||
else self._resume_offset
|
||||
)
|
||||
if skip > 0:
|
||||
perm = perm.with_skip(skip)
|
||||
permutations.append(perm)
|
||||
|
||||
n = len(permutations)
|
||||
@@ -298,9 +373,19 @@ class StreamingDataset(IterableDataset):
|
||||
if self._transform_parallelism is not None
|
||||
else (os.cpu_count() or 1)
|
||||
)
|
||||
final_transform = (
|
||||
self._transform if self._transform is not None else Transforms.arrow2python
|
||||
)
|
||||
final_transform: Callable[[pa.RecordBatch], Any]
|
||||
if self._pack_sequences is not None:
|
||||
# Packing consumes raw token lists, one per document.
|
||||
def arrow_tokens(batch: pa.RecordBatch) -> list[list[int]]:
|
||||
return cast(list[list[int]], batch.column(0).to_pylist())
|
||||
|
||||
final_transform = arrow_tokens
|
||||
else:
|
||||
final_transform = (
|
||||
self._transform
|
||||
if self._transform is not None
|
||||
else Transforms.arrow2python
|
||||
)
|
||||
|
||||
# Per-split pipeline state.
|
||||
fetch_head = [0] * n
|
||||
@@ -393,6 +478,52 @@ class StreamingDataset(IterableDataset):
|
||||
else:
|
||||
break # split exhausted
|
||||
|
||||
# ── Packing helpers (pack_sequences mode) ─────────────────────────────
|
||||
pack_len = cast(int, self._pack_sequences)
|
||||
eos_id = cast(int, self._eos_id)
|
||||
pad_id = cast(int, self._pad_id)
|
||||
pack_consumed = list(self._pack_consumed)
|
||||
pack_buffers = deepcopy(self._pack_buffers)
|
||||
|
||||
def _pack_buf(i: int) -> dict[str, list[int]]:
|
||||
return pack_buffers.setdefault(my_splits[i], {"tokens": [], "starts": []})
|
||||
|
||||
def _fill_block(i: int) -> None:
|
||||
"""Fill split i's buffer to one block or exhaust the split."""
|
||||
buf = _pack_buf(i)
|
||||
while len(buf["tokens"]) < pack_len:
|
||||
_ensure_cooked(i)
|
||||
if not cooked[i]:
|
||||
return
|
||||
buf["starts"].append(len(buf["tokens"]))
|
||||
buf["tokens"].extend(cooked[i].popleft())
|
||||
buf["tokens"].append(eos_id)
|
||||
pack_consumed[my_splits[i]] += 1
|
||||
local_consumed[i] += 1
|
||||
_advance(i)
|
||||
|
||||
def _emit_block(i: int) -> dict[str, Any]:
|
||||
buf = _pack_buf(i)
|
||||
tokens, starts = buf["tokens"], buf["starts"]
|
||||
# doc_ids label document segments within the block; 0 also covers
|
||||
# the continuation of a document begun in a prior block.
|
||||
doc_ids = torch.zeros(pack_len, dtype=torch.int64)
|
||||
doc_starts = [s for s in starts if 0 < s < pack_len]
|
||||
doc_ids[doc_starts] = 1
|
||||
doc_ids.cumsum_(dim=0) # cumulative sum marks document boundaries
|
||||
block = {
|
||||
"input_ids": torch.tensor(tokens[:pack_len], dtype=torch.int64),
|
||||
"doc_ids": doc_ids,
|
||||
}
|
||||
del tokens[:pack_len]
|
||||
# Shift start boundaries for the next call.
|
||||
buf["starts"] = [s - pack_len for s in starts if s >= pack_len]
|
||||
return block
|
||||
|
||||
def _commit_pack_state() -> None:
|
||||
self._pack_consumed = list(pack_consumed)
|
||||
self._pack_buffers = deepcopy(pack_buffers)
|
||||
|
||||
# ── Main loop ─────────────────────────────────────────────────────────
|
||||
|
||||
with ThreadPoolExecutor(max_workers=n * max_prefetch) as io_pool:
|
||||
@@ -402,10 +533,49 @@ class StreamingDataset(IterableDataset):
|
||||
self._fetch_head_ref = fetch_head
|
||||
self._split_sizes_ref = split_sizes
|
||||
self._local_consumed_ref = local_consumed
|
||||
|
||||
def _update_stats() -> None:
|
||||
# Refresh the shared-memory stats so the main process can
|
||||
# observe pipeline depth even when __iter__ runs in a
|
||||
# worker process.
|
||||
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)
|
||||
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)
|
||||
|
||||
try:
|
||||
for i in range(n):
|
||||
_fill_io(i)
|
||||
|
||||
if self._pack_sequences is not None:
|
||||
while True:
|
||||
# Fill every owned split before deciding whether the
|
||||
# siblings still have enough data for one or more blocks.
|
||||
for i in range(n):
|
||||
_fill_block(i)
|
||||
|
||||
# No real tokens remain in any buffer, so all owned
|
||||
# splits are exhausted and the epoch is complete.
|
||||
# otherwise pad the buffers and emit a block
|
||||
if not any(_pack_buf(i)["tokens"] for i in range(n)):
|
||||
break
|
||||
|
||||
for i in range(n):
|
||||
tokens = _pack_buf(i)["tokens"]
|
||||
tokens.extend([pad_id] * (pack_len - len(tokens)))
|
||||
|
||||
for i in range(n):
|
||||
block = _emit_block(i)
|
||||
if i == n - 1:
|
||||
_commit_pack_state()
|
||||
_update_stats()
|
||||
yield block
|
||||
return
|
||||
|
||||
while True:
|
||||
# Stop when any split is exhausted (all exhaust
|
||||
# simultaneously: equal split sizes + round-robin).
|
||||
@@ -424,18 +594,7 @@ class StreamingDataset(IterableDataset):
|
||||
# even when __iter__ runs in a worker process.
|
||||
if i == n - 1:
|
||||
self._resume_offset = initial_offset + local_consumed[i]
|
||||
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
|
||||
)
|
||||
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)
|
||||
_update_stats()
|
||||
|
||||
yield row
|
||||
finally:
|
||||
@@ -583,11 +742,27 @@ class StreamingDataset(IterableDataset):
|
||||
def state_dict(self) -> dict:
|
||||
"""Snapshot the dataset's consumption state.
|
||||
|
||||
The returned dict is topology-independent: at global step boundaries
|
||||
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.
|
||||
In row mode the returned dict is topology-independent: at global step
|
||||
boundaries every split has been consumed the same number of times, so
|
||||
the per-split count is a single uniform value.
|
||||
|
||||
In ``pack_sequences`` mode the state is keyed by global split index
|
||||
(documents consumed plus partial-block buffers per split), so each
|
||||
owned split resumes exactly. A topology-changing packed resume must
|
||||
collect the state of every rank and DataLoader worker; one iterator's
|
||||
state only covers its owned splits.
|
||||
"""
|
||||
if self._pack_sequences is not None:
|
||||
return {
|
||||
"shuffle_seed": self._shuffle_seed,
|
||||
"num_splits": self._num_splits,
|
||||
"epoch": self._epoch,
|
||||
"pack_sequences": self._pack_sequences,
|
||||
"eos_id": self._eos_id,
|
||||
"pad_id": self._pad_id,
|
||||
"samples_consumed_per_split": list(self._pack_consumed),
|
||||
"pack_buffers": deepcopy(self._pack_buffers),
|
||||
}
|
||||
return {
|
||||
"shuffle_seed": self._shuffle_seed,
|
||||
"num_splits": self._num_splits,
|
||||
@@ -600,7 +775,8 @@ class StreamingDataset(IterableDataset):
|
||||
|
||||
Raises ``ValueError`` if ``num_splits`` or ``shuffle_seed`` differ
|
||||
from the checkpoint, since a different split structure or shuffle order
|
||||
makes mid-epoch resumption meaningless.
|
||||
makes mid-epoch resumption meaningless. Packed checkpoints
|
||||
pin ``pack_sequences``, ``eos_id``, ``pad_id``, and ``epoch``.
|
||||
"""
|
||||
if state["num_splits"] != self._num_splits:
|
||||
raise ValueError(
|
||||
@@ -612,6 +788,22 @@ class StreamingDataset(IterableDataset):
|
||||
f"shuffle_seed mismatch: checkpoint has {state['shuffle_seed']}, "
|
||||
f"current dataset has {self._shuffle_seed}"
|
||||
)
|
||||
|
||||
if "pack_buffers" in state or self._pack_sequences is not None:
|
||||
for key in ("pack_sequences", "eos_id", "pad_id", "epoch"):
|
||||
ours = getattr(self, f"_{key}")
|
||||
if state.get(key) != ours:
|
||||
raise ValueError(
|
||||
f"{key} mismatch: checkpoint has {state.get(key)}, "
|
||||
f"current dataset has {ours}"
|
||||
)
|
||||
self._pack_consumed = [int(c) for c in state["samples_consumed_per_split"]]
|
||||
self._pack_buffers = {
|
||||
int(g): {"tokens": list(b["tokens"]), "starts": list(b["starts"])}
|
||||
for g, b in state["pack_buffers"].items()
|
||||
}
|
||||
return
|
||||
|
||||
consumed = state["samples_consumed_per_split"]
|
||||
# All entries are equal at step boundaries; use the first.
|
||||
if isinstance(consumed, list):
|
||||
|
||||
@@ -1524,6 +1524,124 @@ def test_shuffle_seed_none_generates_stable_seed(lance_table):
|
||||
assert first == second, "Same resolved seed must produce the same ordering"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sequence packing tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _create_token_table(tmp_path, documents):
|
||||
db = lancedb.connect(tmp_path)
|
||||
tokens = pa.array(documents, type=pa.list_(pa.int64()))
|
||||
return db.create_table("tokens", pa.table({"tokens": tokens}))
|
||||
|
||||
|
||||
def _packed_dataset(table, pack_sequences, *, pad_id=0, **kwargs):
|
||||
return StreamingDataset(
|
||||
table,
|
||||
shuffle=False,
|
||||
columns=["tokens"],
|
||||
pack_sequences=pack_sequences,
|
||||
eos_id=9,
|
||||
pad_id=pad_id,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def test_pack_sequences_emits_tokens_and_document_ids(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1, 2], [3, 4]])
|
||||
dataset = _packed_dataset(table, 6)
|
||||
|
||||
blocks = list(dataset)
|
||||
|
||||
assert len(blocks) == 1
|
||||
assert blocks[0]["input_ids"].tolist() == [1, 2, 9, 3, 4, 9]
|
||||
assert blocks[0]["doc_ids"].tolist() == [0, 0, 0, 1, 1, 1]
|
||||
assert blocks[0]["input_ids"].dtype == torch.int64
|
||||
assert blocks[0]["doc_ids"].dtype == torch.int64
|
||||
|
||||
|
||||
def test_pack_sequences_pads_final_short_tail(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1, 2]])
|
||||
dataset = _packed_dataset(table, 4)
|
||||
|
||||
blocks = list(dataset)
|
||||
|
||||
assert len(blocks) == 1
|
||||
assert blocks[0]["input_ids"].tolist() == [1, 2, 9, 0]
|
||||
|
||||
|
||||
def test_pack_sequences_pads_lagging_splits_to_complete_cycles(tmp_path):
|
||||
# split 0 has four real tokens including EOS markers, while split 1 has
|
||||
# eleven. Packing must emit three complete two-split cycles rather than
|
||||
# stopping at the short split at the beginning of the first/second cycle.
|
||||
table = _create_token_table(
|
||||
tmp_path,
|
||||
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
|
||||
)
|
||||
dataset = _packed_dataset(table, 5, num_splits=2)
|
||||
|
||||
input_ids = [block["input_ids"].tolist() for block in dataset]
|
||||
|
||||
assert input_ids == [
|
||||
[1, 9, 2, 9, 0],
|
||||
[10, 11, 12, 13, 14],
|
||||
[0, 0, 0, 0, 0],
|
||||
[15, 16, 17, 9, 20],
|
||||
[0, 0, 0, 0, 0],
|
||||
[9, 0, 0, 0, 0],
|
||||
]
|
||||
|
||||
|
||||
def test_pack_sequences_checkpoint_resumes_partial_buffers(tmp_path):
|
||||
table = _create_token_table(
|
||||
tmp_path,
|
||||
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
|
||||
)
|
||||
dataset = _packed_dataset(table, 5, num_splits=2)
|
||||
iterator = iter(dataset)
|
||||
first_cycle = [next(iterator), next(iterator)]
|
||||
checkpoint = dataset.state_dict()
|
||||
expected_remaining = list(iterator)
|
||||
|
||||
resumed = _packed_dataset(table, 5, num_splits=2)
|
||||
resumed.load_state_dict(checkpoint)
|
||||
actual_remaining = list(resumed)
|
||||
|
||||
assert [block["input_ids"].tolist() for block in first_cycle] == [
|
||||
[1, 9, 2, 9, 0],
|
||||
[10, 11, 12, 13, 14],
|
||||
]
|
||||
assert [block["input_ids"].tolist() for block in actual_remaining] == [
|
||||
block["input_ids"].tolist() for block in expected_remaining
|
||||
]
|
||||
assert [block["doc_ids"].tolist() for block in actual_remaining] == [
|
||||
block["doc_ids"].tolist() for block in expected_remaining
|
||||
]
|
||||
|
||||
|
||||
def test_pack_sequences_checkpoint_rejects_different_padding(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1, 2]])
|
||||
dataset = _packed_dataset(table, 4)
|
||||
checkpoint = dataset.state_dict()
|
||||
resumed = _packed_dataset(table, 4, pad_id=8)
|
||||
|
||||
with pytest.raises(ValueError, match="pad_id mismatch"):
|
||||
resumed.load_state_dict(checkpoint)
|
||||
|
||||
|
||||
def test_pack_sequences_requires_padding_id(tmp_path):
|
||||
table = _create_token_table(tmp_path, [[1, 2]])
|
||||
|
||||
with pytest.raises(ValueError, match="pad_id is required"):
|
||||
StreamingDataset(
|
||||
table,
|
||||
shuffle=False,
|
||||
columns=["tokens"],
|
||||
pack_sequences=4,
|
||||
eos_id=9,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Doc examples — each test mirrors the code snippet in index.mdx so that
|
||||
# broken doc examples are caught before they ship.
|
||||
|
||||
Reference in New Issue
Block a user