TensorFlow.js天气预测模型始终预测'Haze'问题排查与优化咨询
问题排查与优化方案
1. 优先排查数据问题(核心根源)
- 类别不平衡验证:统计数据中各类别的样本占比,如果'Haze'样本占比过高(比如超70%),模型会默认预测该类别来获取“伪高准确率”。解决方式:
- 对少数类别过采样(复制样本)或对'Haze'类欠采样(随机删除部分样本)
- 训练时传入
classWeight参数,给少数类别分配更高权重,比如classWeight: { 'Haze': 1, 'Rain': 6, 'Cloudy': 4 }
- 数据预处理校验:
- 所有数值特征必须做标准化/归一化,可使用
tf.layers.normalization()层自动适配,或手动计算均值标准差后缩放到[0,1]区间 - 多分类标签需做独热编码(用
tf.oneHot()实现),若用整数标签,损失函数需改用sparseCategoricalCrossentropy - 检查训练/验证集划分是否随机,避免某类别集中在单一集合中
- 所有数值特征必须做标准化/归一化,可使用
2. 模型结构调整
- 输入形状匹配:输入层
inputShape必须与特征维度一致,比如单样本含8个特征则设为inputShape: [8] - 密集层配置:
- 先从基线模型开始:1个隐藏层(单元数设为特征数的1.5-2倍,比如16个单元),激活函数用
relu;输出层单元数等于类别总数,激活函数用softmax - 若基线效果差,逐步增加隐藏层(最多2-3层,避免过拟合),每层单元数递减(比如16→8→类别数)
- 先从基线模型开始:1个隐藏层(单元数设为特征数的1.5-2倍,比如16个单元),激活函数用
- 防过拟合措施:在隐藏层后加入
tf.layers.dropout({rate: 0.2}),或给密集层添加L2正则化kernelRegularizer: tf.regularizers.l2(0.001)
3. 优化器与损失函数选择
- 损失函数:
- 独热编码标签用
categoricalCrossentropy,整数标签用sparseCategoricalCrossentropy - 类别不平衡严重时,可自定义
focalLoss降低易分类样本的权重
- 独热编码标签用
- 优化器:
- 优先用
Adam优化器,默认学习率0.001,若训练停滞可微调至0.0001或0.01 - 避免直接用
SGD,除非搭配动量参数(momentum: 0.9),否则收敛速度慢易陷入局部最优
- 优先用
4. 训练过程优化
- 早停机制:加入
tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 5 }),验证损失连续5轮不下降则停止训练,避免无效迭代 - 日志分析:打印每轮的训练/验证准确率、损失,若训练准确率高但验证准确率低是过拟合,两者都低则是欠拟合
- 训练轮数调整:若41轮后停滞,先解决数据/模型结构问题,再考虑增加轮数
示例TensorFlow.js代码片段
// 特征归一化层 const normLayer = tf.layers.normalization({ inputShape: [featureNum] }); normLayer.adapt(trainFeatures); // 构建模型 const model = tf.sequential(); model.add(normLayer); model.add(tf.layers.dense({ units: 16, activation: 'relu', kernelRegularizer: tf.regularizers.l2(0.001) })); model.add(tf.layers.dropout({ rate: 0.2 })); model.add(tf.layers.dense({ units: classNum, activation: 'softmax' })); // 编译模型(独热编码标签) model.compile({ optimizer: tf.train.adam(0.001), loss: 'categoricalCrossentropy', metrics: ['accuracy'] }); // 启动训练 const history = await model.fit(trainFeatures, trainLabels, { epochs: 100, validationSplit: 0.2, callbacks: [tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 5 })], classWeight: precomputedClassWeights });
内容的提问来源于stack exchange,提问作者sribasu
相关产品推荐
相关产品推荐

