You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 10:20:18