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

Julia实现GridWorld的SARSA(λ)代码出现KeyError问题求助

解决Julia中SARSA(λ)算法的KeyError问题

问题描述

在Julia语言中基于GridWorld环境实现带资格迹的SARSA(λ)强化学习算法时,运行代码触发KeyError,提示特定状态键未找到。调整字典结构后问题仍未解决,需排查代码错误并说明原因。

原代码

using ReinforcementLearningBase, GridWorlds
using PyPlot

world = GridWorlds.GridRoomsDirectedModule.GridRoomsDirected();
env = GridWorlds.RLBaseEnv(world)

mutable struct Agent
    env::AbstractEnv
    algo::Symbol
    ϵ::Float64 # exploration coefficient
    ϵ_decay::Float64
    ϵ_min::Float64
    λ::Float64 # parametr lambda
    β::Float64 # discount factor
    α::Float64 # learning rate
    Q::Dict
    score::Int # number of times the agent reached the goal
    steps_per_episode::Vector{Float64} # average number of steps per episode
    E::Dict 
end

function Agent(env, algo; ϵ = 1.0, ϵ_decay = 0.9975, ϵ_min = 0.005, λ=0.9,
        β = 0.99, α = 0.1) 
    if algo != :SARSA && algo != :Qlearning
        @error "unknown algorithm"
    end
    Agent(env, algo,
        ϵ, ϵ_decay, ϵ_min,λ, β, α, 
        Dict(), 0, [0.0,],Dict())
end

function learn!(agent, S, A, r, S′,A′)
    if !haskey(agent.Q, S)
        agent.E[S] = zeros(length(action_space(agent.env)))
        agent.Q[S] = zeros(length(action_space(agent.env)))
        agent.Q[S][A] = r
        agent.E[S][A]=1
    else
        Q_S′ = 0.0
        haskey(agent.Q, S′) && (Q_S′ += agent.Q[S′][A′])
        Δ = r + agent.β * agent.Q[S′][A′] - agent.Q[S][A]
        agent.E[S][A]=agent.β*agent.λ*agent.E[S][A]+1
        agent.Q[S][A] += agent.α * Δ*agent.E[S][A]

    end
end

function run_learning!(agent, steps; burning = true, animated = nothing) 
    step = 1.0
    steps_per_episode = 1.0
    episode = 1.0
    if !isnothing(animated)
        global str = ""
        global str = str * "FRAME_START_DELIMITER"
        global str = str * "step: $(step)\n"
        global str = str * "episode: $(episode)\n"
        global str = str * repr(MIME"text/plain"(), env)
        global str = str * "\ntotal_reward: 0"
    end
    while step <= steps
        if (burning && step < 0.1*steps) || rand() < agent.ϵ || !haskey(agent.Q, state(agent.env))
            A = rand(1:length(action_space(agent.env)))
        else 
            A = argmax(agent.Q[state(agent.env)])
        end
        S = deepcopy(state(agent.env))
        agent.env(action_space(agent.env)[A])
        r = reward(agent.env)
        S′ = deepcopy(state(agent.env))
        if agent.algo == :SARSA
            if (burning && step < 0.1 * steps) || rand() < agent.ϵ || !haskey(agent.Q, state(agent.env))
                A′ = rand(1:length(action_space(agent.env)))
            else 
                A′ = argmax(agent.Q[state(agent.env)])
            end
            learn!(agent, S, A, r, S′,A′)
        else
            learn!(agent, S, A, r, S′)
        end
        if !isnothing(animated) 
            global str = str * "FRAME_START_DELIMITER"
            global str = str * "step: $(step)\n"
            global str = str * "episode: $(episode)\n"
            global str = str * repr(MIME"text/plain"(), env)
            global str = str * "\ntotal_reward: $(agent.score)"
        end
        if is_terminated(agent.env)
            eps = agent.ϵ * agent.ϵ_decay
            agent.ϵ = max(agent.ϵ_min, eps)
            agent.score += 1.0
            push!(agent.steps_per_episode, 
                agent.steps_per_episode[end] + (steps_per_episode - agent.steps_per_episode[end])/episode)
            episode += 1.0
            steps_per_episode = 0
            reset!(agent.env)
        end
        step += 1.0 
        steps_per_episode += 1.0
    end
    if !isnothing(animated) 
        write(animated * ".txt", str)
    end
end

agent_SARSA = Agent(env,:SARSA);

run_learning!(agent_SARSA, 2500)
@info "agent score: $(agent_SARSA.score)"

原报错信息

KeyError: key ([0 0 … 0 0; 1 1 … 1 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 1 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 1 … 1 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 0 … 0 1; 0 0 … 0 0;;; 0 0 … 0 0; 1 1 … 1 1; 0 0 … 0 0], 1) not found

Stacktrace:
 [1] getindex(h::Dict{Any, Any}, key::Tuple{BitArray{3}, Int64})
   @ Base .\\dict.jl:498
 [2] learn!(agent::Agent, S::Tuple{BitArray{3}, Int64}, A::Int64, r::Float32, S′::Tuple{BitArray{3}, Int64}, A′::Int64)
   @ Main .\\In[44]:10
 [3] run_learning!(agent::Agent, steps::Int64; burning::Bool, animated::Nothing)
   @ Main .\\In[45]:31
 [4] run_learning!(agent::Agent, steps::Int64)
   @ Main .\\In[45]:1
 [5] top-level scope
   @ In[51]:1
 [6] eval
   @ .\\boot.jl:368 [inlined]
 [7] include_string(mapexpr::typeof(REPL.softscope), mod::Module, code::String, filename::String)
   @ Base .\\loading.jl:1428

错误原因说明

  1. 未初始化后续状态的Q/E表:原代码在learn!的else分支中直接访问agent.Q[S′][A′],但S′可能从未被探索过,不存在于Q字典中,导致KeyError。
  2. 资格迹更新逻辑错误:SARSA(λ)要求对所有状态动作对的资格迹进行衰减更新,原代码仅更新当前状态动作对的E值,不符合算法的资格迹传播机制。
  3. 终止状态处理缺失:当环境进入终止状态时,未重置资格迹,也未将终止状态的目标值设为0(终止状态无后续动作),导致错误的时序差分计算。

修复后的代码

using ReinforcementLearningBase, GridWorlds
using PyPlot

world = GridWorlds.GridRoomsDirectedModule.GridRoomsDirected();
env = GridWorlds.RLBaseEnv(world)

mutable struct Agent
    env::AbstractEnv
    algo::Symbol
    ϵ::Float64 # exploration coefficient
    ϵ_decay::Float64
    ϵ_min::Float64
    λ::Float64 # parametr lambda
    β::Float64 # discount factor
    α::Float64 # learning rate
    Q::Dict
    score::Int # number of times the agent reached the goal
    steps_per_episode::Vector{Float64} # average number of steps per episode
    E::Dict 
end

function Agent(env, algo; ϵ = 1.0, ϵ_decay = 0.9975, ϵ_min = 0.005, λ=0.9,
        β = 0.99, α = 0.1) 
    if algo != :SARSA && algo != :Qlearning
        @error "unknown algorithm"
    end
    Agent(env, algo,
        ϵ, ϵ_decay, ϵ_min,λ, β, α, 
        Dict(), 0, [0.0,],Dict())
end

function learn!(agent, S, A, r, S′,A′, done)
    # 初始化当前状态的Q和E表(如果不存在)
    if !haskey(agent.Q, S)
        agent.Q[S] = zeros(length(action_space(agent.env)))
        agent.E[S] = zeros(length(action_space(agent.env)))
    end
    # 初始化后续状态的Q和E表(如果不存在)
    if !haskey(agent.Q, S′)
        agent.Q[S′] = zeros(length(action_space(agent.env)))
        agent.E[S′] = zeros(length(action_space(agent.env)))
    end

    # 计算时序差分误差,终止状态目标值为0
    target = done ? r : r + agent.β * agent.Q[S′][A′]
    Δ = target - agent.Q[S][A]

    # 更新当前状态动作对的资格迹
    agent.E[S][A] += 1.0

    # 遍历所有状态动作对,更新Q值并衰减资格迹
    for (state, _) in agent.Q
        agent.Q[state] .+= agent.α * Δ .* agent.E[state]
        agent.E[state] .*= agent.β * agent.λ
    end

    # 终止状态下重置资格迹
    done && (agent.E = Dict())
end

function run_learning!(agent, steps; burning = true, animated = nothing) 
    step = 1.0
    steps_per_episode = 1.0
    episode = 1.0
    if !isnothing(animated)
        global str = ""
        global str = str * "FRAME_START_DELIMITER"
        global str = str * "step: $(step)\n"
        global str = str * "episode: $(episode)\n"
        global str = str * repr(MIME"text/plain"(), env)
        global str = str * "\ntotal_reward: 0"
    end
    while step <= steps
        current_state = state(agent.env)
        # 选择当前动作
        if (burning && step < 0.1*steps) || rand() < agent.ϵ || !haskey(agent.Q, current_state)
            A = rand(1:length(action_space(agent.env)))
        else 
            A = argmax(agent.Q[current_state])
        end
        S = deepcopy(current_state)
        agent.env(action_space(agent.env)[A])
        r = reward(agent.env)
        S′ = deepcopy(state(agent.env))
        done = is_terminated(agent.env)
        
        if agent.algo == :SARSA
            # 选择下一动作,终止状态下随机选动作不影响结果
            if done || (burning && step < 0.1 * steps) || rand() < agent.ϵ || !haskey(agent.Q, S′)
                A′ = rand(1:length(action_space(agent.env)))
            else 
                A′ = argmax(agent.Q[S′])
            end
            learn!(agent, S, A, r, S′, A′, done)
        else
            # Q-learning逻辑(此处保留结构,需完善可自行补充)
            continue
        end

        if !isnothing(animated) 
            global str = str * "FRAME_START_DELIMITER"
            global str = str * "step: $(step)\n"
            global str = str * "episode: $(episode)\n"
            global str = str * repr(MIME"text/plain"(), env)
            global str = str * "\ntotal_reward: $(agent.score)"
        end

        if done
            eps = agent.ϵ * agent.ϵ_decay
            agent.ϵ = max(agent.ϵ_min, eps)
            agent.score += 1
            push!(agent.steps_per_episode, 
                agent.steps_per_episode[end] + (steps_per_episode - agent.steps_per_episode[end])/episode)
            episode += 1.0
            steps_per_episode = 0
            reset!(agent.env)
        end
        step += 1.0 
        steps_per_episode += 1.0
    end
    if !isnothing(animated) 
        write(animated * ".txt", str)
    end
end

agent_SARSA = Agent(env,:SARSA);

run_learning!(agent_SARSA, 2500)
@info "agent score: $(agent_SARSA.score)"

修复要点

  • 提前初始化当前状态和后续状态的Q/E表,确保访问时不会触发KeyError。
  • 遍历所有状态动作对更新Q值和衰减资格迹,符合SARSA(λ)的算法要求。
  • 增加终止状态标记done,在终止时重置资格迹,并将目标值设为即时奖励(无后续动作)。
  • 选择下一动作时,先判断环境是否终止,避免在终止状态无效访问Q表。

内容的提问来源于stack exchange,提问作者przel123

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 13:45:28