强化学习代码在Jupyter Notebook报IndexError问题排查
问题原因与修复方案
错误根源
代码在Jupyter Notebook报错,是因为新版本Gym的API变更:env.reset()不再只返回单个状态值,而是返回包含状态和额外信息的元组(state, info)。如果直接执行state = env.reset(),state会被赋值为元组,后续用它索引q_table时,numpy会因为元组不是合法索引类型而抛出IndexError。Google Colab使用的Gym版本较低,仍沿用旧API,因此不会触发该错误。
修复方法
只需修改初始化状态的代码行,从元组中提取出真正的状态值:
将原代码中的:
state = env.reset()
改为:
state, _ = env.reset()
用下划线_接收不需要的info参数,确保state是整数类型的状态值,即可正常索引q_table。
修改后的完整代码
import numpy as np import gym # Define the environment env = gym.make("Taxi-v3").env # Initialize the q-table with zero values q_table = np.zeros([env.observation_space.n, env.action_space.n]) # Hyperparameters alpha = 0.1 # learning-rate gamma = 0.7 # discount-factor epsilon = 0.1 # explor vs exploit # Random generator rng = np.random.default_rng() # Perform 10,000 episodes for i in range(10_000): # Reset the environment and initialize total_reward state, _ = env.reset() # 此处为修改内容 done = False total_reward = 0 # Loop as long as the game is not over, i.e. done is not True while not done: if rng.random() < epsilon: action = env.action_space.sample() # Explore the action space else: action = np.argmax(q_table[state]) # Exploit learned values # Apply the action and see what happens next_state, reward, done, info = env.step(action) current_value = q_table[state, action] # current Q-value for the state/action couple next_max = np.max(q_table[next_state]) # next best Q-value # Compute the new Q-value with the Bellman equation q_table[state, action] = (1 - alpha) * current_value + alpha * (reward + gamma * next_max) # Update the current state and total_reward state = next_state total_reward += reward # Print the total reward earned in this episode print(f"Episode {i+1}: Total reward = {total_reward}")
内容的提问来源于stack exchange,提问作者Cloud23
相关产品推荐
相关产品推荐

