{"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 --quiet attrdict","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:24.075160Z","iopub.execute_input":"2021-08-16T07:41:24.075677Z","iopub.status.idle":"2021-08-16T07:41:33.483471Z","shell.execute_reply.started":"2021-08-16T07:41:24.075561Z","shell.execute_reply":"2021-08-16T07:41:33.482154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from attrdict import AttrDict\nfrom tqdm import tqdm\nimport os\nimport glob \nimport numpy as np\nfrom PIL import Image\nimport cv2 as cv\nfrom matplotlib import pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:33.487998Z","iopub.execute_input":"2021-08-16T07:41:33.488340Z","iopub.status.idle":"2021-08-16T07:41:33.779210Z","shell.execute_reply.started":"2021-08-16T07:41:33.488302Z","shell.execute_reply":"2021-08-16T07:41:33.778164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision\nimport torchvision.transforms as transforms\nimport torch.nn as nn\nimport torch.nn.functional as F","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:33.781273Z","iopub.execute_input":"2021-08-16T07:41:33.781696Z","iopub.status.idle":"2021-08-16T07:41:35.105102Z","shell.execute_reply.started":"2021-08-16T07:41:33.781642Z","shell.execute_reply":"2021-08-16T07:41:35.103864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = AttrDict(\n    metadata=AttrDict(\n        images_path=\"/kaggle/input/rsna-miccai-png\",\n        labels_path='/kaggle/input/training-labels/train_labels.csv',\n        val_mod=3,        # 1/3 of training images are kept for validation\n        limit_first_k=None, # Load only 10 images of train/val/test\n    ),\n    dataset=AttrDict(\n        input_keys=['FLAIR', 'T1w', 'T1wCE', 'T2w'],\n        img_size=128,     # Resize images to smaller size\n    ),\n    dataloader=AttrDict(\n        batch_size=1,\n        num_workers=4,\n    )\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.107173Z","iopub.execute_input":"2021-08-16T07:41:35.107499Z","iopub.status.idle":"2021-08-16T07:41:35.116873Z","shell.execute_reply.started":"2021-08-16T07:41:35.107469Z","shell.execute_reply":"2021-08-16T07:41:35.115621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_metadata(images_path: str, labels_path: str, val_mod=3, limit_first_k=None):\n    result = {'train': [], 'val': [], 'test': []}\n    \n    scan_id_to_label = {}\n    with open(labels_path, 'r') as f:\n        for i, line in enumerate(f):\n            if i > 0:\n                idx, label = line.split(',')\n                scan_id_to_label[idx.strip()] = int(label.strip())\n    \n    for items_key in ['train', 'test']:\n        all_files = list(os.listdir(f\"{images_path}/{items_key}\"))\n        if limit_first_k:\n            all_files = all_files[:limit_first_k]\n            \n        for scan_id in tqdm(all_files):\n            scan_slices = {}\n\n            for filepath in glob.glob(f\"{images_path}/{items_key}/{scan_id}/*/*.png\"):\n                kind = filepath.split('/')[-2]\n                slices = scan_slices.get(kind, [])\n                slice_id = filepath.split('/')[-1].split('-')[-1].split('.')[0]\n                slices.append((slice_id, filepath))\n                scan_slices[kind] = slices\n            \n            for key in scan_slices:\n                slices = scan_slices[key]\n                slices.sort()\n                scan_slices[key] = [path for _, path in slices]\n            \n            key = items_key\n            if hash(scan_id) % val_mod == 0 and items_key == 'train':\n                key = 'val'\n                \n            result[key].append({\n                'scan': scan_id,\n                'label': scan_id_to_label.get(scan_id),\n                **scan_slices,\n            })\n    return AttrDict(result)","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.119201Z","iopub.execute_input":"2021-08-16T07:41:35.119496Z","iopub.status.idle":"2021-08-16T07:41:35.133293Z","shell.execute_reply.started":"2021-08-16T07:41:35.119468Z","shell.execute_reply":"2021-08-16T07:41:35.132147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metadata = load_metadata(**cfg.metadata)\nprint(f\"{len(metadata.train)} train | {len(metadata.val)} val | {len(metadata.test)} test\")\n# print(metadata.train[1])","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.134396Z","iopub.execute_input":"2021-08-16T07:41:35.134771Z","iopub.status.idle":"2021-08-16T07:41:35.266702Z","shell.execute_reply.started":"2021-08-16T07:41:35.134735Z","shell.execute_reply":"2021-08-16T07:41:35.264181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(list(metadata.values())[0])","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.267824Z","iopub.status.idle":"2021-08-16T07:41:35.268281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset3d(torch.utils.data.Dataset):\n    def __init__(self, metadata, input_keys, img_size=None):\n        super().__init__()\n        \n        self.metadata = metadata\n        self.input_keys = input_keys\n        self.img_size = img_size\n\n        # self.load()\n\n    def load(self, idx, prop):\n        img_size = self.img_size\n        filenames = self.metadata[idx].get(prop)\n        \n        result = []\n        for filename in filenames:\n            img = np.array(Image.open(filename))\n            img = (img / 255 - 0.5) * 2\n            img = cv.resize(img, (img_size, img_size),\n                            interpolation=cv.INTER_NEAREST)\n            result.append(img)\n        return np.array(result)\n\n    def __len__(self):\n        return len(self.metadata)\n    \n    def __getitem__(self, idx):\n        img_size = self.img_size\n        \n        result = self.metadata[idx].copy()\n        for key in self.input_keys:\n            result[key] = self.load(idx, key)\n            \n        return result","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.269695Z","iopub.status.idle":"2021-08-16T07:41:35.270135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = AttrDict(\n  train=Dataset3d(metadata.train, **cfg.dataset),\n  val=Dataset3d(metadata.val, **cfg.dataset),\n  test=Dataset3d(metadata.test, **cfg.dataset),\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.271457Z","iopub.status.idle":"2021-08-16T07:41:35.271920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item = dataset.train[2]\nplt.figure(figsize=(5 * len(cfg.dataset.input_keys), 5))\nprint(f\"Scan: {item['scan']} (label: {item['label']})\")\nfor i, key in enumerate(cfg.dataset.input_keys):\n    print(f\"  {key}: {item[key].shape}\")\n    plt.subplot(1, len(cfg.dataset.input_keys), i + 1)\n    plt.imshow(item[key][10])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.272811Z","iopub.status.idle":"2021-08-16T07:41:35.273223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader = AttrDict(\n    train=torch.utils.data.DataLoader(dataset.train, shuffle=True, **cfg.dataloader),\n    val=torch.utils.data.DataLoader(dataset.train, shuffle=False, **cfg.dataloader),\n    test=torch.utils.data.DataLoader(dataset.train, shuffle=False, **cfg.dataloader),\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.274079Z","iopub.status.idle":"2021-08-16T07:41:35.274501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for batch in dataloader.train:\n#     print(batch)\n#     break","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.275725Z","iopub.status.idle":"2021-08-16T07:41:35.276179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.val[0]['label']","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.277045Z","iopub.status.idle":"2021-08-16T07:41:35.277449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### T1\nFat is depicted in white and water in black.<br/>\nThe shape of the brain can be clearly seen, and morphological abnormalities are easy to detect (Atrophy, tumors, etc.)<br/>\n\n### T2\nWater is painted white.<br/>\nLesions appear white. Suitable for lesion evaluation.<br/>\n\n### FLAIR\nIn T2, the spinal fluid (water) is white and the lesion is also white, so you have to look for the white in the white, which is difficult to understand.<br/>\nFLAIR can be roughly thought of as T2, in which the water is also black, making it easier to find the lesion.<br/>\n","metadata":{}},{"cell_type":"markdown","source":"### Observatie\n - T1w este T1 weighted pre-contrast\n - T2wCE este T1 weighted post-contrast\n - T2w este T2 weighted\n - FLAIR = Fluid Attenuated Inversion Recovery\n - fiecare folder contine tipuri diferite de RMN (contrastul difera)\n - NU exista o regula de orientare a scanarilor (de ex. T1w contine rmn in plan sagital, dar si in plan coronal sau orizontal)\n - Plan Sagital = stanga-dreapta\n - Plan Coronal = fata-spate\n - Plan Orizontal = sus-jos\n<br/>\n<br/>\n\n[Link](https://case.edu/med/neurology/NR/MRI%20Basics.htm) explicatii la ce inseamna T1w, T2w, FLAIR.<br/>","metadata":{}},{"cell_type":"markdown","source":"## #1 try - FLAIR folder only","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(6 * len(cfg.dataset.input_keys),6))\nfor i in range(4):\n    plt.subplot(1, len(cfg.dataset.input_keys), i + 1)\n    plt.imshow(dataset.train[i]['FLAIR'][20])\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.278443Z","iopub.status.idle":"2021-08-16T07:41:35.278854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient = 10\nlen(dataset.train[patient]['T2w'])","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.279716Z","iopub.status.idle":"2021-08-16T07:41:35.280130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_images(dataset, patient: int, folder: str, train=True):\n    images = []\n    if train == True:\n        for img in dataset.train[patient][folder]:\n            # Exclude the blank images\n            if np.max(img)!=0:\n                images.append(img)\n            else:\n                pass\n    else:\n        for img in dataset.test[patient][folder]:\n            # Exclude the blank images\n            if np.max(img)!=0:\n                images.append(img)\n            else:\n                pass\n    \n    return images","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.281250Z","iopub.status.idle":"2021-08-16T07:41:35.281679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient = 6\n\nimages = get_images(dataset, patient, 'FLAIR')\nprint('No of images:', len(images))\n\nfig = plt.figure(figsize=(50,50))\n\nc = 1\nfor image in images:\n    ax = fig.add_subplot(len(images)//10+1, 10, c)\n    ax.imshow(image, cmap='gray')\n    c+=1\n    \n    plt.axis('off')\n    \nfig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.282711Z","iopub.status.idle":"2021-08-16T07:41:35.283123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainX, testX, trainY, testY = [], [], [], []\n\nfor i in dataset.train:\n    trainX += i['FLAIR'][40]\n    trainY += i['label']\n    \nfor i in dataset.test:\n    testX += i['FLAIR'][40]\n    testY += i['label']\n    ","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.284202Z","iopub.status.idle":"2021-08-16T07:41:35.284646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\n\nmodel = Sequential()\nmodel.add(Conv2D(32, kernel_size = (3, 3), input_shape=(128, 128, 1)))\nmodel.add(Activation(\"relu\"))\nmodel.add(Conv2D(64, (3, 3)))\nmodel.add(Activation(\"relu\"))\nmodel.add(MaxPooling2D(pool_size = (2, 2)))\nmodel.add(Flatten())\nmodel.add(Dense(100))\nmodel.add(Activation(\"relu\"))\nmodel.add(Dropout(0.5))\nmodel.add(Dense(10))\nmodel.add(Activation('softmax'))\n\nmodel.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])\n\nmodel.fit(trainX, trainY, batch_size = 32, epochs = 10, verbose=2, validation_data=(testX,testY))\n","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.285687Z","iopub.status.idle":"2021-08-16T07:41:35.286107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item = dataset.train[0]['T2w']\ninputs = torch.FloatTensor(item).reshape([1, 1, *item.shape])\npredictions = net(inputs)\nprint(inputs.shape, predictions.shape, predictions.item())","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.287242Z","iopub.status.idle":"2021-08-16T07:41:35.287675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"raw","source":"import torch.optim as optim\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)","metadata":{"execution":{"iopub.status.busy":"2021-08-10T13:30:13.262756Z","iopub.execute_input":"2021-08-10T13:30:13.263122Z","iopub.status.idle":"2021-08-10T13:30:13.268334Z","shell.execute_reply.started":"2021-08-10T13:30:13.263089Z","shell.execute_reply":"2021-08-10T13:30:13.267098Z"}}},{"cell_type":"code","source":"def minMaxNormalize(volume):\n    # values between 0 and 1\n    min = 0\n    max = 255\n    volume[volume < min] = min\n    volume[volume > max] = max\n    volume = (volume - min) / (max - min)\n    volume = volume.astype(\"float32\")\n    return volume","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.288722Z","iopub.status.idle":"2021-08-16T07:41:35.289162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# each train entry has 4 3D images\n# ","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.290059Z","iopub.status.idle":"2021-08-16T07:41:35.290463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for data in dataset.train:\n    for i,image in enumerate(data['T2w'],0):\n        plt.imshow(image)\n        break","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.291315Z","iopub.status.idle":"2021-08-16T07:41:35.291737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(2):\nrunning_loss = 0.0\nfor i, data in enumerate(trainloader, 0):\n  inputs, labels = data\n  optimizer.zero_grad()\n\n  # forward + backward + optimize\n  outputs = net(inputs)\n  loss = criterion(outputs, labels)\n  loss.backward()\n  optimizer.step()\n\n  # print statistics\n  running_loss += loss.item()\n  if i % 2000 == 1999:    # print every 2000 mini-batches\n      print('[%d, %5d] loss: %.3f' %\n            (epoch + 1, i + 1, running_loss / 2000))\n      running_loss = 0.0","metadata":{"execution":{"iopub.status.busy":"2021-08-16T07:41:35.292729Z","iopub.status.idle":"2021-08-16T07:41:35.293136Z"},"trusted":true},"execution_count":null,"outputs":[]}]}