| 410 | return x + skip |
| 411 | |
| 412 | class NeighborhoodSelfAttentionBlock(nn.Module): |
| 413 | def __init__(self, d_model, d_head, cond_features, kernel_size, dropout=0.0): |
| 414 | super().__init__() |
| 415 | self.d_head = d_head |
| 416 | self.n_heads = d_model // d_head |
| 417 | self.kernel_size = kernel_size |
| 418 | self.norm = AdaRMSNorm(d_model, cond_features) |
| 419 | self.qkv_proj = apply_wd(Linear(d_model, d_model * 3, bias=False)) |
| 420 | self.scale = nn.Parameter(torch.full([self.n_heads], 10.0)) |
| 421 | self.pos_emb = AxialRoPE(d_head // 2, self.n_heads) |
| 422 | self.dropout = nn.Dropout(dropout) |
| 423 | self.out_proj = apply_wd(zero_init(Linear(d_model, d_model, bias=False))) |
| 424 | |
| 425 | def extra_repr(self): |
| 426 | return f"d_head={self.d_head}, kernel_size={self.kernel_size}" |
| 427 | |
| 428 | def forward(self, x, pos, cond): |
| 429 | skip = x |
| 430 | x = self.norm(x, cond) |
| 431 | qkv = self.qkv_proj(x) |
| 432 | if natten is None: |
| 433 | raise ModuleNotFoundError("natten is required for neighborhood attention") |
| 434 | if natten.has_fused_na(): |
| 435 | q, k, v = rearrange(qkv, "n h w (t nh e) -> t n h w nh e", t=3, e=self.d_head) |
| 436 | # print("q", q.shape) #[1, 32, 32, 4, 64]) |
| 437 | # print("k", k.shape) |
| 438 | q, k = scale_for_cosine_sim(q, k, self.scale[:, None], 1e-6) |
| 439 | theta = self.pos_emb(pos) |
| 440 | q = apply_rotary_emb_(q, theta) |
| 441 | k = apply_rotary_emb_(k, theta) |
| 442 | flops.op(flops.op_natten, q.shape, k.shape, v.shape, self.kernel_size) |
| 443 | x = natten.functional.na2d(q, k, v, self.kernel_size, scale=1.0) |
| 444 | x = rearrange(x, "n h w nh e -> n h w (nh e)") |
| 445 | else: |
| 446 | q, k, v = rearrange(qkv, "n h w (t nh e) -> t n nh h w e", t=3, e=self.d_head) |
| 447 | q, k = scale_for_cosine_sim(q, k, self.scale[:, None, None, None], 1e-6) |
| 448 | theta = self.pos_emb(pos).movedim(-2, -4) |
| 449 | q = apply_rotary_emb_(q, theta) |
| 450 | k = apply_rotary_emb_(k, theta) |
| 451 | flops.op(flops.op_natten, q.shape, k.shape, v.shape, self.kernel_size) |
| 452 | qk = natten.functional.na2d_qk(q, k, self.kernel_size) |
| 453 | a = torch.softmax(qk, dim=-1).to(v.dtype) |
| 454 | x = natten.functional.na2d_av(a, v, self.kernel_size) |
| 455 | x = rearrange(x, "n nh h w e -> n h w (nh e)") |
| 456 | x = self.dropout(x) |
| 457 | x = self.out_proj(x) |
| 458 | # exit(1) |
| 459 | return x + skip |
| 460 | |
| 461 | class ShiftedWindowSelfAttentionBlock(nn.Module): |
| 462 | def __init__(self, d_model, d_head, cond_features, window_size, window_shift, dropout=0.0): |