# Dataloader for abdominal images import glob import numpy as np from .m4_utils import niftiio as nio from .m4_utils import transform_utils as trans from .m4_utils.abd_dataset_utils import get_normalize_op from .m4_utils.transform_albu import get_albu_transforms, get_resize_transforms import torch import os from pdb import set_trace from multiprocessing import Process from .basic import BasicDataset LABEL_NAME = ["bg", "NCR", "ED", "ET"] def normalize(data, mean, stddev, eps=0.0): """ Normalize the given tensor. Applies the formula (data - mean) / (stddev + eps). Args: data (torch.Tensor): Input data to be normalized. mean (float): Mean value. stddev (float): Standard deviation. eps (float, default=0.0): Added to stddev to prevent dividing by zero. Returns: torch.Tensor: Normalized tensor """ return (data - mean) / (stddev + eps) def normalize_instance(data, mean=None, std=None, eps=0.0): """ Normalize the given tensor with instance norm/ Applies the formula (data - mean) / (stddev + eps), where mean and stddev are computed from the data itself. Args: data (torch.Tensor): Input data to be normalized eps (float): Added to stddev to prevent dividing by zero Returns: torch.Tensor: Normalized tensor """ if mean is None: mean = data.mean() if std is None: std = data.std() return normalize(data, mean, std, eps), mean, std class BrainDataset(BasicDataset): def __init__(self, mode, base_dir, image_size, nclass, domains, aux_modality, **kwargs): """ Args: mode: 'train', 'val', 'test', 'test_all' transforms: naive data augmentations used by default. Photometric transformations slightly better than those configured by Zhang et al. (bigaug) idx_pct: train-val-test split for source domain extern_norm_fn: feeding data normalization functions from external, only concerns CT-MR cross domain scenario """ self.dataset_key = "brain" transforms = get_albu_transforms(mode, image_size) if isinstance(domains, str): domains = [domains] super(BrainDataset, self).__init__(image_size, mode, transforms, base_dir, domains, aux_modality, nclass=nclass, LABEL_NAME=LABEL_NAME, filter_non_labeled=True, **kwargs) def hwc_to_chw(self,img): img = np.float32(img) img = np.transpose(img, (2, 0, 1)) # [C, H, W] img = torch.from_numpy( img.copy() ) return img def perform_trans(self, img, mask, aux): T = self.albu_transform if self.is_train else self.test_resizer buffer = T(image = img, mask=mask, image2=aux) # [0 - 255] img, mask, aux = buffer['image'], buffer['mask'], buffer['image2'] if len(mask.shape) == 2: mask = mask[..., None] # if self.is_train: # img, mask, aux = self.get_patch_from_img(img, mask, aux, crop_size=self.crop_size) # 192 return img, mask, aux def __getitem__(self, index): index = index % len(self.actual_dataset) curr_dict = self.actual_dataset[index] # numpy # ----------------------- Extract Slice ----------------------- img, mask, aux = curr_dict["img"], curr_dict["lb"], curr_dict["aux"] # H, W, C, [0 - 255] domain, pid = curr_dict["domain"], curr_dict["pid"] mean, std = curr_dict['mean'], curr_dict['std'] aux_mean, aux_std = curr_dict['aux_mean'], curr_dict['aux_std'] # max, min = img.max(), img.min() std = 1 if std < 1e-3 else std # img = (img - mean) / std ### 对input image和target image都做(x-mean)/std的归一化操作 img, img_mean, img_std = normalize_instance(img, eps=1e-6) # mean=mean, std=std, aux, aux_mean, aux_std = normalize_instance(aux, eps=1e-6) # mean=aux_mean, std=aux_std, ### clamp input to ensure training stability. img = np.clip(img, -6, 6) aux = np.clip(aux, -6, 6) mask = mask[..., 0] img, mask, aux = self.perform_trans(img, mask, aux) img, mask, aux = map(lambda arr: self.hwc_to_chw(arr), [img, mask, aux]) img = np.clip(img, -6, 6) aux = np.clip(aux, -6, 6) if self.tile_z_dim > 1 and self.input_window == 1 and self.num_channels == 3 : img = img.repeat( [ self.tile_z_dim, 1, 1] ) assert img.ndimension() == 3 data = {"img": img, "lb": mask, "aux": aux, "img_mean": img_mean, "img_std": img_std, "aux_mean": aux_mean, "aux_std": aux_std, "is_start": curr_dict["is_start"], "is_end": curr_dict["is_end"], "nframe": np.int32(curr_dict["nframe"]), "scan_id": curr_dict["scan_id"], "z_id": curr_dict["z_id"], "file_id": curr_dict["file_id"] } return data