# 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 torch
from torchrl.render.artifacts import write_render_artifact
from torchrl.render.backends import MujocoStateReader
from torchrl.render.checkpoint import (
checkpoint_hash,
infer_state_dict,
load_checkpoint,
save_render_checkpoint,
)
from torchrl.render.config import (
CameraLayout,
EnvBackendName,
ExplorationMode,
FrameBundle,
key_to_string,
NotebookRenderBackendName,
NotebookRolloutMode,
parse_nested_key,
RenderBackendName,
RenderConfig,
RenderEnvSpec,
RenderFormat,
RenderPolicySpec,
RenderResult,
)
from torchrl.render.env import (
add_step_counter,
make_render_env,
normalize_env,
seed_env,
)
from torchrl.render.import_utils import call_with_supported_kwargs, import_from_string
from torchrl.render.mujoco_wasm import (
display_mujoco_wasm_viewer,
extract_qpos_trajectory,
play_mujoco_wasm_trajectory,
send_mujoco_wasm_qpos,
write_mujoco_wasm_viewer,
)
from torchrl.render.policy import (
load_render_policy,
normalize_policy,
TensorDictPolicyAdapter,
)
from torchrl.render.rollout import collect_render_rollouts
__all__ = [
"CameraLayout",
"EnvBackendName",
"ExplorationMode",
"FrameBundle",
"MujocoStateReader",
"NotebookRenderBackendName",
"NotebookRolloutMode",
"RenderBackendName",
"RenderConfig",
"RenderEnvSpec",
"RenderFormat",
"RenderPolicySpec",
"RenderResult",
"TensorDictPolicyAdapter",
"add_step_counter",
"call_with_supported_kwargs",
"checkpoint_hash",
"collect_render_rollouts",
"display_mujoco_wasm_viewer",
"extract_qpos_trajectory",
"import_from_string",
"infer_state_dict",
"key_to_string",
"load_checkpoint",
"load_render_policy",
"make_render_env",
"normalize_env",
"normalize_policy",
"parse_nested_key",
"play_mujoco_wasm_trajectory",
"render_policy",
"save_render_checkpoint",
"send_mujoco_wasm_qpos",
"seed_env",
"write_render_artifact",
"write_mujoco_wasm_viewer",
]
[docs]
def render_policy(config: RenderConfig) -> RenderResult:
"""Renders a policy according to ``config`` and writes the requested artifact.
Args:
config: Render configuration.
Returns:
The render result with trajectories, metadata, and artifact paths.
"""
if _defer_rollout_to_notebook(config):
return write_render_artifact(_empty_notebook_result(config), config)
device = torch.device(config.policy_device or config.device)
checkpoint = load_checkpoint(config.ckpt, map_location=device)
digest = checkpoint_hash(config.ckpt)
env = make_render_env(config, checkpoint=checkpoint)
try:
policy = load_render_policy(
config, env, checkpoint=checkpoint, checkpoint_digest=digest
)
result = collect_render_rollouts(env, policy, config)
result.metadata["checkpoint"] = {
"path": str(config.ckpt),
"sha256": digest,
}
return write_render_artifact(result, config)
finally:
close = getattr(env, "close", None)
if callable(close):
close()
def _defer_rollout_to_notebook(config: RenderConfig) -> bool:
return config.format == "ipynb" and config.notebook_rollout_mode == "live"
def _empty_notebook_result(config: RenderConfig) -> RenderResult:
metadata = {
"format": config.format,
"num_trajs": 0,
"requested_num_trajs": config.num_trajs,
"max_steps": config.max_steps,
"fps": config.fps,
"render_backend": config.render_backend,
"env_backend": config.env_backend,
"notebook_rollout_mode": config.notebook_rollout_mode,
"trajectories": [],
"warnings": [],
}
return RenderResult(
artifact_path=None,
trajectories=[],
frame_paths=[],
metadata=metadata,
warnings=[],
frames=[],
)