{"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":"markdown","source":"## Libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nfrom torchvision import transforms, models\nfrom torchvision.datasets import ImageFolder\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch import nn\nfrom PIL import Image\nimport pydicom\nfrom sklearn.model_selection import train_test_split\nimport torchvision.models as modelst\nimport matplotlib.pyplot as plt\nimport timm","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.706072Z","iopub.execute_input":"2023-10-02T14:28:52.706418Z","iopub.status.idle":"2023-10-02T14:28:52.712442Z","shell.execute_reply.started":"2023-10-02T14:28:52.70639Z","shell.execute_reply":"2023-10-02T14:28:52.711548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class Config:\n    BASE_DIR = '/kaggle/input/rsna-atd-512x512-png-v2-dataset'\n    SEED = 12\n    IMAGE_SIZE = (224, 224)\n    BATCH_SIZE = 2\n    TARGET_COLUMNS = ['bowel_healthy', 'bowel_injury',\n                      'extravasation_healthy', 'extravasation_injury',\n                      'kidney_healthy', 'kidney_low', 'kidney_high',\n                      'liver_healthy', 'liver_low', 'liver_high',\n                      'spleen_healthy', 'spleen_low', 'spleen_high',\n                     ]\n    NUM_EPOCHS = 5\n    \nconfig = Config()","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.71424Z","iopub.execute_input":"2023-10-02T14:28:52.714853Z","iopub.status.idle":"2023-10-02T14:28:52.723906Z","shell.execute_reply.started":"2023-10-02T14:28:52.714822Z","shell.execute_reply":"2023-10-02T14:28:52.722931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(config.BASE_DIR, 'train.csv'))\ndf.head(5)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.725618Z","iopub.execute_input":"2023-10-02T14:28:52.726441Z","iopub.status.idle":"2023-10-02T14:28:52.780638Z","shell.execute_reply.started":"2023-10-02T14:28:52.726411Z","shell.execute_reply":"2023-10-02T14:28:52.779644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## train and validation set","metadata":{}},{"cell_type":"code","source":"def split_group(group, test_size=0.2):\n    if len(group) == 1:\n        return (group, pd.DataFrame()) if np.random.rand() < test_size else (pd.DataFrame(), group)\n    else:\n        return train_test_split(group, test_size=test_size, random_state=config.SEED)\n\ntrain_set = pd.DataFrame()\nvalidation_set = pd.DataFrame()\n\nfor _, group in df.groupby(config.TARGET_COLUMNS):\n    train_group, val_group = split_group(group)\n    train_set = pd.concat([train_set, train_group], ignore_index=True)\n    validation_set = pd.concat([validation_set, val_group], ignore_index=True)\n    \nprint(train_set.shape, validation_set.shape)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.78276Z","iopub.execute_input":"2023-10-02T14:28:52.783137Z","iopub.status.idle":"2023-10-02T14:28:52.797656Z","shell.execute_reply.started":"2023-10-02T14:28:52.783102Z","shell.execute_reply":"2023-10-02T14:28:52.796261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = train_set\ntrain_set.sample(3)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.79932Z","iopub.execute_input":"2023-10-02T14:28:52.800025Z","iopub.status.idle":"2023-10-02T14:28:52.822104Z","shell.execute_reply.started":"2023-10-02T14:28:52.799991Z","shell.execute_reply":"2023-10-02T14:28:52.821027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_set.sample(3)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.824394Z","iopub.execute_input":"2023-10-02T14:28:52.825401Z","iopub.status.idle":"2023-10-02T14:28:52.842704Z","shell.execute_reply.started":"2023-10-02T14:28:52.825368Z","shell.execute_reply":"2023-10-02T14:28:52.841629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset and Dataloader","metadata":{}},{"cell_type":"code","source":"train_transforms =  transforms.Compose([            \n        transforms.Resize(config.IMAGE_SIZE, interpolation=Image.NEAREST),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomAffine(degrees=0, shear=10),\n        transforms.RandomAffine(degrees=0, scale=(0.8, 1.2)),\n        transforms.ToTensor(),\n        transforms.Lambda(lambda x : x / 255),\n    ])\n\nval_transforms = transforms.Compose([\n        transforms.Resize(config.IMAGE_SIZE, interpolation=Image.NEAREST),\n        transforms.ToTensor(),\n        transforms.Lambda(lambda x : x / 255),\n    ])\n\nclass Dataset(Dataset):\n    def __init__(self, df, transform=None, labeled=True):\n        self.df = df\n        self.transform = transform\n        self.labeled = labeled\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        \n        image_path = self.df['image_path'][idx]\n        dicom_file = pydicom.dcmread(image_path)\n        pixel_array = dicom_file.pixel_array.astype(np.int16)\n        image = Image.fromarray(pixel_array)\n        image = self.transform(image)\n        \n        if self.labeled:\n            target = self.df[config.TARGET_COLUMNS].iloc[idx]\n            target = torch.tensor(target.values, dtype=torch.float32)\n            return image, target[:2], target[2:4], target[4:7], target[7:10], target[10:13]\n        else:\n            patient_id = self.df['patient_id'][idx]\n            return patient_id, image\n        \ntrain_set = Dataset(train_set, train_transforms)\nvalidation_set = Dataset(validation_set, val_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.844229Z","iopub.execute_input":"2023-10-02T14:28:52.844707Z","iopub.status.idle":"2023-10-02T14:28:52.85691Z","shell.execute_reply.started":"2023-10-02T14:28:52.844673Z","shell.execute_reply":"2023-10-02T14:28:52.855716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(\n    train_set,\n    batch_size=config.BATCH_SIZE,\n    shuffle=True\n)\n\nval_dataloader = DataLoader(\n    validation_set,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.858381Z","iopub.execute_input":"2023-10-02T14:28:52.859242Z","iopub.status.idle":"2023-10-02T14:28:52.871283Z","shell.execute_reply.started":"2023-10-02T14:28:52.859182Z","shell.execute_reply":"2023-10-02T14:28:52.870313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self):\n        super(Model, self).__init__()\n        \n        base = timm.create_model('tf_efficientnet_b8',\n                                      checkpoint_path='/kaggle/input/tf-efficientnet/pytorch/tf-efficientnet-b8/1/tf_efficientnet_b8_ra-572d5dd9.pth')\n        conv_stem_weights = base.state_dict()['conv_stem.weight'].sum(dim=1, keepdim=True)\n        base.conv_stem = nn.Conv2d(1, 72, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        base.state_dict()['conv_stem.weight'] = conv_stem_weights\n        self.base = base\n        \n        self.hidden1 = torch.nn.Linear(1000, 512)\n        self.bn1 = torch.nn.BatchNorm1d(512)\n        self.hidden2 = torch.nn.Linear(512, 256)\n        self.bn2 = torch.nn.BatchNorm1d(256)\n        self.hidden3 = torch.nn.Linear(256, 128)\n        self.bn3 = torch.nn.BatchNorm1d(128)\n        \n        self.out_bowel = nn.Linear(128, 2)\n        self.out_extravasation = nn.Linear(128, 2)\n        self.out_kidney = nn.Linear(128, 3)\n        self.out_liver = nn.Linear(128, 3)\n        self.out_spleen = nn.Linear(128, 3)\n        \n        self.dropout = nn.Dropout(0.5)        \n        \n    def forward(self, x):\n        x = self.base(x)\n        x = self.hidden1(x)\n        x = self.dropout(x)\n        x = self.bn1(x)\n        x = self.hidden2(x)\n        x = self.dropout(x)\n        x = self.bn2(x)\n        x = self.hidden3(x)\n        x = self.dropout(x)\n        x = self.bn3(x)\n        \n        return self.out_bowel(x), self.out_extravasation(x), self.out_kidney(x), self.out_liver(x), self.out_spleen(x)\n    \nmodel = Model()","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:52.874035Z","iopub.execute_input":"2023-10-02T14:28:52.874845Z","iopub.status.idle":"2023-10-02T14:28:54.42981Z","shell.execute_reply.started":"2023-10-02T14:28:52.874812Z","shell.execute_reply":"2023-10-02T14:28:54.428828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = torch.optim.Adam(model.parameters())\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=5, factor=0.5, verbose=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:54.431237Z","iopub.execute_input":"2023-10-02T14:28:54.431835Z","iopub.status.idle":"2023-10-02T14:28:54.578041Z","shell.execute_reply.started":"2023-10-02T14:28:54.431792Z","shell.execute_reply":"2023-10-02T14:28:54.577003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"train_losses = []\nval_losses = []\n\nfor epoch in range(config.NUM_EPOCHS):\n        \n        print(f'EPOCH: {epoch + 1}/{config.NUM_EPOCHS}')\n        \n        train_loss = 0.0\n        validation_loss = 0.0\n        \n        model.train()\n        for batch_num, data in enumerate(train_dataloader):\n            inputs, labels_b, labels_e, labels_k, labels_l, labels_s = data\n            inputs = inputs.to(device)\n            labels_b = labels_b.to(device)\n            labels_e = labels_e.to(device)\n            labels_k = labels_k.to(device)\n            labels_l = labels_l.to(device)\n            labels_s = labels_s.to(device)\n            \n            optimizer.zero_grad()\n\n            out_b, out_e, out_k, out_l, out_s = model(inputs)\n\n            loss_b = criterion(out_b, labels_b)\n            loss_e = criterion(out_e, labels_e)\n            loss_k = criterion(out_k, labels_k)\n            loss_l = criterion(out_l, labels_l)\n            loss_s = criterion(out_s, labels_s)\n            \n            total_loss = loss_b + loss_e + loss_k + loss_l + loss_s\n            total_loss.backward()\n            \n            optimizer.step()\n            \n            train_loss += total_loss.item()\n            \n        train_loss = train_loss/len(train_dataloader)\n        train_losses.append(train_loss)\n        print(f'train loss: {train_loss}')\n        \n        model.eval()\n        with torch.no_grad():\n            for batch_num, data in enumerate(val_dataloader):\n                inputs, labels_b, labels_e, labels_k, labels_l, labels_s = data\n                inputs = inputs.to(device)\n                labels_b = labels_b.to(device)\n                labels_e = labels_e.to(device)\n                labels_k = labels_k.to(device)\n                labels_l = labels_l.to(device)\n                labels_s = labels_s.to(device)\n                \n                out_b, out_e, out_k, out_l, out_s = model(inputs)\n\n                loss_b = criterion(out_b, labels_b)\n                loss_e = criterion(out_e, labels_e)\n                loss_k = criterion(out_k, labels_k)\n                loss_l = criterion(out_l, labels_l)\n                loss_s = criterion(out_s, labels_s)\n\n                total_loss = loss_b + loss_e + loss_k + loss_l + loss_s\n                validation_loss += total_loss.item()\n                \n        validation_loss = validation_loss/len(val_dataloader)\n        val_losses.append(validation_loss)\n        print(f'validation loss: {validation_loss}')\n                ","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:28:54.580012Z","iopub.execute_input":"2023-10-02T14:28:54.580568Z","iopub.status.idle":"2023-10-02T14:29:02.620519Z","shell.execute_reply.started":"2023-10-02T14:28:54.580536Z","shell.execute_reply":"2023-10-02T14:29:02.619541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training curves","metadata":{}},{"cell_type":"code","source":"plt.subplot(1, 2, 1)\nepochs = range(1, config.NUM_EPOCHS + 1)\nplt.plot(epochs, train_losses, label='Training loss')\nplt.plot(epochs, val_losses, label='Validation loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:29:02.621735Z","iopub.execute_input":"2023-10-02T14:29:02.622547Z","iopub.status.idle":"2023-10-02T14:29:02.840808Z","shell.execute_reply.started":"2023-10-02T14:29:02.622512Z","shell.execute_reply":"2023-10-02T14:29:02.839935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predictions","metadata":{}},{"cell_type":"code","source":"test_set = pd.read_csv(os.path.join(config.BASE_DIR, 'test.csv'))\ntest_set","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:29:02.842171Z","iopub.execute_input":"2023-10-02T14:29:02.842707Z","iopub.status.idle":"2023-10-02T14:29:02.855828Z","shell.execute_reply.started":"2023-10-02T14:29:02.842671Z","shell.execute_reply":"2023-10-02T14:29:02.854916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_set = Dataset(test_set, transform=val_transforms, labeled=False)\n\ntest_dataloader = DataLoader(\n    test_set,\n    batch_size=1,\n    shuffle=False\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:29:02.857095Z","iopub.execute_input":"2023-10-02T14:29:02.857616Z","iopub.status.idle":"2023-10-02T14:29:02.862547Z","shell.execute_reply.started":"2023-10-02T14:29:02.857588Z","shell.execute_reply":"2023-10-02T14:29:02.861495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nindex = []\n\nfor patient_id, image in test_dataloader:\n    image = image.to(device)\n    pred_row = []\n    for out in model(image):\n        out = torch.nn.Softmax()(out.cpu())\n        pred_row.extend(out[0].tolist())\n    print(pred_row)\n    predictions.append(pred_row)\n    index.append(patient_id.item())\n    \npredictions = pd.DataFrame(predictions, index=index, columns=config.TARGET_COLUMNS)\npredictions = predictions.rename_axis('patient_id')\npredictions.sort_index(inplace=True)\npredictions","metadata":{"execution":{"iopub.status.busy":"2023-10-02T14:29:02.863881Z","iopub.execute_input":"2023-10-02T14:29:02.864236Z","iopub.status.idle":"2023-10-02T14:29:02.967083Z","shell.execute_reply.started":"2023-10-02T14:29:02.864204Z","shell.execute_reply":"2023-10-02T14:29:02.965733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions.to_csv('submission.csv', index=True)","metadata":{},"execution_count":null,"outputs":[]}]}