Reloc3r: 大规模相对相机姿态回归训练实现泛化、快速、精确的视觉定位
1. 简介
视觉定位是计算机视觉中的一个重要任务,它能够确定一张照片在三维空间中的精确位置和方向。想象一下,当你用手机拍照时,如果能够知道这张照片是在哪里拍摄的,以及相机当时朝向哪个方向,这就是视觉定位要解决的问题。

Reloc3r是一个革命性的视觉定位框架,它通过深度学习技术,能够在几毫秒内准确估计相机的6自由度姿态(位置和方向)。这个方法的独特之处在于,它不需要针对特定场景进行训练,可以泛化到全新的环境中,同时保持了极高的精度和速度。

论文链接: https://arxiv.org/abs/2412.08376
代码仓库: https://github.com/ffrivera0/reloc3r
2. 研究背景与挑战
传统的视觉定位方法主要依赖结构运动(SfM)技术,通过重建三维模型来定位相机。这种方法虽然精度很高,但计算过程复杂,需要大量的时间和计算资源,难以在实时应用中使用。近年来,深度学习方法的出现为视觉定位带来了新的可能性。绝对姿态回归(APR)方法可以直接从图像中预测相机姿态,速度很快,但通常只能应用于训练过的特定场景。相对姿态回归(RPR)方法虽然具有跨场景泛化能力,但在精度上往往不如APR方法。因此,如何在保持高精度的同时,实现跨场景泛化和实时性能,成为了视觉定位领域的主要挑战。
Reloc3r的核心思想非常巧妙:它不直接预测相机的绝对位置,而是先预测查询图像与已知图像之间的相对关系,然后通过几何计算得到最终的绝对位置。这种方法的好处是,相对关系更容易学习,因为它不依赖于具体的场景尺度。比如,无论你是在室内还是室外,两张相邻照片之间的相对位置关系都是相似的。Reloc3r包含两个主要组件:一个相对姿态回归网络和一个运动平均模块。前者负责学习图像对之间的相对关系,后者负责将这些相对关系转换为绝对位置。

3. 网络架构设计
Reloc3r采用了全对称的Vision Transformer(ViT)架构,这是其设计的一大亮点。全对称意味着网络的两个分支(处理两张图像)共享相同的权重,这样可以消除图像输入顺序带来的偏差。
让我们看看核心的网络架构代码:
class Reloc3rRelpose(nn.Module, PyTorchModelHubMixin):
def __init__(self,
img_size=512, # 输入图像尺寸
patch_size=16, # 图像分块大小
enc_embed_dim=1024, # 编码器特征维度
enc_depth=24, # 编码器层数
enc_num_heads=16, # 注意力头数
dec_embed_dim=768, # 解码器特征维度
dec_depth=12, # 解码器层数
dec_num_heads=12, # 解码器注意力头数
pos_embed='RoPE100', # 位置编码类型
):
这个架构的关键参数包括图像尺寸、分块大小、编码器和解码器的深度等。较大的模型(如512像素输入)通常能提供更好的精度,但需要更多的计算资源。
3.1 图像处理流程
Reloc3r首先将输入图像分割成固定大小的块(patches),每个块的大小为16×16像素。这些块经过线性变换后,被转换为向量序列,作为Transformer的输入。
def _encode_image(self, image, true_shape):
# 将图像分割成patches并嵌入
x, pos = self.patch_embed(image, true_shape=true_shape)
# 通过Transformer编码器处理
for blk in self.enc_blocks:
x = blk(x, pos)
x = self.enc_norm(x)
return x, pos, None
编码器通过多层自注意力机制,学习每个patch的特征表示。位置编码(RoPE)帮助模型理解patch之间的空间关系,这对于后续的姿态估计至关重要。
3.2 交叉注意力机制
解码器是Reloc3r的核心创新之一。它使用交叉注意力机制,让两张图像的特征进行交互,学习它们之间的对应关系。
def _decoder(self, f1, pos1, f2, pos2):
# 投影到解码器维度
f1 = self.decoder_embed(f1)
f2 = self.decoder_embed(f2)
for blk in self.dec_blocks:
# 图像1侧:使用f1作为查询,f2作为键值
f1, _ = blk(*final_output[-1][::+1], pos1, pos2)
# 图像2侧:使用f2作为查询,f1作为键值
f2, _ = blk(*final_output[-1][::-1], pos2, pos1)
final_output.append((f1, f2))
这种设计使得模型能够学习两张图像之间的空间对应关系,为后续的相对姿态估计提供基础。交叉注意力机制让模型能够"看到"两张图像中相同物体的不同视角。
3.3 姿态回归头
经过解码器处理后,模型需要将学习到的特征转换为具体的相机姿态。这是通过姿态回归头实现的:
class PoseHead(nn.Module):
def __init__(self, net, num_resconv_block=2, rot_representation='9D'):
super().__init__()
output_dim = 4*self.patch_size**2
# 投影层
self.proj = nn.Linear(net.dec_embed_dim, output_dim)
# 残差卷积块
self.res_conv = nn.ModuleList([
copy.deepcopy(ResConvBlock(output_dim, output_dim))
for _ in range(self.num_resconv_block)
])
# 全局平均池化
self.avgpool = nn.AdaptiveAvgPool2d(1)
# 多层感知器
self.more_mlps = nn.Sequential(
nn.Linear(output_dim, output_dim),
nn.ReLU(),
nn.Linear(output_dim, output_dim),
nn.ReLU()
)
# 输出层:平移和旋转
self.fc_t = nn.Linear(output_dim, 3) # 平移向量
self.fc_rot = nn.Linear(output_dim, 9) # 旋转矩阵
姿态回归头首先将特征投影到更高维度,然后通过残差卷积块进行特征提取,最后通过全连接层输出平移向量和旋转矩阵。
3.4 旋转表示与SVD正交化
旋转的表示是姿态估计中的一个关键问题。Reloc3r使用9维向量来表示旋转,然后通过SVD(奇异值分解)将其转换为有效的旋转矩阵:
def svd_orthogonalize(self, m):
"""使用SVD将9D表示转换为有效的旋转矩阵"""
if m.dim() < 3:
m = m.reshape((-1, 3, 3))
# 归一化并转置
m_transpose = torch.transpose(
torch.nn.functional.normalize(m, p=2, dim=-1),
dim0=-1, dim1=-2
)
# SVD分解
u, s, v = torch.svd(m_transpose)
# 计算行列式
det = torch.det(torch.matmul(v, u.transpose(-2, -1)))
# 构建旋转矩阵
r = torch.matmul(
torch.cat([v[:, :, :-1], v[:, :, -1:] * det.view(-1, 1, 1)], dim=2),
u.transpose(-2, -1)
)
return r
这个过程确保了输出的旋转矩阵满足旋转矩阵的数学性质(正交且行列式为1),这对于准确的姿态估计至关重要。
3.5 损失函数设计
Reloc3r使用角度损失来训练模型,这种方法比直接使用欧几里得距离更适合姿态估计:
class Reloc3rPoseLoss(nn.Module):
def transl_ang_loss(self, t, tgt, eps=1e-6):
"""平移方向角度损失"""
# 归一化平移向量
t_norm = torch.norm(t, dim=1, keepdim=True)
t_normed = t / (t_norm + eps)
tgt_norm = torch.norm(tgt, dim=1, keepdim=True)
tgt_normed = tgt / (tgt_norm + eps)
# 计算余弦相似度
cosine = torch.sum(t_normed * tgt_normed, dim=1)
T_err = torch.acos(torch.clamp(cosine, -1.0 + eps, 1.0 - eps))
return T_err.mean()
def rot_ang_loss(self, R, Rgt, eps=1e-6):
"""旋转角度损失"""
# 计算相对旋转
residual = torch.matmul(R.transpose(1, 2), Rgt)
trace = torch.diagonal(residual, dim1=-2, dim2=-1).sum(-1)
cosine = (trace - 1) / 2
R_err = torch.acos(torch.clamp(cosine, -1.0 + eps, 1.0 - eps))
return R_err.mean()
平移损失关注方向而非距离,旋转损失通过矩阵迹计算角度差异。这种设计使得模型能够学习到更准确的相对姿态关系。
…详情请参照古月居
更多推荐
所有评论(0)