PyTorch中net.cuda()失效:GPU训练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

