SwinUNETR / Segmentation /dataloading /dataloader2D_testing.py
deboraJ23's picture
uploaded files from https://github.com/smriti-joshi/bcnaim-odelia-challenge (except Readme, Licence and .gitignore)
361b108 verified
Raw
History Blame Contribute Delete
3.12 kB
import os
import pandas as pd
import torch
from torch.utils.data import Dataset
from PIL import Image
import albumentations as A
from albumentations.pytorch import ToTensorV2
import numpy as np
class NiftiSegmentationDataset(Dataset):
def __init__(self, csv_file, channel_keys = ['first_path'], augment=True):
"""
Args:
csv_file (str): CSV file with columns: 'image_path', 'label_path'
use_mask (bool): Whether to concatenate mask as extra channel (not typical for segmentation targets)
augment (bool): Apply data augmentation
"""
self.df = pd.read_csv(csv_file)
self.augment = augment
source_directory = '/workspace/data/mipsodeliadefault'
breast_mask_source_directory = '/workspace/Segmentation/breast_masks/mamamia/breastmaskmamamia'
self.samples = []
for _, row in self.df.iterrows():
image_paths = []
for key in channel_keys:
image_paths.append(os.path.join(source_directory, row['patient_id'], key + '.png'))
mask_path = os.path.join(source_directory, row['patient_id'], 'label.png')
breast_mask_path = os.path.join(breast_mask_source_directory, row['patient_id']+'.png')
self.samples.append({
"image_path": image_paths,
"breast_mask_path": breast_mask_path,
"mask_path": mask_path,
})
self.transform = self.build_transforms()
def build_transforms(self):
additional_targets = {'breast_mask': 'mask'}
if self.augment:
return A.Compose([
A.HorizontalFlip(p=0.5),
A.VerticalFlip(p=0.5),
A.RandomRotate90(p=0.5),
A.Rotate(limit=15, p=0.5),
A.Affine(
scale=(0.95, 1.05),
shear=5,
translate_percent=(0.05, 0.05),
p=0.5
),
# A.ColorJitter(brightness=0.1, contrast=0.1, p=0.5),
A.GaussianBlur(blur_limit=3, p=0.3),
A.GaussNoise(var_limit=(10.0, 50.0), p=0.3),
A.Normalize(mean=(0.5,), std=(0.5,)),
ToTensorV2()
], additional_targets=additional_targets)
else:
return A.Compose([
A.Normalize(mean=(0.5,), std=(0.5,)),
ToTensorV2()
], additional_targets=additional_targets)
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
sample = self.samples[idx]
# Load images as numpy arrays
image_channels = []
for image_path in sample["image_path"]:
img = np.array(Image.open(image_path).convert('L')) # Convert to grayscale
image_channels.append(img)
image = np.stack(image_channels, axis=-1) # Shape: [H, W, C]
augmented = self.transform(image=image)
image = augmented['image'] # Tensor [C, H, W]
return image