Download harness/scripts/extract_rft_data.py from agentic-ptb/kimi-record: direct link, hf CLI and curl.
- Browser
- Download file 5.81 kB
-
https://huggingface.co/agentic-ptb/kimi-record/resolve/main/harness/scripts/extract_rft_data.py
- Command line
-
hf download hf://agentic-ptb/kimi-record/harness/scripts/extract_rft_data.py
-
curl -L -o extract_rft_data.py https://huggingface.co/agentic-ptb/kimi-record/resolve/main/harness/scripts/extract_rft_data.py
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() | |