| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| try: |
| from spatial_correlation_sampler import SpatialCorrelationSampler |
| except ImportError: |
| SpatialCorrelationSampler = None |
|
|
| class ConvBlock(nn.Module): |
| def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, dilation=1, activation=True): |
| super().__init__() |
| self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride=stride, padding=dilation, dilation=dilation, bias=True) |
| self.activation = nn.LeakyReLU(0.1, inplace=True) if activation else nn.Identity() |
|
|
| def forward(self, x): |
| return self.activation(self.conv(x)) |
|
|
| class FeatureExtractor(nn.Module): |
| def __init__(self): |
| super().__init__() |
| |
| self.conv1_1 = ConvBlock(3, 16, stride=2) |
| self.conv1_2 = ConvBlock(16, 16) |
| |
| self.conv2_1 = ConvBlock(16, 32, stride=2) |
| self.conv2_2 = ConvBlock(32, 32) |
| |
| self.conv3_1 = ConvBlock(32, 64, stride=2) |
| self.conv3_2 = ConvBlock(64, 64) |
| |
| self.conv4_1 = ConvBlock(64, 96, stride=2) |
| self.conv4_2 = ConvBlock(96, 96) |
| |
| self.conv5_1 = ConvBlock(96, 128, stride=2) |
| self.conv5_2 = ConvBlock(128, 128) |
| |
| self.conv6_1 = ConvBlock(128, 192, stride=2) |
| self.conv6_2 = ConvBlock(192, 192) |
|
|
| def forward(self, x): |
| c1 = self.conv1_2(self.conv1_1(x)) |
| c2 = self.conv2_2(self.conv2_1(c1)) |
| c3 = self.conv3_2(self.conv3_1(c2)) |
| c4 = self.conv4_2(self.conv4_1(c3)) |
| c5 = self.conv5_2(self.conv5_1(c4)) |
| c6 = self.conv6_2(self.conv6_1(c5)) |
| return { |
| 'conv2_2': c2, |
| 'conv3_2': c3, |
| 'conv4_2': c4, |
| 'conv5_2': c5, |
| 'conv6_2': c6 |
| } |
|
|
| class ContextNetwork(nn.Module): |
| def __init__(self, in_channels): |
| super().__init__() |
| self.conv1 = ConvBlock(in_channels, 128, dilation=1) |
| self.conv2 = ConvBlock(128, 128, dilation=2) |
| self.conv3 = ConvBlock(128, 128, dilation=4) |
| self.conv4 = ConvBlock(128, 96, dilation=8) |
| self.conv5 = ConvBlock(96, 64, dilation=16) |
| self.conv6 = ConvBlock(64, 32, dilation=1) |
| self.conv7 = ConvBlock(32, 2, dilation=1, activation=False) |
|
|
| def forward(self, x, flow): |
| x_in = torch.cat([x, flow], dim=1) |
| out = self.conv1(x_in) |
| out = self.conv2(out) |
| out = self.conv3(out) |
| out = self.conv4(out) |
| out = self.conv5(out) |
| out = self.conv6(out) |
| delta_flow = self.conv7(out) |
| return flow + delta_flow |
|
|
| class FlowEstimator(nn.Module): |
| def __init__(self, ch_in, ch_out=2): |
| super().__init__() |
| self.conv1 = ConvBlock(ch_in, 128) |
| self.conv2 = ConvBlock(128, 128) |
| self.conv3 = ConvBlock(128, 96) |
| self.conv4 = ConvBlock(96, 64) |
| self.conv5 = ConvBlock(64, 32) |
| self.conv6 = ConvBlock(32, ch_out, activation=False) |
|
|
| def forward(self, x): |
| x = self.conv1(x) |
| x = self.conv2(x) |
| x = self.conv3(x) |
| x = self.conv4(x) |
| features = self.conv5(x) |
| flow = self.conv6(features) |
| return flow, features |
|
|
| class DDFlowNet(nn.Module): |
| def __init__(self, args=None): |
| super().__init__() |
| self.feature_extractor = FeatureExtractor() |
| |
| |
| self.search_range = 4 |
| if SpatialCorrelationSampler is not None: |
| self.corr = SpatialCorrelationSampler(1, 9, 1, 0, 1) |
| else: |
| print("WARNING: SpatialCorrelationSampler not found. DDFlow will likely fail or require slow fallback.") |
| self.corr = None |
|
|
| |
| |
| |
| |
| |
| self.est6 = FlowEstimator(81 + 192) |
| self.est5 = FlowEstimator(81 + 128 + 2) |
| self.est4 = FlowEstimator(81 + 96 + 2) |
| self.est3 = FlowEstimator(81 + 64 + 2) |
| self.est2 = FlowEstimator(81 + 32 + 2) |
| |
| self.context = ContextNetwork(32 + 2) |
|
|
| def warp(self, x, flo): |
| """ |
| warp an image/tensor (im2) back to im1, according to the optical flow |
| x: [B, C, H, W] (im2) |
| flo: [B, 2, H, W] flow |
| """ |
| B, C, H, W = x.size() |
| |
| xx = torch.arange(0, W).view(1, -1).repeat(H, 1) |
| yy = torch.arange(0, H).view(-1, 1).repeat(1, W) |
| xx = xx.view(1, 1, H, W).repeat(B, 1, 1, 1) |
| yy = yy.view(1, 1, H, W).repeat(B, 1, 1, 1) |
| grid = torch.cat((xx, yy), 1).float().to(x.device) |
| vgrid = grid + flo |
| |
| |
| vgrid[:, 0, :, :] = 2.0 * vgrid[:, 0, :, :] / max(W - 1, 1) - 1.0 |
| vgrid[:, 1, :, :] = 2.0 * vgrid[:, 1, :, :] / max(H - 1, 1) - 1.0 |
| |
| vgrid = vgrid.permute(0, 2, 3, 1) |
| output = F.grid_sample(x, vgrid, mode='bilinear', align_corners=True, padding_mode="border") |
| return output |
|
|
| def forward(self, x): |
| |
| img1 = x[:, :3] |
| img2 = x[:, 3:6] |
| |
| f1 = self.feature_extractor(img1) |
| f2 = self.feature_extractor(img2) |
| |
| |
| |
| |
| if self.corr is not None: |
| out_corr6 = self.corr(f1['conv6_2'], f2['conv6_2']) |
| out_corr6 = out_corr6.view(out_corr6.shape[0], -1, out_corr6.shape[3], out_corr6.shape[4]) |
| else: |
| B, _, H, W = f1['conv6_2'].shape |
| out_corr6 = torch.zeros(B, 81, H, W).to(img1.device) |
| |
| inp6 = torch.cat([out_corr6, f1['conv6_2']], dim=1) |
| flow6, _ = self.est6(inp6) |
| |
| |
| flow6_up = F.interpolate(flow6, scale_factor=2, mode='bilinear', align_corners=True) * 2.0 |
| f2_5_warp = self.warp(f2['conv5_2'], flow6_up) |
| if self.corr is not None: |
| out_corr5 = self.corr(f1['conv5_2'], f2_5_warp) |
| out_corr5 = out_corr5.view(out_corr5.shape[0], -1, out_corr5.shape[3], out_corr5.shape[4]) |
| else: |
| B, _, H, W = f1['conv5_2'].shape |
| out_corr5 = torch.zeros(B, 81, H, W).to(img1.device) |
| inp5 = torch.cat([out_corr5, f1['conv5_2'], flow6_up], dim=1) |
| flow5, _ = self.est5(inp5) |
| |
| |
| flow5_up = F.interpolate(flow5, scale_factor=2, mode='bilinear', align_corners=True) * 2.0 |
| f2_4_warp = self.warp(f2['conv4_2'], flow5_up) |
| if self.corr is not None: |
| out_corr4 = self.corr(f1['conv4_2'], f2_4_warp) |
| out_corr4 = out_corr4.view(out_corr4.shape[0], -1, out_corr4.shape[3], out_corr4.shape[4]) |
| else: |
| B, _, H, W = f1['conv4_2'].shape |
| out_corr4 = torch.zeros(B, 81, H, W).to(img1.device) |
| inp4 = torch.cat([out_corr4, f1['conv4_2'], flow5_up], dim=1) |
| flow4, _ = self.est4(inp4) |
|
|
| |
| flow4_up = F.interpolate(flow4, scale_factor=2, mode='bilinear', align_corners=True) * 2.0 |
| f2_3_warp = self.warp(f2['conv3_2'], flow4_up) |
| if self.corr is not None: |
| out_corr3 = self.corr(f1['conv3_2'], f2_3_warp) |
| out_corr3 = out_corr3.view(out_corr3.shape[0], -1, out_corr3.shape[3], out_corr3.shape[4]) |
| else: |
| B, _, H, W = f1['conv3_2'].shape |
| out_corr3 = torch.zeros(B, 81, H, W).to(img1.device) |
| inp3 = torch.cat([out_corr3, f1['conv3_2'], flow4_up], dim=1) |
| flow3, _ = self.est3(inp3) |
| |
| |
| flow3_up = F.interpolate(flow3, scale_factor=2, mode='bilinear', align_corners=True) * 2.0 |
| f2_2_warp = self.warp(f2['conv2_2'], flow3_up) |
| if self.corr is not None: |
| out_corr2 = self.corr(f1['conv2_2'], f2_2_warp) |
| out_corr2 = out_corr2.view(out_corr2.shape[0], -1, out_corr2.shape[3], out_corr2.shape[4]) |
| else: |
| B, _, H, W = f1['conv2_2'].shape |
| out_corr2 = torch.zeros(B, 81, H, W).to(img1.device) |
| inp2 = torch.cat([out_corr2, f1['conv2_2'], flow3_up], dim=1) |
| flow2_raw, feat2 = self.est2(inp2) |
| |
| |
| flow2 = self.context(feat2, flow2_raw) |
| |
| if self.training: |
| return flow2, flow3, flow4, flow5, flow6 |
| else: |
| |
| flow_full = F.interpolate(flow2, scale_factor=4, mode='bilinear', align_corners=True) * 4.0 |
| return flow_full |
|
|