rosdiff / api /simulate.py
Chandra Kiran
Free debugger on OpenCode Zen; separate code and scene model lists
536e671 unverified
Raw History Blame Contribute Delete
8.23 kB
"""Scenario description -> RAG-grounded MJCF scene, validated by loading it in MuJoCo."""
from __future__ import annotations
import json
import logging
import re
from dataclasses import dataclass, field
from functools import lru_cache
from pathlib import Path
from rag.store import Hit, RagStore, format_context
from .config import Settings
from .llm import LLM, LLMError, Message
from .mjcf_check import ValidationResult, extract_mjcf, validate_mjcf
from .robots import ROBOTS, robot_dir, robot_facts
log = logging.getLogger("rosdiff.simulate")
# Elements almost every scene uses; their exact attribute lists are always in the prompt.
CORE_ELEMENTS = (
"mujoco",
"option",
"visual-headlight",
"asset-texture",
"asset-material",
"asset-hfield",
"default",
"body",
"body-geom",
"body-joint",
"body-freejoint",
"body-inertial",
"body-site",
"body-light",
"body-camera",
"equality-weld",
)
SYSTEM_PROMPT = """You are an expert MuJoCo modeller. You write MJCF scene files for robotics simulation.
Rules:
- Output exactly one MJCF document in a single ```xml fenced block, then at most three short lines of notes.
- The root element is <mujoco model="...">.
- {robot_rule}
- Use only elements and attributes listed in the MJCF SCHEMA / reference context. Do not invent attributes.
- Build the environment from primitive geoms (plane, box, sphere, capsule, cylinder, ellipsoid) or a builtin hfield.
Do not reference external mesh/texture/hfield files; builtin textures (builtin="checker"/"gradient"/"flat") are fine.
- Give every body and geom you add a descriptive name. Add a floor, a light, and sensible friction/mass values.
- Objects meant to be pushed, carried or to fall get a <freejoint/>; fixed terrain and furniture do not.
- Angles: the robot file sets <compiler angle="radian"/>; give euler/axisangle values in radians.
"""
ROBOT_RULE = (
'Include the robot with <include file="{file}"/> as the first child of <mujoco>. Never copy, redefine, or rename '
"the robot's bodies, joints, actuators, meshes, compiler or keyframes."
)
NO_ROBOT_RULE = "There is no robot to include; do not use <include>."
@dataclass
class Attempt:
raw_output: str
mjcf: str | None
validation: dict
latency_ms: int
model: str = ""
usage: dict = field(default_factory=dict)
@dataclass
class SimulateResult:
valid: bool
mjcf: str | None
validation: ValidationResult
attempts: list[Attempt] = field(default_factory=list)
sources: list[str] = field(default_factory=list)
model: str = ""
@lru_cache
def load_schema(path: str) -> dict:
p = Path(path)
return json.loads(p.read_text()) if p.is_file() else {}
def schema_lines(schema: dict, keys) -> str:
rows = []
for key in keys:
entry = schema.get(key)
if entry and entry["attributes"]:
rows.append(f"<{entry['element']}> ({entry['path']}): {', '.join(entry['attributes'])}")
return "\n".join(rows)
def schema_hint(schema: dict, error: str) -> str:
"""For 'unrecognized attribute ... Element 'geom'' style errors: the element's valid attributes."""
m = re.search(r"Element '(\w+)'", error)
if not m:
return ""
element = m.group(1)
keys = [k for k, e in schema.items() if e["element"] == element]
return schema_lines(schema, keys)
def retrieve(store: RagStore, description: str) -> list[Hit]:
"""Scenario-specific sections of the MuJoCo docs (the robot and the schema are pinned separately)."""
return store.query("mujoco", description, k=6, where={"kind": "doc"})
def robot_example(settings: Settings, robot: str) -> str:
"""Menagerie's own scene.xml for the robot: the canonical way to include it."""
path = robot_dir(settings.menagerie_dir, ROBOTS[robot]) / "scene.xml"
return path.read_text() if path.is_file() else ""
def build_messages(
description: str, robot: str | None, context: str, facts: str, example: str = "", schema: str = ""
) -> list[Message]:
robot_rule = ROBOT_RULE.format(file=ROBOTS[robot].model_file) if robot else NO_ROBOT_RULE
user = []
if facts:
user.append(f"ROBOT FACTS (read from the model file):\n{facts}")
if example:
user.append(f"MENAGERIE EXAMPLE (scene.xml shipped with the robot):\n```xml\n{example.strip()}\n```")
if schema:
user.append(f"MJCF SCHEMA (the only valid attributes of common elements):\n{schema}")
user.append(f"REFERENCE CONTEXT (MuJoCo documentation):\n{context}")
user.append(f"SCENARIO:\n{description.strip()}")
user.append("Write the MJCF scene now.")
return [
Message("system", SYSTEM_PROMPT.format(robot_rule=robot_rule)),
Message("user", "\n\n".join(user)),
]
def generate_scene(
description: str,
robot: str | None,
llm: LLM,
store: RagStore,
settings: Settings,
max_attempts: int | None = None,
plan=None,
) -> SimulateResult:
"""`plan` (api.routing.Plan) picks the model per attempt: the planned tier first, one tier up per repair."""
from .routing import candidates_for_attempt
if robot is not None and robot not in ROBOTS:
raise ValueError(f"unknown robot {robot!r}; available: {sorted(ROBOTS)}")
hits = retrieve(store, description)
facts = robot_facts(str(settings.menagerie_dir), robot) if robot else ""
example = robot_example(settings, robot) if robot else ""
schema = load_schema(str(settings.mjcf_schema_file))
messages = build_messages(
description, robot, format_context(hits, max_chars=10000), facts, example, schema_lines(schema, CORE_ELEMENTS)
)
sources = ([f"mujoco_menagerie:{ROBOTS[robot].directory}/scene.xml"] if example else []) + [h.source for h in hits]
attempts: list[Attempt] = []
validation = ValidationResult(False, "extract", "no attempt made")
mjcf = None
for i in range(max(1, max_attempts or settings.mjcf_max_attempts)):
completion, model = _complete(llm, messages, [m.id for m in candidates_for_attempt(plan, i)] if plan else [])
mjcf = extract_mjcf(completion.text)
if mjcf is None:
validation = ValidationResult(False, "extract", "the response did not contain an MJCF document")
else:
validation = validate_mjcf(
mjcf, robot, settings.menagerie_dir, settings.mjcf_sim_seconds, settings.validate_timeout
)
attempts.append(
Attempt(
completion.text,
mjcf,
validation.to_dict(),
completion.latency_ms,
model or completion.model,
dict(completion.usage or {}),
)
)
if validation.valid:
break
# Repair: show the model its own output and exactly what MuJoCo said.
messages = messages + [
Message("assistant", completion.text),
Message(
"user",
f"{validation.feedback()}{_hint_block(schema, validation.error)}\n\n"
"Fix the problem and output the complete corrected MJCF document in one ```xml block. "
"Keep everything else the same.",
),
]
return SimulateResult(
valid=validation.valid,
mjcf=mjcf,
validation=validation,
attempts=attempts,
sources=sources,
model=attempts[-1].model if attempts else llm.model,
)
def _complete(llm: LLM, messages, candidates: list[str]):
"""Try each candidate model in turn (a free debugger first, then its paid fallback); raise the last error."""
if not candidates:
return llm.complete(messages), None
for n, model in enumerate(candidates):
try:
return llm.complete(messages, model=model), model
except LLMError as e:
if n == len(candidates) - 1:
raise
log.info("model %s failed (%s); falling back to %s", model, e, candidates[n + 1])
raise AssertionError("unreachable")
def _hint_block(schema: dict, error: str) -> str:
hint = schema_hint(schema, error)
return f"\nValid attributes for that element:\n{hint}" if hint else ""