inital pass

This commit is contained in:
Ayush Chaurasia
2026-08-11 15:25:31 +05:30
parent 12405a4077
commit 508621cb38
2 changed files with 333 additions and 23 deletions
+215 -23
View File
@@ -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.