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

PyTorch数据增强后组织图像显示异常问题求助

Fixing Augmented Tissue Image Display Issues in PyTorch Nucleus Segmentation

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:

  1. All-black images when using img.astype(np.uint8):
    torchvision.transforms.ToTensor() converts your 0-255 uint8 images to 0-1 range float32 tensors. If you directly cast this 0-1 float array to uint8, every value gets truncated to 0 (since all values are <1), hence the all-black result.

  2. 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_train images 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 17:12:47