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