Merge pull request #3079 from aftersomemath:sysid-pr

PiperOrigin-RevId: 868229512
Change-Id: I790bc08fc8b0745583a2f92d9ee2c5a19ba558ea
This commit is contained in:
Copybara-Service
2026-02-10 11:04:12 -08:00
49 changed files with 11554 additions and 0 deletions
@@ -0,0 +1,15 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
@@ -0,0 +1,63 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Abstract base class for report sections."""
import abc
from collections.abc import Iterable
from typing import Any
class ReportSection(abc.ABC):
"""Abstract base class for all report sections."""
@property
@abc.abstractmethod
def template_filename(self) -> str:
"""The filename of the Jinja2 template (e.g., 'parameters.html')."""
pass
@abc.abstractmethod
def get_context(self) -> dict[str, Any]:
"""Returns data needed by the template."""
pass
def __init__(self, collapsible: bool = True, is_open: bool = True):
self._collapsible = collapsible
self._is_open = is_open
self._anchor = ""
@property
def title(self) -> str:
return ""
@property
def anchor(self) -> str:
"""Returns a unique HTML anchor string."""
# Auto-generate a safe anchor from title if not provided
if not hasattr(self, "_anchor") or not self._anchor:
return self.title.lower().replace(" ", "-")
return self._anchor
@property
def collapsible(self) -> bool:
return self._collapsible
@property
def is_open(self) -> bool:
return self._is_open
def header_includes(self) -> Iterable[str]:
"""Returns strings (like <script> tags) to add to <head>."""
return set()
@@ -0,0 +1,177 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Covariance and correlation matrix report sections."""
from collections.abc import Callable
from typing import Any
import matplotlib
from matplotlib import cm
from matplotlib import colors as mpl_colors
from mujoco.sysid._src import parameter
from mujoco.sysid.report.sections.base import ReportSection
from mujoco.sysid.report.utils import get_text_color
import numpy as np
from numpy import typing as npt
def _compute_correlation(cov: npt.ArrayLike) -> np.ndarray:
"""Calculates the correlation matrix from a covariance matrix.
Formula: A[i,j] = cov[i,j] / sqrt(cov[i,i] * cov[j,j])
Args:
cov: The covariance matrix.
Returns:
The correlation matrix.
"""
cov = np.array(cov)
sqrt_diag_cov = np.sqrt(np.diag(cov))
# denom[i, j] = sqrt_diag_cov[i] * sqrt_diag_cov[j]
denom = np.outer(sqrt_diag_cov, sqrt_diag_cov)
# Safe division (handles division by zero by returning 0)
return np.divide(
cov, denom, out=np.zeros_like(cov, dtype=float), where=denom != 0
)
class Covariance(ReportSection):
"""Displays values of the covariance and correlation matrices.
Assumed to be symmetric.
"""
def __init__(
self,
title: str,
covariance: npt.ArrayLike,
parameter_dict: parameter.ParameterDict,
anchor: str = "",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._anchor = anchor
self._covariance = np.array(covariance)
self._parameter_dict = parameter_dict
@property
def template_filename(self) -> str:
return "covariance.html"
@property
def title(self):
return self._title
@property
def anchor(self) -> str:
return self._anchor
def get_context(self) -> dict[str, Any]:
dim_names = []
for param_name, param in self._parameter_dict.parameters.items():
if param.frozen:
continue
if param.size == 1:
dim_names.append(param_name)
else:
for i in range(param.size):
dim_names.append(f"{param_name}[{i}]")
context = {
"title": self._title,
"covariance_data": self._covariance_table_data(),
"correlation_data": self._correlation_table_data(),
"dim_names": dim_names,
}
if self._covariance.size == 0:
context["message"] = (
"Covariance matrix is empty. This usually means there are no"
" parameters to optimize or all parameters are frozen."
)
return context
def header_includes(self) -> set[str]:
return {""}
def _create_table_data(
self,
matrix: np.ndarray,
norm: mpl_colors.Normalize,
cmap: mpl_colors.Colormap,
# Function to get the value used for coloring based on (row, col, value)
get_color_input_value: Callable[[int, int, float], float],
) -> list[list[dict[str, Any]]]:
"""Helper method to generate formatted table data with colors."""
scalar_map = cm.ScalarMappable(norm=norm, cmap=cmap)
table_data = []
for i in range(matrix.shape[0]):
row_data = []
for j in range(matrix.shape[1]):
value = matrix[i, j]
color_value = get_color_input_value(i, j, value)
# Get RGBA color, convert to HEX.
rgba_color = scalar_map.to_rgba(color_value) # pyright: ignore[reportArgumentType] # type: ignore[arg-type]
hex_color = mpl_colors.rgb2hex(rgba_color) # pyright: ignore[reportArgumentType] # type: ignore[arg-type]
# Make sure the text is visible over the cell color.
text_color = get_text_color(hex_color)
row_data.append({
"value": value,
"bgcolor": hex_color,
"textcolor": text_color,
})
table_data.append(row_data)
return table_data
def _covariance_table_data(self):
# Use the positive side of bwr that goes from white to red
cmap = matplotlib.colormaps["bwr"]
if self._covariance.size == 0:
max_abs_val = 1.0 # Default value to avoid errors
else:
max_abs_val = np.max(np.abs(self._covariance))
if max_abs_val == 0:
max_abs_val = 1
norm = mpl_colors.Normalize(vmin=-max_abs_val, vmax=max_abs_val)
# Color based on the actual value
def get_color_input(_i, _j, val):
return np.abs(val)
return self._create_table_data(
self._covariance, norm, cmap, get_color_input
)
def _correlation_table_data(self):
correlation = _compute_correlation(self._covariance)
# Use the positive side of bwr that goes from white to red
cmap = matplotlib.colormaps["bwr"]
norm = mpl_colors.Normalize(vmin=-1, vmax=1)
# Color based on the abs value, but use 0 for the diagonal (white)
def get_color_input(i, j, val):
if i == j:
return 0.0
return np.abs(val)
return self._create_table_data(correlation, norm, cmap, get_color_input)
@@ -0,0 +1,59 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Group report section that contains multiple child sections."""
from typing import Any
from mujoco.sysid.report.sections.base import ReportSection
class GroupSection(ReportSection):
"""A report section that groups multiple other sections together."""
def __init__(
self,
title: str,
sections: list[ReportSection],
anchor: str = "",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self.sections = sections
self._anchor = anchor
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
@property
def template_filename(self) -> str:
return "group.html"
def header_includes(self) -> set[str]:
includes = set()
for section in self.sections:
includes.update(section.header_includes())
return includes
def get_context(self) -> dict[str, Any]:
return {
"title": self.title,
"anchor": self.anchor,
}
@@ -0,0 +1,113 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Automated insights report section."""
from typing import Any
from mujoco.sysid._src import parameter
from mujoco.sysid.report.sections.base import ReportSection
class AutomatedInsights(ReportSection):
"""Analyzes identification results and generates automated insights/suggestions.
Currently checks for parameters stuck at boundaries.
"""
def __init__(
self,
title: str,
parameter_dict: parameter.ParameterDict,
anchor: str = "insights",
collapsible: bool = True,
):
# Default to collapsed
super().__init__(collapsible=collapsible, is_open=False)
self._anchor = anchor
self._parameter_dict = parameter_dict
self._x_hat = parameter_dict.as_vector()
self._insights = self._generate_insights()
self._n_warnings = len(
[i for i in self._insights if i["type"] == "warning"]
)
self._title = f"Log - {self._n_warnings} warnings"
def _generate_insights(self) -> list[dict[str, Any]]:
insights = []
# Logic: Check for boundary hits
param_names = self._parameter_dict.get_non_frozen_parameter_names()
bounds = self._parameter_dict.get_bounds()
lower_bounds, upper_bounds = bounds
if len(self._x_hat) != len(param_names):
raise ValueError("Parameter count mismatch")
for i, name in enumerate(param_names):
val = self._x_hat[i]
lb = lower_bounds[i]
ub = upper_bounds[i]
rng = max(ub - lb, 1e-9)
# Threshold: 0.1% of range or 1e-6 absolute
threshold = rng * 1e-3
if threshold < 1e-8:
threshold = 1e-8
if abs(val - lb) < threshold:
insights.append({
"type": "warning",
"title": "Lower Bound Hit",
"message": (
f"Parameter <b>{name}</b> ({val:.4g}) is at its lower bound"
f" ({lb:.4g})."
),
})
elif abs(val - ub) < threshold:
insights.append({
"type": "warning",
"title": "Upper Bound Hit",
"message": (
f"Parameter <b>{name}</b> ({val:.4g}) is at its upper bound"
f" ({ub:.4g})."
),
})
if not insights:
insights.append({
"type": "success",
"title": "No Issues Detected",
"message": "All parameters are within their bounds.",
})
return insights
@property
def template_filename(self) -> str:
return "insights.html"
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
def get_context(self) -> dict[str, Any]:
return {"title": self._title, "insights": self._insights}
@@ -0,0 +1,471 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Optimization trace report section."""
from collections.abc import Sequence
import math
from typing import Any
from mujoco.sysid.report.sections.base import ReportSection
from mujoco.sysid.report.utils import plotly_script_tag
import numpy as np
from numpy import typing as npt
from plotly import subplots as plt_subplots
import plotly.graph_objects as go
class OptimizationTrace(ReportSection):
"""Displays plots summarizing the optimization process."""
def __init__(
self,
title: str,
objective: Sequence[float],
candidate: Sequence[npt.ArrayLike],
bounds: (
tuple[Sequence[float] | np.ndarray, Sequence[float] | np.ndarray]
| None
) = None,
param_names: Sequence[str] | None = None,
log_diff: bool = True,
bound_eps: float = 1e-3,
dims_per_page: int = 6, # For candidate plot paging
anchor: str = "",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._objective = objective
self._candidate = candidate # List of vectors
self._bounds = bounds
self._param_names = param_names
self._log_diff = log_diff
self._bound_eps = bound_eps
self._dims_per_page = dims_per_page
self._anchor = anchor
@property
def template_filename(self) -> str:
return "optimization_trace.html"
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
def header_sections(self) -> set[str]:
return {plotly_script_tag()}
def get_context(self) -> dict[str, Any]:
objective_fig = self._get_objective_figure()
candidate_figs = self._get_candidate_figures()
candidate_heatmap_fig = self._get_candidate_heatmap_figure()
config = {
"displayModeBar": True,
"displaylogo": False,
"toImageButtonOptions": {
"format": "svg",
"filename": f"{self._title}_optimization",
"height": 800,
"width": 1200,
"scale": 1,
},
"responsive": True,
}
return {
"title": self._title,
"objective_plot_html": (
objective_fig.to_html(
full_html=False, include_plotlyjs=False, config=config
)
if objective_fig
else None
),
"candidate_plots_html": (
[
fig.to_html(
full_html=False, include_plotlyjs=False, config=config
)
for fig in candidate_figs
]
if candidate_figs
else None
),
"candidate_heatmap_html": (
candidate_heatmap_fig.to_html(
full_html=False, include_plotlyjs=False, config=config
)
if candidate_heatmap_fig
else None
),
}
def _get_objective_figure(self) -> go.Figure | None:
"""Generates the objective function plot."""
if not self._objective:
return None
fig = go.Figure()
fig.add_trace(
go.Scatter(
y=self._objective,
mode="lines+markers",
marker=dict(size=4),
line=dict(width=2),
name="Objective",
)
)
final_value = self._objective[-1]
if abs(final_value) < 1e-3 or abs(final_value) > 1e3:
final_str = f"{final_value:.2e}"
else:
final_str = f"{final_value:.4f}"
fig.update_layout(
title=f"Objective Over Time (Final: {final_str})",
xaxis_title="Iteration",
yaxis_title="Objective",
hovermode="x unified",
height=400,
autosize=True,
margin=dict(l=50, r=50, t=80, b=50),
template="plotly_white",
)
fig.update_xaxes(
showgrid=True, gridwidth=1, gridcolor="rgba(211, 211, 211, 0.7)"
)
fig.update_yaxes(
showgrid=True, gridwidth=1, gridcolor="rgba(211, 211, 211, 0.7)"
)
return fig
def _get_candidate_figures(self) -> list[go.Figure]:
"""Generates the parameter candidate plots (paged)."""
if not self._candidate:
return []
values = np.array(self._candidate).T # shape: (n_dim, n_iter)
n_dim, n_iter = values.shape
if n_iter <= 1:
return []
diffs = np.diff(values, axis=1) # shape: (n_dim, n_iter-1)
if self._bounds is not None:
mins = np.array(self._bounds[0])
maxs = np.array(self._bounds[1])
if not (mins.shape == (n_dim,) and maxs.shape == (n_dim,)):
raise ValueError("Bounds dimensions do not match parameter dimensions.")
else:
mins = np.full(n_dim, -np.inf)
maxs = np.full(n_dim, np.inf)
param_names = (
self._param_names
if self._param_names is not None
else [f"Dim {i}" for i in range(n_dim)]
)
if len(param_names) != n_dim:
raise ValueError(
"Number of parameter names does not match parameter dimensions."
)
n_pages = math.ceil(n_dim / self._dims_per_page)
figures = []
iterations = np.arange(n_iter)
iterations_diff = np.arange(1, n_iter)
for page in range(n_pages):
start_dim = page * self._dims_per_page
end_dim = min((page + 1) * self._dims_per_page, n_dim)
dims_in_page = end_dim - start_dim
page_param_names = param_names[start_dim:end_dim]
fig = plt_subplots.make_subplots(
rows=dims_in_page,
cols=2,
shared_xaxes=True,
subplot_titles=[
title
for name in page_param_names
for title in (f"{name} Value", f"{name} Δ")
],
vertical_spacing=max(0.02, 0.1 / dims_in_page),
)
for i, dim in enumerate(range(start_dim, end_dim)):
row_idx = i + 1
vals = values[dim, :]
diff_vals = diffs[dim, :]
lower, upper = mins[dim], maxs[dim]
# --- Value Plot (Col 1) ---
# Check for near-bound points
near_lower = np.abs(vals - lower) < self._bound_eps
near_upper = np.abs(vals - upper) < self._bound_eps
near_bound = near_lower | near_upper
# Plot segments with different colors if near bounds
for t in range(1, n_iter):
is_near_prev = near_bound[t - 1]
is_near_curr = near_bound[t]
color = (
"#d62728" if (is_near_prev or is_near_curr) else "#1f77b4"
) # Red if current or prev near bound
fig.add_trace(
go.Scatter(
x=iterations[t - 1 : t + 1],
y=vals[t - 1 : t + 1],
mode="lines",
line=dict(color=color, width=2),
showlegend=False,
),
row=row_idx,
col=1,
)
fig.add_trace(
go.Scatter(
x=iterations,
y=vals,
mode="markers",
marker=dict(
size=5,
color=["#d62728" if nb else "#1f77b4" for nb in near_bound],
symbol=[
"triangle-down"
if nl
else "triangle-up"
if nu
else "circle"
for nl, nu in zip(near_lower, near_upper, strict=True)
], # Triangles for bounds.
),
name=f"{param_names[dim]}",
showlegend=False,
hoverinfo="x+y+name",
),
row=row_idx,
col=1,
)
# Add bound lines if finite.
if np.isfinite(lower):
fig.add_hline(
y=lower,
line_dash="dash",
line_color="gray",
row=row_idx, # pyright: ignore[reportArgumentType]
col=1, # pyright: ignore[reportArgumentType]
opacity=0.5,
)
if np.isfinite(upper):
fig.add_hline(
y=upper,
line_dash="dash",
line_color="gray",
row=row_idx, # pyright: ignore[reportArgumentType]
col=1, # pyright: ignore[reportArgumentType]
opacity=0.5,
)
# Add final value annotation.
final_val = vals[-1]
final_str = (
f"{final_val:.2e}"
if abs(final_val) < 1e-3 or abs(final_val) > 1e3
else f"{final_val:.4f}"
)
fig.add_annotation(
x=iterations[-1],
y=final_val,
text=final_str,
showarrow=True,
arrowhead=1,
ax=20,
ay=-30,
row=row_idx,
col=1,
font=dict(color="blue", size=9),
)
# --- Diff Plot (Col 2) ---
if self._log_diff:
eps = 1e-12
plot_diff_vals = np.log10(np.abs(diff_vals) + eps)
yaxis_title = "log |Δ|"
else:
plot_diff_vals = diff_vals
yaxis_title = "Δ"
fig.add_trace(
go.Scatter(
x=iterations_diff,
y=plot_diff_vals,
mode="lines+markers",
marker=dict(symbol="x", size=5, color="orange"),
line=dict(width=1.5, color="orange"),
name=f"Δ {param_names[dim]}",
showlegend=False,
hoverinfo="x+y+name",
),
row=row_idx,
col=2,
)
fig.update_yaxes(
title_text=yaxis_title, row=row_idx, col=2, title_font_size=10
)
# --- Layout Updates for the Page Figure ---
fig.update_layout(
title=(
f"Candidate Values and Changes (Page {page + 1}/{n_pages}, Dims"
f" {start_dim}-{end_dim - 1})"
),
height=max(400, 200 * dims_in_page), # Adjust height based on dims
autosize=True,
margin=dict(l=60, r=30, t=100, b=50),
hovermode="x unified",
template="plotly_white",
)
fig.update_xaxes(
showgrid=True,
gridwidth=1,
gridcolor="rgba(211, 211, 211, 0.7)",
zeroline=False,
)
fig.update_yaxes(
showgrid=True,
gridwidth=1,
gridcolor="rgba(211, 211, 211, 0.7)",
zeroline=False,
)
# Add common x-axis label to the bottom row.
fig.update_xaxes(title_text="Iteration", row=dims_in_page, col=1)
fig.update_xaxes(title_text="Iteration", row=dims_in_page, col=2)
for annotation in fig.layout.annotations:
annotation.font.size = 10
figures.append(fig)
return figures
def _get_candidate_heatmap_figure(self) -> go.Figure | None:
"""Generates the parameter candidate heatmap."""
if not self._candidate:
return None
data = np.array(self._candidate).T # shape: (n_dim, n_iter)
n_dim, n_iter = data.shape
param_names = (
self._param_names
if self._param_names is not None
else [f"Dim {i}" for i in range(n_dim)]
)
if len(param_names) != n_dim:
raise ValueError(
"Number of parameter names does not match parameter dimensions."
)
# Normalize data for heatmap colors if bounds are provided.
heatmap_data = data.copy()
normalize = self._bounds is not None
if normalize and self._bounds is not None:
min_bounds, max_bounds = self._bounds
if not (len(min_bounds) == len(max_bounds) == n_dim):
raise ValueError("Bounds dimensions do not match parameter dimensions.")
for i in range(n_dim):
min_val, max_val = min_bounds[i], max_bounds[i]
denom = max_val - min_val if max_val > min_val else 1.0
clipped_vals = np.clip(data[i], min_val, max_val)
heatmap_data[i] = (
(clipped_vals - min_val) / denom if denom != 0 else 0.5
) # Center if range is zero
else:
row_mins = np.min(data, axis=1, keepdims=True)
row_maxs = np.max(data, axis=1, keepdims=True)
row_ranges = row_maxs - row_mins
row_ranges[row_ranges == 0] = 1.0 # Avoid division by zero
heatmap_data = (data - row_mins) / row_ranges
fig = go.Figure(
data=go.Heatmap(
z=heatmap_data[::-1],
x=np.arange(n_iter),
y=list(reversed(param_names)),
colorscale="RdBu",
colorbar=dict(
title="Normalized Value"
if normalize
else "Row-Normalized Value"
),
hovertemplate=(
"Iter: %{x}<br>Param: %{y}<br>Value:"
" %{customdata:.4f}<extra></extra>"
),
customdata=data[::-1].tolist(),
)
)
# Add markers for bound hits if bounds exist.
bound_markers_x = []
bound_markers_y = []
if self._bounds is not None:
min_bounds, max_bounds = self._bounds
for dim in range(n_dim):
min_val, max_val = min_bounds[dim], max_bounds[dim]
for iter_idx, val in enumerate(data[dim]):
if (
abs(val - min_val) < self._bound_eps
or abs(val - max_val) < self._bound_eps
):
bound_markers_x.append(iter_idx)
bound_markers_y.append(param_names[dim])
if bound_markers_x:
fig.add_trace(
go.Scatter(
x=bound_markers_x,
y=bound_markers_y,
mode="markers",
marker=dict(color="black", size=6, symbol="x"),
name="At Bound",
showlegend=False,
hoverinfo="skip",
)
)
fig.update_layout(
title="Candidate Parameter Heatmap",
xaxis_title="Iteration",
yaxis_title="Parameter",
height=max(400, 30 * n_dim),
autosize=True,
margin=dict(l=150, r=50, t=80, b=50),
yaxis=dict(
tickmode="array", tickvals=param_names, ticktext=param_names
),
template="plotly_white",
)
return fig
@@ -0,0 +1,428 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Parameter distribution report section."""
import math
from typing import Any
from mujoco.sysid._src import parameter
from mujoco.sysid.report.sections.base import ReportSection
from mujoco.sysid.report.utils import plotly_script_tag
import numpy as np
from plotly import subplots as plt_subplots
import plotly.graph_objects as go
class ParameterDistribution(ReportSection):
"""Displays the identified parameters relative to their bounds and nominal values.
Visualization: One horizontal track per parameter with markers for Nominal,
Identified, and Bounds.
"""
def __init__(
self,
title: str,
opt_params: parameter.ParameterDict,
initial_params: parameter.ParameterDict,
confidence_intervals: np.ndarray | None = None,
height_per_param: int = 60,
anchor: str = "",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._anchor = anchor
self._opt_params = opt_params
self._initial_params = initial_params
self._nominal_params = initial_params.copy()
self._nominal_params.reset()
self._x_nominal = self._nominal_params.as_vector()
self._x_hat = self._opt_params.as_vector()
self._x_initial = self._initial_params.as_vector()
self._confidence_intervals = confidence_intervals
self._height_per_param = height_per_param
@property
def template_filename(self) -> str:
return "plot_generic.html"
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
def header_includes(self) -> set[str]:
return {plotly_script_tag()}
def get_context(self) -> dict[str, Any]:
fig = self._build_figure()
config = {
"displayModeBar": True,
"displaylogo": False,
"responsive": True,
"toImageButtonOptions": {
"format": "svg",
"filename": f"{self._title}_distribution",
"height": max(400, len(self._x_hat) * 100),
"width": 1200,
"scale": 1,
},
}
return {
"title": self._title,
"plot_div": fig.to_html(
full_html=False, include_plotlyjs=False, config=config
),
"caption": (
"<b>Visualization Guide:</b><br><b>Error Bars ( — )</b>: 95%"
" Confidence Interval. Smaller is better.<br><span"
" style='color:green'>◆ High Confidence</span>: Interval is &lt;"
" 0.5% of the parameter range (error bar may be"
" invisible).<br><span style='color:blue'>◆ Identified</span>:"
" Standard confidence interval.<br><span style='color:red'>◆"
" Unconstrained</span>: Interval is infinite or larger than range"
" (Red error bar).<br><span style='color:red'>x Nominal</span>,"
" <span style='color:blue'>- - Bounds</span>: Reference"
" values.<br>* parameter is frozen."
),
}
def _build_figure(self) -> go.Figure:
plot_items = []
non_frozen_idx = 0
for param_name, param in self._opt_params.parameters.items():
# Get bounds for this specific parameter
p_min, p_max = param.get_bounds()
if param.size == 1:
flat_value = (
param.value if np.isscalar(param.value) else param.value.item()
)
lb = p_min.item()
ub = p_max.item()
if param.frozen:
plot_items.append({
"name": param_name,
"value": flat_value,
"is_frozen": True,
"lb": lb,
"ub": ub,
"nominal": flat_value if self._x_nominal is not None else None,
})
else:
conf = (
self._confidence_intervals[non_frozen_idx]
if self._confidence_intervals is not None
else None
)
plot_items.append({
"name": param_name,
"value": self._x_hat[non_frozen_idx],
"is_frozen": False,
"lb": lb,
"ub": ub,
"nominal": (
self._x_nominal[non_frozen_idx]
if self._x_nominal is not None
else None
),
"conf": conf,
})
non_frozen_idx += 1
else:
for i in range(param.size):
if param.shape == (param.size,):
element_name = f"{param_name}[{i}]"
val_frozen = param.value[i]
lb = p_min[i]
ub = p_max[i]
else:
multi_idx = np.unravel_index(i, param.shape)
idx_str = ",".join(str(x) for x in multi_idx)
element_name = f"{param_name}[{idx_str}]"
val_frozen = param.value[multi_idx]
lb = p_min[i]
ub = p_max[i]
if param.frozen:
plot_items.append({
"name": element_name,
"value": val_frozen,
"is_frozen": True,
"lb": lb,
"ub": ub,
"nominal": val_frozen if self._x_nominal is not None else None,
})
else:
conf = (
self._confidence_intervals[non_frozen_idx]
if self._confidence_intervals is not None
else None
)
plot_items.append({
"name": element_name,
"value": self._x_hat[non_frozen_idx],
"is_frozen": False,
"lb": lb,
"ub": ub,
"nominal": (
self._x_nominal[non_frozen_idx]
if self._x_nominal is not None
else None
),
"conf": conf,
})
non_frozen_idx += 1
n_params = len(plot_items)
n_cols = 8 if n_params > 1 else 1
n_rows = math.ceil(n_params / n_cols)
# Prepare subplot titles with conditional formatting
subplot_titles = []
for p in plot_items:
if p["is_frozen"]:
subplot_titles.append(f"<span style='color: grey;'>{p['name']}*</span>")
else:
subplot_titles.append(p["name"])
# Create subplots
fig = plt_subplots.make_subplots(
rows=n_rows,
cols=n_cols,
subplot_titles=subplot_titles,
vertical_spacing=0.05,
horizontal_spacing=0.02,
)
# Add dummy trace for Bounds legend
fig.add_trace(
go.Scatter(
x=[None],
y=[None],
mode="lines",
line=dict(color="blue", width=3, dash="dash"),
name="Bounds",
)
)
shown_legends = set()
for i, item in enumerate(plot_items):
row = (i // n_cols) + 1
col = (i % n_cols) + 1
val = item["value"]
lb, ub = item["lb"], item["ub"]
# 1. Bounds lines (For items with valid bounds)
if lb is not None and ub is not None:
# Lower Bound (Horizontal line at y=lb)
fig.add_trace(
go.Scatter(
x=[-1, 1],
y=[lb, lb],
mode="lines",
line=dict(color="blue", width=3, dash="dash"),
showlegend=False,
hoverinfo="skip",
),
row=row,
col=col,
)
# Upper Bound (Horizontal line at y=ub)
fig.add_trace(
go.Scatter(
x=[-1, 1],
y=[ub, ub],
mode="lines",
line=dict(color="blue", width=3, dash="dash"),
showlegend=False,
hoverinfo="skip",
),
row=row,
col=col,
)
# Calculate margin based on bounds
dist = ub - lb
if dist > 0:
margin = dist * 0.2
y_range_min = lb
y_range_max = ub
else:
margin = abs(val) * 0.2 if val != 0 else 1.0
y_range_min = val
y_range_max = val
else:
# Fallback
margin = abs(val) * 0.2 if val != 0 else 1.0
y_range_min = val
y_range_max = val
# 2. Add Traces (Frozen vs Optimized)
if item["is_frozen"]:
# Frozen Parameter
fig.add_trace(
go.Scatter(
x=[0],
y=[val],
mode="markers",
marker=dict(
symbol="circle", size=10, color="gray", opacity=0.7
),
name="Frozen",
showlegend=False,
hoverinfo="skip",
),
row=row,
col=col,
)
else:
# Optimized Parameter
# Determine Confidence Status
trace_name = "Identified"
marker_color = "blue"
legend_group = "identified"
error_y = None
if item.get("conf") is not None:
interval = item["conf"]
min_bound = item["lb"]
max_bound = item["ub"]
error_val = interval
rng = 0
if min_bound is not None and max_bound is not None:
rng = max_bound - min_bound
if not np.isfinite(interval) or (
rng > 0 and 2.0 * interval > 1.0 * rng
):
# Unconstrained
trace_name = "Unconstrained"
marker_color = "red"
legend_group = "unconstrained"
if rng > 0:
error_val = rng
else:
error_val = interval # or large value call fallback
elif rng > 0 and interval <= 0.005 * rng:
trace_name = "High Confidence"
marker_color = "green"
legend_group = "high_conf"
error_bar_color = marker_color
error_y = dict(
type="data",
array=[error_val],
visible=True,
thickness=1.5,
width=3,
color=error_bar_color,
)
show_leg = trace_name not in shown_legends
if show_leg:
shown_legends.add(trace_name)
fig.add_trace(
go.Scatter(
x=[0],
y=[val],
mode="markers",
marker=dict(symbol="diamond", size=12, color=marker_color),
name=trace_name,
showlegend=show_leg,
legendgroup=legend_group,
hovertemplate=(
f"{trace_name}: %{{y}} ±"
f" {item.get('conf', 0):.4g}<extra></extra>"
),
error_y=error_y,
),
row=row,
col=col,
)
# Nominal Value (if exists)
if item["nominal"] is not None:
show_leg_nom = "Nominal" not in shown_legends
if show_leg_nom:
shown_legends.add("Nominal")
fig.add_trace(
go.Scatter(
x=[0],
y=[item["nominal"]],
mode="markers",
marker=dict(symbol="x", size=10, color="red"),
name="Nominal",
showlegend=show_leg_nom,
legendgroup="nominal",
hovertemplate="Nominal: %{y}<extra></extra>",
),
row=row,
col=col,
)
fig.update_yaxes(
range=[y_range_min - margin, y_range_max + margin],
showgrid=True,
zeroline=False,
row=row,
col=col,
)
fig.update_xaxes(
range=[-1, 1],
showgrid=False,
zeroline=False,
showticklabels=False,
row=row,
col=col,
)
# Global Layout
total_height = max(
400, n_rows * 180
) # Increased height per row for better spacing
fig.update_layout(
height=total_height,
template="plotly_white",
margin=dict(l=60, r=60, t=150, b=60), # Increased side margins
autosize=True,
legend=dict(
orientation="h", yanchor="bottom", y=1.02, xanchor="center", x=0.5
),
hovermode="closest",
)
return fig
@@ -0,0 +1,240 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Parameter table report section."""
from typing import Any
from mujoco.sysid._src import parameter
from mujoco.sysid.report.sections.base import ReportSection
import numpy as np
class ParametersTable(ReportSection):
"""Displays a table of identified, initial, and nominal parameters."""
def __init__(
self,
title: str,
opt_params: parameter.ParameterDict,
initial_params: parameter.ParameterDict,
anchor: str = "",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._anchor = anchor
self._opt_params = opt_params
self._initial_params = initial_params
self._nominal_params = initial_params.copy()
self._nominal_params.reset()
self._x_nominal = self._nominal_params.as_vector()
self._x_hat = self._opt_params.as_vector()
self._x_initial = self._initial_params.as_vector()
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
@property
def template_filename(self) -> str:
"""Tells the builder to look for 'parameters_table.html'."""
return "parameters_table.html"
def header_includes(self) -> set[str]:
return {""}
def get_context(self) -> dict[str, Any]:
# Get the vector of non-frozen parameters
non_frozen_vector = self._opt_params.as_vector()
if len(self._x_hat) != non_frozen_vector.size:
raise ValueError(
"Parameter vector lengths don't match. Expected"
f" {non_frozen_vector.size}, got {len(self._x_hat)}"
)
if self._x_nominal is not None:
if len(self._x_nominal) != non_frozen_vector.size:
raise ValueError(
"Parameter vector lengths don't match. Expected"
f" {non_frozen_vector.size}, got {len(self._x_nominal)}"
)
if self._x_initial is not None:
if len(self._x_initial) != non_frozen_vector.size:
raise ValueError(
"Parameter vector lengths don't match. Expected"
f" {non_frozen_vector.size}, got {len(self._x_initial)}"
)
# Compute error metrics if nominal exists
if self._x_nominal is not None:
if self._x_hat.size > 0:
rel_errors = np.abs(self._x_hat - self._x_nominal) / (
np.abs(self._x_nominal) + 1e-8
)
overall_rmse = np.sqrt(np.mean((self._x_hat - self._x_nominal) ** 2))
abs_errors = np.abs(self._x_hat - self._x_nominal)
else:
rel_errors = np.array([])
overall_rmse = 0.0
abs_errors = np.array([])
else:
rel_errors = None
overall_rmse = None
abs_errors = None
def create_table_row(param_name, val, lb, ub, idx=None, is_frozen=False):
"""Create a dict for a table row."""
error_class = ""
# Bounds check
if not is_frozen:
# Check if on boundary
if (abs(val - lb) < 1e-8 + 1e-3 * abs(lb)) or (
abs(val - ub) < 1e-8 + 1e-3 * abs(ub)
):
error_class = "pt_on_boundary"
# If not on boundary, check errors if nominal exists
elif rel_errors is not None and idx is not None:
rel_error = rel_errors[idx]
if rel_error < 0.02:
error_class = "pt_small_error"
elif rel_error < 0.1:
error_class = "pt_medium_error"
else:
error_class = "pt_large_error"
else:
error_class = "pt_frozen"
row = {
"name": param_name + ("*" if is_frozen else ""),
"pred": val,
"lower_bound": lb,
"upper_bound": ub,
"error_class": error_class,
"is_frozen": is_frozen,
}
if self._x_initial is not None:
if is_frozen:
# Initial for frozen is assumed same as val
row["initial"] = val
elif idx is not None:
row["initial"] = self._x_initial[idx]
if self._x_nominal is not None:
if is_frozen:
# Nominal for frozen is assumed same as val
row.update({
"nominal": val,
"abs_err": 0.0,
"rel_err": 0.0,
})
elif (
idx is not None
and abs_errors is not None
and rel_errors is not None
):
row.update({
"nominal": self._x_nominal[idx],
"abs_err": abs_errors[idx],
"rel_err": rel_errors[idx],
})
return row
# Build table data.
table_data = []
non_frozen_idx = 0 # Index for non-frozen parameters in the arrays
for param_name, param in self._opt_params.parameters.items():
# Get bounds for this specific parameter
p_min, p_max = param.get_bounds()
if param.size == 1:
flat_value = (
param.value if np.isscalar(param.value) else param.value.item()
)
lb = p_min.item()
ub = p_max.item()
if param.frozen:
table_data.append(
create_table_row(param_name, flat_value, lb, ub, is_frozen=True)
)
else:
val = self._x_hat[non_frozen_idx]
table_data.append(
create_table_row(param_name, val, lb, ub, idx=non_frozen_idx)
)
non_frozen_idx += 1
else:
for i in range(param.size):
if param.shape == (param.size,):
element_name = f"{param_name}[{i}]"
val_frozen = param.value[i]
lb = p_min[i]
ub = p_max[i]
else:
multi_idx = np.unravel_index(i, param.shape)
idx_str = ",".join(str(x) for x in multi_idx)
element_name = f"{param_name}[{idx_str}]"
val_frozen = param.value[multi_idx]
lb = p_min[multi_idx]
ub = p_max[multi_idx]
if param.frozen:
table_data.append(
create_table_row(
element_name, val_frozen, lb, ub, is_frozen=True
)
)
else:
val = self._x_hat[non_frozen_idx]
table_data.append(
create_table_row(element_name, val, lb, ub, idx=non_frozen_idx)
)
non_frozen_idx += 1
headers = ["Parameter"]
if self._x_initial is not None:
headers.append("Initial")
if self._x_nominal is not None:
headers.append("Nominal")
headers.append("Identified")
headers.extend(["Lower Bound", "Upper Bound"])
if self._x_nominal is not None:
headers.extend(["Absolute Change", "Relative Change"])
return {
"title": self._title,
"headers": headers,
"table_data": table_data,
"overall_rmse": overall_rmse,
"has_initial": self._x_initial is not None,
"has_nominal": self._x_nominal is not None,
}
@@ -0,0 +1,63 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Row layout report section."""
from typing import Any
from mujoco.sysid.report.sections.base import ReportSection
class RowSection(ReportSection):
"""A report section that displays multiple other sections side-by-side."""
def __init__(
self,
title: str,
sections: list[ReportSection],
anchor: str = "",
collapsible: bool = True,
description: str = "",
):
super().__init__(collapsible=collapsible)
self._title = title
self.sections = sections
self._anchor = anchor
self._description = description
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
@property
def template_filename(self) -> str:
return "row.html"
def header_includes(self) -> set[str]:
includes = set()
for section in self.sections:
includes.update(section.header_includes())
return includes
def get_context(self) -> dict[str, Any]:
return {
"title": self.title,
"anchor": self.anchor,
"description": self._description,
}
@@ -0,0 +1,231 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Signal comparison report section."""
import math
from typing import Any
import mujoco
from mujoco.sysid._src import timeseries
from mujoco.sysid.report.sections.base import ReportSection
from mujoco.sysid.report.utils import plotly_script_tag
import numpy as np
from plotly import colors as plt_colors
from plotly import subplots as plt_subplots
import plotly.graph_objects as go
class SignalReport(ReportSection):
"""A report section comparing predicted vs measured observation data."""
def __init__(
self,
title: str,
model: mujoco.MjModel,
size_factor: float = 1.0,
title_prefix: str = "",
max_datapoints: int | None = 300,
resample_to_frequency: float | None = None,
anchor: str = "",
ts_dict=None,
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._anchor = anchor
self._model = model
self._ts_dict = ts_dict if ts_dict is not None else {}
self._size_factor = size_factor
self._title_prefix = title_prefix
self._max_datapoints = max_datapoints
self._resample_to_frequency = resample_to_frequency
self._figure: go.Figure | None = None
if resample_to_frequency is not None and resample_to_frequency <= 0:
raise ValueError(
f"Invalid resample_to_frequency: {resample_to_frequency}"
)
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
@property
def template_filename(self) -> str:
"""Tells the builder to look for 'sensors.html'."""
return "signals.html"
def header_includes(self) -> set[str]:
"""Ensures the Plotly Javascript library is loaded in <head>."""
return {plotly_script_tag()}
def get_context(self) -> dict[str, Any]:
if self._figure is None:
self._figure = self._build_figure()
config = {
"displayModeBar": True,
"displaylogo": False,
"toImageButtonOptions": {
"format": "svg",
"filename": f"{self._title}_signals",
"height": 800,
"width": 1200,
"scale": 1,
},
"responsive": True,
}
context = {"title": self._title}
if self._figure:
plot_html = self._figure.to_html(
full_html=False, include_plotlyjs=False, config=config
)
context["plot_div"] = plot_html
return context
def _build_figure(self) -> go.Figure | None:
"""Internal helper to construct the Plotly figure."""
mapping = None
signal_dict = {}
for ts_name in self._ts_dict:
ts = self._ts_dict[ts_name]
if ts.signal_mapping:
mapping = ts.signal_mapping
signal_dict[ts_name] = self._resample_if_needed(ts)
if not mapping and self._ts_dict:
# Fallback: create mapping from the first time series assuming identity
first_ts_name = next(iter(self._ts_dict))
first_ts = self._ts_dict[first_ts_name]
n_dim = first_ts.data.shape[1]
mapping = {
f"{first_ts_name}_{i}": (first_ts_name, [i]) for i in range(n_dim)
}
if not mapping:
return
signal_names = []
for name in mapping:
signal_dim = len(mapping[name][1])
if signal_dim > 1:
for j in range(signal_dim):
signal_names.append(f"{name}[{j}]")
else:
signal_names.append(name)
n_plots = len(signal_names)
n_cols = 3 if n_plots > 1 else 1
n_rows = (n_plots + n_cols - 1) // n_cols
n_rows += 1
n_params = len(mapping)
n_rows = math.ceil(n_params / n_cols)
fig = plt_subplots.make_subplots(
rows=n_rows,
cols=n_cols,
shared_xaxes=True,
subplot_titles=signal_names,
vertical_spacing=0.5 / n_rows if n_rows > 1 else 0.2,
)
colors = plt_colors.DEFAULT_PLOTLY_COLORS
plot_idx = 0
curr_row = 1
curr_col = 1
for key in mapping:
indices = mapping[key][1]
signal_dim = len(indices)
for j in range(signal_dim):
for color_id, ts_name in enumerate(signal_dict):
predicted_signal = signal_dict[ts_name].data[:, indices]
predicted_times = signal_dict[ts_name].times
curr_row = (plot_idx // n_cols) + 1
curr_col = (plot_idx % n_cols) + 1
fig.add_trace(
go.Scatter(
x=predicted_times,
y=predicted_signal[:, j],
mode="lines",
line=dict(width=2, color=colors[color_id]),
opacity=0.8,
name=ts_name,
legendgroup=ts_name,
showlegend=(plot_idx == 0),
),
row=curr_row,
col=curr_col,
)
fig.update_xaxes(title_text="Time (s)", row=curr_row, col=curr_col)
plot_idx += 1
fig.update_layout(
title_text=f"{self._title_prefix} Signals",
height=max(400, 220 * n_rows * self._size_factor),
autosize=True,
legend=dict(
orientation="h", yanchor="bottom", y=1.15, xanchor="center", x=0.5
),
margin=dict(l=60, r=60, t=150, b=60),
template="plotly_white",
hovermode="x unified",
)
fig.update_xaxes(
showgrid=True,
gridcolor="rgba(211, 211, 211, 0.7)",
showspikes=True,
spikemode="across",
spikesnap="cursor",
showline=True,
linewidth=1,
linecolor="black",
matches="x", # Critical for zooming all subplots together
)
fig.update_yaxes(
showgrid=True,
gridcolor="rgba(211, 211, 211, 0.7)",
showline=True,
linewidth=1,
linecolor="black",
)
return fig
def _resample_if_needed(
self, data: timeseries.TimeSeries
) -> timeseries.TimeSeries:
"""Helper to downsample data for faster/lighter plotting."""
if self._resample_to_frequency:
new_times = np.arange(
data.times[0], data.times[-1], 1.0 / self._resample_to_frequency
)
data = data.resample(new_times)
if self._max_datapoints and len(data.times) > self._max_datapoints:
new_times = np.linspace(
data.times[0], data.times[-1], self._max_datapoints
)
data = data.resample(new_times)
return data
@@ -0,0 +1,216 @@
# Copyright 2026 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Video report section for system identification results."""
from collections.abc import Callable
import os
import pathlib
from typing import Any
import mujoco
import mujoco.rollout
from mujoco.sysid._src import model_modifier
from mujoco.sysid._src import parameter
from mujoco.sysid._src.plotting import render_rollout
from mujoco.sysid._src.trajectory import SystemTrajectory
from mujoco.sysid.report.sections.base import ReportSection
def spec_apply(spec, attrs, values):
def apply_to_geoms_recursive(body):
for g in body.geoms:
for attr, value in zip(attrs, values, strict=True):
setattr(g, attr, value)
for child_body in body.bodies:
apply_to_geoms_recursive(child_body)
for top_body in spec.worldbody.bodies:
apply_to_geoms_recursive(top_body)
def generate_video_from_trajectories(
initial_params: parameter.ParameterDict,
opt_params: parameter.ParameterDict,
_build_model: Callable[
[parameter.ParameterDict, mujoco.MjSpec], mujoco.MjModel
],
trajectories: list[SystemTrajectory],
model_spec: mujoco.MjSpec,
output_filepath: os.PathLike[str],
render_initial: bool = True,
render_nominal: bool = True,
render_opt: bool = True,
height: int = 480,
width: int = 640,
fovy: float = 60,
camera: str | int = -1,
fps: int = 60,
) -> pathlib.Path:
"""Render trajectories and concatenate into a single video.
Each trajectory is rendered with initial/nominal/optimized parameters
overlaid, then all frames are concatenated.
Args:
initial_params: Initial parameter values.
opt_params: Optimized parameter values.
_build_model: Callable to build a model from parameters and spec.
trajectories: List of trajectories to render.
model_spec: MuJoCo model specification.
output_filepath: Path to save the output video.
render_initial: Whether to render with initial parameters.
render_nominal: Whether to render with nominal parameters.
render_opt: Whether to render with optimized parameters.
height: Frame height in pixels.
width: Frame width in pixels.
fovy: Vertical field of view in degrees.
camera: Camera index or name.
fps: Frames per second.
Returns:
Path to the saved video file.
"""
import imageio
all_frames = []
for traj in trajectories:
# Build models for this trajectory
models = []
datas = []
nominal_params = initial_params.copy()
nominal_params.reset()
# initial
if render_initial:
initial_spec = model_spec.copy()
initial_spec = model_modifier.apply_param_modifiers_spec(
initial_params, initial_spec
)
spec_apply(initial_spec, ["rgba"], [[1, 0, 0, 0.5]])
initial_model = initial_spec.compile()
initial_data = mujoco.MjData(initial_model)
models.append(initial_model)
datas.append(initial_data)
# nominal
if render_nominal:
nominal_spec = model_spec.copy()
nominal_spec = model_modifier.apply_param_modifiers_spec(
nominal_params, nominal_spec
)
spec_apply(nominal_spec, ["rgba"], [[0, 1, 0, 0.4]])
nominal_model = nominal_spec.compile()
nominal_data = mujoco.MjData(nominal_model)
models.append(nominal_model)
datas.append(nominal_data)
# pred
if render_opt:
pred_spec = model_spec.copy()
pred_spec = model_modifier.apply_param_modifiers_spec(
opt_params, pred_spec
)
spec_apply(pred_spec, ["rgba"], [[0, 0, 1, 1.0]])
pred_model = pred_spec.compile()
pred_data = mujoco.MjData(pred_model)
models.append(pred_model)
datas.append(pred_data)
control_ts = traj.control.resample(target_dt=models[0].opt.timestep)
state, _ = mujoco.rollout.rollout(
models, datas, traj.initial_state, control_ts.data
)
models[0].vis.global_.fovy = fovy
models[0].vis.global_.offwidth = width
models[0].vis.global_.offheight = height
frames = render_rollout(
models,
datas[0],
state,
framerate=fps,
height=height,
width=width,
camera=camera,
)
all_frames.extend(list(frames))
output_filepath_str = str(output_filepath)
writer = imageio.get_writer(output_filepath_str, fps=fps, quality=8)
for frame in all_frames:
writer.append_data(frame)
writer.close()
return pathlib.Path(output_filepath_str)
class VideoPlayer(ReportSection):
"""A report section to embed and display a video file."""
def __init__(
self,
title: str,
video_filepath: pathlib.Path,
anchor: str = "",
width: int | str = 800,
height: int | None = 450,
autoplay: bool = False,
controls: bool = True,
muted: bool = False,
loop: bool = True,
caption: str = "<b>Legend:</b> <span class='color-initial'>Initial</span>, <span class='color-nominal'>Nominal</span>, <span class='color-optimized'>Optimized</span>",
collapsible: bool = True,
):
super().__init__(collapsible=collapsible)
self._title = title
self._anchor = anchor
self._video_filepath = video_filepath
self._width = width
self._height = height
self._autoplay = autoplay
self._controls = controls
self._muted = muted
self._loop = loop
self._caption = caption
@property
def title(self) -> str:
return self._title
@property
def anchor(self) -> str:
return self._anchor
@property
def template_filename(self) -> str:
"""Tells the builder to look for 'video.html'."""
return "video.html"
def header_includes(self) -> set[str]:
return set()
def get_context(self) -> dict[str, Any]:
"""Returns the data needed to render the video player in the template."""
return {
"title": self._title,
"video_filepath": self._video_filepath.name,
"width": self._width,
"height": self._height,
"autoplay": "autoplay" if self._autoplay else "",
"controls": "controls" if self._controls else "",
"muted": "muted" if self._muted else "",
"loop": "loop" if self._loop else "",
"caption": self._caption,
}