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

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
+50
View File
@@ -0,0 +1,50 @@
import os,sys
code_dir = os.path.dirname(os.path.abspath(__file__))
sys.path.append(code_dir+'/../')
from foundation_stereo_ori.submodule import FeatureAtt
import torch
import torch.nn as nn
import Utils as U
import pickle
class ForwardHelper(nn.Module):
def __init__(self, layers:list):
super().__init__()
self.layers = nn.ModuleList(layers)
def forward(self, x, left_feat=None):
for layer in self.layers:
if isinstance(layer, FeatureAtt):
x = layer(x, left_feat)
else:
x = layer(x)
return x
class PostForwardHelper(nn.Module):
def __init__(self, layers:list):
super().__init__()
for pos in range(len(layers)):
if layers[pos] in ['sum', 'concat']:
self.op = layers[pos]
break
self.upsample = nn.Sequential(*layers[:pos])
self.out = nn.ModuleList(layers[pos+1:])
def forward(self, conv2, conv3, left_feat=None):
conv3_up = self.upsample(conv3)
if self.op == 'sum':
x = conv3_up + conv2
elif self.op == 'concat':
x = torch.cat((conv3_up, conv2), dim=1)
else:
raise ValueError(f"Unknown operation: {self.op}")
for layer in self.out:
if isinstance(layer, FeatureAtt):
x = layer(x, left_feat)
else:
x = layer(x)
return x
+78
View File
@@ -0,0 +1,78 @@
import torch,os,sys
import torch.nn as nn
import torch.nn.functional as F
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
from core.submodule import Conv2x_IN
import timm
class ContextNetSharedBackbone(nn.Module):
def __init__(self, args, c04, c08, c16, output_dim=[(128,128,128), (128,128,128)], norm_fn='batch', downsample=3):
super().__init__()
self.args = args
self.conv04 = nn.ModuleList([
nn.Conv2d(c04, output_dim[0][0], kernel_size=3, padding=1),
nn.Conv2d(c04, output_dim[1][0], kernel_size=3, padding=1),
])
def forward(self, x4, x8, x16):
outputs04 = []
for i in range(len(self.conv04)):
outputs04.append(self.conv04[i](x4))
return (outputs04,)
class DepthAnythingFeature:
model_configs = {
'vitl': {'encoder': 'vitl', 'features': 256, 'out_channels': [256, 512, 1024, 1024]},
'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]},
'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]}
}
class Feature(nn.Module):
def __init__(self, args):
super(Feature, self).__init__()
self.args = args
model = timm.create_model('edgenext_small', pretrained=True, features_only=False)
self.stem = model.stem
self.stages = model.stages
chans = [48, 96, 160, 304]
self.chans = chans
vit_feat_dim = DepthAnythingFeature.model_configs[self.args.vit_size]['features']//2
self.deconv32_16 = Conv2x_IN(chans[3], chans[2], deconv=True, concat=True)
self.deconv16_8 = Conv2x_IN(chans[2]*2, chans[1], deconv=True, concat=True)
self.deconv8_4 = Conv2x_IN(chans[1]*2, chans[0], deconv=True, concat=True)
self.conv4 = nn.Conv2d(chans[0]*2, self.chans[0]*2+vit_feat_dim, kernel_size=1, stride=1, padding=0)
self.d_out = [self.chans[0]*2+vit_feat_dim, self.chans[1]*2, self.chans[2]*2, self.chans[3]]
def forward(self, x):
B,C,H,W = x.shape
if hasattr(self, 'stem'):
x = self.stem(x)
x4 = self.stages[0](x)
x8 = self.stages[1](x4)
x16 = self.stages[2](x8)
x32 = self.stages[3](x16)
else:
intermediates = self.model.forward_intermediates(x, intermediates_only=True)
x4, x8, x16, x32 = intermediates[-4:]
with torch.profiler.record_function("feature_deconv"):
x16 = self.deconv32_16(x32, x16)
x8 = self.deconv16_8(x16, x8)
x4 = self.deconv8_4(x8, x4)
x4 = self.conv4(x4)
if hasattr(self, 'conv8'):
x8 = self.conv8(x8)
x16 = self.conv16(x16)
x32 = self.conv32(x32)
return [x4, x8, x16, x32]
+445
View File
@@ -0,0 +1,445 @@
import torch,pdb,logging,timm
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import sys,os
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
from core.update import BasicSelectiveMultiUpdateBlock
from core.extractor import ContextNetSharedBackbone, Feature
from core.geometry import Combined_Geo_Encoding_Volume
from core.submodule import (
BasicConv, Conv3dNormActReduced, ResnetBasicBlock3D, BasicConv_IN, Conv2x,
FeatureAtt, CostVolumeDisparityAttention, SpatialAttentionExtractor,
ChannelAttentionEnhancement, disparity_regression, context_upsample,
build_gwc_volume_optimized_pytorch1, build_gwc_volume_triton,
build_concat_volume_optimized_pytorch1, build_concat_volume_optimized_pytorch,
)
from core.utils.utils import InputPadder
import Utils as U
import time
sys.modules['foundation_stereo_ori'] = sys.modules['core']
sys.modules['foundation_stereo_ori.submodule'] = sys.modules['core.submodule']
sys.modules['foundation_stereo_ori.extractor'] = sys.modules['core.extractor']
sys.modules['foundation_stereo_ori.update'] = sys.modules['core.update']
sys.modules['foundation_stereo_ori.foundation_stereo'] = sys.modules['core.foundation_stereo']
class FoundationStereo(nn.Module):
pass
def normalize_image(img):
'''
@img: (B,C,H,W) in range 0-255, RGB order
'''
mean = img.new_tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
std = img.new_tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
return (img/255.0 - mean) / std
class hourglass(nn.Module):
def __init__(self, cfg, in_channels, feat_dims=None):
super().__init__()
self.cfg = cfg
self.conv1 = nn.Sequential(BasicConv(in_channels, in_channels*2, is_3d=True, bn=True, relu=True, kernel_size=3,
padding=1, stride=2, dilation=1),
Conv3dNormActReduced(in_channels*2, in_channels*2, kernel_size=3, kernel_disp=17))
self.conv2 = nn.Sequential(BasicConv(in_channels*2, in_channels*4, is_3d=True, bn=True, relu=True, kernel_size=3,
padding=1, stride=2, dilation=1),
Conv3dNormActReduced(in_channels*4, in_channels*4, kernel_size=3, kernel_disp=17))
self.conv3 = nn.Sequential(BasicConv(in_channels*4, in_channels*6, is_3d=True, bn=True, relu=True, kernel_size=3,
padding=1, stride=2, dilation=1),
Conv3dNormActReduced(in_channels*6, in_channels*6, kernel_size=3, kernel_disp=17))
self.conv3_up = BasicConv(in_channels*6, in_channels*4, deconv=True, is_3d=True, bn=True,
relu=True, kernel_size=(4, 4, 4), padding=(1, 1, 1), stride=(2, 2, 2))
self.conv2_up = BasicConv(in_channels*4, in_channels*2, deconv=True, is_3d=True, bn=True,
relu=True, kernel_size=(4, 4, 4), padding=(1, 1, 1), stride=(2, 2, 2))
self.conv1_up = BasicConv(in_channels*2, in_channels, deconv=True, is_3d=True, bn=True,
relu=True, kernel_size=(4, 4, 4), padding=(1, 1, 1), stride=(2, 2, 2))
self.conv_out = nn.Sequential(
Conv3dNormActReduced(in_channels, in_channels, kernel_size=3, kernel_disp=17),
Conv3dNormActReduced(in_channels, in_channels, kernel_size=3, kernel_disp=17),
)
self.agg_0 = nn.Sequential(BasicConv(in_channels*8, in_channels*4, is_3d=True, kernel_size=1, padding=0, stride=1),
Conv3dNormActReduced(in_channels*4, in_channels*4, kernel_size=3, kernel_disp=17),
Conv3dNormActReduced(in_channels*4, in_channels*4, kernel_size=3, kernel_disp=17),)
self.agg_1 = nn.Sequential(BasicConv(in_channels*4, in_channels*2, is_3d=True, kernel_size=1, padding=0, stride=1),
Conv3dNormActReduced(in_channels*2, in_channels*2, kernel_size=3, kernel_disp=17),
Conv3dNormActReduced(in_channels*2, in_channels*2, kernel_size=3, kernel_disp=17))
self.atts = nn.ModuleDict({
"4": CostVolumeDisparityAttention(d_model=in_channels, nhead=4, dim_feedforward=in_channels, norm_first=False, num_transformer=4, max_len=self.cfg['max_disp']//16),
})
self.conv_patch = nn.Sequential(
nn.Conv3d(in_channels, in_channels, kernel_size=4, stride=4, padding=0, groups=in_channels),
nn.BatchNorm3d(in_channels),
)
self.feature_att_8 = FeatureAtt(in_channels*2, feat_dims[1])
self.feature_att_16 = FeatureAtt(in_channels*4, feat_dims[2])
self.feature_att_32 = FeatureAtt(in_channels*6, feat_dims[3])
self.feature_att_up_16 = FeatureAtt(in_channels*4, feat_dims[2])
self.feature_att_up_8 = FeatureAtt(in_channels*2, feat_dims[1])
self.post32_to_16 = None
self.post16_to_8 = None
self.post8_to_4 = None
def forward(self, x, features):
conv1 = self.conv1(x)
conv1 = self.feature_att_8(conv1, features[1])
conv2 = self.conv2(conv1)
conv2 = self.feature_att_16(conv2, features[2])
conv3 = self.conv3(conv2)
conv3 = self.feature_att_32(conv3, features[3])
if not hasattr(self, 'post32_to_16') or self.post32_to_16 is None:
conv3_up = self.conv3_up(conv3)
conv2 = torch.cat((conv3_up, conv2), dim=1)
conv2 = self.agg_0(conv2)
conv2 = self.feature_att_up_16(conv2, features[2])
else:
conv2 = self.post32_to_16(conv2, conv3, features[2])
if not hasattr(self, 'post16_to_8') or self.post16_to_8 is None:
conv2_up = self.conv2_up(conv2)
conv1 = torch.cat((conv2_up, conv1), dim=1)
conv1 = self.agg_1(conv1)
conv1 = self.feature_att_up_8(conv1, features[1])
else:
conv1 = self.post16_to_8(conv1, conv2, features[1])
conv = self.conv1_up(conv1)
if not hasattr(self, 'post8_to_4') or self.post8_to_4 is None:
x = self.conv_patch(x)
x = self.atts["4"](x)
x = F.interpolate(x, scale_factor=4, mode='trilinear', align_corners=False)
conv = conv + x
conv = self.conv_out(conv)
else:
conv = self.post8_to_4(x, conv)
return conv
class FastFoundationStereo(nn.Module):
def __init__(self, args):
super().__init__()
self.args = args
self.dtype = torch.float32
context_dims = args.hidden_dims
self.cv_group = args.get('cv_group', 8)
self.concat_channel = 24
volume_dim = args.get('volume_dim', 28)
self.volume_dim = volume_dim
self.update_block = BasicSelectiveMultiUpdateBlock(self.args, self.args.hidden_dims[0], volume_dim=volume_dim)
self.sam = SpatialAttentionExtractor()
self.cam = ChannelAttentionEnhancement(self.args.hidden_dims[0])
self.context_zqr_convs = nn.ModuleList([nn.Conv2d(context_dims[i], args.hidden_dims[i]*3, kernel_size=3, padding=3//2) for i in range(self.args.n_gru_layers)])
self.feature = Feature(args)
self.proj_cmb = nn.Conv2d(self.feature.d_out[0], self.concat_channel//2, kernel_size=1, padding=0)
self.cnet = ContextNetSharedBackbone(args, c04=self.feature.d_out[0], c08=self.feature.d_out[1], c16=self.feature.d_out[2], output_dim=[args.hidden_dims, context_dims])
self.stem_2 = nn.Sequential(
BasicConv_IN(3, 32, kernel_size=3, stride=2, padding=1),
nn.Conv2d(32, 32, 3, 1, 1, bias=False),
nn.InstanceNorm2d(32), nn.ReLU()
)
self.spx_2_gru = Conv2x(32, 32, deconv=True, bn=False, concat=True)
self.spx_gru = nn.Sequential(
nn.ConvTranspose2d(2*32, 9, kernel_size=4, stride=2, padding=1),
)
self.corr_stem = nn.Sequential(
nn.Conv3d(self.proj_cmb.out_channels*2+self.cv_group, volume_dim, kernel_size=1),
BasicConv(volume_dim, volume_dim, kernel_size=3, padding=1, is_3d=True),
ResnetBasicBlock3D(volume_dim, volume_dim, kernel_size=3, stride=1, padding=1),
ResnetBasicBlock3D(volume_dim, volume_dim, kernel_size=3, stride=1, padding=1),
)
self.corr_feature_att = FeatureAtt(volume_dim, self.feature.d_out[0])
self.cost_agg = hourglass(cfg=self.args, in_channels=volume_dim, feat_dims=self.feature.d_out)
self.classifier = nn.Sequential(
BasicConv(volume_dim, volume_dim//2, kernel_size=3, padding=1, is_3d=True),
ResnetBasicBlock3D(volume_dim//2, volume_dim//2, kernel_size=3, stride=1, padding=1),
nn.Conv3d(volume_dim//2, 1, kernel_size=7, padding=3),
)
r = self.args.corr_radius
dx = torch.arange(-r, r+1, requires_grad=False, dtype=torch.int8).reshape(1, 1, 2*r+1, 1)
self.register_buffer("dx", dx)
def upsample_disp(self, disp, mask_feat_4, stem_2x):
with torch.amp.autocast('cuda', enabled=self.args.mixed_precision, dtype=U.AMP_DTYPE):
xspx = self.spx_2_gru(mask_feat_4, stem_2x) # 1/2 resolution
spx_pred = self.spx_gru(xspx)
spx_pred = F.softmax(spx_pred, 1)
up_disp = context_upsample(disp*4., spx_pred).unsqueeze(1)
return up_disp.to(self.dtype)
def forward(self, image1, image2, iters=12, test_mode=False, low_memory=False, init_disp=None, profile=False, optimize_build_volume='pytorch1'):
""" Estimate disparity between pair of frames """
B,C,H,W = image1.shape
low_memory = low_memory or (self.args.get('low_memory', False))
image1 = normalize_image(image1)
image2 = normalize_image(image2)
with torch.amp.autocast('cuda', enabled=self.args.mixed_precision, dtype=U.AMP_DTYPE):
out = self.feature(torch.cat([image1, image2], dim=0))
features_left = [o[:B] for o in out]
features_right = [o[B:] for o in out]
stem_2x = self.stem_2(image1)
if optimize_build_volume=='pytorch1':
gwc_volume = build_gwc_volume_optimized_pytorch1(features_left[0], features_right[0], self.args.max_disp//4, self.cv_group, normalize=self.args.normalize)
elif optimize_build_volume=='triton':
gwc_volume = build_gwc_volume_triton(features_left[0], features_right[0], self.args.max_disp//4, self.cv_group, normalize=self.args.normalize)
else:
raise RuntimeError(f"Invalid optimize_build_volume: {optimize_build_volume}")
left_tmp = self.proj_cmb(features_left[0])
right_tmp = self.proj_cmb(features_right[0])
concat_volume = build_concat_volume_optimized_pytorch1(left_tmp, right_tmp, maxdisp=self.args.max_disp//4)
del left_tmp, right_tmp
comb_volume = torch.cat([gwc_volume, concat_volume], dim=1)
del concat_volume, gwc_volume
comb_volume = self.corr_stem(comb_volume)
comb_volume = self.corr_feature_att(comb_volume, features_left[0])
comb_volume = self.cost_agg(comb_volume, features_left)
# Init disp from geometry encoding volume
logits = self.classifier(comb_volume).squeeze(1)
prob = F.softmax(logits, dim=1)
if init_disp is None:
init_disp = disparity_regression(prob, self.args.max_disp//4)
cnet_list = self.cnet(features_left[0], features_left[1], features_left[2])
cnet_list = list(cnet_list)
net_list = [torch.tanh(x[0]) for x in cnet_list]
inp_list = [torch.relu(x[1]) for x in cnet_list]
inp_list = [self.cam(x) * x for x in inp_list]
att = [self.sam(x) for x in inp_list]
geo_fn = Combined_Geo_Encoding_Volume(features_left[0].to(self.dtype), features_right[0].to(self.dtype), comb_volume.to(self.dtype), num_levels=self.args.corr_levels)
b, c, h, w = features_left[0].shape
coords = torch.arange(w, dtype=torch.float, device=init_disp.device).reshape(1,1,w,1).repeat(b, h, 1, 1)
disp = init_disp.to(self.dtype)
disp_preds = []
del comb_volume, features_left, features_right, cnet_list
# GRUs iterations to update disparity (1/4 resolution)
for itr in range(iters):
disp = disp.detach()
geo_feat = geo_fn(disp, coords, dx=self.dx, low_memory=low_memory)
with torch.amp.autocast('cuda', enabled=self.args.mixed_precision, dtype=U.AMP_DTYPE):
net_list, mask_feat_4, delta_disp = self.update_block(net_list, inp_list, geo_feat.to(self.dtype), disp, att)
disp = disp + delta_disp.to(self.dtype)
if test_mode and itr < iters-1:
continue
# upsample predictions
disp_up = self.upsample_disp(disp.to(self.dtype), mask_feat_4.to(self.dtype), stem_2x.to(self.dtype))
disp_preds.append(disp_up)
if test_mode:
return disp_up
return init_disp, disp_preds
def run_hierachical(self, image1, image2, iters=12, test_mode=False, low_memory=False, small_ratio=0.5):
B,_,H,W = image1.shape
img1_small = F.interpolate(image1, scale_factor=small_ratio, align_corners=False, mode='bilinear')
img2_small = F.interpolate(image2, scale_factor=small_ratio, align_corners=False, mode='bilinear')
padder = InputPadder(img1_small.shape[-2:], divis_by=32, force_square=False)
img1_small, img2_small = padder.pad(img1_small, img2_small)
disp_small = self.forward(img1_small, img2_small, test_mode=True, iters=iters, low_memory=low_memory)
disp_small = padder.unpad(disp_small)
disp_small_up = F.interpolate(disp_small, size=(H,W), mode='bilinear', align_corners=True) * 1/small_ratio
disp_small_up = disp_small_up.clip(0, None)
padder = InputPadder(image1.shape[-2:], divis_by=32, force_square=False)
image1, image2, disp_small_up = padder.pad(image1, image2, disp_small_up)
disp_small_up += padder._pad[0]
init_disp = F.interpolate(disp_small_up, scale_factor=0.25, mode='bilinear', align_corners=True) * 0.25 # Init disp will be 1/4
disp = self.forward(image1, image2, iters=iters, test_mode=test_mode, low_memory=low_memory, init_disp=init_disp)
disp = padder.unpad(disp)
return disp
FoundationStereoLite = FastFoundationStereo
class TrtFeatureRunner(nn.Module):
def __init__(self, model):
super().__init__()
self.feature = model.feature
self.stem_2 = model.stem_2
def forward(self, image1, image2):
image1 = normalize_image(image1)
image2 = normalize_image(image2)
B = len(image1)
out = self.feature(torch.cat([image1, image2], dim=0))
features_left = [o[:B] for o in out]
features_right = [o[B:] for o in out]
stem_2x = self.stem_2(image1)
return *features_left, features_right[0], stem_2x
class TrtPostRunner(nn.Module):
def __init__(self, model):
super().__init__()
self.args = model.args
self.dtype = model.dtype
self.register_buffer("dx", model.dx)
self.proj_cmb = model.proj_cmb
self.corr_stem = model.corr_stem
self.corr_feature_att = model.corr_feature_att
self.cost_agg = model.cost_agg
self.classifier = model.classifier
self.update_block = model.update_block
self.sam = model.sam
self.cam = model.cam
self.feature = model.feature
self.spx_2_gru = model.spx_2_gru
self.spx_gru = model.spx_gru
self.cnet = model.cnet
def upsample_disp(self, disp, mask_feat_4, stem_2x):
with torch.amp.autocast('cuda', enabled=self.args.mixed_precision, dtype=U.AMP_DTYPE):
xspx = self.spx_2_gru(mask_feat_4, stem_2x) # 1/2 resolution
spx_pred = self.spx_gru(xspx)
spx_pred = F.softmax(spx_pred, 1)
up_disp = context_upsample(disp*4., spx_pred).unsqueeze(1)
return up_disp.to(self.dtype)
def forward(self, features_left_04, features_left_08, features_left_16, features_left_32, features_right_04, stem_2x, gwc_volume):
features_left = [features_left_04, features_left_08, features_left_16, features_left_32]
left_tmp = self.proj_cmb(features_left_04)
right_tmp = self.proj_cmb(features_right_04)
concat_volume = build_concat_volume_optimized_pytorch(left_tmp, right_tmp, maxdisp=self.args.max_disp//4)
del left_tmp, right_tmp
comb_volume = torch.cat([gwc_volume, concat_volume], dim=1)
del concat_volume, gwc_volume
comb_volume = self.corr_stem(comb_volume)
comb_volume = self.corr_feature_att(comb_volume, features_left_04)
comb_volume = self.cost_agg(comb_volume, features_left)
# Init disp from geometry encoding volume
logits = self.classifier(comb_volume).squeeze(1)
prob = F.softmax(logits, dim=1)
init_disp = disparity_regression(prob, self.args.max_disp//4)
cnet_list = self.cnet(features_left[0], features_left[1], features_left[2])
cnet_list = list(cnet_list)
net_list = [torch.tanh(x[0]) for x in cnet_list]
inp_list = [torch.relu(x[1]) for x in cnet_list]
inp_list = [self.cam(x) * x for x in inp_list]
att = [self.sam(x) for x in inp_list]
geo_fn = Combined_Geo_Encoding_Volume(features_left_04.to(self.dtype), features_right_04.to(self.dtype), comb_volume.to(self.dtype), num_levels=self.args.corr_levels)
b, c, h, w = features_left[0].shape
coords = torch.arange(w, dtype=torch.float, device=init_disp.device).reshape(1,1,w,1).repeat(b, h, 1, 1)
disp = init_disp.to(self.dtype)
# GRUs iterations to update disparity (1/4 resolution)
for itr in range(self.args.valid_iters):
disp = disp.detach()
geo_feat = geo_fn(disp, coords, dx=self.dx, low_memory=True)
net_list, mask_feat_4, delta_disp = self.update_block(net_list, inp_list, geo_feat.to(self.dtype), disp, att)
disp = disp + delta_disp.to(self.dtype)
if itr < self.args.valid_iters-1:
continue
disp_up = self.upsample_disp(disp.to(self.dtype), mask_feat_4.to(self.dtype), stem_2x.to(self.dtype))
return disp_up
class TrtRunner(nn.Module):
def __init__(self, args, feature_runner_engine_path, post_runner_engine_path):
super().__init__()
import tensorrt as trt
self.args = args
with open(feature_runner_engine_path, 'rb') as file:
engine_data = file.read()
self.trt_logger = trt.Logger(trt.Logger.WARNING)
self.feature_engine = trt.Runtime(self.trt_logger).deserialize_cuda_engine(engine_data)
self.feature_context = self.feature_engine.create_execution_context()
with open(post_runner_engine_path, 'rb') as file:
engine_data = file.read()
self.post_engine = trt.Runtime(self.trt_logger).deserialize_cuda_engine(engine_data)
self.post_context = self.post_engine.create_execution_context()
self.max_disp = args.max_disp
self.cv_group = args.get('cv_group', 8)
def trt_dtype_to_torch(self, dt):
import tensorrt as trt
if dt==trt.DataType.FLOAT: return torch.float32
if dt==trt.DataType.HALF: return torch.float16
if dt==trt.DataType.BF16: return torch.bfloat16
if dt==trt.DataType.INT32: return torch.int32
if dt==trt.DataType.INT8: return torch.int8
if dt==trt.DataType.BOOL: return torch.bool
raise RuntimeError(f"Unsupported TRT dtype: {dt}")
def get_io_tensor_names(self, engine, mode):
names = []
n = engine.num_io_tensors
for i in range(n):
name = engine.get_tensor_name(i)
if engine.get_tensor_mode(name)==mode:
names.append(name)
return names
def run_trt(self, engine, context, inputs_by_name:dict):
import tensorrt as trt
for name, tensor in list(inputs_by_name.items()):
expected_dtype = self.trt_dtype_to_torch(engine.get_tensor_dtype(name))
if tensor.dtype != expected_dtype: inputs_by_name[name] = tensor.to(expected_dtype)
if not inputs_by_name[name].is_contiguous(): inputs_by_name[name] = inputs_by_name[name].contiguous()
context.set_input_shape(name, tuple(inputs_by_name[name].shape))
outputs = {}
out_names = [n for n in self.get_io_tensor_names(engine, trt.TensorIOMode.OUTPUT)]
for name in out_names:
shp = tuple(context.get_tensor_shape(name))
dtype = self.trt_dtype_to_torch(engine.get_tensor_dtype(name))
outputs[name] = torch.empty(shp, device='cuda', dtype=dtype)
for name, tensor in inputs_by_name.items(): context.set_tensor_address(name, int(tensor.data_ptr()))
for name, tensor in outputs.items(): context.set_tensor_address(name, int(tensor.data_ptr()))
stream = torch.cuda.current_stream().cuda_stream
ok = context.execute_async_v3(stream)
assert ok
return outputs
def forward(self, image1, image2):
import tensorrt as trt
feat_out = self.run_trt(self.feature_engine, self.feature_context, {'left': image1, 'right': image2})
gwc_volume = build_gwc_volume_triton(feat_out['features_left_04'].half(), feat_out['features_right_04'].half(), self.args.max_disp//4, self.cv_group, normalize=self.args.normalize)
post_inputs = feat_out
post_inputs['gwc_volume'] = gwc_volume
in_names = self.get_io_tensor_names(self.post_engine, trt.TensorIOMode.INPUT)
tmp_keys = list(post_inputs.keys())
for k in tmp_keys:
if k not in in_names:
del post_inputs[k]
out = self.run_trt(self.post_engine, self.post_context, post_inputs)
disp = out['disp']
return disp
+81
View File
@@ -0,0 +1,81 @@
import torch,os,sys
import torch.nn.functional as F
from core.utils.utils import bilinear_sampler, bilinear_sampler1d
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
class Combined_Geo_Encoding_Volume:
def __init__(self, init_fmap1, init_fmap2, geo_volume, num_levels=2):
self.num_levels = num_levels
self.geo_volume_pyramid = []
self.init_corr_pyramid = []
# all pairs correlation
init_corr = Combined_Geo_Encoding_Volume.corr(init_fmap1, init_fmap2)
b, h, w, _, w2 = init_corr.shape
b, c, d, h, w = geo_volume.shape
geo_volume = geo_volume.permute(0, 3, 4, 1, 2).reshape(b*h*w, c, 1, d)
init_corr = init_corr.view(b*h*w, 1, 1, w2)
self.geo_volume_pyramid.append(geo_volume)
self.init_corr_pyramid.append(init_corr)
for _ in range(self.num_levels-1):
geo_volume = F.avg_pool2d(geo_volume, [1,2], stride=[1,2])
self.geo_volume_pyramid.append(geo_volume)
for _ in range(self.num_levels-1):
init_corr = F.avg_pool2d(init_corr, [1,2], stride=[1,2])
self.init_corr_pyramid.append(init_corr)
def __call__(self, disp, coords, dx, low_memory=True):
b, _, h, w = disp.shape
out_pyramid = []
for i in range(self.num_levels):
with torch.profiler.record_function(f"make disp_lvl {i}"):
geo_volume = self.geo_volume_pyramid[i]
x0 = dx + disp.view(b*h*w, 1, 1, 1) / 2**i
with torch.profiler.record_function(f"bilinear_sampler geo_volume {i}"):
if low_memory:
geo_volume = bilinear_sampler1d(geo_volume, x0, mode='bilinear', align_corners=True)
else:
y0 = torch.zeros_like(x0)
disp_lvl = torch.cat([x0,y0], dim=-1)
geo_volume = bilinear_sampler(geo_volume, disp_lvl, low_memory=low_memory)
geo_volume = geo_volume.view(b, h, w, -1) #(b, h, h, 3x3xC)
with torch.profiler.record_function(f"make init_coords_lvl {i}"):
init_corr = self.init_corr_pyramid[i] # (B*H*W, 1, 1, W2)
init_x0 = coords.view(b*h*w, 1, 1, 1)/2**i - disp.view(b*h*w, 1, 1, 1) / 2**i + dx # X on right image
with torch.profiler.record_function(f"bilinear_sampler init_corr {i}"):
if low_memory:
init_corr = bilinear_sampler1d(init_corr, init_x0, mode='bilinear', align_corners=True)
else:
init_coords_lvl = torch.cat([init_x0,y0], dim=-1)
init_corr = bilinear_sampler(init_corr, init_coords_lvl, low_memory=low_memory)
init_corr = init_corr.view(b, h, w, -1)
out_pyramid.append(geo_volume)
out_pyramid.append(init_corr)
with torch.profiler.record_function(f"make out_pyramid"):
out_pyramid = torch.cat(out_pyramid, dim=-1)
return out_pyramid.permute(0, 3, 1, 2) #(B,C,H,W)
@staticmethod
def corr(fmap1, fmap2, normalize=True):
with torch.profiler.record_function("build corr"):
B, D, H, W1 = fmap1.shape
_, _, _, W2 = fmap2.shape
fmap1 = fmap1.view(B, D, H, W1)
fmap2 = fmap2.view(B, D, H, W2)
if normalize:
with torch.cuda.amp.autocast(enabled=False):
corr = torch.einsum('aijk,aijh->ajkh', F.normalize(fmap1.float(), dim=1), F.normalize(fmap2.float(), dim=1))
else:
corr = corr.view(B, H, W1, 1, W2).to(fmap1.dtype)
corr = corr.view(B, H, W1, 1, W2).to(fmap1.dtype)
return corr
+675
View File
@@ -0,0 +1,675 @@
import torch,pdb,os,sys
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
from Utils import AMP_DTYPE
import Utils as U
try:
import triton
import triton.language as tl
except Exception:
triton = None
tl = None
def _is_contiguous(tensor: torch.Tensor) -> bool:
if torch.jit.is_scripting():
return tensor.is_contiguous()
else:
return tensor.is_contiguous(memory_format=torch.contiguous_format)
class LayerNorm2d(nn.LayerNorm):
r""" https://huggingface.co/spaces/Roll20/pet_score/blob/b258ef28152ab0d5b377d9142a23346f863c1526/lib/timm/models/convnext.py#L85
LayerNorm for channels_first tensors with 2d spatial dimensions (ie N, C, H, W).
"""
def __init__(self, normalized_shape, eps=1e-6):
"""
@normalized_shape: channel dim
"""
super().__init__(normalized_shape, eps=eps)
def forward(self, x) -> torch.Tensor:
"""
@x: (B,C,H,W)
"""
if _is_contiguous(x):
return F.layer_norm(x.permute(0, 2, 3, 1), self.normalized_shape, self.weight, self.bias, self.eps).permute(0, 3, 1, 2).contiguous()
else:
s, u = torch.var_mean(x, dim=1, keepdim=True)
x = (x - u) * torch.rsqrt(s + self.eps)
x = x * self.weight[:, None, None] + self.bias[:, None, None]
return x
class BasicConv(nn.Module):
def __init__(self, in_channels, out_channels, deconv=False, is_3d=False, bn=True, relu=True, norm='batch', **kwargs):
super(BasicConv, self).__init__()
self.relu = nn.LeakyReLU(inplace=True) if relu else nn.Identity()
self.use_bn = bn
self.bn = nn.Identity()
if is_3d:
if deconv:
self.conv = nn.ConvTranspose3d(in_channels, out_channels, bias=False, **kwargs)
else:
self.conv = nn.Conv3d(in_channels, out_channels, bias=False, **kwargs)
if self.use_bn:
if norm=='batch':
self.bn = nn.BatchNorm3d(out_channels)
elif norm=='instance':
self.bn = nn.InstanceNorm3d(out_channels)
else:
if deconv:
self.conv = nn.ConvTranspose2d(in_channels, out_channels, bias=False, **kwargs)
else:
self.conv = nn.Conv2d(in_channels, out_channels, bias=False, **kwargs)
if self.use_bn:
if norm=='batch':
self.bn = nn.BatchNorm2d(out_channels)
elif norm=='instance':
self.bn = nn.InstanceNorm2d(out_channels)
def forward(self, x):
x = self.conv(x)
if self.use_bn:
x = self.bn(x)
if isinstance(self.relu, bool):
if self.relu:
self.relu = nn.LeakyReLU(inplace=True)
else:
self.relu = nn.Identity()
x = self.relu(x)
return x
class Conv3dNormActReduced(nn.Module):
def __init__(self, C_in, C_out, hidden=None, kernel_size=3, kernel_disp=None, stride=1, norm=nn.BatchNorm3d):
super().__init__()
if kernel_disp is None:
kernel_disp = kernel_size
if hidden is None:
hidden = C_out
self.conv1 = nn.Sequential(
nn.Conv3d(C_in, hidden, kernel_size=(1,kernel_size,kernel_size), padding=(0, kernel_size//2, kernel_size//2), stride=(1, stride, stride)),
norm(hidden),
nn.ReLU(),
)
self.conv2 = nn.Sequential(
nn.Conv3d(hidden, C_out, kernel_size=(kernel_disp, 1, 1), padding=(kernel_disp//2, 0, 0), stride=(stride, 1, 1)),
norm(C_out),
nn.ReLU(),
)
def forward(self, x):
"""
@x: (B,C,D,H,W)
"""
x = self.conv1(x)
x = self.conv2(x)
return x
class ResnetBasicBlock(nn.Module):
def __init__(self, inplanes, planes, kernel_size=3, stride=1, padding=1, downsample=None, groups=1, base_width=64, dilation=1, norm_layer=nn.BatchNorm2d, bias=False):
super().__init__()
self.norm_layer = norm_layer
if groups != 1 or base_width != 64:
raise ValueError('BasicBlock only supports groups=1 and base_width=64')
if dilation > 1:
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
# Both self.conv1 and self.downsample layers downsample the input when stride != 1
self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=kernel_size, stride=stride, bias=bias, padding=padding)
if self.norm_layer is not None:
self.bn1 = norm_layer(planes)
self.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=kernel_size, bias=bias, padding=padding)
if self.norm_layer is not None:
self.bn2 = norm_layer(planes)
self.downsample = downsample
self.stride = stride
def forward(self, x):
identity = x
out = self.conv1(x)
if self.norm_layer is not None:
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
if self.norm_layer is not None:
out = self.bn2(out)
if self.downsample is not None:
identity = self.downsample(x)
out += identity
out = self.relu(out)
return out
class ResnetBasicBlock3D(nn.Module):
def __init__(self, inplanes, planes, kernel_size=3, stride=1, padding=1, downsample=None, groups=1, base_width=64, dilation=1, norm_layer=nn.BatchNorm3d, bias=False):
super().__init__()
self.norm_layer = norm_layer
if groups != 1 or base_width != 64:
raise ValueError('BasicBlock only supports groups=1 and base_width=64')
if dilation > 1:
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
# Both self.conv1 and self.downsample layers downsample the input when stride != 1
self.conv1 = nn.Conv3d(inplanes, planes, kernel_size=kernel_size, stride=stride, bias=bias, padding=padding)
if self.norm_layer is not None:
self.bn1 = norm_layer(planes)
self.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv3d(planes, planes, kernel_size=kernel_size, bias=bias, padding=padding)
if self.norm_layer is not None:
self.bn2 = norm_layer(planes)
self.downsample = downsample
self.stride = stride
def forward(self, x):
identity = x
out = self.conv1(x)
if self.norm_layer is not None:
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
if self.norm_layer is not None:
out = self.bn2(out)
if self.downsample is not None:
identity = self.downsample(x)
out += identity
out = self.relu(out)
return out
class FlashMultiheadAttention(nn.Module):
def __init__(self, embed_dim, num_heads):
super().__init__()
self.num_heads = num_heads
self.embed_dim = embed_dim
self.head_dim = embed_dim // num_heads
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
def forward(self, query, key, value, attn_mask=None, window_size=(-1,-1)):
"""
@query: (B,L,C)
"""
B,L,C = query.shape
Q = self.q_proj(query)
K = self.k_proj(key)
V = self.v_proj(value)
Q = Q.view(Q.size(0), Q.size(1), self.num_heads, self.head_dim)
K = K.view(K.size(0), K.size(1), self.num_heads, self.head_dim)
V = V.view(V.size(0), V.size(1), self.num_heads, self.head_dim)
attn_output = F.scaled_dot_product_attention(Q, K, V)
attn_output = attn_output.reshape(B,L,-1)
output = self.out_proj(attn_output)
return output
class FlashAttentionTransformerEncoderLayer(nn.Module):
def __init__(self, embed_dim, num_heads, dim_feedforward, dropout=0.1, act=nn.GELU, norm=nn.LayerNorm):
super().__init__()
self.self_attn = FlashMultiheadAttention(embed_dim, num_heads)
self.act = act()
self.linear1 = nn.Linear(embed_dim, dim_feedforward)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(dim_feedforward, embed_dim)
self.norm1 = norm(embed_dim)
self.norm2 = norm(embed_dim)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
def forward(self, src, src_mask=None, window_size=(-1, -1)):
dtype = src.dtype
src2 = self.self_attn(src, src, src, src_mask, window_size=window_size)
src = src + self.dropout1(src2)
src = self.norm1(src).to(dtype)
src2 = self.linear2(self.dropout(self.act(self.linear1(src))))
src = src + self.dropout2(src2)
src = self.norm2(src).to(dtype)
return src
class Conv2x(nn.Module):
def __init__(self, in_channels, out_channels, deconv=False, is_3d=False, concat=True, keep_concat=True, bn=True, relu=True, keep_dispc=False):
super(Conv2x, self).__init__()
self.concat = concat
self.is_3d = is_3d
if deconv and is_3d:
kernel = (4, 4, 4)
elif deconv:
kernel = 4
else:
kernel = 3
if deconv and is_3d and keep_dispc:
kernel = (1, 4, 4)
stride = (1, 2, 2)
padding = (0, 1, 1)
self.conv1 = BasicConv(in_channels, out_channels, deconv, is_3d, bn=bn, relu=True, kernel_size=kernel, stride=stride, padding=padding)
else:
self.conv1 = BasicConv(in_channels, out_channels, deconv, is_3d, bn=bn, relu=True, kernel_size=kernel, stride=2, padding=1)
if self.concat:
mul = 2 if keep_concat else 1
self.conv2 = BasicConv(out_channels*2, out_channels*mul, False, is_3d, bn, relu, kernel_size=3, stride=1, padding=1)
else:
self.conv2 = BasicConv(out_channels, out_channels, False, is_3d, bn, relu, kernel_size=3, stride=1, padding=1)
def forward(self, x, rem):
x = self.conv1(x)
if x.shape != rem.shape:
x = F.interpolate(x, size=(rem.shape[-2], rem.shape[-1]), mode='bilinear')
if self.concat:
x = torch.cat((x, rem), 1)
else:
x = x + rem
x = self.conv2(x)
return x
class BasicConv_IN(nn.Module):
def __init__(self, in_channels, out_channels, deconv=False, is_3d=False, IN=True, relu=True, **kwargs):
super(BasicConv_IN, self).__init__()
if relu:
self.relu = nn.LeakyReLU(inplace=True)
else:
self.relu = nn.Identity()
self.use_in = IN
if is_3d:
if deconv:
self.conv = nn.ConvTranspose3d(in_channels, out_channels, bias=False, **kwargs)
else:
self.conv = nn.Conv3d(in_channels, out_channels, bias=False, **kwargs)
self.IN = nn.InstanceNorm3d(out_channels)
else:
if deconv:
self.conv = nn.ConvTranspose2d(in_channels, out_channels, bias=False, **kwargs)
else:
self.conv = nn.Conv2d(in_channels, out_channels, bias=False, **kwargs)
self.IN = nn.InstanceNorm2d(out_channels)
def forward(self, x):
x = self.conv(x)
if self.use_in:
x = self.IN(x)
if isinstance(self.relu, bool):
if self.relu:
self.relu = nn.LeakyReLU(inplace=True)
else:
self.relu = nn.Identity()
x = self.relu(x)
return x
class Conv2x_IN(nn.Module):
def __init__(self, in_channels, out_channels, c_middle=None, deconv=False, is_3d=False, concat=True, keep_concat=True, IN=True, relu=True, keep_dispc=False):
super(Conv2x_IN, self).__init__()
self.concat = concat
self.is_3d = is_3d
if deconv and is_3d:
kernel = (4, 4, 4)
elif deconv:
kernel = 4
else:
kernel = 3
if c_middle is None:
c_middle = out_channels
if deconv and is_3d and keep_dispc:
kernel = (1, 4, 4)
stride = (1, 2, 2)
padding = (0, 1, 1)
self.conv1 = BasicConv_IN(in_channels, c_middle, deconv, is_3d, IN=True, relu=True, kernel_size=kernel, stride=stride, padding=padding)
else:
self.conv1 = BasicConv_IN(in_channels, c_middle, deconv, is_3d, IN=True, relu=True, kernel_size=kernel, stride=2, padding=1)
if self.concat:
mul = 2 if keep_concat else 1
self.conv2 = ResnetBasicBlock(out_channels*2, out_channels*mul, kernel_size=3, stride=1, padding=1, norm_layer=nn.InstanceNorm2d)
else:
self.conv2 = BasicConv_IN(c_middle, out_channels, False, is_3d, IN, relu, kernel_size=3, stride=1, padding=1)
def forward(self, x, rem):
x = self.conv1(x)
if x.shape != rem.shape:
x = F.interpolate(x, size=(rem.shape[-2], rem.shape[-1]), mode='bilinear')
if self.concat:
x = torch.cat((x, rem), 1)
else:
x = x + rem
x = self.conv2(x)
return x
@torch.compile
def build_gwc_volume_optimized_pytorch1(refimg_fea: torch.Tensor, targetimg_fea: torch.Tensor, maxdisp: int, num_groups: int, normalize=True):
dtype = refimg_fea.dtype
B, C, H, W = refimg_fea.shape
channels_per_group = C // num_groups
ref_volume = refimg_fea.unsqueeze(2).expand(B, C, maxdisp, H, W)
padded_target = F.pad(targetimg_fea, (maxdisp - 1, 0, 0, 0))
unfolded_target = padded_target.unfold(3, W, 1)
target_volume = torch.flip(unfolded_target, [3]).permute(0, 1, 3, 2, 4)
ref_volume = ref_volume.view(B, num_groups, channels_per_group, maxdisp, H, W)
target_volume = target_volume.view(B, num_groups, channels_per_group, maxdisp, H, W)
if normalize:
ref_volume = F.normalize(ref_volume.float(), dim=2).to(dtype)
target_volume = F.normalize(target_volume.float(), dim=2).to(dtype)
cost_volume = (ref_volume * target_volume).sum(dim=2)
return cost_volume.contiguous()
if triton is not None and torch.cuda.is_available():
@triton.autotune(configs=[
triton.Config({'BLOCK_C':4,'BLOCK_W':128,'BLOCK_D':8}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_C':8,'BLOCK_W':128,'BLOCK_D':8}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_C':16,'BLOCK_W':128,'BLOCK_D':8}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_C':64,'BLOCK_W':128,'BLOCK_D':8}, num_warps=8, num_stages=2),
triton.Config({'BLOCK_C':128,'BLOCK_W':64,'BLOCK_D':8}, num_warps=8, num_stages=2),
triton.Config({'BLOCK_C':128,'BLOCK_W':128,'BLOCK_D':8}, num_warps=8, num_stages=2),
], key=['C','W','D','G','K','NORMALIZE'])
@triton.jit
def _gwc_triton_kernel(ref_ptr, tar_ptr, ref_norm_ptr, tar_norm_ptr, out_ptr, BH, C, W, D: tl.constexpr, G: tl.constexpr, K: tl.constexpr,
stride_rn, stride_rw, stride_rc, stride_tn, stride_tw, stride_tc,
stride_nn, stride_ng, stride_nw,
stride_on, stride_og, stride_od, stride_ow,
NORMALIZE: tl.constexpr,
BLOCK_C: tl.constexpr, BLOCK_W: tl.constexpr, BLOCK_D: tl.constexpr):
pid0 = tl.program_id(0)
db = tl.program_id(1)
wb = tl.program_id(2)
bh = pid0 // G
g = pid0 % G
w_off = wb*BLOCK_W + tl.arange(0, BLOCK_W)
d_off = db*BLOCK_D + tl.arange(0, BLOCK_D)
w_mask = w_off < W
w_src = w_off[None, :] - d_off[:, None]
td_mask = (w_src >= 0) & w_mask[None, :]
acc = tl.zeros((BLOCK_D, BLOCK_W), dtype=tl.float32)
for k0 in tl.static_range(0, K, BLOCK_C):
k_off = k0 + tl.arange(0, BLOCK_C)
k_mask = k_off < K
c_idx = g*K + k_off
ref_ptrs = ref_ptr + bh*stride_rn + w_off[None, :]*stride_rw + c_idx[:, None]*stride_rc
ref_vals = tl.load(ref_ptrs, mask=k_mask[:, None] & w_mask[None, :], other=0.).to(tl.float32)
tar_ptrs = tar_ptr + bh*stride_tn + w_src[None, :, :]*stride_tw + c_idx[:, None, None]*stride_tc
tar_vals = tl.load(tar_ptrs, mask=k_mask[:, None, None] & td_mask[None, :, :], other=0.).to(tl.float32)
acc += tl.sum(tar_vals * ref_vals[:, None, :], axis=0)
if NORMALIZE:
norm_offset = bh*stride_nn + g*stride_ng
ref_norm = tl.load(ref_norm_ptr + norm_offset + w_off*stride_nw, mask=w_mask, other=1.0).to(tl.float32)
tar_norm = tl.load(tar_norm_ptr + norm_offset + w_src*stride_nw, mask=td_mask, other=1.0).to(tl.float32)
denom = (ref_norm[None, :] * tar_norm) + 1e-5
acc = acc / denom
out_ptrs = out_ptr + bh*stride_on + g*stride_og + d_off[:, None]*stride_od + w_off[None, :]*stride_ow
tl.store(out_ptrs, acc, mask=w_mask[None, :])
@torch.no_grad()
def build_gwc_volume_triton(refimg_fea: torch.Tensor, targetimg_fea: torch.Tensor, maxdisp: int, num_groups: int, normalize=True):
if triton is None:
raise RuntimeError('Triton is not available. Please install triton to use build_gwc_volume_triton.')
B, C, H, W = refimg_fea.shape
assert maxdisp > 0 and C % num_groups == 0
K = C // num_groups
in_dtype = refimg_fea.dtype if refimg_fea.dtype in (torch.float16, torch.bfloat16, torch.float32) else torch.float32
if normalize:
ref_norm = refimg_fea.float().view(B, num_groups, K, H, W).norm(dim=2)
tar_norm = targetimg_fea.float().view(B, num_groups, K, H, W).norm(dim=2)
ref_norm = ref_norm.permute(0, 2, 1, 3).reshape(B*H, num_groups, W).to(in_dtype).contiguous()
tar_norm = tar_norm.permute(0, 2, 1, 3).reshape(B*H, num_groups, W).to(in_dtype).contiguous()
else:
# Dummy tensors; kernel won't read them when NORMALIZE=False
ref_norm = refimg_fea.new_empty((1, 1, 1), dtype=in_dtype)
tar_norm = refimg_fea.new_empty((1, 1, 1), dtype=in_dtype)
ref = refimg_fea.to(in_dtype)
tar = targetimg_fea.to(in_dtype)
ref_bhwc = ref.permute(0, 2, 3, 1).view(B * H, W, C).contiguous()
tar_bhwc = tar.permute(0, 2, 3, 1).view(B * H, W, C).contiguous()
out_bhw = torch.empty((B * H, num_groups, maxdisp, W), device=ref.device, dtype=in_dtype)
BH = B * H
D_eff = min(maxdisp, W)
grid = lambda META: (BH * num_groups, triton.cdiv(D_eff, META['BLOCK_D']), triton.cdiv(W, META['BLOCK_W']))
_gwc_triton_kernel[grid](ref_bhwc, tar_bhwc, ref_norm, tar_norm, out_bhw, BH, C, W, D_eff, num_groups, K,
ref_bhwc.stride(0), ref_bhwc.stride(1), ref_bhwc.stride(2),
tar_bhwc.stride(0), tar_bhwc.stride(1), tar_bhwc.stride(2),
ref_norm.stride(0), ref_norm.stride(1), ref_norm.stride(2),
out_bhw.stride(0), out_bhw.stride(1), out_bhw.stride(2), out_bhw.stride(3),
NORMALIZE=normalize)
if D_eff < maxdisp: out_bhw[:, :, D_eff:, :] = 0
volume = out_bhw.view(B, H, num_groups, maxdisp, W).permute(0, 2, 3, 1, 4).contiguous()
return volume
@torch.compile
def build_concat_volume_optimized_pytorch(refimg_fea, targetimg_fea, maxdisp:int):
B, C, H, W = refimg_fea.shape
ref_volume = refimg_fea.unsqueeze(2).expand(B, C, maxdisp, H, W)
shifted_target_list = [F.pad(targetimg_fea, (int(d), 0, 0, 0), "constant", 0.0)[:, :, :, :W] for d in range(maxdisp)]
target_volume = torch.stack(shifted_target_list, dim=2)
volume = torch.cat((ref_volume, target_volume), dim=1)
return volume.contiguous()
@torch.compile
def build_concat_volume_optimized_pytorch1(refimg_fea, targetimg_fea, maxdisp:int):
B, C, H, W = refimg_fea.shape
ref_volume = refimg_fea.unsqueeze(2).expand(B, C, maxdisp, H, W)
padded_target = F.pad(targetimg_fea, (maxdisp - 1, 0, 0, 0)) # (B, C, H, W + maxdisp - 1)
unfolded_target = padded_target.unfold(dimension=3, size=W, step=1) # (B, C, H, maxdisp, W)
target_volume = torch.flip(unfolded_target, [3]).permute(0, 1, 3, 2, 4)
volume = torch.cat((ref_volume, target_volume), dim=1)
return volume.contiguous()
def disparity_regression(x, maxdisp):
assert len(x.shape) == 4
disp_values = torch.arange(0, maxdisp, dtype=x.dtype, device=x.device)
disp_values = disp_values.reshape(1, maxdisp, 1, 1)
return torch.sum(x * disp_values, 1, keepdim=True) #(B,1,H,W)
class FeatureAtt(nn.Module):
def __init__(self, cv_chan, feat_chan):
super(FeatureAtt, self).__init__()
self.feat_att = nn.Sequential(
BasicConv(feat_chan, feat_chan//2, kernel_size=1, stride=1, padding=0),
nn.Conv2d(feat_chan//2, cv_chan, 1)
)
def forward(self, cv, feat):
'''
@cv: cost volume (B,C,D,H,W)
@feat: (B,C,H,W)
'''
feat_att = self.feat_att(feat).unsqueeze(2) #(B,C,1,H,W)
cv = torch.sigmoid(feat_att)*cv
return cv
def context_upsample(disp_low, up_weights):
"""
@disp_low: (b,1,h,w) 1/4 resolution
@up_weights: (b,9,4*h,4*w) Image resolution
"""
b, c, h, w = disp_low.shape
disp_unfold = F.unfold(disp_low.reshape(b,c,h,w),3,1,1).reshape(b,-1,h,w)
disp_unfold = F.interpolate(disp_unfold,(h*4,w*4),mode='nearest').reshape(b,9,h*4,w*4)
disp = (disp_unfold*up_weights).sum(1)
return disp
class PositionalEmbedding(nn.Module):
def __init__(self, d_model, max_len=512):
super().__init__()
# Compute the positional encodings once in log space.
pe = torch.zeros(max_len, d_model, dtype=torch.float)
pe.require_grad = False
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) #(N,1)
div_term = (torch.arange(0, d_model, 2, dtype=torch.float) * -(np.log(10000.0) / d_model)).exp()[None]
pe[:, 0::2] = torch.sin(position * div_term) #(N, d_model/2)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0)
self.pe = pe
def forward(self, x, resize_embed=False):
'''
@x: (B,N,D)
'''
dtype = x.dtype
self.pe = self.pe.to(x.device).to(x.dtype)
pe = self.pe
if pe.shape[1]<x.shape[1]:
if resize_embed:
pe = F.interpolate(pe.permute(0,2,1), size=x.shape[1], mode='linear', align_corners=True).permute(0,2,1)
else:
raise RuntimeError(f'x:{x.shape}, pe:{pe.shape}')
return (x + pe[:, :x.size(1)]).to(dtype)
class CostVolumeDisparityAttention(nn.Module):
def __init__(self, d_model, nhead, dim_feedforward, dropout=0.1, act=nn.GELU, norm_first=False, num_transformer=6, max_len=512, resize_embed=False):
super().__init__()
self.resize_embed = resize_embed
self.sa = nn.ModuleList([])
for _ in range(num_transformer):
self.sa.append(FlashAttentionTransformerEncoderLayer(embed_dim=d_model, num_heads=nhead, dim_feedforward=dim_feedforward, act=act, dropout=dropout))
self.pos_embed0 = PositionalEmbedding(d_model, max_len=max_len)
def forward(self, cv, window_size=(-1,-1)):
"""
@cv: (B,C,D,H,W) where D is max disparity
"""
x = cv
B,C,D,H,W = x.shape
x = x.permute(0,3,4,2,1).reshape(B*H*W, D, C)
x = self.pos_embed0(x, resize_embed=self.resize_embed) #!NOTE No resize since disparity is pre-determined
for i in range(len(self.sa)):
x = self.sa[i](x, window_size=window_size)
x = x.reshape(B,H,W,D,C).permute(0,4,3,1,2)
return x
class ChannelAttentionEnhancement(nn.Module):
def __init__(self, in_planes, ratio=16):
"""From selective-IGEV
"""
super(ChannelAttentionEnhancement, self).__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.max_pool = nn.AdaptiveMaxPool2d(1)
self.fc = nn.Sequential(nn.Conv2d(in_planes, in_planes // 16, 1, bias=False),
nn.ReLU(),
nn.Conv2d(in_planes // 16, in_planes, 1, bias=False))
self.sigmoid = nn.Sigmoid()
def forward(self, x):
avg_out = self.fc(self.avg_pool(x))
max_out = self.fc(self.max_pool(x))
out = avg_out + max_out
return self.sigmoid(out)
class SpatialAttentionExtractor(nn.Module):
def __init__(self, kernel_size=7):
"""From selective-IGEV
"""
super(SpatialAttentionExtractor, self).__init__()
self.samconv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2, bias=False)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
avg_out = torch.mean(x, dim=1, keepdim=True)
max_out, _ = torch.max(x, dim=1, keepdim=True)
x = torch.cat([avg_out, max_out], dim=1)
x = self.samconv(x)
return self.sigmoid(x)
class EdgeNextConvEncoder(nn.Module):
def __init__(self, dim, layer_scale_init_value=1e-6, expan_ratio=4, kernel_size=7, norm='layer'):
"""https://github.com/mmaaz60/EdgeNeXt/blob/main/models/conv_encoder.py#L7
"""
super().__init__()
self.dwconv = nn.Conv2d(dim, dim, kernel_size=kernel_size, padding=kernel_size // 2, groups=dim)
if norm=='layer':
self.norm = LayerNorm2d(dim, eps=1e-6)
elif norm=='batch':
self.norm = nn.BatchNorm2d(dim)
else:
self.norm = nn.Identity()
self.pwconv1 = nn.Linear(dim, expan_ratio * dim)
self.act = nn.GELU()
self.pwconv2 = nn.Linear(expan_ratio * dim, dim)
self.gamma = nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True) if layer_scale_init_value > 0 else None
def forward(self, x):
input = x
x = self.dwconv(x)
x = self.norm(x)
x = x.permute(0, 2, 3, 1) # (N, C, H, W) -> (N, H, W, C)
x = self.pwconv1(x)
x = self.act(x)
x = self.pwconv2(x)
if self.gamma is not None:
x = self.gamma * x
x = x.permute(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
x = input + x
return x
+108
View File
@@ -0,0 +1,108 @@
import torch,os,sys
import torch.nn as nn
import torch.nn.functional as F
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
from core.submodule import EdgeNextConvEncoder
class DispHead(nn.Module):
def __init__(self, input_dim=128, hidden_dim=256, output_dim=1):
super(DispHead, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(input_dim, input_dim, kernel_size=3, padding=1),
nn.ReLU(),
EdgeNextConvEncoder(input_dim, expan_ratio=4, kernel_size=7, norm=None),
EdgeNextConvEncoder(input_dim, expan_ratio=4, kernel_size=7, norm=None),
nn.Conv2d(input_dim, output_dim, 3, padding=1),
)
def forward(self, x):
return self.conv(x)
class BasicMotionEncoder(nn.Module):
def __init__(self, args, ngroup=8):
super(BasicMotionEncoder, self).__init__()
self.args = args
cor_planes = args.corr_levels * (2*args.corr_radius + 1) * (ngroup+1)
self.convc1 = nn.Conv2d(cor_planes, 256, kernel_size=1, padding=0)
self.convc2 = nn.Conv2d(256, 256, kernel_size=3, padding=1)
self.convd1 = nn.Conv2d(1, 64, kernel_size=7, padding=3)
self.convd2 = nn.Conv2d(64, 64, kernel_size=3, padding=1)
self.conv = nn.Conv2d(64+256, args.hidden_dims[0]-1, kernel_size=1, padding=0)
def forward(self, disp, corr):
cor = F.relu(self.convc1(corr))
cor = F.relu(self.convc2(cor))
disp_ = F.relu(self.convd1(disp))
disp_ = F.relu(self.convd2(disp_))
cor_disp = torch.cat([cor, disp_], dim=1)
out = F.relu(self.conv(cor_disp))
return torch.cat([out, disp], dim=1)
class RaftConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=256, kernel_size=3):
super().__init__()
self.convz = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2)
self.convr = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2)
self.convq = nn.Conv2d(hidden_dim+input_dim, hidden_dim, kernel_size, padding=kernel_size // 2)
def forward(self, h, x, hx):
z = torch.sigmoid(self.convz(hx))
r = torch.sigmoid(self.convr(hx))
q = torch.tanh(self.convq(torch.cat([r*h, x], dim=1)))
h = (1-z) * h + z * q
return h
class SelectiveConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=256, small_kernel_size=1, large_kernel_size=3, patch_size=None):
super(SelectiveConvGRU, self).__init__()
self.conv0 = nn.Sequential(
nn.Conv2d(input_dim, input_dim, kernel_size=3, padding=1),
nn.ReLU(),
)
self.conv1 = nn.Sequential(
nn.Conv2d(input_dim+hidden_dim, input_dim+hidden_dim, kernel_size=3, padding=1),
nn.ReLU(),
)
self.small_gru = RaftConvGRU(hidden_dim, input_dim, small_kernel_size)
self.large_gru = RaftConvGRU(hidden_dim, input_dim, large_kernel_size)
def forward(self, att, h, *x):
x = torch.cat(x, dim=1)
x = self.conv0(x)
hx = torch.cat([x, h], dim=1)
hx = self.conv1(hx)
h = self.small_gru(h, x, hx) * att + self.large_gru(h, x, hx) * (1 - att)
return h
class BasicSelectiveMultiUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128, volume_dim=8):
super().__init__()
self.args = args
self.encoder = BasicMotionEncoder(args, volume_dim)
self.gru04 = SelectiveConvGRU(hidden_dim, hidden_dim*2)
self.disp_head = DispHead(hidden_dim, 256)
self.mask = nn.Sequential(
nn.Conv2d(hidden_dim, 64, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(64, 32, 3, padding=1),
nn.ReLU(inplace=True),
)
def forward(self, net, inp, corr, disp, att):
motion_features = self.encoder(disp, corr)
motion_features = torch.cat([inp[0], motion_features], dim=1)
net[0] = self.gru04(att[0], net[0], motion_features)
delta_disp = self.disp_head(net[0])
mask = .25 * self.mask(net[0])
return net, mask, delta_disp
View File
+202
View File
@@ -0,0 +1,202 @@
import numpy as np
from PIL import Image
from os.path import basename, exists, splitext
import os,sys
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../../')
import re
import json
import imageio
import cv2
from turbojpeg import TurboJPEG, TJPF_GRAY, TJSAMP_GRAY, TJFLAG_PROGRESSIVE, TJFLAG_FASTUPSAMPLE, TJFLAG_FASTDCT
jpeg = TurboJPEG()
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
# 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 * 4 + d_g / (2**6) + d_b / (2**14))[..., 0]
mask = np.array(Image.open(file_name.replace('disparities', 'occlusions')))
valid = ((mask == 0) & (disp > 0))
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)
fx = intrinsics['camera_settings'][0]['intrinsic_settings']['fx']
disp = (fx * 6.0 * 100) / a.astype(np.float32)
valid = disp > 0
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 > 0
return disp, valid
def readDispMiddlebury(file_name):
assert 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
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 in ['.jpeg','.jpg']:
with open(file_name, 'rb') as ff:
bgr_array = jpeg.decode(ff.read())
img = bgr_array[...,::-1]
return img
elif ext == '.png' or ext == '.ppm':
img = cv2.imread(file_name)[...,:3]
if len(img.shape)==3:
img = img[...,::-1]
elif len(img.shape)==2:
img = np.tile(img[...,None], (1,1,3))
return img
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]
return []
+122
View File
@@ -0,0 +1,122 @@
import torch,os,sys
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../../')
import torch.nn.functional as F
import numpy as np
class InputPadder:
""" Pads images such that dimensions are divisible by 8 """
def __init__(self, dims, mode='sintel', divis_by=8, force_square=False):
self.ht, self.wd = dims[-2:]
if force_square:
max_side = max(self.ht, self.wd)
pad_ht = ((max_side // divis_by) + 1) * divis_by - self.ht
pad_wd = ((max_side // divis_by) + 1) * divis_by - self.wd
else:
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]]
@torch.compile
def bilinear_sampler1d(img, x_coords, mode='bilinear', align_corners=True):
"""
1D bilinear sampling along width dimension only (for stereo applications)
Much faster than grid_sample for stereo where y is constant
Args:
img: (B, C, 1, W) input tensor
x_coords: (B, 1, W_out, 1) x coordinates in pixel space [0, W-1]
mode: interpolation mode ('bilinear' or 'nearest')
align_corners: if True, corner pixels are aligned (like grid_sample)
Returns:
sampled: (B, C, 1, W_coords) sampled tensor
mask: (B, 1, H, W) validity mask (if mask=True)
"""
B, C, H_img, W = img.shape
x = x_coords.reshape(B,-1) # (B, W_out)
if align_corners:
# align_corners=True: coordinate range [0, W-1] maps to pixel centers
# This matches grid_sample with align_corners=True behavior
x_normalized = x
else:
# align_corners=False: coordinate range [0, W-1] maps to pixel edges
# Need to adjust coordinates to match grid_sample with align_corners=False
# grid_sample maps [-1, 1] to [0, W-1] when align_corners=False
# So our [0, W-1] input should be treated as [0.5, W-0.5] in pixel space
x_normalized = x + 0.5
if mode == 'nearest':
# Nearest neighbor sampling with zero padding outside [0, W-1]
if align_corners:
x_nearest = torch.round(x_normalized).long()
else:
x_nearest = torch.floor(x_normalized).long()
valid = (x_nearest >= 0) & (x_nearest < W) # (B, W_out)
x_index = torch.clamp(x_nearest, 0, W-1)
sampled = torch.gather(img, 3, x_index.view(B,1,1,-1).expand(B,C,1,-1))
sampled = sampled * valid.view(B,1,1,-1).to(img.dtype)
else: # bilinear
# Get integer and fractional parts
x_floor = torch.floor(x_normalized)
x_ceil = x_floor + 1
x_frac = x_normalized - x_floor # (B, W_out)
# Zero padding behavior: mark validity and zero-out invalid contributions
valid_floor = (x_floor >= 0) & (x_floor < W)
valid_ceil = (x_ceil >= 0) & (x_ceil < W)
x_floor_clamped = torch.clamp(x_floor, 0, W-1)
x_ceil_clamped = torch.clamp(x_ceil, 0, W-1)
# Create index tensors
batch_idx = torch.arange(B, device=img.device).view(B, 1)
img_floor = torch.gather(img, 3, x_floor_clamped.view(B,1,1,-1).expand(B,C,1,-1).long())
img_ceil = torch.gather(img, 3, x_ceil_clamped.view(B,1,1,-1).expand(B,C,1,-1).long())
# Apply validity masks (zero out-of-bounds samples)
img_floor = img_floor * valid_floor.view(B,1,1,-1).to(img.dtype)
img_ceil = img_ceil * valid_ceil.view(B,1,1,-1).to(img.dtype)
# Linear interpolation
x_frac = x_frac.view(B,1,1,-1)
sampled = img_floor * (1 - x_frac) + img_ceil * x_frac
return sampled
def bilinear_sampler(img, coords, mode='bilinear', mask=False, low_memory=False, use1d=False):
""" Wrapper for grid_sample, uses pixel coordinates """
H, W = img.shape[-2:]
coords[...,0] = 2*coords[...,0]/(W-1) - 1
if low_memory:
B = img.shape[0]
out = []
bs = 102400
for b in np.arange(0,B,bs):
tmp = F.grid_sample(img[b:b+bs], coords[b:b+bs], align_corners=True)
out.append(tmp)
img = torch.cat(out, dim=0)
else:
img = F.grid_sample(img, coords, align_corners=True)
if mask:
mask = (xgrid > -1) & (ygrid > -1) & (xgrid < 1) & (ygrid < 1)
return img, mask.float()
return img