基于PyTorch的RAFT光流场图像重映射异常问题求助
问题根源与修复方案
核心问题:RAFT光流与grid_sample的坐标空间不匹配
RAFT输出的光流是像素级偏移量(如(Δx, Δy),代表当前帧像素对应参考帧的位置偏移),但grid_sample需要的是归一化到[-1, 1]范围的目标坐标,而非原始偏移量。此外还可能存在光流方向搞反的问题,这两个点是导致对齐后图像异常的关键。
具体修复步骤
1. 光流坐标转换(必须步骤)
修改warp_flow函数,将RAFT的像素偏移转换为grid_sample要求的归一化坐标:
def warp_flow(img, flow): B, C, H, W = img.size() # 创建原始归一化坐标网格 (B, H, W, 2) x = torch.linspace(-1, 1, W, device=img.device).view(1, 1, W).repeat(B, H, 1) y = torch.linspace(-1, 1, H, device=img.device).view(1, H, 1).repeat(B, 1, W) grid = torch.stack([x, y], dim=-1) # 将RAFT输出的像素偏移转为归一化偏移量 # 像素偏移除以图像半宽/半高,映射到[-1,1]范围的偏移 flow_norm = flow.permute(0,2,3,1) / torch.tensor([W/2, H/2], device=img.device) # 原始网格 + 归一化偏移 = 目标采样坐标 grid = grid + flow_norm # 执行网格采样,开启align_corners保证边缘对齐 remapped = torch.nn.functional.grid_sample(img, grid, align_corners=True) return remapped
2. 确认光流方向正确性
RAFT的输入mc_model(img1, img2)输出的是从img1到img2的光流,即该光流可将img1 warp到img2视角。如果你的需求是把当前帧对齐到参考帧:
- 若输入RAFT的顺序是
mc_model(curr_transf, ref_transf),则上述函数逻辑正确; - 若对齐后图像仍异常,尝试将光流取反(方向搞反):
flow_norm = -flow.permute(0,2,3,1) / torch.tensor([W/2, H/2], device=img.device)
3. 其他潜在问题修复
- 初始化
aligned_frames时需指定设备,避免CPU/GPU数据不匹配:aligned_frames = torch.zeros_like(noise_img, device=device) - 确保
transforms函数的预处理和RAFT训练时一致(比如图像像素值范围需在[0,1]或[-1,1],不能超出范围)
修复后的完整对齐逻辑示例
def warp_flow(img, flow): B, C, H, W = img.size() x = torch.linspace(-1, 1, W, device=img.device).view(1, 1, W).repeat(B, H, 1) y = torch.linspace(-1, 1, H, device=img.device).view(1, H, 1).repeat(B, 1, W) grid = torch.stack([x, y], dim=-1) flow_norm = flow.permute(0,2,3,1) / torch.tensor([W/2, H/2], device=img.device) grid = grid + flow_norm return torch.nn.functional.grid_sample(img, grid, align_corners=True) # 主逻辑 noise_img = noise_img.to(device) clean_img = clean_img.to(device) return_index = noise_img.size(1) // 2 aligned_frames = torch.zeros_like(noise_img, device=device) aligned_frames[:, return_index] = noise_img[:, return_index] for idx in range(noise_img.size(1)): if idx == return_index: continue curr_frame = noise_img[:, idx] ref_frame = noise_img[:, return_index] curr_transf, ref_transf = transforms(curr_frame, ref_frame) curr_flow = mc_model(curr_transf, ref_transf)[-1] aligned_frames[:, idx] = warp_flow(curr_transf, curr_flow)
内容的提问来源于stack exchange,提问作者Bled Clement
相关产品推荐
相关产品推荐

