MultimodalReasoning3B / data_preprocessing /build_data_object.py
gaaaaaaaaaaa's picture
Multiturn pre-SFT
3330bff verified
Raw
History Blame Contribute Delete
6.14 kB
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()