Skip to content

compute_similarity_matrix

Compute a pairwise similarity matrix for trained GRADIEND models.

Parameters:

Name Type Description Default
models Dict[str, object]

Mapping from display/model id to a trained model. For "topk_overlap", models must provide get_topk_weights. For vector and mass measures, models must provide a .gradiend object with built encoder/decoder weights. Values may also be non-empty lists or tuples of seed models. At least two model ids are required.

required
measure str

Similarity measure. Supported values are "cosine", "cosine_signed", "spearman", "spearman_signed", "topk_overlap", and "mass_overlap". Unsigned cosine and spearman take absolute values.

'cosine'
part Optional[str]

GRADIEND part to compare. Supported values for vector/mass measures are "encoder-weight", "decoder-weight", "decoder-bias", and "decoder-sum". If omitted, "topk_overlap" uses "decoder-weight" and all other measures use "encoder-weight".

None
topk Optional[Union[int, float]]

Optional coordinate selection. Required for "topk_overlap" and "mass_overlap". None uses all coordinates where allowed, an integer selects that many top coordinates, and a float in (0, 1] selects that fraction using each model's own get_topk_weights implementation.

None
value str

Output value for "topk_overlap". "intersection_frac" divides the intersection by the smaller selected set size. "intersection" returns the raw intersection count. Ignored by other measures.

'intersection_frac'
seed_aggregate str

Aggregation used when a model id maps to a list/tuple of seed models. Supported values are "mean", "median", "min", and "max".

'mean'
dispersion str

Optional dispersion metadata for list/tuple seed cells. Supported values are "none", "std", "range", and "minmax". "range" is incompatible with seed_aggregate of "min" or "max".

'none'
seed_pairing_mode str

How seed groups are paired. "all_pairs" compares every seed in the left group with every seed in the right group. "matched" pairs models by position in each selected-seed group and requires equal group sizes. Original seed ids need not match.

'matched'

Returns:

Type Description
Dict[str, Any]

For one model per id, a payload with measure, model_ids,

Dict[str, Any]

matrix, part, topk, and measure-specific metadata such as

Dict[str, Any]

value, resolved_topk, per_model, block_dim,

Dict[str, Any]

full_zero_filled, sparse_exact, or fallback_dense. For

Dict[str, Any]

list/tuple seed groups, every cell aggregates the selected seed-pair

Dict[str, Any]

scores, and the payload also contains

Dict[str, Any]

seed_aggregate, dispersion, n_matrix, cell_stats,

Dict[str, Any]

multi_seed=True, plus global_n or global_n_range when

Dict[str, Any]

available.

Raises:

Type Description
TypeError

If models or topk has an invalid type, or vector/mass measures receive models without a .gradiend attribute.

ValueError

If fewer than two model ids are passed, a model group is empty, measure/part/value is unsupported, required topk is missing, topk is out of range, or seed aggregation and dispersion settings are incompatible.

Source code in gradiend/comparison/similarity.py
def compute_similarity_matrix(
    models: Dict[str, object],
    *,
    measure: str = "cosine",
    part: Optional[str] = None,
    topk: Optional[Union[int, float]] = None,
    value: str = "intersection_frac",
    seed_aggregate: str = "mean",
    dispersion: str = "none",
    seed_pairing_mode: str = "matched",
) -> Dict[str, Any]:
    """Compute a pairwise similarity matrix for trained GRADIEND models.

    Args:
        models: Mapping from display/model id to a trained model. For
            ``"topk_overlap"``, models must provide ``get_topk_weights``. For
            vector and mass measures, models must provide a ``.gradiend`` object
            with built encoder/decoder weights. Values may also be non-empty
            lists or tuples of seed models. At least two model ids are required.
        measure: Similarity measure. Supported values are ``"cosine"``,
            ``"cosine_signed"``, ``"spearman"``, ``"spearman_signed"``,
            ``"topk_overlap"``, and ``"mass_overlap"``. Unsigned ``cosine`` and
            ``spearman`` take absolute values.
        part: GRADIEND part to compare. Supported values for vector/mass
            measures are ``"encoder-weight"``, ``"decoder-weight"``,
            ``"decoder-bias"``, and ``"decoder-sum"``. If omitted,
            ``"topk_overlap"`` uses ``"decoder-weight"`` and all other measures
            use ``"encoder-weight"``.
        topk: Optional coordinate selection. Required for ``"topk_overlap"``
            and ``"mass_overlap"``. ``None`` uses all coordinates where allowed,
            an integer selects that many top coordinates, and a float in
            ``(0, 1]`` selects that fraction using each model's own
            ``get_topk_weights`` implementation.
        value: Output value for ``"topk_overlap"``. ``"intersection_frac"``
            divides the intersection by the smaller selected set size.
            ``"intersection"`` returns the raw intersection count. Ignored by
            other measures.
        seed_aggregate: Aggregation used when a model id maps to a list/tuple of
            seed models. Supported values are ``"mean"``, ``"median"``,
            ``"min"``, and ``"max"``.
        dispersion: Optional dispersion metadata for list/tuple seed cells.
            Supported values are ``"none"``, ``"std"``, ``"range"``, and
            ``"minmax"``. ``"range"`` is incompatible with ``seed_aggregate`` of
            ``"min"`` or ``"max"``.
        seed_pairing_mode: How seed groups are paired. ``"all_pairs"`` compares
            every seed in the left group with every seed in the right group.
            ``"matched"`` pairs models by position in each selected-seed group
            and requires equal group sizes. Original seed ids need not match.

    Returns:
        For one model per id, a payload with ``measure``, ``model_ids``,
        ``matrix``, ``part``, ``topk``, and measure-specific metadata such as
        ``value``, ``resolved_topk``, ``per_model``, ``block_dim``,
        ``full_zero_filled``, ``sparse_exact``, or ``fallback_dense``. For
        list/tuple seed groups, every cell aggregates the selected seed-pair
        scores, and the payload also contains
        ``seed_aggregate``, ``dispersion``, ``n_matrix``, ``cell_stats``,
        ``multi_seed=True``, plus ``global_n`` or ``global_n_range`` when
        available.

    Raises:
        TypeError: If ``models`` or ``topk`` has an invalid type, or vector/mass
            measures receive models without a ``.gradiend`` attribute.
        ValueError: If fewer than two model ids are passed, a model group is
            empty, ``measure``/``part``/``value`` is unsupported, required
            ``topk`` is missing, ``topk`` is out of range, or seed aggregation
            and dispersion settings are incompatible.
    """
    model_groups = _normalize_model_groups(models)
    if seed_pairing_mode not in {"all_pairs", "matched"}:
        raise ValueError("seed_pairing_mode must be 'all_pairs' or 'matched'")
    _validate_topk_optional(topk)
    measure = (measure or "cosine").lower()
    resolved_part = (part or ("decoder-weight" if measure == "topk_overlap" else "encoder-weight")).lower()
    _validate_aggregate_dispersion_combo(seed_aggregate, dispersion)
    if all(len(group) == 1 for group in model_groups.values()):
        flat_models = {mid: group[0] for mid, group in model_groups.items()}
        _validate_models(flat_models)
        if measure == "cosine":
            return _compute_cosine_like_matrix(flat_models, part=resolved_part, topk=topk, take_abs=True, measure_name="cosine")
        if measure == "cosine_signed":
            return _compute_cosine_like_matrix(flat_models, part=resolved_part, topk=topk, take_abs=False, measure_name="cosine_signed")
        if measure == "spearman":
            return _compute_spearman_matrix(flat_models, part=resolved_part, topk=topk, take_abs=True, measure_name="spearman")
        if measure == "spearman_signed":
            return _compute_spearman_matrix(flat_models, part=resolved_part, topk=topk, take_abs=False, measure_name="spearman_signed")
        if measure == "topk_overlap":
            if topk is None:
                raise ValueError("topk must be provided for measure='topk_overlap'")
            return _compute_topk_overlap_similarity_matrix(flat_models, part=resolved_part, topk=topk, value=value)
        if measure == "mass_overlap":
            if topk is None:
                raise ValueError("topk must be provided for measure='mass_overlap'")
            return _compute_mass_overlap_matrix(flat_models, part=resolved_part, topk=topk)
        raise ValueError(
            "measure must be one of 'cosine', 'cosine_signed', 'spearman', 'spearman_signed', 'topk_overlap', or 'mass_overlap'"
        )

    model_ids = list(model_groups.keys())
    matrix = [[0.0] * len(model_ids) for _ in range(len(model_ids))]
    n_matrix = [[0] * len(model_ids) for _ in range(len(model_ids))]
    cell_stats: List[List[Dict[str, Any]]] = []
    all_n: List[int] = []
    for i, mi in enumerate(model_ids):
        stats_row: List[Dict[str, Any]] = []
        for j, mj in enumerate(model_ids):
            if seed_pairing_mode == "matched":
                left_group = model_groups[mi]
                right_group = model_groups[mj]
                if len(left_group) != len(right_group):
                    raise ValueError(
                        "seed_pairing_mode='matched' requires equal selected-seed group sizes; "
                        f"{mi!r} has {len(left_group)} and {mj!r} has {len(right_group)}"
                    )
                model_pairs = list(zip(left_group, right_group))
            else:
                model_pairs = [
                    (left_model, right_model)
                    for left_model in model_groups[mi]
                    for right_model in model_groups[mj]
                ]
            scores = [
                _score_similarity_pair(left_model, right_model, measure=measure, part=resolved_part, topk=topk, value=value)
                for left_model, right_model in model_pairs
            ]
            stats = _aggregate_seed_scores(scores, seed_aggregate=seed_aggregate, dispersion=dispersion)
            matrix[i][j] = float(stats["aggregate"])
            n_matrix[i][j] = int(stats["n"])
            stats_row.append(stats)
            all_n.append(int(stats["n"]))
        cell_stats.append(stats_row)
    payload: Dict[str, Any] = {
        "measure": measure,
        "model_ids": model_ids,
        "matrix": matrix,
        "part": resolved_part,
        "topk": topk,
        "value": value,
        "seed_aggregate": seed_aggregate,
        "dispersion": dispersion,
        "seed_pairing_mode": seed_pairing_mode,
        "n_matrix": n_matrix,
        "cell_stats": cell_stats,
        "multi_seed": True,
    }
    if all_n:
        if min(all_n) == max(all_n):
            payload["global_n"] = int(all_n[0])
        else:
            payload["global_n_range"] = [int(min(all_n)), int(max(all_n))]
    return payload