Skip to content

plot_gradiend_feature_cross_encoding_heatmap

Plot a dense GRADIEND × feature-class cross-encoding matrix.

Each row is one trained GRADIEND; each column is one feature class. Cell (i, j) is the mean encoded value when GRADIEND i encodes test-split snippets for class j (shared eval pool merged across trainers).

Parameters:

Name Type Description Default
trainers Dict[str, object]

Mapping of ids to trainers.

required
feature_classes Sequence[str]

Feature classes used as columns.

required
trainer_order Optional[Sequence[str]]

Optional explicit trainer row order.

None
split str

Encoder split to evaluate.

'test'
max_size Optional[int]

Optional evaluation row cap.

None
aggregate str

"mean" for encoded values or "count" for counts.

'mean'
order Union[str, List[str]]

Heatmap row order.

'input'
cluster bool

Whether to cluster rows/columns.

False
pretty_groups Optional[Dict[str, List[str]]]

Optional groups shown as brackets.

None
highlight_non_convergence bool

Whether labels mark non-converged runs.

True
**plot_kwargs Any

Additional options forwarded to plot_comparison_heatmap.

{}
Source code in gradiend/visualizer/heatmaps/encoding.py
def plot_gradiend_feature_cross_encoding_heatmap(
    trainers: Dict[str, object],
    feature_classes: Sequence[str],
    *,
    trainer_order: Optional[Sequence[str]] = None,
    split: str = "test",
    max_size: Optional[int] = None,
    aggregate: str = "mean",
    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 a dense GRADIEND × feature-class cross-encoding matrix.

    Each row is one trained GRADIEND; each column is one feature class. Cell
    *(i, j)* is the mean encoded value when GRADIEND *i* encodes test-split
    snippets for class *j* (shared eval pool merged across trainers).

    Args:
        trainers: Mapping of ids to trainers.
        feature_classes: Feature classes used as columns.
        trainer_order: Optional explicit trainer row order.
        split: Encoder split to evaluate.
        max_size: Optional evaluation row cap.
        aggregate: ``"mean"`` for encoded values or ``"count"`` for counts.
        order: Heatmap row order.
        cluster: Whether to cluster rows/columns.
        pretty_groups: Optional groups shown as brackets.
        highlight_non_convergence: Whether labels mark non-converged runs.
        **plot_kwargs: Additional options forwarded to ``plot_comparison_heatmap``.
    """
    if aggregate not in {"mean", "count"}:
        raise ValueError("aggregate must be 'mean' or 'count'")
    comparison_data = compute_gradiend_feature_cross_encoding_matrix(
        trainers,
        feature_classes,
        trainer_order=trainer_order,
        split=split,
        max_size=max_size,
    )
    if aggregate == "count":
        comparison_data = dict(comparison_data)
        comparison_data["matrix"] = comparison_data["n_matrix"]
        comparison_data["measure"] = "gradiend_feature_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),
    )