File size: 2,465 Bytes
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
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()