{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":6399540,"sourceType":"datasetVersion","datasetId":3689401}],"dockerImageVersionId":30554,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -qq monai","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-02T16:06:10.685617Z","iopub.execute_input":"2023-09-02T16:06:10.686459Z","iopub.status.idle":"2023-09-02T16:06:27.142065Z","shell.execute_reply.started":"2023-09-02T16:06:10.686424Z","shell.execute_reply":"2023-09-02T16:06:27.140716Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport gc\nimport glob\nimport tqdm\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\n\nfrom torchvision.utils import make_grid\n\nfrom sklearn.model_selection import train_test_split\n\nimport monai\nfrom monai import transforms\nfrom monai.networks.nets import UNet\nfrom monai.losses import FocalLoss\nfrom monai.metrics import ConfusionMatrixMetric\nfrom monai.data import Dataset, DataLoader\nfrom monai.visualize import blend_images\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Available device: ', device)","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:06:31.139698Z","iopub.execute_input":"2023-09-02T16:06:31.140063Z","iopub.status.idle":"2023-09-02T16:07:17.741417Z","shell.execute_reply.started":"2023-09-02T16:06:31.140032Z","shell.execute_reply":"2023-09-02T16:07:17.740444Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_transform = transforms.Compose([\n    transforms.LoadImaged(keys=['image', 'mask'], image_only=True),\n    transforms.ScaleIntensityd(keys=['image']),\n    transforms.EnsureChannelFirstd(keys=['image', 'mask']),\n    transforms.ToTensord(keys=['image'], dtype=torch.float32),\n    transforms.ToTensord(keys=['mask'], dtype=torch.long),\n])","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:17.744232Z","iopub.execute_input":"2023-09-02T16:07:17.744976Z","iopub.status.idle":"2023-09-02T16:07:17.754467Z","shell.execute_reply.started":"2023-09-02T16:07:17.744947Z","shell.execute_reply":"2023-09-02T16:07:17.753573Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_dir = '/kaggle/input/rsnasegdataset/dataset/dataset/images'\nseries_list = os.listdir(series_dir)\n\ntrain_series, val_series = train_test_split(series_list, test_size=0.2, random_state=42)\n\ntrain_paths = []\nfor series in train_series:\n    train_paths.extend(glob.glob(os.path.join(series_dir, series, '*.png')))\n    \nval_paths = []\nfor series in val_series:\n    val_paths.extend(glob.glob(os.path.join(series_dir, series, '*.png')))","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:25.724383Z","iopub.execute_input":"2023-09-02T16:07:25.724755Z","iopub.status.idle":"2023-09-02T16:07:32.902911Z","shell.execute_reply.started":"2023-09-02T16:07:25.724723Z","shell.execute_reply":"2023-09-02T16:07:32.901911Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seg_paths = [path.replace('images', 'segmentations') for path in train_paths]\nval_seg_paths = [path.replace('images', 'segmentations') for path in val_paths]\n\ntrain_dict = [{'image': ct_image, 'mask': seg_image} for ct_image, seg_image in zip(train_paths, train_seg_paths)]\nval_dict = [{'image': ct_image, 'mask': seg_image} for ct_image, seg_image in zip(val_paths, val_seg_paths)]","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:34.299323Z","iopub.execute_input":"2023-09-02T16:07:34.299704Z","iopub.status.idle":"2023-09-02T16:07:34.315148Z","shell.execute_reply.started":"2023-09-02T16:07:34.299673Z","shell.execute_reply":"2023-09-02T16:07:34.314116Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = Dataset(train_dict, transform=data_transform)\ntrain_loader = DataLoader(train_ds, shuffle=True, batch_size=16, num_workers=2)\n\nval_ds = Dataset(val_dict, transform=data_transform)\nval_loader = DataLoader(val_ds, shuffle=False, batch_size=16, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:34.457045Z","iopub.execute_input":"2023-09-02T16:07:34.457363Z","iopub.status.idle":"2023-09-02T16:07:34.465166Z","shell.execute_reply.started":"2023-09-02T16:07:34.457334Z","shell.execute_reply":"2023-09-02T16:07:34.464223Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = []\nbatch = next(iter(train_loader))\n\nfor image, label in zip(batch['image'], batch['mask']):\n    ret = blend_images(image, label=label, alpha=0.5, cmap='hsv')\n    results.append(ret)\n    \nplt.figure(figsize=(6, 6))\nplt.imshow(make_grid(results, nrow=4).permute(1, 2, 0))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:35.638879Z","iopub.execute_input":"2023-09-02T16:07:35.639227Z","iopub.status.idle":"2023-09-02T16:07:39.355101Z","shell.execute_reply.started":"2023-09-02T16:07:35.639199Z","shell.execute_reply":"2023-09-02T16:07:39.353112Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(\n    spatial_dims=2,\n    in_channels=1,\n    out_channels=6,\n    channels=(32, 64, 128, 256, 512),\n    strides=(2, 2, 2, 2),\n    num_res_units=2,\n).to(device)\n\ncriterion = FocalLoss(to_onehot_y=True, use_softmax=True)\nmetric = ConfusionMatrixMetric(include_background=False, metric_name=['f1_score'])\noptimizer = torch.optim.Adam(model.parameters(), 1e-3)","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:07:39.35698Z","iopub.execute_input":"2023-09-02T16:07:39.357316Z","iopub.status.idle":"2023-09-02T16:07:45.112861Z","shell.execute_reply.started":"2023-09-02T16:07:39.357283Z","shell.execute_reply":"2023-09-02T16:07:45.111815Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"n_epochs = 10\nbest_score = 0\n\nfor epoch in range(n_epochs):\n    progress_bar = tqdm.tqdm(enumerate(train_loader, start=1), total=len(train_loader), ncols=100)\n    progress_bar.set_description(f'Epoch {epoch}')\n    epoch_loss, epoch_score = 0, 0\n    model.train()\n    for step, batch in progress_bar:\n        inputs, labels = batch['image'].to(device), batch['mask'].to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        \n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        epoch_loss += loss.item()\n        metric(outputs.argmax(1), labels[:, 0])\n        epoch_score = metric.aggregate()[0].item()\n        \n        progress_bar.set_postfix({\n            'train_loss': epoch_loss / step,\n            'train_score': epoch_score,\n        })\n    metric.reset()\n        \n    progress_bar = tqdm.tqdm(enumerate(val_loader, start=1), total=len(val_loader), ncols=100)\n    progress_bar.set_description(f'Epoch {epoch}')\n    epoch_val_loss, epoch_val_score = 0, 0\n    model.eval()\n    with torch.no_grad():\n        for step, batch in progress_bar:\n            inputs, labels = batch['image'].to(device), batch['mask'].to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            \n            epoch_val_loss += loss.item()\n            metric(outputs.argmax(1), labels[:, 0])\n            epoch_val_score = metric.aggregate()[0].item()\n            \n            progress_bar.set_postfix({\n                'val_loss': epoch_val_loss / step,\n                'val_score': epoch_val_score,\n            })\n        metric.reset()\n            \n    if epoch_val_score > best_score:\n        best_score = epoch_val_score\n        torch.save(model.state_dict(), 'best_model_segmentation2d_dict.pth')\n        print('saved new best metric model')","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:08:56.986209Z","iopub.execute_input":"2023-09-02T16:08:56.986599Z","iopub.status.idle":"2023-09-02T16:44:31.858233Z","shell.execute_reply.started":"2023-09-02T16:08:56.986546Z","shell.execute_reply":"2023-09-02T16:44:31.856965Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"state_dict = torch.load('/kaggle/working/best_model_segmentation2d_dict.pth', map_location=device)\nmodel.load_state_dict(state_dict)","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:45:19.999701Z","iopub.execute_input":"2023-09-02T16:45:20.000083Z","iopub.status.idle":"2023-09-02T16:45:20.043234Z","shell.execute_reply.started":"2023-09-02T16:45:20.000051Z","shell.execute_reply":"2023-09-02T16:45:20.042163Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(val_loader))\ninputs, labels = batch['image'].to(device), batch['mask'].to(device)\nmodel.eval()\n\nwith torch.no_grad():\n    outputs = model(inputs)\n    \npreds = outputs.argmax(1)","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:45:23.754429Z","iopub.execute_input":"2023-09-02T16:45:23.755164Z","iopub.status.idle":"2023-09-02T16:45:24.516039Z","shell.execute_reply.started":"2023-09-02T16:45:23.755126Z","shell.execute_reply":"2023-09-02T16:45:24.514809Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = []\n\nfor image, label in zip(inputs, preds):\n    image, label = image.cpu(), label.cpu()\n    ret = blend_images(image, label=label.unsqueeze(0), alpha=0.5, cmap='hsv')\n    results.append(ret)\n\nfor image, label in zip(inputs, labels):\n    image, label = image.cpu(), label.cpu()\n    ret = blend_images(image, label=label, alpha=0.5, cmap='hsv')\n    results.append(ret)\n\nplt.figure(figsize=(10, 6))\nplt.imshow(make_grid(results, nrow=8).permute(1, 2, 0))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-02T16:45:26.123926Z","iopub.execute_input":"2023-09-02T16:45:26.124339Z","iopub.status.idle":"2023-09-02T16:45:28.532397Z","shell.execute_reply.started":"2023-09-02T16:45:26.124299Z","shell.execute_reply":"2023-09-02T16:45:28.531474Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}