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
错误原因分析
Gym版本兼容问题:
env.reset()返回值变化
在Gym 0.26及以上版本中,env.reset()不再直接返回整数状态值,而是返回元组(observation, info)。原代码直接state = env.reset()会让state变成元组,用元组作为索引访问qtable数组时,就会触发IndexError。变量名冲突
代码开头用state = env.observation_space.n定义了状态总数,但在循环中又用state = env.reset()覆盖了该变量,虽然这不是直接报错原因,但会导致代码逻辑混淆,增加调试难度。env.step()返回值解构错误
新版本Gym中env.step(action)的返回值是(observation, reward, terminated, truncated, info),原代码里的done实际对应truncated,后续判断结束条件的逻辑也会出现偏差。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

