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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 07:45:22