Source code for arcade.stages.visualization

"""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