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

PyTorch DataLoader加载PCN点云 完整GT点云仅首次迭代正常后续异常

PCN点云补全训练GT点云加载异常问题

问题现象

使用PCN网络开展点云补全训练时,需同时加载两类数据:

  • 物体部分点云(模型输入Input)
  • 对应完整点云(训练真值Ground Truth,简称GT)
    当前加载存在固定异常:部分点云全程加载正常,但GT完整点云仅在训练首次迭代时加载正确,首次迭代后加载的GT数据存在异常,疑似数据损坏:首次迭代可视化结果正常,第二次及后续迭代的GT点云显示异常,不同运行轮次中出现异常的点云对象随机,经校验磁盘存储的GT点云文件本身无损坏。
    首次迭代与第二次迭代的可视化结果如下(其他运行轮次中异常迭代的点云可对应任意加载对象,正常GT点云应与磁盘存储文件完全一致):
    首次迭代正常结果
    第二次迭代异常结果

已做的适配与前置排查

参考项目内datasets/shapenet.py的加载逻辑做了自定义数据集适配:单物体共400份点云,划分为280份训练集、120份验证集,路径拼接时使用%280等取模逻辑实现索引循环,点云文件按0.pcd、1.pcd的序号规则依次命名。
最初单物体仅存储1份完整点云文件,为排除重复打开同一文件导致异常的可能,已为每个部分点云复制了对应同名的完整点云副本,通过逐文件可视化校验确认磁盘上存储的所有点云文件本身均无问题。

相关代码

数据路径加载逻辑

def _load_data(self):

    with open(os.path.join(self.dataroot, '{}.list').format(self.split), 'r') as f:
        lines = f.read().splitlines()
    if self.category != 'all':
        lines = list(filter(lambda x: x.startswith(self.cat2id2[self.category]), lines))

    partial_paths, complete_paths = list(), list()
    train_counter = 0
    test_counter = 280 # 待清理调整
    total_counter = 0

    for line in lines:
        print(line)
        category, model_id = line.split('/')
        if self.split == 'train':
            for i in range(280):
                partial_paths.append(os.path.join(self.dataroot, self.split, 'partial', category,str(train_counter%280)+'.pcd'))
                complete_paths.append(os.path.join(self.dataroot, self.split, 'complete', category, str(train_counter%280)+'.pcd'))
                train_counter+=1

        else:
            for i in range(120):
                partial_paths.append(os.path.join(self.dataroot, self.split, 'partial', category, str(test_counter)+'.pcd'))
                complete_paths.append(os.path.join(self.dataroot, self.split, 'complete', category, str(test_counter)+'.pcd'))

                test_counter+=1
                if(test_counter == 400):
                    test_counter = 280

                total_counter+=1

    # 逐文件校验逻辑(已注释)
    # for i in range (len(complete):
    #     print(i)
    #     count = count+50
    #     d = complete_paths[count]
    #     pcd = o3d.io.read_point_cloud(d)
    #     o3d.visualization.draw_geometries([pcd])

    return partial_paths, complete_paths

数据集与DataLoader初始化(train.py内)

train_dataset = ShapeNet('data/PCN', 'train', params.category)
train_dataloader = DataLoader(train_dataset, batch_size=params.batch_size, shuffle=True, num_workers=params.num_workers)

训练循环核心代码

经训练循环内校验排查,从DataLoader取出的GT完整点云数值存储异常,疑似PyTorch DataLoader处理逻辑存在问题。

for epoch in range(1, params.epochs + 1):
    # 超参数alpha调整逻辑
    if train_step < 10000:
        alpha = 0.01
    elif train_step < 20000:
        alpha = 0.1
    elif epoch < 50000:
        alpha = 0.5
    else:
        alpha = 1.0

    # 训练阶段
    model.train()
    for i, (p, c) in enumerate(train_dataloader):
        p, c = p.to(params.device), c.to(params.device)

        optimizer.zero_grad()

        # 前向传播
        coarse_pred, dense_pred = model(p)

其中变量p对应部分点云、c对应完整点云,将张量转numpy后调用open3d可视化验证发现:c仅在首个epoch迭代时数值正确,后续迭代取出的c类似空白点云,无法定位异常原因。校验所用可视化代码如下:

d = p[i].cpu().numpy()
        e = c[i].cpu().numpy()
        
        print('当前DataLoader迭代序号:', i)
        pcd = o3d.geometry.PointCloud()
        pcd.points = o3d.utility.Vector3dVector(d)
        o3d.visualization.draw_geometries([pcd])

        pcb = o3d.geometry.PointCloud()
        pcd.points = o3d.utility.Vector3dVector(e)
        o3d.visualization.draw_geometries([pcb])

问题原因与修复方案

1. 可视化代码存在两处错误,直接导致异常现象

该问题完全匹配描述的“首次迭代正常、后续迭代随机异常”特征:

  • 索引使用错误:循环中i是DataLoader的批次序号,不是批次内的样本索引。p、c的第一维长度为批次大小batch_size,当批次序号i大于等于batch_size时,p[i]、c[i]会越界访问张量内存,读取到随机垃圾值,表现为点云空白、形状异常。首次迭代i=0,只要batch_size>=1就不会越界,因此首次结果正常。
  • GT点云对象赋值错误:新建GT点云对象pcb后,错误将点数据赋值给了部分点云对象pcd的points属性,pcb本身未写入任何点数据,可视化结果必然异常。
    修正后的校验代码:
# 固定取批次内第一个样本可视化即可,禁止用批次序号i做样本索引
d = p[0].cpu().numpy()
e = c[0].cpu().numpy()

print('当前批次点云可视化')
pcd = o3d.geometry.PointCloud()
pcd.points = o3d.utility.Vector3dVector(d)
o3d.visualization.draw_geometries([pcd])

pcb = o3d.geometry.PointCloud()
# 注意赋值给pcb的points属性,不要写错对象
pcb.points = o3d.utility.Vector3dVector(e)
o3d.visualization.draw_geometries([pcb])

2. 多Worker加载内存覆盖问题排查

如果修正可视化代码后仍存在GT数据异常,可按以下步骤排查:

  • 先将DataLoader的num_workers参数设为0测试,若设为0后数据恢复正常,即可确认是多worker内存共享导致的覆盖问题。
  • 检查自定义Dataset的__getitem__方法:读取点云转为PyTorch张量时,调用.clone()生成独立张量副本再返回,禁止直接返回类成员变量的引用;每次读取pcd文件时新建独立的open3d点云对象,不要复用全局/类层面的点云对象。

3. 训练时获取样本索引的方法

如需在训练循环中获取当前批次对应的样本索引,只需修改Dataset的__getitem__返回值,增加索引项即可:

# Dataset类__getitem__方法末尾,返回值增加idx
return partial_points, complete_points, idx

对应训练循环修改为:

for i, (p, c, idx) in enumerate(train_dataloader):
    # idx为当前批次所有样本的索引列表,可直接匹配对应GT文件路径

内容的提问来源于stack exchange,提问作者M.K

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 00:36:20