{"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":"### Importing Libraries\n","metadata":{}},{"cell_type":"code","source":"%matplotlib inline\n%config InlineBackend.figure_format = 'retina'\n\nimport numpy as np\nimport pandas as pd\n\nimport shutil\nimport os\n\nfrom PIL import Image\n\nimport pydicom as di\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset, dataloader\nfrom torchvision import transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelBinarizer\n\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\nfrom torchvision.models import resnet50, ResNet50_Weights\n\nimport gc","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:43.438520Z","iopub.execute_input":"2023-04-01T16:14:43.439216Z","iopub.status.idle":"2023-04-01T16:14:47.501169Z","shell.execute_reply.started":"2023-04-01T16:14:43.439179Z","shell.execute_reply":"2023-04-01T16:14:47.499910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Helper Functions\n","metadata":{}},{"cell_type":"code","source":"def imshow(image, ax=None, title=None, normalize=True, **kwargs):\n    \"\"\"Imshow for Tensor.\"\"\"\n    if ax is None:\n        fig, ax = plt.subplots()\n    image = image.numpy().transpose((1, 2, 0))\n\n    if normalize:\n        mean = np.array([0.485, 0.456, 0.406])\n        std = np.array([0.229, 0.224, 0.225])\n        image = std * image + mean\n        image = np.clip(image, 0, 1)\n\n    ax.imshow(image, **kwargs)\n    ax.spines[\"top\"].set_visible(False)\n    ax.spines[\"right\"].set_visible(False)\n    ax.spines[\"left\"].set_visible(False)\n    ax.spines[\"bottom\"].set_visible(False)\n    ax.tick_params(axis=\"both\", length=0)\n    ax.set_xticklabels(\"\")\n    ax.set_yticklabels(\"\")\n\n    return ax\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.507721Z","iopub.execute_input":"2023-04-01T16:14:47.508237Z","iopub.status.idle":"2023-04-01T16:14:47.519422Z","shell.execute_reply.started":"2023-04-01T16:14:47.508195Z","shell.execute_reply":"2023-04-01T16:14:47.517674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_default_device():\n    \"\"\"Sets device name\"\"\"\n    if torch.cuda.is_available():\n        return torch.device(\"cuda\")\n\n    if hasattr(torch.backends, \"mps\"):  # type: ignore\n        if torch.backends.mps.is_available() and torch.backends.mps.is_built():  # type: ignore\n            return torch.device(\"mps\")\n\n    return torch.device(\"cpu\")\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.521286Z","iopub.execute_input":"2023-04-01T16:14:47.521832Z","iopub.status.idle":"2023-04-01T16:14:47.529483Z","shell.execute_reply.started":"2023-04-01T16:14:47.521795Z","shell.execute_reply":"2023-04-01T16:14:47.528155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_ckp(\n    state,\n    base_checkpoint_dir=\"./checkpoints\",\n    is_best=False,\n    best_model_dir=None,\n    checkpoint_name=\"checkpoint.pt\",\n    checkpoint_name_best=\"best_model.pt\",\n):\n    \"\"\"Saves Checkpoint\"\"\"\n    checkpoint_dir = os.path.join(\n        os.path.join(base_checkpoint_dir, \"saved_models\"), \"models\"\n    )\n    os.makedirs(checkpoint_dir, exist_ok=True)\n\n    f_path = os.path.join(checkpoint_dir, checkpoint_name)\n    torch.save(state, f_path)\n    if is_best:\n        if best_model_dir is None:\n            best_model_dir = checkpoint_dir.replace(\"models\", \"best_models\")\n\n        os.makedirs(best_model_dir, exist_ok=True)\n\n        best_fpath = os.path.join(best_model_dir, checkpoint_name_best)\n        shutil.copyfile(f_path, best_fpath)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.533053Z","iopub.execute_input":"2023-04-01T16:14:47.533925Z","iopub.status.idle":"2023-04-01T16:14:47.544721Z","shell.execute_reply.started":"2023-04-01T16:14:47.533871Z","shell.execute_reply":"2023-04-01T16:14:47.542187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_ckp(checkpoint_path, model, optimizer, checkpoint_name=\"checkpoint.pt\"):\n    \"\"\"Loads Checkpoint\"\"\"\n    checkpoint_path = os.path.join(\n        os.path.join(checkpoint_path, \"saved_models\"), \"models\"\n    )\n    checkpoint_fpath = os.path.join(checkpoint_path, checkpoint_name)\n\n    try:\n        checkpoint = torch.load(checkpoint_fpath)\n    except FileNotFoundError:\n        print(\"No checkpoint found at \" + checkpoint_fpath)\n        return model, optimizer, 0\n\n    print(\"Loading checkpoint at\" + checkpoint_fpath)\n    model.load_state_dict(checkpoint[\"state_dict\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer\"])\n    return model, optimizer, checkpoint[\"epoch\"]\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.548017Z","iopub.execute_input":"2023-04-01T16:14:47.548922Z","iopub.status.idle":"2023-04-01T16:14:47.569863Z","shell.execute_reply.started":"2023-04-01T16:14:47.548863Z","shell.execute_reply":"2023-04-01T16:14:47.559117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_best(current_loss, best_loss):\n    \"\"\"Returns if the loss is best or not along with the loss itself\"\"\"\n    if current_loss < best_loss:\n        return current_loss, True\n    return best_loss, False","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.571514Z","iopub.execute_input":"2023-04-01T16:14:47.571922Z","iopub.status.idle":"2023-04-01T16:14:47.585992Z","shell.execute_reply.started":"2023-04-01T16:14:47.571868Z","shell.execute_reply":"2023-04-01T16:14:47.584799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plt_null_graph(df, figsize, palette):\n    \"\"\"Plot a bar grpah indicating all the fields with null values\"\"\"\n    plt.figure(figsize=figsize)\n    cols_w_null = df.isnull().sum().sort_values(ascending=False)\n    sns.barplot(x=cols_w_null.index, y=cols_w_null, palette=palette)\n    plt.xticks(rotation=90)\n    plt.show()\n    return cols_w_null","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.587691Z","iopub.execute_input":"2023-04-01T16:14:47.588781Z","iopub.status.idle":"2023-04-01T16:14:47.601114Z","shell.execute_reply.started":"2023-04-01T16:14:47.588714Z","shell.execute_reply":"2023-04-01T16:14:47.599210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def safe_nav(obj, attr, return_val = None):\n    return obj[attr] if attr in obj else return_val","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.604340Z","iopub.execute_input":"2023-04-01T16:14:47.605914Z","iopub.status.idle":"2023-04-01T16:14:47.614190Z","shell.execute_reply.started":"2023-04-01T16:14:47.605873Z","shell.execute_reply":"2023-04-01T16:14:47.612200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_image_path(\n    df: pd.DataFrame,\n    base_path: str,\n    image_col_name: str,\n    image_extention: str = \"jpg\",\n    has_train_test_folder: bool = None,\n    is_test: bool = None,\n    drop_obj = None,\n    rename_obj = None,\n):\n    \"\"\"Set the path of the image for image_ids\"\"\"\n    path = base_path\n    \n    if has_train_test_folder is not None:\n        train_or_test = \"test\" if is_test else \"train\"\n        path = f\"{path}/{train_or_test}\"\n    \n    def handle_image_directory_set(image_id):\n        \n        _path = f\"{path}/{image_id}.{image_extention}\"\n        \n        if os.path.exists(_path):\n            return _path\n\n        print(\"Image not found\")\n        return None\n    \n    df[image_col_name] = df[image_col_name].apply(handle_image_directory_set)\n\n    if drop_obj is not None:\n        df = df.drop(safe_nav(drop_obj, 'drop_cols'), **safe_nav(drop_obj, 'other_args', {}))\n\n    if rename_obj is not None:\n        df = df.rename(\n            columns=safe_nav(rename_obj, 'rename_cols'),\n            index=safe_nav(rename_obj, 'rename_rows'),\n            **safe_nav(rename_obj, 'other_args', {})\n        )\n    \n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.616624Z","iopub.execute_input":"2023-04-01T16:14:47.618008Z","iopub.status.idle":"2023-04-01T16:14:47.635683Z","shell.execute_reply.started":"2023-04-01T16:14:47.617969Z","shell.execute_reply":"2023-04-01T16:14:47.634643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Import to Dataframe\n","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(\"../input/siim-isic-melanoma-classification/train.csv\")\ndf_train\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.658478Z","iopub.execute_input":"2023-04-01T16:14:47.663736Z","iopub.status.idle":"2023-04-01T16:14:47.788336Z","shell.execute_reply.started":"2023-04-01T16:14:47.663695Z","shell.execute_reply":"2023-04-01T16:14:47.787196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ans = pd.read_csv(\"../input/siim-isic-melanoma-classification/test.csv\")\ndf_ans","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.792886Z","iopub.execute_input":"2023-04-01T16:14:47.793279Z","iopub.status.idle":"2023-04-01T16:14:47.841319Z","shell.execute_reply.started":"2023-04-01T16:14:47.793241Z","shell.execute_reply":"2023-04-01T16:14:47.840217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = \"ISIC_0074311\"\npath = f\"/kaggle/input/siim-isic-melanoma-classification/train/{image_id}.dcm\"\nval = di.dcmread(path).pixel_array\nplt.imshow(val, cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:47.845526Z","iopub.execute_input":"2023-04-01T16:14:47.845865Z","iopub.status.idle":"2023-04-01T16:14:52.785170Z","shell.execute_reply.started":"2023-04-01T16:14:47.845828Z","shell.execute_reply":"2023-04-01T16:14:52.784216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['diagnosis'] = df_train.loc[:, 'diagnosis'].map(lambda x: None if x == 'unknown' else x)\npalette = sns.hls_palette(1, h=336 / 360, l=50 / 100, s=100 / 100)\ncols_w_null_train = plt_null_graph(df_train, (6, 3), palette)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:52.786198Z","iopub.execute_input":"2023-04-01T16:14:52.786539Z","iopub.status.idle":"2023-04-01T16:14:53.085055Z","shell.execute_reply.started":"2023-04-01T16:14:52.786503Z","shell.execute_reply":"2023-04-01T16:14:53.084064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cols_w_null_train","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.086730Z","iopub.execute_input":"2023-04-01T16:14:53.087462Z","iopub.status.idle":"2023-04-01T16:14:53.095711Z","shell.execute_reply.started":"2023-04-01T16:14:53.087423Z","shell.execute_reply":"2023-04-01T16:14:53.094550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cols_w_null_test = plt_null_graph(df_ans, (6, 3), palette)","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.097718Z","iopub.execute_input":"2023-04-01T16:14:53.098085Z","iopub.status.idle":"2023-04-01T16:14:53.357125Z","shell.execute_reply.started":"2023-04-01T16:14:53.098050Z","shell.execute_reply":"2023-04-01T16:14:53.356099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cols_w_null_test","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.358988Z","iopub.execute_input":"2023-04-01T16:14:53.359416Z","iopub.status.idle":"2023-04-01T16:14:53.368307Z","shell.execute_reply.started":"2023-04-01T16:14:53.359372Z","shell.execute_reply":"2023-04-01T16:14:53.367202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.loc[:, 'age_approx'].fillna(df_train.loc[:, 'age_approx'].median(), inplace=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.369445Z","iopub.execute_input":"2023-04-01T16:14:53.370728Z","iopub.status.idle":"2023-04-01T16:14:53.380301Z","shell.execute_reply.started":"2023-04-01T16:14:53.370690Z","shell.execute_reply":"2023-04-01T16:14:53.379411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.loc[:, 'anatom_site_general_challenge'].fillna(\"\", inplace=True)\ndf_train.loc[:, 'sex'].fillna(\"\", inplace=True)\ndf_ans.loc[:, 'anatom_site_general_challenge'].fillna(\"\", inplace=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.383025Z","iopub.execute_input":"2023-04-01T16:14:53.383330Z","iopub.status.idle":"2023-04-01T16:14:53.397617Z","shell.execute_reply.started":"2023-04-01T16:14:53.383302Z","shell.execute_reply":"2023-04-01T16:14:53.396590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_sample_size(labels, num_samples):\n    df_temp = pd.DataFrame(\n        {\n            \"label\": labels,\n            \"n_samples\": num_samples\n        }\n    )\n    df_temp.sort_values(by=\"n_samples\", inplace=True)\n\n    plt.figure()\n    sns.set(style=\"whitegrid\", font_scale=0.6, color_codes=True)\n    palette = sns.color_palette(\"ch:s=-.2,r=.6\", len(df_temp[\"n_samples\"]))\n    rank = df_temp[\"n_samples\"].argsort().argsort()\n    sns.barplot(data=df_temp, x=\"label\", y=\"n_samples\", palette=np.array(palette[::-1])[rank])\n    plt.xticks(rotation=90)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.400182Z","iopub.execute_input":"2023-04-01T16:14:53.402347Z","iopub.status.idle":"2023-04-01T16:14:53.409570Z","shell.execute_reply.started":"2023-04-01T16:14:53.402311Z","shell.execute_reply":"2023-04-01T16:14:53.408640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = df_train[\"benign_malignant\"].unique()\nnum_samples = df_train[\"benign_malignant\"].value_counts().values\nplot_sample_size(labels, num_samples)","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.411009Z","iopub.execute_input":"2023-04-01T16:14:53.411462Z","iopub.status.idle":"2023-04-01T16:14:53.713189Z","shell.execute_reply.started":"2023-04-01T16:14:53.411426Z","shell.execute_reply":"2023-04-01T16:14:53.712134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['target'].value_counts()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.714843Z","iopub.execute_input":"2023-04-01T16:14:53.715270Z","iopub.status.idle":"2023-04-01T16:14:53.724610Z","shell.execute_reply.started":"2023-04-01T16:14:53.715229Z","shell.execute_reply":"2023-04-01T16:14:53.723342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[\"benign_malignant\"].value_counts()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.726604Z","iopub.execute_input":"2023-04-01T16:14:53.727008Z","iopub.status.idle":"2023-04-01T16:14:53.739936Z","shell.execute_reply.started":"2023-04-01T16:14:53.726966Z","shell.execute_reply":"2023-04-01T16:14:53.738685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_path = \"/kaggle/input/siim-isic-melanoma-classification\"\nimage_extention = \"dcm\"\n\nimage_col_name = \"image_name\"\n\ndrop_obj = {\n    \"drop_cols\": [\"diagnosis\", \"target\"],\n    \"other_args\":{\n        \"axis\": 1\n    }\n}\n\nrename_obj = {\n    \"rename_cols\": {\n        \"benign_malignant\": \"label\",\n        \"image_name\": \"image_path\",\n        \"sex\": \"gender\",\n        \"anatom_site_general_challenge\": \"anatom_site\"\n    },\n}\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.741642Z","iopub.execute_input":"2023-04-01T16:14:53.742773Z","iopub.status.idle":"2023-04-01T16:14:53.749574Z","shell.execute_reply.started":"2023-04-01T16:14:53.742738Z","shell.execute_reply":"2023-04-01T16:14:53.748451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = set_image_path(\n    df_train, \n    base_path,\n    image_col_name,\n    image_extention = image_extention,\n    has_train_test_folder = True,\n    drop_obj = drop_obj,\n    rename_obj = rename_obj\n)\n\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:14:53.750949Z","iopub.execute_input":"2023-04-01T16:14:53.751971Z","iopub.status.idle":"2023-04-01T16:16:01.312970Z","shell.execute_reply.started":"2023-04-01T16:14:53.751853Z","shell.execute_reply":"2023-04-01T16:16:01.311901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ans = set_image_path(\n    df_ans,\n    base_path,\n    image_col_name,\n    image_extention = image_extention,\n    has_train_test_folder = True,\n    is_test = True,\n    rename_obj = rename_obj\n)\n\ndf_ans.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:01.314361Z","iopub.execute_input":"2023-04-01T16:16:01.314999Z","iopub.status.idle":"2023-04-01T16:16:25.826900Z","shell.execute_reply.started":"2023-04-01T16:16:01.314960Z","shell.execute_reply":"2023-04-01T16:16:25.825818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train, df_test = train_test_split(df_train, test_size=0.2, random_state=42)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:25.828536Z","iopub.execute_input":"2023-04-01T16:16:25.828903Z","iopub.status.idle":"2023-04-01T16:16:25.845316Z","shell.execute_reply.started":"2023-04-01T16:16:25.828865Z","shell.execute_reply":"2023-04-01T16:16:25.844269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.DataFrame(df_train)\ndf_test = pd.DataFrame(df_test)\n\n# Used custom encoding instead of 'target' as will give the controls provided by \n# sklearn.LabelBinarizer() and thus can be used directly.\nlb_1 = LabelBinarizer()\nlb_2 = LabelBinarizer()\nlb_3 = LabelBinarizer()\n\ndf_train[\"enc_label\"] = lb_1.fit_transform(df_train[\"label\"]).tolist()  # type: ignore\n\ndf_test[\"enc_label\"] = lb_1.transform(df_test[\"label\"]).tolist()  # type: ignore\n\ndf_train[\"enc_anatom_site\"] = lb_2.fit_transform(df_train[\"anatom_site\"]).tolist()  # type: ignore\n\ndf_ans[\"enc_anatom_site\"] = lb_2.transform(df_ans[\"anatom_site\"]).tolist()  # type: ignore\n\ndf_test[\"enc_anatom_site\"] = lb_2.transform(df_test[\"anatom_site\"]).tolist()  # type: ignore\n\ndf_train[\"enc_gender\"] = lb_3.fit_transform(df_train[\"gender\"]).tolist()  # type: ignore\n\ndf_ans[\"enc_gender\"] = lb_3.transform(df_ans[\"gender\"]).tolist()  # type: ignore\n\ndf_test[\"enc_gender\"] = lb_3.transform(df_test[\"gender\"]).tolist()  # type: ignore\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:25.847073Z","iopub.execute_input":"2023-04-01T16:16:25.847446Z","iopub.status.idle":"2023-04-01T16:16:26.224668Z","shell.execute_reply.started":"2023-04-01T16:16:25.847410Z","shell.execute_reply":"2023-04-01T16:16:26.223669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_enc(el):\n    if el[0] == 0:\n        return [0, 1]\n    \n    if el[0] == 1:\n        return [1, 0]\n    \n    raise ValueError(\"Value not 0 or 1\")","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.232816Z","iopub.execute_input":"2023-04-01T16:16:26.233111Z","iopub.status.idle":"2023-04-01T16:16:26.238566Z","shell.execute_reply.started":"2023-04-01T16:16:26.233084Z","shell.execute_reply":"2023-04-01T16:16:26.237470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[\"enc_label\"] = df_train[\"enc_label\"].apply(fix_enc)\n\ndf_test[\"enc_label\"] = df_test[\"enc_label\"].apply(fix_enc)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.239805Z","iopub.execute_input":"2023-04-01T16:16:26.240428Z","iopub.status.idle":"2023-04-01T16:16:26.393399Z","shell.execute_reply.started":"2023-04-01T16:16:26.240391Z","shell.execute_reply":"2023-04-01T16:16:26.392196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.head()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.394741Z","iopub.execute_input":"2023-04-01T16:16:26.395204Z","iopub.status.idle":"2023-04-01T16:16:26.418675Z","shell.execute_reply.started":"2023-04-01T16:16:26.395168Z","shell.execute_reply":"2023-04-01T16:16:26.416657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ans.head()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.420293Z","iopub.execute_input":"2023-04-01T16:16:26.420794Z","iopub.status.idle":"2023-04-01T16:16:26.440593Z","shell.execute_reply.started":"2023-04-01T16:16:26.420752Z","shell.execute_reply":"2023-04-01T16:16:26.439697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Setting the Device\n","metadata":{}},{"cell_type":"code","source":"device = get_default_device()\n\nprint(f\"Using device: {device}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.443686Z","iopub.execute_input":"2023-04-01T16:16:26.443969Z","iopub.status.idle":"2023-04-01T16:16:26.505893Z","shell.execute_reply.started":"2023-04-01T16:16:26.443943Z","shell.execute_reply":"2023-04-01T16:16:26.504670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Preprocessing the Data\n","metadata":{}},{"cell_type":"code","source":"class MelanomaClassificationLoader(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def collate_fn(self, batch):\n        batch = list(filter(lambda x: x is not None, batch))\n        return dataloader.default_collate(batch)\n\n    def __getitem__(self, idx):\n        try:\n            image = di.dcmread(self.df.iloc[idx][\"image_path\"]).pixel_array\n        except IOError:\n            print(f\"Ah Snap: {self.df.iloc[idx]['image_path']}\")\n            return None\n\n        if self.transform:\n            image = self.transform(image)\n\n        enc_label = torch.tensor(self.df.iloc[idx][\"enc_label\"])\n\n        return image.float(), enc_label.float()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.507727Z","iopub.execute_input":"2023-04-01T16:16:26.508548Z","iopub.status.idle":"2023-04-01T16:16:26.517018Z","shell.execute_reply.started":"2023-04-01T16:16:26.508510Z","shell.execute_reply":"2023-04-01T16:16:26.516190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transformer = transforms.Compose(\n    [\n        transforms.ToPILImage(),\n        transforms.Resize((224, 224)),\n        transforms.RandomRotation(30),\n        # transforms.RandomResizedCrop(224),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomVerticalFlip(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ]\n)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.518698Z","iopub.execute_input":"2023-04-01T16:16:26.519486Z","iopub.status.idle":"2023-04-01T16:16:26.530571Z","shell.execute_reply.started":"2023-04-01T16:16:26.519450Z","shell.execute_reply":"2023-04-01T16:16:26.529678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_train = MelanomaClassificationLoader(df_train, transform=transformer)\ndataset_test = MelanomaClassificationLoader(df_ans, transform=transformer)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.531973Z","iopub.execute_input":"2023-04-01T16:16:26.532780Z","iopub.status.idle":"2023-04-01T16:16:26.540808Z","shell.execute_reply.started":"2023-04-01T16:16:26.532743Z","shell.execute_reply":"2023-04-01T16:16:26.539966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# batch_size = 64\nbatch_size = 32\ntrain_loader = DataLoader(\n    dataset_train,\n    batch_size=batch_size,\n    shuffle=True,\n    collate_fn=dataset_train.collate_fn,\n)\ntest_loader = DataLoader(\n    dataset_test,\n    batch_size=batch_size,\n    shuffle=True,\n    collate_fn=dataset_test.collate_fn,\n)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.542168Z","iopub.execute_input":"2023-04-01T16:16:26.542907Z","iopub.status.idle":"2023-04-01T16:16:26.551962Z","shell.execute_reply.started":"2023-04-01T16:16:26.542868Z","shell.execute_reply":"2023-04-01T16:16:26.551165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def revert_back_fix_enc(el):\n    \n    if int(el[0]) == 1 and int(el[1]) == 0:\n        return 1\n    \n    if int(el[0]) == 0 and int(el[1]) == 1:\n        return 0\n    \n    raise ValueError(\"Unexpected Value\")","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.553272Z","iopub.execute_input":"2023-04-01T16:16:26.554418Z","iopub.status.idle":"2023-04-01T16:16:26.561863Z","shell.execute_reply.started":"2023-04-01T16:16:26.554383Z","shell.execute_reply":"2023-04-01T16:16:26.560875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images_iter = iter(train_loader)\n\nimages, label = next(images_iter)\n\nlabel = label.tolist()\nlabel = [revert_back_fix_enc(x) for x in label]\nlabel = torch.Tensor(label)\n\nlabel = lb_1.inverse_transform(label)\n\nprint(images.shape)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:26.563317Z","iopub.execute_input":"2023-04-01T16:16:26.563792Z","iopub.status.idle":"2023-04-01T16:16:41.458530Z","shell.execute_reply.started":"2023-04-01T16:16:26.563757Z","shell.execute_reply":"2023-04-01T16:16:41.457400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(figsize=(10, 4), ncols=4)\n\nfor i in range(4):\n    ax = axes[i]\n    ax.set_title(f\"{label[i]}\", fontsize=8)\n    imshow(images[i], ax=ax)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:41.459854Z","iopub.execute_input":"2023-04-01T16:16:41.461400Z","iopub.status.idle":"2023-04-01T16:16:42.283774Z","shell.execute_reply.started":"2023-04-01T16:16:41.461360Z","shell.execute_reply":"2023-04-01T16:16:42.279745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training Model\n","metadata":{}},{"cell_type":"code","source":"class BaseNet(nn.Module):\n    def __init__(self, output_class, pretrained_model):\n        super(BaseNet, self).__init__()\n        self.backbone = pretrained_model\n        self.top_layer_processing = nn.Sequential(\n            nn.Dropout(0.5),\n            nn.BatchNorm1d(1000),\n        )\n        self.output_layer = nn.Sequential(\n            nn.Linear(in_features=1000, out_features=output_class),\n            nn.Softmax(dim=1)\n        )\n    \n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.top_layer_processing(x)\n        x = self.output_layer(x)\n        return x\n        ","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:42.285284Z","iopub.execute_input":"2023-04-01T16:16:42.285905Z","iopub.status.idle":"2023-04-01T16:16:42.294118Z","shell.execute_reply.started":"2023-04-01T16:16:42.285868Z","shell.execute_reply":"2023-04-01T16:16:42.292865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resnet50_model = resnet50(weights=ResNet50_Weights.DEFAULT).to(device)","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:42.295477Z","iopub.execute_input":"2023-04-01T16:16:42.296477Z","iopub.status.idle":"2023-04-01T16:16:47.097080Z","shell.execute_reply.started":"2023-04-01T16:16:42.296442Z","shell.execute_reply":"2023-04-01T16:16:47.095878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(list(resnet50_model.children()))","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.098883Z","iopub.execute_input":"2023-04-01T16:16:47.099631Z","iopub.status.idle":"2023-04-01T16:16:47.104201Z","shell.execute_reply.started":"2023-04-01T16:16:47.099592Z","shell.execute_reply":"2023-04-01T16:16:47.102958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 2\nmodel = BaseNet(num_classes, resnet50_model).to(device)\ngc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.105790Z","iopub.execute_input":"2023-04-01T16:16:47.106536Z","iopub.status.idle":"2023-04-01T16:16:47.314931Z","shell.execute_reply.started":"2023-04-01T16:16:47.106501Z","shell.execute_reply":"2023-04-01T16:16:47.313974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learning_rate = 0.001\ncriterion = nn.CrossEntropyLoss().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=0.001)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.316553Z","iopub.execute_input":"2023-04-01T16:16:47.317238Z","iopub.status.idle":"2023-04-01T16:16:47.328346Z","shell.execute_reply.started":"2023-04-01T16:16:47.317200Z","shell.execute_reply":"2023-04-01T16:16:47.327428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(\"Model's state_dict:\")\n# for param_tensor in model.state_dict():\n#     print(param_tensor, \"\\t\", model.state_dict()[param_tensor].size())\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.330471Z","iopub.execute_input":"2023-04-01T16:16:47.331968Z","iopub.status.idle":"2023-04-01T16:16:47.337646Z","shell.execute_reply.started":"2023-04-01T16:16:47.331932Z","shell.execute_reply":"2023-04-01T16:16:47.336671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_accuracy(predictions, labels):\n    \"\"\"Calculates accuracy of the model\"\"\"\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for prediction, label in zip(predictions, labels):\n            _, predicted = torch.max(prediction.data, 1)\n            _, label = torch.max(label.data, 1)\n            total += label.size(0)\n            correct += (predicted == label).sum().item()\n\n    return 100 * correct / total, correct\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.339661Z","iopub.execute_input":"2023-04-01T16:16:47.341090Z","iopub.status.idle":"2023-04-01T16:16:47.352894Z","shell.execute_reply.started":"2023-04-01T16:16:47.341056Z","shell.execute_reply":"2023-04-01T16:16:47.351903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test(model, test_loader, criterion, device):\n    model.eval()\n    test_loss = 0\n    correct = 0\n    with torch.no_grad():\n        predictions = []\n        labels = []\n        for data, target in test_loader:\n            data, target = data.to(device), target.to(device)\n            output = model(data)\n            loss = criterion(output, target)\n            test_loss += loss.item()\n            predictions.append(output)\n            labels.append(target)\n\n    test_loss /= len(test_loader.dataset)\n    accuracy, correct = calculate_accuracy(predictions, labels)\n\n    print(\n        f\"\\nTest set: Average loss: {test_loss:.4f}, Accuracy: {correct}/{len(test_loader.dataset)} ({(accuracy):.0f}%)\\n\"\n    )\n\n    return test_loss, accuracy\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.355486Z","iopub.execute_input":"2023-04-01T16:16:47.356840Z","iopub.status.idle":"2023-04-01T16:16:47.371271Z","shell.execute_reply.started":"2023-04-01T16:16:47.356804Z","shell.execute_reply":"2023-04-01T16:16:47.370196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_network(\n    model,\n    train_loader, \n    criterion, \n    optimizer, \n    num_epochs, \n    device, \n    print_every=10,\n    base_checkpoint_dir=\"./checkpoint_dir\"\n):\n    \n    train_losses = []\n    test_losses = []\n    best_accuracy = 0\n    # train_accuracies = []\n    # test_accuracies = []\n    # Set the model to training mode\n    \n    model, optimizer, start_epoch = load_ckp(base_checkpoint_dir, model, optimizer)\n    \n    model.train()\n\n    for epoch in range(start_epoch, num_epochs):\n        running_loss = 0.0\n        predictions = []\n        all_labels = []\n        accuracy = 0\n        r_loss = 0\n        \n        model, optimizer, start_step = load_ckp(\"./checkpoint_steps\", model, optimizer)\n        \n        for i, (inputs, labels) in enumerate(train_loader):\n            if i < start_step:\n                continue\n            # Move input and label tensors to the default device\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            # Forward pass\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            # Save predictions and labels\n            all_labels.append(labels)\n            predictions.append(outputs)\n\n            # Backward and optimize\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            \n            running_loss += loss.item()\n\n            # Print statistics\n            if (i + 1) % 100 == 0:\n                checkpoint = {\n                    \"step\": i + 1,\n                    \"state_dict\": model.state_dict(),\n                    \"optimizer\": optimizer.state_dict(),\n                }\n                save_ckp(checkpoint, \"./checkpoint_steps\")\n                \n                r_loss = running_loss / 100 - r_loss\n                print(\n                    f\"Step [{i+1}/{len(train_loader)}], Loss: {r_loss:.4f}\"\n                )\n        \n        train_losses.append(running_loss)\n        running_loss /= len(train_loader.dataset)\n\n\n        # Calculate accuracy\n        accuracy, correct = calculate_accuracy(predictions, all_labels)\n\n        if (epoch + 1) % print_every == 0:\n            print(\n                f\"Epoch [{epoch+1}/{num_epochs}], Loss: {running_loss / 100:.4f}, Accuracy: {correct}/{len(train_loader.dataset)} ({(accuracy):.0f}%)\"\n            )\n            running_loss = 0.0\n            \n        # Testing the model\n        loss, acc = test(model, test_loader, criterion, device)\n        \n        # Checking if this epoch has the best Accuracy\n        best_accuracy, is_acc_best = is_best(acc, best_accuracy)\n        \n        # Saving the epoch\n        checkpoint = {\n            \"epoch\": epoch + 1,\n            \"state_dict\": model.state_dict(),\n            \"optimizer\": optimizer.state_dict(),\n        }\n        save_ckp(checkpoint, base_checkpoint_dir, is_acc_best)\n        \n        test_losses.append(loss)\n\n        # if acc > 0.99:\n        #     print(\"Accuracy is greater than 99%\")\n        #     break\n        \n    return model, train_losses, test_losses\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.373419Z","iopub.execute_input":"2023-04-01T16:16:47.374501Z","iopub.status.idle":"2023-04-01T16:16:47.399613Z","shell.execute_reply.started":"2023-04-01T16:16:47.374466Z","shell.execute_reply":"2023-04-01T16:16:47.398590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 5\nprint_every = 1\n\ntry:\n    trained_model, train_losses, test_losses = train_network(\n        model, train_loader, criterion, optimizer, num_epochs, device, print_every\n    )\nexcept KeyboardInterrupt:\n    print(\"Training stopped\")\n","metadata":{"execution":{"iopub.status.busy":"2023-04-01T16:16:47.404119Z","iopub.execute_input":"2023-04-01T16:16:47.407061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree(\"./checkpoint_dir\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}