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

Class Gemma2DecoderLayer

model/modeling_gemma2.py:436–506  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

434
435
436class Gemma2DecoderLayer(nn.Module):
437 def __init__(self, config: Gemma2Config, layer_idx: int):
438 super().__init__()
439 self.hidden_size = config.hidden_size
440 self.config = config
441 self.is_sliding = not bool(layer_idx % 2)
442 self.self_attn = Gemma2Attention(config=config, layer_idx=layer_idx)
443 self.mlp = Gemma2MLP(config)
444 self.input_layernorm = Gemma2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
445 self.post_attention_layernorm = Gemma2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
446
447 self.pre_feedforward_layernorm = Gemma2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
448 self.post_feedforward_layernorm = Gemma2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
449 self.sliding_window = config.sliding_window
450
451 def forward(
452 self,
453 hidden_states: torch.Tensor,
454 attention_mask: Optional[torch.Tensor] = None,
455 position_ids: Optional[torch.LongTensor] = None,
456 past_key_value: Optional[Cache] = None,
457 output_attentions: Optional[bool] = False,
458 use_cache: Optional[bool] = False,
459 cache_position: Optional[torch.LongTensor] = None,
460 ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
461 if self.is_sliding and attention_mask is not None: # efficient SDPA and no padding
462 # Flash-attn is a 2D tensor
463 if self.config._attn_implementation == "flash_attention_2":
464 if past_key_value is not None: # when decoding
465 attention_mask = attention_mask[:, -self.sliding_window :]
466 else:
467 min_dtype = torch.finfo(hidden_states.dtype).min
468 sliding_window_mask = torch.tril(
469 torch.ones_like(attention_mask, dtype=torch.bool), diagonal=-self.sliding_window
470 )
471 attention_mask = torch.where(sliding_window_mask, min_dtype, attention_mask)
472 if attention_mask.shape[-1] <= 1: # when decoding
473 attention_mask = attention_mask[:, :, :, -self.sliding_window :]
474
475 residual = hidden_states
476
477 hidden_states = self.input_layernorm(hidden_states)
478
479 # Self Attention
480 hidden_states, self_attn_weights, present_key_value = self.self_attn(
481 hidden_states=hidden_states,
482 attention_mask=attention_mask,
483 position_ids=position_ids,
484 past_key_value=past_key_value,
485 output_attentions=output_attentions,
486 use_cache=use_cache,
487 cache_position=cache_position,
488 )
489 hidden_states = self.post_attention_layernorm(hidden_states)
490 hidden_states = residual + hidden_states
491
492 residual = hidden_states
493 hidden_states = self.pre_feedforward_layernorm(hidden_states)

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected