You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将PyTorch DQN教程的图像输入改为1×9列表/数组输入

Adapting PyTorch DQN from Image Input to 1×9 Vector Input

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, add torch.flatten() in preprocess_state to fix the shape.
  • Double-check that the input_size in 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 07:15:13