MCPcopy Create free account
hub / github.com/SooLab/CGFormer / SpatialImageLanguageAttention

Class SpatialImageLanguageAttention

model/backbone.py:678–747  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

676 return mm
677
678class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected