# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
from torchrl.render.config import RenderConfig, RenderResult
__all__ = ["build_notebook", "write_render_notebook"]
[docs]
def build_notebook(result: RenderResult, config: RenderConfig) -> dict[str, Any]:
"""Builds a minimal reproducible render report notebook."""
metadata_json = json.dumps(result.metadata, indent=2, sort_keys=True)
config_payload = config.to_dict()
config_payload["asset_dir"] = result.metadata.get("asset_dir", ".")
if result.metadata.get("asset_dir_absolute") is not None:
config_payload["asset_dir_absolute"] = result.metadata["asset_dir_absolute"]
if result.metadata.get("working_dir") is not None:
config_payload["working_dir"] = result.metadata["working_dir"]
config_json = json.dumps(config_payload, indent=2, sort_keys=True)
command = result.metadata.get("command", "rlrender ...")
cells = [
_markdown_cell(f"# TorchRL render report\n\nGenerated by `{command}`."),
_markdown_cell("## Metadata\n\n```json\n" + metadata_json + "\n```"),
_code_cell(
"import json\n\nrender_config = json.loads(" + repr(config_json) + ")"
),
_code_cell(
"from pathlib import Path\n"
"import json\n"
"import torch\n\n"
"asset_dir = Path(render_config.get('asset_dir', '.'))\n"
"if not (asset_dir / 'metadata.json').exists() and render_config.get('asset_dir_absolute'):\n"
" asset_dir = Path(render_config['asset_dir_absolute'])\n"
"metadata = json.loads((asset_dir / 'metadata.json').read_text())\n"
"metadata"
),
_code_cell(
"# Reconstruct the renderer config. The import paths execute trusted code.\n"
"from pathlib import Path\n"
"from torchrl.render import RenderConfig\n\n"
"def _resolve_file_import_spec(spec):\n"
" if not isinstance(spec, str) or ':' not in spec:\n"
" return spec\n"
" module, attr = spec.split(':', 1)\n"
" path = Path(module)\n"
" if path.is_absolute():\n"
" return spec\n"
" if path.suffix == '.py' or '/' in module or '\\\\' in module:\n"
" candidate = Path(render_config.get('working_dir', '.')) / path\n"
" if candidate.exists():\n"
" return f'{candidate}:{attr}'\n"
" return spec\n\n"
"for key in ('policy', 'env'):\n"
" render_config[key] = _resolve_file_import_spec(render_config[key])\n"
"working_dir = Path(render_config.get('working_dir', '.'))\n"
"ckpt_path = Path(render_config['ckpt']).expanduser()\n"
"if not ckpt_path.is_absolute():\n"
" ckpt_path = working_dir / ckpt_path\n"
"render_config['ckpt'] = str(ckpt_path)\n"
"runtime_config_keys = {'asset_dir', 'asset_dir_absolute', 'working_dir'}\n"
"cfg = RenderConfig(**{k: v for k, v in render_config.items() if k not in runtime_config_keys})\n"
"cfg"
),
_code_cell(
"# Load saved rollouts when present.\n"
"rollout_dir = asset_dir / 'rollouts'\n"
"rollouts = []\n"
"if rollout_dir.exists():\n"
" rollouts = [torch.load(path, weights_only=False) for path in sorted(rollout_dir.glob('traj_*.pt'))]\n"
"len(rollouts)"
),
]
if config.notebook_rollout_mode in ("live", "both"):
cells.extend(_live_rollout_cells())
else:
cells.append(
_code_cell(
"# Rerun the full renderer if you want to regenerate this artifact.\n"
"from torchrl.render import render_policy\n\n"
"# result = render_policy(cfg)"
)
)
video_paths = [path for path in result.frame_paths if path.suffix.lower() == ".mp4"]
if video_paths:
rel_paths = [str(path.name) for path in video_paths]
cells.append(_markdown_cell("## Static video previews"))
cells.append(
_code_cell(
"from IPython.display import Video, display\n"
f"video_paths = {rel_paths!r}\n"
"for video_path in video_paths:\n"
" display(Video(str(asset_dir / video_path), embed=False))"
)
)
if "mujoco_wasm" in result.metadata:
cells.extend(_mujoco_wasm_cells())
cells.append(
_markdown_cell(
"## Troubleshooting\n\n"
"If no frames were captured, construct Gymnasium environments with "
"`render_mode='rgb_array'` or enable TorchRL pixel observations with "
"`from_pixels=True`. For MuJoCo WASM notebook playback, provide "
"`--mujoco-model-path` and `--mujoco-qpos-key`, then rerun the "
"viewer cell before sending qpos trajectories."
)
)
return {
"cells": cells,
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3",
},
"language_info": {"name": "python", "pygments_lexer": "ipython3"},
},
"nbformat": 4,
"nbformat_minor": 5,
}
[docs]
def write_render_notebook(
result: RenderResult, config: RenderConfig, path: str | Path
) -> Path:
"""Writes a render report notebook as plain ipynb JSON."""
path = Path(path)
notebook = build_notebook(result, config)
path.write_text(
json.dumps(notebook, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
)
return path
def _markdown_cell(source: str) -> dict[str, Any]:
return {"cell_type": "markdown", "metadata": {}, "source": source.splitlines(True)}
def _code_cell(source: str) -> dict[str, Any]:
return {
"cell_type": "code",
"execution_count": None,
"metadata": {},
"outputs": [],
"source": source.splitlines(True),
}
def _mujoco_wasm_cells() -> list[dict[str, Any]]:
return [
_markdown_cell(
"## Interactive MuJoCo WASM playback\n\n"
"This notebook includes a generated local browser viewer for the MJCF "
"scene. The helper functions live in `torchrl.render.mujoco_wasm` so "
"the notebook stays small and reusable. The playback cell sends a "
"saved or live-generated qpos trajectory to the live viewer iframe."
),
_code_cell(
"from torchrl.render.mujoco_wasm import (\n"
" display_mujoco_wasm_viewer,\n"
" extract_qpos_trajectory,\n"
" play_mujoco_wasm_trajectory,\n"
" send_mujoco_wasm_qpos,\n"
")\n\n"
"mujoco_wasm = metadata['mujoco_wasm']\n"
"viewer_dir = asset_dir / mujoco_wasm['viewer_dir']\n"
"viewer_port = int(mujoco_wasm.get('viewer_port') or 5178)\n"
"viewer_process = display_mujoco_wasm_viewer(viewer_dir, port=viewer_port)\n"
"viewer_port = int(getattr(viewer_process, 'torchrl_mujoco_wasm_port', viewer_port))\n"
"viewer_origin = getattr(viewer_process, 'torchrl_mujoco_wasm_origin', f'http://127.0.0.1:{viewer_port}')\n"
),
_code_cell(
"qpos_key = mujoco_wasm.get('qpos_key') or render_config.get('mujoco_qpos_key')\n"
"if qpos_key is None:\n"
" raise ValueError('Set --mujoco-qpos-key to play saved rollouts in the WASM viewer.')\n"
"if not rollouts:\n"
" raise ValueError('No rollouts are available. Execute the saved-rollout loading cell or the in-notebook rollout cell.')\n\n"
"trajectory = extract_qpos_trajectory(rollouts[0], qpos_key=qpos_key)\n"
"play_mujoco_wasm_trajectory(\n"
" trajectory,\n"
" fps=float(render_config.get('fps', 30.0)),\n"
" port=viewer_port,\n"
" viewer_origin=viewer_origin,\n"
" wait=True,\n"
")\n"
),
_code_cell(
"# Send one pose interactively. The viewer accepts full model qpos vectors.\n"
"# send_mujoco_wasm_qpos(trajectory[0], port=viewer_port, viewer_origin=viewer_origin, wait=True)\n"
),
]
def _live_rollout_cells() -> list[dict[str, Any]]:
return [
_markdown_cell(
"## Generate rollouts in this notebook\n\n"
"Run this cell to generate trajectories from the checkpoint. "
"It constructs the configured environment and policy inside the kernel, "
"collects rollouts, and replaces the `rollouts` variable used by the "
"playback cells below. Edit `live_env_kwargs` before collecting to "
"change environment inputs without regenerating the notebook."
),
_code_cell(
"from dataclasses import replace\n\n"
"from torchrl.render import (\n"
" checkpoint_hash,\n"
" collect_render_rollouts,\n"
" load_checkpoint,\n"
" load_render_policy,\n"
" make_render_env,\n"
")\n\n"
"def collect_rollouts_in_notebook(config=None, *, env_kwargs=None):\n"
" config = cfg if config is None else config\n"
" if env_kwargs is not None:\n"
" config = replace(\n"
" config, env_kwargs={**config.env_kwargs, **env_kwargs}\n"
" )\n"
" checkpoint = load_checkpoint(\n"
" config.ckpt, map_location=config.policy_device or config.device\n"
" )\n"
" digest = checkpoint_hash(config.ckpt)\n"
" env = make_render_env(config, checkpoint=checkpoint)\n"
" try:\n"
" policy = load_render_policy(\n"
" config, env, checkpoint=checkpoint, checkpoint_digest=digest\n"
" )\n"
" return collect_render_rollouts(env, policy, config)\n"
" finally:\n"
" close = getattr(env, 'close', None)\n"
" if callable(close):\n"
" close()"
),
_code_cell(
"# Edit these values before collecting another rollout.\n"
"live_env_kwargs = dict(cfg.env_kwargs)\n"
"live_env_kwargs"
),
_code_cell(
"live_result = collect_rollouts_in_notebook(\n"
" env_kwargs=live_env_kwargs\n"
")\n"
"rollouts = live_result.trajectories\n"
"live_result.metadata"
),
]