MCPcopy Create free account
hub / github.com/EverMind-AI/MSA / MSADeocoderLayer

Class MSADeocoderLayer

src/msa/model.py:51–135  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

49 train_router_metrics: Optional[Dict] = None
50
51class 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,

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected