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

在TensorFlow的DNNClassifier中正确构建input_fn的技术咨询

解决DNNClassifier中input_fn构建及特征列的常见问题

我看了你用TensorFlow的DNNClassifier构建多分类模型的代码,确实input_fn和特征列的配置很容易踩坑,我帮你梳理下代码里的问题,再给出完整的修正方案:

核心问题梳理

  1. 特征列的列名错误:你定义特征列时写了"df.col1",这是错误的——pandas DataFrame的列名直接写列本身的名称(比如你的列叫col1就写"col1"),不需要加df.前缀。
  2. 数值列的错误转换:indicator_column是用来处理分类列的,数值列不需要转成indicator类型,直接用numeric_column即可。
  3. 数据集划分的问题:你把整个df传给了train_test_split的第一个参数,这样训练数据会包含标签列,应该只传入特征列组成的DataFrame。
  4. 标签类型不匹配:DNNClassifier要求标签是0到n_classes-1的整数,如果你的原始标签是字符串(比如D、BBB等),必须先映射为整数。

修正后的完整代码

1. 数据导入与预处理

import pandas as pd
import numpy as np
import tensorflow as tf
from sklearn.model_selection import train_test_split

# 读取数据
df = pd.read_csv('chunk.csv')

# 处理标签:将字符串标签映射为整数(多分类要求标签为0~n-1的整数)
# 假设你的标签列是"MoreClass",先获取所有唯一标签并建立映射
unique_labels = df["MoreClass"].unique()
label_mapping = {label: idx for idx, label in enumerate(unique_labels)}
df["label"] = df["MoreClass"].map(label_mapping)

# 定义特征列和标签列
FEATURES = ["col1", "col2", ...]  # 替换成你的实际特征列列表(不要包含标签列)
LABEL = "label"

# 划分训练集和测试集:x是特征数据,y是标签数据
x_train, x_test, y_train, y_test = train_test_split(df[FEATURES], df[LABEL], test_size=0.2)

2. 构建特征列

# 分类列和数值列的定义
CATEGORICAL_COLUMNS = ["col1", ...]  # 替换成你的分类列
CONTINUOUS_COLUMNS = ["col2", ...]   # 替换成你的数值列

feature_columns = []

# 处理分类列:高基数分类列用哈希桶+嵌入层
for col in CATEGORICAL_COLUMNS:
    # 哈希桶大小根据你的类别基数调整,一般比实际类别数大一些
    cat_col = tf.feature_column.categorical_column_with_hash_bucket(col, hash_bucket_size=1000)
    # 嵌入层维度根据需求调整,一般是类别基数的平方根左右
    emb_col = tf.feature_column.embedding_column(cat_col, dimension=30)
    feature_columns.append(emb_col)

# 处理数值列:直接使用numeric_column即可
for col in CONTINUOUS_COLUMNS:
    num_col = tf.feature_column.numeric_column(col)
    feature_columns.append(num_col)

3. 定义DNNClassifier

classifier = tf.estimator.DNNClassifier(
    feature_columns=feature_columns,
    hidden_units=[10, 10],  # 隐藏层神经元数量,可根据任务调整
    n_classes=len(unique_labels),  # 类别数用实际的唯一标签数量,这里是10
    dropout=0.1  # dropout比例,防止过拟合
)

4. 正确构建input_fn

# 训练用input_fn:num_epochs=None表示无限循环(适合训练时迭代),shuffle=True打乱数据
def create_train_input_fn():
    return tf.estimator.inputs.pandas_input_fn(
        x=x_train,
        y=y_train,
        batch_size=32,  # 批次大小,根据内存调整
        num_epochs=None,
        shuffle=True
    )

# 测试用input_fn:num_epochs=1表示只遍历一次数据,shuffle=False不需要打乱
def create_test_input_fn():
    return tf.estimator.inputs.pandas_input_fn(
        x=x_test,
        y=y_test,
        batch_size=32,
        num_epochs=1,
        shuffle=False
    )

5. 训练与评估模型

# 训练模型:steps表示训练步数,可根据任务调整
classifier.train(input_fn=create_train_input_fn(), steps=1000)

# 评估模型
eval_result = classifier.evaluate(input_fn=create_test_input_fn())
print(f"Test set accuracy: {eval_result['accuracy']:.4f}")

关键注意事项

  • input_fn的本质:input_fn是一个返回(features, labels)元组或TensorFlow Dataset对象的函数,pandas_input_fn已经帮你封装了从pandas数据到TF数据集的转换,只需要确保传入的x和y是纯特征和纯标签数据。
  • 分类列的选择:如果是低基数分类列(比如类别数<10),可以用categorical_column_with_vocabulary_list替代哈希桶,再转成indicator_column,效果更稳定;高基数分类列才适合用哈希桶+嵌入。
  • 标签必须是整数:如果你的标签已经是整数但范围不对(比如从1开始),记得转成从0开始的索引。

内容的提问来源于stack exchange,提问作者user9149532

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:59:05