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

TensorFlow.js训练模型时MSE持续返回NaN问题求助

TensorFlow.js回归模型训练MSE全为NaN问题解决

问题现象

使用TensorFlow.js训练回归模型时,Epoch阶段的loss、MSE均返回NaN,模型评估时也得到NaN值。已确认输入特征(xs)和标签(ys)张量中无NaN或无穷大值,但模型无法正常训练。

相关代码

主训练流程

// 检查张量
checkForNaNs(xs, 'Features');
// 返回结果:
// Features does not contains NaNs
// Features does not contains Infinities
checkForNaNs(ys, 'Labels');
// 返回结果:
// Labels does not contains NaNs
// Labels does not contains Infinities

xs = normalizeTensor(xs);

const optimizerType = 'sgd';
const lossFunction = 'meanSquaredError';
let model = createRegressionModel(features[0].length);
const metrics = 'mse';
await trainModel(model, xs, ys, optimizerType, lossFunction, metrics);

const evalOutput = await model.evaluate(xs, ys);
console.log(`Debug: evalOutput: ${evalOutput}`)
// 返回: debug: evalOutput: Tensor
// NaN, Tensor 
// NaN

const mse = evalOutput[0].dataSync()[0]; // 获取MSE的第一个元素
console.log(`Mean Squared Error (MSE): ${mse}`);
// 返回: Mean Squared Error (MSE): NaN

工具函数定义

function createRegressionModel(inputShape) {
    return tf.sequential({
        layers: [
            tf.layers.dense({ inputShape: [inputShape], units: 10, activation: 'relu' }),
            tf.layers.dense({ units: 1, activation: 'linear' })
    ]
});
}

async function trainModel(model, xTrain, yTrain, xValidation, yValidation, optimizerType, lossFunction, metrics) {
model.compile({ optimizer: optimizerType, loss: lossFunction, metrics: metrics });
console.log(`Training model using metrix: ${metrics}`);
// 返回: Training model using metrix: mse

await model.fit(xTrain, yTrain, {
    epochs: 10,
    validationData: validationData,
    callbacks: {
        onEpochEnd: (epoch, logs) => {
            console.log(logs);
            console.log(`Epoch ${epoch + 1}: loss = ${logs.loss}, MSE = ${logs.mse}, val_loss = ${logs.val_loss}, val_MSE = ${logs.val_mse}`);
            // 示例返回: Epoch 1: loss = NaN, MSE = NaN, val_loss = undefined, val_MSE = undefined

        }
    }
 });
 }

function checkForNaNs(tensor, tensorName) {
    if (tensor.isNaN().any().dataSync()[0]) {
      console.log(`${tensorName} contains NaNs`);
    } else {
       console.log(`${tensorName} does not contains NaNs`);
    }

    if (tensor.isInf().any().dataSync()[0]) {
       console.log(`${tensorName} contains Infinities`);
    } else {
       console.log(`${tensorName} does not contains Infinities`);
    }
  }

function normalizeTensor(tensor) {
    const mean = tensor.mean(0);
    const std = tensor.sub(mean).square().mean(0).sqrt();
    return tensor.sub(mean).div(std);
  }

数据读取代码

let xs = tf.tensor2d(features, [features.length, features[0].length]);
let ys = tf.tensor2d(labels, [labels.length, 1]); 

问题排查与修复方案

1. 训练函数参数顺序完全错误

trainModel的参数定义为(model, xTrain, yTrain, xValidation, yValidation, optimizerType, lossFunction, metrics),但调用时传入的是trainModel(model, xs, ys, optimizerType, lossFunction, metrics),导致xValidation被赋值为optimizerType,optimizerType实际为undefined,模型编译时使用了无效的优化器,直接导致训练失败。

修复代码:
调整参数定义与调用顺序,简化无验证数据的场景:

// 修改trainModel参数定义
async function trainModel(model, xTrain, yTrain, optimizerType, lossFunction, metrics) {
  model.compile({ optimizer: optimizerType, loss: lossFunction, metrics: [metrics] });
  console.log(`Training model using metric: ${metrics}`);

  await model.fit(xTrain, yTrain, {
      epochs: 10,
      callbacks: {
          onEpochEnd: (epoch, logs) => {
              console.log(logs);
              console.log(`Epoch ${epoch + 1}: loss = ${logs.loss}, MSE = ${logs.mse}`);
          }
      }
   });
 }

// 调用时匹配参数顺序
await trainModel(model, xs, ys, optimizerType, lossFunction, metrics);

2. 归一化函数存在除以0风险

如果某个特征的所有值相同,其标准差std为0,执行div(std)会产生无穷大值,后续训练会出现NaN。

修复代码:
给标准差添加极小值避免除以0:

function normalizeTensor(tensor) {
    const mean = tensor.mean(0);
    const std = tensor.sub(mean).square().mean(0).sqrt();
    // 添加极小值防止除以0
    const safeStd = std.add(tf.scalar(1e-8));
    return tensor.sub(mean).div(safeStd);
  }

3. SGD默认学习率过高导致梯度爆炸

TensorFlow.js中SGD默认学习率为0.01,对于部分数据集可能过大,导致梯度爆炸出现NaN。

修复代码:
手动设置更小的学习率:

// 替换原optimizerType定义
const optimizer = tf.train.sgd(0.001);
// 编译时传入自定义优化器
model.compile({ optimizer: optimizer, loss: lossFunction, metrics: [metrics] });

4. 验证数据未正确处理

原代码中validationData未定义就传入model.fit,导致日志中val_*字段为undefined,如果需要验证数据,需正确传入[xValidation, yValidation],否则移除该参数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 06:37:05