{"cells":[{"metadata":{"_uuid":"1b961b61d1e8353e64c4714f1b3c4889a72f8a2d"},"cell_type":"markdown","source":"## Important Notice\nThe code is (obviously) not runnable, but it is part of a starter kit for pytorch users to begin their experiments.\nSome features inclue:\n1. Pre-process images into .npy for faster loading during training/validation\n2. self.probweights that can be used in WeightedRandomSampler to provide more balanced batching"},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import sys\nimport os\nprint(sys.path)\nsys.path.append('../../')\n#from dependencies import DATA, CODE\nfrom tqdm import tqdm\n#from vision.augmentations import compute_center_pad, do_gamma\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data.dataset import Dataset\nimport torch\nimport gc\nDATA = '.'\nCODE = '.'","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","collapsed":true,"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":false},"cell_type":"code","source":"class ProteinDatasetLoad(Dataset):\n\n    def __init__(self, split, augment=null_augment, mode='train',size=256):\n        super(ProteinDatasetLoad, self).__init__()\n        self.split   = split\n        self.mode   = mode\n        self.augment = augment\n        self.size = size\n        self.truth = pd.read_csv(DATA+'train.csv')\n        self.problist = [(27, 0.125),\n                         (15, 0.05555555555555555),\n                         (10, 0.043478260869565216),\n                         (9, 0.02702702702702703),\n                         (8, 0.022222222222222223),\n                         (20, 0.007633587786259542),\n                         (17, 0.006024096385542169),\n                         (24, 0.003875968992248062),\n                         (26, 0.003816793893129771),\n                         (13, 0.002352941176470588),\n                         (16, 0.002288329519450801),\n                         (12, 0.0017482517482517483),\n                         (22, 0.001579778830963665),\n                         (18, 0.0013679890560875513),\n                         (6, 0.001221001221001221),\n                         (11, 0.001141552511415525),\n                         (14, 0.0011312217194570137),\n                         (1, 0.0010183299389002036),\n                         (19, 0.0008665511265164644),\n                         (3, 0.0008285004142502071),\n                         (4, 0.0006844626967830253),\n                         (5, 0.0004935834155972359),\n                         (7, 0.0004434589800443459),\n                         (23, 0.0004187604690117253),\n                         (2, 0.00034916201117318437),\n                         (21, 0.00033090668431502316),\n                         (25, 0.00015130882130428205),\n                         (0, 9.693679720822024e-05)]\n        self.label_dict = {0:  'Nucleoplasm',\n                        1:  'Nuclear membrane',\n                        2:  'Nucleoli',\n                        3:  'Nucleoli fibrillar center',\n                        4:  'Nuclear speckles',\n                        5:  'Nuclear bodies',\n                        6:  'Endoplasmic reticulum',\n                        7:  'Golgi apparatus',\n                        8:  'Peroxisomes',\n                        9:  'Endosomes',\n                        10:  'Lysosomes',\n                        11:  'Intermediate filaments',\n                        12:  'Actin filaments',\n                        13:  'Focal adhesion sites',\n                        14:  'Microtubules',\n                        15:  'Microtubule ends',\n                        16:  'Cytokinetic bridge',\n                        17:  'Mitotic spindle',\n                        18:  'Microtubule organizing center',\n                        19:  'Centrosome',\n                        20:  'Lipid droplets',\n                        21:  'Plasma membrane',\n                        22:  'Cell junctions',\n                        23:  'Mitochondria',\n                        24:  'Aggresome',\n                        25:  'Cytosol',\n                        26:  'Cytoplasmic bodies',\n                        27:  'Rods & rings' }\n\n        split_file =  CODE + '/datasets/Protein/splits/' + split\n        lines = read_list_from_file(split_file)\n\n        self.ids    = []\n\n        for l in tqdm(lines):                                                                                                                                                                 \n            folder, name = l.split('/')\n            self.ids.append(name)\n\n        print(len(lines))\n        def save_to_dir(i):\n           if(i%1000==0): print(\"loaded sample {}\".format(i))\n           folder, name = lines[i].split('/')\n           image_file = DATA+folder+\"/\" + name\n           if not(os.path.isfile(DATA+folder+ \"_np\"+ str(size) +\"/\"+name+\".npy\")):\n              r = cv2.imread(image_file+'_red.png',cv2.IMREAD_GRAYSCALE).astype(\"uint8\")\n              g = cv2.imread(image_file+'_green.png',cv2.IMREAD_GRAYSCALE).astype(\"uint8\")\n              b = cv2.imread(image_file+'_blue.png',cv2.IMREAD_GRAYSCALE).astype(\"uint8\")\n              y = cv2.imread(image_file+'_yellow.png',cv2.IMREAD_GRAYSCALE).astype(\"uint8\")\n              r = cv2.resize(r,(size,size))\n              g = cv2.resize(g,(size,size))\n              b = cv2.resize(b,(size,size))\n              y = cv2.resize(y,(size,size))\n\n              image = np.dstack((r,g,b,y))\n              np.save(DATA+folder+ \"_np\"+ str(size) +\"/\"+name,image)\n\n           else:\n\n              image = np.load(DATA+folder+ \"_np\"+ str(size) +\"/\"+name+\".npy\")\n\n        num_cores = multiprocessing.cpu_count()\n        if self.split.find(\"valid\")==-1:\n            Parallel(n_jobs=num_cores, prefer=\"threads\")(delayed(save_to_dir)(i) for i in range(len(lines)))\n        else:\n            import random\n            self.images=Parallel(n_jobs=num_cores, prefer=\"threads\")(delayed(save_to_dir)(i) for i in range(len(lines)))\n\n\n        self.annotations  = []\n        self.probweights = []\n        if self.mode in ['train','valid']:\n            for l in tqdm(lines):\n                folder, file = l.split('/')\n                self.folder = folder\n\n                label_encode = self.truth.loc[self.truth['Id']==file]['Target']\n                label_list = [[int(i) for i in s.split()] for s in label_encode]\n\n                for tuplea in self.problist:\n                    if tuplea[0] in label_list[0]:\n                        self.probweights.append(tuplea[1])\n                        break\n                self.annotations.append( np.eye(len(self.label_dict),dtype=np.float)[label_list].sum(axis=0))\n        elif self.mode in ['test']:\n            self.annotations  = [[] for l in lines]\n\n        #-------\n        print('\\tProteinDataset')\n        print('\\tsplit            = %s'%split)\n        print('\\tlen(self.ids) = %d'%len(self.ids))\n        print('')\n\n\n    def __getitem__(self, index):\n        image = np.load(DATA+self.folder+\"_np\"+str(self.size)+\"/\"+self.ids[index]+\".npy\")\n        annotations  = self.annotations[index]\n\n        return self.augment(image, annotations, index)\n\n    def __len__(self):\n        return len(self.ids)\n\n                                                     ","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}