Skip to content

Commit 8ae2585

Browse files
committed
Add LLaMA weight initialization
1 parent 6ce120c commit 8ae2585

8 files changed

Lines changed: 144 additions & 35 deletions

File tree

src/fairseq2/models/jepa/_factory.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
66

77
from __future__ import annotations
88

9-
import math
109
from typing import cast
1110

1211
import torch
@@ -183,12 +182,12 @@ def create_encoder(self) -> TransformerEncoder:
183182
layer_norm_factory=self.create_layer_norm,
184183
)
185184

186-
def create_encoder_layer(self, idx: int) -> TransformerEncoderLayer:
185+
def create_encoder_layer(self, layer_idx: int) -> TransformerEncoderLayer:
187186
config = self._config
188187

189-
self_attn = self.create_attention(idx)
188+
self_attn = self.create_attention(layer_idx)
190189

191-
ffn = self.create_ffn(idx)
190+
ffn = self.create_ffn(layer_idx)
192191

193192
drop_path = DropPathResidualConnect(drop_p=config.droppath_p)
194193

@@ -226,7 +225,7 @@ def init_projection(proj: Linear) -> None:
226225
_init_truncated_normal(proj.weight, proj.bias, std=init_std)
227226

228227
with torch.no_grad():
229-
proj.weight.div_(math.sqrt(2.0 * (layer_idx + 1)))
228+
proj.weight.div_((2.0 * (layer_idx + 1)) ** 0.5)
230229

231230
return Linear(
232231
config.model_dim, config.model_dim, bias=True, init_fn=init_projection
@@ -241,7 +240,7 @@ def init_projection(proj: Linear) -> None:
241240
_init_truncated_normal(proj.weight, proj.bias, std=init_std)
242241

243242
with torch.no_grad():
244-
proj.weight.div_(math.sqrt(2.0 * (layer_idx + 1)))
243+
proj.weight.div_((2.0 * (layer_idx + 1)) ** 0.5)
245244

246245
inner_dim = int(config.model_dim * config.ffn_inner_dim_ratio)
247246

@@ -282,7 +281,7 @@ def init_layer_norm(m: LayerNorm) -> None:
282281
def _init_truncated_normal(
283282
weight: Tensor, bias: Tensor | None, *, std: float = 1.0
284283
) -> None:
285-
nn.init.trunc_normal_(weight, std=std)
284+
nn.init.trunc_normal_(weight, mean=0.0, std=std)
286285

287286
if bias is not None:
288287
nn.init.zeros_(bias)

src/fairseq2/models/llama/_config.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from __future__ import annotations
88

99
from dataclasses import dataclass, field
10-
from typing import Final
10+
from typing import Final, Literal
1111

1212
from fairseq2.context import RuntimeContext
1313
from fairseq2.data import VocabularyInfo
@@ -81,6 +81,20 @@ class LLaMAConfig:
8181
encoder, aiming to increase the context length.
8282
"""
8383

84+
init_std: float | None = None
85+
"""
86+
If not ``None``, the standard deviation to initialize input embeddings and
87+
projection weights; otherwise, ``model_dim ** -0.5`` will be used instead.
88+
"""
89+
90+
init_std_scale: Literal["none", "layer", "stack"] = "layer"
91+
"""
92+
The method to use to scale ``init_std`` per layer. If 'none', no scaling
93+
will be applied. If 'layer', ``init_std`` will be scaled by the depth of
94+
the layer. If 'stack', ``init_std`` will be scaled by the total depth of
95+
the decoder.
96+
"""
97+
8498
dropout_p: float = 0.1
8599
"""The dropout probability on outputs of Transformer layers."""
86100

src/fairseq2/models/llama/_factory.py

Lines changed: 83 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,13 @@
1010
from functools import partial
1111

1212
import torch
13+
import torch.nn as nn
1314
from torch import Tensor
1415

1516
from fairseq2.models.llama._config import LLaMAConfig, LLaMARopeScalingConfig
1617
from fairseq2.models.transformer import (
1718
TransformerEmbeddingFrontend,
1819
TransformerFrontend,
19-
init_final_projection,
2020
)
2121
from fairseq2.models.transformer_decoder import TransformerDecoderModel
2222
from 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+
195268
def init_llama_scaled_freqs(
196269
pos_encoder: RotaryEncoder, rope_scaling: LLaMARopeScalingConfig
197270
) -> Tensor:

src/fairseq2/models/transformer/__init__.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,16 +22,16 @@
2222
from fairseq2.models.transformer._factory import (
2323
create_transformer_model as create_transformer_model,
2424
)
25+
from fairseq2.models.transformer._factory import (
26+
init_final_projection as init_final_projection,
27+
)
2528
from fairseq2.models.transformer._frontend import (
2629
TransformerEmbeddingFrontend as TransformerEmbeddingFrontend,
2730
)
2831
from fairseq2.models.transformer._frontend import (
2932
TransformerFrontend as TransformerFrontend,
3033
)
3134
from fairseq2.models.transformer._model import TransformerModel as TransformerModel
32-
from fairseq2.models.transformer._model import (
33-
init_final_projection as init_final_projection,
34-
)
3535

3636
# isort: split
3737

src/fairseq2/models/transformer/_factory.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66

77
from __future__ import annotations
88

9+
import torch.nn as nn
10+
911
from fairseq2.models.transformer._config import TransformerConfig
1012
from fairseq2.models.transformer._frontend import (
1113
TransformerEmbeddingFrontend,
@@ -174,3 +176,10 @@ def create_final_proj(self, embed: Embedding) -> Projection:
174176
return TiedProjection(embed.weight, bias=None)
175177

176178
return Linear(config.model_dim, config.vocab_info.size, bias=False)
179+
180+
181+
def init_final_projection(proj: Linear) -> None:
182+
nn.init.normal_(proj.weight, std=proj.input_dim**-0.5)
183+
184+
if proj.bias is not None:
185+
nn.init.zeros_(proj.bias)

src/fairseq2/models/transformer/_model.py

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -8,15 +8,14 @@
88

99
from typing import final
1010

11-
import torch.nn as nn
1211
from torch import Tensor
1312
from typing_extensions import override
1413

1514
from fairseq2.data import VocabularyInfo
1615
from fairseq2.models.encoder_decoder import EncoderDecoderModel
1716
from fairseq2.models.sequence import SequenceModelOutput
1817
from fairseq2.models.transformer._frontend import TransformerFrontend
19-
from fairseq2.nn import IncrementalStateBag, Linear, Projection
18+
from fairseq2.nn import IncrementalStateBag, Projection
2019
from fairseq2.nn.padding import PaddingMask
2120
from fairseq2.nn.transformer import TransformerDecoder, TransformerEncoder
2221

@@ -106,10 +105,3 @@ def project(
106105
logits = self.final_proj(decoder_output)
107106

108107
return SequenceModelOutput(logits, self.target_vocab_info.pad_idx)
109-
110-
111-
def init_final_projection(proj: Linear) -> None:
112-
nn.init.normal_(proj.weight, std=proj.input_dim**-0.5)
113-
114-
if proj.bias is not None:
115-
nn.init.zeros_(proj.bias)

src/fairseq2/models/wav2vec2/asr/_model.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,7 @@ def __init__(
7474
self.model_dim,
7575
target_vocab_info.size,
7676
bias=True,
77-
init_fn=init_final_projection,
77+
init_fn=_init_final_projection,
7878
device=device,
7979
dtype=dtype,
8080
)
@@ -105,7 +105,7 @@ def forward(self, batch: SequenceBatch) -> AsrModelOutput:
105105
return AsrModelOutput(logits, padding_mask)
106106

107107

108-
def init_final_projection(proj: Linear) -> None:
108+
def _init_final_projection(proj: Linear) -> None:
109109
"""Initialize ``proj`` as the final projection of a wav2vec 2.0 ASR model."""
110110
nn.init.xavier_uniform_(proj.weight)
111111

0 commit comments

Comments
 (0)