Spaces:
Paused
Paused
Download api/simulate.py from chandrakiran06/rosdiff: direct link, hf CLI and curl.
- Browser
- Download file 8.23 kB
-
https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/api/simulate.py
- Command line
-
hf download hf://spaces/chandrakiran06/rosdiff/api/simulate.py
-
curl -L -o simulate.py https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/api/simulate.py
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>." | |
| class Attempt: | |
| raw_output: str | |
| mjcf: str | None | |
| validation: dict | |
| latency_ms: int | |
| model: str = "" | |
| usage: dict = field(default_factory=dict) | |
| 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 = "" | |
| 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 "" | |