Download train_data.py from FineEnvs/DesignGym: direct link, hf CLI and curl.
- Browser
- Download file 4.84 kB
-
https://huggingface.co/spaces/FineEnvs/DesignGym/resolve/main/train_data.py
- Command line
-
hf download hf://spaces/FineEnvs/DesignGym/train_data.py
-
curl -L -o train_data.py https://huggingface.co/spaces/FineEnvs/DesignGym/resolve/main/train_data.py
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 | |