Spaces:
Running on Zero
Running on Zero
| 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 | |