【问题标题】:tf.tensor_scatter_nd_add function did worktf.tensor_scatter_nd_add 函数确实有效
【发布时间】:2019-12-09 09:58:25
【问题描述】:

我想在张量的第一个维度中插入两个切片,其中包含两个新值矩阵,我使用的是 tensor_scatter_add 方法,但它给了我一个错误

indices = tf.constant([[0], [2]])
updates = tf.constant([[[5, 5, 5, 5], [6, 6, 6, 6],
                        [7, 7, 7, 7], [8, 8, 8, 8]],
                       [[5, 5, 5, 5], [6, 6, 6, 6],
                        [7, 7, 7, 7], [8, 8, 8, 8]]])
tensor = tf.ones([4, 5, 4])
updated = tf.tensor_scatter_add(tensor, indices, updates)
with tf.Session() as se:
  print(ses.run(scatter))

【问题讨论】:

    标签: python tensorflow tensor


    【解决方案1】:

    tensor 的内部 2 维必须与 updates 的内部 2 维匹配。 Dimension 0 in both shapes must be equal, but are 5 and 4.

    tensor 必须与 dtype 相同,但您的代码不同。

    有错误:

    with tf.Session() as se:
      print(ses.run(scatter))
    

    您将tf.Session() 别名为se,但调用ses 而不是se,并且您的传递分散到ses.run() 但它没有在任何地方定义; se.run(updated) 应该是正确的函数调用。

    带有代码修复的片段:
    这对你应该没问题。

    indices = tf.constant([[0], [2]])
    updates = tf.constant([[[5, 5, 5, 5], [6, 6, 6, 6],
                            [7, 7, 7, 7], [8, 8, 8, 8]],
                           [[5, 5, 5, 5], [6, 6, 6, 6],
                            [7, 7, 7, 7], [8, 8, 8, 8]]])
    tensor = tf.ones([4, 4, 4], dtype=tf.int32)
    updated = tf.tensor_scatter_nd_add(tensor, indices, updates)
    with tf.Session() as se:
      print(se.run(updated))
    

    【讨论】:

      【解决方案2】:

      只需更正这些行,您输入错误,这导致了代码中的问题:

      tensor = tf.ones([4, 4, 4])
      updated = tf.tensor_scatter_add(tensor, indices, updates)
      with tf.Session() as se:
        print(se.run(scatter))
      

      【讨论】:

      • se.run(scatter) 是错误的,因为 scatter 没有在任何地方定义。 updated 应该传递给 se 或者重命名更新为 scatter。
      猜你喜欢
      • 2021-04-20
      • 1970-01-01
      • 2016-04-01
      • 1970-01-01
      • 1970-01-01
      • 2010-09-28
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多