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

