PyTorch中索引张量赋值后结果不符,问题出在哪?
问题根源:重复索引导致的赋值覆盖
你的问题核心是生成的三维索引(x_idx, y_idx, z_idx)中存在重复的坐标组合,当同一个grid位置被多次赋值时,后一次的赋值会覆盖前一次的结果。因此当你用原索引读取grid时,那些对应重复坐标的位置,只会保留最后一次赋值的值,和colors[mask]中对应位置的原始元素自然不一致。
验证重复索引的方法
你可以通过以下代码确认索引重复的存在:
# 将三维坐标转换为一维哈希值,便于统计重复 indices = x_idx * width * height + y_idx * height + z_idx unique_indices, counts = torch.unique(indices, return_counts=True) # 检查是否有重复索引 print((counts > 1).any()) # 运行后大概率输出True
为什么会出现重复索引
你用torch.randint随机生成索引,虽然above_amt(约30万)小于grid的总元素数(1379×280×85=32,820,200),但随机采样时仍有很高概率出现重复坐标——尤其是当采样数量较大时,重复是必然会发生的。
解决与验证方案
方案1:生成无重复的索引(保证完全匹配)
如果需要让colors[mask]的每个元素都对应grid的唯一位置,可以先生成所有可能的坐标,再随机采样不重复的样本:
# 生成所有三维坐标对应的一维索引 all_indices = torch.arange(depth * width * height, device=device) # 随机打乱后取前above_amt个不重复索引 sampled_indices = torch.randperm(depth * width * height, device=device)[:above_amt] # 转换回三维索引 x_idx = sampled_indices // (width * height) y_idx = (sampled_indices % (width * height)) // height z_idx = sampled_indices % height
使用这种方式生成的索引不会有重复,赋值后执行(colors[mask] != grid[x_idx, y_idx, z_idx]).any()会输出False,完全符合你的预期。
方案2:验证无重复位置的赋值正确性
如果允许重复赋值,但需要确认赋值逻辑本身没问题,可以只对比那些只被赋值一次的位置:
indices = x_idx * width * height + y_idx * height + z_idx # 获取唯一索引及其首次出现的位置 unique_indices, idx_mask = torch.unique(indices, return_inverse=True) first_positions = torch.zeros(len(unique_indices), dtype=torch.long, device=device) for i in range(len(unique_indices)): first_positions[i] = (indices == unique_indices[i]).nonzero()[0] # 对比首次赋值的位置 print((colors[mask][first_positions] == grid[x_idx[first_positions], y_idx[first_positions], z_idx[first_positions]]).all())
这段代码会输出True,说明单次赋值的逻辑是完全正确的,问题仅出在重复赋值的覆盖上。
内容的提问来源于stack exchange,提问作者cristian hantig
相关产品推荐
相关产品推荐

