transformer_lens.conversion_utils package

Subpackages

Submodules

Module contents

Model bridge conversion utilities.

This module contains utilities for converting between different model architectures.

class transformer_lens.conversion_utils.TensorConversionSet(fields: dict[str, Any])

Bases: BaseTensorConversion

get_component(model: Any, name: str) → Any

Get a component from the model using the field mapping.

Parameters:
  • model – The model to get the component from.

  • name – The name of the component to get.

Returns:

The requested component.

get_conversion_action(field: str) → BaseTensorConversion
handle_conversion(input_value: Any, *full_context: Any) → dict[str, Any]
process_conversion(input_value: Any, remote_field: str, conversion: BaseTensorConversion, *full_context: Any) → Any
process_conversion_action(input_value: Any, conversion_details: Any, *full_context: Any) → Any
transformer_lens.conversion_utils.convert_tl_checkpoint(state_dict: dict[str, Tensor], cfg: TransformerBridgeConfig) → dict[str, Tensor]

Convert a legacy TL-property-format state dict to the key/tensor format TransformerBridge.boot_native(cfg).load_state_dict accepts.

Parameters:
  • state_dict – A state dict in the old HookedTransformer convention (e.g. from HookedTransformer.state_dict()), with keys like "blocks.0.attn.W_Q" and per-head tensor shapes.

  • cfg – The config the checkpoint was trained/saved under. Used both to reshape per-head attention weights and to validate that the checkpoint’s per-head shapes actually match this cfg — a mismatched cfg would otherwise silently mis-group heads without ever tripping a shape error, since d_model == n_heads * d_head holds for any wrong factoring too.

Returns:

A state dict with modern bridge keys (e.g. "blocks.0.attn.q.weight") and flat nn.Linear-oriented tensor shapes, ready for bridge.load_state_dict(converted, strict=True).