{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### By: Justin Cheigh\n---\n- [GitHub](https://github.com/jcheigh)\n- [LinkedIn](https://www.linkedin.com/in/justin-cheigh/)\n- [Medium (blog)](https://medium.com/@jcheigh)\n---","metadata":{}},{"cell_type":"code","source":"import os \nimport time\nimport random\n\nfrom abc import ABC, abstractmethod, abstractproperty\nfrom collections import defaultdict\nfrom tqdm import tqdm \nfrom PIL import Image\n\nimport wandb \n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport sklearn\nfrom sklearn.model_selection import train_test_split, StratifiedGroupKFold\nfrom sklearn.metrics import accuracy_score \n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport torchvision.models as models\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nfrom timm import create_model\n\nimport lightning.pytorch as pl\nfrom lightning.pytorch.callbacks.early_stopping import EarlyStopping\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:23.074301Z","iopub.execute_input":"2023-09-04T15:08:23.075392Z","iopub.status.idle":"2023-09-04T15:08:23.085507Z","shell.execute_reply.started":"2023-09-04T15:08:23.075357Z","shell.execute_reply":"2023-09-04T15:08:23.084242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"libraries = [\n    \"numpy\", \n    \"pandas\", \n    \"sklearn\", \n    \"torch\", \n    \"wandb\",\n    \"lightning\"\n    ]\n\nversions = [\n    np.__version__, \n    pd.__version__,\n    sklearn.__version__,\n    torch.__version__,\n    wandb.__version__,\n    pl.__version__\n    ]\n\n\nfor lib, ver in zip(libraries, versions):\n    print(f\"{lib} {ver}\")\n!python --version","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:24.217886Z","iopub.execute_input":"2023-09-04T15:08:24.21827Z","iopub.status.idle":"2023-09-04T15:08:25.17889Z","shell.execute_reply.started":"2023-09-04T15:08:24.218237Z","shell.execute_reply":"2023-09-04T15:08:25.177667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def timeit(func):\n    \"decorator to time functions\"\n    def wrapper(*args, **kwargs):\n        start_time = time.time()\n        result = func(*args, **kwargs)\n        end_time = time.time()\n        elapsed_time = end_time - start_time\n        print(f\"{func.__name__} executed in {elapsed_time:.3f} seconds\")\n        return result\n    return wrapper","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:25.549672Z","iopub.execute_input":"2023-09-04T15:08:25.550077Z","iopub.status.idle":"2023-09-04T15:08:25.557404Z","shell.execute_reply.started":"2023-09-04T15:08:25.550044Z","shell.execute_reply":"2023-09-04T15:08:25.555478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_PATH  = os.path.join('/kaggle', 'input')  # to read\nOUTPUT_PATH = os.path.join('/kaggle', 'output') # to write\n\nfor path in os.listdir(INPUT_PATH):\n    print(path)","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:27.276302Z","iopub.execute_input":"2023-09-04T15:08:27.276999Z","iopub.status.idle":"2023-09-04T15:08:27.283347Z","shell.execute_reply.started":"2023-09-04T15:08:27.276963Z","shell.execute_reply":"2023-09-04T15:08:27.282244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We have two input folders here. rsna-2023-abdominal-trauma-detection was the original data; the problem is that the images are in .dcm format. rsna-atd-512x512-png-v2-dataset was created in [this notebook](https://www.kaggle.com/code/awsaf49/rsna-atd-512x512-png-v2-data/notebook). The general idea is to merge the image_levels_labels.csv dataframe and the train.csv dataframe (creating a dataframe with a row for each image) and then use [this notebook](https://www.kaggle.com/code/radek1/how-to-process-dicom-images-to-pngs) to convert the images to .png format. ","metadata":{}},{"cell_type":"code","source":"DICOM_DATA  = 'rsna-2023-abdominal-trauma-detection'\nPNG_DATA    = 'rsna-atd-512x512-png-v2-dataset'\n\nfor path in os.listdir(f\"{INPUT_PATH}/{PNG_DATA}\"):\n    print(path)","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:28.989752Z","iopub.execute_input":"2023-09-04T15:08:28.990506Z","iopub.status.idle":"2023-09-04T15:08:29.015783Z","shell.execute_reply.started":"2023-09-04T15:08:28.990468Z","shell.execute_reply":"2023-09-04T15:08:29.014893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at train.csv first","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{INPUT_PATH}/{PNG_DATA}/train.csv\")\ndf = df.drop_duplicates()\ndf = df.drop([\"series_id\", \"instance_number\"], axis = 1)\n\ndef get_img_path(path):\n    path = path.replace('.dcm', '.png')\n    return path.replace(DICOM_DATA, PNG_DATA)\n\ndf[\"image_path\"] = df[\"image_path\"].apply(lambda path : get_img_path(path))","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:30.189827Z","iopub.execute_input":"2023-09-04T15:08:30.190193Z","iopub.status.idle":"2023-09-04T15:08:30.318889Z","shell.execute_reply.started":"2023-09-04T15:08:30.190163Z","shell.execute_reply":"2023-09-04T15:08:30.317854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"train.csv contains target labels for the train set. Note that patients labeled healthy may still have other medical issues, such as cancer or broken bones, that don't happen to be covered by the competition labels.\n\nColumn Description: <br>\n- patient_id: A unique ID code for each patient\n- (bowel/extravasation)_(healthy/injury): two injury types with binary targets\n- (kidney/liver/spleen)_(healthy/low/high): -three injury types with three target levels\n- any_injury: whether patient had any injury at all\n- injury_name: name of injury\n- image_path: path to .png of ct scan\n- width/height: width/height of .png","metadata":{}},{"cell_type":"markdown","source":"\nAuto EDA with dataprep\n\n!pip install dataprep <br>\nfrom dataprep.eda import create_report <br>\ncreate_report(df)\n\nBasically just showed there's some class imbalance with multiclass labels\n","metadata":{"_kg_hide-output":false,"execution":{"iopub.status.busy":"2023-09-04T15:02:02.794458Z","iopub.execute_input":"2023-09-04T15:02:02.79485Z","iopub.status.idle":"2023-09-04T15:02:02.805282Z","shell.execute_reply.started":"2023-09-04T15:02:02.794818Z","shell.execute_reply":"2023-09-04T15:02:02.804216Z"}}},{"cell_type":"markdown","source":"The following block of code is straight from [this notebook.](https://www.kaggle.com/code/ayushs9020/understanding-the-competition-rsna)","metadata":{}},{"cell_type":"code","source":"train_csv = pd.read_csv(f\"{INPUT_PATH}/{DICOM_DATA}/train.csv\")\norgan_columns = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n\norgan_counts = pd.DataFrame()\norgan_counts['Organ'] = train_csv.columns[1:]\norgan_counts[\"count\"] = [0 for _ in range(organ_counts.shape[0])]\nfor index , column in enumerate(train_csv.columns[1:]):\n    organ_counts['count'][index] = train_csv[column].sum()\n    \nplt.figure(figsize=(10, 3))\nsns.barplot(data=organ_counts.sort_values(by=['count']), x='Organ', y='count')\nplt.xticks(rotation=90)\nplt.title(\"Distribution of Injury\")\nplt.xlabel(\"Injury --->\")\nplt.ylabel(\"Count --->\")\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:32.225714Z","iopub.execute_input":"2023-09-04T15:08:32.226094Z","iopub.status.idle":"2023-09-04T15:08:32.645459Z","shell.execute_reply.started":"2023-09-04T15:08:32.226062Z","shell.execute_reply":"2023-09-04T15:08:32.644459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Weights & Biases Login \ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key      = user_secrets.get_secret(\"WANDB\")\n\n    wandb.login(key=api_key)\n    del api_key\nexcept:\n    raise Exception(\"Your API key must remain anonymous\")","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:35.254987Z","iopub.execute_input":"2023-09-04T15:08:35.255384Z","iopub.status.idle":"2023-09-04T15:08:38.526162Z","shell.execute_reply.started":"2023-09-04T15:08:35.255344Z","shell.execute_reply":"2023-09-04T15:08:38.52518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using {'GPU' if 'cuda' in str(device) else 'CPU'}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:40.31644Z","iopub.execute_input":"2023-09-04T15:08:40.317262Z","iopub.status.idle":"2023-09-04T15:08:40.333082Z","shell.execute_reply.started":"2023-09-04T15:08:40.317225Z","shell.execute_reply":"2023-09-04T15:08:40.331982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    COMPETITION   = 'rsna-atd' \n    COMMENT       = 'Base ResNet18'\n    EXP_NAME      = \"baseline: ResNet18 + multihead\"   \n    \n    DEVICE        = \"CPU\" if device.type == 'cpu' else \"GPU\"   \n    MODEL_NAME    = \"ResNet18\"\n    \n    SEED          = 42                   \n    NUM_FOLDS     = 4           \n    IMG_SIZE      = (256, 256)  \n    BATCH_SIZE    = 64          \n    EPOCHS        = 10 \n    NUM_WORKERS   =  2\n    TEST_SIZE     = .1\n    PIN_MEMORY    = True  \n    LOSS          = \"Cross Entropy\"\n    OPTIMIZER     = \"SGD\"\n    SCHEDULER     = \"StepLR\"\n    \n    TARGET_COLS   = [\n                    \"bowel_injury\", \"extravasation_injury\",\n                    \"kidney_healthy\", \"kidney_low\", \"kidney_high\",\n                    \"liver_healthy\", \"liver_low\", \"liver_high\",\n                    \"spleen_healthy\", \"spleen_low\", \"spleen_high\",\n                    ]","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:43.310605Z","iopub.execute_input":"2023-09-04T15:08:43.311638Z","iopub.status.idle":"2023-09-04T15:08:43.319287Z","shell.execute_reply.started":"2023-09-04T15:08:43.311575Z","shell.execute_reply":"2023-09-04T15:08:43.318221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=Config.SEED):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n\nset_seed()","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:44.459163Z","iopub.execute_input":"2023-09-04T15:08:44.46045Z","iopub.status.idle":"2023-09-04T15:08:44.470943Z","shell.execute_reply.started":"2023-09-04T15:08:44.460404Z","shell.execute_reply":"2023-09-04T15:08:44.469889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Data(Dataset):\n    \"Basic custom data class\"\n    \n    def __init__(self, img_paths, labels, transform):\n        self.paths     = img_paths\n        self.labels    = labels\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        img_path = self.paths[idx]\n        image    = Image.open(img_path)\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        label    = self.labels[idx]\n        labels   = (\n            label[0:1], # bowel\n            label[1:2], # extravasation\n            label[2:5], # kidney\n            label[5:8], # liver\n            label[8:11] # spleen\n            )\n        \n        return image, labels","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:50.196274Z","iopub.execute_input":"2023-09-04T15:08:50.19669Z","iopub.status.idle":"2023-09-04T15:08:50.204663Z","shell.execute_reply.started":"2023-09-04T15:08:50.196657Z","shell.execute_reply":"2023-09-04T15:08:50.20372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transform(phase):\n    \"Returns list of transforms based on phase (train/test)\"\n    if phase == 'train':\n        transform_lst = [\n            transforms.RandomHorizontalFlip(p=.25),\n            transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1)\n            ]\n    else:\n        transform_lst = []\n    \n    transform_lst += [\n            transforms.Resize(size=Config.IMG_SIZE),\n            transforms.ToTensor(),\n            transforms.Normalize([0.445], [0.269])  \n            ]\n    \n    transform  = transforms.Compose(transform_lst)\n    return transform\n\ndef get_dataloader(dataframe, phase='train'):\n    paths     = dataframe.image_path.to_list()\n    labels    = dataframe[Config.TARGET_COLS].values\n    transform = get_transform(phase)\n    data      = Data(\n                    img_paths   = paths, \n                    labels      = labels,\n                    transform   = transform\n                    )\n    \n    return DataLoader(\n                dataset     = data,\n                batch_size  = Config.BATCH_SIZE,\n                shuffle     = True,\n                num_workers = Config.NUM_WORKERS,\n                pin_memory  = Config.PIN_MEMORY,\n                )","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:51.185281Z","iopub.execute_input":"2023-09-04T15:08:51.186007Z","iopub.status.idle":"2023-09-04T15:08:51.195076Z","shell.execute_reply.started":"2023-09-04T15:08:51.185971Z","shell.execute_reply":"2023-09-04T15:08:51.193822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stratify(df=df, num_folds=Config.NUM_FOLDS):\n    \"\"\"\n    Adds new col df[\"fold\"] that's the group for that row\n    \n    We use stra\n    \"\"\"\n    df['stratify'] = ''\n    for col in Config.TARGET_COLS:\n        df['stratify'] += df[col].astype(str)\n\n    df  = df.reset_index(drop=True)\n    skf = StratifiedGroupKFold(\n                n_splits     = num_folds, \n                shuffle      = True,\n                random_state = Config.SEED\n                )\n\n    for fold, (train_idxs, val_idxs) in enumerate(skf.split(df, df['stratify'], df[\"patient_id\"])):\n        df.loc[val_idxs, 'fold'] = fold\n\n    return df\n\ndef kfold_iter(df, k=Config.NUM_FOLDS):\n    \"\"\"\n    Returns an iterator of train, test dataloaders for kfold validation\n    \n    Usage:\n        dl_iter = kfold_iter(df, k)\n\n        for _ range(k):\n            train_dl, test_dl = next(dl_iter)\n            ...\n    \n    Use an iterator for memory reasons\n    \"\"\"\n    \n    df = stratify(df, k)\n\n    def get_dataloaders(df, test_grp):\n        train = df[df['fold']  != test_grp]\n        test  = df[df['fold']  == test_grp]\n\n        train_dl   = get_dataloader(train, phase='train')\n        test_dl    = get_dataloader(test, phase='test')\n\n        return train_dl, test_dl\n    \n    return iter((get_dataloaders(df, i) for i in range(k)))\n","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:52.016201Z","iopub.execute_input":"2023-09-04T15:08:52.016919Z","iopub.status.idle":"2023-09-04T15:08:52.028686Z","shell.execute_reply.started":"2023-09-04T15:08:52.016883Z","shell.execute_reply":"2023-09-04T15:08:52.02739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_iter = kfold_iter(df)\ntrain_dl, test_dl = next(dl_iter)\n\nimages, labels = next(iter(train_dl))\n# B x C x H x W = 64 x 1 x 256 x 256\nprint(f\"Image Shape: {images.shape}\\n\")\nprint(f'Label Shape:')\n\nfor label in labels:\n    print(label.shape)","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:08:52.829946Z","iopub.execute_input":"2023-09-04T15:08:52.830926Z","iopub.status.idle":"2023-09-04T15:09:01.543968Z","shell.execute_reply.started":"2023-09-04T15:08:52.830884Z","shell.execute_reply":"2023-09-04T15:09:01.542071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Images is a batch of images, where each image has 1 color channel and has image size 256x256. Thus, images is a tensor with shape B x C x H x W = 64 x 1 x 256 x 256. Each individual label is a list of 1 dimensional tensors (binary for extra/bowel and 3 outputs otherwise). Thus, labels is a list of 5 tensors.","metadata":{}},{"cell_type":"code","source":"def plot_imgs(images, size=4):\n    fig, axis = plt.subplots(4, 4, figsize = (15, 10))\n    for i, ax in enumerate(axis.flat):\n        with torch.no_grad():\n            img = images[i].numpy()\n            img = np.transpose(img, (1, 2, 0)) # torch/numpy convention\n            ax.imshow(img)\n            ax.axis('off')\n            \nplot_imgs(images)","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:09:03.11652Z","iopub.execute_input":"2023-09-04T15:09:03.116901Z","iopub.status.idle":"2023-09-04T15:09:03.978803Z","shell.execute_reply.started":"2023-09-04T15:09:03.116867Z","shell.execute_reply":"2023-09-04T15:09:03.977794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LitMultiHeadModel(pl.LightningModule, ABC): \n    \"\"\"\n    General MultiHeadedModel\n    \n    Forward pass involves passing though main layers \n    (self.backbone) and then using Global Average \n    Pooling (i.e. get activation maps) then dense \n    layers to 5 different heads.\n    \n    To Fill Out:\n        self.define_backbone() should return a nn.Sequential\n        self.hidden_dim specifies the output dim of self.backbone\n    \"\"\"\n\n    def __init__(self):\n        super().__init__()    \n        self.train_log = defaultdict(float)\n        self.val_log   = defaultdict(float)\n        self.backbone  = self.define_backbone()\n        \n        # Global Average Pooling\n        self.global_avg_pool = nn.AdaptiveAvgPool2d((1, 1))\n        \n        # Break into 5 heads \n        self.fc_bowel  = nn.Linear(self.hidden_dim, 32)\n        self.fc_extra  = nn.Linear(self.hidden_dim, 32)\n        self.fc_liver  = nn.Linear(self.hidden_dim, 32)\n        self.fc_kidney = nn.Linear(self.hidden_dim, 32)\n        self.fc_spleen = nn.Linear(self.hidden_dim, 32)\n        \n        # Prediction heads \n        self.out_bowel  = nn.Linear(32, 1)\n        self.out_extra  = nn.Linear(32, 1)\n        self.out_liver  = nn.Linear(32, 3)\n        self.out_kidney = nn.Linear(32, 3)\n        self.out_spleen = nn.Linear(32, 3)       \n    \n    @abstractmethod\n    def define_backbone(self) -> nn.Sequential:\n        raise NotImplementedError\n    \n    @abstractproperty\n    def hidden_dim(self) -> int:\n        raise NotImplementedError\n\n    def init_wandb(self, fold):\n        return None\n        config = {var : val for var, val in dict(vars(Config)).items() if '__' not in var}\n        config |= {\"fold\" : int(fold)} # use .update if python isn't new enough\n        run    = wandb.init(\n                    project   = \"rsna-atd-public\",\n                    name      = f\"fold-{fold}|dim-{Config.IMG_SIZE[0]}x{Config.IMG_SIZE[1]}|model-{Config.MODEL_NAME}\",\n                    config    = config,\n                    anonymous = None,\n                    group     = Config.EXP_NAME\n                    )\n        return run\n    \n    def criterion(self, outputs, labels, method=\"cross entropy\"):\n        if method == 'cross entropy':\n            # define loss for each head- for now just sum of BCE & CCE\n            loss_fn_bowel  = nn.BCELoss()           \n            loss_fn_extra  = nn.BCELoss()\n            loss_fn_liver  = nn.CrossEntropyLoss()\n            loss_fn_kidney = nn.CrossEntropyLoss()\n            loss_fn_spleen = nn.CrossEntropyLoss()\n\n            labels = [label.float() for label in labels]\n\n            # compute loss for each head\n            loss_bowel  = loss_fn_bowel(outputs[0], labels[0])\n            loss_extra  = loss_fn_extra(outputs[1], labels[1])\n            loss_liver  = loss_fn_liver(outputs[2], labels[2])\n            loss_kidney = loss_fn_kidney(outputs[3], labels[3])\n            loss_spleen = loss_fn_spleen(outputs[4], labels[4])\n\n            return loss_bowel + loss_extra + loss_liver + loss_kidney + loss_spleen\n        \n        elif method == 'weighted cross entropy':\n            raise NotImplementedError\n            \n        raise NotImplementedError \n\n    def update_acc_log(self, outputs, labels, phase):\n        if phase == \"Train\":\n            acc_log = self.train_log\n        else:\n            acc_log = self.val_log\n            \n        get_preds  = lambda i: (outputs[i] > .5).float().squeeze().cpu().numpy()\n        get_labels = lambda i: labels[i].squeeze().cpu().numpy()\n\n        acc_log[f\"{phase} Bowel Accuracy\"]  += accuracy_score(get_preds(0), get_labels(0))\n        acc_log[f\"{phase} Extra Accuracy\"]  += accuracy_score(get_preds(1), get_labels(1)) \n        acc_log[f\"{phase} Liver Accuracy\"]  += accuracy_score(get_preds(2), get_labels(2))\n        acc_log[f\"{phase} Kidney Accuracy\"] += accuracy_score(get_preds(3), get_labels(3))\n        acc_log[f\"{phase} Spleen Accuracy\"] += accuracy_score(get_preds(4), get_labels(4))\n        acc_log[f\"{phase} Avg. Accuracy\"]   += round(np.mean(list(acc_log.values())), 3)\n        \n        return acc_log\n\n    def show_stats(self, train_dl, val_dl):\n        def show(phase):\n            if phase == \"Train\":\n                acc_log  = self.train_log\n                data_len = len(train_dl)\n            else:\n                acc_log  = self.val_log \n                data_len = len(val_dl)\n            \n            to_acc = lambda total_acc : round(float((total_acc / data_len * 100), 3))\n            log    = {head : to_acc(acc) for head, acc in acc_log.items()}\n                                              \n            print('-' * 20)\n            for head, accuracy in log.items():\n                print(f\"{head}: {accuracy}\")\n            print('-' * 20)\n            \n            #wandb.log(log) \n                                              \n        print(f\"Train Stats:\")\n        show(\"Train\")\n        print(f\"Validation Stats:\")\n        show(\"Val\")\n        \n    def finish_run(self):\n        wandb.finish()\n        \n    def forward(self, x):\n        # Backbone + GAP\n        x = self.backbone(x)\n        x = self.global_avg_pool(x)              \n        x = x.view(x.size(0), -1)              \n        \n        # Split into 5 Heads\n        # SiLU is ReLU * x is slower but better convergence/stability\n        x_bowel  = nn.SiLU()(self.fc_bowel(x)) \n        x_extra  = nn.SiLU()(self.fc_extra(x))\n        x_liver  = nn.SiLU()(self.fc_liver(x))\n        x_kidney = nn.SiLU()(self.fc_kidney(x))\n        x_spleen = nn.SiLU()(self.fc_spleen(x))\n        \n        # Prediction heads\n        out_bowel  = torch.sigmoid(self.out_bowel(x_bowel))\n        out_extra  = torch.sigmoid(self.out_extra(x_extra))\n        out_liver  = nn.Softmax(dim=1)(self.out_liver(x_liver))\n        out_kidney = nn.Softmax(dim=1)(self.out_kidney(x_kidney))\n        out_spleen = nn.Softmax(dim=1)(self.out_spleen(x_spleen))\n        \n        return [out_bowel, out_extra, out_liver, out_kidney, out_spleen]\n\n    def training_step(self, batch, batch_idx):\n        x, y    = batch\n        y_hat   = self(x)\n        loss    = self.criterion(y_hat, y)\n        acc_log = self.update_acc_log(y_hat, y, \"Train\")\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        # this is the val loop\n        x, y    = batch\n        y_hat   = self(x)\n        loss    = self.criterion(y_hat, y)\n        acc_log = self.update_acc_log(y_hat, y, \"Val\")\n        return loss \n    \n    def configure_optimizers(self):\n        def get_optimizer():\n            if Config.OPTIMIZER == \"Adam\":\n                return optim.Adam(self.parameters(), lr=1e-4)\n            \n            elif Config.OPTIMIZER == \"SGD\":\n                return optim.SGD(self.parameters(), lr=0.001, momentum=0.9)\n\n            raise NotImplementedError\n            \n        def get_scheduler(optimizer):\n            if Config.SCHEDULER == \"StepLR\":\n                return lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n            \n            raise NotImplementedError\n                \n        optimizer = get_optimizer()\n        scheduler = get_scheduler(optimizer)\n        return {'optimizer': optimizer, 'lr_scheduler': scheduler}","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:11:47.638905Z","iopub.execute_input":"2023-09-04T15:11:47.639267Z","iopub.status.idle":"2023-09-04T15:11:47.673975Z","shell.execute_reply.started":"2023-09-04T15:11:47.639236Z","shell.execute_reply":"2023-09-04T15:11:47.67252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BasicCNN(LitMultiHeadModel):\n    \"\"\"\n    Basic Convolutional Neural Network\n        2 stacks of conv to relu to max pool\n        1 final conv layer\n        then multihead model\n    \"\"\"\n    \n    hidden_dim = 64\n    \n    def define_backbone(self):\n        \n        def create_conv(in_channels, out_channels):\n            return nn.Conv2d(\n                    in_channels  = in_channels,\n                    out_channels = out_channels,\n                    kernel_size  = 3,\n                    padding      = 1,\n                    stride       = 2\n                    )\n        \n        def create_conv_stack(in_channels, out_channels):\n            relu = nn.ReLU()\n            conv = create_conv(in_channels, out_channels)\n            pool = nn.MaxPool2d(2)\n            return [conv, relu, pool]\n        \n        conv1 = create_conv_stack(1, 16)\n        conv2 = create_conv_stack(16, 32)\n        conv3 = create_conv(32, 64)\n        \n        return nn.Sequential(*conv1, *conv2, conv3)","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:11:48.173693Z","iopub.execute_input":"2023-09-04T15:11:48.174081Z","iopub.status.idle":"2023-09-04T15:11:48.182491Z","shell.execute_reply.started":"2023-09-04T15:11:48.174052Z","shell.execute_reply":"2023-09-04T15:11:48.181386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet(LitMultiHeadModel):\n    \"\"\"\n    Basic ResNet18\n    \n    Takes pretrained ResNet18 but changes the first layer\n    (bc only one color channel here)\n    \"\"\"\n    hidden_dim = 512\n    \n    def define_backbone(self): \n        # Load pre-trained ResNet model + higher level layers\n        resnet = models.resnet18(weights='IMAGENET1K_V1')\n\n        # Remove classification head and GAP layer\n        # resnet.children() returns iter of nn.Module, unpack with * then nn.Sequential\n        main_layers = nn.Sequential(*list(resnet.children())[:-2])\n\n        # Change first layer for 1 color channel\n        # Weights initialized as avg of original weights\n        original_weights = resnet.conv1.weight.clone()\n        new_conv1 = nn.Conv2d(\n                        in_channels  = 1, \n                        out_channels = 64, \n                        kernel_size  = 7,\n                        stride       = 2,\n                        padding      = 3,\n                        bias         = False\n                        )\n        new_conv1.weight.data = original_weights.sum(dim=1, keepdim=True) / 3.0\n        main_layers[0] = new_conv1\n\n        return main_layers","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:11:48.724669Z","iopub.execute_input":"2023-09-04T15:11:48.725034Z","iopub.status.idle":"2023-09-04T15:11:48.732795Z","shell.execute_reply.started":"2023-09-04T15:11:48.725002Z","shell.execute_reply":"2023-09-04T15:11:48.731821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EfficientNet(LitMultiHeadModel):\n\n    hidden_dim = 1792  \n    \n    def define_backbone(self):\n        eff_net  = create_model('tf_efficientnet_b4_ns', pretrained = True)\n        features = nn.Sequential(*list(eff_net.children())[:-2])\n        \n        # Change the first layer for 1 color channel\n        original_conv1 = features[0]\n        new_conv1 = nn.Conv2d(\n                        in_channels  = 1,\n                        out_channels = original_conv1.out_channels,\n                        kernel_size  = original_conv1.kernel_size,\n                        stride       = original_conv1.stride,\n                        padding      = original_conv1.padding,\n                        bias         = False\n                    )\n        new_conv1.weight.data = original_conv1.weight.sum(dim=1, keepdim=True) / 3.0\n        features[0] = new_conv1\n        \n        return features","metadata":{"execution":{"iopub.status.busy":"2023-09-04T15:11:49.544757Z","iopub.execute_input":"2023-09-04T15:11:49.545135Z","iopub.status.idle":"2023-09-04T15:11:49.552382Z","shell.execute_reply.started":"2023-09-04T15:11:49.545105Z","shell.execute_reply":"2023-09-04T15:11:49.551178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model_class, epochs=Config.EPOCHS, num_folds=Config.NUM_FOLDS, early_stop=True):\n    \"\"\"\n    Perform k-fold cross-validation on a PyTorch Lightning model\n    \"\"\"\n    data_iter = kfold_iter(df)\n    for fold in range(num_folds):\n        print(f\"======= Training Fold {fold} =======\")\n        model             = model_class()\n        train_dl, val_dl  = next(data_iter)\n        run               = model.init_wandb(fold)\n    \n        trainer = pl.Trainer(\n                max_epochs  = epochs,\n                #callbacks  = [EarlyStopping(monitor='val_loss', patience=3)],\n                devices     =\"auto\",\n                accelerator =\"auto\"\n                )\n        trainer.fit(model, train_dl, val_dl)\n        model.show_stats()\n        model.finish_run()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train(ResNet)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Non-Lightning Code (older)","metadata":{}},{"cell_type":"code","source":"@timeit\ndef run_epoch(model, dataloader, optimizer, phase, device):\n    if phase == 'train':\n        model.train()\n    else:\n        model.eval()\n        \n    running_loss = 0.0\n    accuracy_log = defaultdict(float)\n    \n    for inputs, labels in dataloader:\n        inputs = inputs.to(device)\n        labels = [label.to(device) for label in labels]\n        \n        optimizer.zero_grad() if optimizer else None\n        \n        with torch.set_grad_enabled(phase == 'train'):\n            outputs = model(inputs)                     # forward pass\n            loss = criterion(outputs, labels)           # compute loss\n               \n            if phase == \"train\":\n                loss.backward()                         # autograd \n                optimizer.step()                        # backprop step\n            \n            running_loss += float(loss)                 # update running loss\n            update_log(accuracy_log, outputs, labels)   # update accuracy log\n            \n    epoch_loss    = running_loss / len(dataloader)\n    to_acc = lambda total_acc : round(float((total_acc / len(dataloader)) * 100), 3)\n    epoch_acc_log = {head : to_acc(total_acc) for head, total_acc in accuracy_log.items()}\n    \n    return epoch_loss, epoch_acc_log","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### No KFold\n@timeit\ndef train_model(model, optimizer, scheduler, device, train_dl, test_dl, num_epochs=config.EPOCHS):\n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch}/{num_epochs - 1}')\n        print('=' * 20)\n        \n        loss, acc_log = run_epoch(\n                            model = model,\n                            dataloader = train_dl,\n                            optimizer = optimizer,\n                            phase = 'train',\n                            device = device\n                            )\n        \n        show_stats(acc_log, \"Train\", loss)\n        \n        if scheduler is not None:\n            scheduler.step()\n\n        loss, acc_log = run_epoch(\n                            model = model,\n                            dataloader = test_dl,\n                            optimizer = None,\n                            phase = 'test',\n                            device = device\n                            )\n        show_stats(acc_log, \"Test\", loss)\n\n    avg_acc = round(np.mean(list(acc_log.values())), 3)\n    acc_log |= {f\"Average Accuracy: {average_acc}\"}\n    print(f\"Model Avg. Accuracy: {avg_acc}\")\n    return model, acc_log","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Future Steps\n\nIf I had more time...\n- Go through/impement these [image classification tips/tricks](https://neptune.ai/blog/image-classification-tips-and-tricks-from-13-kaggle-competitions)\n- Actual EDA\n- Additional data augmentation and visualization of all of it\n- Custom logger \n- Trainer customization (callbacks, auto batch etc.)\n- Actually track experiments \n- Hyperparameter tuning (even model tuning)\n- Different loss (weighted cross entropy or even ordinal regression stuff for some parts)\n- Visualization of model focus (what is model looking at)\n- Incorporating real techniques of analyzing ct scans\n- Incorporating segmentations (individualized models for certain things rather than one multiheaded model?\n- Train ensemble \n- Connect TPU","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}