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

Taxi-v3环境Q-learning代码触发IndexError错误,请求技术帮助

Taxi-v3 Q-learning代码IndexError错误排查

我编写了一段基于OpenAI Gym Taxi-v3环境的Q-learning强化学习代码,运行时触发了IndexError错误,错误信息如下:

IndexError: only integers, slices (:), ellipsis (...), numpy.newaxis (None) and integer or boolean arrays are valid indices

以下是我的代码及完整报错栈:

原代码

import numpy as np
import gym
import random

ENV_NAME = "Taxi-v3"
env = gym.make(ENV_NAME)

print("Number of actions: %d" % env.action_space.n)
print("Number of states: %d" % env.observation_space.n)

action = env.action_space.n
state = env.observation_space.n

qtable = np.zeros((state, action),dtype=int)
print(qtable)

total_episodes = 50000
total_test_episodes = 5
max_steps = 99

learning_rate = 0.7
discount_rate = 0.9

epsilon = 1
max_epsilon = 1
min_epsilon = 0.01
decay_rate = 0.01

for episode in range(total_episodes):
    
    state = env.reset()
    step = 0
    done = False
    
    for step in range(max_steps):

        exp_exp_tradeoff = random.uniform(0,1)
        
      
        if exp_exp_tradeoff > epsilon:
            action = np.argmax(qtable[state, :])

        else:
            action = env.action_space.sample()
        
        new_state, reward, terminated, done, info = env.step(action)
        
        qtable[state, action] = qtable[state, action] + learning_rate * (reward + discount_rate * np.max(qtable[new_state, :]) - qtable[state, action])
        state = new_state
        
        if done is True:
            break
    
    episode += 1
  
    epsilon = min_epsilon + (max_epsilon - min_epsilon) * np.exp(-decay_rate * episode)

报错栈

IndexError                                Traceback (most recent call last)
~\AppData\Local\Temp/ipykernel_1092/2436469204.py in <module>
     51         #Update q value for the state based on the formula
     52         #Q(s,a) = Q(s,a) + lr[R(s,a) + gamma * max Q(s',a') - Q(s,a)]
---> 53         qtable[state, action] = qtable[state, action] + learning_rate * (reward + discount_rate * np.max(qtable[new_state, :]) - qtable[state, action])
     54         state = new_state
     55 

IndexError: only integers, slices (`:`), ellipsis (`...`), numpy.newaxis (`None`) and integer or boolean arrays are valid indices

错误原因分析

  1. Gym版本兼容问题:env.reset()返回值变化
    在Gym 0.26及以上版本中,env.reset()不再直接返回整数状态值,而是返回元组(observation, info)。原代码直接state = env.reset()会让state变成元组,用元组作为索引访问qtable数组时,就会触发IndexError。

  2. 变量名冲突
    代码开头用state = env.observation_space.n定义了状态总数,但在循环中又用state = env.reset()覆盖了该变量,虽然这不是直接报错原因,但会导致代码逻辑混淆,增加调试难度。

  3. env.step()返回值解构错误
    新版本Gym中env.step(action)的返回值是(observation, reward, terminated, truncated, info),原代码里的done实际对应truncated,后续判断结束条件的逻辑也会出现偏差。

  4. Q表数据类型不合理
    原代码将qtable的dtype设为int,但Q-learning更新公式会产生浮点数,强制用整数存储会导致精度丢失,影响算法效果。

修复方案

修复后的完整代码

import numpy as np
import gym
import random

ENV_NAME = "Taxi-v3"
env = gym.make(ENV_NAME)

print("Number of actions: %d" % env.action_space.n)
print("Number of states: %d" % env.observation_space.n)

# 重命名变量避免冲突,明确区分状态总数和当前状态
action_num = env.action_space.n
state_num = env.observation_space.n

# Q表改用浮点类型存储,适配Q值的浮点运算需求
qtable = np.zeros((state_num, action_num), dtype=np.float32)
print(qtable)

total_episodes = 50000
total_test_episodes = 5
max_steps = 99

learning_rate = 0.7
discount_rate = 0.9

epsilon = 1
max_epsilon = 1
min_epsilon = 0.01
decay_rate = 0.01

for episode in range(total_episodes):
    # 正确获取reset返回的状态值,忽略info信息
    state, _ = env.reset()
    step = 0
    
    for step in range(max_steps):
        exp_exp_tradeoff = random.uniform(0, 1)
        
        if exp_exp_tradeoff > epsilon:
            action = np.argmax(qtable[state, :])
        else:
            action = env.action_space.sample()
        
        # 按照新版本Gym的返回值顺序解构
        new_state, reward, terminated, truncated, info = env.step(action)
        
        # 更新Q值
        qtable[state, action] = qtable[state, action] + learning_rate * (
            reward + discount_rate * np.max(qtable[new_state, :]) - qtable[state, action]
        )
        state = new_state
        
        # 当任务终止或步数截断时退出循环
        if terminated or truncated:
            break
    
    # 衰减探索率epsilon
    epsilon = min_epsilon + (max_epsilon - min_epsilon) * np.exp(-decay_rate * episode)

关键修复点说明

  • 修正env.reset()赋值:用state, _ = env.reset()提取整数状态值,确保state是合法的数组索引类型。
  • 重命名变量:将表示状态总数/动作总数的变量改为state_num和action_num,避免和循环中表示当前状态/动作的变量重名。
  • 适配env.step()返回值:按照新版本Gym的返回值顺序解构,并使用terminated or truncated作为episode结束的判断条件。
  • 调整Q表数据类型:将dtype改为np.float32,保证Q值运算的精度。

内容的提问来源于stack exchange,提问作者Erdal Tutku YILMAZ

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 22:35:18