如何利用Flux.jl梯度更新PPO的Actor网络参数?梯度返回nothing
Julia PPO实现中Flux梯度为nothing的问题及解决办法
问题描述
作为Julia新手首次实现PPO算法,使用Flux.jl更新Actor网络参数时,发现梯度字典有键但所有值都是nothing。核心代码片段如下:
batch_obs, batch_acts, _, batch_rtgo, _ = rollout(ppo_network) V, curr_log_probs = evaluate(ppo_network, batch_obs, batch_acts) V_scaled = V*maximum(abs.(batch_rtgo)) / maximum(abs.(V)) ratios = exp.(curr_log_probs - transpose(hcat(old_batch_log_probs...))) ratios = ratios .+ 1e-8 A_k = batch_rtgo - deepcopy(V_scaled) A_k = (A_k .- mean(A_k)) ./ (std(A_k) .+ 1e-10) # surrogate objectives surr1 = ratios .* A_k surr2 = clamp.(ratios, 1-clip, 1+clip) .* A_k actor_loss = -mean(min.(surr1, surr2)) actor_opt = Adam(lr) actor_gs = gradient(() -> actor_loss, params(ppo.actor.model, ppo.actor.mean, ppo.actor.logstd)) # update parameters update!(actor_opt, params([ppo.actor.model, ppo.actor.mean, ppo.actor.logstd]), actor_gs)
尝试将rollout和损失计算全部放进gradient闭包后,能得到非零梯度,但每次求梯度都要重新运行rollout,计算成本极高,大批次时还会出现大型命名元组的问题:
# define actor parameters actor_ps = params(ppo.actor.model, ppo.actor.mean, ppo.actor.logstd) # define gradient function actor_gs = gradient(actor_ps) do batch_obs, batch_acts, _, batch_rtgo, _ = rollout(ppo_network) V, curr_log_probs = evaluate(ppo_network, batch_obs, batch_acts) V_scaled = V*maximum(abs.(batch_rtgo)) / maximum(abs.(V)) ratios = exp.(curr_log_probs - transpose(hcat(old_batch_log_probs...))) ratios = ratios .+ 1e-8 A_k = batch_rtgo - deepcopy(V_scaled) A_k = (A_k .- mean(A_k)) ./ (std(A_k) .+ 1e-10) surr1 = ratios .* A_k surr2 = clamp.(ratios, 1-clip, 1+clip) .* A_k actor_loss = -mean(min.(surr1, surr2)) return actor_loss end
疑问:资料说损失函数需要传入参数,否则梯度不知道对哪些参数求导,因此返回nothing,这个说法是否正确?该如何正确定义梯度完成参数更新?
问题分析
你查到的结论是对的:梯度为nothing的核心原因是,第一个代码里的actor_loss是预先计算好的标量,它和你要更新的Actor参数之间没有计算图关联。Flux的gradient函数需要追踪从参数到损失值的计算路径,而你提前把损失算出来,闭包里只返回这个固定标量,梯度自然无法追踪到参数,所以返回nothing。
你第二个代码能得到梯度,是因为把依赖参数的evaluate步骤放进了闭包,让Flux能追踪到参数→curr_log_probs→损失的计算路径,但错误地把不需要求导的rollout也放进了闭包,导致每次求梯度都重复收集数据,完全没必要。
正确解决方案
核心思路:把不需要对参数求导的步骤(rollout收集数据、优势函数A_k的计算)移到gradient闭包外面,只把依赖参数的计算(evaluate得到当前log概率、计算损失)留在闭包内。这样既保留了参数到损失的计算图,又避免重复rollout,大幅降低计算成本。
修改后的代码如下:
# 1. 提前完成rollout和与参数无关的预处理(只做一次) batch_obs, batch_acts, _, batch_rtgo, _ = rollout(ppo_network) # 提前计算Critic的V值(Actor梯度计算不需要依赖Critic参数,可提前算出) V, _ = evaluate(ppo_network, batch_obs, batch_acts) V_scaled = V*maximum(abs.(batch_rtgo)) / maximum(abs.(V)) A_k = batch_rtgo - V_scaled # 不需要deepcopy,V_scaled是独立于参数的标量运算结果 A_k = (A_k .- mean(A_k)) ./ (std(A_k) .+ 1e-10) # 2. 定义Actor参数与优化器 actor_ps = params(ppo.actor.model, ppo.actor.mean, ppo.actor.logstd) actor_opt = Adam(lr) # 3. 梯度计算:仅包含依赖Actor参数的步骤 actor_gs = gradient(actor_ps) do # 仅重新计算Actor输出的当前log概率,无需重复计算Critic的V _, curr_log_probs = evaluate(ppo_network, batch_obs, batch_acts) ratios = exp.(curr_log_probs - transpose(hcat(old_batch_log_probs...))) ratios = ratios .+ 1e-8 surr1 = ratios .* A_k surr2 = clamp.(ratios, 1-clip, 1+clip) .* A_k actor_loss = -mean(min.(surr1, surr2)) return actor_loss end # 4. 更新Actor参数 update!(actor_opt, actor_ps, actor_gs)
关键细节说明
rollout和优势函数A_k的计算完全不依赖Actor参数,仅需在梯度计算前执行一次,无需放进闭包。- 如果
evaluate函数可以拆分,建议单独编写get_actor_log_probs函数,避免冗余计算Critic的V值,进一步降低开销。 - 无需对
V_scaled做deepcopy,它是独立于参数的标量运算结果,不会干扰梯度追踪。
内容的提问来源于stack exchange,提问作者Max Kim
相关产品推荐
相关产品推荐

