将PyTorch张量转换为NumPy数组时全部值变为零的问题排查
问题原因与解决方法
问题根源
你的代码中核心问题出在argmax(dim=1)这一步:
- 原张量
var1的形状是[1, 1, 480, 640],其中dim=1维度的长度仅为1。 argmax(dim=1)的作用是在指定维度上取最大值对应的索引,当该维度只有一个元素时,所有位置的最大值索引必然是0,这直接导致后续得到的数组全为0,和squeeze、转NumPy的操作无关。
正确转换方式
如果只是想将原PyTorch张量转换为NumPy数组,无需调用argmax,直接执行以下代码:
# 两种等价写法,任选其一即可 nump_var1 = var1.squeeze(0).squeeze(0).cpu().numpy() # 或者 nump_var1 = var1[0, 0].cpu().numpy()
执行后就能得到与原张量数值完全一致的NumPy数组,形状为(480, 640)。
补充说明
如果你的argmax调用是出于特定业务需求(比如原本预期dim=1是多分类的通道维度),那说明输入张量var1的维度不符合预期。此时需要检查上游代码,确认是否在生成var1时错误地将通道数设置为1,导致argmax无法输出有效索引。
内容的提问来源于stack exchange,提问作者PScode
相关产品推荐
相关产品推荐

