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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 10:09:15