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

TensorFlow三分类任务中如何添加概率输出约束?

解决方案:TensorFlow 3类分类的概率约束实现

一、自定义损失函数实现约束

模型输出的y_pred就是你需要的概率向量[p0, p1, p2],直接在自定义损失函数里就能访问。我们可以在基础分类损失(比如交叉熵)上添加约束惩罚项,当模型输出不满足要求时,额外增加损失,迫使模型调整输出满足约束。

代码实现

import tensorflow as tf

def constrained_loss(y_true, y_pred):
    # 处理标签格式:如果是one-hot编码,转成类别索引
    y_true = tf.argmax(y_true, axis=1) if len(y_true.shape) > 1 else y_true
    
    # 提取三个类别的概率
    p0 = y_pred[:, 0]
    p1 = y_pred[:, 1]
    p2 = y_pred[:, 2]
    
    # 定义约束惩罚项:不满足条件时产生损失
    # 标签为0时,需要p1 > p2,否则惩罚p2 - p1的差值
    loss_0 = tf.where(y_true == 0, tf.maximum(0.0, p2 - p1), 0.0)
    # 标签为2时,需要p1 > p0,否则惩罚p0 - p1的差值
    loss_2 = tf.where(y_true == 2, tf.maximum(0.0, p0 - p1), 0.0)
    # 标签为1时,需要p1 > min(p0,p2),否则惩罚min(p0,p2) - p1的差值
    loss_1 = tf.where(y_true == 1, tf.maximum(0.0, tf.minimum(p0, p2) - p1), 0.0)
    
    # 基础分类损失用交叉熵,加上约束惩罚(alpha是惩罚权重,可按需调整)
    ce_loss = tf.keras.losses.sparse_categorical_crossentropy(y_true, y_pred) if len(y_true.shape)==1 else tf.keras.losses.categorical_crossentropy(y_true, y_pred)
    alpha = 1.0
    total_loss = ce_loss + alpha * (loss_0 + loss_2 + loss_1)
    
    return total_loss

使用方式

假设你的模型输出层是Dense(3, activation='softmax'),直接把自定义损失传入compile:

model.compile(optimizer='adam', loss=constrained_loss, metrics=['accuracy'])

二、修改标签的间接实现方式

如果不想写复杂的损失函数,可以通过调整目标标签的概率分布,让模型训练时自然贴合约束要求:

  • 标签0的样本,构造目标分布让p0最大、p1次之、p2最小,比如[0.9, 0.09, 0.01]
  • 标签2的样本,构造目标分布让p2最大、p1次之、p0最小,比如[0.01, 0.09, 0.9]
  • 标签1的样本,构造目标分布让p1至少大于其中一个类别概率,比如[0.3, 0.4, 0.3]

代码实现

import tensorflow as tf

def adjust_labels(y_true):
    # 输入是类别索引数组,输出调整后的概率目标
    adjusted = []
    for y in y_true:
        if y == 0:
            adjusted.append([0.9, 0.09, 0.01])
        elif y == 2:
            adjusted.append([0.01, 0.09, 0.9])
        else:
            adjusted.append([0.3, 0.4, 0.3])
    return tf.convert_to_tensor(adjusted, dtype=tf.float32)

# 编译模型时用常规的分类损失
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])

# 训练时传入调整后的标签
train_labels_adjusted = adjust_labels(train_labels)
model.fit(train_data, train_labels_adjusted, epochs=10, validation_split=0.1)

注意点

目标分布的数值可以根据实际情况调整,只要满足你的约束条件即可,这种方式的优点是实现简单,但灵活性不如自定义损失函数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 13:55:16