使用DQN算法运行FrozenLake-v0环境时遇维度不匹配错误求助
Let's break down exactly what's going wrong and how to fix it. Your core issue is a mismatch between the state format FrozenLake outputs and what your DQN model expects:
The Root Problem
FrozenLake-v0 gives you a single integer state (0-15) representing which grid cell you're in, but your model is defined to accept a 16-dimensional input (since input_dim = env.observation_space.n = 16). You need to convert those integer states into one-hot encoded vectors—a 16-element array where only the index matching the state is 1, and all others are 0 (e.g., state 3 becomes [0,0,0,1,0,...0]).
Step-by-Step Fixes
1. Add a Helper Function for One-Hot Encoding
First, write a simple function to convert integer states to the 16-dimensional vectors your model needs:
def state_to_one_hot(state, num_states): one_hot = np.zeros(num_states) one_hot[state] = 1.0 return one_hot
2. Tweak the Model Definition (Optional but Clearer)
While input_dim=16 works, using input_shape=(16,) is more idiomatic for Keras and makes the expected input shape explicit:
model = tf.keras.Sequential() model.add(tf.keras.layers.Dense(64, input_shape=(input_dim,) , activation='relu')) model.add(tf.keras.layers.Dense(32, activation='relu')) model.add(tf.keras.layers.Dense(n_actions, activation='linear')) model.compile(optimizer=tf.keras.optimizers.Adam(), loss='mse')
3. Update State Handling in the Episode Loop
You need to convert states to one-hot before feeding them to the model, and store the encoded states in your replay memory:
for n in range(n_episodes): s = env.reset() # Convert initial state to one-hot s_one_hot = state_to_one_hot(s, input_dim) done=False r_sum = 0 while not done: # Model expects batch inputs, so reshape to (1, 16) qvals_s = model.predict(s_one_hot.reshape(1, -1), verbose=0) if np.random.random() < epsilon: a = env.action_space.sample() else: a = np.argmax(qvals_s[0]) # Grab the max index from the batch output sprime, r, done, info = env.step(a) r_sum += r # Convert next state to one-hot for replay memory sprime_one_hot = state_to_one_hot(sprime, input_dim) # Manage replay memory size if len(replay_memory) > mem_max_size: replay_memory.pop(0) # Store encoded states, not raw integers replay_memory.append({"s": s_one_hot, "a": a, "r": r, "sprime": sprime_one_hot, "done": done}) # Update current state to the encoded next state s_one_hot = sprime_one_hot
4. Safeguard the Replay Function
Add a check to avoid sampling from a replay memory that's too small, and disable verbose logging during predictions to keep output clean:
def replay(replay_memory, minibatch_size=32): # Don't train if we don't have enough samples yet if len(replay_memory) < minibatch_size: return model minibatch = np.random.choice(replay_memory, minibatch_size, replace=True) s_l = np.array(list(map(lambda x: x['s'], minibatch))) a_l = np.array(list(map(lambda x: x['a'], minibatch))) r_l = np.array(list(map(lambda x: x['r'], minibatch))) sprime_l = np.array(list(map(lambda x: x['sprime'], minibatch))) done_l = np.array(list(map(lambda x: x['done'], minibatch))) qvals_sprime_l = model.predict(sprime_l, verbose=0) target_f = model.predict(s_l, verbose=0) for i,(s,a,r,qvals_sprime, done) in enumerate(zip(s_l,a_l,r_l,qvals_sprime_l, done_l)): if not done: target = r + gamma * np.max(qvals_sprime) else: target = r target_f[i][a] = target model.fit(s_l, target_f, epochs=1, verbose=0) return model
5. Adjust Epsilon Decay (Optional)
Your current decay will only bring epsilon down to 0.4 after 500 episodes. If you want to reach 0.1, adjust the decay rate:
if epsilon > 0.1: epsilon = max(0.1, epsilon - 0.0016) # Hits 0.1 after 500 episodes
Full Fixed Code
Putting it all together, here's the complete working version:
import gym import numpy as np import tensorflow as tf def state_to_one_hot(state, num_states): one_hot = np.zeros(num_states) one_hot[state] = 1.0 return one_hot def replay(replay_memory, model, minibatch_size=32, gamma=0.99): if len(replay_memory) < minibatch_size: return model minibatch = np.random.choice(replay_memory, minibatch_size, replace=True) s_l = np.array(list(map(lambda x: x['s'], minibatch))) a_l = np.array(list(map(lambda x: x['a'], minibatch))) r_l = np.array(list(map(lambda x: x['r'], minibatch))) sprime_l = np.array(list(map(lambda x: x['sprime'], minibatch))) done_l = np.array(list(map(lambda x: x['done'], minibatch))) qvals_sprime_l = model.predict(sprime_l, verbose=0) target_f = model.predict(s_l, verbose=0) for i,(s,a,r,qvals_sprime, done) in enumerate(zip(s_l,a_l,r_l,qvals_sprime_l, done_l)): if not done: target = r + gamma * np.max(qvals_sprime) else: target = r target_f[i][a] = target model.fit(s_l, target_f, epochs=1, verbose=0) return model # Initialize environment and model env = gym.make("FrozenLake-v0") n_actions = env.action_space.n input_dim = env.observation_space.n model = tf.keras.Sequential() model.add(tf.keras.layers.Dense(64, input_shape=(input_dim,) , activation='relu')) model.add(tf.keras.layers.Dense(32, activation='relu')) model.add(tf.keras.layers.Dense(n_actions, activation='linear')) model.compile(optimizer=tf.keras.optimizers.Adam(), loss='mse') # Training parameters n_episodes = 500 gamma = 0.99 epsilon = 0.9 minibatch_size = 32 r_sums = [] replay_memory = [] mem_max_size = 100000 # Training loop for n in range(n_episodes): s = env.reset() s_one_hot = state_to_one_hot(s, input_dim) done=False r_sum = 0 while not done: qvals_s = model.predict(s_one_hot.reshape(1, -1), verbose=0) if np.random.random() < epsilon: a = env.action_space.sample() else: a = np.argmax(qvals_s[0]) sprime, r, done, info = env.step(a) r_sum += r sprime_one_hot = state_to_one_hot(sprime, input_dim) if len(replay_memory) > mem_max_size: replay_memory.pop(0) replay_memory.append({"s": s_one_hot, "a": a, "r": r, "sprime": sprime_one_hot, "done": done}) s_one_hot = sprime_one_hot model = replay(replay_memory, model, minibatch_size, gamma) if epsilon > 0.1: epsilon = max(0.1, epsilon - 0.0016) r_sums.append(r_sum) if n % 100 == 0: print(f"Episode {n}, Average Reward: {np.mean(r_sums[-100:])}")
This should resolve the shape mismatch errors and let your DQN train properly on FrozenLake-v0. Even though DQN is overkill for this environment, it's a great way to practice the implementation!
内容的提问来源于stack exchange,提问作者mikanim

