MCPcopy Create free account
hub / github.com/kyegomez/BitNet / __init__

Method __init__

bitnet/bit_llama.py:420–461  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

418
419class 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):

Callers

nothing calls this directly

Calls 4

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

Tested by

no test coverage detected