| 262 | |
| 263 | |
| 264 | class Transformer(nn.Module): |
| 265 | def __init__(self, params): |
| 266 | super().__init__() |
| 267 | self.params = params |
| 268 | self.vocab_size = params.vocab_size |
| 269 | self.n_layers = params.num_hidden_layers |
| 270 | |
| 271 | self.embed_tokens = nn.Embedding(params.vocab_size, params.hidden_size) |
| 272 | |
| 273 | self.layers = torch.nn.ModuleList() |
| 274 | for layer_id in range(params.num_hidden_layers): |
| 275 | self.layers.append(TransformerBlock(layer_id, params)) |
| 276 | |
| 277 | self.norm = RMSNorm(params.hidden_size, eps=params.rms_norm_eps) |
| 278 | |
| 279 | self.freqs_cis = precompute_freqs_cis( |
| 280 | self.params.hidden_size // self.params.num_attention_heads, |
| 281 | self.params.max_position_embeddings * 2, |
| 282 | ) |
| 283 | |
| 284 | @torch.inference_mode() |
| 285 | def forward(self, tokens: torch.Tensor, start_pos: int): |
| 286 | _bsz, seqlen = tokens.shape |
| 287 | h = self.embed_tokens(tokens) |
| 288 | self.freqs_cis = self.freqs_cis.to(h.device) |
| 289 | freqs_cis = self.freqs_cis[start_pos : start_pos + seqlen] |
| 290 | |
| 291 | mask = None |
| 292 | if seqlen > 1: |
| 293 | mask = torch.full( |
| 294 | (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device |
| 295 | ) |
| 296 | mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h) |
| 297 | for layer in self.layers: |
| 298 | h = layer(h, start_pos, freqs_cis, mask) |
| 299 | h = self.norm(h) |
| 300 | return h |
| 301 | |
| 302 | |
| 303 | class LlamaForCausalLM(nn.Module): |