tensorflow里面许多针对数组操作的函数,官方文档又看了没啥卵用,网上帖子直接copy官方文档而不解释,只能自己写个程序测试理解,以3个维度的tensor进行理解
tf.transpose()作为数组的转置函数,原型如下:
def transpose(a, perm=None, name="transpose"): """Transposes `a`. Permutes the dimensions according to `perm`.
a:是传入的数组
perm:控制转置的操作,以perm = [0,1,2] 3个维度的数组为例, 0--代表的是最外层的一维, 1--代表外向内数第二维, 2--代表最内层的一维,这种perm是默认的值.现在以如下输入数组来理解这个函数和参数perm
import tensorflow as tf
import numpy as np
input = [[[1, 2, 3, 4],[5, 6, 7, 8],[9, 10, 11, 12]],[[13, 14, 15, 16],[17, 18, 19, 20],[21, 22, 23, 24]]]
with tf.Session() as sess:
x=tf.transpose(input,[1,0,2])
print(sess.run(x))
input_x 是一个 2x3x4的一个tensor, 假设perm = [1,0,2], 就是将最外2层转置,得到tensor应该是 3x2x4的一个张量,将input_x抽象化,不管第3维度
输出为:
[[[ 1 2 3 4] [13 14 15 16]] [[ 5 6 7 8] [17 18 19 20]] [[ 9 10 11 12] [21 22 23 24]]]