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

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只能追踪显式传入函数的变量,对象属性的可学习参数无法被自动识别为需要追踪的节点。

  • 具体修改步骤:

    1. 调整actorupdate函数参数:将actor、两个critic作为显式参数传入,避免直接访问对象属性:
      function actorgradient = actorupdate(actor, critic_track, critic_obstacle, Observation)
      
    2. 修改learnImpl中的调用逻辑:在dlfeval中传递所需的网络和观测值:
      ActorGradient = dlfeval(@actorupdate, obj.actor, obj.critic_track, obj.critic_obstacle, Obs_reformat);
      
    3. 更新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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 11:27:39