| from collections import defaultdict | |
| import os | |
| import random | |
| 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_odelia= '/workspace/data/mipsodeliadefault' | |
| source_directory_mamamia= '/workspace/data/mips3' | |
| breast_mask_source_directory_odelia = '/workspace/Segmentation/breastmaskodelia' | |
| breast_mask_source_directory_mamamia = '/workspace/Segmentation/breastmaskmamamia' | |
| self.samples = [] | |
| for _, row in self.df.iterrows(): | |
| image_paths = [] | |
| source_directory = source_directory_mamamia if row['source'] == 'public' else source_directory_odelia | |
| breast_mask_source_directory = breast_mask_source_directory_mamamia if row['source'] == 'public' else breast_mask_source_directory_odelia | |
| 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({ | |
| "patient_id": row['patient_id'], | |
| "image_path": image_paths, | |
| "mask_path": mask_path, | |
| "breast_mask_path": breast_mask_path, | |
| "label": int(row['label']) #if not row['source'] == 'public' else None # Assuming label is an integer | |
| }) | |
| self.label_to_indices = defaultdict(list) | |
| for idx, sample in enumerate(self.samples): | |
| self.label_to_indices[sample['label']].append(idx) | |
| self.transform = self.build_transforms() | |
| def build_transforms(self): | |
| additional_targets = {'breast_mask': 'mask'} | |
| if self.augment: | |
| return A.Compose([ | |
| # -------------------- | |
| # Geometric transforms | |
| # -------------------- | |
| A.HorizontalFlip(p=0.5), | |
| A.VerticalFlip(p=0.5), | |
| A.RandomRotate90(p=0.5), | |
| # Larger random rotation | |
| A.Rotate(limit=30, border_mode=0, p=0.8), | |
| # Strong affine transformations | |
| A.Affine( | |
| scale=(0.85, 1.15), | |
| shear=(-15, 15), | |
| translate_percent=(0.15, 0.15), | |
| rotate=(-20, 20), | |
| p=0.8 | |
| ), | |
| # Heavy warping | |
| A.ElasticTransform(alpha=80, sigma=10, p=0.4), | |
| A.GridDistortion(num_steps=5, distort_limit=0.4, p=0.4), | |
| A.OpticalDistortion(distort_limit=0.3, p=0.3), | |
| # Random crops/resizes | |
| A.RandomResizedCrop(size=(256, 256), scale=(0.8, 1.0), p=0.5), | |
| # -------------------- | |
| # Intensity transforms | |
| # -------------------- | |
| A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3, p=0.5), | |
| A.MultiplicativeNoise(multiplier=(0.8, 1.2), per_channel=True, p=0.4), | |
| # Simulate different scanner properties | |
| A.GaussianBlur(blur_limit=(3, 7), p=0.4), | |
| A.MotionBlur(blur_limit=5, p=0.3), | |
| A.GaussNoise(std_range=(0.02, 0.08), p=0.5), | |
| # Simulate local signal loss (cutout) | |
| A.CoarseDropout( | |
| num_holes_range=(4, 12), | |
| hole_height_range=(16, 48), | |
| hole_width_range=(16, 48), | |
| fill=0, | |
| p=0.5 | |
| ), | |
| # Normalize | |
| A.Normalize( | |
| mean=(0.2074, 0.1290, 0.1396, 0.1470), | |
| std=(0.2110, 0.1629, 0.1620, 0.1626) | |
| ), | |
| ToTensorV2() | |
| ], additional_targets=additional_targets) | |
| # return A.Compose([ | |
| # A.HorizontalFlip(p=0.5), | |
| # A.VerticalFlip(p=0.5), | |
| # # Stronger rotations | |
| # A.RandomRotate90(p=0.5), | |
| # A.Rotate(limit=25, p=0.7), | |
| # # Affine: stronger scale, translation | |
| # A.Affine( | |
| # scale=(0.9, 1.1), | |
| # shear=(-10, 10), | |
| # translate_percent=(0.1, 0.1), | |
| # p=0.7 | |
| # ), | |
| # # Elastic deformation & grid distortion | |
| # A.ElasticTransform(alpha=50, sigma=8, alpha_affine=8, p=0.3), | |
| # A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3), | |
| # # Contrast & brightness | |
| # # A.CLAHE(clip_limit=2.0, p=0.3), | |
| # A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.4), | |
| # # Intensity augmentations that work for N-channel | |
| # A.MultiplicativeNoise(multiplier=(0.9, 1.1), per_channel=True, p=0.3), | |
| # # Noise & blur | |
| # A.GaussianBlur(blur_limit=(3, 5), p=0.4), | |
| # A.GaussNoise(var_limit=(10.0, 50.0), p=0.4), | |
| # # Random cutout to hide regions | |
| # A.CoarseDropout(max_holes=8, max_height=32, max_width=32, min_holes=2, | |
| # fill_value=0, mask_fill_value=0, p=0.3), | |
| # # Normalize (keep your mean/std) | |
| # A.Normalize(mean=(0.2074, 0.1290, 0.1396, 0.1470), | |
| # std=(0.2110, 0.1629, 0.1620, 0.1626)), | |
| # ToTensorV2() | |
| # ], additional_targets=additional_targets) | |
| # 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(std_range=(0.2, 0.4), p=0.3), | |
| # A.Normalize(mean=(0.2074, 0.1290, 0.1396, 0.1470), std=(0.2110, 0.1629, 0.1620,0.1626)), | |
| # ToTensorV2() | |
| # ], additional_targets=additional_targets) | |
| else: | |
| return A.Compose([ | |
| A.Normalize(mean=(0.2074, 0.1290, 0.1396, 0.1470), std=(0.2110, 0.1629, 0.1620,0.1626)), | |
| 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] | |
| breast_mask =np.array(Image.open(sample["breast_mask_path"]).convert('L')) | |
| if os.path.exists(sample["mask_path"]): | |
| mask = np.array(Image.open(sample["mask_path"]).convert('L')) | |
| augmented = self.transform(image=image, mask=mask, breast_mask = breast_mask) | |
| breast_mask = augmented['breast_mask'].unsqueeze(0) # [1, H, W] | |
| image = augmented['image'] # Tensor [C, H, W] | |
| image = image * breast_mask | |
| # image = torch.cat([image, breast_mask], dim=0) | |
| mask = augmented['mask'].long() # Tensor [H, W] as long tensor | |
| mask = mask.unsqueeze(0) # Tensor [1, H, W] if needed | |
| else: | |
| augmented = self.transform(image=image, breast_mask = breast_mask) | |
| image = augmented['image'] | |
| breast_mask = augmented['breast_mask'].unsqueeze(0) # [1, H, W] | |
| image = image * breast_mask | |
| # image = torch.cat([image, breast_mask], dim=0) | |
| mask = None | |
| label = sample["label"] | |
| if label is not None: | |
| label = torch.tensor(label, dtype=torch.long) | |
| return { | |
| 'patient_id': sample['patient_id'], | |
| 'image': image, | |
| 'mask': mask, | |
| 'cls_label': label | |
| } | |
| # contrastive | |
| # def load_and_transform(self, image_paths, mask_path=None): | |
| # image_channels = [] | |
| # for path in image_paths: | |
| # img = np.array(Image.open(path).convert('L')) | |
| # image_channels.append(img) | |
| # image = np.stack(image_channels, axis=-1) | |
| # if mask_path: | |
| # mask = np.array(Image.open(mask_path).convert('L')) | |
| # augmented = self.transform(image=image, mask=mask) | |
| # image = augmented['image'] # Tensor [C, H, W] | |
| # mask = augmented['mask'].long() # Tensor [H, W] as long tensor | |
| # mask = mask.unsqueeze(0) # Tensor [1, H, W] | |
| # return image, mask | |
| # else: | |
| # augmented = self.transform(image=image) | |
| # image = augmented['image'] | |
| # return image #self.transform(image=image)['image'] | |
| # def __getitem__(self, idx): | |
| # anchor_sample = self.samples[idx] | |
| # anchor_label = anchor_sample['label'] | |
| # if os.path.exists(anchor_sample['mask_path']): | |
| # anchor_image, anchor_mask = self.load_and_transform(anchor_sample['image_path'], anchor_sample['mask_path']) | |
| # else: | |
| # anchor_image = self.load_and_transform(anchor_sample['image_path']) | |
| # anchor_mask = None | |
| # # Sample positive (same class, different index) | |
| # positive_idx = idx | |
| # while positive_idx == idx: | |
| # positive_idx = random.choice(self.label_to_indices[anchor_label]) | |
| # positive_sample = self.samples[positive_idx] | |
| # positive_image, positive_mask = self.load_and_transform(positive_sample['image_path'], positive_sample['mask_path']) | |
| # # Sample negative (different class) | |
| # negative_label = random.choice([lbl for lbl in self.label_to_indices if lbl != anchor_label]) | |
| # negative_idx = random.choice(self.label_to_indices[negative_label]) | |
| # negative_sample = self.samples[negative_idx] | |
| # negative_image, negative_mask = self.load_and_transform(negative_sample['image_path'], negative_sample['mask_path']) | |
| # anchor_image = anchor_image*anchor_mask | |
| # positive_image = positive_image*positive_mask | |
| # negative_image = negative_image*negative_mask | |
| # anchor_image = torch.cat([anchor_image, anchor_mask], dim=0) | |
| # negative_image = torch.cat([negative_image, negative_mask], dim=0) | |
| # positive_image = torch.cat([positive_image, positive_mask], dim=0) | |
| # return { | |
| # 'patient_id': self.samples[idx]['patient_id'], | |
| # 'anchor': anchor_image, # Tensor [C, H, W] | |
| # 'anchor_mask': anchor_mask, # Tensor [1, H, W] | |
| # 'positive': positive_image, # Tensor [C, H, W] | |
| # 'negative': negative_image, # Tensor [C, H, W] | |
| # 'anchor_label': torch.tensor(anchor_label, dtype=torch.long) } | |