add blocks per epoch support for elastic resume

This commit is contained in:
Ayush Chaurasia
2026-08-18 16:30:33 +05:30
parent 7c2851dcca
commit d0405bf1db
2 changed files with 123 additions and 95 deletions
+104 -92
View File
@@ -149,15 +149,18 @@ class StreamingDataset(IterableDataset):
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``.
* ``eos_id``, ``pad_id``, and ``columns`` naming a single integer-list
column are required; incompatible with ``transform``.
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.
``pack_sequences``, ignored otherwise. It must be reserved for padding:
padding positions retain the preceding document's ``doc_id`` (or zero
in an all-padding block), so callers must mask them separately using
``input_ids == pad_id``.
blocks_per_epoch:
Total number of packed blocks emitted globally per epoch. Required with
``pack_sequences``. An integer must be divisible by ``num_splits``.
@@ -165,10 +168,10 @@ class StreamingDataset(IterableDataset):
blocks: exhausted splits emit padding, while tokens beyond the budget
are left out of the epoch. This fixed per-split budget keeps packed
iteration and checkpoints independent of rank and worker topology.
Pass ``"auto"`` to calculate a corpus-level budget by scanning only the
token column in bounded batches. For large-scale training, materialize a
``token_count`` column, calculate an explicit block budget once, and pass
that integer instead.
Pass ``"auto"`` to estimate a corpus-level budget from a bounded sample
of token lists. The estimate may be inaccurate. For large-scale training,
materialize a ``token_count`` column, calculate an explicit block budget
once, and pass that integer instead.
on_transform_error:
What to do when the transform raises an exception:
@@ -299,6 +302,11 @@ class StreamingDataset(IterableDataset):
f"pack_sequences requires a list-typed token column; "
f"{columns[0]} has type {field.type}"
)
if not pa.types.is_integer(field.type.value_type):
raise ValueError(
"pack_sequences requires a token column with integer values; "
f"{columns[0]} has value type {field.type.value_type}"
)
elif blocks_per_epoch is not None:
raise ValueError("blocks_per_epoch requires pack_sequences")
if on_transform_error not in ("raise", "skip", "warn") and not callable(
@@ -331,7 +339,7 @@ class StreamingDataset(IterableDataset):
self._connection_factory = connection_factory
self._worker_info_override = worker_info_override
# Packing resume state: documents consumed and partial-block token buffers.
# Packing resume state: permutation positions and partial-block buffers.
self._pack_consumed: list[int] = [0] * num_splits
self._pack_buffers: dict[int, dict[str, list[int]]] = {}
self._pack_blocks_emitted: list[int] = [0] * num_splits
@@ -396,7 +404,7 @@ class StreamingDataset(IterableDataset):
)
def _estimate_blocks_per_epoch(self) -> int:
"""Calculate a fixed packed-block budget from the complete token column."""
"""Estimate a fixed packed-block budget from a bounded token sample."""
if self._pack_sequences is None or not self._columns:
raise RuntimeError(
"packing must be configured before estimating its budget"
@@ -404,47 +412,63 @@ class StreamingDataset(IterableDataset):
pack_len = self._pack_sequences
token_column = self._columns[0]
sample_cap_per_split = max(1, 100_000 // self._num_splits)
sampled_tokens = 0
total_sampled = 0
total_rows = 0
rng = random.Random(self._shuffle_seed)
warnings.warn(
"blocks_per_epoch='auto' scans the complete token column on every "
"distributed rank. For large-scale training, materialize a token_count "
"column, calculate blocks_per_epoch once, and pass that integer instead.",
# TODO: Link to docs example that shows how to do this
"blocks_per_epoch='auto' uses an approximate token-count sample; "
"pass an explicit value for exact epoch sizing",
stacklevel=3,
)
fragments = list(self._table.to_lance().get_fragments())
def fragment_total(fragment) -> tuple[int, int]:
token_count = 0
row_count = 0
scanner = fragment.scanner(
columns=[token_column],
filter=self._filter,
batch_size=2048,
# TODO: Replace this fallback with Lance's dedicated exact token-count
# estimation API once it is available.
for split in range(self._num_splits):
permutation = Permutation.from_tables(
self._table, self._perm_table, split=split
)
for batch in scanner.to_batches():
permutation = permutation.select_columns([token_column])
permutation = permutation.with_transform(Transforms.arrow2arrow)
split_rows = permutation.num_rows
if split_rows == 0:
raise ValueError(
"blocks_per_epoch='auto' cannot estimate an empty dataset"
)
# Sample roughly 1% from each logical split, with at least one row
# per split and a global target cap of 100,000 rows.
sample_rows = min(
split_rows,
max(1, min((split_rows + 99) // 100, sample_cap_per_split)),
)
sample_offsets = sorted(rng.sample(range(split_rows), sample_rows))
sample_batch_size = max(1, self._read_batch_size)
for start in range(0, sample_rows, sample_batch_size):
batch = permutation.__getitems__(
sample_offsets[start : start + sample_batch_size]
)
lengths = pc.list_value_length(batch.column(0))
if lengths.null_count:
raise ValueError("pack_sequences does not support null token lists")
token_count += int(pc.sum(lengths).as_py())
row_count += batch.num_rows
return token_count, row_count
sampled_tokens += int(pc.sum(lengths).as_py())
token_count = 0
row_count = 0
with ThreadPoolExecutor(max_workers=min(16, len(fragments))) as pool:
for fragment_tokens, fragment_rows in pool.map(fragment_total, fragments):
token_count += fragment_tokens
row_count += fragment_rows
total_sampled += sample_rows
total_rows += split_rows
# Each document contributes one EOS token. Round down to a complete
# split cycle, favoring a short unused tail over all-padding blocks.
blocks = (token_count + row_count) // pack_len
blocks_per_epoch = max(
# Pool the samples into one global average. Each document contributes
# one EOS token. Round down to a complete split cycle so an inaccurate
# estimate favors an unused tail over all-padding blocks.
estimated_tokens = (
(sampled_tokens + total_sampled) * total_rows // total_sampled
)
blocks = estimated_tokens // pack_len
return max(
self._num_splits,
blocks - blocks % self._num_splits,
)
return blocks_per_epoch
def _resolve_my_splits(self) -> list[int]:
"""Return the split indices this instance should read in __iter__."""
@@ -497,9 +521,8 @@ class StreamingDataset(IterableDataset):
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(Transforms.arrow2arrow)
# Packing tracks documents consumed per split. Row mode tracks
# absolute permutation positions so transform failures can skip
# rows without making resume repeat them.
# 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
@@ -530,7 +553,12 @@ class StreamingDataset(IterableDataset):
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())
token_column = batch.column(0)
if token_column.null_count or token_column.values.null_count:
raise ValueError(
"pack_sequences does not support null token lists or values"
)
return cast(list[list[int]], token_column.to_pylist())
final_transform = arrow_tokens
else:
@@ -694,7 +722,23 @@ class StreamingDataset(IterableDataset):
else:
break # split exhausted
# sequence packing helpers
def _update_stats(*, idle: bool = False) -> None:
"""Refresh pipeline statistics visible to the parent process."""
ws = self._worker_stats
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
ws[1] = (
0
if idle
else sum(batch.num_rows for q in raw_batches for _, batch in q)
)
ws[2] = 0 if idle else sum(len(q) for q in cooked)
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
# Sequence-packing helpers
pack_len = cast(int, self._pack_sequences)
eos_id = cast(int, self._eos_id)
pad_id = cast(int, self._pad_id)
@@ -707,12 +751,12 @@ class StreamingDataset(IterableDataset):
pack_buffers = deepcopy(self._pack_buffers)
pack_blocks_emitted = list(self._pack_blocks_emitted)
def _pack_buf(i: int) -> dict[str, list[int]]:
def _pack_buffer(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)
buf = _pack_buffer(i)
while len(buf["tokens"]) < pack_len:
_ensure_cooked(i)
if not cooked[i]:
@@ -721,13 +765,12 @@ class StreamingDataset(IterableDataset):
pos, tokens = cooked[i].popleft()
buf["tokens"].extend(tokens)
buf["tokens"].append(eos_id)
pack_consumed[my_splits[i]] += 1
pack_consumed[my_splits[i]] = pos + 1
local_consumed[i] += 1
pos_consumed[i] = pos + 1
_advance(i)
def _emit_block(i: int) -> dict[str, Any]:
buf = _pack_buf(i)
buf = _pack_buffer(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.
@@ -759,20 +802,6 @@ class StreamingDataset(IterableDataset):
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)
ws[7] = self._rows_skipped
try:
for i in range(n):
_fill_io(i)
@@ -797,10 +826,9 @@ class StreamingDataset(IterableDataset):
_fill_block(i)
for i in range(n):
tokens = _pack_buf(i)["tokens"]
tokens.extend([pad_id] * (pack_len - len(tokens)))
for i in range(n):
tokens = _pack_buffer(i)["tokens"]
if len(tokens) < pack_len:
tokens.extend([pad_id] * (pack_len - len(tokens)))
block = _emit_block(i)
pack_blocks_emitted[my_splits[i]] += 1
if i == n - 1:
@@ -850,15 +878,7 @@ class StreamingDataset(IterableDataset):
# when iteration ends mid-cycle (e.g. a split whose rows
# were all skipped before completing a single cycle), so
# counters like rows_skipped would otherwise be stale.
ws = self._worker_stats
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
ws[1] = 0 # queue-depth properties document 0 when idle
ws[2] = 0
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
_update_stats(idle=True)
self._raw_batches_ref = None
self._cooked_ref = None
self._fetch_head_ref = None
@@ -1121,7 +1141,7 @@ class StreamingDataset(IterableDataset):
For row mode, the elementwise maximum of permutation positions recovers
splits advanced by different ranks after transform failures. For packed
mode, the state that emitted the most blocks for each logical split
supplies that split's document count and partial token buffer. Packed
supplies that split's permutation position and partial token buffer. Packed
states must cover every rank or worker at the same global step.
Raises ``ValueError`` if the states are empty, were not produced by
@@ -1167,31 +1187,23 @@ class StreamingDataset(IterableDataset):
raise ValueError("merge_state_dicts requires at least one state dict")
first = states[0]
packed = "pack_buffers" in first
config_keys = ["shuffle_seed", "num_splits", "epoch"]
if packed:
config_keys.extend(
["pack_sequences", "eos_id", "pad_id", "blocks_per_epoch"]
)
for state in states[1:]:
for key in ("shuffle_seed", "num_splits", "epoch"):
if ("pack_buffers" in state) != packed:
raise ValueError("cannot merge packed and unpacked state dicts")
for key in config_keys:
if state[key] != first[key]:
raise ValueError(
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
if ("pack_buffers" in state) != packed:
raise ValueError("cannot merge packed and unpacked state dicts")
if packed:
config_keys = (
"pack_sequences",
"eos_id",
"pad_id",
"blocks_per_epoch",
)
for state in states[1:]:
for key in config_keys:
if state[key] != first[key]:
raise ValueError(
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
num_splits = first["num_splits"]
for state in states:
for key in (
+19 -3
View File
@@ -1998,7 +1998,7 @@ def test_pack_sequences_pads_lagging_splits(tmp_path):
assert sharded == input_ids
def test_pack_sequences_auto_scans_filtered_token_column(tmp_path):
def test_pack_sequences_auto_estimates_filtered_token_column(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"tokens",
@@ -2018,7 +2018,7 @@ def test_pack_sequences_auto_scans_filtered_token_column(tmp_path):
)
)
with pytest.warns(UserWarning, match="complete token column"):
with pytest.warns(UserWarning, match="approximate token-count sample"):
dataset = _packed_dataset(
table,
5,
@@ -2067,7 +2067,7 @@ def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path):
]
def test_pack_sequences_validates_padding_id(tmp_path):
def test_pack_sequences_validates_configuration_and_tokens(tmp_path):
table = _create_token_table(tmp_path, [[1, 2]])
with pytest.raises(ValueError, match="pad_id is required"):
@@ -2100,6 +2100,22 @@ def test_pack_sequences_validates_padding_id(tmp_path):
with pytest.raises(ValueError, match="pad_id mismatch"):
resumed.load_state_dict(checkpoint)
float_db = lancedb.connect(tmp_path / "float")
float_table = float_db.create_table(
"tokens",
pa.table({"tokens": pa.array([[1.5, 2.5]], type=pa.list_(pa.float64()))}),
)
with pytest.raises(ValueError, match="token column with integer values"):
_packed_dataset(float_table, 4, blocks_per_epoch=1)
null_db = lancedb.connect(tmp_path / "null")
null_table = null_db.create_table(
"tokens",
pa.table({"tokens": pa.array([None], type=pa.list_(pa.int64()))}),
)
with pytest.raises(ValueError, match="does not support null token lists"):
list(_packed_dataset(null_table, 4, blocks_per_epoch=1))
# ---------------------------------------------------------------------------
# Doc examples — each test mirrors the code snippet in index.mdx so that