无法打印Tensor对象内容?加载Pandas DataFrame到TF数据集异常
我之前也碰到过一模一样的坑——把Pandas DataFrame转换成TensorFlow Dataset后,遍历打印时输出的全是Tensor的节点元信息,根本看不到实际的数值内容,哪怕完全照搬官方示例代码也没用。
问题重现
先贴一下当时的代码:
import pandas as pd import tensorflow as tf URL = 'https://storage.googleapis.com/applied-dl/heart.csv' df = pd.read_csv(URL) df['thal'] = pd.Categorical(df['thal']) df['thal'] = df.thal.cat.codes target = df.pop('target') dataset = tf.data.Dataset.from_tensor_slices((df.values, target.values)) for feat, targ in dataset.take(5): print('Features: {}, Target: {}'.format(feat, targ))
预期应该输出具体的特征数组和目标值:
Features: [ 63. 1. 1. 145. 233. 1. 2. 150. 0. 2.3 3. 0. 2. ], Target: 0
Features: [ 67. 1. 4. 160. 286. 0. 2. 108. 1. 1.5 2. 3. 3. ], Target: 1
...
但实际输出却是一堆Tensor节点的描述:
Features: Tensor("IteratorGetNext:0", shape=(13,), dtype=float64), Target: Tensor("IteratorGetNext:1", shape=(), dtype=int64)
Features: Tensor("IteratorGetNext_1:0", shape=(13,), dtype=float64), Target: Tensor("IteratorGetNext_1:1", shape=(), dtype=int64)
...
解决方案
这个问题的核心是TensorFlow的执行模式:在TensorFlow 1.x版本中,默认是图执行模式——此时创建的Tensor只是计算图里的一个节点,不会立刻计算出实际数值,必须通过Session.run()才能获取结果。而要直接打印Tensor的内容,我们需要启用即刻执行(Eager Execution)。
解决步骤超简单,只需要在导入TensorFlow之后,立刻添加一行代码:
tf.enable_eager_execution()
如果是使用TensorFlow 2.x的话,其实默认已经开启了即刻执行模式,不过如果是从1.x迁移过来的项目,还是手动加一下更稳妥。
添加这行代码后再运行原有的遍历打印逻辑,就能得到预期的数值输出了!
原因解释
即刻执行模式会让TensorFlow在创建Tensor的同时就立即计算并存储其数值,不需要等到构建完整计算图后再通过会话触发计算。这样在遍历Dataset时,我们拿到的就是已经计算好的具体数值,而不是待执行的计算节点。
内容的提问来源于stack exchange,提问作者IAmSteve

