循环超过张量

Moh*_*hal 13 python tensorflow

我试图以python的方式处理一个可变大小的张量,如下所示:

# X is of shape [m, n]
for x in X:
    process(x)
Run Code Online (Sandbox Code Playgroud)

我试图使用tf.scan,问题是我想处理每个子张量,所以我试图使用嵌套扫描,但是我启用了它,因为tf.scan可以使用累加器,如果没有发现它将把elems的第一个条目作为初始化器,我不想这样做.举个例子,假设我想在张量的每个元素中添加一个(这只是一个例子),我想逐个元素地处理它.如果我运行下面的代码,我将只添加一个子张量,因为扫描将第一个张量视为初始化器,以及每个子张量的第一个元素.

import numpy as np
import tensorflow as tf

batch_x = np.random.randint(0, 10, size=(5, 10))
x = tf.placeholder(tf.float32, shape=[None, 10])

def inner_loop(x_in):
    return tf.scan(lambda _, x_: x_ + 1, x_in)

outer_loop = tf.scan(lambda _, input_: inner_loop(input_), x, back_prop=True)

with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())
    rs = sess.run(outer_loop, feed_dict={x: batch_x})
Run Code Online (Sandbox Code Playgroud)

有什么建议 ?

Dmi*_*kiy 9

大多数tensorflow内置函数都可以逐元素应用。因此,您可以将张量传递给函数。喜欢:

outer_loop = inner_loop(x)
Run Code Online (Sandbox Code Playgroud)

但是,如果您有一些无法通过这种方式应用的功能(确实很容易看到该功能),则可以使用map_fn

假设您的函数只是将张量的每个元素加1(或其他任何值):

inputs = tf.placeholder...

def my_elementwise_func(x):
    return x + 1

def recursive_map(inputs):
   if tf.shape(inputs).ndims > 0:
       return tf.map_fn(recursive_map, inputs)
   else:
       return my_elementwise_func(inputs)

result = recursive_map(inputs)  
Run Code Online (Sandbox Code Playgroud)


Dzj*_*jkb 8

要遍历张量,您可以尝试tf.unstack

将秩-R张量的给定维度解包为秩(R-1)张量.

因此,为每个张量添加1看起来像:

import tensorflow as tf
x = tf.placeholder(tf.float32, shape=(None, 10))
x_unpacked = tf.unstack(x) # defaults to axis 0, returns a list of tensors

processed = [] # this will be the list of processed tensors
for t in x_unpacked:
    # do whatever
    result_tensor = t + 1
    processed.append(result_tensor)

output = tf.concat(processed, 0)

with tf.Session() as sess:
    print(sess.run([output], feed_dict={x: np.zeros((5, 10))}))
Run Code Online (Sandbox Code Playgroud)

显然,您可以从列表中进一步解压缩每个张量以处理它,直到单个元素.为了避免大量的嵌套解包,你可以尝试tf.reshape(x, [-1])先用x展平x ,然后像它一样循环

flattened_unpacked = tf.unstack(tf.reshape(x, [-1])
for elem in flattened_unpacked:
    process(elem)
Run Code Online (Sandbox Code Playgroud)

在这种情况下elem是标量.

  • 哦,对不起,我错过了 `unstack` 不适用于 None 维度。我发现了这个问题,[当变量的第一维为无时使用 tf.unpack()](http://stackoverflow.com/questions/39446313/using-tf-unpack-when-first-dimension-of-variable- is-none)这似乎得到了很好的回答。 (3认同)