训练算法输出值走向异常,求问题排查及代码优化建议
问题描述
点击网站训练按钮时,输出节点值变化方向完全一致,偶尔还会出现方向错误的情况。
尝试通过调试查看权重和节点值排查问题,但无法准确定位故障点。
代码与优化需求
以下是train函数代码,作为编程新手,同时希望获得代码优化建议:
function train() { //Define variables var deltaError = 0; var sum = 0; var errorSum = 0; var valueCounter = 0; var valueCounter2 = 0; var numHiddenLayers = hiddenLayers.length; var hiddenLayerIndex = 1; var deltaWeights_ho = []; var error_k = []; var deltaWeights_ih = []; var error_j = []; //Clean inputs and normalize inputFeildsInputNode.forEach((inputFeild, index) => { cleanedInputs.push((inputFeild.value - min) / (max - min)); inputNodes[index].value = cleanedInputs[index]; }); //Loop through first hidden layer and calculate values for each node of first hidden layer hiddenLayers[0].forEach((node, nodeIndex) => { inputNodes.forEach((inputNode, inputIndex) => { sum += inputNode.value * inputHiddenWeights[valueCounter].value; valueCounter++; }); node.value = sigmoid(sum); sum = 0; }); valueCounter = 0; sum = 0; //Loop through all hidden layers and calculate values for each node of each hidden layer while (hiddenLayerIndex < numHiddenLayers) { hiddenLayers[hiddenLayerIndex].forEach((node, nodeIndex) => { hiddenLayers[hiddenLayerIndex - 1].forEach((hiddenNode, hiddenNodeIndex) => { sum += hiddenNode.value * hiddenHiddenWeights[valueCounter].value; valueCounter++; }); node.value = sigmoid(sum); sum = 0; sum = 0; }); valueCounter = 0; sum = 0; hiddenLayerIndex++; valueCounter2 = 0; } //Loop through output nodes and calculate values for each output node outputNodes.forEach((node, nodeIndex) => { hiddenLayers[hiddenLayers.length - 1].forEach((hiddenNode, hiddenNodeIndex) => { sum += hiddenNode.value * hiddenOutputWeights[valueCounter].value; valueCounter++; }); node.value = sigmoid(sum); console.log(node.value); valueCounter = 0; //Loop through all output nodes and calculate error for each output node error_k[nodeIndex] = (inputFeildsOutputNode[nodeIndex].value - node.value); hiddenLayers[hiddenLayers.length - 1].forEach((hiddenNode, hiddenNodeIndex) => { deltaWeights_ho[valueCounter2] = -2 * error_k[nodeIndex] * hiddenNode.value * node.value * (1 - node.value); hiddenOutputWeights[valueCounter2].value -= deltaWeights_ho[valueCounter2] * learningRate; valueCounter2++; }); }); valueCounter2 = 0; valueCounter = 0; //Loop through all hidden layers and calculate error for each node of each hidden layer using the error of the output nodes and calculate delta weights for each connection between hidden and output nodes hiddenLayers[hiddenLayers.length - 1].forEach((node, nodeIndex) => { error_j[nodeIndex] = 0; error_k.forEach((error, index) => { error_j[nodeIndex] += error * hiddenOutputWeights[valueCounter2].value; valueCounter2++; }); //error_j[nodeIndex] = error_j[nodeIndex] * node.value * (1 - node.value); valueCounter2 = outputNodes.length * valueCounter; deltaError = -2 * error_j[nodeIndex] * node.value * (1 - node.value); }); valueCounter2 = 0; valueCounter = 0; sum = 0; hiddenLayerIndex = 1; //Loop through all hidden layers and calculate delta weights for each connection between hidden and hidden nodes while (hiddenLayerIndex < numHiddenLayers) { hiddenLayers[hiddenLayerIndex].forEach((node, nodeIndex) => { hiddenLayers[hiddenLayerIndex - 1].forEach((hiddenNode, hiddenNodeIndex) => { deltaError = -2 * error_j[nodeIndex] * node.value * (1 - node.value) * hiddenNode.value; hiddenHiddenWeights[valueCounter2].value -= deltaError * learningRate; valueCounter2++; }); }); hiddenLayerIndex++; } valueCounter = 0; sum = 0; outputNodes.forEach((node, nodeIndex) => { console.log(node.value); }); valueCounter2 = 0; valueCounter = 0; sum = 0; console.log("Training the neural network..."); // Add your neural network training logic here }
问题分析与修复建议
核心问题定位
- 输出节点变化方向一致:反向传播时隐藏层误差计算逻辑错误。计算
error_j时,valueCounter2累加后未正确重置,导致权重索引越界,所有隐藏节点复用相同权重值计算误差,最终输出节点更新方向趋同。 - 方向错误:权重更新公式符号逻辑冲突。
deltaWeights_ho已包含负号(-2 * error_k[...]),但权重更新又做减法,导致更新方向与预期相反。
具体修复步骤
- 修正
error_j的权重索引:每个隐藏节点的误差计算起始索引应为nodeIndex * outputNodes.length,确保对应正确的输出权重:hiddenLayers[hiddenLayers.length - 1].forEach((node, nodeIndex) => { error_j[nodeIndex] = 0; valueCounter2 = nodeIndex * outputNodes.length; // 修正索引起始位置 error_k.forEach((error, index) => { error_j[nodeIndex] += error * hiddenOutputWeights[valueCounter2].value; valueCounter2++; }); error_j[nodeIndex] *= node.value * (1 - node.value); // 恢复sigmoid导数计算 deltaError = -2 * error_j[nodeIndex]; // 简化计算,导数已整合到error_j }); - 调整权重更新符号:将减法改为加法,或去掉
deltaWeights_ho中的负号,推荐后者让公式更直观:deltaWeights_ho[valueCounter2] = 2 * error_k[nodeIndex] * hiddenNode.value * node.value * (1 - node.value); hiddenOutputWeights[valueCounter2].value += deltaWeights_ho[valueCounter2] * learningRate; - 恢复隐藏层误差的导数计算:取消注释
error_j[nodeIndex] = error_j[nodeIndex] * node.value * (1 - node.value);,这是反向传播中隐藏层误差计算的必要步骤。
代码优化建议
- 减少全局变量依赖:将
inputNodes、hiddenLayers等变量作为函数参数传入,避免全局污染,提升可维护性。 - 用索引计算替代计数器:通过节点索引直接计算权重位置,比如输入到第一层隐藏层的权重索引为
nodeIndex * inputNodes.length + inputIndex,避免计数器出错。 - 提取重复逻辑为函数:将前向传播计算节点值的逻辑封装成函数,减少代码冗余:
function computeLayerValues(prevLayer, currLayer, weights) { currLayer.forEach((node, nodeIdx) => { let sum = 0; prevLayer.forEach((prevNode, prevIdx) => { sum += prevNode.value * weights[nodeIdx * prevLayer.length + prevIdx].value; }); node.value = sigmoid(sum); }); } - 变量声明优化:用
let/const替代var,避免变量提升问题;删除未使用的变量(如errorSum)。 - 添加关键注释:对反向传播的公式各部分添加注释,说明误差项、导数项的含义,方便后续调试。
内容的提问来源于stack exchange,提问作者Christian Phillips
相关产品推荐
相关产品推荐

