MCPcopy Create free account
hub / github.com/SpatialVLA/SpatialVLA / __init__

Method __init__

model/modeling_gemma2.py:655–671  ·  view source on GitHub ↗
(self, config: Gemma2Config)

Source from the content-addressed store, hash-verified

653 """
654
655 def __init__(self, config: Gemma2Config):
656 super().__init__(config)
657 self.padding_idx = config.pad_token_id
658 self.vocab_size = config.vocab_size
659
660 self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
661 self.layers = nn.ModuleList(
662 [Gemma2DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
663 )
664 self.norm = Gemma2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
665
666 self.gradient_checkpointing = False
667 if getattr(config, "pretraining_tp", 1) != 1:
668 logger.warn("`pretraining_tp` is deprecated, please use `model.tensor_parallel` instead.")
669
670 # Initialize weights and apply final processing
671 self.post_init()
672
673 def get_input_embeddings(self):
674 return self.embed_tokens

Callers

nothing calls this directly

Calls 3

Gemma2DecoderLayerClass · 0.85
Gemma2RMSNormClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected