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

多分类神经网络过拟合+训练缓慢求助:精度未达90%要求

多分类神经网络优化方案与改进建议

一、数据预处理优化

  • 特征标准化/归一化:输入仅9维特征,若特征尺度差异较大,会导致tanh激活函数效果受限(tanh对输入范围敏感)。建议对trainX和testX做标准化处理:
    from sklearn.preprocessing import StandardScaler
    scaler = StandardScaler()
    trainX = scaler.fit_transform(trainX)
    testX = scaler.transform(testX)
    
  • 分层数据集划分:当前train_test_split未指定stratify参数,可能导致训练/测试集类别分布失衡,影响验证精度。修改为:
    (trainX, testX, trainY, testY) = train_test_split(dataset, values, test_size=0.25, random_state=42, stratify=values)
    
  • 数据集诊断:检查样本总量与类别分布:
    • 若样本量小于1000,复杂模型极易过拟合,需大幅简化结构;
    • 若类别不平衡(某类样本占比过低),在model.fit()中添加class_weight参数,或采用SMOTE过采样、欠采样优化。

二、模型结构调整

当前模型容量远超输入维度(9维特征对应4层全连接+大神经元数),是过拟合的核心原因之一,建议调整如下:

  1. 简化模型结构:减少层数与神经元数量,匹配输入特征规模:
    visible = layers.Input(shape=(9,))
    hidden0 = layers.Dense(32, activation="LeakyReLU")(visible)
    batch0 = layers.BatchNormalization()(hidden0)
    drop0 = layers.Dropout(0.2)(batch0)
    hidden1 = layers.Dense(16, activation="LeakyReLU")(drop0)
    batch1 = layers.BatchNormalization()(hidden1)
    output = layers.Dense(4, activation="softmax")(batch1)
    model = tf.keras.Model(inputs=visible, outputs=output)
    
  2. 替换激活函数:tanh易出现梯度消失问题,改用LeakyReLU或ReLU提升梯度传递效率,加快拟合速度。
  3. 调整BN与Dropout位置:将BatchNormalization移至激活函数之前(部分场景下能提升稳定性):
    hidden0 = layers.Dense(32)(visible)
    batch0 = layers.BatchNormalization()(hidden0)
    act0 = layers.LeakyReLU()(batch0)
    drop0 = layers.Dropout(0.2)(act0)
    

三、训练策略优化

  1. 添加早停与学习率衰减:
    • 早停避免过度训练,节省时间同时防止过拟合;
    • 学习率衰减在训练后期自动降低学习率,提升收敛精度。
    callbacks = [
        tf.keras.callbacks.EarlyStopping(monitor='val_precision', patience=30, mode='max', restore_best_weights=True),
        tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=15, min_lr=1e-6)
    ]
    
    并在fit()中加入callbacks=callbacks,同时将epochs从5000降至200-300(早停会自动终止)。
  2. 调整学习率与Batch Size:
    • 当前学习率0.0001过小,导致拟合速度慢,初始改为0.001;
    • 若样本量不大,将batch_size从256调整为32/64,小批量数据的噪声能起到天然正则化效果。
  3. 确认损失函数适配性:若trainY是整数标签(非独热编码),将损失函数改为sparse_categorical_crossentropy,避免编码错误。

四、增强正则化

  • 添加L2正则化:在Dense层中加入权重正则化,限制模型权重规模:
    hidden0 = layers.Dense(32, kernel_regularizer=tf.keras.regularizers.l2(0.001), activation="LeakyReLU")(visible)
    
  • 表格数据增强:对输入特征添加轻微高斯噪声,提升模型泛化能力:
    def add_noise(x):
        noise = tf.random.normal(shape=tf.shape(x), mean=0.0, stddev=0.01, dtype=tf.float32)
        return x + noise
    
    trainX_noisy = add_noise(trainX)
    # 训练时可混合原始与噪声数据
    model.fit(tf.concat([trainX, trainX_noisy], axis=0), tf.concat([trainY, trainY], axis=0), ...)
    

改进后训练代码示例

from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import train_test_split
import tensorflow as tf
from tensorflow.keras import layers

# 数据预处理
scaler = StandardScaler()
dataset_scaled = scaler.fit_transform(dataset)
(trainX, testX, trainY, testY) = train_test_split(dataset_scaled, values, test_size=0.25, random_state=42, stratify=values)

# 模型构建
visible = layers.Input(shape=(9,))
hidden0 = layers.Dense(32, kernel_regularizer=tf.keras.regularizers.l2(0.001))(visible)
batch0 = layers.BatchNormalization()(hidden0)
act0 = layers.LeakyReLU()(batch0)
drop0 = layers.Dropout(0.2)(act0)
hidden1 = layers.Dense(16, kernel_regularizer=tf.keras.regularizers.l2(0.001))(drop0)
batch1 = layers.BatchNormalization()(hidden1)
act1 = layers.LeakyReLU()(batch1)
output = layers.Dense(4, activation="softmax")(act1)
model = tf.keras.Model(inputs=visible, outputs=output)

# 编译与训练
model.compile(optimizer=tf.keras.optimizers.Adam(0.001),
              loss='categorical_crossentropy',
              metrics=['Precision'])

callbacks = [
    tf.keras.callbacks.EarlyStopping(monitor='val_precision', patience=30, mode='max', restore_best_weights=True),
    tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=15, min_lr=1e-6)
]

history = model.fit(trainX, trainY, 
                    validation_data=(testX, testY), 
                    epochs=300, 
                    batch_size=64,
                    callbacks=callbacks)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 22:45:40