Rate this Page

Source code for torchrl.render.notebook

# 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" ), ]