TensorFlow flat_map报错:缺失op、value_index、dtype参数的解决方法
TypeError 问题及解决方法
问题描述
运行以下代码时出现 TypeError: __init__() missing 3 required positional arguments: 'op', 'value_index' and 'dtype':
from tensorflow.python.ops.numpy_ops import np_config np_config.enable_numpy_behavior() import pandas as pd df = pd.DataFrame( {'x':[1.,2.,3.,4.], 'y':[1.59,4.24,2.38,0.53]} ) data = tf.data.Dataset.from_tensor_slices(df.to_numpy()) data = data.flat_map(lambda x: x.reshape((2,1)))
需求是使用flat_map生成shape=(1,)、dtype=tf.float64的张量,预期输出如下:
for item in data: print(item) tf.Tensor([1.], shape=(1,), dtype=float64) tf.Tensor([2.], shape=(1,), dtype=float64) tf.Tensor([3.], shape=(1,), dtype=float64) tf.Tensor([4.], shape=(1,), dtype=float64) tf.Tensor([1.59], shape=(1,), dtype=float64) tf.Tensor([4.24], shape=(1,), dtype=float64) tf.Tensor([2.38], shape=(1,), dtype=float64) tf.Tensor([0.53], shape=(1,), dtype=float64)
解决方案
错误核心是flat_map要求传入的函数必须返回tf.data.Dataset对象,而非直接返回张量;同时启用numpy行为后,操作返回的是numpy数组而非TensorFlow张量,需显式转换并构造数据集。
修正后的代码如下:
import tensorflow as tf from tensorflow.python.ops.numpy_ops import np_config np_config.enable_numpy_behavior() import pandas as pd df = pd.DataFrame( {'x':[1.,2.,3.,4.], 'y':[1.59,4.24,2.38,0.53]} ) # 显式指定数据类型为float64,构造初始数据集 data = tf.data.Dataset.from_tensor_slices(tf.cast(df.to_numpy(), tf.float64)) # 在flat_map中返回拆分后的数据集 data = data.flat_map(lambda x: tf.data.Dataset.from_tensor_slices(tf.reshape(x, (-1, 1)))) # 验证输出 for item in data: print(item)
关键说明
tf.cast(df.to_numpy(), tf.float64):强制将数据转为tf.float64类型,确保最终张量符合要求tf.reshape(x, (-1, 1)):将每个样本的2个元素拆分为shape=(1,)的独立张量tf.data.Dataset.from_tensor_slices(...):在flat_map的lambda函数内返回Dataset对象,满足flat_map的输入要求
内容的提问来源于stack exchange,提问作者gülsemin
相关产品推荐
相关产品推荐

