File size: 4,552 Bytes
e8f2c80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Generate model outputs for a VERL SFT validation parquet file."""

from __future__ import annotations

import argparse
import json
import re
from pathlib import Path

import pandas as pd
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

THINK_BLOCK_RE = re.compile(r"(?is)<think>(.*?)</think>")
ANSWER_BLOCK_RE = re.compile(r"(?is)<answer>(.*?)</answer>")


def split_tagged_response(text: str) -> dict[str, object]:
    think_match = THINK_BLOCK_RE.search(text)
    answer_match = ANSWER_BLOCK_RE.search(text)
    return {
        "think": think_match.group(1).strip() if think_match else "",
        "answer": answer_match.group(1).strip() if answer_match else "",
        "has_think": think_match is not None,
        "has_answer": answer_match is not None,
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", required=True, help="HF model directory or model id.")
    parser.add_argument("--adapter", help="Optional PEFT/LoRA adapter directory.")
    parser.add_argument("--tokenizer", help="Optional tokenizer path. Defaults to --model.")
    parser.add_argument("--data", default="data/hep_sft/val.parquet")
    parser.add_argument("--output", default="data/hep_sft/validation_outputs.jsonl")
    parser.add_argument("--limit", type=int, default=32)
    parser.add_argument("--max-new-tokens", type=int, default=512)
    args = parser.parse_args()

    model_path = Path(args.model)
    data_path = Path(args.data)
    output_path = Path(args.output)
    output_path.parent.mkdir(parents=True, exist_ok=True)

    tokenizer = AutoTokenizer.from_pretrained(args.tokenizer or model_path, trust_remote_code=True)
    model = AutoModelForCausalLM.from_pretrained(
        model_path,
        torch_dtype=torch.bfloat16,
        device_map="auto",
        trust_remote_code=True,
    )
    if args.adapter:
        from peft import PeftModel

        model = PeftModel.from_pretrained(model, args.adapter)
    model.eval()

    df = pd.read_parquet(data_path)
    if args.limit > 0:
        df = df.head(args.limit)

    with output_path.open("w") as out:
        for row in df.to_dict("records"):
            messages = row["messages"]
            prompt_messages = [m for m in messages if m["role"] != "assistant"]
            reference = next((m["content"] for m in messages if m["role"] == "assistant"), "")

            inputs = tokenizer.apply_chat_template(
                prompt_messages,
                add_generation_prompt=True,
                tokenize=True,
                return_tensors="pt",
            ).to(model.device)

            with torch.no_grad():
                generated = model.generate(
                    inputs,
                    max_new_tokens=args.max_new_tokens,
                    do_sample=False,
                    temperature=None,
                    top_p=None,
                    pad_token_id=tokenizer.eos_token_id,
                )

            output_ids = generated[0, inputs.shape[-1] :]
            prediction = tokenizer.decode(output_ids, skip_special_tokens=True).strip()
            prediction_parts = split_tagged_response(prediction)
            reference_parts = split_tagged_response(reference)
            record = {
                "id": row.get("id"),
                "arxiv_id": row.get("arxiv_id"),
                "task_type": row.get("task_type"),
                "target_type": row.get("target_type"),
                "target_category": row.get("target_category"),
                "target_process_id": row.get("target_process_id"),
                "target_background": row.get("target_background"),
                "prompt": prompt_messages[-1]["content"],
                "prediction": prediction,
                "reference": reference,
                "prediction_think": prediction_parts["think"],
                "prediction_answer": prediction_parts["answer"],
                "prediction_has_think": prediction_parts["has_think"],
                "prediction_has_answer": prediction_parts["has_answer"],
                "reference_think": reference_parts["think"],
                "reference_answer": reference_parts["answer"],
                "reference_has_think": reference_parts["has_think"],
                "reference_has_answer": reference_parts["has_answer"],
            }
            out.write(json.dumps(record, ensure_ascii=False) + "\n")
            print(f"wrote {record['id']}")

    print(f"Wrote {output_path}")


if __name__ == "__main__":
    main()