import torch,os,sys import torch.nn as nn import torch.nn.functional as F code_dir = os.path.dirname(os.path.realpath(__file__)) sys.path.append(f'{code_dir}/../') from core.submodule import EdgeNextConvEncoder class DispHead(nn.Module): def __init__(self, input_dim=128, hidden_dim=256, output_dim=1): super(DispHead, self).__init__() self.conv = nn.Sequential( nn.Conv2d(input_dim, input_dim, kernel_size=3, padding=1), nn.ReLU(), EdgeNextConvEncoder(input_dim, expan_ratio=4, kernel_size=7, norm=None), EdgeNextConvEncoder(input_dim, expan_ratio=4, kernel_size=7, norm=None), nn.Conv2d(input_dim, output_dim, 3, padding=1), ) def forward(self, x): return self.conv(x) class BasicMotionEncoder(nn.Module): def __init__(self, args, ngroup=8): super(BasicMotionEncoder, self).__init__() self.args = args cor_planes = args.corr_levels * (2*args.corr_radius + 1) * (ngroup+1) self.convc1 = nn.Conv2d(cor_planes, 256, kernel_size=1, padding=0) self.convc2 = nn.Conv2d(256, 256, kernel_size=3, padding=1) self.convd1 = nn.Conv2d(1, 64, kernel_size=7, padding=3) self.convd2 = nn.Conv2d(64, 64, kernel_size=3, padding=1) self.conv = nn.Conv2d(64+256, args.hidden_dims[0]-1, kernel_size=1, padding=0) def forward(self, disp, corr): cor = F.relu(self.convc1(corr)) cor = F.relu(self.convc2(cor)) disp_ = F.relu(self.convd1(disp)) disp_ = F.relu(self.convd2(disp_)) cor_disp = torch.cat([cor, disp_], dim=1) out = F.relu(self.conv(cor_disp)) return torch.cat([out, disp], dim=1) class RaftConvGRU(nn.Module): def __init__(self, hidden_dim=128, input_dim=256, kernel_size=3): super().__init__() self.convz = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2) self.convr = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2) self.convq = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2) def forward(self, h, x, hx): z = torch.sigmoid(self.convz(hx)) r = torch.sigmoid(self.convr(hx)) q = torch.tanh(self.convq(torch.cat([r*h, x], dim=1))) h = (1-z) * h + z * q return h class SelectiveConvGRU(nn.Module): def __init__(self, hidden_dim=128, input_dim=256, small_kernel_size=1, large_kernel_size=3, patch_size=None): super(SelectiveConvGRU, self).__init__() self.conv0 = nn.Sequential( nn.Conv2d(input_dim, input_dim, kernel_size=3, padding=1), nn.ReLU(), ) self.conv1 = nn.Sequential( nn.Conv2d(input_dim+hidden_dim, input_dim+hidden_dim, kernel_size=3, padding=1), nn.ReLU(), ) self.small_gru = RaftConvGRU(hidden_dim, input_dim, small_kernel_size) self.large_gru = RaftConvGRU(hidden_dim, input_dim, large_kernel_size) def forward(self, att, h, *x): x = torch.cat(x, dim=1) x = self.conv0(x) hx = torch.cat([x, h], dim=1) hx = self.conv1(hx) h = self.small_gru(h, x, hx) * att + self.large_gru(h, x, hx) * (1 - att) return h class BasicSelectiveMultiUpdateBlock(nn.Module): def __init__(self, args, hidden_dim=128, volume_dim=8): super().__init__() self.args = args self.encoder = BasicMotionEncoder(args, volume_dim) self.gru04 = SelectiveConvGRU(hidden_dim, hidden_dim*2) self.disp_head = DispHead(hidden_dim, 256) self.mask = nn.Sequential( nn.Conv2d(hidden_dim, 64, 3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(64, 32, 3, padding=1), nn.ReLU(inplace=True), ) def forward(self, net, inp, corr, disp, att): motion_features = self.encoder(disp, corr) motion_features = torch.cat([inp[0], motion_features], dim=1) net[0] = self.gru04(att[0], net[0], motion_features) delta_disp = self.disp_head(net[0]) mask = .25 * self.mask(net[0]) return net, mask, delta_disp