Initialize a Transformer model. Args: params (ModelArgs): Model configuration parameters. Attributes: params (ModelArgs): Model configuration parameters. vocab_size (int): Vocabulary size. n_layers (int): Number of layers in
(self, params: ModelArgs)
| 418 | |
| 419 | class Transformer(nn.Module): |
| 420 | def __init__(self, params: ModelArgs): |
| 421 | """ |
| 422 | Initialize a Transformer model. |
| 423 | |
| 424 | Args: |
| 425 | params (ModelArgs): Model configuration parameters. |
| 426 | |
| 427 | Attributes: |
| 428 | params (ModelArgs): Model configuration parameters. |
| 429 | vocab_size (int): Vocabulary size. |
| 430 | n_layers (int): Number of layers in the model. |
| 431 | tok_embeddings (ParallelEmbedding): Token embeddings. |
| 432 | layers (torch.nn.ModuleList): List of Transformer blocks. |
| 433 | norm (RMSNorm): Layer normalization for the model output. |
| 434 | output (ColumnParallelLinear): Linear layer for final output. |
| 435 | freqs_cis (torch.Tensor): Precomputed cosine and sine frequencies. |
| 436 | |
| 437 | """ |
| 438 | super().__init__() |
| 439 | self.params = params |
| 440 | self.vocab_size = params.vocab_size |
| 441 | self.n_layers = params.n_layers |
| 442 | |
| 443 | self.tok_embeddings = ParallelEmbedding( |
| 444 | params.vocab_size, params.dim, init_method=lambda x: x |
| 445 | ) |
| 446 | |
| 447 | self.layers = torch.nn.ModuleList() |
| 448 | for layer_id in range(params.n_layers): |
| 449 | self.layers.append(TransformerBlock(layer_id, params)) |
| 450 | |
| 451 | self.norm = RMSNorm(params.dim, eps=params.norm_eps) |
| 452 | self.output = ColumnParallelLinear( |
| 453 | params.dim, params.vocab_size, bias=False, init_method=lambda x: x |
| 454 | ) |
| 455 | |
| 456 | self.freqs_cis = precompute_freqs_cis( |
| 457 | # Note that self.params.max_seq_len is multiplied by 2 because the token limit for the Llama 2 generation of models is 4096. |
| 458 | # Adding this multiplier instead of using 4096 directly allows for dynamism of token lengths while training or fine-tuning. |
| 459 | self.params.dim // self.params.n_heads, |
| 460 | self.params.max_seq_len * 2, |
| 461 | ) |
| 462 | |
| 463 | @torch.inference_mode() |
| 464 | def forward(self, tokens: torch.Tensor, start_pos: int): |
nothing calls this directly
no test coverage detected