models_animerun / DDFlow /models /DDFlowNet_pt.py
doanh25032004's picture
Upload folder using huggingface_hub
872b0a0 verified
Raw
History Blame Contribute Delete
8.79 kB
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__()
# conv1
self.conv1_1 = ConvBlock(3, 16, stride=2)
self.conv1_2 = ConvBlock(16, 16)
# conv2
self.conv2_1 = ConvBlock(16, 32, stride=2)
self.conv2_2 = ConvBlock(32, 32)
# conv3
self.conv3_1 = ConvBlock(32, 64, stride=2)
self.conv3_2 = ConvBlock(64, 64)
# conv4
self.conv4_1 = ConvBlock(64, 96, stride=2)
self.conv4_2 = ConvBlock(96, 96)
# conv5
self.conv5_1 = ConvBlock(96, 128, stride=2)
self.conv5_2 = ConvBlock(128, 128)
# conv6
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()
# PWC-Net style Cost Volume parameters
self.search_range = 4
if SpatialCorrelationSampler is not None:
self.corr = SpatialCorrelationSampler(1, 9, 1, 0, 1) # kernel=1, max_disp=4 -> patch 9
else:
print("WARNING: SpatialCorrelationSampler not found. DDFlow will likely fail or require slow fallback.")
self.corr = None
# Estimators for levels 6, 5, 4, 3, 2
# Input to estimator: CostVolume + Features + Flow(upsampled)
# Channels:
# cv_ch = (search_range*2+1)**2 = 81
# Level 6: 192 ch. In: 81 + 192 + 0(flow? no flow) = 273
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) # Feat from est2 + flow2
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()
# mesh grid
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
# scale grid to [-1,1]
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):
# x: [B, 6, H, W]
img1 = x[:, :3]
img2 = x[:, 3:6]
f1 = self.feature_extractor(img1)
f2 = self.feature_extractor(img2)
# Level 6
# No initial flow
# Correlation
if self.corr is not None:
out_corr6 = self.corr(f1['conv6_2'], f2['conv6_2']) # [B, 81, H, W]
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)
# Level 5
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)
# Level 4
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)
# Level 3
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)
# Level 2
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)
# Context Network
flow2 = self.context(feat2, flow2_raw)
if self.training:
return flow2, flow3, flow4, flow5, flow6
else:
# Upsample to full res
flow_full = F.interpolate(flow2, scale_factor=4, mode='bilinear', align_corners=True) * 4.0
return flow_full