"""Extract successful on-policy RL rollouts into an RFT SFT dataset. Scans runs/rl_v*/run_default/rollouts/step_*/train/all/traces.jsonl, keeps traces with rewards.solved.score == 1.0 and clean completion, dedups per task (up to 2 shortest solves), converts node messages to plain OpenAI chat messages, length-filters with the Qwen3.5 tokenizer, and writes data/rft_v1_parquet/train.parquet in the same shape as data/sft_v2_parquet (messages, tools, source, n_tokens). """ import glob import json import os import re import sys sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from pi_prompt import TOOLS OUT = "/mnt/pvc/users/simon/agentptb/runs/d/workspace/data/rft_v1_parquet" MAX_TOKENS = 15000 MAX_PER_TASK = 2 MIN_NODES = 6 MAX_NODES = 200 def clean_messages(nodes): msgs = [] for n in nodes: m = n["message"] role = m.get("role") if role not in ("system", "user", "assistant", "tool"): return None out = {"role": role, "content": m.get("content")} if isinstance(out["content"], list): # user content arrives as [{type: text, text: ...}] parts; flatten to text out["content"] = "".join( p.get("text", "") if isinstance(p, dict) else str(p) for p in out["content"] ) if out["content"] is None: out["content"] = "" if role == "assistant" and m.get("tool_calls"): # traces store flat {id, name, arguments}; convert to OAI shape used by sft_v2 tcs = [] for tc in m["tool_calls"]: if "function" in tc: tcs.append({"id": tc.get("id", ""), "type": "function", "function": tc["function"]}) else: tcs.append({ "id": tc.get("id", ""), "type": "function", "function": {"name": tc.get("name", ""), "arguments": tc.get("arguments", "")}, }) out["tool_calls"] = tcs if role == "tool": out["tool_call_id"] = m.get("tool_call_id", "") out["name"] = m.get("name", "") msgs.append(out) if not msgs or msgs[0]["role"] != "system": return None if not any(m["role"] == "assistant" and m.get("tool_calls") for m in msgs): return None return msgs def args_to_dict(messages): out = [] for m in messages: m = dict(m) if m.get("tool_calls"): tcs = [] for tc in m["tool_calls"]: tc = dict(tc) fn = dict(tc["function"]) if isinstance(fn["arguments"], str): fn["arguments"] = json.loads(fn["arguments"]) tc["function"] = fn tcs.append(tc) m["tool_calls"] = tcs out.append(m) return out def main(): best = {} # (type, name) -> list of (n_nodes, run, step, messages) n_traces = n_solved = 0 for path in sorted(glob.glob( "/mnt/pvc/users/simon/agentptb/runs/d/workspace/runs/rl_v*/run_default/rollouts/step_*/train/all/traces.jsonl" )): run = re.search(r"rl_v\d+", path).group(0) step = path.split("/step_")[1].split("/")[0] with open(path) as f: for line in f: d = json.loads(line) n_traces += 1 score = ((d.get("rewards") or {}).get("solved") or {}).get("score", 0.0) if score != 1.0 or not d.get("ok"): continue if d.get("stop_condition") != "agent_completed": continue nodes = d.get("nodes") or [] if not (MIN_NODES <= len(nodes) <= MAX_NODES): continue msgs = clean_messages(nodes) if msgs is None: continue task = d.get("task") or {} tdata = task.get("data") or {} key = (task.get("type"), tdata.get("name") or tdata.get("instance_id") or d.get("id")) n_solved += 1 best.setdefault(key, []).append((len(nodes), run, int(step), msgs)) samples = [] for key, lst in best.items(): lst.sort(key=lambda x: x[0]) for n_nodes, run, step, msgs in lst[:MAX_PER_TASK]: samples.append({ "messages": msgs, "tools": json.dumps(TOOLS), "source": f"rft_{run}", }) print(f"traces={n_traces} solved={n_solved} unique_tasks={len(best)} samples={len(samples)}", flush=True) from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained("Qwen/Qwen3.5-9B-Base") keep, dropped, err = [], 0, 0 for s in samples: try: r = tok.apply_chat_template( args_to_dict(s["messages"]), tools=json.loads(s["tools"]), add_generation_prompt=False ) n = len(r["input_ids"]) except Exception: err += 1 continue if n <= MAX_TOKENS: s["n_tokens"] = n keep.append(s) else: dropped += 1 print(f"kept {len(keep)}, dropped_long {dropped}, template_errors {err}", flush=True) import random random.seed(0) random.shuffle(keep) from datasets import Dataset ds = Dataset.from_list(keep) os.makedirs(OUT, exist_ok=True) ds.to_parquet(os.path.join(OUT, "train.parquet")) import numpy as np from collections import Counter lens = ds["n_tokens"] print(Counter(ds["source"])) print(f"tokens: mean {np.mean(lens):.0f} p50 {np.percentile(lens,50):.0f} p90 {np.percentile(lens,90):.0f} max {max(lens)}") print("wrote", OUT, flush=True) if __name__ == "__main__": main()