TensorFlow中图执行与即时执行的测试适配问题
没问题!我帮你把基于Graph与Session的TensorFlow测试用例转换成Eager Mode版本,同时附上原代码对比和关键差异说明,方便你快速验证功能:
原Graph & Session模式测试用例
import numpy as np import tensorflow as tf def test_normal_execution(): matrix_2x4 = np.array([[1, 2, 3, 4], [6, 7, 8, 9]]) dataset = tf.data.Dataset.from_tensor_slices(matrix_2x4) iterator = dataset.make_one_shot_iterator() first_elem = iterator.get_next() with tf.Session() as sess: result = sess.run(first_elem) assert (result == [1, 2, 3, 4]).all() sess.close()
对应的Eager Mode测试用例
import numpy as np import tensorflow as tf # TensorFlow 2.x 默认启用Eager Mode,若使用TF1.x需手动开启:tf.enable_eager_execution() def test_normal_execution_eager(): matrix_2x4 = np.array([[1, 2, 3, 4], [6, 7, 8, 9]]) dataset = tf.data.Dataset.from_tensor_slices(matrix_2x4) # Eager模式下直接通过迭代器获取元素,无需创建Session first_elem = next(iter(dataset)) # 调用.numpy()将Tensor转为numpy数组后断言(也可直接对比Tensor与数组,TF会自动转换) assert (first_elem.numpy() == [1, 2, 3, 4]).all()
关键差异说明
- 无需Session管理:Eager Mode默认即时执行,不用手动创建
tf.Session(),也不需要sess.run()来触发计算 - 简化数据集迭代:直接通过
iter(dataset)获取迭代器,用next()取元素,或者直接用for elem in dataset:遍历,省去了make_one_shot_iterator()的步骤 - 直观的结果操作:Eager模式下Tensor对象可以直接调用
.numpy()转为numpy数组,断言逻辑和原代码保持一致,更符合Python原生编程习惯
内容的提问来源于stack exchange,提问作者3voC
相关产品推荐
相关产品推荐

