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

Class PWAM

model/backbone.py:648–676  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

646
647
648class 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
678class SpatialImageLanguageAttention(nn.Module):
679 def __init__(self, v_in_channels, l_in_channels, key_channels, value_channels, out_channels=None, num_heads=1):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected