transformer_lens.utilities.matrix module

matrix.

This module contains utility functions related to the transformer lens implementation of factored matrices.

transformer_lens.utilities.matrix.composition_scores(left: FactoredMatrix, right: FactoredMatrix, broadcast_dims=True) → Float[Tensor, '*leading_dims'] | Float[Tensor, '*leading_dims_left_and_right']

Composition scores between two factored matrices.

Returns ||left @ right||_F / (||left||_F * ||right||_F), computed from the factored forms so the full products are never materialized. With broadcast_dims, left and right leading dims are broadcast against each other (left dims first), scoring every left/right pair. See TransformerBridge.all_composition_scores.

transformer_lens.utilities.matrix.get_matrix_corner(matrix: FactoredMatrix, n=3)