(self, model_args: ModelArgs)
| 365 | """ |
| 366 | |
| 367 | def __init__(self, model_args: ModelArgs): |
| 368 | super().__init__() |
| 369 | self.model_args = model_args |
| 370 | self.vocab_size = model_args.vocab_size |
| 371 | self.n_layers = model_args.n_layers |
| 372 | self.model_dim = model_args.dim |
| 373 | |
| 374 | self.tok_embeddings = nn.Embedding(model_args.vocab_size, model_args.dim) |
| 375 | self.register_buffer( |
| 376 | "freqs_cis", |
| 377 | precompute_freqs_cis( |
| 378 | model_args.dim // model_args.n_heads, |
| 379 | # Need to compute until at least the max token limit for generation |
| 380 | # (use 2x max sequence length to be safe) |
| 381 | model_args.max_seq_len * 2, |
| 382 | ), |
| 383 | ) |
| 384 | self.layers = torch.nn.ModuleList() |
| 385 | for layer_id in range(model_args.n_layers): |
| 386 | self.layers.append(TransformerBlock(layer_id, model_args)) |
| 387 | |
| 388 | self.norm = RMSNorm( |
| 389 | dim=model_args.dim, eps=model_args.norm_eps |
| 390 | ) |
| 391 | |
| 392 | self.output = nn.Linear(model_args.dim, model_args.vocab_size, bias=False) |
| 393 | self.init_weights() |
| 394 | |
| 395 | def init_weights(self): |
| 396 | """ |
nothing calls this directly
no test coverage detected