MCPcopy Create free account
hub / github.com/pytorch/examples / __init__

Method __init__

distributed/tensor_parallelism/llama2_model.py:367–393  ·  view source on GitHub ↗
(self, model_args: ModelArgs)

Source from the content-addressed store, hash-verified

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 """

Callers

nothing calls this directly

Calls 5

init_weightsMethod · 0.95
precompute_freqs_cisFunction · 0.85
RMSNormClass · 0.85
TransformerBlockClass · 0.70
__init__Method · 0.45

Tested by

no test coverage detected