{"cells":[{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.sampler import Sampler\nimport torch\nimport pandas as pd\nfrom sklearn.preprocessing import MultiLabelBinarizer\nimport pathlib\nimport torchvision.transforms as transforms\nimport torch\nimport PIL\n\n\nLABEL_MAP = {\n0: \"Nucleoplasm\" ,\n1: \"Nuclear membrane\"   ,\n2: \"Nucleoli\"   ,\n3: \"Nucleoli fibrillar center\",   \n4: \"Nuclear speckles\"   ,\n5: \"Nuclear bodies\"   ,\n6: \"Endoplasmic reticulum\"   ,\n7: \"Golgi apparatus\"  ,\n8: \"Peroxisomes\"   ,\n9:  \"Endosomes\"   ,\n10: \"Lysosomes\"   ,\n11: \"Intermediate filaments\"  , \n12: \"Actin filaments\"   ,\n13: \"Focal adhesion sites\"  ,\n14: \"Microtubules\"   ,\n15: \"Microtubule ends\"   ,\n16: \"Cytokinetic bridge\"   ,\n17: \"Mitotic spindle\"  ,\n18: \"Microtubule organizing center\",  \n19: \"Centrosome\",\n20: \"Lipid droplets\"   ,\n21: \"Plasma membrane\"  ,\n22: \"Cell junctions\"   ,\n23: \"Mitochondria\"   ,\n24: \"Aggresome\"   ,\n25: \"Cytosol\" ,\n26: \"Cytoplasmic bodies\",\n27: \"Rods & rings\"}\n\n\n\nimage_transform = transforms.Compose([\n            transforms.ToTensor(),\n    \n        ])\n\n\nclass MultiBandMultiLabelDataset(Dataset):\n    BANDS_NAMES = ['_red.png','_green.png','_blue.png','_yellow.png']\n    \n    def __len__(self):\n        return len(self.images_df)\n    \n    def __init__(self, images_df, base_path, image_transform=image_transform, augmentator=None):\n        if not isinstance(base_path, pathlib.Path):\n            base_path = pathlib.Path(base_path)\n            \n        self.images_df = images_df.copy()\n        self.image_transform = image_transform\n        self.augmentator = augmentator\n        self.images_df.Id = self.images_df.Id.apply(lambda x: base_path / x)\n        self.mlb = MultiLabelBinarizer(classes=list(LABEL_MAP.keys()))\n                                      \n        \n    def __getitem__(self, index):\n        X = self._load_multiband_image(index)\n        y = self._load_multilabel_target(index)\n        \n        # augmentator can be for instance imgaug augmentation object\n        if self.augmentator is not None:\n            X = self.augmentator(X)\n            \n        X = self.image_transform(X)\n            \n        return X, y \n        \n    def _load_multiband_image(self, index):\n        row = self.images_df.iloc[index]\n        image_bands = []\n        for band_name in self.BANDS_NAMES:\n            p = str(row.Id.absolute()) + band_name\n            pil_channel = PIL.Image.open(p)\n            image_bands.append(pil_channel)\n            \n        # lets pretend its a RBGA image to support 4 channels\n        band4image = PIL.Image.merge('RGBA', bands=image_bands)\n        return band4image\n    \n    \n    def _load_multilabel_target(self, index):\n        return list(map(int, self.images_df.iloc[index].Target.split(' ')))\n    \n        \n    def collate_func(self, batch):\n        \n        images = [x[0] for x in batch]\n        labels = [x[1] for x in batch]\n        \n        labels_one_hot  = self.mlb.fit_transform(labels)\n        \n        return torch.stack(images), torch.FloatTensor(labels_one_hot)\n                                               \n\ndf = pd.read_csv('../input/train.csv')\ng = MultiBandMultiLabelDataset(df, base_path='../input/train')\ng.collate_func([g[i] for i in range(16)])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"c9967735da18108a205273faf55a56cb840c1a0a"},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"1d176908b1705080c510d25c342fbb9a41136e25"},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"8122a97807d46a8a082f752dc8b97db61369c9d4"},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}