Merge pull request #3079 from aftersomemath:sysid-pr
PiperOrigin-RevId: 868229512 Change-Id: I790bc08fc8b0745583a2f92d9ee2c5a19ba558ea
This commit is contained in:
@@ -0,0 +1,122 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""HTML report builder using Jinja2 templates."""
|
||||
|
||||
# report/builder.py
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import jinja2
|
||||
from mujoco.sysid.report.sections.base import ReportSection
|
||||
|
||||
# Path to the templates directory
|
||||
TEMPLATE_DIR = os.path.join(os.path.dirname(__file__), "templates")
|
||||
|
||||
|
||||
class ReportBuilder:
|
||||
"""Assembles report sections into a complete HTML report."""
|
||||
|
||||
def __init__(self, title: str, global_context: dict[str, Any] | None = None):
|
||||
self._title = title
|
||||
self._sections: list[ReportSection] = []
|
||||
self._global_context = global_context or {}
|
||||
|
||||
# Setup Jinja to load from files
|
||||
self._env = jinja2.Environment(
|
||||
loader=jinja2.FileSystemLoader(TEMPLATE_DIR),
|
||||
autoescape=jinja2.select_autoescape(["html", "xml"]),
|
||||
)
|
||||
|
||||
def add_section(self, section: ReportSection):
|
||||
self._sections.append(section)
|
||||
|
||||
def build(self) -> str:
|
||||
"""Render all sections into a complete HTML report string."""
|
||||
# Load the main layout
|
||||
layout_template = self._env.get_template("layout.html")
|
||||
|
||||
# Render each section individually
|
||||
rendered_sections = []
|
||||
all_header_includes = set()
|
||||
|
||||
for section in self._sections:
|
||||
# Get the specific template for this section
|
||||
sec_template = self._env.get_template(section.template_filename)
|
||||
|
||||
# Check if this is a GroupSection (has 'sections' attribute)
|
||||
extra_context = {}
|
||||
child_sections: list[ReportSection] = getattr(section, "sections", [])
|
||||
if child_sections:
|
||||
child_sections_content = []
|
||||
for child in child_sections:
|
||||
child_template = self._env.get_template(child.template_filename)
|
||||
child_html = child_template.render(child.get_context())
|
||||
all_header_includes.update(child.header_includes())
|
||||
|
||||
# Fallback anchor for child
|
||||
child_anchor = child.anchor
|
||||
if not child_anchor:
|
||||
child_anchor = (
|
||||
child.title.lower()
|
||||
.replace(" ", "-")
|
||||
.replace("[", "")
|
||||
.replace("]", "")
|
||||
.replace("(", "")
|
||||
.replace(")", "")
|
||||
)
|
||||
|
||||
child_sections_content.append({
|
||||
"title": child.title,
|
||||
"content": child_html,
|
||||
"anchor": child_anchor,
|
||||
})
|
||||
extra_context["child_sections"] = child_sections_content
|
||||
|
||||
# Render the section HTML (main wrapper)
|
||||
html_content = sec_template.render(section.get_context() | extra_context)
|
||||
# Collect header requirements (scripts/css)
|
||||
all_header_includes.update(section.header_includes())
|
||||
|
||||
# Fallback anchor generation
|
||||
anchor = section.anchor
|
||||
if not anchor:
|
||||
anchor = (
|
||||
section.title.lower()
|
||||
.replace(" ", "-")
|
||||
.replace("[", "")
|
||||
.replace("]", "")
|
||||
.replace("(", "")
|
||||
.replace(")", "")
|
||||
)
|
||||
|
||||
rendered_sections.append({
|
||||
"title": section.title,
|
||||
"anchor": anchor,
|
||||
"collapsible": section.collapsible,
|
||||
"is_open": section.is_open,
|
||||
"content": html_content,
|
||||
})
|
||||
|
||||
# Render final report
|
||||
return layout_template.render(
|
||||
report_title=self._title,
|
||||
sections=rendered_sections,
|
||||
header_includes=all_header_includes,
|
||||
**self._global_context,
|
||||
)
|
||||
|
||||
def save(self, path: str):
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
f.write(self.build())
|
||||
@@ -0,0 +1,396 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Default report generation for system identification results."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
from mujoco.sysid._src import model_modifier
|
||||
from mujoco.sysid._src import parameter
|
||||
from mujoco.sysid._src import plotting
|
||||
from mujoco.sysid._src.optimize import calculate_intervals
|
||||
from mujoco.sysid._src.residual import BuildModelFn
|
||||
from mujoco.sysid._src.trajectory import ModelSequences
|
||||
from mujoco.sysid.report.builder import ReportBuilder
|
||||
from mujoco.sysid.report.sections.covariance import Covariance
|
||||
from mujoco.sysid.report.sections.optimization_trace import OptimizationTrace
|
||||
from mujoco.sysid.report.sections.parameters import ParametersTable
|
||||
from mujoco.sysid.report.sections.signals import SignalReport
|
||||
import numpy as np
|
||||
import scipy.optimize as scipy_optimize
|
||||
|
||||
|
||||
def default_report(
|
||||
models_sequences: Sequence[ModelSequences],
|
||||
initial_params: parameter.ParameterDict,
|
||||
opt_params: parameter.ParameterDict,
|
||||
residual_fn,
|
||||
opt_result: scipy_optimize.OptimizeResult,
|
||||
title="SysID",
|
||||
save_path=None,
|
||||
build_model: BuildModelFn | None = model_modifier.apply_param_modifiers,
|
||||
generate_videos=True,
|
||||
) -> ReportBuilder:
|
||||
"""Returns a ReportBuilder containing experiment results.
|
||||
|
||||
Users needing a custom report can copy and modify this code.
|
||||
"""
|
||||
from mujoco.sysid.report.sections.group import GroupSection
|
||||
from mujoco.sysid.report.sections.insights import AutomatedInsights
|
||||
from mujoco.sysid.report.sections.parameter_distribution import ParameterDistribution
|
||||
from mujoco.sysid.report.sections.row import RowSection
|
||||
from mujoco.sysid.report.sections.video import generate_video_from_trajectories
|
||||
from mujoco.sysid.report.sections.video import VideoPlayer
|
||||
|
||||
####################################
|
||||
# Build report
|
||||
# Sections:
|
||||
# Fit
|
||||
# Parameter tables
|
||||
# Confidence intervals
|
||||
# Extras: Optimization trace
|
||||
####################################
|
||||
rb = ReportBuilder(title)
|
||||
|
||||
if generate_videos:
|
||||
# 1. Video Player
|
||||
if save_path is None:
|
||||
raise ValueError("save_path is required when generate_videos=True")
|
||||
if build_model is None:
|
||||
raise ValueError("build_model is required when generate_videos=True")
|
||||
|
||||
# Collect ALL trajectories from all model sequences
|
||||
all_trajectories = []
|
||||
for model_sequences in models_sequences:
|
||||
for traj in model_sequences.measured_rollout:
|
||||
all_trajectories.append(traj)
|
||||
|
||||
# Use first model's spec for rendering
|
||||
model_spec_to_render = models_sequences[0].spec
|
||||
|
||||
video_dir = pathlib.Path(save_path)
|
||||
video_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Video 1: All (Initial + Nominal + Optimized)
|
||||
# all trajectories concatenated
|
||||
video_all_path = video_dir / "video_all.mp4"
|
||||
generate_video_from_trajectories(
|
||||
initial_params=initial_params,
|
||||
opt_params=opt_params,
|
||||
build_model=build_model,
|
||||
trajectories=all_trajectories,
|
||||
model_spec=model_spec_to_render,
|
||||
output_filepath=video_all_path,
|
||||
fps=60,
|
||||
)
|
||||
|
||||
# Video 2: Initial + Nominal (no optimized)
|
||||
video_init_path = video_dir / "video_init.mp4"
|
||||
generate_video_from_trajectories(
|
||||
initial_params=initial_params,
|
||||
opt_params=opt_params,
|
||||
build_model=build_model,
|
||||
trajectories=all_trajectories,
|
||||
model_spec=model_spec_to_render,
|
||||
output_filepath=video_init_path,
|
||||
render_opt=False,
|
||||
fps=60,
|
||||
)
|
||||
|
||||
# Video 3: Optimized + Nominal (no initial)
|
||||
video_opt_path = video_dir / "video_opt.mp4"
|
||||
generate_video_from_trajectories(
|
||||
initial_params=initial_params,
|
||||
opt_params=opt_params,
|
||||
build_model=build_model,
|
||||
trajectories=all_trajectories,
|
||||
model_spec=model_spec_to_render,
|
||||
output_filepath=video_opt_path,
|
||||
render_initial=False,
|
||||
fps=60,
|
||||
)
|
||||
|
||||
video_all_section = VideoPlayer(
|
||||
title="All Models",
|
||||
video_filepath=video_all_path,
|
||||
anchor="visual_run_all",
|
||||
autoplay=True,
|
||||
muted=True,
|
||||
width="100%",
|
||||
height=None,
|
||||
caption=(
|
||||
"<span class='color-initial'>Initial</span>, <span"
|
||||
" class='color-nominal'>Nominal</span>, <span"
|
||||
" class='color-optimized'>Optimized</span>"
|
||||
),
|
||||
)
|
||||
|
||||
video_init_section = VideoPlayer(
|
||||
title="Initial vs Nominal",
|
||||
video_filepath=video_init_path,
|
||||
anchor="visual_run_init",
|
||||
autoplay=True,
|
||||
muted=True,
|
||||
width="100%",
|
||||
height=None,
|
||||
caption=(
|
||||
"<span class='color-initial'>Initial</span>, <span"
|
||||
" class='color-nominal'>Nominal</span>"
|
||||
),
|
||||
)
|
||||
|
||||
video_opt_section = VideoPlayer(
|
||||
title="Optimized vs Nominal",
|
||||
video_filepath=video_opt_path,
|
||||
anchor="visual_run_opt",
|
||||
autoplay=True,
|
||||
muted=True,
|
||||
width="100%",
|
||||
height=None,
|
||||
caption=(
|
||||
"<span class='color-nominal'>Nominal</span>, <span"
|
||||
" class='color-optimized'>Optimized</span>"
|
||||
),
|
||||
)
|
||||
|
||||
rb.add_section(
|
||||
RowSection(
|
||||
title="Visual Comparison",
|
||||
sections=[video_all_section, video_init_section, video_opt_section],
|
||||
anchor="visual_comparison",
|
||||
description=(
|
||||
"Visual comparison of the system identification results. The"
|
||||
" nominal model is shown in green, the initial model in red,"
|
||||
" and the optimized model in blue."
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
# 2. Automated Insights (Logs)
|
||||
rb.add_section(AutomatedInsights("Automated Insights", opt_params))
|
||||
|
||||
# 3. Parameters Table (Unified)
|
||||
rb.add_section(
|
||||
ParametersTable(
|
||||
"Parameters", opt_params, initial_params, anchor="Parameters"
|
||||
)
|
||||
)
|
||||
|
||||
# 4. Control Signals (per sequence, grouped like observations)
|
||||
# Get predictions for initial solution.
|
||||
names = [
|
||||
f"{model_sequences.name}\n{sequence}"
|
||||
for model_sequences in models_sequences
|
||||
for sequence in model_sequences.sequence_name
|
||||
]
|
||||
_, pred0s, _ = residual_fn(
|
||||
initial_params.as_vector(), initial_params, return_pred_all=True
|
||||
)
|
||||
|
||||
residuals_star, preds_star, records_star = residual_fn(
|
||||
opt_params.as_vector(), opt_params, return_pred_all=True
|
||||
)
|
||||
|
||||
assert build_model is not None
|
||||
model_hat = build_model(initial_params, models_sequences[0].spec)
|
||||
|
||||
# Build control signal reports for each sequence
|
||||
control_reports = []
|
||||
seq_idx = 0
|
||||
for model_sequences in models_sequences:
|
||||
for i, seq_name in enumerate(model_sequences.sequence_name):
|
||||
ctrl_ts = model_sequences.control[i]
|
||||
name = f"{model_sequences.name}\n{seq_name}"
|
||||
control_reports.append(
|
||||
SignalReport(
|
||||
f"Sequence: {name}",
|
||||
model_hat,
|
||||
title_prefix="",
|
||||
ts_dict={"control": ctrl_ts},
|
||||
collapsible=True,
|
||||
)
|
||||
)
|
||||
seq_idx += 1
|
||||
|
||||
rb.add_section(
|
||||
GroupSection("Control Signals", control_reports, anchor="control_signals")
|
||||
)
|
||||
|
||||
# 5. Observation Signals
|
||||
observation_reports = []
|
||||
for name, pred, record, pred0 in zip(
|
||||
names, preds_star, records_star, pred0s, strict=True
|
||||
):
|
||||
obs_dict = {"initial": pred0[0], "nominal": record[0], "fitted": pred[0]}
|
||||
observation_reports.append(
|
||||
SignalReport(
|
||||
f"Sequence: {name}",
|
||||
model_hat,
|
||||
title_prefix="",
|
||||
ts_dict=obs_dict,
|
||||
collapsible=True,
|
||||
)
|
||||
)
|
||||
|
||||
rb.add_section(
|
||||
GroupSection(
|
||||
"Observation Signals", observation_reports, anchor="observations"
|
||||
)
|
||||
)
|
||||
|
||||
covariance, intervals = calculate_intervals(residuals_star, opt_result.jac)
|
||||
|
||||
# 6. Parameter Distribution
|
||||
rb.add_section(
|
||||
ParameterDistribution(
|
||||
title="Parameter Distribution",
|
||||
opt_params=opt_params,
|
||||
initial_params=initial_params,
|
||||
confidence_intervals=intervals,
|
||||
anchor="param_dist",
|
||||
)
|
||||
)
|
||||
|
||||
rb.add_section(
|
||||
Covariance(
|
||||
title="Covariance and Correlation",
|
||||
anchor="cov",
|
||||
covariance=covariance,
|
||||
parameter_dict=opt_params,
|
||||
)
|
||||
)
|
||||
|
||||
# Add diagnostic optimization trace plots.
|
||||
if "extras" in opt_result:
|
||||
# Add to the report.
|
||||
rb.add_section(
|
||||
OptimizationTrace(
|
||||
title="Optimization Trace",
|
||||
anchor="opt",
|
||||
objective=opt_result.extras.get("objective"),
|
||||
candidate=opt_result.extras.get("candidate"),
|
||||
bounds=opt_params.get_bounds(),
|
||||
param_names=opt_params.get_non_frozen_parameter_names(),
|
||||
)
|
||||
)
|
||||
|
||||
rb.build()
|
||||
if save_path:
|
||||
rb.save(save_path / "report.html")
|
||||
return rb
|
||||
|
||||
|
||||
# TODO(nimrod): Consider deleting this function, given we can export plots from
|
||||
# plotly either on the web or with fig.write_image.
|
||||
def default_report_matplotlib(
|
||||
experiment_results_folder: os.PathLike[str],
|
||||
models_sequences: Sequence[ModelSequences],
|
||||
params: parameter.ParameterDict,
|
||||
sysid_residual,
|
||||
x0: np.ndarray,
|
||||
opt_result: scipy_optimize.OptimizeResult,
|
||||
build_model: BuildModelFn | None = model_modifier.apply_param_modifiers,
|
||||
):
|
||||
"""Outputs PNG plots to the experiment results folder."""
|
||||
experiment_results_folder = pathlib.Path(experiment_results_folder)
|
||||
if not experiment_results_folder.exists():
|
||||
experiment_results_folder.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
x_hat = opt_result.x
|
||||
params.update_from_vector(x_hat)
|
||||
|
||||
# Save the ID'd models out
|
||||
assert build_model is not None
|
||||
model_hat = None
|
||||
for model_sequences in models_sequences:
|
||||
model_hat = build_model(params, model_sequences.spec)
|
||||
assert model_hat is not None
|
||||
|
||||
# Get predictions for initial solution.
|
||||
params.update_from_vector(x0)
|
||||
names = [
|
||||
f"{model_sequences.name}\n{sequence}"
|
||||
for model_sequences in models_sequences
|
||||
for sequence in model_sequences.sequence_name
|
||||
]
|
||||
_, pred0s, record0s = sysid_residual(x0, return_pred_all=True)
|
||||
|
||||
for name, pred0, record0 in zip(names, pred0s, record0s, strict=True):
|
||||
plotting.plot_sensor_comparison(
|
||||
model_hat,
|
||||
predicted_times=pred0[0].times,
|
||||
predicted_data=pred0[0].data,
|
||||
real_times=record0[0].times,
|
||||
real_data=record0[0].data,
|
||||
title_prefix=f"x0 {name}",
|
||||
size_factor=0.5,
|
||||
)
|
||||
name_fig = name.replace("/", " ")
|
||||
name_fig = name_fig.replace("\n", " ")
|
||||
plt.savefig(os.path.join(experiment_results_folder, f"x0 {name_fig}.png"))
|
||||
|
||||
residuals_star, preds_star, records_star = sysid_residual(
|
||||
x_hat, return_pred_all=True
|
||||
)
|
||||
for name, pred, record, _ in zip(
|
||||
names, preds_star, records_star, pred0s, strict=True
|
||||
):
|
||||
plotting.plot_sensor_comparison(
|
||||
model_hat,
|
||||
predicted_times=pred[0].times,
|
||||
predicted_data=pred[0].data,
|
||||
real_times=record[0].times,
|
||||
real_data=record[0].data,
|
||||
title_prefix=f"x* {name}",
|
||||
size_factor=0.5,
|
||||
)
|
||||
name_fig = name.replace("/", " ")
|
||||
name_fig = name_fig.replace("\n", " ")
|
||||
plt.savefig(experiment_results_folder / f"xstar {name_fig}.png")
|
||||
|
||||
# Add diagnostic optimization trace plots.
|
||||
if "extras" in opt_result:
|
||||
# Objective value over iterations.
|
||||
objective = opt_result.extras["objective"]
|
||||
plotting.plot_objective(objective)
|
||||
plt.savefig(experiment_results_folder / "loss.png", dpi=300)
|
||||
|
||||
# Candidate parameter values over iterations.
|
||||
candidate = opt_result.extras["candidate"]
|
||||
|
||||
# Candidate parameter values over iterations.
|
||||
# Candidate heatmap over iterations.
|
||||
plotting.plot_candidate_heatmap(
|
||||
candidate,
|
||||
param_names=params.get_non_frozen_parameter_names(),
|
||||
bounds=params.get_bounds(),
|
||||
)
|
||||
plt.savefig(experiment_results_folder / "candidate_heatmap.png", dpi=300)
|
||||
|
||||
plotting.plot_candidate(
|
||||
candidate,
|
||||
bounds=params.get_bounds(),
|
||||
param_names=params.get_non_frozen_parameter_names(),
|
||||
)
|
||||
plt.savefig(experiment_results_folder / "candidate.png", dpi=300)
|
||||
|
||||
_, intervals = calculate_intervals(residuals_star, opt_result.jac)
|
||||
plotting.parameter_confidence(
|
||||
all_exp_names=["trial"], all_params=[params], all_intervals=[intervals]
|
||||
)
|
||||
# plotting.parameter_confidence(["trial"], [params], [x_hat], [intervals])
|
||||
plt.savefig(experiment_results_folder / "params.png")
|
||||
@@ -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 <"
|
||||
" 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,
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
<style>
|
||||
.cov-and-corr {
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
flex-wrap: wrap;
|
||||
gap: 2rem;
|
||||
justify-content: flex-start;
|
||||
margin-bottom: 1rem;
|
||||
}
|
||||
|
||||
/* Container for a single matrix (covariance or correlation) */
|
||||
.cov-matrix-container {
|
||||
border: 1px solid var(--border-color);
|
||||
padding: 1rem;
|
||||
border-radius: 5px;
|
||||
background-color: var(--card-bg);
|
||||
box-shadow: 2px 2px 5px rgba(0, 0, 0, 0.05);
|
||||
max-width: 100%;
|
||||
overflow-x: auto;
|
||||
}
|
||||
|
||||
.cov-matrix-container h4 {
|
||||
margin-top: 0;
|
||||
margin-bottom: 0.8rem;
|
||||
text-align: left;
|
||||
color: var(--heading-color);
|
||||
font-weight: 600;
|
||||
font-size: 1.1rem;
|
||||
}
|
||||
|
||||
.cov-matrix-container p.explanation {
|
||||
font-size: 0.85rem;
|
||||
color: var(--text-muted);
|
||||
margin-top: 1rem;
|
||||
max-width: 600px;
|
||||
line-height: 1.4;
|
||||
text-align: left;
|
||||
}
|
||||
|
||||
.cov-matrix-container p.explanation strong {
|
||||
color: var(--text-color);
|
||||
}
|
||||
|
||||
table.cov {
|
||||
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, sans-serif;
|
||||
font-size: 0.9rem;
|
||||
color: var(--text-color);
|
||||
text-wrap: nowrap;
|
||||
margin-left: 0;
|
||||
}
|
||||
|
||||
table.cov tbody tr td {
|
||||
padding: 0 4px;
|
||||
text-align: left;
|
||||
border-top: 1px solid var(--border-color);
|
||||
vertical-align: middle;
|
||||
/* Ensure text is readable on colored backgrounds - might need text shadow or specific color logic in python generation */
|
||||
}
|
||||
|
||||
table.cov thead,
|
||||
table.cov thead th {
|
||||
background-color: var(--bg-color);
|
||||
/* Very light grey header background */
|
||||
color: var(--heading-color);
|
||||
font-weight: 600;
|
||||
border-top: none;
|
||||
border-bottom: 2px solid var(--border-color);
|
||||
/* Keep a stronger bottom border for header */
|
||||
text-align: right;
|
||||
}
|
||||
|
||||
/* Body row styles */
|
||||
table.cov tbody tr {
|
||||
background-color: var(--card-bg);
|
||||
}
|
||||
|
||||
table.cov {
|
||||
border-collapse: collapse;
|
||||
margin: auto;
|
||||
margin-bottom: 1rem;
|
||||
}
|
||||
|
||||
table.cov th,
|
||||
table.cov td {
|
||||
border: 1px solid var(--border-color);
|
||||
padding: 0 4px;
|
||||
/* Very compact */
|
||||
text-align: center;
|
||||
font-size: 0.9em;
|
||||
}
|
||||
|
||||
/* Row Headers (first cell in each body row) */
|
||||
table.cov tbody th {
|
||||
background-color: var(--bg-color);
|
||||
font-weight: bold;
|
||||
text-align: left;
|
||||
padding-left: 10px;
|
||||
position: sticky;
|
||||
/* Keep row headers visible */
|
||||
left: 0;
|
||||
z-index: 1;
|
||||
color: var(--heading-color);
|
||||
}
|
||||
|
||||
/* Top-left corner cell */
|
||||
table.cov thead th:first-child {
|
||||
background-color: var(--border-color);
|
||||
/* Slightly darker than bg-color */
|
||||
z-index: 3;
|
||||
/* Ensure it's above row/col headers */
|
||||
position: sticky;
|
||||
left: 0;
|
||||
/* Also sticky */
|
||||
}
|
||||
|
||||
table.cov td span {
|
||||
display: inline-block;
|
||||
min-width: 100%;
|
||||
}
|
||||
|
||||
table.cov td.blank {
|
||||
border: 0;
|
||||
background-color: var(--card-bg);
|
||||
}
|
||||
</style>
|
||||
<div class="cov-and-corr">
|
||||
{% if message %}
|
||||
<div class="cov-matrix-container" style="width: 100%;">
|
||||
<h4>Covariance & Correlation</h4>
|
||||
<p class="explanation" style="color: var(--text-color);">
|
||||
{{ message }}
|
||||
</p>
|
||||
</div>
|
||||
{% else %}
|
||||
<div class="cov-matrix-container">
|
||||
<h4>Covariance Matrix</h4>
|
||||
<table class="cov">
|
||||
<tbody>
|
||||
{% for row in covariance_data %}
|
||||
{% set outer_loop = loop %}
|
||||
<tr>
|
||||
<th>{{dim_names[outer_loop.index0] }}</th>
|
||||
{% for cell in row %}
|
||||
{%if loop.index0 <= outer_loop.index0 %} <td
|
||||
style="background-color: {{ cell.bgcolor }}; color: {{ cell.textcolor }};" title="{{ cell.value }}" {%if
|
||||
outer_loop.index0==loop.index0 %}class="diag" {% endif %}>
|
||||
<span>
|
||||
{{ "%.1e" | format(cell.value) }}
|
||||
</span>
|
||||
</td>
|
||||
{% else %}
|
||||
<td class="blank" />
|
||||
{% endif %}
|
||||
{% endfor %}
|
||||
</tr>
|
||||
{% endfor %}
|
||||
</tbody>
|
||||
</table>
|
||||
<p class="explanation">
|
||||
The covariance matrix shows how pairs of parameters vary together.
|
||||
The diagonal elements (Variance) show the variance of each parameter (how much it varies on its own).
|
||||
Larger positive values indicate the parameter estimate is less certain.
|
||||
The off-diagonal elements (Covariance) show the covariance between pairs of parameters, better viewed in the
|
||||
normalized correlation matrix.
|
||||
</p>
|
||||
</div> {# End cov-matrix-container for covariance #}
|
||||
|
||||
<div class="cov-matrix-container">
|
||||
<h4>Correlation Matrix</h4>
|
||||
<table class="cov">
|
||||
<tbody>
|
||||
{% for row in correlation_data %}
|
||||
{% set outer_loop = loop %}
|
||||
<tr>
|
||||
<th>{{ dim_names[outer_loop.index0] }}</th>
|
||||
{% for cell in row %}
|
||||
{%if loop.index0 <= outer_loop.index0 %} <td
|
||||
style="background-color: {{ cell.bgcolor }}; color: {{ cell.textcolor }};" title="{{ cell.value }}" {%if
|
||||
outer_loop.index0==loop.index0 %}class="diag" {% endif %}>
|
||||
<span>
|
||||
{{ "%.2f" | format(cell.value) }}
|
||||
</span>
|
||||
</td>
|
||||
{% else %}
|
||||
<td class="blank" />
|
||||
{% endif %}
|
||||
{% endfor %}
|
||||
</tr>
|
||||
{% endfor %}
|
||||
</tbody>
|
||||
</table>
|
||||
<p class="explanation">
|
||||
The correlation matrix is a normalized version of the covariance matrix.
|
||||
Values range from -1 to +1.
|
||||
The diagonal elements are always 1, representing perfect correlation of a parameter with itself.
|
||||
Off-diagonal elements (Correlation Coefficient) show the linear correlation between pairs of parameters.
|
||||
+1 indicates perfect positive correlation, -1 indicates perfect negative correlation, and 0 indicates no linear
|
||||
correlation.
|
||||
Values close to +1 or -1 suggests that the parameters are highly dependent. Thus, if their variance is also high,
|
||||
their confidence intervals will be large.
|
||||
</p>
|
||||
</div>
|
||||
{% endif %}
|
||||
</div>
|
||||
@@ -0,0 +1,13 @@
|
||||
<div class="group-section">
|
||||
{% for child in child_sections %}
|
||||
<details class="group-item" open
|
||||
style="margin-bottom: 1rem; border: 1px solid var(--border-color); border-radius: 4px; padding: 1rem;">
|
||||
<summary style="font-size: 1.2rem; font-weight: 600; margin-bottom: 0.5rem; border-bottom: none;">
|
||||
{{ child.title }}
|
||||
</summary>
|
||||
<div class="group-content">
|
||||
{{ child.content | safe }}
|
||||
</div>
|
||||
</details>
|
||||
{% endfor %}
|
||||
</div>
|
||||
@@ -0,0 +1,90 @@
|
||||
<div class="insights-container">
|
||||
{% for insight in insights %}
|
||||
<div class="insight-card insight-{{ insight.type }}">
|
||||
<div class="insight-icon">
|
||||
{% if insight.type == 'warning' %}
|
||||
⚠️
|
||||
{% elif insight.type == 'success' %}
|
||||
✅
|
||||
{% else %}
|
||||
ℹ️
|
||||
{% endif %}
|
||||
</div>
|
||||
<div class="insight-content">
|
||||
<span class="insight-message">
|
||||
<strong>{{ insight.title }}:</strong> {{ insight.message | safe }}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
{% endfor %}
|
||||
</div>
|
||||
|
||||
<style>
|
||||
.insights-container {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 0.25rem;
|
||||
/* Reduced gap */
|
||||
}
|
||||
|
||||
.insight-card {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
/* Center Vertically */
|
||||
padding: 0.25rem 0.5rem;
|
||||
/* Smaller padding */
|
||||
border-radius: 0.25rem;
|
||||
border: 1px solid transparent;
|
||||
background-color: var(--bg-color);
|
||||
font-size: 0.9rem;
|
||||
/* Smaller font */
|
||||
min-height: 1.5rem;
|
||||
/* Ensure roughly one line height */
|
||||
}
|
||||
|
||||
.insight-icon {
|
||||
font-size: 1rem;
|
||||
margin-right: 0.5rem;
|
||||
flex-shrink: 0;
|
||||
line-height: 1;
|
||||
}
|
||||
|
||||
.insight-content {
|
||||
flex-grow: 1;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
}
|
||||
|
||||
.insight-message {
|
||||
margin: 0;
|
||||
color: var(--text-color);
|
||||
/* Use standard text color */
|
||||
white-space: nowrap;
|
||||
/* Keep on one line if possible, or truncate? User said 'could just be one lines height' */
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
}
|
||||
|
||||
/* Allow wrapping if very long but keeping compact verticality */
|
||||
.insight-message {
|
||||
white-space: normal;
|
||||
}
|
||||
|
||||
|
||||
/* Types */
|
||||
.insight-warning {
|
||||
background-color: rgba(255, 193, 7, 0.1);
|
||||
border-color: rgba(255, 193, 7, 0.4);
|
||||
}
|
||||
|
||||
[data-theme="dark"] .insight-warning {
|
||||
background-color: rgba(255, 193, 7, 0.05);
|
||||
/* slightly lighter in dark mode */
|
||||
}
|
||||
|
||||
|
||||
.insight-success {
|
||||
background-color: rgba(25, 135, 84, 0.1);
|
||||
border-color: rgba(25, 135, 84, 0.4);
|
||||
}
|
||||
</style>
|
||||
@@ -0,0 +1,699 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>{{ report_title }}</title>
|
||||
<style>
|
||||
@import url('https://fonts.googleapis.com/css2?family=DM+Sans:ital,wght@0,400;0,500;0,700;1,400;1,500;1,700&family=JetBrains+Mono:wght@400;700&display=swap');
|
||||
|
||||
:root {
|
||||
/* Light Theme Defaults */
|
||||
--bg-color: #f4f6f8;
|
||||
--card-bg: #ffffff;
|
||||
--text-color: #212529;
|
||||
--text-muted: #6c757d;
|
||||
--heading-color: #111;
|
||||
--primary-color: #0d6efd;
|
||||
--border-color: #dee2e6;
|
||||
--shadow-sm: 0 0.125rem 0.25rem rgba(0, 0, 0, 0.075);
|
||||
--toc-link-bg: #f4f6f8;
|
||||
--toc-link-hover: #0d6efd;
|
||||
--toc-link-hover-text: #fff;
|
||||
/* Semantic Colors - Light Default */
|
||||
--color-initial: #d32f2f;
|
||||
--color-nominal: #388e3c;
|
||||
--color-optimized: #1976d2;
|
||||
--color-success: #2e7d32;
|
||||
--color-warning: #f57c00;
|
||||
--color-error: #c62828;
|
||||
--color-boundary: #7b1fa2;
|
||||
}
|
||||
|
||||
[data-theme="dark"] {
|
||||
/* Dark Theme Overrides */
|
||||
--bg-color: #121212;
|
||||
--card-bg: #1e1e1e;
|
||||
--text-color: #e0e0e0;
|
||||
--text-muted: #a0a0a0;
|
||||
--heading-color: #fff;
|
||||
--primary-color: #6ea8fe;
|
||||
--border-color: #333;
|
||||
--shadow-sm: 0 0.125rem 0.25rem rgba(0, 0, 0, 0.3);
|
||||
--toc-link-bg: #2d2d2d;
|
||||
--toc-link-hover: #6ea8fe;
|
||||
--toc-link-hover-text: #121212;
|
||||
/* Semantic Colors - Dark Default */
|
||||
--color-initial: #ef5350;
|
||||
--color-nominal: #66bb6a;
|
||||
--color-optimized: #42a5f5;
|
||||
--color-success: #4caf50;
|
||||
--color-warning: #ff9800;
|
||||
--color-error: #f44336;
|
||||
--color-boundary: #ab47bc;
|
||||
}
|
||||
|
||||
/* Colorblind-friendly overrides (applied via data-colorscheme) */
|
||||
[data-colorscheme="colorblind"] {
|
||||
--color-initial: #d55e00;
|
||||
--color-nominal: #0072b2;
|
||||
--color-optimized: #cc79a7;
|
||||
--color-success: #009e73;
|
||||
--color-warning: #f0e442;
|
||||
--color-error: #d55e00;
|
||||
--color-boundary: #cc79a7;
|
||||
}
|
||||
|
||||
/* Semantic color helper classes for inline use */
|
||||
.color-initial {
|
||||
color: var(--color-initial);
|
||||
}
|
||||
|
||||
.color-nominal {
|
||||
color: var(--color-nominal);
|
||||
}
|
||||
|
||||
.color-optimized {
|
||||
color: var(--color-optimized);
|
||||
}
|
||||
|
||||
.color-success {
|
||||
color: var(--color-success);
|
||||
}
|
||||
|
||||
.color-warning {
|
||||
color: var(--color-warning);
|
||||
}
|
||||
|
||||
.color-error {
|
||||
color: var(--color-error);
|
||||
}
|
||||
|
||||
.color-boundary {
|
||||
color: var(--color-boundary);
|
||||
}
|
||||
|
||||
* {
|
||||
box-sizing: border-box;
|
||||
}
|
||||
|
||||
body {
|
||||
font-family: 'DM Sans', -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
|
||||
font-size: 16px;
|
||||
line-height: 1.6;
|
||||
color: var(--text-color);
|
||||
background-color: var(--bg-color);
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
-webkit-font-smoothing: antialiased;
|
||||
transition: background-color 0.3s, color 0.3s;
|
||||
}
|
||||
|
||||
h1,
|
||||
h2,
|
||||
h3,
|
||||
h4,
|
||||
h5,
|
||||
h6 {
|
||||
margin-top: 0;
|
||||
margin-bottom: 0.5rem;
|
||||
font-weight: 700;
|
||||
line-height: 1.2;
|
||||
color: var(--heading-color);
|
||||
}
|
||||
|
||||
h1 {
|
||||
font-size: 2.25rem;
|
||||
margin-bottom: 0.5rem;
|
||||
}
|
||||
|
||||
h2 {
|
||||
font-size: 1.75rem;
|
||||
margin-bottom: 1rem;
|
||||
padding-bottom: 0.5rem;
|
||||
border-bottom: 2px solid var(--bg-color);
|
||||
}
|
||||
|
||||
h3 {
|
||||
font-size: 1.5rem;
|
||||
}
|
||||
|
||||
a {
|
||||
color: var(--primary-color);
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
a:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
.container {
|
||||
max-width: var(--max-width);
|
||||
margin: 0 auto;
|
||||
padding: 2rem;
|
||||
}
|
||||
|
||||
header.report-header {
|
||||
background-color: var(--card-bg);
|
||||
padding: 1.5rem 0;
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
margin-bottom: 2rem;
|
||||
box-shadow: var(--shadow-sm);
|
||||
transition: background-color 0.3s, border-color 0.3s;
|
||||
}
|
||||
|
||||
header .container {
|
||||
padding-top: 0;
|
||||
padding-bottom: 0;
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.report-meta {
|
||||
color: var(--text-muted);
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
/* CARD STYLES */
|
||||
section {
|
||||
background-color: var(--card-bg);
|
||||
border-radius: 0.5rem;
|
||||
box-shadow: var(--shadow-sm);
|
||||
margin-bottom: 2rem;
|
||||
padding: 2rem;
|
||||
border: 1px solid var(--border-color);
|
||||
transition: background-color 0.3s, border-color 0.3s;
|
||||
overflow: hidden;
|
||||
/* Prevent spillover */
|
||||
}
|
||||
|
||||
/* Make content responsive */
|
||||
section>div,
|
||||
section>table,
|
||||
section>.plotly-graph-div {
|
||||
overflow-x: auto;
|
||||
max-width: 100%;
|
||||
}
|
||||
|
||||
/* TOC STYLES */
|
||||
.toc-card {
|
||||
background-color: var(--card-bg);
|
||||
border-radius: 0.5rem;
|
||||
padding: 1.5rem;
|
||||
margin-bottom: 2rem;
|
||||
border: 1px solid var(--border-color);
|
||||
transition: background-color 0.3s, border-color 0.3s;
|
||||
}
|
||||
|
||||
.toc-card h3 {
|
||||
margin-top: 0;
|
||||
font-size: 1.25rem;
|
||||
color: var(--text-muted);
|
||||
text-transform: uppercase;
|
||||
letter-spacing: 0.05em;
|
||||
}
|
||||
|
||||
.toc-card ul {
|
||||
list-style: none;
|
||||
padding-left: 0;
|
||||
margin: 0;
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 0.5rem;
|
||||
}
|
||||
|
||||
.toc-card li a {
|
||||
display: inline-block;
|
||||
padding: 0.5rem 1rem;
|
||||
background: var(--toc-link-bg);
|
||||
border-radius: 2rem;
|
||||
font-weight: 500;
|
||||
font-size: 0.9rem;
|
||||
transition: all 0.2s;
|
||||
color: var(--text-color);
|
||||
}
|
||||
|
||||
.toc-card li a:hover {
|
||||
background: var(--toc-link-hover);
|
||||
color: var(--toc-link-hover-text);
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
/* DETAILS / SUMMARY overrides */
|
||||
details {
|
||||
margin-top: 0;
|
||||
border: none;
|
||||
background: none;
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
details>summary {
|
||||
font-size: 1.75rem;
|
||||
font-weight: 700;
|
||||
cursor: pointer;
|
||||
margin-bottom: 1rem;
|
||||
padding-bottom: 0.5rem;
|
||||
border-bottom: 2px solid var(--bg-color);
|
||||
list-style: none;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
color: var(--heading-color);
|
||||
}
|
||||
|
||||
details>summary::-webkit-details-marker {
|
||||
display: none;
|
||||
}
|
||||
|
||||
details>summary::after {
|
||||
content: '+';
|
||||
margin-left: auto;
|
||||
font-weight: 400;
|
||||
color: var(--text-muted);
|
||||
}
|
||||
|
||||
details[open]>summary::after {
|
||||
content: '−';
|
||||
}
|
||||
|
||||
/* Code/Mono font improvements */
|
||||
code,
|
||||
pre {
|
||||
font-family: 'JetBrains Mono', source-code-pro, Menlo, Monaco, Consolas, "Courier New", monospace;
|
||||
}
|
||||
|
||||
/* Toggle Button */
|
||||
#theme-toggle,
|
||||
#colorscheme-toggle {
|
||||
background: none;
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: 2rem;
|
||||
padding: 0.5rem 1rem;
|
||||
cursor: pointer;
|
||||
font-family: inherit;
|
||||
font-size: 0.9rem;
|
||||
color: var(--text-color);
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 0.5rem;
|
||||
}
|
||||
|
||||
#theme-toggle:hover,
|
||||
#colorscheme-toggle:hover {
|
||||
background-color: var(--bg-color);
|
||||
}
|
||||
|
||||
/* Header Controls */
|
||||
.header-controls {
|
||||
display: flex;
|
||||
gap: 0.5rem;
|
||||
align-items: center;
|
||||
}
|
||||
|
||||
/* Color Scheme Dropdown */
|
||||
.color-scheme-dropdown {
|
||||
position: relative;
|
||||
}
|
||||
|
||||
.color-panel {
|
||||
position: absolute;
|
||||
top: calc(100% + 0.5rem);
|
||||
right: 0;
|
||||
background: var(--card-bg);
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: 0.75rem;
|
||||
box-shadow: 0 10px 40px rgba(0, 0, 0, 0.2);
|
||||
padding: 1rem;
|
||||
z-index: 1000;
|
||||
min-width: 220px;
|
||||
opacity: 0;
|
||||
transform: translateY(-10px) scale(0.95);
|
||||
pointer-events: none;
|
||||
transition: opacity 0.2s ease, transform 0.2s ease;
|
||||
}
|
||||
|
||||
.color-panel.visible {
|
||||
opacity: 1;
|
||||
transform: translateY(0) scale(1);
|
||||
pointer-events: auto;
|
||||
}
|
||||
|
||||
.color-panel-title {
|
||||
font-size: 0.75rem;
|
||||
font-weight: 600;
|
||||
text-transform: uppercase;
|
||||
letter-spacing: 0.05em;
|
||||
color: var(--text-muted);
|
||||
margin-bottom: 0.75rem;
|
||||
}
|
||||
|
||||
.color-option {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 0.75rem;
|
||||
padding: 0.625rem 0.75rem;
|
||||
border: none;
|
||||
background: none;
|
||||
border-radius: 0.5rem;
|
||||
cursor: pointer;
|
||||
width: 100%;
|
||||
text-align: left;
|
||||
color: var(--text-color);
|
||||
font-family: inherit;
|
||||
font-size: 0.9rem;
|
||||
transition: background-color 0.15s ease;
|
||||
}
|
||||
|
||||
.color-option:hover {
|
||||
background-color: var(--bg-color);
|
||||
}
|
||||
|
||||
.color-option.active {
|
||||
background-color: var(--primary-color);
|
||||
color: white;
|
||||
}
|
||||
|
||||
.color-swatch {
|
||||
display: flex;
|
||||
gap: 3px;
|
||||
}
|
||||
|
||||
.color-swatch span {
|
||||
width: 14px;
|
||||
height: 14px;
|
||||
border-radius: 3px;
|
||||
}
|
||||
|
||||
.swatch-default .c1 {
|
||||
background: #ef5350;
|
||||
}
|
||||
|
||||
.swatch-default .c2 {
|
||||
background: #66bb6a;
|
||||
}
|
||||
|
||||
.swatch-default .c3 {
|
||||
background: #42a5f5;
|
||||
}
|
||||
|
||||
.swatch-colorblind .c1 {
|
||||
background: #d55e00;
|
||||
}
|
||||
|
||||
.swatch-colorblind .c2 {
|
||||
background: #0072b2;
|
||||
}
|
||||
|
||||
.swatch-colorblind .c3 {
|
||||
background: #cc79a7;
|
||||
}
|
||||
|
||||
/* Back to Top */
|
||||
#back-to-top {
|
||||
position: fixed;
|
||||
bottom: 2rem;
|
||||
right: 2rem;
|
||||
background-color: var(--primary-color);
|
||||
color: #fff;
|
||||
border: none;
|
||||
border-radius: 50%;
|
||||
width: 3rem;
|
||||
height: 3rem;
|
||||
display: none;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
cursor: pointer;
|
||||
box-shadow: var(--shadow-md);
|
||||
z-index: 1000;
|
||||
font-size: 1.5rem;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
|
||||
#back-to-top:hover {
|
||||
transform: translateY(-3px);
|
||||
}
|
||||
|
||||
/* Responsive */
|
||||
@media (max-width: 768px) {
|
||||
.container {
|
||||
padding: 1rem;
|
||||
}
|
||||
|
||||
h1 {
|
||||
font-size: 1.75rem;
|
||||
}
|
||||
|
||||
section {
|
||||
padding: 1.5rem;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
|
||||
{% for include in header_includes %}
|
||||
{{ include | safe }}
|
||||
{% endfor %}
|
||||
</head>
|
||||
|
||||
<body>
|
||||
|
||||
<header class="report-header">
|
||||
<div class="container">
|
||||
<div>
|
||||
<h1>{{ report_title }}</h1>
|
||||
{% if report_timestamp %}
|
||||
<div class="report-meta">Generated at {{ report_timestamp }}</div>
|
||||
{% endif %}
|
||||
</div>
|
||||
<div class="header-controls">
|
||||
<button id="theme-toggle" onclick="toggleTheme()">
|
||||
<span>🌗</span> Theme
|
||||
</button>
|
||||
<div class="color-scheme-dropdown">
|
||||
<button id="colorscheme-toggle" onclick="toggleColorSchemePanel()">
|
||||
<span>🎨</span> Colors
|
||||
</button>
|
||||
<div id="color-panel" class="color-panel">
|
||||
<div class="color-panel-title">Color Scheme</div>
|
||||
<button class="color-option active" data-scheme="default" onclick="setColorScheme('default')">
|
||||
<div class="color-swatch swatch-default">
|
||||
<span class="c1"></span>
|
||||
<span class="c2"></span>
|
||||
<span class="c3"></span>
|
||||
</div>
|
||||
Default
|
||||
</button>
|
||||
<button class="color-option" data-scheme="colorblind" onclick="setColorScheme('colorblind')">
|
||||
<div class="color-swatch swatch-colorblind">
|
||||
<span class="c1"></span>
|
||||
<span class="c2"></span>
|
||||
<span class="c3"></span>
|
||||
</div>
|
||||
Colorblind Safe
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<div class="container">
|
||||
|
||||
<nav class="toc-card">
|
||||
<h3>Contents</h3>
|
||||
<ul>
|
||||
{% for sec in sections %}
|
||||
<li><a href="#{{ sec.anchor }}">{{ sec.title }}</a></li>
|
||||
{% endfor %}
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
<main>
|
||||
{% for sec in sections %}
|
||||
<section id="{{ sec.anchor }}">
|
||||
{% if sec.collapsible %}
|
||||
<details {% if sec.is_open %}open{% endif %}>
|
||||
<summary>{{ sec.title }}</summary>
|
||||
{{ sec.content | safe }}
|
||||
</details>
|
||||
{% else %}
|
||||
<h2>{{ sec.title }}</h2>
|
||||
{{ sec.content | safe }}
|
||||
{% endif %}
|
||||
</section>
|
||||
{% endfor %}
|
||||
</main>
|
||||
|
||||
</div>
|
||||
|
||||
<button id="back-to-top" onclick="window.scrollTo({top: 0, behavior: 'smooth'})" title="Back to Top">
|
||||
↑
|
||||
</button>
|
||||
|
||||
<script>
|
||||
// Theme Logic
|
||||
function toggleTheme() {
|
||||
const body = document.body;
|
||||
const currentTheme = body.getAttribute('data-theme');
|
||||
const newTheme = currentTheme === 'dark' ? 'light' : 'dark';
|
||||
body.setAttribute('data-theme', newTheme);
|
||||
localStorage.setItem('theme', newTheme);
|
||||
updatePlotsTheme(newTheme);
|
||||
}
|
||||
|
||||
// Update Plotly plots
|
||||
function updatePlotsTheme(theme) {
|
||||
const isDark = theme === 'dark';
|
||||
const bgColor = isDark ? '#1e1e1e' : '#ffffff';
|
||||
const fontColor = isDark ? '#e0e0e0' : '#212529';
|
||||
const gridColor = isDark ? '#444' : 'rgba(211, 211, 211, 0.7)';
|
||||
const lineColor = isDark ? '#666' : 'black';
|
||||
|
||||
const update = {
|
||||
'paper_bgcolor': bgColor,
|
||||
'plot_bgcolor': bgColor,
|
||||
'font.color': fontColor
|
||||
};
|
||||
|
||||
const plots = document.getElementsByClassName('plotly-graph-div');
|
||||
for (let i = 0; i < plots.length; i++) {
|
||||
const div = plots[i];
|
||||
// dynamic update object
|
||||
const layoutUpdate = { ...update };
|
||||
|
||||
// Iterate over layout keys to find axes
|
||||
if (div.layout) {
|
||||
Object.keys(div.layout).forEach(key => {
|
||||
if (key.match(/^[xy]axis\d*$/)) {
|
||||
layoutUpdate[`${key}.gridcolor`] = gridColor;
|
||||
layoutUpdate[`${key}.linecolor`] = lineColor;
|
||||
layoutUpdate[`${key}.zerolinecolor`] = gridColor;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
Plotly.relayout(div, layoutUpdate);
|
||||
}
|
||||
}
|
||||
|
||||
// Load saved theme
|
||||
const savedTheme = localStorage.getItem('theme');
|
||||
let startTheme = 'light';
|
||||
if (savedTheme) {
|
||||
document.body.setAttribute('data-theme', savedTheme);
|
||||
startTheme = savedTheme;
|
||||
} else if (window.matchMedia && window.matchMedia('(prefers-color-scheme: dark)').matches) {
|
||||
document.body.setAttribute('data-theme', 'dark');
|
||||
startTheme = 'dark';
|
||||
}
|
||||
|
||||
// Color Scheme Logic
|
||||
function toggleColorSchemePanel() {
|
||||
const panel = document.getElementById('color-panel');
|
||||
panel.classList.toggle('visible');
|
||||
}
|
||||
|
||||
function setColorScheme(scheme) {
|
||||
// Update body attribute
|
||||
if (scheme === 'colorblind') {
|
||||
document.body.setAttribute('data-colorscheme', 'colorblind');
|
||||
} else {
|
||||
document.body.removeAttribute('data-colorscheme');
|
||||
}
|
||||
localStorage.setItem('colorscheme', scheme);
|
||||
|
||||
// Update active state on buttons
|
||||
document.querySelectorAll('.color-option').forEach(btn => {
|
||||
btn.classList.toggle('active', btn.dataset.scheme === scheme);
|
||||
});
|
||||
|
||||
// Close panel
|
||||
document.getElementById('color-panel').classList.remove('visible');
|
||||
}
|
||||
|
||||
// Close panel when clicking outside
|
||||
document.addEventListener('click', (e) => {
|
||||
const dropdown = document.querySelector('.color-scheme-dropdown');
|
||||
if (dropdown && !dropdown.contains(e.target)) {
|
||||
document.getElementById('color-panel').classList.remove('visible');
|
||||
}
|
||||
});
|
||||
|
||||
// Load saved color scheme
|
||||
const savedColorScheme = localStorage.getItem('colorscheme');
|
||||
if (savedColorScheme === 'colorblind') {
|
||||
document.body.setAttribute('data-colorscheme', 'colorblind');
|
||||
document.querySelectorAll('.color-option').forEach(btn => {
|
||||
btn.classList.toggle('active', btn.dataset.scheme === 'colorblind');
|
||||
});
|
||||
}
|
||||
|
||||
// Apply plot theme on load (wait for plotly to init?)
|
||||
// Plotly scripts usually run inline, so we might need to wait a tick or use an event.
|
||||
// However, the report generates static HTML where the Plotly.newPlot is called in script tags.
|
||||
// Those script tags might run after this one depending on placement.
|
||||
// layout.html ends with this script, and content is inserted before it?
|
||||
// StartTheme application:
|
||||
window.addEventListener('DOMContentLoaded', () => {
|
||||
// Observers to catch plots as they load if they are async,
|
||||
// but typically jinja renders them inline.
|
||||
// We'll try running it immediately and also after a short delay to be safe.
|
||||
updatePlotsTheme(startTheme);
|
||||
setTimeout(() => updatePlotsTheme(startTheme), 500);
|
||||
});
|
||||
|
||||
// Back to Top Logic
|
||||
const backToTopBtn = document.getElementById('back-to-top');
|
||||
window.addEventListener('scroll', () => {
|
||||
if (window.scrollY > 300) {
|
||||
backToTopBtn.style.display = 'flex';
|
||||
} else {
|
||||
backToTopBtn.style.display = 'none';
|
||||
}
|
||||
});
|
||||
|
||||
// Video Synchronization
|
||||
const videos = document.querySelectorAll('video');
|
||||
let ignoreEvents = false;
|
||||
|
||||
videos.forEach(video => {
|
||||
video.addEventListener('play', (e) => {
|
||||
if (ignoreEvents) return;
|
||||
ignoreEvents = true;
|
||||
videos.forEach(v => {
|
||||
if (v !== e.target && v.paused) {
|
||||
v.play();
|
||||
}
|
||||
});
|
||||
setTimeout(() => ignoreEvents = false, 10);
|
||||
});
|
||||
|
||||
video.addEventListener('pause', (e) => {
|
||||
if (ignoreEvents) return;
|
||||
ignoreEvents = true;
|
||||
videos.forEach(v => {
|
||||
if (v !== e.target && !v.paused) {
|
||||
v.pause();
|
||||
}
|
||||
});
|
||||
setTimeout(() => ignoreEvents = false, 10);
|
||||
});
|
||||
|
||||
video.addEventListener('seeking', (e) => {
|
||||
if (ignoreEvents) return;
|
||||
ignoreEvents = true;
|
||||
const currentTime = e.target.currentTime;
|
||||
videos.forEach(v => {
|
||||
if (v !== e.target && Math.abs(v.currentTime - currentTime) > 0.05) {
|
||||
v.currentTime = currentTime;
|
||||
}
|
||||
});
|
||||
setTimeout(() => ignoreEvents = false, 10);
|
||||
});
|
||||
});
|
||||
</script>
|
||||
|
||||
</body>
|
||||
|
||||
</html>
|
||||
@@ -0,0 +1,20 @@
|
||||
{% if objective_plot_html %}
|
||||
<h3>Objective Function</h3>
|
||||
<div>
|
||||
{{ objective_plot_html | safe }}
|
||||
</div>
|
||||
{% endif %}
|
||||
{% if candidate_heatmap_html %}
|
||||
<h3>Parameter Candidate Heatmap</h3>
|
||||
<div>
|
||||
{{ candidate_heatmap_html | safe }}
|
||||
</div>
|
||||
{% endif %}
|
||||
{% if candidate_plots_html %}
|
||||
<h3>Parameter Candidates</h3>
|
||||
{% for plot_html in candidate_plots_html %}
|
||||
<div>
|
||||
{{ plot_html | safe }}
|
||||
</div>
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
@@ -0,0 +1,3 @@
|
||||
<div>
|
||||
{{ plot_html | safe }}
|
||||
</div>
|
||||
@@ -0,0 +1,147 @@
|
||||
<style>
|
||||
table.pt {
|
||||
width: 100%;
|
||||
border-collapse: collapse;
|
||||
margin-bottom: 1rem;
|
||||
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, sans-serif;
|
||||
font-size: 0.9rem;
|
||||
color: var(--text-color);
|
||||
}
|
||||
|
||||
table.pt thead tr th,
|
||||
table.pt tbody tr td {
|
||||
padding: 10px 12px;
|
||||
text-align: left;
|
||||
border-top: 1px solid var(--border-color);
|
||||
vertical-align: middle;
|
||||
}
|
||||
|
||||
table.pt thead,
|
||||
table.pt thead th {
|
||||
background-color: var(--bg-color);
|
||||
/* Very light grey header background */
|
||||
color: var(--heading-color);
|
||||
font-weight: 600;
|
||||
border-top: none;
|
||||
border-bottom: 2px solid var(--border-color);
|
||||
/* Keep a stronger bottom border for header */
|
||||
text-align: right;
|
||||
}
|
||||
|
||||
table.pt tbody tr:first-child td {
|
||||
border-top: none;
|
||||
}
|
||||
|
||||
table.pt th,
|
||||
table.pt td {
|
||||
border-left: 1px solid var(--border-color);
|
||||
border-right: 1px solid var(--border-color);
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
}
|
||||
|
||||
/* Body row styles */
|
||||
table.pt tbody tr {
|
||||
background-color: var(--card-bg);
|
||||
}
|
||||
|
||||
/* Alternating row colors (need a variable or opacity) */
|
||||
/* Using a slight transparency for alternating rows to adapt to dark/light */
|
||||
table.pt tbody tr:nth-child(even) {
|
||||
background-color: rgba(128, 128, 128, 0.05);
|
||||
}
|
||||
|
||||
table.pt tbody tr:hover {
|
||||
background-color: rgba(128, 128, 128, 0.1);
|
||||
}
|
||||
|
||||
table.pt thead th:first-child {
|
||||
text-align: left;
|
||||
/* Align 'Parameter' header left */
|
||||
}
|
||||
|
||||
table.pt tbody tr td {
|
||||
text-align: right;
|
||||
font-family: "SFMono-Regular", Menlo, Monaco, Consolas, "Liberation Mono", "Courier New", monospace;
|
||||
}
|
||||
|
||||
table.pt tbody tr td.pt_param {
|
||||
font-weight: 500;
|
||||
white-space: nowrap;
|
||||
text-align: left;
|
||||
color: var(--primary-color);
|
||||
}
|
||||
|
||||
/* Update these to use variable or standard colors suitable for both modes */
|
||||
.pt_param {
|
||||
color: var(--primary-color);
|
||||
}
|
||||
|
||||
.pt_nominal {
|
||||
color: var(--primary-color);
|
||||
/* Was #1f1fb4 */
|
||||
}
|
||||
|
||||
.pt_bound {
|
||||
color: var(--primary-color);
|
||||
}
|
||||
|
||||
.pt_small_error {
|
||||
color: var(--color-success);
|
||||
}
|
||||
|
||||
.pt_medium_error {
|
||||
color: var(--color-warning);
|
||||
}
|
||||
|
||||
.pt_large_error {
|
||||
color: var(--color-error);
|
||||
}
|
||||
|
||||
.pt_on_boundary {
|
||||
color: var(--color-boundary);
|
||||
}
|
||||
</style>
|
||||
<div class="parameters">
|
||||
<table class="pt">
|
||||
<thead>
|
||||
<tr>
|
||||
{% for header in headers %}
|
||||
<th>{{ header }}</th>
|
||||
{% endfor %}
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{% for row in table_data %}
|
||||
<tr>
|
||||
<td class="pt_param" {% if row.is_frozen %}style="color: grey;" {% endif %}>{{ row.name }}</td> {# Parameter
|
||||
name #}
|
||||
{% if row.initial is defined %}
|
||||
<td class="pt_nominal" style="color: var(--text-muted);">{{ "%.4f" | format(row.initial) }}</td>
|
||||
{% endif %}
|
||||
{% if row.nominal is defined %}
|
||||
<td class="pt_nominal">{{ "%.4f" | format(row.nominal) }}</td>
|
||||
{% endif %}
|
||||
<td class="{{row.error_class}}">{{ "%.4f" | format(row.pred) }}</td>
|
||||
{% if row.is_frozen %}
|
||||
<td class="pt_bound">-</td>
|
||||
<td class="pt_bound">-</td>
|
||||
{% else %}
|
||||
<td class="pt_bound">{{ "%.4f" | format(row.lower_bound) }}</td>
|
||||
<td class="pt_bound">{{ "%.4f" | format(row.upper_bound) }}</td>
|
||||
{% endif %}
|
||||
{% if row.abs_err is defined %}
|
||||
<td class="{{row.error_class}}">{{ "%.4f" | format(row.abs_err) }}</td>
|
||||
{% endif %}
|
||||
{# Format relative error as percentage #}
|
||||
{% if row.rel_err is defined %}
|
||||
<td class="{{row.error_class}}">{{ "%.1f%%" | format(row.rel_err * 100) }}</td>
|
||||
{% endif %}
|
||||
</tr>
|
||||
{% endfor %}
|
||||
</tbody>
|
||||
</table>
|
||||
{% if overall_rmse is not none %}
|
||||
<p class="pt_mse">Overall Parameters RMSE: {{ "%.4f" | format(overall_rmse) }}</p>
|
||||
{% endif %}
|
||||
<p style="color: grey; font-size: 0.85em; margin-top: 0.5rem;">* parameter is frozen</p>
|
||||
</div>
|
||||
@@ -0,0 +1,8 @@
|
||||
<div class="plot-container">
|
||||
{{ plot_div | safe }}
|
||||
{% if caption %}
|
||||
<div class="plot-caption" style="margin-top: 0.5rem; color: var(--text-muted); font-size: 0.9em;">
|
||||
<em>{{ caption | safe }}</em>
|
||||
</div>
|
||||
{% endif %}
|
||||
</div>
|
||||
@@ -0,0 +1,13 @@
|
||||
<div class="row-section">
|
||||
{% if description %}
|
||||
<p class="section-description" style="color: var(--text-muted); margin-bottom: 1rem;">{{ description | safe }}</p>
|
||||
{% endif %}
|
||||
<div class="row-content"
|
||||
style="display: flex; flex-direction: row; flex-wrap: wrap; gap: 1rem; justify-content: space-around;">
|
||||
{% for child in child_sections %}
|
||||
<div class="row-item" style="flex: 1; min-width: 300px;">
|
||||
{{ child.content | safe }}
|
||||
</div>
|
||||
{% endfor %}
|
||||
</div>
|
||||
</div>
|
||||
@@ -0,0 +1,3 @@
|
||||
<div class="sensor-comparison-section">
|
||||
{{ plot_div | safe }}
|
||||
</div>
|
||||
@@ -0,0 +1,11 @@
|
||||
<div style="text-align: center; margin: 20px 0;">
|
||||
<video src="{{ video_filepath }}" width="{{ width }}" height="{{ height }}" {{ controls }} {{ autoplay }} {{ muted
|
||||
}} {{ loop }} preload="metadata" style="max-width: 100%; height: auto;">
|
||||
Your browser does not support the video tag.
|
||||
<a href="{{ video_filepath }}" target="_blank">Download the video here.</a>
|
||||
</video>
|
||||
{% if caption %}
|
||||
<p style="margin-top: 10px; font-size: 0.9em; color: var(--text-muted);" class="video-caption">{{ caption | safe }}
|
||||
</p>
|
||||
{% endif %}
|
||||
</div>
|
||||
@@ -0,0 +1,61 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
|
||||
"""Utility functions for report rendering."""
|
||||
|
||||
import math
|
||||
|
||||
from plotly import offline as plt_offline
|
||||
|
||||
|
||||
def plotly_script_tag() -> str:
|
||||
"""Returns an HTML <script> tag for Plotly.
|
||||
|
||||
currently installed in the environment.
|
||||
"""
|
||||
plotly_js_version = plt_offline.get_plotlyjs_version()
|
||||
return (
|
||||
"<script"
|
||||
f' src="https://cdn.plot.ly/plotly-{plotly_js_version}.min.js"></script>'
|
||||
)
|
||||
|
||||
|
||||
def get_text_color(bg_color_hex: str) -> str:
|
||||
"""Returns 'black' or 'white' for text over a background color.
|
||||
|
||||
luminance of the given background hex color.
|
||||
Useful for heatmaps and colored data tables.
|
||||
"""
|
||||
# Convert hex to RGB
|
||||
hex_color = bg_color_hex.lstrip("#")
|
||||
if len(hex_color) != 6:
|
||||
return "black" # Fallback
|
||||
|
||||
rgb = tuple(int(hex_color[i : i + 2], 16) for i in (0, 2, 4))
|
||||
r, g, b = [x / 255.0 for x in rgb]
|
||||
|
||||
# Calculate luminance (per WCAG guidelines)
|
||||
# https://www.w3.org/TR/WCAG20/#relativeluminancedef
|
||||
def lum_component(c):
|
||||
return c / 12.92 if c <= 0.03928 else math.pow((c + 0.055) / 1.055, 2.4)
|
||||
|
||||
luminance = (
|
||||
0.2126 * lum_component(r)
|
||||
+ 0.7152 * lum_component(g)
|
||||
+ 0.0722 * lum_component(b)
|
||||
)
|
||||
|
||||
# Return 'black' for light backgrounds, 'white' for dark backgrounds
|
||||
return "black" if luminance > 0.4 else "white"
|
||||
Reference in New Issue
Block a user