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

Julia使用Crux.jl运行DQN求解MDP时出现logpdf方法匹配错误

错误原因及修复方案

核心错误点

    1. Q网络输出维度配置错误
      你的代码中定义最后一层全连接层时写了layer3 = Dense(64, length(3)),length(3)的返回值是1,但你共有3个离散动作,需要输出3维的Q值向量,维度不匹配直接导致动作采样失败。
    1. MDP类型参数定义错误
      你声明MDP类型时写了struct SliderMDP <: MDP{Array{Float32}, Array{Float32}},第二个泛型参数代表动作类型,你使用的是离散标量动作而非数组,应该改为MDP{Vector{Float32}, Float64},否则Crux采样动作时会出现类型不匹配。
    1. 未暴露动作空间接口
      Crux依赖POMDPs.actions标准接口获取动作集合,你仅在结构体内部定义了actions字段,没有对外实现接口,导致动作采样时返回空值Nothing,这就是你报错no method matching logpdf(::Categorical, ::Nothing)的直接原因。
    1. 环境逻辑修改动作破坏训练逻辑
      你在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 15:06:08