如何为TensorFlow Dataset中含tf.py_func的map操作返回结果命名?
解决tf.py_func返回元组转命名字典的问题
我刚好处理过类似的场景,其实核心思路很简单:既然tf.py_func不能直接返回字典,那我们可以先拿到它返回的张量元组,再手动把这些张量包装成带键名的字典返回给数据集就可以了。
修改后的完整代码
import tensorflow as tf filelist = ['fileA_6', 'fileB_10', 'fileC_7'] def map_element_counts(fname): # perform operations outside of tensorflow return 'test', 10 ds = tf.data.Dataset.from_tensor_slices({'filenames': filelist}) # 关键修改:把tf.py_func返回的元组包装成字典(高效写法,避免重复调用tf.py_func) ds = ds.map(map_func=lambda x: { 'elementA': elem[0], 'elementB': elem[1] } for elem in [tf.py_func( func=map_element_counts, inp=[x['filenames']], Tout=[tf.string, tf.int64] )]) element = ds.make_one_shot_iterator().get_next() with tf.Session() as sess: print(sess.run(element))
代码说明
这里我用了更高效的写法:先把tf.py_func返回的元组存到临时变量elem里,再从elem中按索引取出对应的张量,分别赋值给字典的elementA和elementB键。这样tf.py_func只会被调用一次,避免了重复执行带来的性能损耗。
当然你也可以用更直白的写法(虽然会重复调用tf.py_func,适合简单场景):
ds = ds.map(map_func=lambda x: { 'elementA': tf.py_func( func=map_element_counts, inp=[x['filenames']], Tout=[tf.string, tf.int64] )[0], 'elementB': tf.py_func( func=map_element_counts, inp=[x['filenames']], Tout=[tf.string, tf.int64] )[1] })
运行结果
执行后会输出你期望的格式:
{'elementA': b'test', 'elementB': 10}
内容的提问来源于stack exchange,提问作者David Parks
相关产品推荐
相关产品推荐

