multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
740d966 verified
Raw
History Blame Contribute Delete
4.2 kB
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