使用OpenAI Gym进行Q-Learning时遇IndexError问题求助
错误信息
IndexError Traceback (most recent call last)
~\AppData\Local\Temp\ipykernel_10800\268253893.py in
15 next_state, reward, done,trauncated,info = env.step(action)
16 #if state == int:
---> 17 q[state,action] = q[state,action] + LEARNING_RATE*(reward + GAMMA*np.max(q[next_state,:]) - q[state,action])
18 state = next_state
19IndexError: only integers, slices (
:), ellipsis (...), numpy.newaxis (None) and integer or boolean arrays are valid indices
问题背景
使用OpenAI Gym结合Q-Learning算法时触发上述错误,已知初始state是元组类型,但直接用整数替代元组会导致Q-Learning无法正常工作。代码参考自FreeCodeCamp的TensorFlow2.0教程,教程中运行正常,但本地运行报错。
复现代码
rewards = [] for episode in range(EPISODES): state = env.reset() for _ in range(MAX_STEPS): if RENDER: env.render() if np.random.uniform(0,1) < epsilon: action = env.action_space.sample() else: #if state == int: action = np.argmax(q[state,:]) next_state, reward, done,trauncated,info = env.step(action) #if state == int: q[state,action] = q[state,action] + LEARNING_RATE*(reward + GAMMA*np.max(q[next_state,:]) - q[state,action]) state = next_state if done: rewards.append(reward) epsilon -= 0.001 break print(q) print("Score over time: " + str(sum(rewards)/EPISODES))
问题核心是numpy数组仅支持整数、切片等类型作为索引,而你的state是元组,直接用元组索引numpy数组就会触发该错误。教程运行正常,大概率是因为教程使用的Gym环境返回的state是整数类型,而你当前使用的环境返回的是元组(比如CartPole这类多维状态空间的环境)。
针对元组类型的state,有两种可行处理方式:
1. 将元组状态离散化,转换为整数索引
如果state是连续值组成的元组(比如CartPole的状态为(位置, 速度, 角度, 角速度)),可将每个维度的连续值划分成若干区间,把元组映射为唯一的整数索引:
import numpy as np # 假设state是4维元组,每个维度划分10个区间 DISCRETE_BINS = 10 state_bounds = env.observation_space.high - env.observation_space.low state_bins = [np.linspace(env.observation_space.low[i], env.observation_space.high[i], DISCRETE_BINS) for i in range(4)] def discretize_state(state): indices = [] for i in range(len(state)): # 将维度值映射到区间索引,digitize返回1起始索引,转成0起始 idx = np.digitize(state[i], state_bins[i]) - 1 # 防止索引越界 idx = max(0, min(idx, DISCRETE_BINS - 1)) indices.append(idx) # 将多维索引转为一维整数(如4维10区间,转成0~9999的整数) return np.ravel_multi_index(indices, [DISCRETE_BINS]*4)
修改原代码中的state使用逻辑:
# 初始化Q表,维度对应离散化后的状态数和动作数 Q_TABLE_SIZE = DISCRETE_BINS ** len(env.observation_space.high) q = np.zeros((Q_TABLE_SIZE, env.action_space.n)) rewards = [] for episode in range(EPISODES): raw_state = env.reset() state = discretize_state(raw_state) # 转成整数索引 for _ in range(MAX_STEPS): if RENDER: env.render() if np.random.uniform(0,1) < epsilon: action = env.action_space.sample() else: action = np.argmax(q[state,:]) raw_next_state, reward, done, truncated, info = env.step(action) next_state = discretize_state(raw_next_state) q[state,action] = q[state,action] + LEARNING_RATE*(reward + GAMMA*np.max(q[next_state,:]) - q[state,action]) state = next_state if done: rewards.append(reward) epsilon -= 0.001 break
2. 使用字典存储Q表,直接用元组作为键
如果状态空间不大,可直接用字典存储Q值,元组可直接作为字典的键:
# 初始化空字典,按需生成对应状态的Q值 q = {} rewards = [] for episode in range(EPISODES): state = env.reset() # 确保初始状态在字典中存在 if state not in q: q[state] = np.zeros(env.action_space.n) for _ in range(MAX_STEPS): if RENDER: env.render() if np.random.uniform(0,1) < epsilon: action = env.action_space.sample() else: action = np.argmax(q[state]) next_state, reward, done, truncated, info = env.step(action) # 确保next_state在字典中存在 if next_state not in q: q[next_state] = np.zeros(env.action_space.n) q[state][action] = q[state][action] + LEARNING_RATE*(reward + GAMMA*np.max(q[next_state]) - q[state][action]) state = next_state if done: rewards.append(reward) epsilon -= 0.001 break
这种方式无需离散化,适合格子世界等状态空间有限的环境;但如果是连续状态空间,字典会变得异常庞大、效率低下,此时推荐使用第一种离散化方案。
内容的提问来源于stack exchange,提问作者Arjun Prakash

