File size: 6,140 Bytes
61449ba
 
 
 
 
 
 
 
 
 
 
3330bff
61449ba
 
 
 
 
 
3330bff
61449ba
 
 
 
 
 
 
 
 
 
 
3330bff
 
 
61449ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3330bff
 
 
 
 
 
 
61449ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
from __future__ import annotations

import argparse
from pathlib import Path

from .common import render_records
from .data_object import DataObject
from .process_complex_reasoning_bespoke_parquet import load_records as load_bespoke_records
from .process_limo_jsonl import load_records as load_limo_records
from .process_multimodal_jsonl import load_records as load_multimodal_records
from .process_reasoning_dpo_jsonl import load_records as load_dpo_records
from .process_sft_math_katex import load_records as load_sft_math_katex_records

DATASET_NAMES = (
    "multimodal",
    "reasoning_dpo",
    "limo",
    "complex_bespoke",
    'sft_math_katex',
    "all",
)


def _multimodal_paths(data_dir: Path) -> list[Path]:
    return sorted((data_dir / "MultimodalReasoning").glob("*.jsonl"))


def _reasoning_dpo_paths(data_dir: Path) -> list[Path]:
    return sorted((data_dir / "ReasoningDPO").glob("*.jsonl"))

def _sft_math_katex_path(data_dir: Path) -> list[Path]:
    return sorted((data_dir / "SFT_MATH_MULTI_TURN").glob("*.jsonl"))


def load_data_object(
    name: str = "all",
    *,
    data_dir: str | Path = "data",
    include_source: bool = False,
    include_system: bool = False,
    parquet_limit: int | None = None,
    skip_unavailable: bool = True,
) -> DataObject:
    data_path = Path(data_dir)

    if name == "multimodal":
        records = load_multimodal_records(
            _multimodal_paths(data_path),
            include_source=include_source,
        )
        return DataObject(records, name=name)
    
    if name == 'sft_math_katex':
        records = load_sft_math_katex_records(
            _sft_math_katex_path(data_path),
            include_source=include_source,
        )
        return DataObject(records, name=name)

    if name == "reasoning_dpo":
        records = load_dpo_records(
            _reasoning_dpo_paths(data_path),
            include_source=include_source,
        )
        return DataObject(records, name=name)

    if name == "limo":
        records = load_limo_records(
            data_path / "limo.jsonl",
            include_source=include_source,
        )
        return DataObject(records, name=name)

    if name == "complex_bespoke":
        try:
            records = load_bespoke_records(
                data_path / "ComplexReasoningBespoke.parquet",
                include_system=include_system,
                include_source=include_source,
                limit=parquet_limit,
            )
            return DataObject(records, name=name)
        except ImportError as exc:
            if not skip_unavailable:
                raise

            return DataObject([], name=name, warnings=[str(exc)])

    if name == "all":
        datasets = [
            load_data_object(
                "multimodal",
                data_dir=data_path,
                include_source=include_source,
                include_system=include_system,
                parquet_limit=parquet_limit,
                skip_unavailable=skip_unavailable,
            ),
            load_data_object(
                "reasoning_dpo",
                data_dir=data_path,
                include_source=include_source,
                include_system=include_system,
                parquet_limit=parquet_limit,
                skip_unavailable=skip_unavailable,
            ),
            load_data_object(
                "limo",
                data_dir=data_path,
                include_source=include_source,
                include_system=include_system,
                parquet_limit=parquet_limit,
                skip_unavailable=skip_unavailable,
            ),
            load_data_object(
                "complex_bespoke",
                data_dir=data_path,
                include_source=include_source,
                include_system=include_system,
                parquet_limit=parquet_limit,
                skip_unavailable=skip_unavailable,
            ),
        ]
        return DataObject.concat(datasets, name="all")

    raise ValueError(f"Unknown dataset name: {name}. Expected one of {DATASET_NAMES}.")


def load_data_objects(
    *,
    data_dir: str | Path = "data",
    include_source: bool = False,
    include_system: bool = False,
    parquet_limit: int | None = None,
    skip_unavailable: bool = True,
) -> dict[str, DataObject]:
    return {
        name: load_data_object(
            name,
            data_dir=data_dir,
            include_source=include_source,
            include_system=include_system,
            parquet_limit=parquet_limit,
            skip_unavailable=skip_unavailable,
        )
        for name in DATASET_NAMES
        if name != "all"
    }


def main() -> None:
    parser = argparse.ArgumentParser(
        description="Build a normalized DataObject from the local reasoning datasets."
    )
    parser.add_argument(
        "--dataset",
        default="all",
        choices=DATASET_NAMES,
        help="Dataset family to load.",
    )
    parser.add_argument("--data-dir", default="data")
    parser.add_argument("--output", help="Optional output JSONL path.")
    parser.add_argument("--take", type=int, default=0, help="Print the first N rows.")
    parser.add_argument("--shuffle", action="store_true")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--include-source", action="store_true")
    parser.add_argument("--include-system", action="store_true")
    parser.add_argument("--parquet-limit", type=int)
    parser.add_argument("--strict", action="store_true", help="Fail if parquet cannot load.")
    args = parser.parse_args()

    dataset = load_data_object(
        args.dataset,
        data_dir=args.data_dir,
        include_source=args.include_source,
        include_system=args.include_system,
        parquet_limit=args.parquet_limit,
        skip_unavailable=not args.strict,
    )

    if args.shuffle:
        dataset = dataset.shuffle(seed=args.seed)

    if args.output:
        dataset.to_jsonl(args.output)

    print(dataset.summary())

    if args.take:
        print(render_records(dataset.take(args.take)))


if __name__ == "__main__":
    main()