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),
)