| 49 | train_router_metrics: Optional[Dict] = None |
| 50 | |
| 51 | class MSADeocoderLayer(Qwen3DecoderLayer): |
| 52 | def __init__(self, config: Qwen3Config, layer_idx: int, attn_type: str = "sparse_attention"): |
| 53 | super().__init__(config=config, layer_idx=layer_idx) |
| 54 | self.layer_idx = layer_idx |
| 55 | self.attn_type = attn_type |
| 56 | self.hidden_size = config.hidden_size |
| 57 | if attn_type == "full_attention": |
| 58 | self.self_attn = Qwen3Attention(config=config, layer_idx=layer_idx) |
| 59 | elif attn_type == "sparse_attention": |
| 60 | self.self_attn = MemorySparseAttention(config=config, layer_idx=layer_idx) |
| 61 | |
| 62 | self.mlp = Qwen3MLP(config) |
| 63 | self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| 64 | self.post_attention_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| 65 | config.sliding_window = False |
| 66 | |
| 67 | def forward( |
| 68 | self, |
| 69 | hidden_states: torch.Tensor, |
| 70 | attention_mask: Optional[torch.Tensor] = None, |
| 71 | position_ids: Optional[torch.LongTensor] = None, |
| 72 | past_key_value: Optional[Cache] = None, |
| 73 | output_attentions: Optional[bool] = False, |
| 74 | output_docs_score: Optional[bool] = False, |
| 75 | use_cache: Optional[bool] = False, |
| 76 | cache_position: Optional[torch.LongTensor] = None, |
| 77 | position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC |
| 78 | doc_ids: Optional[torch.Tensor] = None, |
| 79 | input_ids: Optional[torch.LongTensor] = None, |
| 80 | **kwargs: Unpack[FlashAttentionKwargs], |
| 81 | ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]: |
| 82 | residual = hidden_states |
| 83 | |
| 84 | hidden_states = self.input_layernorm(hidden_states) |
| 85 | |
| 86 | # Self Attention |
| 87 | if self.attn_type == "full_attention": |
| 88 | hidden_states, self_attn_weights = self.self_attn( |
| 89 | hidden_states=hidden_states, |
| 90 | attention_mask=attention_mask, |
| 91 | position_ids=position_ids, |
| 92 | past_key_value=past_key_value, |
| 93 | output_attentions=output_attentions, |
| 94 | use_cache=use_cache, |
| 95 | cache_position=cache_position, |
| 96 | position_embeddings=position_embeddings, |
| 97 | doc_ids=doc_ids, |
| 98 | input_ids=input_ids, |
| 99 | **kwargs, |
| 100 | ) |
| 101 | else: |
| 102 | hidden_states, self_attn_weights = self.self_attn( |
| 103 | hidden_states=hidden_states, |
| 104 | attention_mask=attention_mask, |
| 105 | position_ids=position_ids, |
| 106 | past_key_value=past_key_value, |
| 107 | output_attentions=output_attentions, |
| 108 | use_cache=use_cache, |