TF2中tf.print在tf.data.Dataset管道内无法输出张量的解决问询
解决tf.data.Dataset中tf.print无法打印张量的问题
你的代码里有两个核心问题导致看不到tf.print的输出,我帮你梳理并修改:
问题1:Dataset的map操作没有赋值回变量
TensorFlow的tf.data.Dataset是不可变对象,所有变换操作(比如map)都会返回一个新的Dataset实例,不会修改原对象。你原来只调用了dataset.map(myFunc)但没把结果存下来,遍历的还是原始的、未经过myFunc处理的Dataset,自然不会触发tf.print。
问题2:stdout重定向到文件(若你想打印到控制台)
你把sys.stdout重定向到了tf.log文件,如果目标是打印到控制台,这一步需要去掉;如果确实要输出到文件,建议用更安全的with语句管理文件句柄。
修改后的完整代码
import tensorflow as tf import sys def main(): ## Dataset generator numRows= 100 indx = tf.reshape([i+1 for i in range(numRows)], [numRows,1]) features = tf.random.uniform([numRows, 2], minval=1, maxval=10, dtype=tf.int32) myData = tf.concat([indx, features], 1) ## tf.data.Dataset dataset = tf.data.Dataset.from_tensor_slices(myData) ## Pipeline:将map的结果重新赋值给dataset dataset = dataset.map(myFunc) # 这行是关键修改! ## Run pipeline for d in dataset: print('--') def myFunc(t): # tf.print默认输出到控制台,无需额外指定output_stream tf.print(t) return t if __name__ == "__main__": main()
额外说明
- 输出到文件的正确方式(如果需要):
如果你想把tf.print的内容写入文件,建议用with语句自动管理文件,避免资源泄漏:import tensorflow as tf import sys def main(): numRows= 100 indx = tf.reshape([i+1 for i in range(numRows)], [numRows,1]) features = tf.random.uniform([numRows, 2], minval=1, maxval=10, dtype=tf.int32) myData = tf.concat([indx, features], 1) dataset = tf.data.Dataset.from_tensor_slices(myData) dataset = dataset.map(myFunc) with open('tf.log', 'w') as f: sys.stdout = f for d in dataset: print('--') def myFunc(t): tf.print(t, output_stream=sys.stdout) return t if __name__ == "__main__": main() - 图模式下的打印差异:
普通的Pythonprint在TensorFlow图模式下可能不会按预期执行(因为图构建阶段不会运行),而tf.print是图内操作,会在Dataset迭代时被触发执行,这部分你选的工具是对的。
内容的提问来源于stack exchange,提问作者Kleyson Rios
相关产品推荐
相关产品推荐

