You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 21:47:45