DesignGym / train_data.py
AdithyaSK's picture
AdithyaSK HF Staff
Serve Crello corpus aafdb0e86816 from the mounted bucket
df84e45 verified
Raw History Blame Contribute Delete
4.84 kB
"""SFT export: expert trajectories as tool-calling chat transcripts.
Each Crello template becomes one conversation per (surface, mode): the
environment's own instructions as the system message, the target image (or
the brief) as the user turn, then the expert's tool calls one per assistant
turn, each followed by the tool result the environment actually returned.
Results are real because the trajectory is replayed through the environment
in-process, which also asserts that every exported trajectory scores 1.0.
The format is the OpenAI-style `messages` + `tools` that TRL's SFTTrainer and
most chat templates accept; images are written next to the JSONL and
referenced by relative path.
"""
from __future__ import annotations
import json
from pathlib import Path
from openenv.core.env_server.mcp_types import CallToolAction, ListToolsAction
try:
from .core import expert
from .core.pack import TaskPack
from .server.environment import DesignCanvasEnvironment
except ImportError:
from core import expert
from core.pack import TaskPack
from server.environment import DesignCanvasEnvironment
def _tool_schema(t) -> dict:
return {"type": "function", "function": {"name": t.name, "description": t.description,
"parameters": t.input_schema}}
def _result_text(obs) -> str:
if obs.error is not None:
return f"Error: {obs.error.message}"
res = obs.result
# Over the wire this is a dict; in-process it is FastMCP's CallToolResult.
content = (res.get("content") if isinstance(res, dict) else getattr(res, "content", None)) or []
texts = []
for c in content:
kind = c.get("type") if isinstance(c, dict) else getattr(c, "type", None)
if kind == "text":
texts.append(c.get("text", "") if isinstance(c, dict) else c.text)
return "\n".join(texts) or "(image)"
def trajectory(pack: TaskPack, task: dict, surface: str, mode: str, image_dir: Path) -> dict:
env = DesignCanvasEnvironment(pack)
obs = env.reset(task_id=task["task_id"], surface=surface, mode=mode)
meta = obs.metadata
tools = [t for t in env.step(ListToolsAction()).tools if t.name != "start_task"]
gt = pack.gt_doc(task)
user: list[dict] = []
images: list[str] = []
if mode == "reference":
image_dir.mkdir(parents=True, exist_ok=True)
path = image_dir / f"{task['task_id']}.jpg"
if not path.exists():
path.write_bytes(pack.target_jpeg(task, 768))
images.append(f"{image_dir.name}/{path.name}")
user.append({"type": "image", "image": images[-1]})
user.append({"type": "text", "text": "This is the target design. Recreate it."})
else:
user.append({"type": "text", "text": "Build the design from the brief in your instructions."})
messages = [{"role": "system", "content": meta["instructions"]}, {"role": "user", "content": user}]
calls = expert.tool_calls(gt) if surface == "tools" else [
{"name": "write_html", "arguments": {"html": expert.html_page(gt)}},
{"name": "submit", "arguments": {}},
]
reward = None
for n, call in enumerate(calls):
call_id = f"call_{n}"
messages.append({"role": "assistant", "content": "", "tool_calls": [
{"id": call_id, "type": "function",
"function": {"name": call["name"], "arguments": json.dumps(call["arguments"], ensure_ascii=False)}}]})
out = env.step(CallToolAction(tool_name=call["name"], arguments=call["arguments"]))
messages.append({"role": "tool", "tool_call_id": call_id, "name": call["name"],
"content": _result_text(out)})
if out.done:
reward = out.reward
env.close()
if reward is None or reward < 0.9999:
raise RuntimeError(f"expert trajectory for {task['task_id']} scored {reward}")
return {"messages": messages, "tools": [_tool_schema(t) for t in tools], "images": images,
"meta": {"task_id": task["task_id"], "surface": surface, "mode": mode,
"reward": reward, "prompt_version": meta.get("prompt_version")}}
def export(pack: TaskPack, split: str, out: Path, surfaces=("tools", "html"),
modes=("reference", "description"), limit: int | None = None) -> int:
out = Path(out)
out.parent.mkdir(parents=True, exist_ok=True)
image_dir = out.parent / f"{out.stem}_images"
n = 0
with open(out, "w") as fh:
for i in range(min(pack.count(split), limit or 10**9)):
task = pack.at(split, i)
for surface in surfaces:
for mode in modes:
fh.write(json.dumps(trajectory(pack, task, surface, mode, image_dir),
ensure_ascii=False) + "\n")
n += 1
return n