Source code for zephon.work.base

# Copyright 2025 DatologyAI
# SPDX-License-Identifier: Apache-2.0

"""Abstract base definitions for work sources and chunks."""

import copy
from abc import ABC
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Iterator, Mapping, MutableMapping, Sequence

from zephon._internal.checkpoint import (
    WORK_CHUNK_VERSION as _WORK_CHUNK_VERSION,
)
from zephon._internal.checkpoint import (
    WorkChunkStateV2 as _WorkChunkStateV2,
)
from zephon._internal.utils.swrr import swrr_iterate as _swrr_iterate
from zephon.io.dataset import Dataset
from zephon.types import SampleId
from zephon.work.mixture import MixtureSpec

MixtureComponent = str
SourcedSampleId = tuple[SampleId, MixtureComponent]
SamplesPerComponent = MutableMapping[MixtureComponent, list[SampleId]]


[docs] class MixtureReadMode(str, Enum): """Supported sampling strategies when iterating over a ``WorkChunk``.""" WEIGHTED_RANDOM = "weighted_random" WEIGHTED_ROUND_ROBIN = "weighted_round_robin"
[docs] class ComponentOrder(str, Enum): """How to traverse samples within an individual mixture component.""" AS_IS = "as_is" SHUFFLE = "shuffle"
[docs] @dataclass(frozen=True) class MixtureReadConfig: """Reader configuration controlling deterministic traversal of a chunk.""" mode: MixtureReadMode = MixtureReadMode.WEIGHTED_ROUND_ROBIN seed: int | None = None precompute: bool = False # Whether to pre-compute sample order. Might cause performance spikes when requesting the first item. within_component: ComponentOrder = ( ComponentOrder.AS_IS ) # how to yield samples within the same component.
@dataclass class _Bucket: name: str items: Sequence[SampleId] # for trivial concatenate it: Iterator[SampleId] # for streaming weight: float
[docs] @dataclass class WorkChunk: """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 :attr:`mixture` stays composition-derived for within-chunk interleaving. """ components: SamplesPerComponent seed: int | None = None target_mixture: Mapping[str, float] | None = None ### INTERNAL ATTRIBUTES ### _order_cache: list[SourcedSampleId] | None = field( init=False, default=None, repr=False ) _order_cache_key: tuple | None = field(init=False, default=None, repr=False) _component_order: tuple[MixtureComponent, ...] = field(init=False, repr=False) _total_samples: int = field(init=False, repr=False) def __post_init__(self) -> None: self._component_order = tuple(self.components.keys()) self._total_samples = sum(len(items) for items in self.components.values()) if self.target_mixture is not None: # Normalize to canonical ratios (matches the `mixture` property) and # validate (non-empty, positive) in one step. self.target_mixture = MixtureSpec(self.target_mixture).normalized def __len__(self) -> int: return self._total_samples def __iter__(self) -> Iterator[SourcedSampleId]: yield from self.iter_samples() @property def mixture(self) -> Mapping[str, float]: """Return normalized mixture weights for the *active* components.""" comps = [c for c in self._component_order if self.components.get(c)] if not comps: return {} # Otherwise derive from counts (skip empty buckets so MixtureSpec stays >0). counts = {c: len(self.components[c]) for c in comps} return MixtureSpec(counts).normalized def _resolve_config(self, config: MixtureReadConfig | None) -> MixtureReadConfig: """Return an effective config with defaults applied (no mutation of input).""" base = config or MixtureReadConfig() # Fill seed from chunk if missing. seed = base.seed if base.seed is not None else self.seed # If you want a dynamic default (None for single-component), keep mode as-is here # and change MixtureReadConfig.mode to Optional[MixtureReadMode] with default None. return MixtureReadConfig( mode=base.mode, seed=seed, precompute=base.precompute, within_component=base.within_component, )
[docs] def iter_samples( self, config: MixtureReadConfig | None = None ) -> Iterator[SourcedSampleId]: """Iterate samples in mixture order, yielding (sample_id, component_name) tuples.""" cfg = self._resolve_config(config) if cfg.precompute: yield from self.materialize_order(cfg) return if self._total_samples == 0: return yield from self._iter_streaming(cfg)
[docs] def materialize_order(self, config: MixtureReadConfig) -> list[SourcedSampleId]: key = (config.mode, config.seed, config.within_component) if self._order_cache is not None and self._order_cache_key == key: return self._order_cache # Populate cache no_precompute_cfg = MixtureReadConfig( mode=config.mode, seed=config.seed, precompute=False, within_component=config.within_component, ) # avoid recursion order = list(self.iter_samples(no_precompute_cfg)) self._order_cache = order self._order_cache_key = key return self._order_cache
[docs] def sample_at( self, index: int, config: MixtureReadConfig | None = None ) -> SourcedSampleId: cfg = self._resolve_config(config) return self.materialize_order(cfg)[index]
def _iter_streaming(self, config: MixtureReadConfig) -> Iterator[SourcedSampleId]: if not self._total_samples: return buckets = self._build_buckets(config.seed, config.within_component) if not buckets: return # Single-bucket fast path (covers both None/WRR/Random cases) if len(buckets) == 1: name = buckets[0].name for sample_id in buckets[0].items: yield (sample_id, name) return if config.mode is MixtureReadMode.WEIGHTED_RANDOM: yield from self._emit_weighted_random(buckets, config.seed) elif config.mode is MixtureReadMode.WEIGHTED_ROUND_ROBIN: yield from self._emit_weighted_round_robin(buckets) else: raise ValueError(f"Unsupported mixture read mode: {config.mode}") def _build_buckets( self, seed: int | None, within_component: ComponentOrder, ) -> list[_Bucket]: from random import Random mix = self.mixture if not mix: return [] rng = Random() buckets: list[_Bucket] = [] for pos, name in enumerate(self._component_order): items = self.components.get(name, []) if not items: continue seq: list[SampleId] if within_component is ComponentOrder.SHUFFLE: seq = list(items) effective_seed = seed if seed is not None else self.seed if effective_seed is not None: rng.seed((effective_seed << 16) + pos) rng.shuffle(seq) else: seq = items buckets.append( _Bucket(name=name, items=seq, it=iter(seq), weight=mix[name]) ) return buckets # further modes: just random next sample (random without weights), trivial round robin def _emit_weighted_random( self, buckets: list[_Bucket], seed: int | None ) -> Iterator[SourcedSampleId]: from random import Random active = list(buckets) if not active: return weights = [b.weight for b in active] rng = Random(seed) while active: # pick a bucket index according to its weight idx = rng.choices(range(len(active)), weights=weights, k=1)[0] b = active[idx] try: yield (next(b.it), b.name) except StopIteration: # drop exhausted bucket and its weight del active[idx] del weights[idx] def _emit_weighted_round_robin( self, buckets: list[_Bucket] ) -> Iterator[SourcedSampleId]: """Smooth Weighted Round Robin (SWRR) using shared implementation. Delegates to swrr_iterate() which provides deterministic, proportional emission matching target weights over time. """ yield from _swrr_iterate( components={b.name: b.items for b in buckets}, weights={b.name: b.weight for b in buckets}, order=[b.name for b in buckets], )
[docs] def state_dict(self) -> dict[str, Any]: """Portable, JSON-friendly snapshot of this chunk (always the current version).""" comps_serial: list[tuple[str, list[list[int]]]] = [] for name in self._component_order: items = self.components.get(name, []) comps_serial.append((name, [list(sid) for sid in items])) state = _WorkChunkStateV2( version=_WORK_CHUNK_VERSION, seed=None if self.seed is None else int(self.seed), components=comps_serial, component_order=list(self._component_order), total_samples=int(self._total_samples), target_mixture=( None if self.target_mixture is None else {k: float(v) for k, v in self.target_mixture.items()} ), ) return state.to_dict()
[docs] @classmethod def from_state(cls, payload: Mapping[str, Any]) -> "WorkChunk": """Rebuild a WorkChunk from state_dict().""" ckpt = _WorkChunkStateV2.load(payload) comps: dict[str, list[SampleId]] = {} for name, items in ckpt.components: restored: list[SampleId] = [] for raw in items: if len(raw) != 3: raise ValueError(f"Bad SampleId for component {name}: {raw!r}") a, b, c = int(raw[0]), int(raw[1]), int(raw[2]) restored.append((a, b, c)) comps[name] = restored chunk = cls( components=comps, seed=ckpt.seed, target_mixture=ckpt.target_mixture, ) if ( ckpt.component_order and tuple(ckpt.component_order) != chunk._component_order ): ordered = {name: comps[name] for name in ckpt.component_order} chunk.components = ordered chunk.__post_init__() if ( ckpt.total_samples is not None and ckpt.total_samples != chunk._total_samples ): raise ValueError( f"WorkChunk total_samples mismatch: payload={ckpt.total_samples}, " f"computed={chunk._total_samples}" ) return chunk
[docs] class WorkSource(ABC): """Abstract producer of `WorkChunk` instances for the engine.""" def __init__(self): self._lane: int | None = None self._canon: int | None = None self._cloned = False
[docs] def next_chunk(self) -> WorkChunk | None: raise NotImplementedError()
@property def requires_token_priming(self) -> bool: """Whether :meth:`prime` must run before this source produces chunks.""" return False
[docs] def prime( self, *, io_options: Any = None, counting_spec: Any = None, pre_tokenize_replay: Any = None, mp_context: Any = None, ) -> None: """Driver-side calibration hook; no-op by default."""
[docs] def state_dict(self) -> dict[str, Any]: return { "lane_id": self._lane, "canonical_replicas": self._canon, "chunk_size_hint": self.chunk_size_hint(), }
[docs] def load_state_dict(self, state: dict[str, Any]) -> None: self._verify_base_state(int(state["lane_id"]), int(state["canonical_replicas"]))
def _verify_base_state(self, lane_id: int, canonical_replicas: int) -> None: """Verify lane/canonical-replicas identity against this instance. Subclasses that route state through a typed schema should call this directly (passing schema fields) instead of ``super().load_state_dict(state)`` so a future rename in the schema does not desync from the base. """ if lane_id != self._lane: raise RuntimeError("Lane mismatch loading LaneWorkSource state.") if canonical_replicas != self._canon: raise RuntimeError("canonical_replicas changed; migration required.") def _bind_lane(self, lane_id: int, canonical_replicas: int) -> None: """Bind the clone to a specific lane. Update any internal seed/cursors using lane_id if needed. """ if not self._cloned: raise RuntimeError( "State Error: The WorkSource should have been cloned internally before binding it." ) self._lane = int(lane_id) self._canon = int(canonical_replicas)
[docs] def clone_for_lane(self, lane_id: int, canonical_replicas: int) -> "WorkSource": """Default clone strategy: config-based if available, else deepcopy.""" if self._cloned: raise RuntimeError( "State Error: clone_for_lane should only be called on user-defined WorkSource instances." ) ws = copy.deepcopy(self) ws._cloned = True ws._bind_lane(lane_id, canonical_replicas) # TODO(MaxiBoether): Implement an alternative to deepcopy if worksources support it. # try: # cfg = self.config_dict() # ws = type(self).from_config(cfg) # except Exception: # fall back to deepcopy; ensure your subclass is deepcopy-safe # ws = copy.deepcopy(self) return ws
[docs] def supports_indexing(self) -> bool: raise NotImplementedError()
def __len__(self) -> int: raise NotImplementedError()
[docs] def sample_id_at(self, index: int) -> SampleId: raise NotImplementedError()
@property def datasets_by_id(self) -> Mapping[int, Dataset]: raise NotImplementedError()
[docs] def component_ids(self) -> Mapping[str, int]: """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. """ raise NotImplementedError()
[docs] def chunk_size_hint(self) -> int | None: """Return fixed chunk size if constant. Used for deterministic resume validation. Default None. """ return None
__all__ = [ "ComponentOrder", "MixtureComponent", "MixtureReadConfig", "MixtureReadMode", "SamplesPerComponent", "SourcedSampleId", "WorkChunk", "WorkSource", ]