Coverage for transformer_lens/model_bridge/supported_architectures/phi.py: 100%
19 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Phi architecture adapter."""
3from typing import Any
5from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion
6from transformer_lens.conversion_utils.param_processing_conversion import (
7 ParamProcessingConversion,
8)
9from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
10from transformer_lens.model_bridge.generalized_components import (
11 EmbeddingBridge,
12 LinearBridge,
13 MLPBridge,
14 NormalizationBridge,
15 ParallelBlockBridge,
16 PositionEmbeddingsAttentionBridge,
17 RotaryEmbeddingBridge,
18 UnembeddingBridge,
19)
22class PhiArchitectureAdapter(ArchitectureAdapter):
23 """Architecture adapter for Phi models."""
25 _testing_eager = None
27 default_cfg = {"use_fast": False}
29 def __init__(self, cfg: Any) -> None:
30 """Initialize the Phi architecture adapter.
32 Args:
33 cfg: The configuration object.
34 """
35 super().__init__(cfg)
37 # Set config variables for weight processing
38 self.cfg.normalization_type = "LN"
39 self.cfg.positional_embedding_type = "rotary"
40 self.cfg.final_rms = False
41 self.cfg.gated_mlp = False
42 self.cfg.attn_only = False
43 self.cfg.parallel_attn_mlp = True
45 self.cfg.default_prepend_bos = False
47 self.weight_processing_conversions = {
48 "blocks.{i}.attn.q.weight": ParamProcessingConversion(
49 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads),
50 ),
51 "blocks.{i}.attn.k.weight": ParamProcessingConversion(
52 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads),
53 ),
54 "blocks.{i}.attn.v.weight": ParamProcessingConversion(
55 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads),
56 ),
57 "blocks.{i}.attn.q.bias": ParamProcessingConversion(
58 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=self.cfg.n_heads),
59 ),
60 "blocks.{i}.attn.k.bias": ParamProcessingConversion(
61 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=self.cfg.n_heads),
62 ),
63 "blocks.{i}.attn.v.bias": ParamProcessingConversion(
64 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=self.cfg.n_heads),
65 ),
66 "blocks.{i}.attn.o.weight": ParamProcessingConversion(
67 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=self.cfg.n_heads),
68 ),
69 }
71 # Set up component mapping
72 self.component_mapping = {
73 "embed": EmbeddingBridge(name="model.embed_tokens"),
74 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
75 "blocks": ParallelBlockBridge(
76 name="model.layers",
77 submodules={
78 "ln1": NormalizationBridge(
79 name="input_layernorm",
80 config=self.cfg,
81 use_native_layernorm_autograd=True,
82 ),
83 "attn": PositionEmbeddingsAttentionBridge(
84 name="self_attn",
85 config=self.cfg,
86 submodules={
87 "q": LinearBridge(name="q_proj"),
88 "k": LinearBridge(name="k_proj"),
89 "v": LinearBridge(name="v_proj"),
90 "o": LinearBridge(name="dense"),
91 },
92 requires_attention_mask=True,
93 requires_position_embeddings=True,
94 ),
95 "mlp": MLPBridge(
96 name="mlp",
97 submodules={
98 "in": LinearBridge(name="fc1"),
99 "out": LinearBridge(name="fc2"),
100 },
101 ),
102 },
103 ),
104 "ln_final": NormalizationBridge(
105 name="model.final_layernorm",
106 config=self.cfg,
107 use_native_layernorm_autograd=True,
108 ),
109 "unembed": UnembeddingBridge(name="lm_head"),
110 }