1
Fork 0

Merge branch 'dev' into docs/mkdocs

This commit is contained in:
Tibo De Peuter 2026-05-20 15:07:08 +02:00
commit 6d4a05e1bb
21 changed files with 893 additions and 89 deletions

View file

@ -3,7 +3,7 @@ import glob
import re import re
import shutil import shutil
folders_to_copy = ["src", "scripts", "configs" ] folders_to_copy = ["src", "scripts", "configs"]
for folder in folders_to_copy: for folder in folders_to_copy:
if os.path.exists(folder): if os.path.exists(folder):
shutil.copytree(folder, f"docs/{folder}", dirs_exist_ok=True) shutil.copytree(folder, f"docs/{folder}", dirs_exist_ok=True)
@ -14,7 +14,9 @@ for filepath in glob.glob("docs/**/*.md", recursive=True):
# RULE A: Fix links pointing OUT to src/, scripts/, or configs/ # RULE A: Fix links pointing OUT to src/, scripts/, or configs/
# Logic: Because the folders were moved one level deeper, we remove exactly ONE '../' # Logic: Because the folders were moved one level deeper, we remove exactly ONE '../'
content = re.sub(r"\]\(\.\./((?:\.\./)*)(src|scripts|configs)/([^)]*)\)", r"](\1\2/\3)", content) content = re.sub(
r"\]\(\.\./((?:\.\./)*)(src|scripts|configs)/([^)]*)\)", r"](\1\2/\3)", content
)
# RULE B: Fix links pointing FROM the copied files back TO the original docs/ folder # RULE B: Fix links pointing FROM the copied files back TO the original docs/ folder
# Logic: Since these files are now inside docs/, the 'docs/' segment in the path is redundant. # Logic: Since these files are now inside docs/, the 'docs/' segment in the path is redundant.

4
.gitignore vendored
View file

@ -6,6 +6,7 @@ outputs/
multirun/ multirun/
metrics/ metrics/
adjacency_debug.txt adjacency_debug.txt
vids/
# Python-generated files # Python-generated files
__pycache__/ __pycache__/
@ -523,3 +524,6 @@ Network Trash Folder
Temporary Items Temporary Items
.apdisk .apdisk
*.pdf *.pdf
# plot directory
poster_plots/

View file

@ -23,7 +23,7 @@ morphology:
morph_mode: CENTRALIZED morph_mode: CENTRALIZED
experiment: experiment:
exp_name: "final-models/centralized/" exp_name: "final-models-v2/centralized/"
seed: 42 seed: 42
torch_deterministic: true torch_deterministic: true
cuda: true cuda: true
@ -34,8 +34,8 @@ logging:
save_checkpoints: true save_checkpoints: true
upload_final_model: true upload_final_model: true
upload_checkpoints: true upload_checkpoints: true
checkpoint_frequency: 20 checkpoint_frequency: 10
wandb_project_name: "final-models" wandb_project_name: "final-models-v2"
evaluation: evaluation:
evaluate_checkpoints: true evaluate_checkpoints: true

View file

@ -9,12 +9,14 @@ eval_seed: 0
# Cross-model comparison settings # Cross-model comparison settings
# We use 10 episodes to get a more robust average for the final poster results. # We use 10 episodes to get a more robust average for the final poster results.
comparison_base_seed: 0 comparison_base_seed: 0
comparison_num_episodes: 2 comparison_num_episodes: 10
comparison_output_csv: "runs/evaluation/comparison.csv" comparison_output_csv: "runs/evaluation/comparison.csv"
# Paths to the .cleanrl_model files to be compared (relative to workspace root). # Paths to the .cleanrl_model files to be compared (relative to workspace root).
comparison_models: comparison_models:
- "runs/input-space-2-arms/2026-05-02/08-14-58/final_model.flax" - "runs/final-v2-centralized/artifacts/12-19-01_checkpoint_v22/checkpoint_step_230.flax"
- "runs/final-v2-fully-conn/artifacts/14-02-00_checkpoint_v17/checkpoint_step_180.flax"
- "runs/final-v2-ring/artifacts/15-27-03_checkpoint_v21/checkpoint_step_220.flax"
# Path to the morphologies to evaluate against. # Path to the morphologies to evaluate against.
comparison_morphologies: comparison_morphologies:

View file

@ -26,7 +26,7 @@ morphology:
morph_mode: FULLY_CONNECTED morph_mode: FULLY_CONNECTED
experiment: experiment:
exp_name: "final-models/fully-connected/" exp_name: "final-models-v2/fully-connected/"
seed: 42 seed: 42
torch_deterministic: true torch_deterministic: true
cuda: true cuda: true
@ -37,8 +37,8 @@ logging:
save_checkpoints: true save_checkpoints: true
upload_final_model: true upload_final_model: true
upload_checkpoints: true upload_checkpoints: true
checkpoint_frequency: 20 checkpoint_frequency: 10
wandb_project_name: "final-models" wandb_project_name: "final-models-v2"
evaluation: evaluation:
evaluate_checkpoints: true evaluate_checkpoints: true

View file

@ -26,7 +26,7 @@ morphology:
morph_mode: RING morph_mode: RING
experiment: experiment:
exp_name: "final-models/ring/" exp_name: "final-models-v2/ring/"
seed: 42 seed: 42
torch_deterministic: true torch_deterministic: true
cuda: true cuda: true
@ -37,8 +37,8 @@ logging:
save_checkpoints: true save_checkpoints: true
upload_final_model: true upload_final_model: true
upload_checkpoints: true upload_checkpoints: true
checkpoint_frequency: 20 checkpoint_frequency: 10
wandb_project_name: "final-models" wandb_project_name: "final-models-v2"
evaluation: evaluation:
evaluate_checkpoints: true evaluate_checkpoints: true

View file

@ -21,6 +21,10 @@ video_output_path: null
# Camera ID to use for video recording (1 is usually the close-up camera) # Camera ID to use for video recording (1 is usually the close-up camera)
camera_id: 1 camera_id: 1
video_width: 640
video_height: 80
video_fps: 60
# Optional override for the metadata YAML file path. # Optional override for the metadata YAML file path.
# If null, the script looks for `<model_name>_metadata.yaml` alongside the model_path. # If null, the script looks for `<model_name>_metadata.yaml` alongside the model_path.
metadata_path: null metadata_path: null

View file

@ -38,4 +38,17 @@ uv run scripts/simulate.py \
Videos and evaluation metadata are stored in timestamped folders alongside the model: Videos and evaluation metadata are stored in timestamped folders alongside the model:
`runs/your_run/final_model_evaluations/eval_<timestamp>/simulation.mp4` `runs/your_run/final_model_evaluations/eval_<timestamp>/simulation.mp4`
### Top-Down and Follow Cameras
Using the following script, you can render a top-down and follow camera view for multiple models at once:
```bash
uv run scripts/poster_visualisations/render_poster_videos.py \
runs/final-models/centralized/.../final_model.flax \
runs/final-models/fully-connected/.../final_model.flax \
runs/final-models/ring/.../final_model.flax \
--max-steps 10000 --width 640 --height 480 --fps 60 \
--output-root vids/poster/
```
For batch evaluation, checkpoint analysis, and cross-model architecture comparisons, see the **[Checkpoint & Model Evaluation Guide](./evaluation.md)**. For batch evaluation, checkpoint analysis, and cross-model architecture comparisons, see the **[Checkpoint & Model Evaluation Guide](./evaluation.md)**.

View file

@ -6,18 +6,17 @@ Rate, Distance Remaining).
""" """
import os import os
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from plot_config import ( from plot_config import (
COLORS,
apply_style,
BEST_PERFORMER_MARKER,
BEST_PERFORMER_TEXT,
BEST_PERFORMER_COLOR, BEST_PERFORMER_COLOR,
create_common_parser, BEST_PERFORMER_TEXT,
COLORS,
LEGEND_KWARGS, LEGEND_KWARGS,
apply_style,
create_common_parser,
) )
@ -84,7 +83,8 @@ def plot_grouped_bar(
fig, ax = plt.subplots(figsize=figsize) fig, ax = plt.subplots(figsize=figsize)
bar_width = 0.35 bar_width = 0.35
x_indices = np.arange(len(morphologies)) group_spacing = 1.3
x_indices = np.arange(len(morphologies)) * group_spacing
all_bars = {} all_bars = {}
all_means = [] all_means = []
@ -118,7 +118,7 @@ def plot_grouped_bar(
) )
all_bars[arch] = (x_pos, means, stds, bars) all_bars[arch] = (x_pos, means, stds, bars)
for m_idx, m in enumerate(morphologies): for m_idx, _ in enumerate(morphologies):
m_means = {arch: all_bars[arch][1][m_idx] for arch in architectures} m_means = {arch: all_bars[arch][1][m_idx] for arch in architectures}
best_arch = ( best_arch = (
max(m_means, key=m_means.get) if higher_is_better else min(m_means, key=m_means.get) max(m_means, key=m_means.get) if higher_is_better else min(m_means, key=m_means.get)
@ -144,12 +144,13 @@ def plot_grouped_bar(
x_ticks_pos = ( x_ticks_pos = (
x_indices x_indices
+ bar_width # center the label in the 3 bars
+ (bar_width / 2 if len(architectures) % 2 == 0 else 0) + (bar_width / 2 if len(architectures) % 2 == 0 else 0)
- (bar_width / 2 if len(architectures) == 2 else 0) - (bar_width / 2 if len(architectures) == 2 else 0)
) )
ax.set_xticks(x_ticks_pos) ax.set_xticks(x_ticks_pos)
ax.set_xticklabels([f"{m} Arms" for m in morphologies]) ax.set_xticklabels([f"{m} Arms" for m in morphologies])
ax.tick_params(axis="x", pad=25) # More padding for the squares ax.tick_params(axis="x") # More padding for the squares
# X-axis at zero # X-axis at zero
ax.axhline(0, color="black", linewidth=1.5) ax.axhline(0, color="black", linewidth=1.5)
@ -172,20 +173,7 @@ def plot_grouped_bar(
plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}")
) )
_add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) ax.legend(**LEGEND_KWARGS, ncol=len(architectures))
# Add custom legend entry for best performer
ax.plot(
[],
[],
marker=BEST_PERFORMER_MARKER,
color="w",
markerfacecolor=BEST_PERFORMER_COLOR,
markersize=15,
label="Best Performance",
ls="",
)
ax.legend(**LEGEND_KWARGS, ncol=len(architectures) + 1)
ax.set_facecolor("white") ax.set_facecolor("white")
fig.patch.set_facecolor("white") fig.patch.set_facecolor("white")
@ -305,22 +293,7 @@ def plot_grouped_bar_alt(
plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}")
) )
# In this alt plot, placeholders might be per architecture ax.legend(**LEGEND_KWARGS, ncol=len(morphologies))
_add_square_placeholders(
ax, x_indices, [arch.replace("_", "\n").title() for arch in architectures]
)
ax.plot(
[],
[],
marker=BEST_PERFORMER_MARKER,
color="w",
markerfacecolor=BEST_PERFORMER_COLOR,
markersize=15,
label="Best Performance",
ls="",
)
ax.legend(**LEGEND_KWARGS, ncol=len(morphologies) + 1)
ax.set_facecolor("white") ax.set_facecolor("white")
fig.patch.set_facecolor("white") fig.patch.set_facecolor("white")
@ -358,8 +331,8 @@ if __name__ == "__main__":
plot_grouped_bar( plot_grouped_bar(
df=df, df=df,
metric_col="approx_max_velocity", metric_col="approx_max_velocity",
ylabel="Max Forward Velocity (cm/s)", ylabel="",
title="Graceful Degradation: Velocity Across Morphologies", title="Maximal forward velocity (in cm/s)",
output_filename="poster_plot_velocity.png", output_filename="poster_plot_velocity.png",
output_dir=OUTPUT_DIR, output_dir=OUTPUT_DIR,
higher_is_better=True, higher_is_better=True,

View file

@ -44,20 +44,24 @@ class Columns(str, Enum):
# ... (rest of the file remains same, just need to update plotting functions and obtain_data) # ... (rest of the file remains same, just need to update plotting functions and obtain_data)
"""Column names expected in every evaluation CSV.""" """Column names expected in every evaluation CSV."""
CHECKPOINT = "checkpoint"
ARCH = "architecture" ARCH = "architecture"
TIMESTEPS = "total_trained_timesteps" TIMESTEPS = "trained_timesteps"
REWARD = "accumulated_reward" REWARD = "eval_return"
VELOCITY = "velocity" VELOCITY = "velocity"
EVAL_STEPS = "eval_steps"
FINAL_XY_DIST = "final_xy_dist"
INITIAL_XY_DIST = "initial_xy_dist"
REACHED_TARGET = "reached_target"
# Maps architecture display names to the path of their evaluation CSV. # Maps architecture display names to the path of their evaluation CSV.
# Update these paths once real evaluation data is available. # Update these paths once real evaluation data is available.
FILE_MAPPING: dict[str, str] = { FILE_MAPPING: dict[str, str] = {
"centralized 2 arms": "runs/dummy/dummy_centralized_2_arms.csv", # "centralized 2 arms": "runs/dummy/dummy_centralized_2_arms.csv",
"centralized 5 arms": "runs/dummy/dummy_centralized_5_arms.csv", "centralized 5 arms": "runs/final-v2-centralized/checkpoint_evaluation.csv",
"decentralized fully connected": "runs/dummy/dummy_decentralized_fully_connected.csv", "decentralized fully connected": "runs/final-v2-fully-conn/checkpoint_evaluation.csv",
"decentralized ring-level": "runs/dummy/dummy_decentralized_ring-level.csv", "decentralized ring-level": "runs/final-v2-ring/checkpoint_evaluation.csv",
"decentralized segment-level": "runs/dummy/dummy_decentralized_segment-level.csv",
} }
# Architecture profiles for dummy data generation: (max_reward, max_velocity, sigmoid_speed) # Architecture profiles for dummy data generation: (max_reward, max_velocity, sigmoid_speed)
@ -108,7 +112,14 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame:
Loads one CSV per architecture, injects the architecture name as a column, Loads one CSV per architecture, injects the architecture name as a column,
and returns the combined DataFrame with only the required columns. and returns the combined DataFrame with only the required columns.
""" """
required = [Columns.TIMESTEPS, Columns.REWARD, Columns.VELOCITY] required = [
Columns.CHECKPOINT,
Columns.TIMESTEPS,
Columns.REWARD,
Columns.INITIAL_XY_DIST,
Columns.FINAL_XY_DIST,
Columns.EVAL_STEPS,
]
dfs = [] dfs = []
for arch_name, filepath in file_mapping.items(): for arch_name, filepath in file_mapping.items():
@ -124,17 +135,30 @@ def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame:
continue continue
df = df[required].copy() df = df[required].copy()
df[Columns.VELOCITY] = (df[Columns.INITIAL_XY_DIST] - df[Columns.FINAL_XY_DIST]) / df[
Columns.EVAL_STEPS
]
df[Columns.ARCH] = arch_name df[Columns.ARCH] = arch_name
df[Columns.VELOCITY] = (df[Columns.INITIAL_XY_DIST] - df[Columns.FINAL_XY_DIST]) / df[
Columns.EVAL_STEPS
]
dfs.append(df) dfs.append(df)
return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame() return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame()
def _convergence_timestep(series: pd.Series, timesteps: pd.Series) -> float: def _convergence_timestep(
series: pd.Series, timesteps: pd.Series, checkpoints: pd.Series
) -> tuple[float, int, int]:
"""Returns the first timestep where the smoothed series reaches 95% of its peak.""" """Returns the first timestep where the smoothed series reaches 95% of its peak."""
smoothed = series.rolling(window=SMOOTHING_WINDOW, min_periods=1).mean() smoothed = series.rolling(window=SMOOTHING_WINDOW, min_periods=1).mean()
threshold = smoothed.max() * CONVERGENCE_THRESHOLD threshold = smoothed.max() * CONVERGENCE_THRESHOLD
return timesteps[smoothed >= threshold].iloc[0]
mask = smoothed >= threshold
first_idx = mask.idxmax()
return timesteps.loc[first_idx], first_idx, checkpoints.loc[first_idx]
def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame:
@ -144,21 +168,40 @@ def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame:
""" """
results = [] results = []
centralized_base = 0
for arch in df[Columns.ARCH].unique(): for arch in df[Columns.ARCH].unique():
arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS) arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS)
reward_timestep, reward_checkpoint_idx, reward_checkpoint = _convergence_timestep(
arch_data[Columns.REWARD],
arch_data[Columns.TIMESTEPS],
arch_data[Columns.CHECKPOINT],
)
velocity_timestep, velocity_checkpoint_idx, velocity_checkpoint = _convergence_timestep(
arch_data[Columns.VELOCITY],
arch_data[Columns.TIMESTEPS],
arch_data[Columns.CHECKPOINT],
)
results.append( results.append(
{ {
"Architecture": arch, "Architecture": arch,
"Reward_Convergence_Timestep": _convergence_timestep( "Reward_Convergence_Timestep": reward_timestep,
arch_data[Columns.REWARD], arch_data[Columns.TIMESTEPS] "Reward_Convergence_Checkpoint_Idx": reward_checkpoint_idx,
), "Reward_Convergence_Checkpoint": reward_checkpoint,
"Velocity_Convergence_Timestep": _convergence_timestep( "Velocity_Convergence_Timestep": velocity_timestep,
arch_data[Columns.VELOCITY], arch_data[Columns.TIMESTEPS] "Velocity_Convergence_Checkpoint_Idx": velocity_checkpoint_idx,
), "Velocity_Convergence_Checkpoint": velocity_checkpoint,
} }
) )
if arch == "centralized 5 arms":
centralized_base = reward_checkpoint
else:
print(arch, "speedup:", 1 - reward_checkpoint / centralized_base)
return pd.DataFrame(results) return pd.DataFrame(results)
@ -303,6 +346,7 @@ def plot_results(df: pd.DataFrame, results: pd.DataFrame, output_dir: str, **kwa
def obtain_data() -> pd.DataFrame: def obtain_data() -> pd.DataFrame:
"""Resolves the file mapping, falling back to generated dummy CSVs if needed.""" """Resolves the file mapping, falling back to generated dummy CSVs if needed."""
global USING_DUMMY_DATA global USING_DUMMY_DATA
if not any(os.path.exists(p) for p in FILE_MAPPING.values()): if not any(os.path.exists(p) for p in FILE_MAPPING.values()):
logger.info("No real evaluation files found. Generating dummy CSVs at expected locations.") logger.info("No real evaluation files found. Generating dummy CSVs at expected locations.")
generate_dummy_csvs(FILE_MAPPING) generate_dummy_csvs(FILE_MAPPING)
@ -319,6 +363,12 @@ def run_analysis(output_dir: str, **kwargs):
return return
results = analyze_convergence(df) results = analyze_convergence(df)
print(
results[
["Architecture", "Reward_Convergence_Checkpoint_Idx", "Reward_Convergence_Checkpoint"]
]
)
plot_results(df, results, output_dir, **kwargs) plot_results(df, results, output_dir, **kwargs)
logger.info("Analysis complete. Plots saved to disk.") logger.info("Analysis complete. Plots saved to disk.")

View file

@ -4,26 +4,24 @@ import matplotlib.pyplot as plt
# Shared Color Palette (Colorblind friendly, high contrast) # Shared Color Palette (Colorblind friendly, high contrast)
# Matches poster design # Matches poster design
COLORS = { COLORS = {
"CENTRALIZED": "#2B4162", # Deep Slate Blue "CENTRALIZED": "#0D567C", # Blue
"FULLY_CONNECTED": "#FA9F42", # Vibrant Orange "FULLY_CONNECTED": "#8C0E0F", # Reddish
"RING_LEVEL": "#4E937A", # Muted Teal "RING": "#FCB305", # Pale Yellow
"SEGMENT_LEVEL": "#B4436C", # Soft Red
"DECENTRALIZED": "#4E937A", # Default decentralized fallback
} }
def apply_style(font_size=28): def apply_style(font_size=36):
""" """
Applies the shared typography and aesthetic settings to Matplotlib. Applies the shared typography and aesthetic settings to Matplotlib.
""" """
plt.rcParams.update( plt.rcParams.update(
{ {
"font.size": font_size, "font.size": font_size,
"axes.labelsize": font_size + 4, "axes.labelsize": font_size,
"axes.titlesize": font_size + 8, "axes.titlesize": font_size,
"xtick.labelsize": font_size - 4, "xtick.labelsize": font_size,
"ytick.labelsize": font_size - 4, "ytick.labelsize": font_size,
"legend.fontsize": font_size - 6, "legend.fontsize": font_size,
"axes.linewidth": 2, "axes.linewidth": 2,
"axes.spines.top": False, "axes.spines.top": False,
"axes.spines.right": False, "axes.spines.right": False,
@ -44,7 +42,7 @@ BEST_PERFORMER_COLOR = "#D4AF37" # Gold
# Centralized Legend Configuration # Centralized Legend Configuration
LEGEND_KWARGS = { LEGEND_KWARGS = {
"loc": "upper center", "loc": "upper center",
"bbox_to_anchor": (0.5, -0.5), "bbox_to_anchor": (0.5, -0.12),
"frameon": False, "frameon": False,
} }

View file

@ -0,0 +1,150 @@
from __future__ import annotations
import argparse
from pathlib import Path
from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs
from brittle_star_project.evaluation.eval_env_builder import build_eval_env
from brittle_star_project.evaluation.video import record_episode_multi_camera
_ARCH_DIR_MAP = {
"CENTRALIZED": "centralized",
"FULLY_CONNECTED": "fully-connected",
"RING": "ring",
"SEGMENT": "segment",
}
ROBOT_COLOR_MAP = {
"CENTRALIZED": "#0D567C", # Blue
"FULLY_CONNECTED": "#8C0E0F", # Reddish
"RING": "#FCB304", # Pale Yellow
}
def _arch_dir(name: str) -> str:
return _ARCH_DIR_MAP.get(name, name.lower())
def _resolve_overrides(overrides: list[str], count: int) -> list[str | None]:
if not overrides:
return [None] * count
if len(overrides) == 1 and count > 1:
return overrides * count
if len(overrides) != count:
raise ValueError("morphology overrides must match the number of models")
return overrides
def main() -> None:
parser = argparse.ArgumentParser(description="Render top-down and follow videos for poster.")
parser.add_argument("models", nargs="+", help="Paths to .flax checkpoints")
parser.add_argument(
"--morphology-override",
action="append",
default=[],
help="Override morphology YAML path (repeat to match models)",
)
parser.add_argument("--output-root", default="vids/poster")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--max-steps", type=int, default=5000)
parser.add_argument("--topdown-camera", type=int, default=0)
parser.add_argument("--follow-camera", type=int, default=1)
parser.add_argument("--topdown-camera-x", type=float, default=-3.0)
parser.add_argument("--topdown-camera-y", type=float, default=0.0)
parser.add_argument("--topdown-camera-z", type=float, default=4.5)
parser.add_argument("--topdown-camera-fovy", type=float, default=None)
parser.add_argument("--target-x", type=float, default=-6.0)
parser.add_argument("--target-y", type=float, default=0.0)
parser.add_argument("--width", type=int, default=2160)
parser.add_argument("--height", type=int, default=960)
parser.add_argument("--fps", type=int, default=60)
parser.add_argument(
"--robot-color",
default="#2B4162",
help="Hex color for the brittle star robot",
)
args = parser.parse_args()
if (args.target_x is None) != (args.target_y is None):
raise ValueError("target-x and target-y must be provided together")
target_xy = None
if args.target_x is not None:
target_xy = (float(args.target_x), float(args.target_y))
camera_fovy = None
if args.topdown_camera_fovy is not None:
camera_fovy = {args.topdown_camera: float(args.topdown_camera_fovy)}
camera_x = None
if args.topdown_camera_x is not None:
camera_x = {args.topdown_camera: float(args.topdown_camera_x)}
camera_y = None
if args.topdown_camera_y is not None:
camera_y = {args.topdown_camera: float(args.topdown_camera_y)}
camera_z = None
if args.topdown_camera_z is not None:
camera_z = {args.topdown_camera: float(args.topdown_camera_z)}
camera_xyz = (camera_x, camera_y, camera_z)
overrides = _resolve_overrides(args.morphology_override, len(args.models))
output_root = Path(args.output_root)
for model_path_str, override in zip(args.models, overrides):
model_path = Path(model_path_str)
metadata = load_metadata(model_path, None)
training = metadata_to_configs(metadata)
bundle = build_eval_env(
model_path=model_path,
training=training,
metadata=metadata,
morphology_override_path=override,
)
arch_dir = _arch_dir(bundle.architecture)
arms_dir = f"{bundle.num_active_arms}arms"
out_dir = output_root / arms_dir / arch_dir
out_dir.mkdir(parents=True, exist_ok=True)
output_paths = {
args.topdown_camera: out_dir / "topdown.mp4",
args.follow_camera: out_dir / "follow.mp4",
}
print(bundle.architecture)
color = ROBOT_COLOR_MAP.get(bundle.architecture, args.robot_color)
result = record_episode_multi_camera(
env=bundle.env,
policy=bundle.policy,
seed=args.seed,
max_steps=args.max_steps,
action_low=bundle.action_low,
action_high=bundle.action_high,
action_mask=bundle.action_mask,
output_paths=output_paths,
camera_ids=[args.topdown_camera, args.follow_camera],
camera_fovy=camera_fovy,
camera_xyz=camera_xyz,
target_xy=target_xy,
robot_color=color,
width=args.width,
height=args.height,
fps=args.fps,
)
final_dist = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}"
print(
f"{arms_dir}/{arch_dir}: return={result.return_:.6f}, len={result.length}, "
f"target_reached={result.reached_target}, final_xy_dist={final_dist}"
)
bundle.env.close()
if __name__ == "__main__":
main()

View file

@ -0,0 +1,10 @@
#!/usr/bin/env bash
# Multi-camera renders per model -> vids/poster/{arms}arms/{arch}/topdown.mp4 + follow.mp4
path=$1
uv run scripts/poster_visualisations/render_poster_videos.py \
"$path"/centralized.flax \
"$path"/fully-connected.flax \
"$path"/ring.flax \

View file

@ -0,0 +1,241 @@
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs
from brittle_star_project.evaluation.eval_env_builder import build_eval_env
from brittle_star_project.evaluation.rollout import (
_get_observations,
_maybe_clip_action,
_target_reached,
)
from brittle_star_project.evaluation.video import (
_apply_camera_overrides,
_ensure_offscreen_size,
hex_to_rgba,
)
ROBOT_COLOR_MAP = {
"CENTRALIZED": "#0D567C", # Blue
"FULLY_CONNECTED": "#8C0E0F", # Reddish
"RING": "#FCB304", # Pale Yellow
}
def _enum_value(enum_obj, *names: str) -> int:
for name in names:
if hasattr(enum_obj, name):
return int(getattr(enum_obj, name))
raise AttributeError(f"Could not find any of {names!r} on {enum_obj!r}")
def _append_sphere(scene, mujoco, center: np.ndarray, radius: float, rgba: np.ndarray) -> None:
geom = scene.geoms[scene.ngeom]
mujoco.mjv_initGeom(
geom,
mujoco.mjtGeom.mjGEOM_SPHERE,
np.asarray([radius, 0.0, 0.0], dtype=np.float32),
center,
np.eye(3, dtype=np.float32).reshape(-1),
rgba,
)
scene.ngeom += 1
def main() -> None:
parser = argparse.ArgumentParser(description="Render a static path image from a rollout.")
parser.add_argument("model", help="Path to .flax checkpoint")
parser.add_argument("--morphology-override", default=None)
parser.add_argument("--output-path", required=True)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--max-steps", type=int, default=5000)
parser.add_argument("--body-name", default="BrittleStarMorphology/central_disk")
parser.add_argument("--camera-id", type=int, default=0)
parser.add_argument("--camera-x", type=float, default=-3.0)
parser.add_argument("--camera-y", type=float, default=0.0)
parser.add_argument("--camera-z", type=float, default=4.5)
parser.add_argument("--camera-fovy", type=float, default=None)
parser.add_argument("--target-x", type=float, default=-6.0)
parser.add_argument("--target-y", type=float, default=0.0)
parser.add_argument("--width", type=int, default=2160)
parser.add_argument("--height", type=int, default=960)
parser.add_argument("--frame-stride", type=int, default=15)
parser.add_argument(
"--path-color", default="#FA9F42"
) # ring = #888888, centralized = #2B4162, fully connected = FA9F42
parser.add_argument(
"--robot-color",
default="#FA9F42",
help="Hex color for brittle star robot (e.g. #ff0000)",
)
args = parser.parse_args()
model_path = Path(args.model)
metadata = load_metadata(model_path, None)
training = metadata_to_configs(metadata)
bundle = build_eval_env(
model_path=model_path,
training=training,
metadata=metadata,
morphology_override_path=args.morphology_override,
)
if (args.target_x is None) != (args.target_y is None):
raise ValueError("target-x and target-y must be provided together")
target_xy = None
if args.target_x is not None:
target_xy = (float(args.target_x), float(args.target_y))
try:
import imageio
import mujoco
except ImportError as e:
raise ImportError(
"Static image rendering requires 'mujoco' and 'imageio'. "
"Please install the evaluation dependencies: `uv pip install .[evaluation]`"
) from e
reset_kwargs = {}
if target_xy is not None:
reset_kwargs["target_position"] = (target_xy[0], target_xy[1], 0.0)
state = bundle.env.reset(seed=args.seed, **reset_kwargs)
model = state.mj_model
data = state.mj_data
_apply_camera_overrides(
model,
camera_fovy={args.camera_id: float(args.camera_fovy)}
if args.camera_fovy is not None
else None,
camera_xyz=(
{args.camera_id: float(args.camera_x)} if args.camera_x is not None else None,
{args.camera_id: float(args.camera_y)} if args.camera_y is not None else None,
{args.camera_id: float(args.camera_z)} if args.camera_z is not None else None,
),
)
_ensure_offscreen_size(model, args.width, args.height)
body_id = mujoco.mj_name2id(
model, mujoco.mjtObj.mjOBJ_BODY, "BrittleStarMorphology/central_disk"
)
robot_rgba = hex_to_rgba(ROBOT_COLOR_MAP.get(bundle.architecture, args.robot_color), 1.0)
# Optionally override robot color by recoloring geoms belonging to the robot's body subtree.
if args.robot_color is not None:
# Collect body IDs in the subtree rooted at `body_id` by walking parent links.
nbody = int(model.nbody)
body_parent = model.body_parentid
robot_body_ids = set([int(body_id)])
for i in range(1, nbody):
cur = int(i)
# walk up until root (0) or until we hit the robot root
while cur not in (-1, 0, int(body_id)):
cur = int(body_parent[cur])
if cur == int(body_id):
robot_body_ids.add(i)
# Recolor geoms whose body id is in the robot subtree
for g in range(int(model.ngeom)):
if int(model.geom_bodyid[g]) in robot_body_ids:
model.geom_rgba[g][:] = robot_rgba
# Optionally override robot color by recoloring geoms belonging to the robot's body subtree.
if args.robot_color is not None:
# Collect body IDs in the subtree rooted at `body_id` by walking parent links.
nbody = int(model.nbody)
body_parent = model.body_parentid
robot_body_ids = set([int(body_id)])
for i in range(1, nbody):
cur = int(i)
# walk up until root (0) or until we hit the robot root
while cur not in (-1, 0, int(body_id)):
cur = int(body_parent[cur])
if cur == int(body_id):
robot_body_ids.add(i)
# Recolor geoms whose body id is in the robot subtree
for g in range(int(model.ngeom)):
if int(model.geom_bodyid[g]) in robot_body_ids:
model.geom_rgba[g][:] = robot_rgba
positions = []
observations = _get_observations(state)
for _ in range(int(args.max_steps)):
positions.append(np.asarray(data.xpos[body_id], dtype=np.float32))
obs_dict = observations or {}
action = bundle.policy.act(observations=obs_dict)
if bundle.action_mask is not None:
action = action[bundle.action_mask]
action = _maybe_clip_action(action, bundle.action_low, bundle.action_high)
state = bundle.env.step(state=state, action=action)
data = state.mj_data
observations = _get_observations(state)
if _target_reached(state=state):
break
positions_arr = np.vstack(positions)
if len(positions_arr) < 2:
raise ValueError("Need at least two rollout positions to render a path")
path_points = positions_arr.copy()
path_points[:, 2] -= 0.02
path_step = max(1, int(args.frame_stride))
path_points_visible = path_points[::path_step]
path_rgba = hex_to_rgba(ROBOT_COLOR_MAP.get(bundle.architecture, args.path_color), 0.92)
ctx = mujoco.GLContext(args.width, args.height)
ctx.make_current()
try:
catmask = _enum_value(mujoco.mjtCatBit, "mjCAT_ALL")
camera_type = _enum_value(mujoco.mjtCamera, "mjCAMERA_FIXED")
font_scale = _enum_value(mujoco.mjtFontScale, "mjFONTSCALE_100")
maxgeom = int(model.ngeom + len(path_points_visible) + 8)
scene = mujoco.MjvScene(model, maxgeom=maxgeom)
option = mujoco.MjvOption()
perturb = mujoco.MjvPerturb()
camera = mujoco.MjvCamera()
mujoco.mjv_defaultOption(option)
mujoco.mjv_defaultPerturb(perturb)
mujoco.mjv_defaultCamera(camera)
camera.type = camera_type
camera.fixedcamid = int(args.camera_id)
if hasattr(camera, "trackbodyid"):
camera.trackbodyid = -1
context = mujoco.MjrContext(model, font_scale)
viewport = mujoco.MjrRect(0, 0, args.width, args.height)
mujoco.mjv_updateScene(model, data, option, perturb, camera, catmask, scene)
for idx, path_point in enumerate(path_points_visible):
path_rgba[3] = 0.10 + 0.70 * (idx / max(len(path_points_visible) - 1, 1))
_append_sphere(scene, mujoco, path_point, 0.03, path_rgba)
rgb = np.empty((args.height, args.width, 3), dtype=np.uint8)
depth = np.empty((args.height, args.width), dtype=np.float32)
mujoco.mjr_render(viewport, scene, context)
mujoco.mjr_readPixels(rgb, depth, viewport, context)
imageio.imwrite(args.output_path, np.flipud(rgb))
context.free()
finally:
ctx.free()
bundle.env.close()
if __name__ == "__main__":
main()

View file

@ -0,0 +1,10 @@
#!/usr/bin/env bash
# Static path image + optional ghost render
path=$1 # path to .flax model with metadata.yaml alongside it
uv run scripts/poster_visualisations/render_static_path_image.py \
"$path" \
--output-path vids/poster/5arms/centralized/path.png \
--ghost-overlay

View file

@ -107,6 +107,9 @@ def main(dict_cfg: DictConfig) -> None:
action_mask=action_mask, action_mask=action_mask,
output_path=output_path, output_path=output_path,
camera_id=sim_cfg.camera_id, camera_id=sim_cfg.camera_id,
width=sim_cfg.video_width,
height=sim_cfg.video_height,
fps=sim_cfg.video_fps,
) )
save_evaluation_metadata( save_evaluation_metadata(

View file

@ -0,0 +1,142 @@
from pathlib import Path
from concurrent.futures import ThreadPoolExecutor, as_completed
import threading
import wandb
import argparse
# tune these depending on network / W&B limits
MAX_RUN_WORKERS = 8
MAX_FILE_WORKERS = 16
MAX_ARTIFACT_WORKERS = 8
api = wandb.Api()
print_lock = threading.Lock()
def safe_print(*args, **kwargs):
with print_lock:
print(*args, **kwargs)
def download_file(file, run_dir):
target = run_dir / file.name
try:
# skip existing files
if target.exists():
return f"SKIP FILE {target}"
target.parent.mkdir(parents=True, exist_ok=True)
file.download(root=run_dir, replace=False)
return f"DONE FILE {target}"
except Exception as e:
return f"FAIL FILE {target}: {e}"
def sanitize_artifact_name(name: str):
return name.replace(":", "_")
def download_artifact(artifact, artifact_root):
try:
artifact_name = sanitize_artifact_name(artifact.name)
artifact_dir = artifact_root / artifact_name
if artifact_dir.exists() and any(artifact_dir.iterdir()):
return f"SKIP ARTIFACT {artifact.name}"
artifact_dir.mkdir(parents=True, exist_ok=True)
artifact.download(root=artifact_dir)
return f"DONE ARTIFACT {artifact.name}"
except Exception as e:
return f"FAIL ARTIFACT {artifact.name}: {e}"
def download_run(run, root):
run_dir = root / f"{run.name}"
run_dir.mkdir(parents=True, exist_ok=True)
safe_print(f"\n=== {run.name} ({run.id}) ===")
# -------------------------
# Download regular run files
# -------------------------
files = list(run.files())
with ThreadPoolExecutor(max_workers=MAX_FILE_WORKERS) as executor:
futures = [executor.submit(download_file, file, run_dir) for file in files]
for future in as_completed(futures):
safe_print(future.result())
# -------------------------
# Download logged artifacts
# -------------------------
artifact_root = run_dir / "artifacts"
try:
artifacts = list(run.logged_artifacts())
safe_print(f"Found {len(artifacts)} artifacts for {run.name}")
with ThreadPoolExecutor(max_workers=MAX_ARTIFACT_WORKERS) as executor:
futures = [
executor.submit(download_artifact, artifact, artifact_root)
for artifact in artifacts
]
for future in as_completed(futures):
safe_print(future.result())
except Exception as e:
safe_print(f"Artifact download failed for {run.name}: {e}")
# -------------------------
# OPTIONAL: download used/input artifacts
# -------------------------
# try:
# used_artifacts = list(run.used_artifacts())
# used_root = run_dir / "used_artifacts"
#
# for artifact in used_artifacts:
# download_artifact(artifact, used_root)
# except Exception as e:
# safe_print(f"Used artifact download failed: {e}")
safe_print(f"Finished {run.name}")
def main(entity: str, project: str, root: Path):
root.mkdir(exist_ok=True)
runs = list(api.runs(f"{entity}/{project}"))
safe_print(f"Found {len(runs)} runs")
with ThreadPoolExecutor(max_workers=MAX_RUN_WORKERS) as executor:
futures = [executor.submit(download_run, run, root) for run in runs]
for future in as_completed(futures):
try:
future.result()
except Exception as e:
safe_print("RUN FAILED:", e)
safe_print("\nAll downloads complete.")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--entity", type=str, default="SEL3-2026-Groep-4")
parser.add_argument("--project", type=str, required=True)
parser.add_argument("--root", type=str, default="runs")
args = parser.parse_args()
root = Path(args.root)
main(entity=args.entity, project=args.project, root=root)

View file

@ -26,6 +26,9 @@ class SimulationSettings:
video_output_path: Optional[str] = None video_output_path: Optional[str] = None
# Camera ID to use for video recording (1 is usually the close-up camera) # Camera ID to use for video recording (1 is usually the close-up camera)
camera_id: int = 1 camera_id: int = 1
video_width: int = 640
video_height: int = 480
video_fps: int = 60
# Optional override for the sidecar metadata YAML file. # Optional override for the sidecar metadata YAML file.
# If None, it defaults to the model_path with a `_metadata.yaml` suffix. # If None, it defaults to the model_path with a `_metadata.yaml` suffix.

View file

@ -65,11 +65,21 @@ class BrittleStarJaxEnvWrapper:
def single_observation_space(self): def single_observation_space(self):
return self._env.observation_space return self._env.observation_space
def reset(self, seed: int = 0): def reset(self, seed: int = 0, target_position: tuple[float, float] | None = None):
self.logger.info(f"Resetting vectorized environment environments with seed {seed}") self.logger.info(f"Resetting vectorized environment environments with seed {seed}")
self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2) self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2)
env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs)) env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs))
state = self._vectorized_reset(rng=env_rngs)
# If a target_position is provided, pass it through to the underlying env.reset
if target_position is None:
state = jax.jit(jax.vmap(lambda rng: self._env.reset(rng=rng)))(env_rngs)
else:
tp = jnp.asarray(target_position, dtype=jnp.float32)
tp_batched = jnp.tile(tp[None, :], (self._num_envs, 1))
state = jax.jit(jax.vmap(lambda rng, t: self._env.reset(rng=rng, target_position=t)))(
env_rngs, tp_batched
)
return state return state
def sample_actions(self): def sample_actions(self):

View file

@ -62,9 +62,12 @@ class BrittleStarEnv:
return jax.random.PRNGKey(seed) return jax.random.PRNGKey(seed)
def reset(self, *, seed: int = 0): def reset(self, *, seed: int = 0, target_position: tuple[float, float, float] | None = None):
rng = self.make_rng(seed) rng = self.make_rng(seed)
state = self._env.reset(rng=rng) if target_position is not None:
state = self._env.reset(rng=rng, target_position=target_position)
else:
state = self._env.reset(rng=rng)
return state return state
def render(self, *, state: Any): def render(self, *, state: Any):

View file

@ -25,6 +25,59 @@ def create_evaluation_dir(model_path: Path) -> Path:
return eval_dir return eval_dir
def _ensure_offscreen_size(model, width: int, height: int) -> None:
vis_global = getattr(getattr(model, "vis", None), "global_", None)
if vis_global is None:
return
vis_global.offwidth = int(max(width, vis_global.offwidth))
vis_global.offheight = int(max(height, vis_global.offheight))
def _apply_camera_overrides(
model,
*,
camera_fovy: dict[int, float] | None = None,
camera_xyz: tuple[
dict[int, float] | None,
dict[int, float] | None,
dict[int, float] | None,
] = (None, None, None),
) -> None:
if not camera_fovy and not (camera_xyz[0] or camera_xyz[1] or camera_xyz[2]):
return
ncam = int(getattr(model, "ncam", 0))
for cam_id, fovy in (camera_fovy or {}).items():
if cam_id < 0 or cam_id >= ncam:
raise ValueError(f"Camera id {cam_id} is out of range")
model.cam_fovy[cam_id] = float(fovy)
for cam_id, x in (camera_xyz[0] or {}).items():
if cam_id < 0 or cam_id >= ncam:
raise ValueError(f"Camera id {cam_id} is out of range")
model.cam_pos[cam_id][0] = float(x)
for cam_id, y in (camera_xyz[1] or {}).items():
if cam_id < 0 or cam_id >= ncam:
raise ValueError(f"Camera id {cam_id} is out of range")
model.cam_pos[cam_id][1] = float(y)
for cam_id, z in (camera_xyz[2] or {}).items():
if cam_id < 0 or cam_id >= ncam:
raise ValueError(f"Camera id {cam_id} is out of range")
model.cam_pos[cam_id][2] = float(z)
def hex_to_rgba(hex_color: str, alpha: float) -> np.ndarray:
color = hex_color.lstrip("#")
if len(color) != 6:
raise ValueError(f"Expected a 6-digit hex color, got {hex_color!r}")
red = int(color[0:2], 16) / 255.0
green = int(color[2:4], 16) / 255.0
blue = int(color[4:6], 16) / 255.0
return np.asarray([red, green, blue, float(alpha)], dtype=np.float32)
def save_evaluation_metadata( def save_evaluation_metadata(
eval_dir: Path, eval_dir: Path,
*, *,
@ -66,6 +119,7 @@ def record_episode(
fps: int = 60, fps: int = 60,
width: int = 640, width: int = 640,
height: int = 480, height: int = 480,
target_xy: tuple[float, float] | None = None,
) -> EpisodeResult: ) -> EpisodeResult:
"""Run an episode headlessly and record a video using MuJoCo's Renderer and imageio. """Run an episode headlessly and record a video using MuJoCo's Renderer and imageio.
@ -92,10 +146,12 @@ def record_episode(
"Please install the evaluation dependencies: `uv pip install .[evaluation]`" "Please install the evaluation dependencies: `uv pip install .[evaluation]`"
) from e ) from e
state = env.reset(seed=seed) state = env.reset(seed=seed, target_position=target_xy)
model = state.mj_model model = state.mj_model
data = state.mj_data data = state.mj_data
_ensure_offscreen_size(model, width, height)
renderer = mujoco.Renderer(model, width=width, height=height) renderer = mujoco.Renderer(model, width=width, height=height)
ep_return = 0.0 ep_return = 0.0
observations = _get_observations(state) observations = _get_observations(state)
@ -147,3 +203,133 @@ def record_episode(
final_xy_dist=final_dist, final_xy_dist=final_dist,
initial_target_distance=initial_dist, initial_target_distance=initial_dist,
) )
def record_episode_multi_camera(
*,
env: BrittleStarEnv,
policy: ControlPolicy,
seed: int,
max_steps: int,
action_low: np.ndarray | None,
action_high: np.ndarray | None,
output_paths: dict[int, Path],
action_mask: np.ndarray | None = None,
camera_ids: list[int] | None = None,
camera_fovy: dict[int, float] | None = None,
camera_xyz: tuple[
dict[int, float] | None,
dict[int, float] | None,
dict[int, float] | None,
] = (None, None, None),
target_xy: tuple[float, float] | None = None,
robot_color: str | None = None,
fps: int = 60,
width: int = 640,
height: int = 480,
) -> EpisodeResult:
"""Run one episode and render multiple camera views to separate files."""
try:
import imageio
import mujoco
except ImportError as e:
raise ImportError(
"Video recording requires 'imageio' and 'mujoco'. "
"Please install the evaluation dependencies: `uv pip install .[evaluation]`"
) from e
if camera_ids is None:
camera_ids = list(output_paths.keys())
for cam_id in camera_ids:
if cam_id not in output_paths:
raise ValueError(f"Missing output path for camera {cam_id}")
output_paths = {cam_id: output_paths[cam_id] for cam_id in camera_ids}
for path in output_paths.values():
path.parent.mkdir(parents=True, exist_ok=True)
state = env.reset(seed=seed, target_position=(target_xy[0], target_xy[1], 0.0))
model = state.mj_model
data = state.mj_data
_apply_camera_overrides(model, camera_fovy=camera_fovy, camera_xyz=camera_xyz)
if robot_color is not None:
robot_body_id = mujoco.mj_name2id(
model, mujoco.mjtObj.mjOBJ_BODY, "BrittleStarMorphology/central_disk"
)
if robot_body_id < 0:
raise ValueError("Body 'BrittleStarMorphology/central_disk' not found in the model")
robot_rgba = hex_to_rgba(robot_color, 1.0)
body_parent = model.body_parentid
robot_body_ids = {int(robot_body_id)}
for body_id in range(1, int(model.nbody)):
current_body_id = int(body_id)
while current_body_id not in (-1, 0, int(robot_body_id)):
current_body_id = int(body_parent[current_body_id])
if current_body_id == int(robot_body_id):
robot_body_ids.add(body_id)
for geom_id in range(int(model.ngeom)):
if int(model.geom_bodyid[geom_id]) in robot_body_ids:
model.geom_rgba[geom_id][:] = robot_rgba
_ensure_offscreen_size(model, width, height)
renderer = mujoco.Renderer(model, width=width, height=height)
writers = {
cam_id: imageio.get_writer(str(path), fps=fps) for cam_id, path in output_paths.items()
}
ep_return = 0.0
observations = _get_observations(state)
prev_dist = _get_xy_distance_to_target(observations) if observations else None
initial_dist = prev_dist
reached_target = _target_reached(state=state)
steps = 0
try:
for _ in range(int(max_steps)):
for cam_id in camera_ids:
renderer.update_scene(data, camera=cam_id)
writers[cam_id].append_data(renderer.render())
obs_dict = observations or {}
action = policy.act(observations=obs_dict)
if action_mask is not None:
action = action[action_mask]
action = _maybe_clip_action(action, action_low, action_high)
state = env.step(state=state, action=action)
steps += 1
observations = _get_observations(state)
cur_dist = _get_xy_distance_to_target(observations) if observations else None
if prev_dist is not None and cur_dist is not None:
ep_return += prev_dist - cur_dist
prev_dist = cur_dist
reached_target = _target_reached(state=state)
if reached_target:
break
for cam_id in camera_ids:
renderer.update_scene(data, camera=cam_id)
writers[cam_id].append_data(renderer.render())
finally:
renderer.close()
for writer in writers.values():
writer.close()
final_dist = _get_xy_distance_to_target(observations) if observations else None
return EpisodeResult(
return_=ep_return,
length=steps,
reached_target=reached_target,
final_xy_dist=final_dist,
initial_target_distance=initial_dist,
)