TensorFlow.js聊天机器人模型预测返回随机概率问题排查
问题根因
- 核心致命错误:
model.fit()是异步执行方法,代码中未加await等待训练完成就直接返回模型,此时返回的是权重随机初始化的未训练模型,自然输出随机概率结果。 - 流程逻辑错误:每次调用
getPrediction执行预测时,都会重新触发trainModel()从头训练模型,不仅性能极差,每次随机初始化权重训练的结果也会不稳定,完全不符合模型一次训练、多次预测的常规流程。 - 冗余细节问题:词袋特征重复调用了2次
bagOfWords生成,属于无效代码;设置的Adam优化器学习率0.01对于这类小样本任务偏大,容易出现训练震荡不收敛的问题。
修复方案
- 给
model.fit()添加await关键字,确保模型完全训练收敛后再进入后续流程。 - 拆分训练与预测逻辑,程序启动时仅执行一次模型训练,训练完成后缓存模型实例,预测阶段直接复用训练好的模型,禁止每次预测都重新训练。
- 将Adam学习率调整为更通用的0.001,训练时开启数据打乱,提升收敛稳定性;增加张量内存释放逻辑,避免长时间运行出现内存泄漏。
修正后核心代码
// 全局缓存训练完成的模型,避免重复训练 let trainedModel = null; async function trainModel() { const XandY = await createTrainingData() const X = XandY[0]; const y = XandY[1]; const XTensor = tf.tensor2d(X) const yTensor = tf.tensor2d(y) const model = tf.sequential(); model.add(tf.layers.dense( { units: 128, activation: 'relu', inputShape: [108] })); model.add(tf.layers.dense( { units: 64, activation: 'relu' })); model.add(tf.layers.dense( { units: 32, activation: 'relu' })); model.add(tf.layers.dense( { units: 15, activation: 'softmax' })); model.compile({ optimizer: tf.train.adam(0.001), loss: 'categoricalCrossentropy', metrics: ['accuracy'] }); // 等待训练完成,开启训练数据打乱 await model.fit(XTensor, yTensor, { epochs: 150, validationData: [XTensor, yTensor], shuffle: true }); // 释放临时张量内存 XTensor.dispose(); yTensor.dispose(); return model } async function getPrediction(message, wordset, model) { const bow = bagOfWords(message, wordset); const input = tf.tensor2d([bow]); const prediction = model.predict(input); const prediction_values = prediction.dataSync(); const prediction_array = Array.from(prediction_values); console.log(prediction_array) // 释放预测阶段生成的临时张量 input.dispose(); prediction.dispose(); let greatestProba = 0; prediction_array.forEach((element) => { if (greatestProba < element) { greatestProba = element; } }); if (greatestProba > 0.02) { return intents[prediction_array.indexOf(greatestProba)]; } else { return 'undefined' } } async function main(){ const wordset = await getWordset(); // 启动时仅训练一次模型并缓存 trainedModel = await trainModel(); const message = "Meu pagamento não caiu"; // 预测时直接传入训练好的模型 console.log(await getPrediction(message, wordset, trainedModel)); } main()
额外优化建议
- 模型训练完成后可以调用
model.save('localstorage://chatbot-model')将模型持久化存储到本地,下次启动直接加载预训练模型,不需要每次打开页面都重新训练。 - 如果训练样本量较小,可以适当减少训练轮次,或者添加早停回调
tf.callbacks.earlyStopping避免模型过拟合。
内容的提问来源于stack exchange,提问作者Eduardo Pitta
相关产品推荐
相关产品推荐

