Skip to content

plot_gradiend_transition_cross_encoding_heatmap

Plot GRADIEND × directed-transition cross-encoding (pre-anchor aggregation).

Each row is one trained GRADIEND; each column is an input transition source->target from the shared cross-task test pool. This is the standard matrix before anchor sign alignment and anchor-class aggregation.

Source code in gradiend/visualizer/heatmaps/encoding.py
def plot_gradiend_transition_cross_encoding_heatmap(
    trainers: Dict[str, object],
    *,
    trainer_order: Optional[Sequence[str]] = None,
    transition_order: Optional[Sequence[str]] = None,
    encoder_summary: Optional[Dict[str, Any]] = None,
    split: str = "test",
    max_size: Optional[int] = None,
    aggregate: str = "mean",
    seed_selection: Optional[str] = None,
    seed_aggregate: str = "mean",
    dispersion: Optional[str] = None,
    order: Union[str, List[str]] = "input",
    cluster: bool = False,
    pretty_groups: Optional[Dict[str, List[str]]] = None,
    highlight_non_convergence: bool = True,
    **plot_kwargs: Any,
) -> Dict[str, Any]:
    """
    Plot GRADIEND × directed-transition cross-encoding (pre-anchor aggregation).

    Each row is one trained GRADIEND; each column is an input transition
    ``source->target`` from the shared cross-task test pool. This is the standard
    matrix before anchor sign alignment and anchor-class aggregation.
    """
    if aggregate not in {"mean", "count"}:
        raise ValueError("aggregate must be 'mean' or 'count'")
    seed_selection = resolve_seed_selection_for_trainers(trainers, seed_selection)
    if dispersion is None:
        dispersion = resolve_dispersion_for_trainers(trainers, None)
    comparison_data = compute_gradiend_transition_cross_encoding_matrix(
        trainers,
        trainer_order=trainer_order,
        transition_order=transition_order,
        encoder_summary=encoder_summary,
        split=split,
        max_size=max_size,
        seed_selection=seed_selection,
        seed_aggregate=seed_aggregate,
        dispersion=dispersion,
    )
    if aggregate == "count":
        comparison_data = dict(comparison_data)
        comparison_data["matrix"] = comparison_data["n_matrix"]
        comparison_data["measure"] = "gradiend_transition_cross_encoding_count"
    resolved_order = (
        list(order)
        if isinstance(order, list)
        else list(comparison_data["model_ids"])
        if order == "input"
        else order
    )
    plot_kwargs = dict(plot_kwargs)
    plot_kwargs.setdefault("cbar_label", CROSS_ENCODING_CBAR_LABEL)
    return plot_comparison_heatmap(
        comparison_data,
        order=resolved_order,
        cluster=cluster,
        pretty_groups=pretty_groups,
        models=trainers,
        highlight_non_convergence=highlight_non_convergence,
        **filter_comparison_heatmap_plot_kwargs(plot_kwargs),
    )