| 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() |
|
|