【问题标题】:How to keep vectors in a dictionary in tensorflow?如何将向量保存在张量流的字典中?
【发布时间】:2020-09-16 03:27:45
【问题描述】:

tf.lookup.experimental.DenseHashTable 似乎无法保存向量,我找不到如何使用它的示例。

【问题讨论】:

    标签: tensorflow tensorflow2.x


    【解决方案1】:

    您可以在下面找到 Tensorflow 中向量字典的简单实现。这也是tf.lookup.experimental.DenseHashTabletf.TensorArray的用法示例。

    如上所述,向量不能保存在tf.lookup.experimental.DenseHashTable 中,因此tf.TensorArray 用于保存实际向量。

    当然,这是一个简单的例子,它不包括删除字典中的条目——这个操作需要对数组的空闲单元进行一些管理。此外,您应该阅读tf.lookup.experimental.DenseHashTabletf.TensorArray 各自的API 页面,了解如何根据您的需要调整它们。

    import tensorflow as tf
    
    
    class DictionaryOfVectors:
    
      def __init__(self, dtype):
        empty_key = tf.constant('')
        deleted_key = tf.constant('deleted')
    
        self.ht = tf.lookup.experimental.DenseHashTable(key_dtype=tf.string,
                                                        value_dtype=tf.int32,
                                                        default_value=-1,
                                                        empty_key=empty_key,
                                                        deleted_key=deleted_key)
        self.ta = tf.TensorArray(dtype, size=0, dynamic_size=True, clear_after_read=False)
        self.inserts_counter = 0
    
      @tf.function
      def insertOrAssign(self, key, vec):
        # Insert the vector to the TensorArray. The write() method returns a new
        # TensorArray object with flow that ensures the write occurs. It should be 
        # used for subsequent operations.
        with tf.init_scope():
          self.ta = self.ta.write(self.inserts_counter, vec)
    
          # Insert the same counter value to the hash table
          self.ht.insert_or_assign(key, self.inserts_counter)
          self.inserts_counter += 1
    
      @tf.function
      def lookup(self, key):
        with tf.init_scope():
          index = self.ht.lookup(key)
          return self.ta.read(index)
    
    dictionary_of_vectors = DictionaryOfVectors(dtype=tf.float32)
    dictionary_of_vectors.insertOrAssign('first', [1,2,3,4,5])
    print(dictionary_of_vectors.lookup('first'))
    

    这个例子有点复杂,因为插入和查找方法都用@tf.function 修饰。因为这些方法更改了在它们之外定义的变量,所以使用了tf.init_scope()。您可能会问lookup() 方法中发生了什么变化,因为它实际上只从哈希表和数组中读取。原因是在图形模式下,lookup() 调用返回的索引是一个张量,而在 TensorArray 实现中,有一行包含if index < 0: 失败:

    OperatorNotAllowedInGraphError:不允许将tf.Tensor 用作Python bool

    当我们使用 tf.init_scope() 时,正如其 API 文档中所解释的那样,“即使在跟踪 tf.function 时,init_scope 块内的代码也会在启用急切执行的情况下运行”。所以在那种情况下,索引不是张量而是标量。

    【讨论】:

      猜你喜欢
      • 2018-02-13
      • 2020-04-24
      • 1970-01-01
      • 2021-03-14
      • 1970-01-01
      • 2018-07-10
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多