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

使用tf.estimator训练模型遇ValueError:features需为Tensor字典

解决tf.estimator训练时的ValueError及形状不匹配问题

嘿,作为第一次上手tf.estimator的新手,碰到这些问题真的太正常了!我帮你一步步拆解,把这两个坑都填上~

第一个错误:ValueError: features should be a dictionary of Tensors. Given type: `

这个错误的核心原因很直接:tf.estimator要求输入函数返回的features必须是字典格式——键是你的特征名称,值是对应的Tensor张量。很多新手容易直接把数组、张量列表丢进去,这就会触发这个报错。

正确的输入函数写法

不管你用numpy数组还是其他数据源,输入函数必须返回(features_dict, labels)的结构。举个例子,假设你的3个输入特征是x1、x2、x3(都是形状为[样本数, 1]的数值型数组),标签是二分类的labels,输入函数可以这么写:

import tensorflow as tf

def train_input_fn(x1, x2, x3, labels, batch_size=32):
    # 把3个特征打包成字典,键是自定义的特征名,值转成Tensor
    features = {
        'feature1': tf.convert_to_tensor(x1, dtype=tf.float32),
        'feature2': tf.convert_to_tensor(x2, dtype=tf.float32),
        'feature3': tf.convert_to_tensor(x3, dtype=tf.float32)
    }
    # 构建Dataset,自动处理分批、打乱等操作
    dataset = tf.data.Dataset.from_tensor_slices((features, labels))
    dataset = dataset.shuffle(buffer_size=len(labels)).repeat().batch(batch_size)
    # 返回迭代器的元素
    return dataset.make_one_shot_iterator().get_next()

第二个错误:输入形状不匹配

这个问题通常是因为你模型里的输入层维度,和传入的特征维度没对齐。你的模型是3个输入+单神经元二分类,所以需要把3个特征的张量拼接成一个[batch_size, 3]的张量,再喂给输出层。

正确的模型函数写法

自定义model_fn时,要从features字典里取出每个特征,拼接后再传入神经元:

def model_fn(features, labels, mode):
    # 从字典中取出每个特征的张量
    feature1 = features['feature1']
    feature2 = features['feature2']
    feature3 = features['feature3']
    
    # 关键一步:把3个特征拼接成[batch_size, 3]的张量,避免形状不匹配
    input_layer = tf.concat([feature1, feature2, feature3], axis=1)
    
    # 单神经元二分类输出层,用sigmoid激活输出概率
    logits = tf.layers.dense(inputs=input_layer, units=1, activation=None)
    probabilities = tf.sigmoid(logits)
    
    # 预测模式:返回预测结果
    if mode == tf.estimator.ModeKeys.PREDICT:
        predictions = {
            'probabilities': probabilities,
            'class': tf.cast(probabilities > 0.5, tf.int32)
        }
        return tf.estimator.EstimatorSpec(mode=mode, predictions=predictions)
    
    # 训练模式:定义损失和优化器
    loss = tf.losses.sigmoid_cross_entropy(multi_class_labels=labels, logits=logits)
    if mode == tf.estimator.ModeKeys.TRAIN:
        optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.01)
        train_op = optimizer.minimize(loss, global_step=tf.train.get_global_step())
        return tf.estimator.EstimatorSpec(mode=mode, loss=loss, train_op=train_op)
    
    # 评估模式:计算准确率等指标
    if mode == tf.estimator.ModeKeys.EVAL:
        eval_metric_ops = {
            'accuracy': tf.metrics.accuracy(
                labels=tf.cast(labels > 0.5, tf.int32), 
                predictions=tf.cast(probabilities > 0.5, tf.int32)
            )
        }
        return tf.estimator.EstimatorSpec(mode=mode, loss=loss, eval_metric_ops=eval_metric_ops)

完整的训练、评估、预测流程

把上面的代码串起来,替换成你自己的真实数据就可以跑通了:

import numpy as np

# 模拟你的训练数据(替换成你自己的数据)
sample_num = 1000
x1 = np.random.randn(sample_num, 1)
x2 = np.random.randn(sample_num, 1)
x3 = np.random.randn(sample_num, 1)
# 生成二分类标签
labels = np.where(x1 + x2 + x3 > 0, 1, 0).reshape(-1, 1)

# 初始化estimator
estimator = tf.estimator.Estimator(model_fn=model_fn, model_dir='./my_model')

# 启动训练
estimator.train(input_fn=lambda: train_input_fn(x1, x2, x3, labels), steps=1000)

# 评估(假设你有测试数据x1_test, x2_test, x3_test, labels_test)
# x1_test = np.random.randn(200, 1)
# x2_test = np.random.randn(200, 1)
# x3_test = np.random.randn(200, 1)
# labels_test = np.where(x1_test + x2_test + x3_test > 0, 1, 0).reshape(-1, 1)
eval_results = estimator.evaluate(input_fn=lambda: train_input_fn(x1_test, x2_test, x3_test, labels_test, batch_size=32), steps=100)
print('评估结果:', eval_results)

# 预测(假设你有预测数据x1_pred, x2_pred, x3_pred)
# x1_pred = np.random.randn(50, 1)
# x2_pred = np.random.randn(50, 1)
# x3_pred = np.random.randn(50, 1)
predictions = estimator.predict(input_fn=lambda: train_input_fn(x1_pred, x2_pred, x3_pred, np.zeros(len(x1_pred)), batch_size=32))
for idx, pred in enumerate(predictions):
    print(f"样本{idx+1}: 预测类别={pred['class'][0]}, 概率={pred['probabilities'][0]:.4f}")

最后再提醒几个容易踩的小坑

  • 输入函数一定要返回字典格式的features,别直接传张量列表或者单个张量!
  • 特征拼接时要注意axis=1,保证是按特征维度拼接,而不是样本维度。
  • 标签的形状要和logits一致(都是[batch_size, 1]),别用一维数组,不然也会触发形状不匹配。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 10:32:36