Добавлены пропсы конвейера и стереодвижки, задействованные в прогоне
assets/conveyors (274 МБ) - ленты и угловая секция NVIDIA, на которые ссылается сцена относительным путём. Раньше исключались как перекачиваемые, но без них сцена не композится из коробки. cv/ - код стереодвижков, которые вызывает control_test, без весов: * defom-stereo - рабочий бейзлайн (DEFOM vitl, вход 480, iters 24) * crestereo - второй движок, точнее по габаритам (MAE 23.5 против 32.8 мм) * fast-foundationstereo - проверялся, в бейзлайн не вошёл * circular_section.py - показатель кругового сечения, перенесён в measure_plane.py: выравнивает облако по СОБСТВЕННЫМ главным осям и режет на пяти высотах вдоль каждой. Три самодельные версии (мировые оси, одно сечение) давали хуже; результаты проверки на эталонной геометрии - в circular_section_results.json Веса по-прежнему не в репозитории - источники в MODELS.md. Наборы кадров прежних прогонов (cv/flow_*, 1.26 ГБ) исключены: это выход, а не исходники. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,212 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from core.utils.utils import bilinear_sampler
|
||||
|
||||
try:
|
||||
import corr_sampler
|
||||
except:
|
||||
pass
|
||||
|
||||
try:
|
||||
import alt_cuda_corr
|
||||
except:
|
||||
# alt_cuda_corr is not compiled
|
||||
pass
|
||||
|
||||
class CorrSampler(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, volume, coords, radius):
|
||||
ctx.save_for_backward(volume,coords)
|
||||
ctx.radius = radius
|
||||
corr, = corr_sampler.forward(volume, coords, radius)
|
||||
return corr
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
volume, coords = ctx.saved_tensors
|
||||
grad_output = grad_output.contiguous()
|
||||
grad_volume, = corr_sampler.backward(volume, coords, grad_output, ctx.radius)
|
||||
return grad_volume, None, None
|
||||
|
||||
|
||||
class CorrBlockFast1D:
|
||||
def __init__(self, fmap1, fmap2, num_levels=4, radius=4, **kwargs):
|
||||
self.num_levels = num_levels
|
||||
self.radius = radius
|
||||
self.corr_pyramid = []
|
||||
# all pairs correlation
|
||||
corr = CorrBlockFast1D.corr(fmap1, fmap2)
|
||||
batch, h1, w1, dim, w2 = corr.shape
|
||||
corr = corr.reshape(batch*h1*w1, dim, 1, w2)
|
||||
for i in range(self.num_levels):
|
||||
self.corr_pyramid.append(corr.view(batch, h1, w1, -1, w2//2**i))
|
||||
corr = F.avg_pool2d(corr, [1, 2], stride=[1, 2])
|
||||
|
||||
def __call__(self, coords):
|
||||
out_pyramid = []
|
||||
bz, _, ht, wd = coords.shape
|
||||
coords = coords[:, [0]]
|
||||
for i in range(self.num_levels):
|
||||
corr = CorrSampler.apply(self.corr_pyramid[i].squeeze(3), coords/2**i, self.radius)
|
||||
out_pyramid.append(corr.view(bz, -1, ht, wd))
|
||||
return torch.cat(out_pyramid, dim=1)
|
||||
|
||||
@staticmethod
|
||||
def corr(fmap1, fmap2):
|
||||
B, D, H, W1 = fmap1.shape
|
||||
_, _, _, W2 = fmap2.shape
|
||||
fmap1 = fmap1.view(B, D, H, W1)
|
||||
fmap2 = fmap2.view(B, D, H, W2)
|
||||
corr = torch.einsum('aijk,aijh->ajkh', fmap1, fmap2)
|
||||
corr = corr.reshape(B, H, W1, 1, W2).contiguous()
|
||||
return corr / torch.sqrt(torch.tensor(D).float())
|
||||
|
||||
|
||||
class PytorchAlternateCorrBlock1D:
|
||||
def __init__(self, fmap1, fmap2, num_levels=4, radius=4, **kwargs):
|
||||
self.num_levels = num_levels
|
||||
self.radius = radius
|
||||
self.corr_pyramid = []
|
||||
self.fmap1 = fmap1
|
||||
self.fmap2 = fmap2
|
||||
|
||||
def corr(self, fmap1, fmap2, coords):
|
||||
B, D, H, W = fmap2.shape
|
||||
# map grid coordinates to [-1,1]
|
||||
xgrid, ygrid = coords.split([1,1], dim=-1)
|
||||
xgrid = 2*xgrid/(W-1) - 1
|
||||
ygrid = 2*ygrid/(H-1) - 1
|
||||
|
||||
grid = torch.cat([xgrid, ygrid], dim=-1)
|
||||
output_corr = []
|
||||
for grid_slice in grid.unbind(3):
|
||||
fmapw_mini = F.grid_sample(fmap2, grid_slice, align_corners=True)
|
||||
corr = torch.sum(fmapw_mini * fmap1, dim=1)
|
||||
output_corr.append(corr)
|
||||
corr = torch.stack(output_corr, dim=1).permute(0,2,3,1)
|
||||
|
||||
return corr / torch.sqrt(torch.tensor(D).float())
|
||||
|
||||
def __call__(self, coords):
|
||||
r = self.radius
|
||||
coords = coords.permute(0, 2, 3, 1)
|
||||
batch, h1, w1, _ = coords.shape
|
||||
fmap1 = self.fmap1
|
||||
fmap2 = self.fmap2
|
||||
out_pyramid = []
|
||||
for i in range(self.num_levels):
|
||||
dx = torch.zeros(1)
|
||||
dy = torch.linspace(-r, r, 2*r+1)
|
||||
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(coords.device)
|
||||
centroid_lvl = coords.reshape(batch, h1, w1, 1, 2).clone()
|
||||
centroid_lvl[..., 0] = centroid_lvl[..., 0] / 2**i
|
||||
coords_lvl = centroid_lvl + delta.view(-1, 2)
|
||||
corr = self.corr(fmap1, fmap2, coords_lvl)
|
||||
fmap2 = F.avg_pool2d(fmap2, [1, 2], stride=[1, 2])
|
||||
out_pyramid.append(corr)
|
||||
out = torch.cat(out_pyramid, dim=-1)
|
||||
return out.permute(0, 3, 1, 2).contiguous().float()
|
||||
|
||||
|
||||
class CorrBlock1D:
|
||||
def __init__(self, fmap1, fmap2, coords, num_levels=4, radius=4,
|
||||
scale_list=[0.25, 0.5, 2.0, 4.0], scale_corr_radius=4):
|
||||
self.num_levels = num_levels
|
||||
self.radius = radius
|
||||
self.scale_list = scale_list
|
||||
self.scale_corr_radius = scale_corr_radius
|
||||
self.corr_pyramid = []
|
||||
self.coords_pyramid = []
|
||||
dx = torch.linspace(-radius, radius, 2*radius+1)
|
||||
self.dx = dx[:, None].to(coords.device)
|
||||
|
||||
sdx = torch.linspace(-scale_corr_radius, scale_corr_radius, 2*scale_corr_radius+1)
|
||||
self.sdx = sdx[:, None].to(coords.device)
|
||||
|
||||
# all pairs correlation
|
||||
corr = CorrBlock1D.corr(fmap1, fmap2)
|
||||
|
||||
batch, h1, w1, _, w2 = corr.shape
|
||||
self.batch = batch
|
||||
self.h1 = h1
|
||||
self.w1 = w1
|
||||
self.w2 = w2
|
||||
corr = corr.reshape(batch*h1*w1, 1, 1, w2)
|
||||
self.coords = coords.reshape(batch*h1*w1, 1, 1, 1)
|
||||
|
||||
self.corr_pyramid.append(corr)
|
||||
for i in range(1, self.num_levels):
|
||||
corr = F.avg_pool2d(corr, [1, 2], stride=[1, 2])
|
||||
self.corr_pyramid.append(corr)
|
||||
|
||||
def __call__(self, disp, scaling=False):
|
||||
batch, _, h1, w1 = disp.shape
|
||||
|
||||
disp = disp.reshape(self.batch*self.h1*self.w1, 1, 1, 1)
|
||||
out_pyramid = []
|
||||
|
||||
if scaling:
|
||||
corr = self.corr_pyramid[0]
|
||||
for scale in self.scale_list:
|
||||
x0 = self.sdx + self.coords - scale * disp
|
||||
y0 = torch.zeros_like(x0)
|
||||
coords_lvl = torch.cat([x0, y0], dim=-1)
|
||||
corr_s = bilinear_sampler(corr, coords_lvl)
|
||||
corr_s = corr_s.view(self.batch, self.h1, self.w1, -1)
|
||||
out_pyramid.append(corr_s)
|
||||
else:
|
||||
coords = self.coords - disp
|
||||
for i in range(self.num_levels):
|
||||
corr = self.corr_pyramid[i]
|
||||
x0 = self.dx + coords / 2**i
|
||||
y0 = torch.zeros_like(x0)
|
||||
coords_lvl = torch.cat([x0, y0], dim=-1)
|
||||
corr_s = bilinear_sampler(corr, coords_lvl)
|
||||
corr_s = corr_s.view(self.batch, self.h1, self.w1, -1)
|
||||
out_pyramid.append(corr_s)
|
||||
|
||||
out = torch.cat(out_pyramid, dim=-1)
|
||||
return out.permute(0, 3, 1, 2).contiguous().float()
|
||||
|
||||
@staticmethod
|
||||
def corr(fmap1, fmap2):
|
||||
B, D, H, W1 = fmap1.shape
|
||||
_, _, _, W2 = fmap2.shape
|
||||
fmap1 = fmap1.view(B, D, H, W1)
|
||||
fmap2 = fmap2.view(B, D, H, W2)
|
||||
corr = torch.einsum('aijk,aijh->ajkh', fmap1, fmap2)
|
||||
corr = corr.reshape(B, H, W1, 1, W2).contiguous()
|
||||
return corr / torch.sqrt(torch.tensor(D).float())
|
||||
|
||||
|
||||
class AlternateCorrBlock:
|
||||
def __init__(self, fmap1, fmap2, num_levels=4, radius=4, **kwargs):
|
||||
raise NotImplementedError
|
||||
self.num_levels = num_levels
|
||||
self.radius = radius
|
||||
|
||||
self.pyramid = [(fmap1, fmap2)]
|
||||
for i in range(1, self.num_levels):
|
||||
fmap1 = F.avg_pool2d(fmap1, 2, stride=2)
|
||||
fmap2 = F.avg_pool2d(fmap2, 2, stride=2)
|
||||
self.pyramid.append((fmap1, fmap2))
|
||||
|
||||
def __call__(self, coords):
|
||||
coords = coords.permute(0, 2, 3, 1)
|
||||
B, H, W, _ = coords.shape
|
||||
dim = self.pyramid[0][0].shape[1]
|
||||
|
||||
corr_list = []
|
||||
for i in range(self.num_levels):
|
||||
r = self.radius
|
||||
fmap1_i = self.pyramid[0][0].permute(0, 2, 3, 1).contiguous()
|
||||
fmap2_i = self.pyramid[i][1].permute(0, 2, 3, 1).contiguous()
|
||||
|
||||
coords_i = (coords / 2**i).reshape(B, 1, H, W, 2).contiguous()
|
||||
corr, = alt_cuda_corr.forward(fmap1_i, fmap2_i, coords_i, r)
|
||||
corr_list.append(corr.squeeze(1))
|
||||
|
||||
corr = torch.stack(corr_list, dim=1)
|
||||
corr = corr.reshape(B, -1, H, W)
|
||||
return corr / torch.sqrt(torch.tensor(dim).float())
|
||||
@@ -0,0 +1,142 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from core.update import BasicMultiUpdateBlock, ScaleBasicMultiUpdateBlock
|
||||
from core.extractor import BasicEncoder, MultiBasicEncoder, ResidualBlock, DefomEncoder
|
||||
from core.corr import CorrBlock1D, PytorchAlternateCorrBlock1D, CorrBlockFast1D, AlternateCorrBlock
|
||||
from core.utils.utils import coords_grid, upflow, get_danv2_io_size
|
||||
|
||||
|
||||
try:
|
||||
autocast = torch.cuda.amp.autocast
|
||||
except:
|
||||
# dummy autocast for PyTorch < 1.6
|
||||
class autocast:
|
||||
def __init__(self, enabled):
|
||||
pass
|
||||
def __enter__(self):
|
||||
pass
|
||||
def __exit__(self, *args):
|
||||
pass
|
||||
|
||||
|
||||
class DEFOMStereo(nn.Module):
|
||||
def __init__(self, args):
|
||||
super(DEFOMStereo, self).__init__()
|
||||
self.args = args
|
||||
|
||||
self.register_buffer('mean', torch.tensor([[0.485, 0.456, 0.406]])[..., None, None] * 255)
|
||||
self.register_buffer('std', torch.tensor([[0.229, 0.224, 0.225]])[..., None, None] * 255)
|
||||
|
||||
self.defomencoder = DefomEncoder(args.dinov2_encoder, idepth_scale=args.idepth_scale)
|
||||
|
||||
context_dims = args.hidden_dims
|
||||
|
||||
self.fnet = BasicEncoder(self.defomencoder.out_dim, output_dim=256, norm_fn='instance', downsample=args.n_downsample)
|
||||
|
||||
self.context_zqr_convs = nn.ModuleList([nn.Conv2d(context_dims[i], args.hidden_dims[i]*3, 3, padding=3//2) for i in range(self.args.n_gru_layers)])
|
||||
|
||||
self.update_block = BasicMultiUpdateBlock(self.args, hidden_dims=args.hidden_dims)
|
||||
self.scale_update_block = ScaleBasicMultiUpdateBlock(self.args, hidden_dims=args.hidden_dims)
|
||||
|
||||
self.cnet = MultiBasicEncoder(self.defomencoder.out_dim, output_dim=[args.hidden_dims, context_dims],
|
||||
norm_fn=args.context_norm, downsample=args.n_downsample)
|
||||
|
||||
def freeze_bn(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.BatchNorm2d):
|
||||
m.eval()
|
||||
|
||||
def initialize_coords(self, img):
|
||||
""" Disparity is represented as difference between two vertical coordinate grids disp
|
||||
= coords0[:, :1] - coords1[:, :1] """
|
||||
N, _, H, W = img.shape
|
||||
|
||||
coords = coords_grid(N, H, W)[:, :1].to(img.device)
|
||||
|
||||
return coords
|
||||
|
||||
def upsample_flow(self, flow, mask):
|
||||
""" Upsample disparity field [H/scale, W/scale, 1] -> [H, W, 1] using convex combination """
|
||||
N, D, H, W = flow.shape
|
||||
factor = 2 ** self.args.n_downsample
|
||||
mask = mask.view(N, 1, 9, factor, factor, H, W)
|
||||
mask = torch.softmax(mask, dim=2)
|
||||
|
||||
up_flow = F.unfold(factor * flow, [3, 3], padding=1)
|
||||
up_flow = up_flow.view(N, D, 9, 1, 1, H, W)
|
||||
|
||||
up_flow = torch.sum(mask * up_flow, dim=2)
|
||||
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
|
||||
return up_flow.reshape(N, D, factor * H, factor * W)
|
||||
|
||||
def forward(self, image1, image2, iters=12, scale_iters=3, test_mode=False):
|
||||
""" Estimate optical flow between pair of frames """
|
||||
|
||||
image1 = ((image1 - self.mean)/self.std).contiguous()
|
||||
image2 = ((image2 - self.mean)/self.std).contiguous()
|
||||
|
||||
bs, _, h, w = image1.shape
|
||||
danv2_io_sizes = get_danv2_io_size(h, w, self.args.n_downsample)
|
||||
|
||||
# run the context network
|
||||
with autocast(enabled=self.args.mixed_precision):
|
||||
d_features, dfeat1, dfeat2, disp = self.defomencoder([image1, image2], danv2_io_sizes)
|
||||
|
||||
cnet_list = self.cnet(image1, d_features)
|
||||
fmap1, fmap2 = self.fnet([image1, image2], [dfeat1, dfeat2])
|
||||
net_list = [torch.tanh(x[0]) for x in cnet_list]
|
||||
inp_list = [torch.relu(x[1]) for x in cnet_list]
|
||||
# Rather than running the GRU's conv layers on the context features multiple times, we do it once at the beginning
|
||||
inp_list = [list(conv(i).split(split_size=conv.out_channels//3, dim=1)) for i, conv in zip(inp_list, self.context_zqr_convs)]
|
||||
|
||||
coords = self.initialize_coords(net_list[0])
|
||||
|
||||
fmap1, fmap2 = fmap1.float(), fmap2.float()
|
||||
disp = disp.float()
|
||||
corr_fn = CorrBlock1D(fmap1, fmap2, coords, radius=self.args.corr_radius, num_levels=self.args.corr_levels,
|
||||
scale_list=self.args.scale_list, scale_corr_radius=self.args.scale_corr_radius)
|
||||
|
||||
disp_predictions = []
|
||||
for itr in range(iters):
|
||||
disp = disp.detach()
|
||||
|
||||
if itr < scale_iters:
|
||||
corr = corr_fn(disp, scaling=True) # index correlation volume
|
||||
with autocast(enabled=self.args.mixed_precision):
|
||||
net_list, up_mask, scale_disp = self.scale_update_block(net_list, inp_list, corr, disp,
|
||||
iter32=self.args.n_gru_layers == 3,
|
||||
iter16=self.args.n_gru_layers >= 2)
|
||||
|
||||
# F(t+1) = \Scale(t) x F(t)
|
||||
disp = scale_disp * disp
|
||||
else:
|
||||
corr = corr_fn(disp, scaling=False) # index correlation volume
|
||||
with autocast(enabled=self.args.mixed_precision):
|
||||
net_list, up_mask, delta_disp = self.update_block(net_list, inp_list, corr, disp,
|
||||
iter32=self.args.n_gru_layers == 3,
|
||||
iter16=self.args.n_gru_layers >= 2)
|
||||
|
||||
# To avoid unstability, we limit the disparity update within the searching range.
|
||||
delta_disp = torch.clip(delta_disp, min=-2**(self.args.corr_levels-1)*self.args.corr_radius,
|
||||
max=2**(self.args.corr_levels-1)*self.args.corr_radius)
|
||||
|
||||
# F(t+1) = F(t) + \Delta(t)
|
||||
disp = disp + delta_disp
|
||||
|
||||
# We do not need to upsample or output intermediate results in test_mode
|
||||
if test_mode and itr < iters - 1:
|
||||
continue
|
||||
|
||||
# upsample predictions
|
||||
if up_mask is None:
|
||||
disp_up = upflow(disp, factor=2 ** self.n_downsample)
|
||||
else:
|
||||
disp_up = self.upsample_flow(disp, up_mask)
|
||||
|
||||
disp_predictions.append(disp_up)
|
||||
|
||||
if test_mode:
|
||||
return disp_up
|
||||
|
||||
return disp_predictions
|
||||
@@ -0,0 +1,388 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from timm.models.layers import DropPath
|
||||
|
||||
from depth_anything_v2.dpt import DepthAnythingV2
|
||||
|
||||
|
||||
class ConvBlock(nn.Module):
|
||||
def __init__(self, in_planes, planes, norm_fn='group', stride=1):
|
||||
super(ConvBlock, self).__init__()
|
||||
|
||||
self.conv = nn.Conv2d(in_planes, planes, kernel_size=3, padding=1, stride=stride)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
num_groups = planes // 8
|
||||
|
||||
if norm_fn == 'group':
|
||||
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
|
||||
elif norm_fn == 'batch':
|
||||
self.norm1 = nn.BatchNorm2d(planes)
|
||||
self.norm2 = nn.BatchNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.BatchNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(planes)
|
||||
self.norm2 = nn.InstanceNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.InstanceNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
self.norm2 = nn.Sequential()
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.Sequential()
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
return self.relu(self.norm1(self.conv(x)))
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, in_planes, planes, norm_fn='group', stride=1):
|
||||
super(ResidualBlock, self).__init__()
|
||||
|
||||
self.conv1 = nn.Conv2d(in_planes, planes, kernel_size=3, padding=1, stride=stride)
|
||||
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
num_groups = planes // 8
|
||||
|
||||
if norm_fn == 'group':
|
||||
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
|
||||
elif norm_fn == 'batch':
|
||||
self.norm1 = nn.BatchNorm2d(planes)
|
||||
self.norm2 = nn.BatchNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.BatchNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(planes)
|
||||
self.norm2 = nn.InstanceNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.InstanceNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
self.norm2 = nn.Sequential()
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm3 = nn.Sequential()
|
||||
|
||||
if stride == 1 and in_planes == planes:
|
||||
self.downsample = None
|
||||
|
||||
else:
|
||||
self.downsample = nn.Sequential(
|
||||
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3)
|
||||
|
||||
def forward(self, x):
|
||||
y = x
|
||||
y = self.conv1(y)
|
||||
y = self.norm1(y)
|
||||
y = self.relu(y)
|
||||
y = self.conv2(y)
|
||||
y = self.norm2(y)
|
||||
y = self.relu(y)
|
||||
|
||||
if self.downsample is not None:
|
||||
x = self.downsample(x)
|
||||
|
||||
return self.relu(x+y)
|
||||
|
||||
|
||||
class BottleneckBlock(nn.Module):
|
||||
def __init__(self, in_planes, planes, norm_fn='group', stride=1, ratio=4):
|
||||
super(BottleneckBlock, self).__init__()
|
||||
|
||||
self.conv1 = nn.Conv2d(in_planes, planes // ratio, kernel_size=1, padding=0)
|
||||
self.conv2 = nn.Conv2d(planes // ratio, planes // ratio, kernel_size=3, padding=1, stride=stride)
|
||||
self.conv3 = nn.Conv2d(planes // ratio, planes, kernel_size=1, padding=0)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
num_groups = planes // 8
|
||||
|
||||
if norm_fn == 'group':
|
||||
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // ratio)
|
||||
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // ratio)
|
||||
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
||||
|
||||
elif norm_fn == 'batch':
|
||||
self.norm1 = nn.BatchNorm2d(planes // ratio)
|
||||
self.norm2 = nn.BatchNorm2d(planes // ratio)
|
||||
self.norm3 = nn.BatchNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm4 = nn.BatchNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(planes // ratio)
|
||||
self.norm2 = nn.InstanceNorm2d(planes // ratio)
|
||||
self.norm3 = nn.InstanceNorm2d(planes)
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm4 = nn.InstanceNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
self.norm2 = nn.Sequential()
|
||||
self.norm3 = nn.Sequential()
|
||||
if not (stride == 1 and in_planes == planes):
|
||||
self.norm4 = nn.Sequential()
|
||||
|
||||
if stride == 1 and in_planes == planes:
|
||||
self.downsample = None
|
||||
|
||||
else:
|
||||
self.downsample = nn.Sequential(
|
||||
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4)
|
||||
|
||||
def forward(self, x):
|
||||
y = x
|
||||
y = self.relu(self.norm1(self.conv1(y)))
|
||||
y = self.relu(self.norm2(self.conv2(y)))
|
||||
y = self.relu(self.norm3(self.conv3(y)))
|
||||
|
||||
if self.downsample is not None:
|
||||
x = self.downsample(x)
|
||||
|
||||
return self.relu(x + y)
|
||||
|
||||
|
||||
class BasicEncoder(nn.Module):
|
||||
def __init__(self, d_dim, output_dim=128, norm_fn='batch', downsample=3):
|
||||
super(BasicEncoder, self).__init__()
|
||||
self.norm_fn = norm_fn
|
||||
self.downsample = downsample
|
||||
|
||||
if self.norm_fn == 'group':
|
||||
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=64)
|
||||
|
||||
elif self.norm_fn == 'batch':
|
||||
self.norm1 = nn.BatchNorm2d(64)
|
||||
|
||||
elif self.norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(64)
|
||||
|
||||
elif self.norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
|
||||
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=1 + (downsample > 2), padding=3)
|
||||
self.relu1 = nn.ReLU(inplace=True)
|
||||
|
||||
self.in_planes = 64
|
||||
self.layer1 = self._make_layer(64, stride=1)
|
||||
self.layer2 = self._make_layer(96, stride=1 + (downsample > 1))
|
||||
self.layer3 = self._make_layer(128, stride=1 + (downsample > 0))
|
||||
|
||||
# depth feat convolution
|
||||
self.convd = ConvBlock(d_dim, 128, self.norm_fn)
|
||||
|
||||
# output convolution
|
||||
self.conv2 = nn.Conv2d(128, output_dim, kernel_size=1)
|
||||
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
|
||||
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
|
||||
if m.weight is not None:
|
||||
nn.init.constant_(m.weight, 1)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def _make_layer(self, dim, stride=1):
|
||||
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride)
|
||||
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1)
|
||||
layers = (layer1, layer2)
|
||||
|
||||
self.in_planes = dim
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x, dfeats):
|
||||
|
||||
# if input is list, combine batch dimension
|
||||
is_list = isinstance(x, tuple) or isinstance(x, list)
|
||||
if is_list:
|
||||
batch_dim = x[0].shape[0]
|
||||
x = torch.cat(x, dim=0)
|
||||
|
||||
is_list = isinstance(dfeats, tuple) or isinstance(dfeats, list)
|
||||
if is_list:
|
||||
batch_dim = dfeats[0].shape[0]
|
||||
dfeats = torch.cat(dfeats, dim=0)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = self.relu1(x)
|
||||
|
||||
x = self.layer1(x)
|
||||
x = self.layer2(x)
|
||||
x = self.layer3(x)
|
||||
|
||||
x = x + self.convd(dfeats)
|
||||
|
||||
x = self.conv2(x)
|
||||
|
||||
if is_list:
|
||||
x = x.split(split_size=batch_dim, dim=0)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class MultiBasicEncoder(nn.Module):
|
||||
def __init__(self, d_dim, output_dim=[128, 128, 128], norm_fn='batch', downsample=3, drop_path_rate=0.2):
|
||||
super(MultiBasicEncoder, self).__init__()
|
||||
self.d_dim = d_dim
|
||||
self.norm_fn = norm_fn
|
||||
self.downsample = downsample
|
||||
|
||||
if self.norm_fn == 'group':
|
||||
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=64)
|
||||
|
||||
elif self.norm_fn == 'batch':
|
||||
self.norm1 = nn.BatchNorm2d(64)
|
||||
|
||||
elif self.norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(64)
|
||||
|
||||
elif self.norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
|
||||
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=1 + (downsample > 2), padding=3)
|
||||
self.relu1 = nn.ReLU(inplace=True)
|
||||
|
||||
self.in_planes = 64
|
||||
self.layer1 = self._make_layer(64, stride=1)
|
||||
self.layer2 = self._make_layer(96, stride=1 + (downsample > 1))
|
||||
self.layer3 = self._make_layer(128, stride=1 + (downsample > 0))
|
||||
self.layer4 = self._make_layer(128, stride=2)
|
||||
self.layer5 = self._make_layer(128, stride=2)
|
||||
|
||||
self.drop_path = DropPath(drop_path_rate)
|
||||
|
||||
self.conv08 = ConvBlock(d_dim, 128, self.norm_fn)
|
||||
output_list = []
|
||||
for dim in output_dim:
|
||||
conv_out = nn.Sequential(
|
||||
ResidualBlock(128, 128, self.norm_fn, stride=1),
|
||||
nn.Conv2d(128, dim[2], 3, padding=1))
|
||||
output_list.append(conv_out)
|
||||
|
||||
self.outputs08 = nn.ModuleList(output_list)
|
||||
|
||||
self.conv16 = ConvBlock(d_dim, 128, self.norm_fn)
|
||||
output_list = []
|
||||
for dim in output_dim:
|
||||
conv_out = nn.Sequential(
|
||||
ResidualBlock(128, 128, self.norm_fn, stride=1),
|
||||
nn.Conv2d(128, dim[1], 3, padding=1))
|
||||
output_list.append(conv_out)
|
||||
|
||||
self.outputs16 = nn.ModuleList(output_list)
|
||||
|
||||
self.conv32 = ConvBlock(d_dim, 128, self.norm_fn)
|
||||
output_list = []
|
||||
for dim in output_dim:
|
||||
conv_out = nn.Conv2d(128, dim[0], 3, padding=1)
|
||||
output_list.append(conv_out)
|
||||
|
||||
self.outputs32 = nn.ModuleList(output_list)
|
||||
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
|
||||
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
|
||||
if m.weight is not None:
|
||||
nn.init.constant_(m.weight, 1)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def _make_layer(self, dim, stride=1):
|
||||
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride)
|
||||
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1)
|
||||
layers = (layer1, layer2)
|
||||
|
||||
self.in_planes = dim
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x, d_feats, num_layers=3):
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = self.relu1(x)
|
||||
|
||||
x = self.layer1(x)
|
||||
x = self.layer2(x)
|
||||
x = self.layer3(x)
|
||||
|
||||
feat = x + self.drop_path(self.conv08(d_feats[0]))
|
||||
outputs08 = [f(feat) for f in self.outputs08]
|
||||
if num_layers == 1:
|
||||
return (outputs08,)
|
||||
|
||||
y = self.layer4(x)
|
||||
feat = y + self.drop_path(self.conv16(d_feats[1]))
|
||||
outputs16 = [f(feat) for f in self.outputs16]
|
||||
|
||||
if num_layers == 2:
|
||||
return (outputs08, outputs16)
|
||||
|
||||
z = self.layer5(y)
|
||||
feat = z + self.drop_path(self.conv32(d_feats[2]))
|
||||
outputs32 = [f(feat) for f in self.outputs32]
|
||||
|
||||
return (outputs08, outputs16, outputs32)
|
||||
|
||||
|
||||
class DefomEncoder(nn.Module):
|
||||
def __init__(self, dinov2_encoder, pretrained=True, freeze=True, idepth_scale=0.25):
|
||||
super(DefomEncoder, self).__init__()
|
||||
self.dinov2_encoder = dinov2_encoder
|
||||
self.idepth_scale = idepth_scale
|
||||
self.pretrained = pretrained
|
||||
self.freeze = freeze
|
||||
|
||||
model_configs = {
|
||||
'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]},
|
||||
'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]},
|
||||
'vitl': {'encoder': 'vitl', 'features': 256, 'out_channels': [256, 512, 1024, 1024]},
|
||||
'vitg': {'encoder': 'vitg', 'features': 384, 'out_channels': [1536, 1536, 1536, 1536]}
|
||||
}
|
||||
|
||||
self.depth_anything = DepthAnythingV2(**model_configs[self.dinov2_encoder])
|
||||
|
||||
if pretrained and os.path.exists(f'./checkpoints/depth_anything_v2_{dinov2_encoder}.pth'):
|
||||
self.depth_anything.load_state_dict(
|
||||
torch.load(f'./checkpoints/depth_anything_v2_{dinov2_encoder}.pth', map_location='cpu'), strict=False)
|
||||
if freeze:
|
||||
for param in self.depth_anything.pretrained.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.depth_anything.depth_head.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
self.out_dim = model_configs[self.dinov2_encoder]['features']
|
||||
|
||||
def forward(self, x, danv2_io_sizes):
|
||||
|
||||
x = torch.cat(x, dim=0)
|
||||
ih, iw, oh, ow = danv2_io_sizes
|
||||
x = F.interpolate(x, (ih, iw), mode="bilinear", align_corners=True)
|
||||
|
||||
features, left_feat, right_feat, idepth = self.depth_anything(x, oh, ow)
|
||||
|
||||
bs = idepth.shape[0]
|
||||
max_idepth, _ = torch.max(idepth.view(bs, -1), dim=1)
|
||||
max_idepth = max_idepth.detach().view(bs, 1, 1, 1) + 1e-8
|
||||
idepth = idepth / max_idepth * self.idepth_scale * ow + 0.01
|
||||
|
||||
return features, left_feat, right_feat, idepth
|
||||
@@ -0,0 +1,583 @@
|
||||
# Data loading based on https://github.com/NVIDIA/flownet2-pytorch
|
||||
|
||||
import numpy as np
|
||||
from numpy import linalg as LA
|
||||
import torch
|
||||
import torch.utils.data as data
|
||||
import torch.nn.functional as F
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
from pathlib import Path
|
||||
from glob import glob
|
||||
import os.path as osp
|
||||
|
||||
from core.utils import frame_utils
|
||||
from core.utils.augmentor import DispAugmentor, SparseDispAugmentor
|
||||
|
||||
|
||||
class StereoDataset(data.Dataset):
|
||||
def __init__(self, aug_params=None, sparse=False, reader=None, is_eval=False, is_test=False):
|
||||
self.augmentor = None
|
||||
self.sparse = sparse
|
||||
if aug_params is not None and "crop_size" in aug_params:
|
||||
if sparse:
|
||||
self.augmentor = SparseDispAugmentor(**aug_params)
|
||||
else:
|
||||
self.augmentor = DispAugmentor(**aug_params)
|
||||
|
||||
if reader is None:
|
||||
self.disparity_reader = frame_utils.read_gen
|
||||
else:
|
||||
self.disparity_reader = reader
|
||||
|
||||
self.is_eval = is_eval
|
||||
self.is_test = is_test
|
||||
self.init_seed = False
|
||||
self.disparity_list = []
|
||||
self.image_list = []
|
||||
|
||||
# number of copies of the datasets
|
||||
self.v = 1
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
if self.is_test:
|
||||
img1 = frame_utils.read_gen(self.image_list[index][0])
|
||||
img2 = frame_utils.read_gen(self.image_list[index][1])
|
||||
img1 = np.array(img1).astype(np.uint8)
|
||||
img2 = np.array(img2).astype(np.uint8)
|
||||
if len(img1.shape) == 2:
|
||||
img1 = np.tile(img1[..., None], (1, 1, 3))
|
||||
img2 = np.tile(img2[..., None], (1, 1, 3))
|
||||
else:
|
||||
img1 = img1[..., :3]
|
||||
img2 = img2[..., :3]
|
||||
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
|
||||
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
|
||||
return img1, img2, self.image_list[index][0]
|
||||
|
||||
if not self.init_seed:
|
||||
worker_info = torch.utils.data.get_worker_info()
|
||||
if worker_info is not None:
|
||||
torch.manual_seed(worker_info.id)
|
||||
np.random.seed(worker_info.id)
|
||||
random.seed(worker_info.id)
|
||||
self.init_seed = True
|
||||
|
||||
index = index % (len(self.image_list)*self.v)
|
||||
index = index % len(self.image_list)
|
||||
|
||||
if not self.is_eval and len(self.disparity_list[index]) > 1 and np.random.rand() > 0.5:
|
||||
disp = self.disparity_reader(self.disparity_list[index][1])
|
||||
if isinstance(disp, tuple):
|
||||
disp, valid = disp
|
||||
else:
|
||||
valid = disp < 1024
|
||||
img1 = frame_utils.read_gen(self.image_list[index][1])
|
||||
img2 = frame_utils.read_gen(self.image_list[index][0])
|
||||
|
||||
img1 = np.array(img1).astype(np.uint8)[:, ::-1]
|
||||
img2 = np.array(img2).astype(np.uint8)[:, ::-1]
|
||||
disp = np.array(disp).astype(np.float32)[:, ::-1]
|
||||
valid = np.array(valid).astype(np.bool_)[:, ::-1]
|
||||
|
||||
else:
|
||||
disp = self.disparity_reader(self.disparity_list[index][0])
|
||||
if isinstance(disp, tuple):
|
||||
disp, valid = disp
|
||||
else:
|
||||
valid = disp < 1024
|
||||
|
||||
img1 = frame_utils.read_gen(self.image_list[index][0])
|
||||
img2 = frame_utils.read_gen(self.image_list[index][1])
|
||||
|
||||
img1 = np.array(img1).astype(np.uint8)
|
||||
img2 = np.array(img2).astype(np.uint8)
|
||||
disp = np.array(disp).astype(np.float32)
|
||||
valid = np.array(valid).astype(np.bool_)
|
||||
|
||||
# grayscale images
|
||||
if len(img1.shape) == 2:
|
||||
img1 = np.tile(img1[..., None], (1, 1, 3))
|
||||
img2 = np.tile(img2[..., None], (1, 1, 3))
|
||||
else:
|
||||
img1 = img1[..., :3]
|
||||
img2 = img2[..., :3]
|
||||
|
||||
if self.augmentor is not None:
|
||||
if self.sparse:
|
||||
img1, img2, disp, valid = self.augmentor(img1, img2, disp, valid)
|
||||
else:
|
||||
img1, img2, disp = self.augmentor(img1, img2, disp)
|
||||
|
||||
img1 = torch.from_numpy(img1.copy()).permute(2, 0, 1).float()
|
||||
img2 = torch.from_numpy(img2.copy()).permute(2, 0, 1).float()
|
||||
disp = torch.from_numpy(disp[..., np.newaxis].copy()).permute(2, 0, 1).float()
|
||||
if self.sparse:
|
||||
valid = torch.from_numpy(valid[..., np.newaxis].astype(np.bool_).copy()).permute(2, 0, 1)
|
||||
else:
|
||||
valid = disp < 512
|
||||
|
||||
return {"img1": img1, "img2": img2, "disp": disp, "valid": valid, "imageL_file": self.image_list[index][0], "disp_file": self.disparity_list[index][0]}
|
||||
|
||||
def __mul__(self, v):
|
||||
self.v = v
|
||||
return self
|
||||
|
||||
def __len__(self):
|
||||
return len(self.image_list)*self.v
|
||||
|
||||
|
||||
class SceneFlowDatasets(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/SceneFlow/', dstype='frames_cleanpass', things_test=False):
|
||||
super(SceneFlowDatasets, self).__init__(aug_params, is_eval=things_test)
|
||||
self.root = root
|
||||
self.dstype = dstype
|
||||
|
||||
if things_test:
|
||||
self._add_things("TEST")
|
||||
else:
|
||||
self._add_things("TRAIN")
|
||||
self._add_monkaa()
|
||||
self._add_driving()
|
||||
|
||||
def _add_things(self, split='TRAIN'):
|
||||
""" Add FlyingThings3D data """
|
||||
|
||||
original_length = len(self.disparity_list)
|
||||
root = osp.join(self.root, 'FlyingThings3D')
|
||||
left_images = sorted(glob(osp.join(root, self.dstype, split, '*/*/left/*.png')))
|
||||
right_images = [im.replace('left', 'right') for im in left_images]
|
||||
disparity_images = [im.replace(self.dstype, 'disparity').replace('.png', '.pfm') for im in left_images]
|
||||
|
||||
# Choose a random subset of 400 images for validation
|
||||
state = np.random.get_state()
|
||||
np.random.seed(1000)
|
||||
val_idxs = set(np.random.permutation(len(left_images))[:400])
|
||||
np.random.set_state(state)
|
||||
|
||||
for idx, (img1, img2, disp) in enumerate(zip(left_images, right_images, disparity_images)):
|
||||
if (split == 'TEST' and idx in val_idxs) or split == 'TRAIN':
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
logging.info(f"Added {len(self.disparity_list) - original_length} from FlyingThings {self.dstype}")
|
||||
|
||||
def _add_monkaa(self):
|
||||
""" Add FlyingThings3D data """
|
||||
|
||||
original_length = len(self.disparity_list)
|
||||
root = osp.join(self.root, 'Monkaa')
|
||||
left_images = sorted(glob(osp.join(root, self.dstype, '*/left/*.png')) )
|
||||
right_images = [image_file.replace('left', 'right') for image_file in left_images ]
|
||||
disparity_images = [im.replace(self.dstype, 'disparity').replace('.png', '.pfm') for im in left_images ]
|
||||
|
||||
for img1, img2, disp in zip(left_images, right_images, disparity_images):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
logging.info(f"Added {len(self.disparity_list) - original_length} from Monkaa {self.dstype}")
|
||||
|
||||
def _add_driving(self):
|
||||
""" Add FlyingThings3D data """
|
||||
|
||||
original_length = len(self.disparity_list)
|
||||
root = osp.join(self.root, 'Driving')
|
||||
left_images = sorted(glob(osp.join(root, self.dstype, '*/*/*/left/*.png')) )
|
||||
right_images = [image_file.replace('left', 'right') for image_file in left_images ]
|
||||
disparity_images = [im.replace(self.dstype, 'disparity').replace('.png', '.pfm') for im in left_images ]
|
||||
|
||||
for img1, img2, disp in zip(left_images, right_images, disparity_images):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
logging.info(f"Added {len(self.disparity_list) - original_length} from Driving {self.dstype}")
|
||||
|
||||
|
||||
class ETH3D(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/ETH3D', split='training', is_eval=False, is_test=False):
|
||||
super(ETH3D, self).__init__(aug_params, sparse=True, is_eval=is_eval, is_test=is_test)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, f'two_view_{split}/*/im0.png')))
|
||||
image2_list = sorted(glob(osp.join(root, f'two_view_{split}/*/im1.png')))
|
||||
disp_list = sorted(glob(osp.join(root, 'two_view_training_gt/*/disp0GT.pfm'))) if split == 'training'\
|
||||
else [osp.join(root, 'two_view_training_gt/playground_1l/disp0GT.pfm')]*len(image1_list)
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp]]
|
||||
|
||||
|
||||
class KITTI(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/KITTI', split='15', image_set='training', is_eval=False, is_test=False):
|
||||
super(KITTI, self).__init__(aug_params, sparse=True, reader=frame_utils.readDispKITTI, is_eval=is_eval, is_test=is_test)
|
||||
assert split in ["12", "15"]
|
||||
root = root + split
|
||||
assert os.path.exists(root)
|
||||
|
||||
if split == '15':
|
||||
image1_list = sorted(glob(os.path.join(root, image_set, 'image_2/*_10.png')))
|
||||
image2_list = sorted(glob(os.path.join(root, image_set, 'image_3/*_10.png')))
|
||||
disp_list = sorted(
|
||||
glob(os.path.join(root, 'training', 'disp_occ_0/*_10.png'))) if image_set == 'training' else [osp.join(
|
||||
root, 'training/disp_occ_0/000085_10.png')]*len(image1_list)
|
||||
else:
|
||||
image1_list = sorted(glob(os.path.join(root, image_set, 'colored_0/*_10.png')))
|
||||
image2_list = sorted(glob(os.path.join(root, image_set, 'colored_1/*_10.png')))
|
||||
disp_list = sorted(
|
||||
glob(os.path.join(root, 'training', 'disp_occ/*_10.png'))) if image_set == 'training' else [osp.join(
|
||||
root, 'training/disp_occ/000085_10.png')] * len(image1_list)
|
||||
|
||||
for idx, (img1, img2, disp) in enumerate(zip(image1_list, image2_list, disp_list)):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp]]
|
||||
|
||||
|
||||
class Middlebury(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/Middlebury', split='F', image_set='training', is_eval=False, is_test=False):
|
||||
super(Middlebury, self).__init__(aug_params, sparse=True, reader=frame_utils.readDispMiddlebury, is_eval=is_eval, is_test=is_test)
|
||||
assert os.path.exists(root)
|
||||
assert split in ["F", "H", "Q", "2005", "2006", "2014", "2021"]
|
||||
assert image_set in ["training", "test"]
|
||||
|
||||
if split == "2005":
|
||||
scenes = list((Path(root) / "2005").glob("*"))
|
||||
for scene in scenes:
|
||||
self.image_list += [[str(scene / "view1.png"), str(scene / "view5.png")]]
|
||||
self.disparity_list += [[str(scene / "disp1.png"), str(scene / "disp5.png")]]
|
||||
for illum in ["1", "2", "3"]:
|
||||
for exp in ["0", "1", "2"]:
|
||||
self.image_list += [[str(scene / f"Illum{illum}/Exp{exp}/view1.png"), str(scene / f"Illum{illum}/Exp{exp}/view5.png")]]
|
||||
self.disparity_list += [[str(scene / "disp1.png"), str(scene / "disp5.png")]]
|
||||
elif split == "2006":
|
||||
scenes = list((Path(root) / "2006").glob("*"))
|
||||
for scene in scenes:
|
||||
self.image_list += [[str(scene / "view1.png"), str(scene / "view5.png")]]
|
||||
self.disparity_list += [[str(scene / "disp1.png"), str(scene / "disp5.png")]]
|
||||
for illum in ["1", "2", "3"]:
|
||||
for exp in ["0", "1", "2"]:
|
||||
self.image_list += [[str(scene / f"Illum{illum}/Exp{exp}/view1.png"), str(scene / f"Illum{illum}/Exp{exp}/view5.png")]]
|
||||
self.disparity_list += [[str(scene / "disp1.png"), str(scene / "disp5.png")]]
|
||||
elif split == "2014": # datasets/Middlebury/2014/Pipes-perfect/im0.png
|
||||
scenes = list((Path(root) / "2014").glob("*"))
|
||||
for scene in scenes:
|
||||
for s in ["E", "L", ""]:
|
||||
self.image_list += [[str(scene / "im0.png"), str(scene / f"im1{s}.png")]]
|
||||
self.disparity_list += [[str(scene / "disp0.pfm"), str(scene / "disp1.pfm")]]
|
||||
elif split == "2021":
|
||||
scenes = list((Path(root) / "2021/data").glob("*"))
|
||||
for scene in scenes:
|
||||
self.image_list += [[str(scene / "im0.png"), str(scene / "im1.png")]]
|
||||
self.disparity_list += [[str(scene / "disp0.pfm"), str(scene / "disp1.pfm")]]
|
||||
for s in ["0", "1", "2", "3"]:
|
||||
if os.path.exists(str(scene / f"ambient/L0/im0e{s}.png")):
|
||||
self.image_list += [[str(scene / f"ambient/L0/im0e{s}.png"), str(scene / f"ambient/L0/im1e{s}.png")]]
|
||||
self.disparity_list += [[str(scene / "disp0.pfm"), str(scene / "disp1.pfm")]]
|
||||
else:
|
||||
if image_set == 'training':
|
||||
lines = list(map(osp.basename, glob(os.path.join(root, "MiddEval3/trainingF/*"))))
|
||||
if is_eval:
|
||||
lines = list(filter(lambda p: any(s in p.split('/') for s in Path(os.path.join(root, "MiddEval3/official_train.txt")).read_text().splitlines()), lines))
|
||||
else:
|
||||
lines = list(map(osp.basename, glob(os.path.join(root, "MiddEval3/testF/*"))))
|
||||
|
||||
image1_list = sorted([os.path.join(root, "MiddEval3", f'{image_set}{split}', f'{name}/im0.png') for name in lines])
|
||||
image2_list = sorted([os.path.join(root, "MiddEval3", f'{image_set}{split}', f'{name}/im1.png') for name in lines])
|
||||
|
||||
disp_list = sorted([os.path.join(root, "MiddEval3", f'training{split}', f'{name}/disp0GT.pfm') for name in lines]) \
|
||||
if image_set == 'training' else [os.path.join(root, "MiddEval3", f'training{split}', 'Adirondack/disp0GT.pfm')]*len(image1_list)
|
||||
|
||||
assert len(image1_list) == len(image2_list) == len(disp_list) > 0, [image1_list, split]
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp]]
|
||||
|
||||
|
||||
class SintelStereo(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/SintelStereo'):
|
||||
super().__init__(aug_params, reader=frame_utils.readDispSintelStereo)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, 'training/*_left/*/frame_*.png')))
|
||||
image2_list = sorted(glob(osp.join(root, 'training/*_right/*/frame_*.png')))
|
||||
disp_list = sorted(glob(osp.join(root, 'training/disparities/*/frame_*.png'))) * 2
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
assert img1.split('/')[-2:] == disp.split('/')[-2:]
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp]]
|
||||
|
||||
|
||||
class FallingThings(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/FallingThings'):
|
||||
super().__init__(aug_params, reader=frame_utils.readDispFallingThings)
|
||||
assert os.path.exists(root)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, 'fat/single/*/*/*.left.jpg'))) + \
|
||||
sorted(glob(osp.join(root, 'fat/mixed/*/*.left.jpg')))
|
||||
image2_list = [e.replace('left.jpg', 'right.jpg') for e in image1_list]
|
||||
disp_list = [e.replace('left.jpg', 'left.depth.png') for e in image1_list]
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
|
||||
|
||||
class TartanAir(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/TartanAir'):
|
||||
super().__init__(aug_params, reader=frame_utils.readDispTartanAir)
|
||||
assert os.path.exists(root)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, '*/*/*/*/image_left/*_left.png')))
|
||||
image2_list = [e.replace('_left', '_right') for e in image1_list]
|
||||
disp_list = [e.replace('image_left', 'depth_left').replace('left.png', 'left_depth.npy') for e in image1_list]
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
|
||||
|
||||
class CarlaHighres(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/HRVS/carla-highres'):
|
||||
super().__init__(aug_params)
|
||||
assert os.path.exists(root)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, 'trainingF/*/im0.png')))
|
||||
image2_list = [e.replace('im0', 'im1') for e in image1_list]
|
||||
disp1_list = [e.replace('im0.png', 'disp0GT.pfm') for e in image1_list]
|
||||
disp2_list = [e.replace('im1.png', 'disp1GT.pfm') for e in image2_list]
|
||||
|
||||
for img1, img2, disp1, disp2 in zip(image1_list, image2_list, disp1_list, disp2_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp1, disp2]]
|
||||
|
||||
|
||||
class InStereo2K(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/InStereo2K', split='training'):
|
||||
super(InStereo2K, self).__init__(aug_params, sparse=True, reader=frame_utils.readDispInStereo2K, is_eval=split!="training")
|
||||
if split == "training":
|
||||
image1_list = sorted(glob(osp.join(root, 'part*/*/left.png')))
|
||||
else:
|
||||
image1_list = sorted(glob(osp.join(root, 'test/*/left.png')))
|
||||
|
||||
image2_list = [e.replace('left', 'right') for e in image1_list]
|
||||
disp_list = [e.replace('left', 'left_disp') for e in image1_list]
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
|
||||
|
||||
class CreStereo(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/CreStereo'):
|
||||
super(CreStereo, self).__init__(aug_params, reader=frame_utils.readDispCreStereo)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, '*/*_left.jpg')))
|
||||
image2_list = [e.replace('left', 'right') for e in image1_list]
|
||||
disp_list = [e.replace('_left.jpg', '_left.disp.png') for e in image1_list]
|
||||
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp, disp.replace('left', 'right')]]
|
||||
|
||||
|
||||
class IRS(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/IRSDataset'):
|
||||
super().__init__(aug_params)
|
||||
image1_list = sorted(glob(osp.join(root, '*/*/l_*.png')))
|
||||
image2_list = sorted(glob(osp.join(root, '*/*/r_*.png')))
|
||||
disp_list = sorted(glob(osp.join(root, '*/*/d_*.pfm')))
|
||||
for img1, img2, disp in zip(image1_list, image2_list, disp_list):
|
||||
assert img1.split('/')[-2] == disp.split('/')[-2]
|
||||
assert img1.split('.')[0].split('_')[-1] == disp.split('.')[0].split('_')[-1]
|
||||
if 'QAOfficeAndSecurityRoom2_Night' in img1: # bad scenes
|
||||
continue
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp]]
|
||||
|
||||
|
||||
class Booster(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/Booster_Dataset', split='train', is_eval=False, is_test=False):
|
||||
super().__init__(aug_params, sparse=True, reader=frame_utils.readDispBooster, is_eval=is_eval, is_test=is_test)
|
||||
assert os.path.exists(root)
|
||||
|
||||
folder_list = sorted(glob(osp.join(root, split+'/balanced/*')))
|
||||
for folder in folder_list:
|
||||
image1_list = sorted(glob(osp.join(folder, 'camera_00/im*.png')))
|
||||
image2_list = sorted(glob(osp.join(folder, 'camera_02/im*.png')))
|
||||
if split=="train":
|
||||
for img1 in image1_list:
|
||||
for img2 in image2_list:
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[osp.join(folder, 'disp_00.npy'), osp.join(folder, 'disp_02.npy')]]
|
||||
else:
|
||||
for img1, img2 in zip(image1_list, image2_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[osp.join(folder, 'disp_00.npy'), osp.join(folder, 'disp_02.npy')]]
|
||||
|
||||
|
||||
class ThreeDKenBurns(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/3dkenburns'):
|
||||
super().__init__(aug_params, reader=frame_utils.readDisp3DKenBurns)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, '*/*l-image.png')))
|
||||
image2_list = sorted(glob(osp.join(root, '*/*r-image.png')))
|
||||
|
||||
disp1_list = sorted(glob(osp.join(root, '*/*l-depth.exr')))
|
||||
disp2_list = sorted(glob(osp.join(root, '*/*r-depth.exr')))
|
||||
|
||||
for img1, img2, disp1, disp2 in zip(image1_list, image2_list, disp1_list, disp2_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp1, disp2]]
|
||||
|
||||
|
||||
class VKITTI2(StereoDataset):
|
||||
def __init__(self, aug_params=None, root='./datasets/VKITTI2'):
|
||||
super().__init__(aug_params, reader=frame_utils.readDispVKITTI2)
|
||||
|
||||
image1_list = sorted(glob(osp.join(root, 'Scene*/*/frames/rgb/Camera_0/rgb_*.jpg')))
|
||||
image2_list = sorted(glob(osp.join(root, 'Scene*/*/frames/rgb/Camera_1/rgb_*.jpg')))
|
||||
|
||||
disp1_list = sorted(glob(osp.join(root, 'Scene*/*/frames/depth/Camera_0/depth_*.png')))
|
||||
disp2_list = sorted(glob(osp.join(root, 'Scene*/*/frames/depth/Camera_1/depth_*.png')))
|
||||
|
||||
for img1, img2, disp1, disp2 in zip(image1_list, image2_list, disp1_list, disp2_list):
|
||||
self.image_list += [[img1, img2]]
|
||||
self.disparity_list += [[disp1, disp2]]
|
||||
|
||||
|
||||
def fetch_dataloader(args):
|
||||
""" Create the data loader for the corresponding trainign set """
|
||||
|
||||
aug_params = {'crop_size': args.image_size, 'min_scale': args.spatial_scale[0],
|
||||
'max_scale': args.spatial_scale[1], 'do_flip': False, 'yjitter': not args.noyjitter}
|
||||
if hasattr(args, "saturation_range") and args.saturation_range is not None:
|
||||
aug_params["saturation_range"] = args.saturation_range
|
||||
if hasattr(args, "img_gamma") and args.img_gamma is not None:
|
||||
aug_params["gamma"] = args.img_gamma
|
||||
if hasattr(args, "do_flip") and args.do_flip is not None:
|
||||
aug_params["do_flip"] = args.do_flip
|
||||
|
||||
assert len(args.train_datasets) == len(args.train_folds)
|
||||
|
||||
train_dataset = None
|
||||
for fold, dataset_name in zip(args.train_folds, args.train_datasets):
|
||||
if dataset_name.startswith("middlebury_"):
|
||||
new_dataset = Middlebury(aug_params, split=dataset_name.replace('middlebury_','')) * fold
|
||||
elif dataset_name == 'sceneflow':
|
||||
clean_dataset = SceneFlowDatasets(aug_params, dstype='frames_cleanpass')
|
||||
final_dataset = SceneFlowDatasets(aug_params, dstype='frames_finalpass')
|
||||
new_dataset = clean_dataset*fold+final_dataset*fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from SceneFlow")
|
||||
elif 'kitti1' in dataset_name:
|
||||
new_dataset = KITTI(aug_params, split=dataset_name[-2:], image_set='training') * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from KITTI"+dataset_name[-2:])
|
||||
elif 'eth3d' in dataset_name:
|
||||
new_dataset = ETH3D(aug_params, split='training') * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from ETH3D")
|
||||
elif dataset_name == 'sintel_stereo':
|
||||
new_dataset = SintelStereo(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Sintel Stereo")
|
||||
elif dataset_name == 'falling_things':
|
||||
new_dataset = FallingThings(aug_params)*fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from FallingThings")
|
||||
elif dataset_name.startswith('tartan_air'):
|
||||
new_dataset = TartanAir(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Tartain Air")
|
||||
elif dataset_name.startswith('carla_highres'):
|
||||
new_dataset = CarlaHighres(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Carla Highres")
|
||||
elif dataset_name.startswith('irs'):
|
||||
new_dataset = IRS(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from IRS")
|
||||
elif dataset_name.startswith('crestereo'):
|
||||
new_dataset = CreStereo(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from CreStereo")
|
||||
elif dataset_name.startswith('instereo2k'):
|
||||
new_dataset = InStereo2K(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from InStereo2K")
|
||||
elif dataset_name.startswith('booster'):
|
||||
new_dataset = Booster(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Booster")
|
||||
elif dataset_name.startswith('3dkenburns'):
|
||||
new_dataset = ThreeDKenBurns(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from 3D Ken Burns")
|
||||
elif dataset_name.startswith('vkitti2'):
|
||||
new_dataset = VKITTI2(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from VKITTI2")
|
||||
|
||||
train_dataset = new_dataset if train_dataset is None else train_dataset + new_dataset
|
||||
|
||||
train_loader = data.DataLoader(train_dataset, batch_size=args.batch_size,
|
||||
pin_memory=True, shuffle=True, num_workers=int(os.environ.get('SLURM_CPUS_PER_TASK', 6))-2, drop_last=True)
|
||||
|
||||
logging.info('Training with %d image pairs' % len(train_dataset))
|
||||
return train_loader
|
||||
|
||||
|
||||
def fetch_dataset(args):
|
||||
""" Create the dataset for the corresponding training set """
|
||||
|
||||
aug_params = {'crop_size': args.image_size, 'min_scale': args.spatial_scale[0],
|
||||
'max_scale': args.spatial_scale[1], 'do_flip': False, 'yjitter': not args.noyjitter}
|
||||
if hasattr(args, "saturation_range") and args.saturation_range is not None:
|
||||
aug_params["saturation_range"] = args.saturation_range
|
||||
if hasattr(args, "img_gamma") and args.img_gamma is not None:
|
||||
aug_params["gamma"] = args.img_gamma
|
||||
if hasattr(args, "do_flip") and args.do_flip is not None:
|
||||
aug_params["do_flip"] = args.do_flip
|
||||
|
||||
assert len(args.train_datasets) == len(args.train_folds)
|
||||
|
||||
train_dataset = None
|
||||
for fold, dataset_name in zip(args.train_folds, args.train_datasets):
|
||||
if dataset_name.startswith("middlebury_"):
|
||||
new_dataset = Middlebury(aug_params, split=dataset_name.replace('middlebury_', '')) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from {dataset_name}")
|
||||
elif 'eth3d' in dataset_name:
|
||||
new_dataset = ETH3D(aug_params, split='training') * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from ETH3D")
|
||||
elif 'kitti1' in dataset_name:
|
||||
new_dataset = KITTI(aug_params, split=dataset_name[-2:], image_set='training') * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from KITTI"+dataset_name[-2:])
|
||||
elif dataset_name == 'sceneflow':
|
||||
clean_dataset = SceneFlowDatasets(aug_params, dstype='frames_cleanpass')
|
||||
final_dataset = SceneFlowDatasets(aug_params, dstype='frames_finalpass')
|
||||
new_dataset = clean_dataset*fold+final_dataset*fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from SceneFlow")
|
||||
elif dataset_name == 'sintel_stereo':
|
||||
new_dataset = SintelStereo(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Sintel Stereo")
|
||||
elif dataset_name == 'falling_things':
|
||||
new_dataset = FallingThings(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from FallingThings")
|
||||
elif dataset_name.startswith('tartan_air'):
|
||||
new_dataset = TartanAir(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Tartain Air")
|
||||
elif dataset_name.startswith('carla_highres'):
|
||||
new_dataset = CarlaHighres(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Carla Highres")
|
||||
elif dataset_name.startswith('irs'):
|
||||
new_dataset = IRS(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from IRS")
|
||||
elif dataset_name.startswith('crestereo'):
|
||||
new_dataset = CreStereo(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from CreStereo")
|
||||
elif dataset_name.startswith('instereo2k'):
|
||||
new_dataset = InStereo2K(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from InStereo2K")
|
||||
elif dataset_name.startswith('booster'):
|
||||
new_dataset = Booster(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from Booster")
|
||||
elif dataset_name.startswith('3dkenburns'):
|
||||
new_dataset = ThreeDKenBurns(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from 3D Ken Burns")
|
||||
elif dataset_name.startswith('vkitti2'):
|
||||
new_dataset = VKITTI2(aug_params) * fold
|
||||
logging.info(f"Adding {len(new_dataset)} samples from VKITTI2")
|
||||
|
||||
train_dataset = new_dataset if train_dataset is None else train_dataset + new_dataset
|
||||
|
||||
logging.info('Training with %d image pairs' % len(train_dataset))
|
||||
return train_dataset
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from opt_einsum import contract
|
||||
|
||||
|
||||
class DispHead(nn.Module):
|
||||
def __init__(self, input_dim=128, hidden_dim=256, output_dim=1):
|
||||
super(DispHead, self).__init__()
|
||||
self.conv1 = nn.Conv2d(input_dim, hidden_dim, 3, padding=1)
|
||||
self.conv2 = nn.Conv2d(hidden_dim, output_dim, 3, padding=1)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv2(self.relu(self.conv1(x)))
|
||||
|
||||
|
||||
class ConvGRU(nn.Module):
|
||||
def __init__(self, hidden_dim, input_dim, kernel_size=3):
|
||||
super(ConvGRU, self).__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, cz, cr, cq, *x_list):
|
||||
x = torch.cat(x_list, dim=1)
|
||||
hx = torch.cat([h, x], dim=1)
|
||||
|
||||
z = torch.sigmoid(self.convz(hx) + cz)
|
||||
r = torch.sigmoid(self.convr(hx) + cr)
|
||||
q = torch.tanh(self.convq(torch.cat([r*h, x], dim=1)) + cq)
|
||||
|
||||
h = (1-z) * h + z * q
|
||||
return h
|
||||
|
||||
|
||||
class SepConvGRU(nn.Module):
|
||||
def __init__(self, hidden_dim=128, input_dim=192+128):
|
||||
super(SepConvGRU, self).__init__()
|
||||
self.convz1 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (1,5), padding=(0,2))
|
||||
self.convr1 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (1,5), padding=(0,2))
|
||||
self.convq1 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (1,5), padding=(0,2))
|
||||
|
||||
self.convz2 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (5,1), padding=(2,0))
|
||||
self.convr2 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (5,1), padding=(2,0))
|
||||
self.convq2 = nn.Conv2d(hidden_dim+input_dim, hidden_dim, (5,1), padding=(2,0))
|
||||
|
||||
def forward(self, h, *x):
|
||||
# horizontal
|
||||
x = torch.cat(x, dim=1)
|
||||
hx = torch.cat([h, x], dim=1)
|
||||
z = torch.sigmoid(self.convz1(hx))
|
||||
r = torch.sigmoid(self.convr1(hx))
|
||||
q = torch.tanh(self.convq1(torch.cat([r*h, x], dim=1)))
|
||||
h = (1-z) * h + z * q
|
||||
|
||||
# vertical
|
||||
hx = torch.cat([h, x], dim=1)
|
||||
z = torch.sigmoid(self.convz2(hx))
|
||||
r = torch.sigmoid(self.convr2(hx))
|
||||
q = torch.tanh(self.convq2(torch.cat([r*h, x], dim=1)))
|
||||
h = (1-z) * h + z * q
|
||||
|
||||
return h
|
||||
|
||||
|
||||
class BasicMotionEncoder(nn.Module):
|
||||
def __init__(self, cor_planes, c1_planes=64, c2_planes=64, f1_planes=64, f2_planes=64, out_planes=128):
|
||||
super(BasicMotionEncoder, self).__init__()
|
||||
|
||||
self.convc1 = nn.Conv2d(cor_planes, c1_planes, 1, padding=0)
|
||||
self.convc2 = nn.Conv2d(c1_planes, c2_planes, 3, padding=1)
|
||||
self.convd1 = nn.Conv2d(1, f1_planes, 7, padding=3)
|
||||
self.convd2 = nn.Conv2d(f1_planes, f2_planes, 3, padding=1)
|
||||
self.conv = nn.Conv2d(c2_planes+f2_planes, out_planes-1, 3, padding=1)
|
||||
|
||||
def forward(self, disp, corr):
|
||||
cor = F.relu(self.convc1(corr))
|
||||
cor = F.relu(self.convc2(cor))
|
||||
dis = F.relu(self.convd1(disp))
|
||||
dis = F.relu(self.convd2(dis))
|
||||
|
||||
cor_dis = torch.cat([cor, dis], dim=1)
|
||||
out = F.relu(self.conv(cor_dis))
|
||||
return torch.cat([out, disp], dim=1)
|
||||
|
||||
|
||||
def pool2x(x):
|
||||
return F.avg_pool2d(x, 3, stride=2, padding=1)
|
||||
|
||||
|
||||
def pool4x(x):
|
||||
return F.avg_pool2d(x, 5, stride=4, padding=1)
|
||||
|
||||
|
||||
def interp(x, dest):
|
||||
interp_args = {'mode': 'bilinear', 'align_corners': True}
|
||||
return F.interpolate(x, dest.shape[2:], **interp_args)
|
||||
|
||||
|
||||
# for RAFT-Stereo
|
||||
class BasicMultiUpdateBlock(nn.Module):
|
||||
def __init__(self, args, hidden_dims=[128, 128, 128]):
|
||||
super().__init__()
|
||||
self.args = args
|
||||
encoder_output_dim = 128
|
||||
cor_planes = args.corr_levels * (2*args.corr_radius + 1)
|
||||
self.encoder = BasicMotionEncoder(cor_planes, out_planes=encoder_output_dim)
|
||||
|
||||
self.gru08 = ConvGRU(hidden_dims[2], encoder_output_dim + hidden_dims[1] * (args.n_gru_layers > 1))
|
||||
self.gru16 = ConvGRU(hidden_dims[1], hidden_dims[0] * (args.n_gru_layers == 3) + hidden_dims[2])
|
||||
self.gru32 = ConvGRU(hidden_dims[0], hidden_dims[1])
|
||||
self.disp_head = DispHead(hidden_dims[2], hidden_dim=256, output_dim=1)
|
||||
|
||||
factor = 2**self.args.n_downsample
|
||||
|
||||
self.mask = nn.Sequential(
|
||||
nn.Conv2d(hidden_dims[2], 256, 3, padding=1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, (factor**2)*9, 1, padding=0))
|
||||
|
||||
def forward(self, net, inp, corr=None, disp=None, iter08=True, iter16=True, iter32=True, update=True):
|
||||
|
||||
if iter32:
|
||||
net[2] = self.gru32(net[2], *(inp[2]), pool2x(net[1]))
|
||||
if iter16:
|
||||
if self.args.n_gru_layers > 2:
|
||||
net[1] = self.gru16(net[1], *(inp[1]), pool2x(net[0]), interp(net[2], net[1]))
|
||||
else:
|
||||
net[1] = self.gru16(net[1], *(inp[1]), pool2x(net[0]))
|
||||
if iter08:
|
||||
motion_features = self.encoder(disp, corr)
|
||||
if self.args.n_gru_layers > 1:
|
||||
net[0] = self.gru08(net[0], *(inp[0]), motion_features, interp(net[1], net[0]))
|
||||
else:
|
||||
net[0] = self.gru08(net[0], *(inp[0]), motion_features)
|
||||
|
||||
if not update:
|
||||
return net
|
||||
|
||||
delta_disp = self.disp_head(net[0])
|
||||
|
||||
# scale mask to balence gradients
|
||||
mask = .25 * self.mask(net[0])
|
||||
return net, mask, delta_disp
|
||||
|
||||
|
||||
class ScaleBasicMultiUpdateBlock(nn.Module):
|
||||
def __init__(self, args, hidden_dims=[128, 128, 128]):
|
||||
super().__init__()
|
||||
self.args = args
|
||||
encoder_output_dim = 128
|
||||
cor_planes = len(args.scale_list) * (2*args.scale_corr_radius + 1)
|
||||
self.encoder = BasicMotionEncoder(cor_planes, out_planes=encoder_output_dim)
|
||||
|
||||
self.gru08 = ConvGRU(hidden_dims[2], encoder_output_dim + hidden_dims[1] * (args.n_gru_layers > 1))
|
||||
self.gru16 = ConvGRU(hidden_dims[1], hidden_dims[0] * (args.n_gru_layers == 3) + hidden_dims[2])
|
||||
self.gru32 = ConvGRU(hidden_dims[0], hidden_dims[1])
|
||||
self.disp_head = DispHead(hidden_dims[2], hidden_dim=256, output_dim=1)
|
||||
|
||||
factor = 2**self.args.n_downsample
|
||||
|
||||
self.mask = nn.Sequential(
|
||||
nn.Conv2d(hidden_dims[2], 256, 3, padding=1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, (factor**2)*9, 1, padding=0))
|
||||
|
||||
def forward(self, net, inp, corr=None, disp=None, iter08=True, iter16=True, iter32=True, update=True):
|
||||
|
||||
if iter32:
|
||||
net[2] = self.gru32(net[2], *(inp[2]), pool2x(net[1]))
|
||||
if iter16:
|
||||
if self.args.n_gru_layers > 2:
|
||||
net[1] = self.gru16(net[1], *(inp[1]), pool2x(net[0]), interp(net[2], net[1]))
|
||||
else:
|
||||
net[1] = self.gru16(net[1], *(inp[1]), pool2x(net[0]))
|
||||
if iter08:
|
||||
motion_features = self.encoder(disp, corr)
|
||||
if self.args.n_gru_layers > 1:
|
||||
net[0] = self.gru08(net[0], *(inp[0]), motion_features, interp(net[1], net[0]))
|
||||
else:
|
||||
net[0] = self.gru08(net[0], *(inp[0]), motion_features)
|
||||
|
||||
if not update:
|
||||
return net
|
||||
|
||||
x_disp = self.disp_head(net[0])
|
||||
scale_disp = F.relu6(torch.exp(.25*x_disp))
|
||||
|
||||
# scale mask to balence gradients
|
||||
mask = .25 * self.mask(net[0])
|
||||
return net, mask, scale_disp
|
||||
@@ -0,0 +1,307 @@
|
||||
import numpy as np
|
||||
import random
|
||||
import warnings
|
||||
import os
|
||||
import time
|
||||
from glob import glob
|
||||
from skimage import color, io
|
||||
from PIL import Image
|
||||
|
||||
import cv2
|
||||
cv2.setNumThreads(0)
|
||||
cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
import torch
|
||||
from torchvision.transforms import ColorJitter, functional, Compose
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def get_middlebury_images():
|
||||
root = "../datasets/Middlebury/MiddEval3"
|
||||
with open(os.path.join(root, "official_train.txt"), 'r') as f:
|
||||
lines = f.read().splitlines()
|
||||
return sorted([os.path.join(root, 'trainingQ', f'{name}/im0.png') for name in lines])
|
||||
|
||||
|
||||
def get_eth3d_images():
|
||||
return sorted(glob('../datasets/ETH3D/two_view_training/*/im0.png'))
|
||||
|
||||
|
||||
def get_kitti_images():
|
||||
return sorted(glob('..datasets/KITTI/training/image_2/*_10.png'))
|
||||
|
||||
|
||||
def transfer_color(image, style_mean, style_stddev):
|
||||
reference_image_lab = color.rgb2lab(image)
|
||||
reference_stddev = np.std(reference_image_lab, axis=(0, 1), keepdims=True)# + 1
|
||||
reference_mean = np.mean(reference_image_lab, axis=(0, 1), keepdims=True)
|
||||
|
||||
reference_image_lab = reference_image_lab - reference_mean
|
||||
lamb = style_stddev/reference_stddev
|
||||
style_image_lab = lamb * reference_image_lab
|
||||
output_image_lab = style_image_lab + style_mean
|
||||
l, a, b = np.split(output_image_lab, 3, axis=2)
|
||||
l = l.clip(0, 100)
|
||||
output_image_lab = np.concatenate((l, a, b), axis=2)
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore", category=UserWarning)
|
||||
output_image_rgb = color.lab2rgb(output_image_lab) * 255
|
||||
return output_image_rgb
|
||||
|
||||
|
||||
class AdjustGamma(object):
|
||||
|
||||
def __init__(self, gamma_min, gamma_max, gain_min=1.0, gain_max=1.0):
|
||||
self.gamma_min, self.gamma_max, self.gain_min, self.gain_max = gamma_min, gamma_max, gain_min, gain_max
|
||||
|
||||
def __call__(self, sample):
|
||||
gain = random.uniform(self.gain_min, self.gain_max)
|
||||
gamma = random.uniform(self.gamma_min, self.gamma_max)
|
||||
return functional.adjust_gamma(sample, gamma, gain)
|
||||
|
||||
def __repr__(self):
|
||||
return f"Adjust Gamma {self.gamma_min}, ({self.gamma_max}) and Gain ({self.gain_min}, {self.gain_max})"
|
||||
|
||||
|
||||
class DispAugmentor:
|
||||
def __init__(self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=True, yjitter=False,
|
||||
saturation_range=[0.6, 1.4], gamma=[1, 1, 1, 1]):
|
||||
|
||||
# spatial augmentation params
|
||||
self.crop_size = crop_size
|
||||
self.min_scale = min_scale
|
||||
self.max_scale = max_scale
|
||||
self.spatial_aug_prob = 1.0
|
||||
self.stretch_prob = 0.8
|
||||
self.max_stretch = 0.2
|
||||
|
||||
# flip augmentation params
|
||||
self.yjitter = yjitter
|
||||
self.do_flip = do_flip
|
||||
self.v_flip_prob = 0.1
|
||||
|
||||
# photometric augmentation params
|
||||
self.photo_aug = Compose([ColorJitter(brightness=0.4, contrast=0.4, saturation=saturation_range, hue=0.5/3.14), AdjustGamma(*gamma)])
|
||||
self.asymmetric_color_aug_prob = 0.2
|
||||
self.eraser_aug_prob = 0.5
|
||||
|
||||
def color_transform(self, img1, img2):
|
||||
""" Photometric augmentation """
|
||||
|
||||
# asymmetric
|
||||
if np.random.rand() < self.asymmetric_color_aug_prob:
|
||||
img1 = np.array(self.photo_aug(Image.fromarray(img1)), dtype=np.uint8)
|
||||
img2 = np.array(self.photo_aug(Image.fromarray(img2)), dtype=np.uint8)
|
||||
|
||||
# symmetric
|
||||
else:
|
||||
image_stack = np.concatenate([img1, img2], axis=0)
|
||||
image_stack = np.array(self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8)
|
||||
img1, img2 = np.split(image_stack, 2, axis=0)
|
||||
|
||||
return img1, img2
|
||||
|
||||
def eraser_transform(self, img1, img2, bounds=[50, 100]):
|
||||
""" Occlusion augmentation """
|
||||
|
||||
ht, wd = img1.shape[:2]
|
||||
if np.random.rand() < self.eraser_aug_prob:
|
||||
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
|
||||
for _ in range(np.random.randint(1, 3)):
|
||||
x0 = np.random.randint(0, wd)
|
||||
y0 = np.random.randint(0, ht)
|
||||
dx = np.random.randint(bounds[0], bounds[1])
|
||||
dy = np.random.randint(bounds[0], bounds[1])
|
||||
img2[y0:y0 + dy, x0:x0 + dx, :] = mean_color
|
||||
|
||||
return img1, img2
|
||||
|
||||
def spatial_transform(self, img1, img2, disp):
|
||||
# randomly sample scale
|
||||
ht, wd = img1.shape[:2]
|
||||
min_scale = np.maximum(
|
||||
(self.crop_size[0] + 8) / float(ht),
|
||||
(self.crop_size[1] + 8) / float(wd))
|
||||
|
||||
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
|
||||
if scale>min_scale:
|
||||
scale = np.random.uniform(min_scale, scale)
|
||||
scale_x = scale
|
||||
scale_y = scale
|
||||
if np.random.rand() < self.stretch_prob:
|
||||
scale_x *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
|
||||
scale_y *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
|
||||
|
||||
scale_x = np.clip(scale_x, min_scale, 2*min_scale)
|
||||
scale_y = np.clip(scale_y, min_scale, 2*min_scale)
|
||||
|
||||
if np.random.rand() < self.spatial_aug_prob or min_scale >= 1.0:
|
||||
# rescale the images
|
||||
img1 = cv2.resize(img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR)
|
||||
img2 = cv2.resize(img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR)
|
||||
disp = cv2.resize(disp, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR)
|
||||
disp = disp * scale_x
|
||||
|
||||
if self.do_flip:
|
||||
if np.random.rand() < self.v_flip_prob and self.do_flip == 'v': # v-flip
|
||||
img1 = img1[::-1, :]
|
||||
img2 = img2[::-1, :]
|
||||
disp = disp[::-1, :]
|
||||
|
||||
if self.yjitter:
|
||||
y0 = np.random.randint(2, img1.shape[0] - self.crop_size[0] - 2)
|
||||
x0 = np.random.randint(0, img1.shape[1] - self.crop_size[1] - 0)
|
||||
|
||||
y1 = y0 + np.random.randint(-2, 2 + 1)
|
||||
y1 = np.clip(y1, 0, img1.shape[0] - self.crop_size[0])
|
||||
img1 = img1[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
img2 = img2[y1:y1 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
disp = disp[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
|
||||
else:
|
||||
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0])
|
||||
x0 = np.random.randint(0, img1.shape[1] - self.crop_size[1])
|
||||
|
||||
img1 = img1[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
img2 = img2[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
disp = disp[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
|
||||
return img1, img2, disp
|
||||
|
||||
def __call__(self, img1, img2, disp):
|
||||
img1, img2 = self.color_transform(img1, img2)
|
||||
img1, img2 = self.eraser_transform(img1, img2)
|
||||
img1, img2, disp = self.spatial_transform(img1, img2, disp)
|
||||
|
||||
img1 = np.ascontiguousarray(img1)
|
||||
img2 = np.ascontiguousarray(img2)
|
||||
disp = np.ascontiguousarray(disp)
|
||||
|
||||
return img1, img2, disp
|
||||
|
||||
|
||||
class SparseDispAugmentor:
|
||||
def __init__(self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=False, yjitter=False,
|
||||
saturation_range=[0.7, 1.3], gamma=[1, 1, 1, 1]):
|
||||
# spatial augmentation params
|
||||
self.crop_size = crop_size
|
||||
self.min_scale = min_scale
|
||||
self.max_scale = max_scale
|
||||
self.spatial_aug_prob = 0.8
|
||||
self.stretch_prob = 0.8
|
||||
self.max_stretch = 0.2
|
||||
|
||||
# flip augmentation params
|
||||
self.do_flip = do_flip
|
||||
self.v_flip_prob = 0.1
|
||||
|
||||
# photometric augmentation params
|
||||
self.photo_aug = Compose(
|
||||
[ColorJitter(brightness=0.3, contrast=0.3, saturation=saturation_range, hue=0.3/3.14),
|
||||
AdjustGamma(*gamma)])
|
||||
self.asymmetric_color_aug_prob = 0.2
|
||||
self.eraser_aug_prob = 0.5
|
||||
|
||||
def color_transform(self, img1, img2):
|
||||
image_stack = np.concatenate([img1, img2], axis=0)
|
||||
image_stack = np.array(self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8)
|
||||
img1, img2 = np.split(image_stack, 2, axis=0)
|
||||
return img1, img2
|
||||
|
||||
def eraser_transform(self, img1, img2):
|
||||
ht, wd = img1.shape[:2]
|
||||
if np.random.rand() < self.eraser_aug_prob:
|
||||
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
|
||||
for _ in range(np.random.randint(1, 3)):
|
||||
x0 = np.random.randint(0, wd)
|
||||
y0 = np.random.randint(0, ht)
|
||||
dx = np.random.randint(50, 100)
|
||||
dy = np.random.randint(50, 100)
|
||||
img2[y0:y0 + dy, x0:x0 + dx, :] = mean_color
|
||||
|
||||
return img1, img2
|
||||
|
||||
def resize_sparse_flow_map(self, disp, valid, fx=1.0, fy=1.0):
|
||||
ht, wd = disp.shape[:2]
|
||||
coords = np.meshgrid(np.arange(wd), np.arange(ht))
|
||||
coords = np.stack(coords, axis=-1)
|
||||
|
||||
coords = coords.reshape(-1, 2).astype(np.float32)
|
||||
disp = disp.reshape(-1).astype(np.float32)
|
||||
valid = valid.reshape(-1).astype(np.float32)
|
||||
|
||||
coords0 = coords[valid >= 1]
|
||||
disp0 = disp[valid >= 1]
|
||||
|
||||
ht1 = int(round(ht * fy))
|
||||
wd1 = int(round(wd * fx))
|
||||
|
||||
coords1 = coords0 * [fx, fy]
|
||||
disp1 = disp0 * fx
|
||||
|
||||
xx = np.round(coords1[:, 0]).astype(np.int32)
|
||||
yy = np.round(coords1[:, 1]).astype(np.int32)
|
||||
|
||||
v = (xx > 0) & (xx < wd1) & (yy > 0) & (yy < ht1)
|
||||
xx = xx[v]
|
||||
yy = yy[v]
|
||||
disp1 = disp1[v]
|
||||
|
||||
disp_img = np.zeros([ht1, wd1], dtype=np.float32)
|
||||
valid_img = np.zeros([ht1, wd1], dtype=np.int32)
|
||||
|
||||
disp_img[yy, xx] = disp1
|
||||
valid_img[yy, xx] = 1
|
||||
|
||||
return disp_img, valid_img
|
||||
|
||||
def spatial_transform(self, img1, img2, disp, valid):
|
||||
# randomly sample scale
|
||||
|
||||
ht, wd = img1.shape[:2]
|
||||
min_scale = np.maximum(
|
||||
(self.crop_size[0] + 1) / float(ht),
|
||||
(self.crop_size[1] + 1) / float(wd))
|
||||
|
||||
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
|
||||
if scale>min_scale:
|
||||
scale = np.random.uniform(min_scale, 2*min_scale)
|
||||
scale_x = scale
|
||||
scale_y = scale
|
||||
|
||||
scale_x = np.clip(scale_x, min_scale, 2*min_scale)
|
||||
scale_y = np.clip(scale_y, min_scale, 2*min_scale)
|
||||
|
||||
if np.random.rand() < self.spatial_aug_prob or min_scale >= 1.0:
|
||||
# rescale the images
|
||||
img1 = cv2.resize(img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR)
|
||||
img2 = cv2.resize(img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR)
|
||||
disp, valid = self.resize_sparse_flow_map(disp, valid, fx=scale_x, fy=scale_y)
|
||||
|
||||
if self.do_flip:
|
||||
if np.random.rand() < self.v_flip_prob and self.do_flip == 'v': # v-flip
|
||||
img1 = img1[::-1, :]
|
||||
img2 = img2[::-1, :]
|
||||
disp = disp[::-1, :]
|
||||
valid = valid[::-1, :]
|
||||
|
||||
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0])
|
||||
x0 = np.random.randint(0, img1.shape[1] - self.crop_size[1])
|
||||
|
||||
img1 = img1[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
img2 = img2[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
disp = disp[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
valid = valid[y0:y0 + self.crop_size[0], x0:x0 + self.crop_size[1]]
|
||||
return img1, img2, disp, valid
|
||||
|
||||
def __call__(self, img1, img2, disp, valid):
|
||||
img1, img2 = self.color_transform(img1, img2)
|
||||
img1, img2 = self.eraser_transform(img1, img2)
|
||||
img1, img2, disp, valid = self.spatial_transform(img1, img2, disp, valid)
|
||||
|
||||
img1 = np.ascontiguousarray(img1)
|
||||
img2 = np.ascontiguousarray(img2)
|
||||
disp = np.ascontiguousarray(disp)
|
||||
valid = np.ascontiguousarray(valid)
|
||||
|
||||
return img1, img2, disp, valid
|
||||
@@ -0,0 +1,105 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
# https://github.com/open-mmlab/mmcv/blob/7540cf73ac7e5d1e14d0ffbd9b6759e83929ecfc/mmcv/runner/dist_utils.py
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
import torch
|
||||
import torch.multiprocessing as mp
|
||||
from torch import distributed as dist
|
||||
|
||||
|
||||
def init_dist(launcher, backend='nccl', **kwargs):
|
||||
if mp.get_start_method(allow_none=True) is None:
|
||||
mp.set_start_method('spawn')
|
||||
if launcher == 'pytorch':
|
||||
_init_dist_pytorch(backend, **kwargs)
|
||||
elif launcher == 'mpi':
|
||||
_init_dist_mpi(backend, **kwargs)
|
||||
elif launcher == 'slurm':
|
||||
_init_dist_slurm(backend, **kwargs)
|
||||
else:
|
||||
raise ValueError(f'Invalid launcher type: {launcher}')
|
||||
|
||||
|
||||
def _init_dist_pytorch(backend, **kwargs):
|
||||
# TODO: use local_rank instead of rank % num_gpus
|
||||
rank = int(os.environ['RANK'])
|
||||
num_gpus = torch.cuda.device_count()
|
||||
torch.cuda.set_device(rank % num_gpus)
|
||||
dist.init_process_group(backend=backend, **kwargs)
|
||||
|
||||
|
||||
def _init_dist_mpi(backend, **kwargs):
|
||||
# TODO: use local_rank instead of rank % num_gpus
|
||||
rank = int(os.environ['OMPI_COMM_WORLD_RANK'])
|
||||
num_gpus = torch.cuda.device_count()
|
||||
torch.cuda.set_device(rank % num_gpus)
|
||||
dist.init_process_group(backend=backend, **kwargs)
|
||||
|
||||
|
||||
def _init_dist_slurm(backend, port=None):
|
||||
"""Initialize slurm distributed training environment.
|
||||
If argument ``port`` is not specified, then the master port will be system
|
||||
environment variable ``MASTER_PORT``. If ``MASTER_PORT`` is not in system
|
||||
environment variable, then a default port ``29500`` will be used.
|
||||
Args:
|
||||
backend (str): Backend of torch.distributed.
|
||||
port (int, optional): Master port. Defaults to None.
|
||||
"""
|
||||
proc_id = int(os.environ['SLURM_PROCID'])
|
||||
ntasks = int(os.environ['SLURM_NTASKS'])
|
||||
node_list = os.environ['SLURM_NODELIST']
|
||||
num_gpus = torch.cuda.device_count()
|
||||
torch.cuda.set_device(proc_id % num_gpus)
|
||||
addr = subprocess.getoutput(
|
||||
f'scontrol show hostname {node_list} | head -n1')
|
||||
# specify master port
|
||||
if port is not None:
|
||||
os.environ['MASTER_PORT'] = str(port)
|
||||
elif 'MASTER_PORT' in os.environ:
|
||||
pass # use MASTER_PORT in the environment variable
|
||||
else:
|
||||
# 29500 is torch.distributed default port
|
||||
os.environ['MASTER_PORT'] = '29500'
|
||||
# use MASTER_ADDR in the environment variable if it already exists
|
||||
if 'MASTER_ADDR' not in os.environ:
|
||||
os.environ['MASTER_ADDR'] = addr
|
||||
os.environ['WORLD_SIZE'] = str(ntasks)
|
||||
os.environ['LOCAL_RANK'] = str(proc_id % num_gpus)
|
||||
os.environ['RANK'] = str(proc_id)
|
||||
dist.init_process_group(backend=backend)
|
||||
|
||||
|
||||
def get_dist_info():
|
||||
# if (TORCH_VERSION != 'parrots'
|
||||
# and digit_version(TORCH_VERSION) < digit_version('1.0')):
|
||||
# initialized = dist._initialized
|
||||
# else:
|
||||
if dist.is_available():
|
||||
initialized = dist.is_initialized()
|
||||
else:
|
||||
initialized = False
|
||||
if initialized:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
else:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
return rank, world_size
|
||||
|
||||
|
||||
# from DETR repo
|
||||
def setup_for_distributed(is_master):
|
||||
"""
|
||||
This function disables printing when not in master process
|
||||
"""
|
||||
import builtins as __builtin__
|
||||
builtin_print = __builtin__.print
|
||||
|
||||
def print(*args, **kwargs):
|
||||
force = kwargs.pop('force', False)
|
||||
if is_master or force:
|
||||
builtin_print(*args, **kwargs)
|
||||
|
||||
__builtin__.print = print
|
||||
@@ -0,0 +1,286 @@
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from os.path import *
|
||||
import re
|
||||
import json
|
||||
import imageio
|
||||
import os
|
||||
import math
|
||||
|
||||
os.environ["OPENCV_IO_ENABLE_OPENEXR"]="1"
|
||||
import cv2
|
||||
cv2.setNumThreads(0)
|
||||
cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
TAG_CHAR = np.array([202021.25], np.float32)
|
||||
|
||||
|
||||
def readFlow(fn):
|
||||
""" Read .flo file in Middlebury format"""
|
||||
# Code adapted from:
|
||||
# http://stackoverflow.com/questions/28013200/reading-middlebury-flow-files-with-python-bytes-array-numpy
|
||||
|
||||
# WARNING: this will work on little-endian architectures (eg Intel x86) only!
|
||||
# print 'fn = %s'%(fn)
|
||||
with open(fn, 'rb') as f:
|
||||
magic = np.fromfile(f, np.float32, count=1)
|
||||
if 202021.25 != magic:
|
||||
print('Magic number incorrect. Invalid .flo file')
|
||||
return None
|
||||
else:
|
||||
w = np.fromfile(f, np.int32, count=1)
|
||||
h = np.fromfile(f, np.int32, count=1)
|
||||
# print 'Reading %d x %d flo file\n' % (w, h)
|
||||
data = np.fromfile(f, np.float32, count=2*int(w)*int(h))
|
||||
# Reshape data into 3D array (columns, rows, bands)
|
||||
# The reshape here is for visualization, the original code is (w,h,2)
|
||||
return np.resize(data, (int(h), int(w), 2))
|
||||
|
||||
|
||||
def readPFM(file):
|
||||
file = open(file, 'rb')
|
||||
|
||||
color = None
|
||||
width = None
|
||||
height = None
|
||||
scale = None
|
||||
endian = None
|
||||
|
||||
header = file.readline().rstrip()
|
||||
if header == b'PF':
|
||||
color = True
|
||||
elif header == b'Pf':
|
||||
color = False
|
||||
else:
|
||||
raise Exception('Not a PFM file.')
|
||||
|
||||
dim_match = re.match(rb'^(\d+)\s(\d+)\s$', file.readline())
|
||||
if dim_match:
|
||||
width, height = map(int, dim_match.groups())
|
||||
else:
|
||||
raise Exception('Malformed PFM header.')
|
||||
|
||||
scale = float(file.readline().rstrip())
|
||||
if scale < 0: # little-endian
|
||||
endian = '<'
|
||||
scale = -scale
|
||||
else:
|
||||
endian = '>' # big-endian
|
||||
|
||||
data = np.fromfile(file, endian + 'f')
|
||||
shape = (height, width, 3) if color else (height, width)
|
||||
|
||||
data = np.reshape(data, shape)
|
||||
data = np.flipud(data)
|
||||
return data
|
||||
|
||||
|
||||
def writePFM(file, array):
|
||||
import os
|
||||
assert type(file) is str and type(array) is np.ndarray and \
|
||||
os.path.splitext(file)[1] == ".pfm"
|
||||
with open(file, 'wb') as f:
|
||||
H, W = array.shape
|
||||
headers = ["Pf\n", f"{W} {H}\n", "-1\n"]
|
||||
for header in headers:
|
||||
f.write(str.encode(header))
|
||||
array = np.flip(array, axis=0).astype(np.float32)
|
||||
f.write(array.tobytes())
|
||||
|
||||
|
||||
def writeFlow(filename,uv,v=None):
|
||||
""" Write optical flow to file.
|
||||
|
||||
If v is None, uv is assumed to contain both u and v channels,
|
||||
stacked in depth.
|
||||
Original code by Deqing Sun, adapted from Daniel Scharstein.
|
||||
"""
|
||||
nBands = 2
|
||||
|
||||
if v is None:
|
||||
assert(uv.ndim == 3)
|
||||
assert(uv.shape[2] == 2)
|
||||
u = uv[:,:,0]
|
||||
v = uv[:,:,1]
|
||||
else:
|
||||
u = uv
|
||||
|
||||
assert(u.shape == v.shape)
|
||||
height,width = u.shape
|
||||
f = open(filename,'wb')
|
||||
# write the header
|
||||
f.write(TAG_CHAR)
|
||||
np.array(width).astype(np.int32).tofile(f)
|
||||
np.array(height).astype(np.int32).tofile(f)
|
||||
# arrange into matrix form
|
||||
tmp = np.zeros((height, width*nBands))
|
||||
tmp[:,np.arange(width)*2] = u
|
||||
tmp[:,np.arange(width)*2 + 1] = v
|
||||
tmp.astype(np.float32).tofile(f)
|
||||
f.close()
|
||||
|
||||
|
||||
def readFlowKITTI(filename):
|
||||
flow = cv2.imread(filename, cv2.IMREAD_ANYDEPTH|cv2.IMREAD_COLOR)
|
||||
flow = flow[:,:,::-1].astype(np.float32)
|
||||
flow, valid = flow[:, :, :2], flow[:, :, 2]
|
||||
flow = (flow - 2**15) / 64.0
|
||||
return flow, valid
|
||||
|
||||
|
||||
def readDispKITTI(filename):
|
||||
disp = cv2.imread(filename, cv2.IMREAD_ANYDEPTH) / 256.0
|
||||
valid = disp > 0.0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispInStereo2K(filename):
|
||||
disp = cv2.imread(filename, cv2.IMREAD_ANYDEPTH) / 100.0
|
||||
valid = disp > 0.0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispVKITTI2(filename):
|
||||
depth = cv2.imread(filename, cv2.IMREAD_ANYCOLOR | cv2.IMREAD_ANYDEPTH).astype(np.float32) / 100.0
|
||||
valid = depth > 0.0
|
||||
baseline = 0.532725
|
||||
focus_length = 725.0087
|
||||
disp = baseline*focus_length/(depth+1e-8)
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispCreStereo(filename):
|
||||
disp = cv2.imread(filename, cv2.IMREAD_ANYDEPTH) / 32
|
||||
valid = disp > -1e-8
|
||||
return disp, valid
|
||||
|
||||
|
||||
# Method taken from /n/fs/raft-depth/RAFT-Stereo/datasets/SintelStereo/sdk/python/sintel_io.py
|
||||
def readDispSintelStereo(file_name):
|
||||
a = np.array(Image.open(file_name))
|
||||
d_r, d_g, d_b = np.split(a, axis=2, indices_or_sections=3)
|
||||
disp = (d_r.astype('float64') * 4 + d_g.astype('float64') / (2**6) + d_b.astype('float64') / (2**14))[..., 0]
|
||||
mask = np.array(Image.open(file_name.replace('disparities', 'occlusions')))
|
||||
valid = ((mask == 0) & (disp > -1e-8))
|
||||
return disp, valid
|
||||
|
||||
|
||||
# Method taken from https://research.nvidia.com/sites/default/files/pubs/2018-06_Falling-Things/readme_0.txt
|
||||
def readDispFallingThings(file_name):
|
||||
a = np.array(Image.open(file_name))
|
||||
with open('/'.join(file_name.split('/')[:-1] + ['_camera_settings.json']), 'r') as f:
|
||||
intrinsics = json.load(f)
|
||||
if 'left' in file_name:
|
||||
idx = 0
|
||||
else:
|
||||
idx = 1
|
||||
fx = intrinsics['camera_settings'][idx]['intrinsic_settings']['fx']
|
||||
disp = (fx * 6.0 * 100) / a.astype(np.float32)
|
||||
valid = disp > -1e-8
|
||||
return disp, valid
|
||||
|
||||
|
||||
# Method taken from https://github.com/castacks/tartanair_tools/blob/master/data_type.md
|
||||
def readDispTartanAir(file_name):
|
||||
depth = np.load(file_name)
|
||||
disp = 80.0 / depth
|
||||
valid = disp > -1e-8
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispBooster(file_name):
|
||||
disp = np.load(file_name)
|
||||
valid = disp > 0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDisp3DKenBurns(file_name):
|
||||
depth = cv2.imread(file_name, cv2.IMREAD_ANYCOLOR | cv2.IMREAD_ANYDEPTH)
|
||||
meta_file_name = file_name.replace('-depth', '')[:-7]+'-meta.json'
|
||||
fltFov = json.loads(open(meta_file_name, 'r').read())['fltFov']
|
||||
fltFocal = 0.5 * 512 * math.tan(math.radians(90.0) - (0.5 * math.radians(fltFov)))
|
||||
fltBaseline = 40.0
|
||||
disp = (fltFocal * fltBaseline) / depth
|
||||
valid = disp > 0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispMiddlebury0(file_name):
|
||||
if basename(file_name) == 'disp0GT.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
assert len(disp.shape) == 2
|
||||
nocc_pix = file_name.replace('disp0GT.pfm', 'mask0nocc.png')
|
||||
assert exists(nocc_pix)
|
||||
nocc_pix = imageio.imread(nocc_pix) == 255
|
||||
assert np.any(nocc_pix)
|
||||
return disp, nocc_pix
|
||||
elif basename(file_name) == 'disp1GT.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
assert len(disp.shape) == 2
|
||||
nocc_pix = file_name.replace('disp1GT.pfm', 'mask1nocc.png')
|
||||
assert exists(nocc_pix)
|
||||
nocc_pix = imageio.imread(nocc_pix) == 255
|
||||
assert np.any(nocc_pix)
|
||||
return disp, nocc_pix
|
||||
elif basename(file_name) == 'disp0.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
valid = disp < 1e3
|
||||
return disp, valid
|
||||
elif basename(file_name) == 'disp1.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
valid = disp < 1e3
|
||||
return disp, valid
|
||||
elif splitext(file_name)[-1] == '.png':
|
||||
disp = np.array(Image.open(file_name)).astype(np.float32)
|
||||
valid = disp > 0.0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def readDispMiddlebury(file_name):
|
||||
if basename(file_name) == 'disp0GT.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
return disp, disp<1e3
|
||||
elif basename(file_name) == 'disp1GT.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
return disp, disp<1e3
|
||||
elif basename(file_name) == 'disp0.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
valid = disp < 1e3
|
||||
return disp, valid
|
||||
elif basename(file_name) == 'disp1.pfm':
|
||||
disp = readPFM(file_name).astype(np.float32)
|
||||
valid = disp < 1e3
|
||||
return disp, valid
|
||||
elif splitext(file_name)[-1] == '.png':
|
||||
disp = np.array(Image.open(file_name)).astype(np.float32)
|
||||
valid = disp > 0.0
|
||||
return disp, valid
|
||||
|
||||
|
||||
def writeFlowKITTI(filename, uv):
|
||||
uv = 64.0 * uv + 2**15
|
||||
valid = np.ones([uv.shape[0], uv.shape[1], 1])
|
||||
uv = np.concatenate([uv, valid], axis=-1).astype(np.uint16)
|
||||
cv2.imwrite(filename, uv[..., ::-1])
|
||||
|
||||
|
||||
def read_gen(file_name, pil=False):
|
||||
ext = splitext(file_name)[-1]
|
||||
if ext == '.png' or ext == '.jpeg' or ext == '.ppm' or ext == '.jpg':
|
||||
return Image.open(file_name)
|
||||
elif ext == '.bin' or ext == '.raw':
|
||||
return np.load(file_name)
|
||||
elif ext == '.flo':
|
||||
return readFlow(file_name).astype(np.float32)
|
||||
elif ext == '.pfm':
|
||||
flow = readPFM(file_name).astype(np.float32)
|
||||
if len(flow.shape) == 2:
|
||||
return flow
|
||||
else:
|
||||
return flow[:, :, :-1]
|
||||
elif ext == '.exr':
|
||||
disp = cv2.imread(file_name, cv2.IMREAD_ANYCOLOR | cv2.IMREAD_ANYDEPTH)
|
||||
if len(disp.shape) > 2:
|
||||
disp = disp[..., 0]
|
||||
return disp
|
||||
return []
|
||||
@@ -0,0 +1,242 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from scipy import interpolate
|
||||
import glob
|
||||
import os.path as osp
|
||||
|
||||
|
||||
def get_danv2_io_size(h, w, nds, max_i_size=2688, multiple_of=14):
|
||||
"""compute the input and output sizes of danv2 network"""
|
||||
danv2_oh, danv2_ow = h//2**nds, w//2**nds
|
||||
danv2_io_factor = 3.5 # more precise, 14/8=3.5
|
||||
ih, iw = danv2_io_factor*danv2_oh, danv2_io_factor*danv2_ow
|
||||
ih = int(np.ceil(ih / multiple_of) * multiple_of)
|
||||
iw = int(np.ceil(iw / multiple_of) * multiple_of)
|
||||
|
||||
max_i_size = int(np.floor(max_i_size / multiple_of) * multiple_of)
|
||||
|
||||
if ih <= max_i_size and iw <= max_i_size:
|
||||
danv2_ih, danv2_iw = ih, iw
|
||||
else:
|
||||
factor_h = max_i_size/ih
|
||||
factor_w = max_i_size/iw
|
||||
|
||||
if factor_w > factor_h:
|
||||
danv2_ih = max_i_size
|
||||
danv2_iw = int(np.ceil(factor_h * iw / multiple_of) * multiple_of)
|
||||
else:
|
||||
danv2_iw = max_i_size
|
||||
danv2_ih = int(np.ceil(factor_w * ih / multiple_of) * multiple_of)
|
||||
|
||||
return danv2_ih, danv2_iw, danv2_oh, danv2_ow
|
||||
|
||||
|
||||
class InputPadder:
|
||||
""" Pads images such that dimensions are divisible by 8 """
|
||||
def __init__(self, dims, mode='sintel', divis_by=8):
|
||||
self.ht, self.wd = dims[-2:]
|
||||
pad_ht = (((self.ht // divis_by) + 1) * divis_by - self.ht) % divis_by
|
||||
pad_wd = (((self.wd // divis_by) + 1) * divis_by - self.wd) % divis_by
|
||||
if mode == 'sintel':
|
||||
self._pad = [pad_wd//2, pad_wd - pad_wd//2, pad_ht//2, pad_ht - pad_ht//2]
|
||||
else:
|
||||
self._pad = [pad_wd//2, pad_wd - pad_wd//2, 0, pad_ht]
|
||||
|
||||
def pad(self, *inputs):
|
||||
assert all((x.ndim == 4) for x in inputs)
|
||||
return [F.pad(x, self._pad, mode='replicate') for x in inputs]
|
||||
|
||||
def unpad(self, x):
|
||||
assert x.ndim == 4
|
||||
ht, wd = x.shape[-2:]
|
||||
c = [self._pad[2], ht-self._pad[3], self._pad[0], wd-self._pad[1]]
|
||||
return x[..., c[0]:c[1], c[2]:c[3]]
|
||||
|
||||
|
||||
def forward_interpolate(flow):
|
||||
flow = flow.detach().cpu().numpy()
|
||||
dx, dy = flow[0], flow[1]
|
||||
|
||||
ht, wd = dx.shape
|
||||
x0, y0 = np.meshgrid(np.arange(wd), np.arange(ht))
|
||||
|
||||
x1 = x0 + dx
|
||||
y1 = y0 + dy
|
||||
|
||||
x1 = x1.reshape(-1)
|
||||
y1 = y1.reshape(-1)
|
||||
dx = dx.reshape(-1)
|
||||
dy = dy.reshape(-1)
|
||||
|
||||
valid = (x1 > 0) & (x1 < wd) & (y1 > 0) & (y1 < ht)
|
||||
x1 = x1[valid]
|
||||
y1 = y1[valid]
|
||||
dx = dx[valid]
|
||||
dy = dy[valid]
|
||||
|
||||
flow_x = interpolate.griddata(
|
||||
(x1, y1), dx, (x0, y0), method='nearest', fill_value=0)
|
||||
|
||||
flow_y = interpolate.griddata(
|
||||
(x1, y1), dy, (x0, y0), method='nearest', fill_value=0)
|
||||
|
||||
flow = np.stack([flow_x, flow_y], axis=0)
|
||||
return torch.from_numpy(flow).float()
|
||||
|
||||
|
||||
def bilinear_sampler(img, coords, mode='bilinear', mask=False):
|
||||
""" Wrapper for grid_sample, uses pixel coordinates """
|
||||
H, W = img.shape[-2:]
|
||||
xgrid, ygrid = coords.split([1, 1], dim=-1)
|
||||
xgrid = 2*xgrid/(W-1) - 1
|
||||
if H > 1:
|
||||
ygrid = 2*ygrid/(H-1) - 1
|
||||
|
||||
grid = torch.cat([xgrid, ygrid], dim=-1)
|
||||
img = F.grid_sample(img, grid, align_corners=True)
|
||||
# img = bilinear_grid_sample(img, grid, align_corners=True)
|
||||
|
||||
if mask:
|
||||
mask = (xgrid > -1) & (ygrid > -1) & (xgrid < 1) & (ygrid < 1)
|
||||
return img, mask.float()
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def coords_grid(batch, ht, wd):
|
||||
coords = torch.meshgrid(torch.arange(ht), torch.arange(wd))
|
||||
coords = torch.stack(coords[::-1], dim=0).float()
|
||||
return coords[None].repeat(batch, 1, 1, 1)
|
||||
|
||||
|
||||
def upflow(flow, factor=8, mode='bilinear', sacle=True):
|
||||
new_size = (factor * flow.shape[2], factor * flow.shape[3])
|
||||
if sacle:
|
||||
return factor * F.interpolate(flow, size=new_size, mode=mode, align_corners=True)
|
||||
else:
|
||||
return F.interpolate(flow, size=new_size, mode=mode, align_corners=True)
|
||||
|
||||
|
||||
def gauss_blur(input, N=5, std=1):
|
||||
B, D, H, W = input.shape
|
||||
x, y = torch.meshgrid(torch.arange(N).float() - N//2, torch.arange(N).float() - N//2)
|
||||
unnormalized_gaussian = torch.exp(-(x.pow(2) + y.pow(2)) / (2 * std ** 2))
|
||||
weights = unnormalized_gaussian / unnormalized_gaussian.sum().clamp(min=1e-4)
|
||||
weights = weights.view(1, 1, N, N).to(input)
|
||||
output = F.conv2d(input.reshape(B*D, 1, H, W), weights, padding=N//2)
|
||||
return output.view(B, D, H, W)
|
||||
|
||||
|
||||
# Ref: https://zenn.dev/pinto0309/scraps/7d4032067d0160
|
||||
def bilinear_grid_sample(im, grid, align_corners=False):
|
||||
"""Given an input and a flow-field grid, computes the output using input
|
||||
values and pixel locations from grid. Supported only bilinear interpolation
|
||||
method to sample the input pixels.
|
||||
|
||||
Args:
|
||||
im (torch.Tensor): Input feature map, shape (N, C, H, W)
|
||||
grid (torch.Tensor): Point coordinates, shape (N, Hg, Wg, 2)
|
||||
align_corners {bool}: If set to True, the extrema (-1 and 1) are
|
||||
considered as referring to the center points of the input’s
|
||||
corner pixels. If set to False, they are instead considered as
|
||||
referring to the corner points of the input’s corner pixels,
|
||||
making the sampling more resolution agnostic.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: A tensor with sampled points, shape (N, C, Hg, Wg)
|
||||
"""
|
||||
n, c, h, w = im.shape
|
||||
gn, gh, gw, _ = grid.shape
|
||||
assert n == gn
|
||||
|
||||
x = grid[:, :, :, 0]
|
||||
y = grid[:, :, :, 1]
|
||||
|
||||
if align_corners:
|
||||
x = ((x + 1) / 2) * (w - 1)
|
||||
y = ((y + 1) / 2) * (h - 1)
|
||||
else:
|
||||
x = ((x + 1) * w - 1) / 2
|
||||
y = ((y + 1) * h - 1) / 2
|
||||
|
||||
x = x.view(n, -1)
|
||||
y = y.view(n, -1)
|
||||
|
||||
x0 = torch.floor(x).long()
|
||||
y0 = torch.floor(y).long()
|
||||
x1 = x0 + 1
|
||||
y1 = y0 + 1
|
||||
|
||||
wa = ((x1 - x) * (y1 - y)).unsqueeze(1)
|
||||
wb = ((x1 - x) * (y - y0)).unsqueeze(1)
|
||||
wc = ((x - x0) * (y1 - y)).unsqueeze(1)
|
||||
wd = ((x - x0) * (y - y0)).unsqueeze(1)
|
||||
|
||||
# Apply default for grid_sample function zero padding
|
||||
im_padded = torch.nn.functional.pad(im, pad=[1, 1, 1, 1], mode='constant', value=0)
|
||||
padded_h = h + 2
|
||||
padded_w = w + 2
|
||||
# save points positions after padding
|
||||
x0, x1, y0, y1 = x0 + 1, x1 + 1, y0 + 1, y1 + 1
|
||||
|
||||
# Clip coordinates to padded image size
|
||||
x0 = torch.where(x0 < 0, torch.tensor(0, device=im.device), x0)
|
||||
x0 = torch.where(x0 > padded_w - 1, torch.tensor(padded_w - 1, device=im.device), x0)
|
||||
x1 = torch.where(x1 < 0, torch.tensor(0, device=im.device), x1)
|
||||
x1 = torch.where(x1 > padded_w - 1, torch.tensor(padded_w - 1, device=im.device), x1)
|
||||
y0 = torch.where(y0 < 0, torch.tensor(0, device=im.device), y0)
|
||||
y0 = torch.where(y0 > padded_h - 1, torch.tensor(padded_h - 1, device=im.device), y0)
|
||||
y1 = torch.where(y1 < 0, torch.tensor(0, device=im.device), y1)
|
||||
y1 = torch.where(y1 > padded_h - 1, torch.tensor(padded_h - 1, device=im.device), y1)
|
||||
|
||||
im_padded = im_padded.view(n, c, -1)
|
||||
|
||||
x0_y0 = (x0 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x0_y1 = (x0 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y0 = (x1 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y1 = (x1 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
|
||||
Ia = torch.gather(im_padded, 2, x0_y0)
|
||||
Ib = torch.gather(im_padded, 2, x0_y1)
|
||||
Ic = torch.gather(im_padded, 2, x1_y0)
|
||||
Id = torch.gather(im_padded, 2, x1_y1)
|
||||
|
||||
return (Ia * wa + Ib * wb + Ic * wc + Id * wd).reshape(n, c, gh, gw)
|
||||
|
||||
|
||||
def read_kitti_calib_file(path):
|
||||
"""Read KITTI calibration file
|
||||
(from https://github.com/hunse/kitti)
|
||||
"""
|
||||
float_chars = set("0123456789.e+- ")
|
||||
data = {}
|
||||
with open(path, 'r') as f:
|
||||
for line in f.readlines():
|
||||
key, value = line.split(':', 1)
|
||||
value = value.strip()
|
||||
data[key] = value
|
||||
if float_chars.issuperset(value):
|
||||
# try to cast to float array
|
||||
try:
|
||||
data[key] = np.array(list(map(float, value.split(' '))))
|
||||
except ValueError:
|
||||
# casting error: data[key] already eq. value, so pass
|
||||
pass
|
||||
|
||||
return data
|
||||
|
||||
|
||||
# from https://github.com/ozendelait/rvc_devkit/blob/master/stereo/stereo_devkit.py
|
||||
def ReadMiddlebury2014CalibFile(path):
|
||||
result = dict()
|
||||
with open(path, 'rb') as calib_file:
|
||||
for line in calib_file.readlines():
|
||||
line = line.decode('UTF-8').rstrip('\n')
|
||||
if len(line) == 0:
|
||||
continue
|
||||
eq_pos = line.find('=')
|
||||
if eq_pos < 0:
|
||||
raise Exception('Cannot parse Middlebury 2014 calib file: ' + path)
|
||||
result[line[:eq_pos]] = line[eq_pos + 1:]
|
||||
return result
|
||||
Reference in New Issue
Block a user