kimi-record / harness /scripts /extract_rft_data.py
simonycl's picture
Upload folder using huggingface_hub
7fde66e verified
Raw
History Blame Contribute Delete
5.81 kB
"""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()