Matlab强化学习:自定义DDPGAgent获取Q值梯度报错问题
自定义DDPG智能体梯度计算错误排查
我在实现自定义强化学习智能体时,尝试计算Q值相对于actor网络参数的梯度,在learnImpl和actorupdate函数中遇到错误。已确认相关变量均为dlarray类型,但调用dlgradient时仍触发如下错误:
Error using dlarray/dlgradient (line 115)
'dlgradient' inputs must be traced dlarray objects or cell arrays, structures or tables containing traced dlarray objects. To enable tracing, use 'dlfeval'.
我已在learnImpl函数中通过dlfeval调用actorupdate,但问题仍未解决,相关代码如下:
classdef CustomDDPGAgent < rl.agent.CustomAgent properties %actor NN actor %critic for tracking target critic_track %critic for obstacle avoidance critic_obstacle %dimensions statesize end methods %constructor function function obj = CustomDDPGAgent(ActorNN,Critic_Track,Critic_Obst,statesize,actionsize) %(required) call abstract class constructor obj = obj@rl.agent.CustomAgent(); %define observation + action space obj.ObservationInfo = rlNumericSpec([statesize 1]); obj.ActionInfo = rlNumericSpec([actionsize 1],LowerLimit = -1,UpperLimit = 1); obj.SampleTime = 0.01; %define the actor and 2 critics obj.actor = ActorNN; obj.critic_track = Critic_Track; obj.critic_obstacle = Critic_Obst; %record observation dimensions obj.statesize = statesize; end end methods (Access = protected) %Actor update based on Q value function actorgradient = actorupdate(obj,Observation) Obs_Obstacle = {dlarray([])}; for index = 1:20 Obs_Obstacle{1}(index) = Observation{1}(index); end disp(Observation); disp(Obs_Obstacle); action = evaluate(obj.actor,Observation,UseForward=true); disp(action); %Obtained combined Q values Qtrack = getValue(obj.critic_track,Observation,action); Qobstacle = getValue(obj.critic_obstacle,Obs_Obstacle,action); Qtotal = Qtrack + Qobstacle; Qtotal = sum(Qtotal); disp(Qtotal); %obtain gradient of Q value wrt parameters of actor network actorgradient = dlgradient(Qtotal,obj.actor.Learnables); end %Action method function action = getActionImpl(obj,Observation) % Given the current state of the system, return an action action = getAction(obj.actor,Observation); end %Action with noise method function action = getActionWithExplorationImpl(obj,Observation) % Given the current observation, select an action action = getAction(obj.actor,Observation); % Add random noise to action end %Learn method function action = learnImpl(obj,Experience) %parse experience Obs = Experience{1}; %reformat in dlarrays Obs_reformat = {dlarray(Obs{1})}; action = getAction(obj.actor,Obs_reformat); %update actor network ActorGradient = dlfeval(@actorupdate,obj,Obs_reformat); end end
问题原因与解决方法
核心问题:在
dlfeval调用的actorupdate函数中,直接访问obj.actor.Learnables不会被自动纳入计算图追踪。dlfeval只能追踪显式传入函数的变量,对象属性的可学习参数无法被自动识别为需要追踪的节点。具体修改步骤:
- 调整
actorupdate函数参数:将actor、两个critic作为显式参数传入,避免直接访问对象属性:function actorgradient = actorupdate(actor, critic_track, critic_obstacle, Observation) - 修改
learnImpl中的调用逻辑:在dlfeval中传递所需的网络和观测值:ActorGradient = dlfeval(@actorupdate, obj.actor, obj.critic_track, obj.critic_obstacle, Obs_reformat); - 更新
actorupdate内部逻辑:使用传入的网络对象计算action和Q值,最后对传入的actor可学习参数求梯度:action = evaluate(actor, Observation, UseForward=true); Qtrack = getValue(critic_track, Observation, action); Qobstacle = getValue(critic_obstacle, Obs_Obstacle, action); Qtotal = sum(Qtrack + Qobstacle); actorgradient = dlgradient(Qtotal, actor.Learnables);
- 调整
额外优化:简化
Obs_Obstacle的构建方式,避免循环赋值,直接利用dlarray的切片特性:Obs_Obstacle = {Observation{1}(1:20)};
内容的提问来源于stack exchange,提问作者Sliferslacker
相关产品推荐
相关产品推荐

