如何访问 tf.data.Dataset.list_files() 收集的文件名?

ROS*_*ROS 5 python tensorflow tensorflow-datasets

我在用

file_data = tf.data.Dataset.list_files("../*.png")
Run Code Online (Sandbox Code Playgroud)

收集图像文件以在 TensorFlow 中进行训练,但想访问收集的文件名列表,以便我可以执行标签查找。

调用 sess.run([file_data]) 已经失败:

TypeError: Fetch argument <TensorSliceDataset shapes: (), types: tf.string> has invalid type <class 'tensorflow.python.data.ops.dataset_ops.TensorSliceDataset'>, must be a string or Tensor. (Can not convert a TensorSliceDataset into a Tensor or Operation.)
Run Code Online (Sandbox Code Playgroud)

我可以使用其他任何方法吗?

ROS*_*ROS 5

通过一些额外的实验,我找到了解决这个问题的方法:

首先,将 Dataset 变成迭代器:

iterator_helper = file_data.make_one_shot_iterator()
Run Code Online (Sandbox Code Playgroud)

然后,遍历 tf Session 中的元素:

with tf.Session() as sess:
    filename_temp = iterator_helper.get_next()
    print(sess.run[filename_temp])
Run Code Online (Sandbox Code Playgroud)


mrr*_*rry 5

APIDataset.list_files()使用tf.matching_files()op 列出与给定模式匹配的文件。您还可以使用该操作获取文件列表tf.Tensor,并将其直接传递给sess.run():

filenames_as_tensor = tf.matching_files("../*.png")
filenames_as_array = sess.run(filenames_as_tensor)

for filename in filenames_as_array:
  print(filename)
Run Code Online (Sandbox Code Playgroud)