Добавлены пропсы конвейера и стереодвижки, задействованные в прогоне

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:
dasha_f
2026-08-01 13:12:07 +00:00
parent 0d32f32db0
commit 6e1a22ba8b
184 changed files with 17666 additions and 3 deletions
View File
+212
View File
@@ -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())
+142
View File
@@ -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
+388
View File
@@ -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
+583
View File
@@ -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
+195
View File
@@ -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
+307
View File
@@ -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
+105
View File
@@ -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
+286
View File
@@ -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 []
+242
View File
@@ -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 inputs
corner pixels. If set to False, they are instead considered as
referring to the corner points of the inputs 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