如何将PyTorch DQN教程的图像输入改为1×9列表/数组输入
Hey there! I’ve been in your exact situation—tweaking a vision-focused DQN tutorial to work with small vector inputs like your 1×9 list is straightforward once you know which parts to modify. Let’s break this down step by step:
1. Replace Convolutional Layers with Fully Connected Layers
The standard PyTorch DQN tutorial uses CNNs to extract image features, but for your 1×9 vector, we can ditch the conv layers entirely and use a simple feed-forward network. Here’s how to redefine the DQN class:
import torch import torch.nn as nn import torch.optim as optim class DQN(nn.Module): def __init__(self, input_size, output_size): super(DQN, self).__init__() # Fully connected network tailored for 1x9 input self.fc1 = nn.Linear(input_size, 128) self.fc2 = nn.Linear(128, 64) self.fc3 = nn.Linear(64, output_size) def forward(self, x): # Input shape: (batch_size, 9) x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) return self.fc3(x) # Initialize models: input_size=9 matches your 1x9 state, output_size = number of actions policy_net = DQN(9, 4) # Example: 4 possible actions in your environment target_net = DQN(9, 4) target_net.load_state_dict(policy_net.state_dict()) target_net.eval()
2. Simplify State Preprocessing
The original tutorial does heavy image-specific processing (grayscale conversion, resizing, frame stacking). For your 1×9 list, you only need to convert it to a PyTorch tensor and adjust the shape for batch compatibility:
def preprocess_state(state): # state is your raw 1x9 list/array, e.g., [0.2, 0.7, 0.1, ..., 0.5] # Convert to float tensor and add a batch dimension return torch.tensor(state, dtype=torch.float32).unsqueeze(0) # Output shape: (1,9)
When collecting experiences in your replay buffer, pass raw 1×9 states through this function instead of the image preprocessing pipeline.
3. Adjust Training Loop for Vector Inputs
Most of the DQN training logic stays the same—experience replay, target network updates, and Q-value calculations work identically. The only key change is that your state batches will now be shaped (batch_size, 9) instead of (batch_size, channels, height, width). Here’s a modified training step snippet:
# Assume your replay buffer stores tuples of (state, action, next_state, reward, done) batch = replay_buffer.sample(batch_size) # Unpack batch and convert to tensors states = torch.cat([preprocess_state(s) for s in batch.state]) # Shape: (batch_size,9) actions = torch.tensor(batch.action, dtype=torch.long).unsqueeze(1) next_states = torch.cat([preprocess_state(s) for s in batch.next_state]) rewards = torch.tensor(batch.reward, dtype=torch.float32).unsqueeze(1) dones = torch.tensor(batch.done, dtype=torch.float32).unsqueeze(1) # Compute current Q-values current_q_values = policy_net(states).gather(1, actions) # Compute target Q-values (Bellman equation) next_q_values = target_net(next_states).max(1)[0].unsqueeze(1) target_q_values = rewards + (gamma * next_q_values * (1 - dones)) # Optimize the model loss = nn.MSELoss()(current_q_values, target_q_values) optimizer.zero_grad() loss.backward() optimizer.step()
4. Quick Checks to Avoid Common Pitfalls
- Ensure your environment returns states as a 1D list/array of length 9 (not a 2D array like
[[...]]). If it does return 2D, addtorch.flatten()inpreprocess_stateto fix the shape. - Double-check that the
input_sizein your DQN class matches the length of your state vector (9 in your case). - If using a GPU, make sure your model and tensors are on the same device—add
.to(device)to your model instances and tensors where needed.
That’s all! The core DQN logic remains intact; we just swap out image-processing components for ones that play nicely with your vector input.
内容的提问来源于stack exchange,提问作者Dae Young Kang

