MCPcopy Create free account
hub / github.com/MeiGen-AI/MultiTalk / SingleStreamAttention

Class SingleStreamAttention

wan/modules/attention.py:191–280  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

189
190
191class SingleStreamAttention(nn.Module):
192 def __init__(
193 self,
194 dim: int,
195 encoder_hidden_states_dim: int,
196 num_heads: int,
197 qkv_bias: bool,
198 qk_norm: bool,
199 norm_layer: nn.Module,
200 attn_drop: float = 0.0,
201 proj_drop: float = 0.0,
202 eps: float = 1e-6,
203 ) -> None:
204 super().__init__()
205 assert dim % num_heads == 0, "dim should be divisible by num_heads"
206 self.dim = dim
207 self.encoder_hidden_states_dim = encoder_hidden_states_dim
208 self.num_heads = num_heads
209 self.head_dim = dim // num_heads
210 self.scale = self.head_dim**-0.5
211 self.qk_norm = qk_norm
212
213 self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
214
215 self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity()
216 self.k_norm = norm_layer(self.head_dim,eps=eps) if qk_norm else nn.Identity()
217
218 self.attn_drop = nn.Dropout(attn_drop)
219 self.proj = nn.Linear(dim, dim)
220 self.proj_drop = nn.Dropout(proj_drop)
221
222 self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
223
224 self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
225 self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
226
227 def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor:
228
229 N_t, N_h, N_w = shape
230 if not enable_sp:
231 x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
232
233 # get q for hidden_state
234 B, N, C = x.shape
235 q = self.q_linear(x)
236 q_shape = (B, N, self.num_heads, self.head_dim)
237 q = q.view(q_shape).permute((0, 2, 1, 3))
238
239 if self.qk_norm:
240 q = self.q_norm(q)
241
242 # get kv from encoder_hidden_states
243 _, N_a, _ = encoder_hidden_states.shape
244 encoder_kv = self.kv_linear(encoder_hidden_states)
245 encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
246 encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
247 encoder_k, encoder_v = encoder_kv.unbind(0)
248

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected