如何在TensorFlow中实现等效PyTorch permute的张量维度置换
TensorFlow实现张量维度置换的正确方法
原有PyTorch实现参考
PyTorch框架下的张量维度置换实现代码如下:
import torch A = torch.rand(1, 2,5) A = A.permute(0,2,1) A.shape
代码运行后输出的张量形状为:
torch.Size([1, 5, 2])
原有TensorFlow代码的错误点
你编写的TensorFlow测试代码存在两处错误,导致无法正常运行:
tf.random.normal生成随机张量时,形状参数需要以列表/元组形式传入,不能直接传入多个独立数值tf.keras.layers.Permute是Keras的网络层类,首先需要实例化后传入张量才能得到计算结果;其次该层的置换参数不包含batch维度,不需要将代表batch维的0写入参数列表
正确实现方案
方案1:使用tf.transpose(与PyTorch permute逻辑完全一致,日常运算优先使用)
tf.transpose的参数逻辑和PyTorch的permute完全对齐,需要传入包含所有维度(含batch维)的置换顺序,不需要做额外的维度偏移,代码如下:
import tensorflow as tf A = tf.random.normal(shape=(1, 2, 5)) A = tf.transpose(A, perm=[0, 2, 1]) print(A.shape)
运行输出:
(1, 5, 2)
和PyTorch实现效果完全一致。
方案2:使用tf.keras.layers.Permute(适合Keras序贯/函数式模型搭建场景)
如果是在构建Keras网络结构时需要做维度置换,使用该层时注意仅传入非batch维度的置换顺序即可,代码如下:
import tensorflow as tf A = tf.random.normal(shape=(1, 2, 5)) # 实例化置换层,参数为非batch维度的置换顺序:交换原第1、2维(非batch维度下索引为1、2,对应层参数写(2,1)) permute_layer = tf.keras.layers.Permute((2, 1)) A = permute_layer(A) print(A.shape)
运行输出同样为:
(1, 5, 2)
内容的提问来源于stack exchange,提问作者Anshuman Sinha
相关产品推荐
相关产品推荐

