PyTorch数据增强后组织图像显示异常问题求助
Hey, let's break down what's causing those weird display issues with your augmented tissue images and fix them step by step. The problems you're seeing (all-black images or distorted 9-grid visuals) are super common when working with PyTorch's torchvision.transforms and mismatched data formats/dimensions.
Let's Diagnose the Root Causes
First, let's map your symptoms to the likely issues:
All-black images when using
img.astype(np.uint8):torchvision.transforms.ToTensor()converts your 0-255uint8images to 0-1 rangefloat32tensors. If you directly cast this 0-1 float array touint8, every value gets truncated to 0 (since all values are <1), hence the all-black result.Color distortion + 9-grid repetition:
PyTorch stores image tensors in(C, H, W)order (channels first), but libraries like Matplotlib/OpenCV expect(H, W, C)(channels last). When you display a(3, 128, 128)tensor directly, the library misinterprets the dimensions—treating the 3 channels as separate rows/columns, leading to that distorted grid look.
Step 1: Fix Your Custom Dataset Class
First, ensure your dataset handles image/mask augmentation correctly (using PIL images, since most torchvision transforms are designed for them) and keeps augmentations synchronized between images and masks.
Here's a corrected version of your Nuc_Seg class:
import torch from torch.utils.data import Dataset from torchvision import transforms from PIL import Image import numpy as np class Nuc_Seg(Dataset): def __init__(self, images, masks, augment=False): self.images = images self.masks = masks self.augment = augment # Image transforms: ToTensor converts uint8 -> float32 (0-1 range) self.img_transform = transforms.Compose([ transforms.ToTensor(), # Optional: Add normalization if needed (remember to reverse it for display!) # transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # Mask transforms: Convert bool mask to float tensor (0=background, 1=nucleus) self.mask_transform = transforms.Compose([ transforms.ToTensor() ]) # Augmentation setup (synchronized for image and mask) if self.augment: self.augmenter = transforms.RandomAffine( degrees=15, translate=(0.1, 0.1), scale=(0.9, 1.1), shear=10 ) def __len__(self): return len(self.images) def __getitem__(self, idx): # Get raw data img = self.images[idx] # (128,128,3) uint8 mask = self.masks[idx] # (128,128,1) bool # Convert numpy arrays to PIL images (required for torchvision transforms) img_pil = Image.fromarray(img) # Mask needs to be squeezed to (128,128) and cast to uint8 for PIL compatibility mask_pil = Image.fromarray(mask.squeeze().astype(np.uint8)) # Apply synchronized augmentation if self.augment: # Use a fixed seed to ensure same transform is applied to image and mask seed = np.random.randint(2147483647) torch.manual_seed(seed) img_pil = self.augmenter(img_pil) torch.manual_seed(seed) # Fill mask with 0 (background) when augmenting mask_pil = self.augmenter(mask_pil, fillcolor=0) # Convert back to tensors img_tensor = self.img_transform(img_pil) # Re-add the channel dimension to mask (back to (1,128,128)) mask_tensor = self.mask_transform(mask_pil).unsqueeze(0) return img_tensor, mask_tensor
Key fixes here:
- Converts numpy arrays to PIL images before applying transforms (avoids dimension mismatches)
- Synchronizes augmentation randomness with a seed so image and mask get identical transforms
- Properly handles mask dimensions (squeezing for PIL, re-adding channel after tensor conversion)
Step 2: Correctly Display Augmented Images
Now, when you pull data from your DataLoader, you need to reverse the tensor formatting to get displayable numpy arrays. Here's how:
import matplotlib.pyplot as plt # Create DataLoader train_loader = torch.utils.data.DataLoader( Nuc_Seg(X_train, Y_train, augment=True), batch_size=1, shuffle=True ) # Grab a sample batch imgs, masks = next(iter(train_loader)) # Process image tensor for display def tensor_to_img(tensor): # Move tensor to CPU, transpose from (C,H,W) to (H,W,C) img_np = tensor.permute(1, 2, 0).cpu().numpy() # If you used normalization, reverse it first: # mean = np.array([0.485, 0.456, 0.406]) # std = np.array([0.229, 0.224, 0.225]) # img_np = (img_np * std + mean) # Clip values to 0-1 to avoid overflow img_np = np.clip(img_np, 0, 1) # Convert back to 0-255 uint8 return (img_np * 255).astype(np.uint8) # Process mask tensor for display def tensor_to_mask(tensor): return tensor.squeeze().cpu().numpy().astype(np.bool) # Display results plt.figure(figsize=(12, 6)) plt.subplot(1, 2, 1) plt.imshow(tensor_to_img(imgs[0])) plt.title("Augmented Tissue Image") plt.axis("off") plt.subplot(1, 2, 2) plt.imshow(tensor_to_mask(masks[0]), cmap="gray") plt.title("Augmented Nucleus Mask") plt.axis("off") plt.show()
Key display fixes:
- Uses
permute(1,2,0)to switch from PyTorch's(C,H,W)to display-friendly(H,W,C) - Converts the 0-1 float tensor back to 0-255
uint8(after reversing normalization if used) - Properly squeezes and casts the mask back to its original boolean format
Quick Checks to Verify
- If you still see issues, double-check that your raw
X_trainimages are indeed in(H,W,C)order (not(C,H,W)). - Ensure you're not accidentally applying transforms twice (e.g., once in the dataset and once in the loader).
内容的提问来源于stack exchange,提问作者TRon710

