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

PyTorch中net.cuda()失效:GPU训练CNN数字识别报错求助

Fixing PyTorch GPU Training Error in Digit Recognition CNN

Hey Daniel, this error almost always pops up when your model and input tensors aren’t on the same device (one’s running on CPU, the other on GPU). Let’s walk through the fixes step by step:

Step 1: Confirm CUDA Availability

First, make sure your system can access the GPU. Add this snippet at the start of your script to check:

import torch
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {device}")

If it prints cuda, your GPU is ready to go. If not, double-check your PyTorch installation or GPU driver setup.

Step 2: Move Your CNN Model to GPU

After initializing your model, send it to the GPU using .to(device):

cnn = Net()  # Your CNN class instance
cnn = cnn.to(device)

Do this before starting training—if you load pre-trained weights later, make sure those get moved to the GPU too.

Step 3: Send Input Tensors to the Same Device

Every batch of inputs you pass to the model needs to match the model’s device. In your training loop, add this line right before feeding inputs to the model:

# Assuming inputs/labels come from your DataLoader
inputs, labels = inputs.to(device), labels.to(device)
outputs = cnn(inputs)  # This should run without errors now!

If you want to automate this, you can add a custom collate function to your DataLoader that moves tensors to the device automatically.

Step 4: Fix Hidden Device Mismatches

Sometimes custom layers or manual tensor creation in your forward pass can accidentally make tensors on CPU. For example:

# Bad: Creates tensor on CPU by default
new_tensor = torch.tensor([1.0, 2.0])
# Good: Uses the same device as your model
new_tensor = torch.tensor([1.0, 2.0], device=device)

Or, you can move existing tensors to the device after creation: new_tensor = new_tensor.to(device)

Quick Troubleshooting Check

If you’re still stuck, verify that your model and inputs are on the same device with these prints:

print(f"Model device: {next(cnn.parameters()).device}")
print(f"Input device: {inputs.device}")

Both should show cuda:0 (or similar) if everything’s set up correctly.

Let me know if you hit any other snags after trying these steps!

内容的提问来源于stack exchange,提问作者Daniel Chen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:06:46