| """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): |
| |
| 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"): |
| |
| 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 = {} |
| 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() |
|
|