带有字符串输入的 Tensorflow 数据集不保留数据类型

des*_*aut 4 numpy tensorflow tensorflow-datasets tensorflow2.0

下面所有可重现的代码都在 Google Colab 上使用 TF 2.2.0-rc2 运行。

改编文档中的简单示例以从简单的 Python 列表创建数据集:

import numpy as np
import tensorflow as tf
tf.__version__
# '2.2.0-rc2'
np.version.version
# '1.18.2'

dataset1 = tf.data.Dataset.from_tensor_slices([1, 2, 3]) 
for element in dataset1: 
  print(element) 
  print(type(element.numpy()))
Run Code Online (Sandbox Code Playgroud)

我们得到结果

tf.Tensor(1, shape=(), dtype=int32)
<class 'numpy.int32'>
tf.Tensor(2, shape=(), dtype=int32)
<class 'numpy.int32'>
tf.Tensor(3, shape=(), dtype=int32)
<class 'numpy.int32'>
Run Code Online (Sandbox Code Playgroud)

int32正如预期的那样,所有数据类型都是。

但是改变这个简单的例子来提供一个字符串列表而不是整数:

dataset2 = tf.data.Dataset.from_tensor_slices(['1', '2', '3']) 
for element in dataset2: 
  print(element) 
  print(type(element.numpy()))
Run Code Online (Sandbox Code Playgroud)

给出结果

tf.Tensor(b'1', shape=(), dtype=string)
<class 'bytes'>
tf.Tensor(b'2', shape=(), dtype=string)
<class 'bytes'>
tf.Tensor(b'3', shape=(), dtype=string)
<class 'bytes'>
Run Code Online (Sandbox Code Playgroud)

令人惊讶的是,尽管张量本身是dtype=string,但它们的评估类型是bytes。

这种行为不仅限于.from_tensor_slices方法;这是情况.list_files(以下代码段在新的 Colab 笔记本中直接运行):

disc_data = tf.data.Dataset.list_files('sample_data/*.csv') # 4 csv files
for element in disc_data: 
  print(element) 
  print(type(element.numpy()))
Run Code Online (Sandbox Code Playgroud)

结果是:

tf.Tensor(b'sample_data/california_housing_test.csv', shape=(), dtype=string)
<class 'bytes'>
tf.Tensor(b'sample_data/mnist_train_small.csv', shape=(), dtype=string)
<class 'bytes'>
tf.Tensor(b'sample_data/california_housing_train.csv', shape=(), dtype=string)
<class 'bytes'>
tf.Tensor(b'sample_data/mnist_test.csv', shape=(), dtype=string)
<class 'bytes'>
Run Code Online (Sandbox Code Playgroud)

再次,评估张量中的文件名返回为bytes,而不是string,尽管张量本身是dtype=string.

使用该.from_generator方法也观察到类似的行为(此处未显示)。

最后的演示:如.as_numpy_iterator方法文档中所示,以下相等条件被评估为True:

dataset3 = tf.data.Dataset.from_tensor_slices({'a': ([1, 2], [3, 4]), 
                                               'b': [5, 6]}) 

list(dataset3.as_numpy_iterator()) == [{'a': (1, 3), 'b': 5}, 
                                       {'a': (2, 4), 'b': 6}] 
# True
Run Code Online (Sandbox Code Playgroud)

但是如果我们将 的元素更改b为字符串,则相等条件现在令人惊讶地评估为False!

dataset4 = tf.data.Dataset.from_tensor_slices({'a': ([1, 2], [3, 4]), 
                                               'b': ['5', '6']})   # change elements of b to strings

list(dataset4.as_numpy_iterator()) == [{'a': (1, 3), 'b': '5'},   # here
                                       {'a': (2, 4), 'b': '6'}]   # also
# False
Run Code Online (Sandbox Code Playgroud)

可能是由于不同的数据类型,因为值本身显然是相同的。


我不是通过学术实验偶然发现这种行为的。我正在尝试使用自定义函数将我的数据传递给 TF 数据集,这些函数从表单的磁盘读取文件对

f = ['filename1', 'filename2']
Run Code Online (Sandbox Code Playgroud)

哪些自定义函数可以很好地独立工作,但通过 TF 数据集映射给出

RuntimeError: not a string
Run Code Online (Sandbox Code Playgroud)

在此挖掘之后,如果返回的数据类型确实是bytes并且不是,那么这似乎至少不是无法解释的string。

那么,这是一个错误(看起来),还是我在这里遗漏了什么?

nes*_*uno 6

这是一个已知的行为:

来自:https : //github.com/tensorflow/tensorflow/issues/5552#issuecomment-260455136

TensorFlow 在大多数地方将 str 转换为字节,包括 sess.run,这不太可能改变。用户可以自由地转换回来,但不幸的是,向核心添加 unicode dtype 的更改太大了。关闭暂时无法解决。

我想 TensorFlow 2.x 没有任何改变 - 仍有一些地方将字符串转换为字节,您必须手动处理。

从您自己打开的问题来看,他们似乎将该主题视为 Numpy 的问题,而不是 Tensorflow 本身的问题。