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

Tensorflow.js中DDPG训练代码问题:梯度变量配对缺失致Loss过高

解决DDPG TensorFlow.js实现中的梯度配对与输入格式问题

你的代码存在三个核心问题,正是这些问题导致训练时Loss过高,下面逐一说明并给出修复方案:

1. 梯度与可训练变量的配对错误

TensorFlow.js的applyGradients方法不能直接传入梯度数组,需要将每个梯度与其对应的可训练变量配对成{grad, variable}结构的对象数组。这对应Python代码中的zip(critic_grad, critic_model.trainable_variables)操作。

修复代码(批评家部分)

const criticGradVars = criticGradient.grads.map((grad, idx) => ({
  grad: grad,
  variable: criticTrainableVars[idx]
}));
this.criticModel.optimizer.applyGradients(criticGradVars);

修复代码(演员部分)

const actorGradVars = actorGradient.grads.map((grad, idx) => ({
  grad: grad,
  variable: actorTrainableVars[idx]
}));
this.actorModel.optimizer.applyGradients(actorGradVars);

2. 批评家模型输入格式不匹配

Python代码中批评家模型接受状态和动作两个独立输入,但你的TF.js代码将两者拼接成单个张量传入,导致模型输入维度与训练逻辑完全不符。

修复代码(所有批评家模型调用)

将拼接操作替换为传入输入数组:

// 目标Q值计算
const targetCriticQs = this.targetCriticModel.apply([nextStates, targetActions], { training: true }) as tf.Tensor;

// 当前Q值计算
const criticQs = this.criticModel.apply([states, actions], { training: true }) as tf.Tensor;

// 演员损失计算中的Q值
const criticQs = this.criticModel.apply([states, policyActions], { training: true }) as tf.Tensor;

3. 未使用训练模式调用模型

Python代码中调用模型时指定了training=True,但TF.js的predict方法默认使用推理模式(会跳过dropout、批量归一化的训练逻辑),需要改用apply方法并显式设置training: true。

修复代码(所有模型调用)

// 目标演员模型预测
const targetActions = this.targetActorModel.apply(nextStates, { training: true }) as tf.Tensor;

// 当前演员模型预测
const policyActions = this.actorModel.apply(states, { training: true }) as tf.Tensor;

完整修复后的训练代码

const batch = this.memory.getMinibatch(this.config.replayBatchSize);
const states = this.actorService.getStateTensor(this.actor, ...batch.map(s => s.state));
const nextStates = this.actorService.getStateTensor(this.actor, ...batch.map(s => s.nextState));
const rewards = tf.tensor2d(batch.map(s => s.reward), [batch.length, 1], 'float32');
const actions = this.actorService.getActionTensor(...batch.map(s => s.action));

// 批评家损失计算与更新
const criticLossFunction = () => tf.tidy(() => {
  let targetQs: tf.Tensor;
  if (this.config.discountRate === 0) {
    targetQs = rewards;
  } else {
    const targetActions = this.targetActorModel.apply(nextStates, { training: true }) as tf.Tensor;
    const targetCriticQs = this.targetCriticModel.apply([nextStates, targetActions], { training: true }) as tf.Tensor;
    targetQs = rewards.add(targetCriticQs.mul(this.config.discountRate));
  }
  const criticQs = this.criticModel.apply([states, actions], { training: true }) as tf.Tensor;
  const criticLoss = tf.losses.meanSquaredError(targetQs, criticQs);
  return criticLoss.asScalar();
});
const criticTrainableVars = this.criticModel.getWeights(true) as tf.Variable<tf.Rank>[];
const criticGradient = tf.variableGrads(criticLossFunction, criticTrainableVars);
const criticGradVars = criticGradient.grads.map((grad, idx) => ({
  grad: grad,
  variable: criticTrainableVars[idx]
}));
this.criticModel.optimizer.applyGradients(criticGradVars);
tf.dispose(criticGradient);

// 演员损失计算与更新
const actorLossFunction = () => tf.tidy(() => {
  const policyActions = this.actorModel.apply(states, { training: true }) as tf.Tensor;
  const criticQs = this.criticModel.apply([states, policyActions], { training: true }) as tf.Tensor;
  const actorLoss = tf.mean(criticQs.mul(-1));
  return actorLoss.asScalar();
});
const actorTrainableVars = this.actorModel.getWeights(true) as tf.Variable<tf.Rank>[];
const actorGradient = tf.variableGrads(actorLossFunction, actorTrainableVars);
const actorGradVars = actorGradient.grads.map((grad, idx) => ({
  grad: grad,
  variable: actorTrainableVars[idx]
}));
this.actorModel.optimizer.applyGradients(actorGradVars);
const actorLoss = actorGradient.value.dataSync()[0];
tf.dispose(actorGradient);

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 02:01:14