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

Method forward

model/modeling_gemma2.py:680–793  ·  view source on GitHub ↗
(
        self,
        input_ids: torch.LongTensor = None,
        attention_mask: Optional[torch.Tensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        past_key_values: Optional[HybridCache] = None,
        inputs_embeds: Optional[torch.FloatTensor] = None,
        use_cache: Optional[bool] = None,
        output_attentions: Optional[bool] = None,
        output_hidden_states: Optional[bool] = None,
        return_dict: Optional[bool] = None,
        cache_position: Optional[torch.LongTensor] = None,
    )

Source from the content-addressed store, hash-verified

678
679 @add_start_docstrings_to_model_forward(GEMMA2_INPUTS_DOCSTRING)
680 def forward(
681 self,
682 input_ids: torch.LongTensor = None,
683 attention_mask: Optional[torch.Tensor] = None,
684 position_ids: Optional[torch.LongTensor] = None,
685 past_key_values: Optional[HybridCache] = None,
686 inputs_embeds: Optional[torch.FloatTensor] = None,
687 use_cache: Optional[bool] = None,
688 output_attentions: Optional[bool] = None,
689 output_hidden_states: Optional[bool] = None,
690 return_dict: Optional[bool] = None,
691 cache_position: Optional[torch.LongTensor] = None,
692 ) -> Union[Tuple, BaseModelOutputWithPast]:
693 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
694 output_hidden_states = (
695 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
696 )
697 use_cache = use_cache if use_cache is not None else self.config.use_cache
698 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
699
700 if (input_ids is None) ^ (inputs_embeds is not None):
701 raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
702
703 if self.gradient_checkpointing and self.training and use_cache:
704 logger.warning_once(
705 "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`."
706 )
707 use_cache = False
708
709 if inputs_embeds is None:
710 inputs_embeds = self.embed_tokens(input_ids)
711
712 if use_cache and past_key_values is None and not self.training:
713 batch_size, seq_len, _ = inputs_embeds.shape
714 past_key_values = HybridCache(
715 self.config,
716 batch_size=batch_size,
717 max_cache_len=seq_len,
718 device=self.device,
719 dtype=inputs_embeds.dtype,
720 )
721
722 if cache_position is None:
723 past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
724 cache_position = torch.arange(
725 past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
726 )
727
728 if position_ids is None:
729 position_ids = cache_position.unsqueeze(0)
730
731 causal_mask = self._update_causal_mask(
732 attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions
733 )
734
735 # embed positions
736 hidden_states = inputs_embeds
737

Callers

nothing calls this directly

Calls 1

_update_causal_maskMethod · 0.95

Tested by

no test coverage detected