Spaces:
Paused
Paused
File size: 8,230 Bytes
f850954 536e671 f850954 536e671 f850954 536e671 f850954 5b063c2 f850954 5b063c2 f850954 5b063c2 536e671 5b063c2 f850954 5b063c2 536e671 f850954 5b063c2 f850954 5b063c2 f850954 536e671 f850954 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 | """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 ""
|