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