| 298 | |
| 299 | |
| 300 | class 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 | |