MCPcopy Create free account
hub / github.com/OpenSparseLLMs/MoM / shared_o

Method shared_o

mom/layers/mom.py:612–697  ·  view source on GitHub ↗
(
        self,
        hidden_states: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        recurrent_state=None,
        use_cache: Optional[bool] = False,
        conv_state_q=[None, None],
        conv_state_k=[None, None],
        conv_state_v=[None, None],
        **kwargs
    )

Source from the content-addressed store, hash-verified

610 return o, None, past_key_values, router_logits.view(-1, self.num_memories)
611
612 def shared_o(
613 self,
614 hidden_states: torch.Tensor,
615 attention_mask: Optional[torch.Tensor] = None,
616 recurrent_state=None,
617 use_cache: Optional[bool] = False,
618 conv_state_q=[None, None],
619 conv_state_k=[None, None],
620 conv_state_v=[None, None],
621 **kwargs
622 ) -> torch.Tensor:
623 if attention_mask is not None:
624 assert len(attention_mask.shape) == 2, (
625 "Expected attention_mask as a 0-1 matrix with shape [batch_size, seq_len] "
626 "for padding purposes (0 indicating padding). "
627 "Arbitrary attention masks of shape [batch_size, seq_len, seq_len] are not allowed."
628 )
629
630 mode = 'fused_recurrent' if hidden_states.shape[1] <= 64 else self.mode
631 if self.training:
632 assert mode == 'chunk', "Only chunk mode is supported in training."
633
634 cu_seqlens = None
635 if attention_mask is not None:
636 batch_size, q_len = hidden_states.shape[0], hidden_states.shape[1]
637 indices, cu_seqlens, _ = get_unpad_data(attention_mask[:, -q_len:])
638 hidden_states = index_first_axis(rearrange(hidden_states, "b s ... -> (b s) ..."), indices).unsqueeze(0)
639
640 if self.use_short_conv:
641 q, conv_state_q[1] = self.q_conv1d(
642 x=self.q_proj(hidden_states),
643 cache=conv_state_q[1],
644 output_final_state=use_cache,
645 cu_seqlens=cu_seqlens
646 )
647 k, conv_state_k[1] = self.k_conv1d(
648 x=self.shared_k(hidden_states),
649 cache=conv_state_k[1],
650 output_final_state=use_cache,
651 cu_seqlens=cu_seqlens
652 )
653 v, conv_state_v[1] = self.v_conv1d(
654 x=self.shared_v(hidden_states),
655 cache=conv_state_v[1],
656 output_final_state=use_cache,
657 cu_seqlens=cu_seqlens
658 )
659 else:
660 q = self.silu(self.q_proj(hidden_states))
661 k = self.silu(self.shared_k(hidden_states))
662 v = self.silu(self.shared_v(hidden_states))
663
664 q, k, v = map(lambda x: rearrange(x, 'b t (h d) -> b t h d', h=self.num_heads), (q, k, v))
665 beta = self.shared_b(hidden_states).sigmoid()
666 g = -self.A_log.float().exp() * F.softplus(self.shared_a(hidden_states).float() + self.dt_bias)
667
668 if mode == 'chunk':
669 o, recurrent_state[-1] = chunk_gated_delta_rule(

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected