如何基于PyTorch SegNet预训练权重实现单张测试图像预测?
Hey there! Let's walk through how to get your trained SegNet model to make predictions on a single test image. I'll break this down step by step with code you can tweak to fit your setup:
Step 1: Import extra required libraries
You'll need tools to load/process images and handle data transformations. Add these imports to your existing code:
from PIL import Image import torchvision.transforms as transforms import numpy as np import matplotlib.pyplot as plt # Optional, for visualizing results
Step 2: Define image preprocessing (critical!)
This must match exactly what you used during training—same image size, normalization values, etc. If you skip this step, your predictions will be unreliable. Here's a common example, replace the values with your training setup:
# Match this to your training preprocessing pipeline! transform = transforms.Compose([ transforms.Resize((256, 256)), # Swap with your training image dimensions transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], # Replace with your dataset's mean std=[0.229, 0.224, 0.225]) # Replace with your dataset's std ])
Step 3: Load and preprocess your test image
Point the code to your test image, then apply the preprocessing we defined:
# Replace with your actual test image path img_path = "./your_test_image.png" # Load image and ensure it's in RGB format (use "L" for grayscale if needed) image = Image.open(img_path).convert("RGB") # Apply transforms and add a batch dimension (models expect batch inputs) input_tensor = transform(image).unsqueeze(0)
Step 4: Run the prediction
Switch the model to evaluation mode (turns off dropout/batch norm training behaviors) and run the forward pass without tracking gradients (saves memory and speeds things up):
# Set model to evaluation mode model.eval() # Disable gradient calculation for inference with torch.no_grad(): output = model(input_tensor)
Step 5: Convert model output to a segmentation mask
SegNet outputs class probabilities for each pixel (shape: [batch_size, num_classes, height, width]). We'll take the class with the highest probability for each pixel to get our final mask:
# Get predicted class for each pixel, remove batch dim, convert to numpy array predicted_mask = torch.argmax(output, dim=1).squeeze(0).cpu().numpy()
Optional: Visualize the results
If you want to see the original image next to your predicted mask, use this code:
plt.figure(figsize=(12, 6)) # Plot original image plt.subplot(1, 2, 1) plt.imshow(image) plt.title("Original Image") plt.axis("off") # Plot predicted mask (use "gray" colormap if it's a 2-class segmentation task) plt.subplot(1, 2, 2) plt.imshow(predicted_mask, cmap="viridis") plt.title("Predicted Segmentation Mask") plt.axis("off") plt.show()
Quick notes to avoid issues:
- If you trained your model on a GPU, move your model and input tensor to GPU with
model = model.to("cuda")andinput_tensor = input_tensor.to("cuda")before inference. Or move the model to CPU withmodel = model.to("cpu")if you're running inference on a CPU. - Double-check that your
.pthfile was saved withtorch.save(model.state_dict(), path)—your current load code works for state dicts, but if you saved the entire model, you'd usemodel = torch.load('./model_segnet_epoch50.pth')instead.
内容的提问来源于stack exchange,提问作者Jimbo

