【问题标题】:Evaluate all pair combinations of rows of two tensors in tensorflow评估张量流中两个张量行的所有对组合
【发布时间】:2017-09-17 22:41:39
【问题描述】:

我正在尝试在 tensorflow 中定义一个自定义操作,其中有一次我需要构造一个矩阵 (z),该矩阵将包含两个矩阵的行对的所有组合的总和 (x 和 @ 987654323@)。一般来说xy的行数是动态的。

在 numpy 中相当简单:

import numpy as np
from itertools import product

rows_x = 4
rows_y = 2
dim = 2

x = np.arange(dim*rows_x).reshape(rows_x, dim)
y = np.arange(dim*rows_y).reshape(rows_y, dim)

print('x:\n{},\ny:\n{}\n'.format(x, y))

z = np.zeros((rows_x*rows_y, dim))
print('for loop:')
for i, (x_id, y_id) in enumerate(product(range(rows_x), range(rows_y))):
    print('row {}: {} + {}'.format(i, x[x_id, ], y[y_id, ]))
    z[i, ] = x[x_id, ] + y[y_id, ]

print('\nz:\n{}'.format(z))

返回:

x:
[[0 1]
 [2 3]
 [4 5]
 [6 7]],
y:
[[0 1]
 [2 3]]

for loop:
row 0: [0 1] + [0 1]
row 1: [0 1] + [2 3]
row 2: [2 3] + [0 1]
row 3: [2 3] + [2 3]
row 4: [4 5] + [0 1]
row 5: [4 5] + [2 3]
row 6: [6 7] + [0 1]
row 7: [6 7] + [2 3]

z:
[[  0.   2.]
 [  2.   4.]
 [  2.   4.]
 [  4.   6.]
 [  4.   6.]
 [  6.   8.]
 [  6.   8.]
 [  8.  10.]]

但是,我不知道如何在 tensorflow 中实现类似的东西。

我主要通过 SO 和 tensorflow API 希望找到一个可以产生两个张量元素组合的函数,或者一个可以给出张量元素排列的函数,但无济于事。

欢迎提出任何建议。

【问题讨论】:

    标签: python numpy tensorflow


    【解决方案1】:

    你可以简单地使用 tensorflow 的广播能力。

    import tensorflow as tf
    
    x = tf.constant([[0, 1],[2, 3],[4, 5],[6, 7]], dtype=tf.float32)
    y = tf.constant([[0, 1],[2, 3]], dtype=tf.float32)
    
    x_ = tf.expand_dims(x, 0)
    y_ = tf.expand_dims(y, 1)
    z = tf.reshape(tf.add(x_, y_), [-1, 2])
    # or more succinctly 
    z = tf.reshape(x[None] + y[:, None], [-1, 2])
    
    sess = tf.Session()
    sess.run(z)
    

    【讨论】:

    • 这很神奇...所以,要做到这一点:您首先扩展 xy,以便 x_ 具有形状 [1, 4, 2] 和 @ 987654326@ 的形状为 [3, 1, 2]。然后,tf.add 的广播能力“弄清楚”如何将维度填入 [3, 4, 2](即tf.add(x_, y_) 的形状),最后,tf.reshape 确保我们有 2 z 中的列。 “弄清楚”是关键部分,正如我正在阅读here:...
    • ... "当遇到两个兼容数组时,结果形状在每个维度索引处的两个输入中具有最大值。",然后:"出现特殊情况,也支持,其中每个输入数组在不同的索引处都有退化维度。在这种情况下,结果是“外部操作”。这很微妙。谢谢你的回答!
    【解决方案2】:

    选项 1

    z 定义为变量并更新其行:

    import tensorflow as tf
    from itertools import product
    
    
    x = tf.constant([[0, 1],[2, 3],[4, 5],[6, 7]],dtype=tf.float32)
    y = tf.constant([[0, 1],[2, 3]],dtype=tf.float32)
    
    rows_x,dim=x.get_shape()
    rows_y=y.get_shape()[0]
    
    z=tf.Variable(initial_value=tf.zeros([rows_x*rows_y,dim]),dtype=tf.float32)
    for i, (x_id, y_id) in enumerate(product(range(rows_x), range(rows_y))):
        z=tf.scatter_update(z,i,x[x_id]+y[y_id])
    
    with tf.Session() as sess:
        tf.global_variables_initializer().run()
        z_val=sess.run(z)
        print(z_val)
    

    打印出来

    [[  0.   2.]
     [  2.   4.]
     [  2.   4.]
     [  4.   6.]
     [  4.   6.]
     [  6.   8.]
     [  6.   8.]
     [  8.  10.]]
    

    选项 2

    创建z throw 列表理解:

    import tensorflow as tf
    from itertools import product
    
    
    x = tf.constant([[0, 1],[2, 3],[4, 5],[6, 7]],dtype=tf.float32)
    y = tf.constant([[0, 1],[2, 3]],dtype=tf.float32)
    
    rows_x,dim=x.get_shape().as_list()
    rows_y=y.get_shape().as_list()[0]
    
    
    z=[x[x_id]+y[y_id] for x_id in range(rows_x) for y_id in range(rows_y)]
    z=tf.reshape(z,(rows_x*rows_y,dim))
    
    with tf.Session() as sess:
        z_val=sess.run(z)
        print(z_val)
    

    比较:第二个解决方案大约快两倍(仅测量两个解决方案中z 的构造)。具体来说,时间安排是: 第一个解:0.211 秒,第二个解:0.137 秒。

    【讨论】:

    • 构建时间通常是一个无效的性能指标,因为它只发生一次;没有人关心
    猜你喜欢
    • 2019-08-02
    • 1970-01-01
    • 2022-07-28
    • 1970-01-01
    • 1970-01-01
    • 2018-05-09
    • 1970-01-01
    • 2019-01-26
    • 1970-01-01
    相关资源
    最近更新 更多