如何用Pandas创建与tf.data.experimental.make_csv_dataset结构一致的Keras监督训练可用TF Dataset?
解决Pandas创建TF Dataset的标签格式问题 & Keras监督训练数据结构说明
一、修正Pandas创建Dataset的代码,匹配make_csv_dataset的结构
你当前代码的问题在于把标签包装成了字典,而tf.data.experimental.make_csv_dataset返回的标签是单个张量(对应label_name指定的列)。只需要把标签部分改成直接取survived列即可,不需要用dict()包裹:
import pandas as pd import tensorflow as tf titanic_file = tf.keras.utils.get_file("titanic_train.csv", "https://storage.googleapis.com/tf-datasets/titanic/train.csv") df = pd.read_csv(titanic_file) # 修正后的代码:标签直接取survived列,而非包装成字典 titanic_from_pandas = tf.data.Dataset.from_tensor_slices(( dict(df.loc[:, df.columns != 'survived']), df['survived'] )) for row in titanic_from_pandas.batch(1).take(1): features = row[0] label = row[1] for feature, value in features.items(): print(f"{feature:20s}: {value}") print(f"label/survived : {label}")
运行后输出就和make_csv_dataset的结果完全一致:
sex : [b'male'] age : [22.] n_siblings_spouses : [1] parch : [0] fare : [7.25] class : [b'Third'] deck : [b'unknown'] embark_town : [b'Southampton'] alone : [b'n'] label/survived : [0]
二、Keras监督训练(features, labels)结构的官方文档定义
这个(features, labels)的数据集结构是TensorFlow Keras中模型训练输入的标准格式,在tf.keras.Model.fit()方法的官方文档里有明确定义:
- 当使用
tf.data.Dataset作为fit()的输入时,数据集的每个元素必须是一个元组(inputs, targets):inputs可以是单个张量、张量列表,或者像你例子中的字典(对应模型的多输入场景,字典的键需要和模型输入层的名称匹配);targets是对应模型输出的标签张量(或张量列表/字典,用于多输出模型)。
这个结构是监督学习的核心输入形式,确保模型能正确区分特征与标签,完成训练流程。
内容的提问来源于stack exchange,提问作者mon
相关产品推荐
相关产品推荐

