import torch import torch.nn.functional as F from torch import nn import numpy as np class SequenceConv(nn.ModuleList): """Sequence conv module. Args: in_channels (int): input tensor channel. out_channels (int): output tensor channel. kernel_size (int): convolution kernel size. sequence_num (int): sequence length. conv_cfg (dict): convolution config dictionary. norm_cfg (dict): normalization config dictionary. act_cfg (dict): activation config dictionary. """ def __init__(self, in_channels, out_channels, kernel_size, sequence_num): super(SequenceConv, self).__init__() self.in_channels = in_channels self.out_channels = out_channels self.kernel_size = kernel_size self.sequence_num = sequence_num for _ in range(sequence_num): self.append( nn.Sequential( nn.Conv2d(self.in_channels, self.out_channels, self.kernel_size, 1, self.kernel_size // 2, bias=False), nn.BatchNorm2d(self.out_channels), nn.ReLU() ) ) def forward(self, sequence_imgs): """ Args: sequence_imgs (Tensor): TxBxCxHxW Returns: sequence conv output: TxBxCxHxW """ sequence_outs = [] assert sequence_imgs.shape[0] == self.sequence_num for i, sequence_conv in enumerate(self): sequence_out = sequence_conv(sequence_imgs[i, ...]) sequence_out = sequence_out.unsqueeze(0) sequence_outs.append(sequence_out) sequence_outs = torch.cat(sequence_outs, dim=0) # TxBxCxHxW return sequence_outs class MemoryModule(nn.Module): """Memory read module. Args: """ def __init__(self, matmul_norm=False): super(MemoryModule, self).__init__() self.matmul_norm = matmul_norm def forward(self, memory_keys, memory_values, query_key, query_value): """ Memory Module forward. Args: memory_keys (Tensor): memory keys tensor, shape: TxBxCxHxW memory_values (Tensor): memory values tensor, shape: TxBxCxHxW query_key (Tensor): query keys tensor, shape: BxCxHxW query_value (Tensor): query values tensor, shape: BxCxHxW Returns: Concat query and memory tensor. """ sequence_num, batch_size, key_channels, height, width = memory_keys.shape _, _, value_channels, _, _ = memory_values.shape assert query_key.shape[1] == key_channels and query_value.shape[1] == value_channels memory_keys = memory_keys.permute(1, 2, 0, 3, 4).contiguous() # BxCxTxHxW memory_keys = memory_keys.view(batch_size, key_channels, sequence_num * height * width) # BxCxT*H*W query_key = query_key.view(batch_size, key_channels, height * width).permute(0, 2, 1).contiguous() # BxH*WxCk key_attention = torch.bmm(query_key, memory_keys) # BxH*WxT*H*W if self.matmul_norm: key_attention = (key_channels ** -.5) * key_attention key_attention = F.softmax(key_attention, dim=-1) # BxH*WxT*H*W memory_values = memory_values.permute(1, 2, 0, 3, 4).contiguous() # BxCxTxHxW memory_values = memory_values.view(batch_size, value_channels, sequence_num * height * width) memory_values = memory_values.permute(0, 2, 1).contiguous() # BxT*H*WxC memory = torch.bmm(key_attention, memory_values) # BxH*WxC memory = memory.permute(0, 2, 1).contiguous() # BxCxH*W memory = memory.view(batch_size, value_channels, height, width) # BxCxHxW query_memory = torch.cat([query_value, memory], dim=1) return query_memory # # class TMAHead(nn.Module): # """TMAHead decoder for video semantic segmentation.""" # # def __init__(self, sequence_num, key_channels, value_channels, num_classes=2, dropout_ratio=0): # super(TMAHead, self).__init__() # # self.sequence_num = sequence_num # self.memory_key_conv = nn.Sequential( # SequenceConv(self.in_channels, key_channels, 1, sequence_num), # SequenceConv(key_channels, key_channels, 3, sequence_num) # ) # self.memory_value_conv = nn.Sequential( # SequenceConv(self.in_channels, value_channels, 1, sequence_num), # SequenceConv(value_channels, value_channels, 3, sequence_num) # ) # self.query_key_conv = nn.Sequential( # nn.Sequential( # nn.Conv2d(self.in_channels, key_channels, 1, 1, 0, bias=False), # nn.BatchNorm2d(key_channels), # nn.ReLU() # ), # nn.Sequential( # nn.Conv2d(key_channels, key_channels, 3, 1, 1, bias=False), # nn.BatchNorm2d(key_channels), # nn.ReLU() # ), # ) # # self.query_value_conv = nn.Sequential( # nn.Sequential( # nn.Conv2d(self.in_channels, value_channels, 1, 1, 0, bias=False), # nn.BatchNorm2d(value_channels), # nn.ReLU() # ), # nn.Sequential( # nn.Conv2d(value_channels, value_channels, 3, 1, 1, bias=False), # nn.BatchNorm2d(value_channels), # nn.ReLU() # ), # ) # self.memory_module = MemoryModule(matmul_norm=False) # self.bottleneck = nn.Sequential( # nn.Conv2d(value_channels * 2, self.channels, 3, 1, 1, bias=False), # nn.BatchNorm2d(value_channels), # nn.ReLU() # ) # # self.conv_seg = nn.Conv2d(self.channels, num_classes, kernel_size=1) # if dropout_ratio > 0: # self.dropout = nn.Dropout2d(dropout_ratio) # else: # self.dropout = None # # def cls_seg(self, feat): # """Classify each pixel.""" # if self.dropout is not None: # feat = self.dropout(feat) # output = self.conv_seg(feat) # return output # # def forward(self, inputs, sequence_imgs): # """ # Forward fuction. # Args: # inputs (list[Tensor]): backbone multi-level outputs. # sequence_imgs (list[Tensor]): len(sequence_imgs) is equal to batch_size, # each element is a Tensor with shape of TxCxHxW. # # Returns: # decoder logits. # """ # x = inputs # sequence_imgs = [y.unsqueeze(0) for y in sequence_imgs] # T, BxCxHxW # sequence_imgs = torch.cat(sequence_imgs, dim=0) # TxBxCxHxW # sequence_num, batch_size, channels, height, width = sequence_imgs.shape # # assert sequence_num == self.sequence_num # memory_keys = self.memory_key_conv(sequence_imgs) # memory_values = self.memory_value_conv(sequence_imgs) # query_key = self.query_key_conv(x) # BxCxHxW # query_value = self.query_value_conv(x) # BxCxHxW # # # memory read # output = self.memory_module(memory_keys, memory_values, query_key, query_value) # output = self.bottleneck(output) # output = self.cls_seg(output) # # return output