| 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): |
| 680 | super(SpatialImageLanguageAttention, self).__init__() |
| 681 | # x shape: (B, H*W, v_in_channels) |
| 682 | # l input shape: (B, l_in_channels, N_l) |
| 683 | # l_mask shape: (B, N_l, 1) |
| 684 | self.v_in_channels = v_in_channels |
| 685 | self.l_in_channels = l_in_channels |
| 686 | self.out_channels = out_channels |
| 687 | self.key_channels = key_channels |
| 688 | self.value_channels = value_channels |
| 689 | self.num_heads = num_heads |
| 690 | if out_channels is None: |
| 691 | self.out_channels = self.value_channels |
| 692 | |
| 693 | # Keys: language features: (B, l_in_channels, #words) |
| 694 | # avoid any form of spatial normalization because a sentence contains many padding 0s |
| 695 | self.f_key = nn.Sequential( |
| 696 | nn.Conv1d(self.l_in_channels, self.key_channels, kernel_size=1, stride=1), |
| 697 | ) |
| 698 | |
| 699 | # Queries: visual features: (B, H*W, v_in_channels) |
| 700 | self.f_query = nn.Sequential( |
| 701 | nn.Conv1d(self.v_in_channels, self.key_channels, kernel_size=1, stride=1), |
| 702 | nn.InstanceNorm1d(self.key_channels), |
| 703 | ) |
| 704 | |
| 705 | # Values: language features: (B, l_in_channels, #words) |
| 706 | self.f_value = nn.Sequential( |
| 707 | nn.Conv1d(self.l_in_channels, self.value_channels, kernel_size=1, stride=1), |
| 708 | ) |
| 709 | |
| 710 | # Out projection |
| 711 | self.W = nn.Sequential( |
| 712 | nn.Conv1d(self.value_channels, self.out_channels, kernel_size=1, stride=1), |
| 713 | nn.InstanceNorm1d(self.out_channels), |
| 714 | ) |
| 715 | |
| 716 | def forward(self, x, l, l_mask): |
| 717 | B, HW = x.size(0), x.size(1) |
| 718 | x = x.permute(0, 2, 1) # (B, key_channels, H*W) |
| 719 | l_mask = l_mask.permute(0, 2, 1) # (B, N_l, 1) -> (B, 1, N_l) |
| 720 | |
| 721 | query = self.f_query(x) # (B, key_channels, H*W) if Conv1D |
| 722 | query = query.permute(0, 2, 1) # (B, H*W, key_channels) |
| 723 | key = self.f_key(l) # (B, key_channels, N_l) |
| 724 | value = self.f_value(l) # (B, self.value_channels, N_l) |
| 725 | key = key * l_mask # (B, key_channels, N_l) |
| 726 | value = value * l_mask # (B, self.value_channels, N_l) |
| 727 | n_l = value.size(-1) |
| 728 | query = query.reshape(B, HW, self.num_heads, self.key_channels//self.num_heads).permute(0, 2, 1, 3) |
| 729 | # (b, num_heads, H*W, self.key_channels//self.num_heads) |
| 730 | key = key.reshape(B, self.num_heads, self.key_channels//self.num_heads, n_l) |
| 731 | # (b, num_heads, self.key_channels//self.num_heads, n_l) |
| 732 | value = value.reshape(B, self.num_heads, self.value_channels//self.num_heads, n_l) |
| 733 | # # (b, num_heads, self.value_channels//self.num_heads, n_l) |
| 734 | l_mask = l_mask.unsqueeze(1) # (b, 1, 1, n_l) |
| 735 | |