zephon.work

Work-source abstractions that feed sample identifiers to the engine.

class ComponentOrder(*values)[source]

Bases: str, Enum

How to traverse samples within an individual mixture component.

AS_IS = 'as_is'
SHUFFLE = 'shuffle'
class MixtureReadConfig(mode=MixtureReadMode.WEIGHTED_ROUND_ROBIN, seed=None, precompute=False, within_component=ComponentOrder.AS_IS)[source]

Bases: object

Reader configuration controlling deterministic traversal of a chunk.

mode: MixtureReadMode = 'weighted_round_robin'
seed: int | None = None
precompute: bool = False
within_component: ComponentOrder = 'as_is'
class MixtureReadMode(*values)[source]

Bases: str, Enum

Supported sampling strategies when iterating over a WorkChunk.

WEIGHTED_RANDOM = 'weighted_random'
WEIGHTED_ROUND_ROBIN = 'weighted_round_robin'
class WorkChunk(components, seed=None, target_mixture=None)[source]

Bases: object

Bundle of sample identifiers handed to the engine for processing.

components stores each mixture component (for example, "German") and the ordered sample identifiers that belong to it. Mapping insertion order is used as a stable tie-breaker whenever behaviour depends on component ordering.

target_mixture, when set, is the per-component target the engine should deliver to downstream mixture correctors (ensure_mixture) instead of the counted composition. Token-aware work sources stamp the user’s token mixture here because their chunks deliberately carry a different sample composition (long-doc components contribute fewer pointers); the counted mixture stays composition-derived for within-chunk interleaving.

components: MutableMapping[str, list[tuple[int, int, int]]]
seed: int | None = None
target_mixture: Mapping[str, float] | None = None
property mixture: Mapping[str, float]

Return normalized mixture weights for the active components.

iter_samples(config=None)[source]

Iterate samples in mixture order, yielding (sample_id, component_name) tuples.

Return type:

Iterator[tuple[tuple[int, int, int], str]]

materialize_order(config)[source]
Return type:

list[tuple[tuple[int, int, int], str]]

sample_at(index, config=None)[source]
Return type:

tuple[tuple[int, int, int], str]

state_dict()[source]

Portable, JSON-friendly snapshot of this chunk (always the current version).

Return type:

dict[str, Any]

classmethod from_state(payload)[source]

Rebuild a WorkChunk from state_dict().

Return type:

WorkChunk

class WorkSource[source]

Bases: ABC

Abstract producer of WorkChunk instances for the engine.

next_chunk()[source]
Return type:

WorkChunk | None

property requires_token_priming: bool

Whether prime() must run before this source produces chunks.

prime(*, io_options=None, counting_spec=None, pre_tokenize_replay=None, mp_context=None)[source]

Driver-side calibration hook; no-op by default.

state_dict()[source]
Return type:

dict[str, Any]

load_state_dict(state)[source]
clone_for_lane(lane_id, canonical_replicas)[source]

Default clone strategy: config-based if available, else deepcopy.

Return type:

WorkSource

supports_indexing()[source]
Return type:

bool

sample_id_at(index)[source]
Return type:

tuple[int, int, int]

property datasets_by_id: Mapping[int, Dataset]
component_ids()[source]

Static component-name -> id vocabulary.

Fixed for the source’s lifetime and covering every component it will ever emit (curriculum stages included), so ids are a pure function of config across ranks, topology changes, and resumes. Ids are unique non-negative ints, not necessarily dense.

Return type:

Mapping[str, int]

chunk_size_hint()[source]

Return fixed chunk size if constant.

Used for deterministic resume validation. Default None.

Return type:

int | None

class MixtureSpec(weights)[source]

Bases: object

Mixture weight specification with shared validation/normalisation logic.

weights: Mapping[str, float]
validate_for(components)[source]

Ensure this mixture matches the provided component names exactly.

normalized_for(components)[source]

Return normalised weights ordered according to components.

Return type:

dict[str, float]

property normalized: dict[str, float]
class StaticMixtureWorkSource(datasets, mixture, chunk_size=16384, seed=0, shuffle_shards=True, shuffle_within_shard=False, shuffle_block_size=None, exhausted_policy=None, reshuffle_on_repeat=_Sentinel.UNSET, max_repeats=None, stop_after_passes=_STOP_AFTER_PASSES_UNSET, lane_assignment='permute', token_estimation=None)[source]

Bases: WorkSource

Emit SampleId triples from one or more datasets according to a mixture.

By default this source uses Bresenham-style fractional accumulators to distribute samples across datasets. Over many chunks the running average converges to the exact requested mixture weights with no 1/chunk_size granularity limitation. Individual chunks may have sparse allocations (a dataset may contribute 0 samples in a given chunk).

When loading a checkpoint that was created by the legacy fixed-quota allocator (pre-accumulator), the source automatically falls back to the original allocation strategy for deterministic continuation.

This source orchestrates ordering only. It builds per-dataset cursors from the provided Dataset descriptors and applies shuffling at three levels: - shard order within a dataset (shuffle_shards) - sample order within each shard (shuffle_within_shard) - optional block-based cross-shard shuffle (shuffle_block_size)

shuffle_block_size accepts:

  • None (default) — cross-shard block shuffle disabled

  • a positive int — explicit block size

  • "auto" — 8 * max_shard across all datasets in the mix

  • "global" — that dataset’s total sample count; the block buffer is O(total_samples)

Resolved per-cursor values are clamped to the dataset’s total.

No IO is performed here; fetching is handled by FetchOp which constructs an internal shard store using the dataset descriptors exposed via datasets_by_id.

The emitted identifiers have the shape (dataset_id, shard_id, sample_idx) where dataset_id is the position of the dataset in the constructor list.

exhausted_policy, reshuffle_on_repeat, and max_repeats each accept either a scalar (broadcast to every dataset) or a per-dataset mapping, so different datasets in the mixture can use different policies (e.g. one repeats forever as padding while another drives termination when exhausted).

stop_after_passes is the minimum number of full passes per dataset: every dataset repeats and the run stops once all have been traversed that many times. The slowest dataset is seen exactly this many times and triggers the stop; faster datasets loop more in the meantime to hold the mixing ratio. It is the default (stop_after_passes=1 — see each dataset once, then stop) when no exhaustion config is passed; giving exhausted_policy or max_repeats opts into the explicit per-dataset API and turns it off. An explicit stop_after_passes requires every dataset to repeat, so it rejects (ValueError) a non-"repeat" exhausted_policy or any max_repeats.

lane_assignment selects chunk->lane routing. "permute" (default) uses a seeded per-block permutation that breaks cadence-sharding resonance (see _lane_for_chunk()); "modulo" is plain g % canonical_replicas and reproduces pre-fix routing.

property total_samples: int | float

Samples this source will produce (float("inf") when unbounded).

clone_for_lane(lane_id, canonical_replicas)[source]

Lightweight clone that avoids deepcopying large cursor buffers.

The default WorkSource implementation performs a full deepcopy, which replicates every per-dataset order list. Those lists can be very large and immutable, so we instead share the order buffers and copy only the mutable cursor state.

Return type:

WorkSource

chunk_size_hint()[source]

Return fixed chunk size if constant.

Used for deterministic resume validation. Default None.

Return type:

int | None

property datasets_by_id: Mapping[int, Dataset]

Mapping of dataset_id to its Dataset descriptor.

The engine forwards this mapping to FetchOp so it can construct a dataset-aware shard store internally. The descriptors themselves do not perform IO.

property dataset_ids: Mapping[str, int]

Mapping of dataset name to the stable dataset_id used in SampleIds.

component_ids()[source]

Component ids are the dataset ids: fixed at construction.

Chunk component names are dataset names, so samples end up labelled with the same id that already identifies their dataset in SampleIds.

Return type:

Mapping[str, int]

property requires_token_priming: bool

True when this source still needs prime() before producing chunks.

The pipeline driver checks this hook before engine construction (and before pickling for DataLoader/MTP workers, so the primed ratios are inherited instead of re-measured). Primed checkpoints restore their ratios and report False; an unprimed template checkpoint restores to True and must be primed before producing.

prime(*, io_options=None, counting_spec=None, pre_tokenize_replay=None, mp_context=None)[source]

Calibrate per-dataset tokens/byte ratios before execution.

This is idempotent, does nothing in sample mode, and preserves ratios restored from a checkpoint.

next_chunk()[source]

Return the next chunk for a given canonical lane and worker.

This implementation follows a compute-everywhere-then-discard strategy:

  • Enumerate the global chunk stream deterministically using the existing chunking logic.

  • Assign each global chunk index g to a lane via _lane_for_chunk().

  • Discard non-matching chunks locally.

Return type:

WorkChunk | None

sample_id_at(index)[source]
Return type:

tuple[int, int, int]

supports_indexing()[source]
Return type:

bool

state_dict()[source]
Return type:

dict[str, Any]

load_state_dict(state)[source]
class TokenEstimation(primer='measure', measure=None, calibration_samples=2048, calibration_shards_min=4, calibration_shards_max=16, fallback_tokens_per_byte=0.25)[source]

Bases: object

Configuration for token-cost estimation (token-aware mixtures).

Passing an instance of this to a work source is what selects the token-aware allocation mode; leaving it None keeps sample-based mixing.

primer — how per-dataset tokens/byte ratios are obtained:

  • measure (default): calibration fetch + tokenize at prime time, deterministic in (seed, datasets, tokenize config).

  • a {dataset_name: tokens_per_byte} mapping: pins the listed datasets and measures the rest (partial pins merge over measured values).

  • a single float: one global tokens/byte for every dataset, no measurement (and no tokenizer needed).

measure — escape hatch replacing text extraction + tokenization entirely: a callable mapping a fetched payload to its delivered token count (weird schemas, VLM cost units). Only consulted when a dataset is actually measured.

calibration_samples — records measured per dataset (size-proportional draws for catalog-backed datasets, scattered offsets otherwise).

calibration_shards_min — shards scanned per dataset when its shards are homogeneous. The count is chosen per dataset from the catalog’s per-shard bytes/row spread (free metadata): census error scales as CV/sqrt(shards), so heterogeneous datasets automatically scan up to calibration_shards_max while uniform ones stay at the minimum. Each scanned shard costs a download + decode at prime time.

calibration_shards_max — upper bound for the adaptive shard count (also used when there is no catalog to read the spread from).

fallback_tokens_per_byte — ratio used when a dataset cannot be measured (no text found, measurement failed, or primer is a float). See DEFAULT_FALLBACK_TOKENS_PER_BYTE.

primer: Literal['measure'] | Mapping[str, float] | float = 'measure'
measure: Callable[[Any], int] | None = None
calibration_samples: int = 2048
calibration_shards_min: int = 4
calibration_shards_max: int = 16
fallback_tokens_per_byte: float = 0.25