1010from functools import partial
1111
1212import torch
13+ import torch .nn as nn
1314from torch import Tensor
1415
1516from fairseq2 .models .llama ._config import LLaMAConfig , LLaMARopeScalingConfig
1617from fairseq2 .models .transformer import (
1718 TransformerEmbeddingFrontend ,
1819 TransformerFrontend ,
19- init_final_projection ,
2020)
2121from fairseq2 .models .transformer_decoder import TransformerDecoderModel
2222from fairseq2 .nn import (
@@ -77,8 +77,19 @@ def create_model(self) -> TransformerDecoderModel:
7777 def create_embedding (self ) -> Embedding :
7878 config = self ._config
7979
80+ init_std = config .init_std
81+
82+ def init_embed (embed : StandardEmbedding ) -> None :
83+ embed_dim = embed .weight .shape [1 ]
84+
85+ std = init_std or (embed_dim ** - 0.5 )
86+
87+ _init_truncated_normal (embed .weight , bias = None , std = std )
88+
8089 return StandardEmbedding (
81- num_embeddings = config .vocab_info .size , embedding_dim = config .model_dim
90+ num_embeddings = config .vocab_info .size ,
91+ embedding_dim = config .model_dim ,
92+ init_fn = init_embed ,
8293 )
8394
8495 def create_decoder_frontend (self , embed : Embedding ) -> TransformerFrontend :
@@ -95,8 +106,8 @@ def create_decoder(self) -> TransformerDecoder:
95106
96107 layers = []
97108
98- for _ in range (config .num_layers ):
99- layer = self .create_decoder_layer (pos_encoder )
109+ for idx in range (config .num_layers ):
110+ layer = self .create_decoder_layer (idx , pos_encoder )
100111
101112 layers .append (layer )
102113
@@ -125,11 +136,11 @@ def create_position_encoder(self) -> PositionEncoder:
125136 )
126137
127138 def create_decoder_layer (
128- self , pos_encoder : PositionEncoder
139+ self , layer_idx : int , pos_encoder : PositionEncoder
129140 ) -> TransformerDecoderLayer :
130- self_attn = self .create_attention (pos_encoder )
141+ self_attn = self .create_attention (layer_idx , pos_encoder )
131142
132- ffn = self .create_ffn ()
143+ ffn = self .create_ffn (layer_idx )
133144
134145 return StandardTransformerDecoderLayer (
135146 self_attn ,
@@ -139,23 +150,49 @@ def create_decoder_layer(
139150 layer_norm_factory = self .create_layer_norm ,
140151 )
141152
142- def create_attention (self , pos_encoder : PositionEncoder ) -> MultiheadAttention :
153+ def create_attention (
154+ self , layer_idx : int , pos_encoder : PositionEncoder
155+ ) -> MultiheadAttention :
143156 config = self ._config
144157
158+ init_std = config .init_std
159+
160+ std_scale_factor = self ._get_std_scale_factor (layer_idx )
161+
162+ def init_projection (proj : Linear ) -> None :
163+ input_dim = proj .weight .shape [1 ]
164+
165+ std = init_std or (input_dim ** - 0.5 )
166+
167+ _init_truncated_normal (proj .weight , proj .bias , std = std / std_scale_factor )
168+
145169 sdpa = create_default_sdpa (attn_dropout_p = config .dropout_p )
146170
147171 return StandardMultiheadAttention (
148172 config .model_dim ,
149173 config .num_attn_heads ,
150174 num_key_value_heads = config .num_key_value_heads ,
175+ qkv_proj_init_fn = init_projection ,
151176 sdpa = sdpa ,
152177 pos_encoder = pos_encoder ,
178+ output_proj_init_fn = init_projection ,
153179 bias = False ,
154180 )
155181
156- def create_ffn (self ) -> FeedForwardNetwork :
182+ def create_ffn (self , layer_idx : int ) -> FeedForwardNetwork :
157183 config = self ._config
158184
185+ init_std = config .init_std
186+
187+ std_scale_factor = self ._get_std_scale_factor (layer_idx )
188+
189+ def init_projection (proj : Linear ) -> None :
190+ input_dim = proj .weight .shape [1 ]
191+
192+ std = init_std or (input_dim ** - 0.5 )
193+
194+ _init_truncated_normal (proj .weight , proj .bias , std = std / std_scale_factor )
195+
159196 ffn_inner_dim = int (config .ffn_inner_dim * config .ffn_inner_dim_multiplier )
160197
161198 return GLUFeedForwardNetwork (
@@ -165,8 +202,26 @@ def create_ffn(self) -> FeedForwardNetwork:
165202 inner_dim_scale = config .ffn_inner_dim_scale ,
166203 inner_dim_to_multiple = config .ffn_inner_dim_to_multiple ,
167204 inner_dropout_p = config .dropout_p ,
205+ proj_init_fn = init_projection ,
168206 )
169207
208+ def _get_std_scale_factor (self , layer_idx : int ) -> float :
209+ config = self ._config
210+
211+ match config .init_std_scale :
212+ case "layer" :
213+ n = layer_idx
214+ case "stack" :
215+ n = config .num_layers
216+ case "none" :
217+ return 1.0
218+ case _:
219+ raise ValueError (
220+ f"`config.init_std_scale` must be 'none', 'layer', or 'stack', but is '{ config .init_std_scale } ' instead."
221+ )
222+
223+ return (2 * (n + 1 )) ** 0.5 # type: ignore[no-any-return]
224+
170225 def create_final_proj (self , embed : Embedding ) -> Projection :
171226 config = self ._config
172227
@@ -178,11 +233,20 @@ def create_final_proj(self, embed: Embedding) -> Projection:
178233
179234 return TiedProjection (embed .weight , bias = None )
180235
236+ init_std = config .init_std
237+
238+ def init_projection (proj : Linear ) -> None :
239+ input_dim = proj .weight .shape [1 ]
240+
241+ std = init_std or (input_dim ** - 0.5 )
242+
243+ _init_truncated_normal (proj .weight , proj .bias , std = std )
244+
181245 return Linear (
182246 config .model_dim ,
183247 config .vocab_info .size ,
184248 bias = False ,
185- init_fn = init_final_projection ,
249+ init_fn = init_projection ,
186250 )
187251
188252 @staticmethod
@@ -192,6 +256,15 @@ def create_layer_norm(
192256 return RMSNorm (model_dim , bias = False , device = device , dtype = dtype )
193257
194258
259+ def _init_truncated_normal (
260+ weight : Tensor , bias : Tensor | None , * , std : float = 1.0
261+ ) -> None :
262+ nn .init .trunc_normal_ (weight , mean = 0.0 , std = std , a = - 3 * std , b = 3 * std )
263+
264+ if bias is not None :
265+ nn .init .zeros_ (bias )
266+
267+
195268def init_llama_scaled_freqs (
196269 pos_encoder : RotaryEncoder , rope_scaling : LLaMARopeScalingConfig
197270) -> Tensor :
0 commit comments