MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / NeighborhoodSelfAttentionBlock

Class NeighborhoodSelfAttentionBlock

k_diffusion/models/modules.py:412–459  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

410 return x + skip
411
412class 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
461class ShiftedWindowSelfAttentionBlock(nn.Module):
462 def __init__(self, d_model, d_head, cond_features, window_size, window_shift, dropout=0.0):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected