{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install fastai==2.5.1","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:25:53.522495Z","iopub.execute_input":"2021-09-22T20:25:53.523252Z","iopub.status.idle":"2021-09-22T20:26:04.991830Z","shell.execute_reply.started":"2021-09-22T20:25:53.523148Z","shell.execute_reply":"2021-09-22T20:26:04.990536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Introduction\nThis notebook is to introduce how to build a customize fastai's DataLoader for a task with item is sequence of images (MRI Images, Video, ...) \n\nWith fastai's DataLoader, you can leverage many of the neat fastai functionalities like *item transformation* or *batch transformation* or my favorite part: *show batch*. This great feature help you to visualize how your raw data is tranformed to the input of your model, and it is especially helpful to debug or present to others. \n\nYou can checkout the official tutorial for *Using fastai on a custom new task* here: https://docs.fast.ai/tutorial.siamese.html","metadata":{}},{"cell_type":"code","source":"from fastai.vision.all import *\nimport fastai\nimport random\nimport PIL","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:04.994090Z","iopub.execute_input":"2021-09-22T20:26:04.994375Z","iopub.status.idle":"2021-09-22T20:26:07.704503Z","shell.execute_reply.started":"2021-09-22T20:26:04.994339Z","shell.execute_reply":"2021-09-22T20:26:07.703343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## For experimenting purpose, only type 'T1W' is chosen and sequence with more than 16 images\nmri_type = 'T1w'\nmin_subset = 16\npatients_seq_folder = [patient/mri_type for patient in Path('../input/rsna-miccai-png/train').ls() if (patient/mri_type).exists() and len((patient/mri_type).ls()) > min_subset]","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-09-22T20:26:07.706062Z","iopub.execute_input":"2021-09-22T20:26:07.706320Z","iopub.status.idle":"2021-09-22T20:26:21.062408Z","shell.execute_reply.started":"2021-09-22T20:26:07.706289Z","shell.execute_reply":"2021-09-22T20:26:21.061306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patients_seq_folder[:5]","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.064584Z","iopub.execute_input":"2021-09-22T20:26:21.064831Z","iopub.status.idle":"2021-09-22T20:26:21.074891Z","shell.execute_reply.started":"2021-09-22T20:26:21.064802Z","shell.execute_reply":"2021-09-22T20:26:21.073616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = pd.read_csv('../input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv')\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.076110Z","iopub.execute_input":"2021-09-22T20:26:21.076387Z","iopub.status.idle":"2021-09-22T20:26:21.111382Z","shell.execute_reply.started":"2021-09-22T20:26:21.076355Z","shell.execute_reply":"2021-09-22T20:26:21.110362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def open_image(fname, size=224):\n    img = PIL.Image.open(fname)\n    img = img.resize((size, size))\n    t = torch.Tensor(np.array(img))\n    return t.float()/255.0","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.112903Z","iopub.execute_input":"2021-09-22T20:26:21.113403Z","iopub.status.idle":"2021-09-22T20:26:21.119807Z","shell.execute_reply.started":"2021-09-22T20:26:21.113367Z","shell.execute_reply":"2021-09-22T20:26:21.118887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To help fastai know how to show your batch, each item (including x and y) need to be an instance of a class which has a *show* function. And by subclassing fastuple, fastai will know how to apply the transformation for each element based on their type (For example applying Resize on to PILImage and not the string)","metadata":{}},{"cell_type":"code","source":"class SeqImage(fastuple):\n    def show(self, ctx=None, **kwargs):\n        *imgs, label = self\n        if not isinstance(imgs[0], Tensor):\n            imgs = [tensor(img).permute(2,0,1) for img in imgs]\n        img_cat = torch.cat(imgs, dim=2)\n        return show_image(img_cat, figsize=(20,20), title=label, ctx=ctx, **kwargs)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.121075Z","iopub.execute_input":"2021-09-22T20:26:21.121581Z","iopub.status.idle":"2021-09-22T20:26:21.130628Z","shell.execute_reply.started":"2021-09-22T20:26:21.121534Z","shell.execute_reply":"2021-09-22T20:26:21.129772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# splits the data to training set and validation set\nsplits = RandomSplitter()(patients_seq_folder)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.131890Z","iopub.execute_input":"2021-09-22T20:26:21.132124Z","iopub.status.idle":"2021-09-22T20:26:21.162976Z","shell.execute_reply.started":"2021-09-22T20:26:21.132096Z","shell.execute_reply":"2021-09-22T20:26:21.161757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_folders, valid_folders = L(patients_seq_folder)[splits[0]], L(patients_seq_folder)[splits[1]]","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.164276Z","iopub.execute_input":"2021-09-22T20:26:21.164552Z","iopub.status.idle":"2021-09-22T20:26:21.170669Z","shell.execute_reply.started":"2021-09-22T20:26:21.164521Z","shell.execute_reply":"2021-09-22T20:26:21.169687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We will try to show an SeqImage as below","metadata":{}},{"cell_type":"code","source":"files_test = train_folders[0].ls()[:16]","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.174713Z","iopub.execute_input":"2021-09-22T20:26:21.175415Z","iopub.status.idle":"2021-09-22T20:26:21.184534Z","shell.execute_reply.started":"2021-09-22T20:26:21.175363Z","shell.execute_reply":"2021-09-22T20:26:21.183866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs_label = [PILImage.create(file) for file in files_test]\nimgs_label.append(0)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.186428Z","iopub.execute_input":"2021-09-22T20:26:21.187086Z","iopub.status.idle":"2021-09-22T20:26:21.360740Z","shell.execute_reply.started":"2021-09-22T20:26:21.187039Z","shell.execute_reply":"2021-09-22T20:26:21.359545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"s = SeqImage(imgs_label)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.362005Z","iopub.execute_input":"2021-09-22T20:26:21.362226Z","iopub.status.idle":"2021-09-22T20:26:21.367403Z","shell.execute_reply.started":"2021-09-22T20:26:21.362199Z","shell.execute_reply":"2021-09-22T20:26:21.366304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tst = Resize(224)(s)\ntst = ToTensor()(tst)\ntst.show();","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.368907Z","iopub.execute_input":"2021-09-22T20:26:21.369118Z","iopub.status.idle":"2021-09-22T20:26:21.767689Z","shell.execute_reply.started":"2021-09-22T20:26:21.369093Z","shell.execute_reply":"2021-09-22T20:26:21.766338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Then encodes function in Transform class used to apply the transformation to each item (similar to *forward* in Pytorch modules) ","metadata":{}},{"cell_type":"code","source":"class SeqTransform(Transform):\n    def encodes(self, folder):\n        files = folder.ls()\n        files = sorted(random.sample(files, min_subset), key=lambda path: int(path.stem.split('-')[1]))\n        imgs = [PILImage.create(file) for file in files]\n        label = labels[labels['BraTS21ID']==int((folder).parent.name)]['MGMT_value'].values[0]\n        return SeqImage(*imgs, label)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.769169Z","iopub.execute_input":"2021-09-22T20:26:21.769426Z","iopub.status.idle":"2021-09-22T20:26:21.777859Z","shell.execute_reply.started":"2021-09-22T20:26:21.769399Z","shell.execute_reply":"2021-09-22T20:26:21.776703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tfm = SeqTransform()","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.779424Z","iopub.execute_input":"2021-09-22T20:26:21.779652Z","iopub.status.idle":"2021-09-22T20:26:21.789542Z","shell.execute_reply.started":"2021-09-22T20:26:21.779627Z","shell.execute_reply":"2021-09-22T20:26:21.788905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tls = TfmdLists(patients_seq_folder, tfm, splits=splits)\n","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.790671Z","iopub.execute_input":"2021-09-22T20:26:21.791004Z","iopub.status.idle":"2021-09-22T20:26:21.909001Z","shell.execute_reply.started":"2021-09-22T20:26:21.790976Z","shell.execute_reply":"2021-09-22T20:26:21.907872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_at(tls.valid, 3);","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:21.910568Z","iopub.execute_input":"2021-09-22T20:26:21.910830Z","iopub.status.idle":"2021-09-22T20:26:22.987464Z","shell.execute_reply.started":"2021-09-22T20:26:21.910800Z","shell.execute_reply":"2021-09-22T20:26:22.986284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = tls.dataloaders(after_item=[Resize(224), ToTensor], \n                      after_batch=[IntToFloatTensor])","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:26:22.989568Z","iopub.execute_input":"2021-09-22T20:26:22.989910Z","iopub.status.idle":"2021-09-22T20:26:23.148611Z","shell.execute_reply.started":"2021-09-22T20:26:22.989865Z","shell.execute_reply":"2021-09-22T20:26:23.147738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define a show_batch function as below to make show_batch works. x is a batch of SeqImage item","metadata":{}},{"cell_type":"code","source":"@typedispatch\ndef show_batch(x:SeqImage, y, samples, ctxs=None, max_n=6, nrows=None, ncols=1, figsize=None, **kwargs):\n    if figsize is None: figsize = (ncols*6, max_n//ncols * 3)\n    if ctxs is None: ctxs = get_grid(min(x[0].shape[0], max_n), nrows=None, ncols=ncols, figsize=figsize)\n    for index,ctx in enumerate(ctxs): \n        imgs_ls = [x[i][index] for i in range(min_subset)]\n        label = int(x[-1][index])\n        SeqImage(*imgs_ls, label).show(ctx=ctx)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:29:57.624710Z","iopub.execute_input":"2021-09-22T20:29:57.625091Z","iopub.status.idle":"2021-09-22T20:29:57.635371Z","shell.execute_reply.started":"2021-09-22T20:29:57.625053Z","shell.execute_reply":"2021-09-22T20:29:57.633914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b = dls.one_batch()","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:29:59.424379Z","iopub.execute_input":"2021-09-22T20:29:59.424756Z","iopub.status.idle":"2021-09-22T20:30:11.105582Z","shell.execute_reply.started":"2021-09-22T20:29:59.424722Z","shell.execute_reply":"2021-09-22T20:30:11.104371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls.show_batch(figsize=(20,20))","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:30:11.110984Z","iopub.execute_input":"2021-09-22T20:30:11.112090Z","iopub.status.idle":"2021-09-22T20:30:25.106150Z","shell.execute_reply.started":"2021-09-22T20:30:11.112031Z","shell.execute_reply":"2021-09-22T20:30:25.104939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}