| 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( |