File size: 4,272 Bytes
dfb775d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""Dataset curation — dispatch on `DataCfg.source` to the right adapter.

Three sources today:

- `hf` (default) — `datasets.load_dataset(streaming=True)`, requires `--extra ml`.
- `local` — JSONL files under `cfg.path`, pure stdlib.
- `mindx_dreams` — mindX dream-cycle JSONL corpus under `cfg.path`, pure stdlib.
  See `mindxtrain.data.sources.mindx_dreams`.

The HF path stays lazy-import so consumers without `--extra ml` can still
read configs that target other sources.
"""

from __future__ import annotations

import json
from collections.abc import Iterator
from pathlib import Path
from typing import Any

from mindxtrain.config.schema import DataCfg


def load_streaming_dataset(cfg: DataCfg) -> Iterator[dict[str, Any]]:
    """Yield rows for the configured `DataCfg.source` one at a time."""
    if cfg.source == "mindx_dreams":
        from mindxtrain.data.sources.mindx_dreams import (
            iter_mindx_dreams,
            iter_mindx_evolutions,
        )

        assert cfg.path is not None  # validated by DataCfg
        # Two-stream yield with a SHARED budget so max_samples caps the
        # combined corpus, not each stream individually. Consolidation rows
        # come first (richer base signal), evolutions follow when opted in.
        cap = cfg.max_samples
        emitted = 0
        remaining = (cap - emitted) if cap is not None else None
        for row in iter_mindx_dreams(cfg.path, max_samples=remaining):
            yield row
            emitted += 1
            if cap is not None and emitted >= cap:
                return
        if cfg.include_evolutions:
            remaining = (cap - emitted) if cap is not None else None
            try:
                for row in iter_mindx_evolutions(cfg.path, max_samples=remaining):
                    yield row
                    emitted += 1
                    if cap is not None and emitted >= cap:
                        return
            except FileNotFoundError:
                # mindx_dreams already succeeded above, so the root exists —
                # this would only fire on a path race. Swallow silently.
                pass
        return

    if cfg.source == "local":
        assert cfg.path is not None  # validated by DataCfg
        yield from _iter_local_jsonl(cfg.path, max_samples=cfg.max_samples)
        return

    if cfg.source == "lighthouse":
        msg = (
            "DataCfg.source='lighthouse' is reserved for shard-tar inputs; "
            "use `mindxtrain.storage.lighthouse.fetch` to materialize them locally first, "
            "then point a `local` source at the resulting directory."
        )
        raise NotImplementedError(msg)

    # source == "hf"
    try:
        from datasets import load_dataset
    except ImportError as exc:
        msg = "datasets not installed; run `uv sync --extra ml`."
        raise RuntimeError(msg) from exc

    kwargs: dict[str, Any] = {"streaming": cfg.streaming}
    revision = getattr(cfg, "revision", None)
    if revision:
        kwargs["revision"] = revision

    ds = load_dataset(cfg.hf_id, split=cfg.split or "train", **kwargs)
    emitted = 0
    for row in ds:
        yield row
        emitted += 1
        if cfg.max_samples is not None and emitted >= cfg.max_samples:
            return


def _iter_local_jsonl(
    path: Path,
    *,
    max_samples: int | None = None,
) -> Iterator[dict[str, Any]]:
    """Walk *.jsonl under `path` and yield parsed rows.

    Skips lines that fail to parse — local datasets are often hand-assembled
    and a single bad line shouldn't fail the run.
    """
    path = Path(path).expanduser()
    if path.is_file():
        files = [path]
    else:
        files = sorted(path.rglob("*.jsonl"))
    emitted = 0
    for f in files:
        with f.open("r", encoding="utf-8") as fh:
            for line in fh:
                line = line.strip()
                if not line:
                    continue
                try:
                    row = json.loads(line)
                except json.JSONDecodeError:
                    continue
                yield row
                emitted += 1
                if max_samples is not None and emitted >= max_samples:
                    return


__all__ = ["load_streaming_dataset"]