MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / SelfAttention

Class SelfAttention

diffsynth/models/wan_video_dit.py:300–382  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

298
299
300class SelfAttention(nn.Module):
301 def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
302 super().__init__()
303 self.dim = dim
304 self.num_heads = num_heads
305 self.head_dim = dim // num_heads
306
307 self.q = nn.Linear(dim, dim)
308 self.k = nn.Linear(dim, dim)
309 self.v = nn.Linear(dim, dim)
310 self.o = nn.Linear(dim, dim)
311 self.norm_q = RMSNorm(dim, eps=eps)
312 self.norm_k = RMSNorm(dim, eps=eps)
313
314 self.attn = AttentionModule(self.num_heads)
315 self.local_attn_mask = None
316
317 def forward(self, x, freqs, f=None, h=None, w=None, local_num=None, topk=None,
318 train_img=False, block_id=None, kv_len=None, is_full_block=False,
319 is_stream=False, pre_cache_k=None, pre_cache_v=None, local_range = 9):
320 B, L, D = x.shape
321 if is_stream and pre_cache_k is not None and pre_cache_v is not None:
322 assert f==2, "f must be 2"
323 if is_stream and (pre_cache_k is None or pre_cache_v is None):
324 assert f==6, " start f must be 6"
325 assert L == f * h * w, "Sequence length mismatch with provided (f,h,w)."
326
327 q = self.norm_q(self.q(x))
328 k = self.norm_k(self.k(x))
329 v = self.v(x)
330 q = rope_apply(q, freqs, self.num_heads)
331 k = rope_apply(k, freqs, self.num_heads)
332
333 win = (2, 8, 8)
334 q = q.view(B, f, h, w, D)
335 k = k.view(B, f, h, w, D)
336 v = v.view(B, f, h, w, D)
337
338 q_w = WindowPartition3D.partition(q, win)
339 k_w = WindowPartition3D.partition(k, win)
340 v_w = WindowPartition3D.partition(v, win)
341
342 seqlen = f//win[0]
343 one_len = k_w.shape[0] // B // seqlen
344 if pre_cache_k is not None and pre_cache_v is not None:
345 k_w = torch.cat([pre_cache_k, k_w], dim=0)
346 v_w = torch.cat([pre_cache_v, v_w], dim=0)
347
348 block_n = q_w.shape[0] // B
349 block_s = q_w.shape[1]
350 block_n_kv = k_w.shape[0] // B
351
352 reorder_q = rearrange(q_w, '(b block_n) (block_s) d -> b (block_n block_s) d', block_n=block_n, block_s=block_s)
353 reorder_k = rearrange(k_w, '(b block_n) (block_s) d -> b (block_n block_s) d', block_n=block_n_kv, block_s=block_s)
354 reorder_v = rearrange(v_w, '(b block_n) (block_s) d -> b (block_n block_s) d', block_n=block_n_kv, block_s=block_s)
355
356 window_size = win[0]*h*w//128
357

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected