TensorFlow与PyTorch张量切片结果不一致,是否由reshape导致?
TensorFlow与PyTorch切片操作结果不一致的原因
先看测试代码:
import tensorflow as tf import numpy as np import torch tf_x = tf.random.uniform((4, 64, 64, 3)) # 第一组:用permute转置维度 pt_x = torch.Tensor(tf_x.numpy()) pt_x = pt_x.permute(0, 3, 1, 2) # 切片后转置回原维度对比 print(np.any(pt_x[:, :, 1:].permute(0, 2, 3, 1).numpy() - tf_x[:, 1:].numpy())) # 输出:False,结果一致 # 第二组:用reshape调整形状 pt_x = torch.Tensor(tf_x.numpy()) b, h, w, c = pt_x.shape pt_x = pt_x.reshape((b, c, h, w)) # 先验证reshape回原形状结果一致 print(np.any(pt_x.view(b, h, w, c).numpy() - tf_x.numpy())) # 输出:False # 切片后reshape回原维度对比 print(np.any(pt_x[:, :, 1:].reshape(4, 63, 64, 3).numpy() - tf_x[:, 1:].numpy())) # 输出:True,结果不一致
核心原因:reshape和permute的本质区别
问题确实出在reshape操作上,它和permute的作用完全不同:
- permute:是真正的维度转置,会重新排列张量在内存中的存储顺序。把
(b,h,w,c)转成(b,c,h,w)时,数据会被重新组织,让每个通道的元素连续存储,此时切片[:, :, 1:]对应的就是原张量height维度从第1行开始的部分,转回去后自然和TensorFlow的切片结果一致。 - reshape:只是改变张量的形状定义,不会修改数据在内存中的存储顺序。原张量
(4,64,64,3)在内存中是按batch→height→width→channel的顺序存储的,reshape成(4,3,64,64)后,内存顺序没变,只是告诉PyTorch按batch→channel→height→width的规则读取数据。这时pt_x[:, :, 1:]取的并不是原height维度的第1行及以后的数据,而是内存中跳过前64*64个元素后的内容,和TensorFlow的切片区域完全不匹配,最终结果自然不一致。
如果把第二组中的reshape换成permute,最后一行的输出就会变成False,和TensorFlow结果一致。
内容的提问来源于stack exchange,提问作者PixelPioneer
相关产品推荐
相关产品推荐

