"""Visualization and reporting stage."""
from __future__ import annotations
from pathlib import Path
from typing import Any
import numpy as np
from .._logging import get_logger
from .base import BaseVisualizationStage
log = get_logger(__name__)
[docs]
class DefaultVisualizationStage(BaseVisualizationStage):
"""Generate plots and summary reports."""
[docs]
def run(
self,
cfg: Any,
data_out: dict,
classify_out: dict,
hardware_out: dict,
) -> Path | None:
"""Write stage plots and a summary report under ``visualization.output_dir``.
Args:
cfg: Pipeline config (``visualization`` block selects stages/format).
data_out: Data-stage output (IQ clusters / averaged traces).
classify_out: Classify-stage output (filters, metrics, sweeps).
hardware_out: Hardware-stage output (resource / DSE plots).
Returns:
Path to the summary report, or ``None`` if no classifier results.
"""
log.info("=== Stage 4: Visualization ===")
vcfg = cfg.visualization
out_dir = Path(vcfg.output_dir)
fmt = vcfg.format if vcfg.format in ("pdf", "png") else "pdf"
dpi = vcfg.dpi
stages_vis = set(vcfg.stages)
import matplotlib
matplotlib.use("Agg")
if "data" in stages_vis:
_viz_data(data_out, out_dir, fmt, dpi, cfg)
if "classify" in stages_vis:
_viz_classify(classify_out, out_dir, fmt, dpi)
if "hardware" in stages_vis:
_viz_hardware(hardware_out, out_dir, fmt, dpi)
return _viz_report(cfg, classify_out, hardware_out, out_dir, vcfg.format, dpi)
[docs]
def run_visualization_stage(
cfg: Any,
data_out: dict,
classify_out: dict,
hardware_out: dict,
) -> Path | None:
"""Run the default visualization stage.
Thin wrapper around :class:`DefaultVisualizationStage`.
Args:
cfg: Pipeline config (``visualization`` block).
data_out: Data-stage output.
classify_out: Classify-stage output.
hardware_out: Hardware-stage output.
Returns:
Path to the summary report, or ``None`` if skipped.
"""
return DefaultVisualizationStage().run(cfg, data_out, classify_out, hardware_out)
def _viz_data(
data_out: dict,
out_dir: Path,
fmt: str,
dpi: int,
cfg: Any,
) -> None:
from arcade.data.transitions.result import TransitionResult, average_traces
from ..viz.iq_scatter import plot_iq_scatter
from ..viz.traces import plot_averaged_traces
vcfg = cfg.visualization
data_dir = out_dir / "data"
tr = data_out.get("transition_result") or data_out.get("relaxation_result")
if tr is not None and vcfg.show_iq_clusters:
plot_iq_scatter(tr, output_dir=data_dir, fmt=fmt, dpi=dpi)
if tr is not None and vcfg.show_averaged_traces:
plot_averaged_traces(tr, output_dir=data_dir, fmt=fmt, dpi=dpi)
if tr is None and vcfg.show_averaged_traces and data_out.get("splits"):
sp = data_out["splits"]
labels = sp.train_labels
train_tr = sp.train_traces
dcfg = cfg.data
av = average_traces(
train_tr,
labels,
dcfg.num_qubits,
dcfg.num_levels,
)
td: dict[int, dict[str, Any]] = {}
for q in range(dcfg.num_qubits):
qd: dict[str, Any] = {}
for s in range(dcfg.num_levels):
traj = av[q][s]
qd[f"traces_{s}"] = traj[np.newaxis, ...]
td[q] = qd
fake = TransitionResult(
num_qubits=dcfg.num_qubits,
num_levels=dcfg.num_levels,
per_qubit=td,
)
plot_averaged_traces(fake, output_dir=data_dir, fmt=fmt, dpi=dpi)
def _viz_classify(classify_out: dict, out_dir: Path, fmt: str, dpi: int) -> None:
from ..viz.filters import plot_filter_envelopes
from ..viz.training import plot_confusion_matrix, plot_training_curves
fs = classify_out.get("filter_set")
if fs is not None:
plot_filter_envelopes(fs, output_dir=out_dir / "filters", fmt=fmt, dpi=dpi)
clf_results = classify_out.get("classifiers", {})
sweep_data: dict[str, Any] = {}
for clf_name, data in clf_results.items():
clf_dir = out_dir / clf_name
history = data.get("history", {})
if history and ("train_loss" in history or "val_accuracy" in history):
plot_training_curves(
history,
output_dir=clf_dir,
classifier_name=clf_name,
fmt=fmt,
dpi=dpi,
)
metrics = data.get("metrics")
if metrics is not None and metrics.confusion_matrix is not None:
plot_confusion_matrix(
metrics.confusion_matrix,
output_dir=clf_dir,
classifier_name=clf_name,
fmt=fmt,
dpi=dpi,
)
if data.get("sweep") is not None:
sweep_data[clf_name] = data["sweep"]
if sweep_data:
from ..viz.sweep import plot_speed_fidelity
plot_speed_fidelity(
sweep_data,
output_dir=out_dir / "sweep",
fmt=fmt,
dpi=dpi,
)
opt_pareto = classify_out.get("optimization", {}).get("pareto", {})
if opt_pareto:
from ..viz.pareto import plot_compression_pareto
for clf_name, points in opt_pareto.items():
if points:
plot_compression_pareto(
points,
output_dir=out_dir / clf_name,
title=f"{clf_name} optimization Pareto",
fmt=fmt,
dpi=dpi,
)
if clf_results and any(r.get("timing") for r in clf_results.values()):
from ..viz.latency import plot_accuracy_vs_latency, plot_latency_breakdown
lat_dir = out_dir / "latency"
try:
plot_accuracy_vs_latency(clf_results, output_dir=lat_dir, fmt=fmt, dpi=dpi)
plot_latency_breakdown(clf_results, output_dir=lat_dir, fmt=fmt, dpi=dpi)
except ValueError:
pass
def _viz_hardware(hardware_out: dict, out_dir: Path, fmt: str, dpi: int) -> None:
from ..viz.hardware import plot_hardware_comparison
reports = hardware_out.get("reports", {})
if reports:
base_reports = {k: v for k, v in reports.items() if not k.endswith("_optimized")}
if base_reports:
plot_hardware_comparison(
base_reports,
output_dir=out_dir / "hardware",
fmt=fmt,
dpi=dpi,
)
dse = hardware_out.get("dse", {})
if dse:
from ..viz.dse import plot_dse_pareto
for clf_name, result in dse.items():
if getattr(result, "pareto_optimal", None):
plot_dse_pareto(
result,
output_dir=out_dir / "hardware" / clf_name,
fmt=fmt,
dpi=dpi,
)
def _viz_report(
cfg: Any,
classify_out: dict,
hardware_out: dict,
out_dir: Path,
fmt: str,
dpi: int,
) -> Path | None:
from ..viz.report import export_results_json, generate_summary_report
clf_results = classify_out.get("classifiers", {})
hw_reports = hardware_out.get("reports", {})
sections = getattr(cfg.visualization, "summary_sections", None)
if not clf_results:
log.info("No classifier results → skipping summary report")
return None
report_fmt = "html" if fmt == "html" else "pdf"
path = generate_summary_report(
clf_results,
hw_reports,
output_dir=out_dir,
fmt=report_fmt,
dpi=dpi,
sections=list(sections) if sections else None,
qubit_bit_order=getattr(cfg.data, "qubit_bit_order", "lsb0"),
)
export_results_json(
clf_results,
hw_reports,
output_dir=out_dir,
optimization=classify_out.get("optimization"),
)
log.info("Summary report: %s", path)
return path