| 621 | |
| 622 | |
| 623 | class Attention1D(nn.Module): |
| 624 | def __init__(self, |
| 625 | dim_in, |
| 626 | dim_out, |
| 627 | num_heads, |
| 628 | qkv_bias=False, |
| 629 | attn_drop=0., |
| 630 | proj_drop=0., |
| 631 | with_cls_token=False, |
| 632 | skip_conv_proj=False, |
| 633 | n_layer=1, |
| 634 | **kwargs |
| 635 | ): |
| 636 | super().__init__() |
| 637 | self.dim = dim_out |
| 638 | self.num_heads = num_heads |
| 639 | self.n_layer = n_layer |
| 640 | # head_dim = self.qkv_dim // num_heads |
| 641 | self.scale = dim_out ** -0.5 |
| 642 | self.with_cls_token = with_cls_token |
| 643 | |
| 644 | self.proj_q = nn.ModuleList() |
| 645 | self.proj_k = nn.ModuleList() |
| 646 | self.proj_v = nn.ModuleList() |
| 647 | self.attn_drop = nn.ModuleList() |
| 648 | self.proj = nn.ModuleList() |
| 649 | self.proj_drop = nn.ModuleList() |
| 650 | |
| 651 | for _ in range(n_layer): |
| 652 | self.proj_q.append(nn.Linear(dim_in, dim_out, bias=qkv_bias)) |
| 653 | self.proj_k.append(nn.Linear(dim_in, dim_out, bias=qkv_bias)) |
| 654 | self.proj_v.append(nn.Linear(dim_in, dim_out, bias=qkv_bias)) |
| 655 | |
| 656 | self.attn_drop.append(nn.Dropout(attn_drop)) |
| 657 | self.proj.append(nn.Linear(dim_out, dim_out)) |
| 658 | self.proj_drop.append(nn.Dropout(proj_drop)) |
| 659 | |
| 660 | dim_in = dim_out |
| 661 | |
| 662 | def forward(self, x, t, h, w): |
| 663 | x = rearrange(x, 'b c t H W -> b t (c H W)') |
| 664 | if self.n_layer == 0: |
| 665 | return None |
| 666 | |
| 667 | for idx in range(self.n_layer): |
| 668 | q = rearrange(self.proj_q[idx](x), 'b t (h d) -> b h t d', h=self.num_heads) |
| 669 | k = rearrange(self.proj_k[idx](x), 'b t (h d) -> b h t d', h=self.num_heads) |
| 670 | v = rearrange(self.proj_v[idx](x), 'b t (h d) -> b h t d', h=self.num_heads) |
| 671 | |
| 672 | attn_score = torch.einsum('bhlk,bhtk->bhlt', [q, k]) * self.scale |
| 673 | attn = F.softmax(attn_score, dim=-1) |
| 674 | attn = self.attn_drop[idx](attn) |
| 675 | |
| 676 | x = torch.einsum('bhlt,bhtv->bhlv', [attn, v]) |
| 677 | x = rearrange(x, 'b h t d -> b t (h d)') |
| 678 | |
| 679 | x = self.proj[idx](x) |
| 680 | x = self.proj_drop[idx](x) |