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

1"""Phi architecture adapter.""" 

2 

3from typing import Any 

4 

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) 

20 

21 

22class PhiArchitectureAdapter(ArchitectureAdapter): 

23 """Architecture adapter for Phi models.""" 

24 

25 _testing_eager = None 

26 

27 default_cfg = {"use_fast": False} 

28 

29 def __init__(self, cfg: Any) -> None: 

30 """Initialize the Phi architecture adapter. 

31 

32 Args: 

33 cfg: The configuration object. 

34 """ 

35 super().__init__(cfg) 

36 

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 

44 

45 self.cfg.default_prepend_bos = False 

46 

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 } 

70 

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 }