Julia使用Crux.jl运行DQN求解MDP时出现logpdf方法匹配错误
错误原因及修复方案
核心错误点
- Q网络输出维度配置错误
你的代码中定义最后一层全连接层时写了layer3 = Dense(64, length(3)),length(3)的返回值是1,但你共有3个离散动作,需要输出3维的Q值向量,维度不匹配直接导致动作采样失败。
- Q网络输出维度配置错误
- MDP类型参数定义错误
你声明MDP类型时写了struct SliderMDP <: MDP{Array{Float32}, Array{Float32}},第二个泛型参数代表动作类型,你使用的是离散标量动作而非数组,应该改为MDP{Vector{Float32}, Float64},否则Crux采样动作时会出现类型不匹配。
- MDP类型参数定义错误
- 未暴露动作空间接口
Crux依赖POMDPs.actions标准接口获取动作集合,你仅在结构体内部定义了actions字段,没有对外实现接口,导致动作采样时返回空值Nothing,这就是你报错no method matching logpdf(::Categorical, ::Nothing)的直接原因。
- 未暴露动作空间接口
- 环境逻辑修改动作破坏训练逻辑
你在gen函数中直接修改输入的动作a为0或0.1,会导致Q网络输出的动作和实际环境执行的动作不一致,破坏DQN的时序差分更新规则,建议改为对非法动作施加惩罚,而非直接修改动作值。
- 环境逻辑修改动作破坏训练逻辑
修复后的完整代码
依赖导入(无修改)
using POMDPs using POMDPModelTools using POMDPPolicies using POMDPSimulators using Parameters using Random using Crux using Flux using Distributions
业务逻辑修复
# 修正MDP类型参数,动作类型改为Float64 @with_kw struct SliderMDP <: MDP{Vector{Float32}, Float64} x0 = Distributions.Uniform(0., 80.)# 初始位置采样分布 v0 = Distributions.Uniform(0., 25.) # 初始速度采样分布 d0 = Distributions.Uniform(0., 2.) # 初始制动力采样分布 m::Float64 = 1. tension::Float64 = 3. dmax::Float64 = 2. target::Float64 = 80. dt::Float64 = .05 γ::Float32 = 1. actions::Vector{Float64} = [-.1, 0., .1] end # 实现POMDPs.actions接口,暴露动作集合 POMDPs.actions(mdp::SliderMDP) = mdp.actions function POMDPs.gen(env::SliderMDP, s, a, rng::AbstractRNG = Random.GLOBAL_RNG) x, ẋ, d = s reward = 0f0 # 改为对越界动作施加惩罚,不直接修改动作值 if d+a >= env.dmax || d+a <= 0 reward -= 10f0 else d += a end # 超过目标也施加惩罚 if x >= env.target reward -= 20f0 end force = (d + env.tension) * -1 ẍ = force/env.m # 状态更新 x_ = x + env.dt * ẋ ẋ_ = ẋ + env.dt * ẍ d_ = d sp = Float32.([x_, ẋ_, d_]) # 主奖励:越接近目标奖励越高 reward -= abs(env.target - x) return (sp=sp, r=reward) end function POMDPs.initialstate(mdp::SliderMDP) ImplicitDistribution((rng) -> Float32.([rand(rng, mdp.x0), rand(rng, mdp.v0), rand(rng, mdp.d0)])) end # 终端条件:速度<=0即将后退 POMDPs.isterminal(mdp::SliderMDP, s) = s[2] <= 0 POMDPs.discount(mdp::SliderMDP) = mdp.γ mdp = SliderMDP(); s = state_space(mdp) # 修正Q网络输出维度,改为3(动作数量) function Q_network() layer1 = Dense(3, 64, relu) layer2 = Dense(64, 64, relu) layer3 = Dense(64, length(actions(mdp))) return DiscreteNetwork(Chain(layer1, layer2, layer3), actions(mdp)) end solver_dqn = DQN(π=Q_network(), S=s, N=30000) policy_dqn = solve(solver_dqn, mdp)
额外优化建议
- 可以调整奖励函数权重,让接近目标的奖励远高于惩罚项,引导智能体优先靠近目标
- DQN训练步数30000可以根据收敛情况调整,也可以调整
buffer_size、batch_size等超参数优化效果 - 训练完成后可以用
simulate函数跑几个episode验证策略效果
内容的提问来源于stack exchange,提问作者Sceptual
相关产品推荐
相关产品推荐

