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

基于概率的深度神经网络作业:反钓鱼算法标签映射问题

解决三分类标签(-1、0、1)映射问题

你的当前模型结构和配置不匹配三分类任务需求,导致输出是连续值而非指定的离散标签。以下是具体修改方案:

核心问题分析

  1. 模型最后层设计错误:你在中间层用了softmax(本该用于分类输出层),又接了Dense(1)无激活层,这会输出连续回归值,无法直接对应三个离散标签。
  2. 损失函数与任务不匹配:mse是回归任务损失,不适合分类场景。
  3. 指标错误:BinaryAccuracy是二分类指标,不适用于三分类。

方案一:标准三分类模型(推荐)

这是最适合离散多标签任务的做法,直接让模型输出三类的概率分布,再映射回-1、0、1。

步骤1:调整标签映射(可选但更规范)

因为Keras的分类损失函数默认要求标签从0开始索引,所以先把原始标签做映射:

  • 原始标签-1 → 0
  • 原始标签0 → 1
  • 原始标签1 → 2

步骤2:修改模型代码

def create_model(self, n_hidden = 1, n_neurons = 30, learning_rate = 3e-3, input_shape = 30):
    model = Sequential()
    model.add(InputLayer(input_shape = input_shape))
    
    for layer in range(n_hidden):
        model.add(Dense(n_neurons, activation = "relu"))
    # 输出3类的概率分布,用softmax激活
    model.add(Dense(3, activation="softmax"))
    optimizer = SGD(learning_rate = learning_rate)
    # 用稀疏分类交叉熵(适用于整数形式的标签)
    model.compile(loss=SparseCategoricalCrossentropy(), 
                  optimizer=optimizer, 
                  metrics=[SparseCategoricalAccuracy()])
    return model

步骤3:预测后映射回原始标签

训练完成后,通过以下代码将模型输出转换为-1、0、1:

import numpy as np

# 获取预测概率
pred_probs = model.predict(X_test)
# 取概率最大的类别索引
pred_classes = pred_probs.argmax(axis=1)
# 映射回原始标签
pred_labels = np.where(pred_classes == 0, -1, np.where(pred_classes == 1, 0, 1))

方案二:直接输出近似值后离散化(不推荐)

如果不想修改模型为分类结构,可以调整最后一层激活函数,再通过阈值将连续值映射为指定标签,但这种方式精度通常低于标准分类模型:

修改模型代码

def create_model(self, n_hidden = 1, n_neurons = 30, learning_rate = 3e-3, input_shape = 30):
    model = Sequential()
    model.add(InputLayer(input_shape = input_shape))
    
    for layer in range(n_hidden):
        model.add(Dense(n_neurons, activation = "relu"))
    # 用tanh将输出限制在[-1,1]区间
    model.add(Dense(1, activation="tanh"))
    optimizer = SGD(learning_rate = learning_rate)
    # 仍用MSE损失,但任务本质是回归转分类
    model.compile(loss="mse", optimizer = optimizer, metrics=["mae"])
    return model

预测后离散化

pred_values = model.predict(X_test).flatten()
# 设定阈值映射标签
pred_labels = np.where(pred_values > 0.5, 1, np.where(pred_values < -0.5, -1, 0))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 06:15:32