为何每次运行TensorFlow模型结果都变化?训练数据预测也不符预期
问题排查:TensorFlow.js模型输出不稳定且预测不符合预期
你的代码存在几个核心问题,导致模型输出波动且无法得到预期的[1,0,0]结果:
1. 数据集严重不平衡
训练数据中仅1个正样本(标签1),其余6个均为负样本(标签0)。这种极端不平衡会让模型倾向于预测0来获取高准确率,同时训练过程极不稳定,每次收敛方向都可能存在差异。
2. 模型初始化的随机性
TensorFlow.js默认随机初始化模型权重,每次运行程序时权重起点不同,导致训练后的模型参数差异较大,最终输出结果波动。
3. 训练配置不合理
- 样本量极少的情况下,带ReLU的中间层容易导致模型过拟合或无法有效学习正样本特征
- 200轮训练可能不足以让模型在极端不平衡的数据下稳定收敛
解决办法
(1)修复数据不平衡问题
- 补充正样本:收集更多用户喜欢的样本(比如
[1,1,1]的变体或其他符合正标签的数据),这是最根本的解决方式 - 使用类别权重:在
model.fit中设置classWeight,给正样本更高权重,让模型重视少数类:model.fit(inputData, labelData, { epochs: 500, classWeight: { 0: 1, 1: 6 // 正样本权重设为负样本的6倍,抵消数量差距 } })
(2)固定随机种子,消除结果波动
在代码开头添加随机种子设置,确保每次模型初始化的权重一致:
tf.setRandomSeed(42); // 固定随机种子,数值可任意选择
(3)调整模型与训练参数
- 简化模型:极小样本集下,去掉中间ReLU层,直接用线性输出层避免过拟合:
const model = tf.sequential(); model.add(tf.layers.dense({ inputShape: [3], units: 1, activation: 'sigmoid' })); - 增加训练轮次:将epochs调到500甚至1000,让模型有足够时间学习正样本特征
- 调整学习率:适当提高Adam优化器的学习率,加快收敛:
model.compile({ optimizer: tf.train.adam(0.01), // 默认学习率0.001,这里提高到0.01 loss: 'binaryCrossentropy', metrics: ['accuracy'] });
修改后的完整示例代码
tf.setRandomSeed(42); // Input data (genre, duration, lead actor) const data = [ [1, 1, 1], // like it [0, 0, 1], // doesn't like it [1, 0, 0], // doesn't like it [0, 1, 0], // doesn't like it [0, 1, 1], // doesn't like it [1, 0, 1], // doesn't like it [1, 1, 0], // doesn't like it ]; // Tags (0 - Dislike, 1 - Like) const labels = [1, 0, 0, 0, 0, 0, 0]; const model = tf.sequential(); model.add(tf.layers.dense({ inputShape: [3], units: 1, activation: 'sigmoid' })); model.compile({ optimizer: tf.train.adam(0.01), loss: 'binaryCrossentropy', metrics: ['accuracy'] }); const inputData = tf.tensor2d(data, [data.length, 3]); const labelData = tf.tensor1d(labels); model.fit(inputData, labelData, { epochs: 1000, classWeight: { 0: 1, 1: 6 } }) .then(() => { const testData = [ [1, 1, 1], // like it [0, 1, 0], // Don't like it [0, 0, 1] // Don't like it ]; const predictions = model.predict(tf.tensor2d(testData)); const predictedLabels = predictions.round().dataSync(); console.log('result:', predictedLabels); // 预期输出[1,0,0] });
内容的提问来源于stack exchange,提问作者Oleksii Shkulipa
相关产品推荐
相关产品推荐

