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 ""