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

如何在tf.estimator.inputs.numpy_input_fn中传入多组标签列表?

嘿,你的思路其实完全没问题——用字典传递多组标签是tf.estimator处理多任务/多标签场景的标准做法,核心是要让你的模型函数(model_fn)能正确接收并处理这个字典格式的标签输入。我给你梳理下具体的实现步骤和代码示例:

1. 你的输入函数写法是对的,先保留

首先,你写的numpy_input_fn完全符合要求,用字典把两组标签传进去就行:

train_input_fn = tf.estimator.inputs.numpy_input_fn(
    x={"x": training_images},
    y={"labels1": training_labels1, "labels2": training_labels2},
    batch_size=BATCH_SIZE,
    num_epochs=None,
    shuffle=True
)

2. 关键:调整模型函数以支持多标签输入

接下来要修改你的my_cnn_model_fn,把labels参数当成字典来处理,分别为两组标签定义输出层、计算损失,最后合并损失(也可以按任务重要性加权)。这里我用类似MNIST的CNN结构举例:

def my_cnn_model_fn(features, labels, mode):
    # 第一步:定义CNN基础网络结构(和MNIST示例类似)
    input_layer = tf.reshape(features["x"], [-1, 28, 28, 1])
    
    conv1 = tf.layers.conv2d(
        inputs=input_layer,
        filters=32,
        kernel_size=[5, 5],
        padding="same",
        activation=tf.nn.relu
    )
    pool1 = tf.layers.max_pooling2d(inputs=conv1, pool_size=[2, 2], strides=2)
    
    conv2 = tf.layers.conv2d(
        inputs=pool1,
        filters=64,
        kernel_size=[5, 5],
        padding="same",
        activation=tf.nn.relu
    )
    pool2 = tf.layers.max_pooling2d(inputs=conv2, pool_size=[2, 2], strides=2)
    
    pool2_flat = tf.reshape(pool2, [-1, 7 * 7 * 64])
    dense = tf.layers.dense(inputs=pool2_flat, units=1024, activation=tf.nn.relu)
    dropout = tf.layers.dropout(
        inputs=dense, rate=0.4, training=mode == tf.estimator.ModeKeys.TRAIN
    )
    
    # 第二步:为两组标签分别定义输出层
    # 这里假设labels1是10分类任务,labels2是5分类任务,你可以根据实际调整units数量
    logits_labels1 = tf.layers.dense(inputs=dropout, units=10)
    logits_labels2 = tf.layers.dense(inputs=dropout, units=5)
    
    # 第三步:预测模式返回两组标签的预测结果
    predictions = {
        "classes_labels1": tf.argmax(input=logits_labels1, axis=1),
        "probabilities_labels1": tf.nn.softmax(logits_labels1, name="softmax_tensor_labels1"),
        "classes_labels2": tf.argmax(input=logits_labels2, axis=1),
        "probabilities_labels2": tf.nn.softmax(logits_labels2, name="softmax_tensor_labels2")
    }
    if mode == tf.estimator.ModeKeys.PREDICT:
        return tf.estimator.EstimatorSpec(mode=mode, predictions=predictions)
    
    # 第四步:计算两组标签的损失并合并
    # 如果你的标签是独热编码,要用tf.losses.softmax_cross_entropy,这里假设是稀疏标签(整数)
    loss_labels1 = tf.losses.sparse_softmax_cross_entropy(labels=labels["labels1"], logits=logits_labels1)
    loss_labels2 = tf.losses.sparse_softmax_cross_entropy(labels=labels["labels2"], logits=logits_labels2)
    # 可以根据任务优先级加权,比如更侧重labels1就用 0.7*loss_labels1 + 0.3*loss_labels2
    total_loss = loss_labels1 + loss_labels2
    
    # 第五步:训练模式定义优化器
    if mode == tf.estimator.ModeKeys.TRAIN:
        optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.001)
        train_op = optimizer.minimize(
            loss=total_loss,
            global_step=tf.train.get_global_step()
        )
        return tf.estimator.EstimatorSpec(mode=mode, loss=total_loss, train_op=train_op)
    
    # 第六步:评估模式定义两组标签的评估指标
    eval_metric_ops = {
        "accuracy_labels1": tf.metrics.accuracy(
            labels=labels["labels1"], predictions=predictions["classes_labels1"]),
        "accuracy_labels2": tf.metrics.accuracy(
            labels=labels["labels2"], predictions=predictions["classes_labels2"])
    }
    return tf.estimator.EstimatorSpec(
        mode=mode, loss=total_loss, eval_metric_ops=eval_metric_ops)

3. 创建Estimator并训练

最后,用修改后的模型函数创建Estimator实例,然后调用train方法就可以了:

# 创建Estimator
my_cnn = tf.estimator.Estimator(
    model_fn=my_cnn_model_fn, model_dir="/tmp/multi_label_cnn_model")

# 启动训练
my_cnn.train(
    input_fn=train_input_fn,
    steps=20000)

一些注意事项

  • 确保你的training_labels1和training_labels2的类型/形状和损失函数匹配:如果是整数标签用sparse_softmax_cross_entropy,独热编码用softmax_cross_entropy
  • 损失合并的权重可以根据任务的重要性调整,让模型更偏向优化某一个任务
  • 评估时可以分别查看两个任务的准确率,方便监控每个任务的表现

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:50:42