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

如何计算TensorFlow.js回归模型预测结果的置信度?

计算TensorFlow.js 0.10.0回归模型的预测置信度

嘿,针对你用TensorFlow.js 0.10.0做回归任务时想计算预测置信度的问题,我来给你梳理几个实用的方案:

首先得明确一点:回归任务里的「置信度」和分类任务里的概率置信度不一样,我们通常用预测区间来衡量回归预测的不确定性——也就是预测值大概率会落在某个范围里。下面的方法都是适配你使用的旧版TF.js API的:

方案1:基于训练残差的预测区间

这是最直接高效的方法,利用训练数据上的预测误差来估计新预测的不确定性:

步骤:

  1. 训练完成后,用模型对训练集trainXs做预测,得到训练阶段的预测值
  2. 计算残差(真实值减去预测值),代表模型在训练数据上的误差
  3. 计算残差的标准差,这个值反映了模型误差的平均波动幅度
  4. 对于新预测值,比如95%置信区间可以用 预测值 ± 1.96 * 残差标准差(1.96是正态分布中对应95%置信度的z值)

适配你的代码示例:

// 新增计算残差标准差的函数
async function calculateResidualStd(model, trainXs, trainYs) {
  const trainPreds = model.predict(trainXs);
  const residuals = trainYs.sub(trainPreds);
  // 计算方差(残差平方的均值)后开根号得到标准差
  const residualVariance = residuals.square().mean().dataSync()[0];
  const residualStd = Math.sqrt(residualVariance);
  return residualStd;
}

// 修改你的训练完成回调
fit(trainXs, trainYs, EPOCHS, BATCH_SIZE)
  .then(async () => {
    console.log('Done training');
    // 计算残差标准差
    const residualStd = await calculateResidualStd(model, trainXs, trainYs);
    console.log(`训练残差标准差: ${residualStd.toFixed(4)}`);
    
    const item = predict.slice([0], 1);
    console.log('待预测输入:', item.dataSync());
    
    const prediction = model.predict(item, true);
    console.log('模型预测值:');
    prediction.print();
    
    console.log('真实值:');
    expect.slice([0], 1).print();
    
    // 计算95%置信区间
    const predValue = prediction.dataSync()[0];
    const lowerBound = predValue - 1.96 * residualStd;
    const upperBound = predValue + 1.96 * residualStd;
    console.log(`95%预测区间: [${lowerBound.toFixed(4)}, ${upperBound.toFixed(4)}]`);
  });

方案2:Bootstrap自助法估计置信区间

如果想要更稳健的不确定性估计(尤其是残差不符合正态分布的场景),可以用Bootstrap方法:多次用有放回采样从训练集中抽取子集,训练多个模型,然后对同一个输入做多次预测,用这些预测结果的分位数来确定置信区间。

适配你的代码示例:

async function bootstrapPredict(trainXs, trainYs, input, bootstrapTimes = 10) {
  const predictions = [];
  const [totalSamples] = trainXs.shape;
  
  for (let i = 0; i < bootstrapTimes; i++) {
    // 生成有放回的采样索引
    const sampleIndices = tf.util.createShuffledIndices(totalSamples).slice(0, totalSamples);
    const bootstrapXs = trainXs.gather(tf.tensor1d(sampleIndices, 'int32'));
    const bootstrapYs = trainYs.gather(tf.tensor1d(sampleIndices, 'int32'));
    
    // 重新构建并训练模型
    const bootstrapModel = sequential();
    bootstrapModel.add(layers.dense({units: 16, inputShape: [37,]}));
    bootstrapModel.add(layers.dense({units: 4}));
    bootstrapModel.add(layers.dense({units: 1}));
    bootstrapModel.compile({ 
      optimizer: train.adam(LEARNING_RATE), 
      loss: 'meanSquaredError' 
    });
    
    // 静默训练(关闭日志)
    await bootstrapModel.fit(bootstrapXs, bootstrapYs, {
      batchSize: BATCH_SIZE,
      epochs: EPOCHS,
      shuffle: true,
      verbose: 0
    });
    
    // 保存预测结果
    const pred = bootstrapModel.predict(input).dataSync()[0];
    predictions.push(pred);
    // 清理内存,避免内存泄漏
    tf.dispose([bootstrapXs, bootstrapYs, bootstrapModel]);
  }
  
  // 排序后取分位数(95%置信区间取2.5%和97.5%分位数)
  predictions.sort((a, b) => a - b);
  const lowerIdx = Math.floor(0.025 * bootstrapTimes);
  const upperIdx = Math.floor(0.975 * bootstrapTimes);
  
  return {
    allPredictions: predictions,
    meanPrediction: predictions.reduce((sum, val) => sum + val, 0) / bootstrapTimes,
    lowerBound: predictions[lowerIdx],
    upperBound: predictions[upperIdx]
  };
}

// 使用方法(修改训练完成回调)
fit(trainXs, trainYs, EPOCHS, BATCH_SIZE)
  .then(async () => {
    console.log('Done training');
    const item = predict.slice([0], 1);
    console.log('待预测输入:', item.dataSync());
    
    // 执行Bootstrap预测
    const bootstrapResult = await bootstrapPredict(trainXs, trainYs, item, 10);
    console.log(`Bootstrap平均预测值: ${bootstrapResult.meanPrediction.toFixed(4)}`);
    console.log(`95% Bootstrap预测区间: [${bootstrapResult.lowerBound.toFixed(4)}, ${bootstrapResult.upperBound.toFixed(4)}]`);
    
    console.log('真实值:');
    expect.slice([0], 1).print();
  });

方案对比:

  • 方案1:计算速度快,适合简单场景,但假设残差是正态分布且方差恒定;
  • 方案2:结果更稳健,不依赖分布假设,但计算成本高(需要多次训练模型),可以根据你的需求调整bootstrapTimes的数量(次数越多结果越可靠,但耗时越长);
  • 由于你使用的是TF.js 0.10.0(较旧版本),没有后来版本中支持概率输出的层(比如distributionLambda),所以上面的两种方法是最适配当前环境的选择。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 04:18:29