| 646 | |
| 647 | |
| 648 | class PWAM(nn.Module): |
| 649 | def __init__(self, dim, v_in_channels, l_in_channels, key_channels, value_channels, num_heads=0, dropout=0.0): |
| 650 | super(PWAM, self).__init__() |
| 651 | # input x shape: (B, H*W, dim) |
| 652 | self.vis_project = nn.Sequential(nn.Conv1d(dim, dim, 1, 1), # the init function sets bias to 0 if bias is True |
| 653 | nn.GELU(), |
| 654 | nn.Dropout(dropout) |
| 655 | ) |
| 656 | |
| 657 | self.image_lang_att = SpatialImageLanguageAttention(v_in_channels, # v_in |
| 658 | l_in_channels, # l_in |
| 659 | key_channels, # key |
| 660 | value_channels, # value |
| 661 | out_channels=value_channels, # out |
| 662 | num_heads=num_heads) |
| 663 | |
| 664 | self.project_mm = nn.Sequential(nn.Conv1d(value_channels, value_channels, 1, 1), |
| 665 | nn.GELU(), |
| 666 | nn.Dropout(dropout) |
| 667 | ) |
| 668 | |
| 669 | def forward(self, x, l, l_mask): |
| 670 | vis = self.vis_project(x.permute(0, 2, 1)) # (B, dim, H*W) |
| 671 | lang = self.image_lang_att(x, l, l_mask) # (B, H*W, dim) |
| 672 | lang = lang.permute(0, 2, 1) # (B, dim, H*W) |
| 673 | mm = torch.mul(vis, lang) |
| 674 | mm = self.project_mm(mm) # (B, dim, H*W) |
| 675 | mm = mm.permute(0, 2, 1) # (B, H*W, dim) |
| 676 | return mm |
| 677 | |
| 678 | class SpatialImageLanguageAttention(nn.Module): |
| 679 | def __init__(self, v_in_channels, l_in_channels, key_channels, value_channels, out_channels=None, num_heads=1): |