MultimodalReasoning3B / data_preprocessing /process_limo_jsonl.py
gaaaaaaaaaaa's picture
Initial commit before training
61449ba verified
Raw History Blame Contribute Delete
2.47 kB
from __future__ import annotations
import argparse
from pathlib import Path
from typing import Iterable, Iterator
from .common import (
clean_text,
make_user_assistant_record,
read_jsonl,
render_records,
resolve_paths,
write_jsonl,
)
FINAL_ANSWER_MARKERS = (
"**Final Answer**",
"Final Answer:",
"Final answer:",
)
def _strip_final_answer_section(solution: str) -> str:
text = clean_text(solution)
for marker in FINAL_ANSWER_MARKERS:
if marker in text:
return text.split(marker, 1)[0].strip()
return text
def _build_assistant(solution: object, answer: object) -> str:
thinking = _strip_final_answer_section(clean_text(solution))
answer_text = clean_text(answer)
if "<think>" in thinking and "</think>" in thinking:
assistant = thinking
else:
assistant = f"<think>\n{thinking}\n</think>"
if answer_text:
assistant = f"{assistant}\n<answer>{answer_text}</answer>"
return assistant
def iter_records(
paths: str | Path | Iterable[str | Path],
*,
include_source: bool = False,
) -> Iterator[dict[str, str]]:
for path in resolve_paths(paths):
for _, row in read_jsonl(path):
record = make_user_assistant_record(
row.get("question"),
_build_assistant(row.get("solution"), row.get("answer")),
source=str(path),
include_source=include_source,
)
if record is not None:
yield record
def load_records(
paths: str | Path | Iterable[str | Path],
*,
include_source: bool = False,
) -> list[dict[str, str]]:
return list(iter_records(paths, include_source=include_source))
def main() -> None:
parser = argparse.ArgumentParser(
description="Normalize LIMO JSONL files to user/assistant rows."
)
parser.add_argument("inputs", nargs="+", help="Input JSONL path(s).")
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("--include-source", action="store_true")
args = parser.parse_args()
records = load_records(args.inputs, include_source=args.include_source)
if args.output:
write_jsonl(records, args.output)
if args.take:
print(render_records(records[: args.take]))
if __name__ == "__main__":
main()