{"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":"<img src=\"https://i.imgur.com/2LmSJWD.jpg\" width=100%>\n<h1><center>Mayo Clinic - STRIP AI</center></h1>\n\n## Goal\nThe goal of this competition is to ***classify*** the blood clot origins in ischemic\nstroke. Using whole slide digital pathology images, you'll build a model that\ndifferentiates between the two major acute ischemic stroke (AIS) etiology subtypes:\n- Cardic (CA)\n- Large artery atherosclerosis (LAA)\n\n## Why❓️\nThe model developed here will enable healthcare providers to better identify the origins of blood clots in deadly strokes, making it easier for physicians to prescribe the best post-stroke therapeutic management and reducing the likelihood of a second stroke.\n\n> A stroke is a medical condition in which poor blood flow to the brain causes cell death. There are two main types of stroke: ischemic, due to lack of blood flow, and hemorrhagic, due to bleeding.\n\n## Submission File\nFor each `patient_id` in the test set, you must predict a probability for each of the two etiology classes. The file should contain a header and have the following format:\n```csv\npatient_id,CE,LAA\n01f2b3,0.5,0.5\n04de22,0.5,0.5\n0a47c9,0.5,0.5\n0af8b6,0.5,0.5\n...\n```\n\n## Evaluation\n> Metric Used: Weighted multi-class logarithmic loss\n\n<center>$\\operatorname{Log Loss}=-\\left(\\frac{\\sum_{i=1}^{M} w_{i} \\cdot \\sum_{j=1}^{N_{i}} \\frac{y_{i j}}{N_{i}} \\cdot \\ln p_{i j}}{\\sum_{i=1}^{M} w_{i}}\\right)$</center>\nwhere,<br>\n<center>\n    $\\operatorname{N} = \\operatorname{Number of images in the class set}$<br>\n    $\\operatorname{M} = \\operatorname{Number of classes}$<br>\n    $\\operatorname{ln} = \\operatorname{Natural Logarithm}$<br>\n    \n</center>\n\n> **Note**: The submitted probabilities for a given image are not required to sum to one because they are rescaled prior to being scored (each row is divided by the row sum).\n\n> In order to avoid the extremes of the log function, each predicted probability   is replaced with $max(min({p}, 1-10^{-15}), 10^{-15})$\n\n## Good to know\n### Blood clot\nBlood clots are gel-like collections of blood that form in your veins or arteries when blood changes from liquid to partially solid. Clotting is normal, but clots can be dangerous when they do not dissolve on their own. Treatments range from medications to surgery.\n\n### Whole Slide Imaging\nWhole slide imaging, also known as virtual microscopy, refers to scanning a complete microscope slide and creating a single high-resolution digital file. This is commonly achieved by capturing many small high-resolution image tiles or strips and then montaging them to create a full image of a histological section.\nFor example: <br>\n<center><img height=250 width=250 src=\"https://external-content.duckduckgo.com/iu/?u=https%3A%2F%2Ftse1.mm.bing.net%2Fth%3Fid%3DOIP.ybYtlmv7bts8vkdAhkMrnwHaE8%26pid%3DApi&f=1\"></center>","metadata":{}},{"cell_type":"code","source":"!pip install -Uq timm","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:03.479615Z","iopub.execute_input":"2022-07-31T17:30:03.480308Z","iopub.status.idle":"2022-07-31T17:30:16.174143Z","shell.execute_reply.started":"2022-07-31T17:30:03.48027Z","shell.execute_reply":"2022-07-31T17:30:16.172974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom glob import glob\nfrom pprint import pprint\nimport random\nimport cv2\nfrom joblib import Parallel, delayed\nfrom sklearn.metrics import accuracy_score\nfrom tqdm.notebook import tqdm\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nimport rasterio\nfrom sklearn.model_selection import train_test_split\n\nimport timm\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, utils\n\ncolors = ['#E7D5E8','#F9659B','#F69581','#F68FBB']\nsns.palplot(sns.color_palette(colors))\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Set Style\nsns.set_style(\"whitegrid\")\nsns.despine(left=True, bottom=True)\n\n# plt.rc('xtick',labelsize=11)\n# plt.rc('ytick',labelsize=11)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:16.17667Z","iopub.execute_input":"2022-07-31T17:30:16.17711Z","iopub.status.idle":"2022-07-31T17:30:23.735067Z","shell.execute_reply.started":"2022-07-31T17:30:16.177066Z","shell.execute_reply":"2022-07-31T17:30:23.733826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = dict(\n    orig_train_dir = os.path.abspath('../input/mayo-clinic-strip-ai/train'),\n    orig_test_dir = os.path.abspath('../input/mayo-clinic-strip-ai/test'),\n    orig_other_dir = os.path.abspath('../input/mayo-clinic-strip-ai/other'),\n    orig_train_csv_path =  os.path.abspath('../input/mayo-clinic-strip-ai/train.csv'),\n    orig_test_csv_path = os.path.abspath('../input/mayo-clinic-strip-ai/test.csv'),\n    orig_sample_submission_path = os.path.abspath('../input/mayo-clinic-strip-ai/sample_submission.csv'),\n    orig_other_csv_path = os.path.abspath('../input/mayo-clinic-strip-ai/other.csv'),\n    \n    seed = 42,\n    device = 'cuda:0' if torch.cuda.is_available() else 'cpu',\n    \n    batch_size = 128,\n    num_epochs=10,\n    lr = 0.0003,\n    \n    use_wandb=False\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.737306Z","iopub.execute_input":"2022-07-31T17:30:23.738204Z","iopub.status.idle":"2022-07-31T17:30:23.818838Z","shell.execute_reply.started":"2022-07-31T17:30:23.738164Z","shell.execute_reply":"2022-07-31T17:30:23.817792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pprint(config)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.822779Z","iopub.execute_input":"2022-07-31T17:30:23.82345Z","iopub.status.idle":"2022-07-31T17:30:23.839494Z","shell.execute_reply.started":"2022-07-31T17:30:23.823408Z","shell.execute_reply":"2022-07-31T17:30:23.838282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(config['seed'])","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.841478Z","iopub.execute_input":"2022-07-31T17:30:23.841986Z","iopub.status.idle":"2022-07-31T17:30:23.851335Z","shell.execute_reply.started":"2022-07-31T17:30:23.84195Z","shell.execute_reply":"2022-07-31T17:30:23.850408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = pd.read_csv(config['orig_train_csv_path'])\ndisplay(train_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.85269Z","iopub.execute_input":"2022-07-31T17:30:23.853153Z","iopub.status.idle":"2022-07-31T17:30:23.877935Z","shell.execute_reply.started":"2022-07-31T17:30:23.853114Z","shell.execute_reply":"2022-07-31T17:30:23.876933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Contains annotations for images in the `train/` folder\n#### Coloums:\n- `image_id`: A unique identifier for this instance having the form `{patient_id}_{image_num}`. Corresponds to the image `{image_id}.tif`\n- `center_id`: Identifies the medical center where the slide was obtained.\n- `patient_id`: Identifies the patient from whom the slide was obtained.\n- `image_num`: Enumerates images of clots obtained from the same patient.\n- `label`: The etiology of the clot, either `CE` or `LAA`. This field is the classification target.","metadata":{}},{"cell_type":"code","source":"print(f'Number of Training Samples: {train_data.shape[0]}')\nprint(f'Are there any missing values?: {train_data.isnull().values.any()}')                                                ","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.879295Z","iopub.execute_input":"2022-07-31T17:30:23.879738Z","iopub.status.idle":"2022-07-31T17:30:23.886776Z","shell.execute_reply.started":"2022-07-31T17:30:23.8797Z","shell.execute_reply":"2022-07-31T17:30:23.885764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = pd.read_csv(config['orig_test_csv_path'])\ndisplay(test_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.888107Z","iopub.execute_input":"2022-07-31T17:30:23.889183Z","iopub.status.idle":"2022-07-31T17:30:23.905324Z","shell.execute_reply.started":"2022-07-31T17:30:23.889063Z","shell.execute_reply":"2022-07-31T17:30:23.904512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_data = pd.read_csv(config['orig_sample_submission_path'])\ndisplay(submission_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.906499Z","iopub.execute_input":"2022-07-31T17:30:23.907388Z","iopub.status.idle":"2022-07-31T17:30:23.923052Z","shell.execute_reply.started":"2022-07-31T17:30:23.907328Z","shell.execute_reply":"2022-07-31T17:30:23.922074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"supplementary_data = pd.read_csv(config['orig_other_csv_path'])\ndisplay(supplementary_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.927457Z","iopub.execute_input":"2022-07-31T17:30:23.928089Z","iopub.status.idle":"2022-07-31T17:30:23.945903Z","shell.execute_reply.started":"2022-07-31T17:30:23.928054Z","shell.execute_reply":"2022-07-31T17:30:23.944933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_to_label, label_to_id = {}, {}\n\nfor id, label in enumerate(train_data['label'].unique()):\n    label_to_id[label] = id\n    id_to_label[id] = label\n\ntrain_data['label'] = train_data['label'].replace(label_to_id)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.947303Z","iopub.execute_input":"2022-07-31T17:30:23.947907Z","iopub.status.idle":"2022-07-31T17:30:23.959844Z","shell.execute_reply.started":"2022-07-31T17:30:23.94787Z","shell.execute_reply":"2022-07-31T17:30:23.958767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🔎 Exploratory data analysis","metadata":{}},{"cell_type":"code","source":"type_distribution = train_data['label'].value_counts()\nplt.figure(figsize=(20, 5))\nsns.barplot(x=type_distribution.values, y=list(id_to_label.values()), palette=colors)\nplt.title('Class Distribution', fontsize=25)\nplt.xlabel('Frequency', fontsize=25)\nplt.ylabel('Label', fontsize=25)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:23.961564Z","iopub.execute_input":"2022-07-31T17:30:23.961963Z","iopub.status.idle":"2022-07-31T17:30:24.185844Z","shell.execute_reply.started":"2022-07-31T17:30:23.961929Z","shell.execute_reply":"2022-07-31T17:30:24.184709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"center_distribution = train_data['center_id'].value_counts().sort_values(ignore_index=True)\nplt.figure(figsize=(20, 5))\nsns.barplot(x=center_distribution.index, y=center_distribution.values, palette=colors)\nplt.title('CenterID Distribution', fontsize=25)\nplt.xlabel('Center ID', fontsize=25)\nplt.ylabel('Frequency', fontsize=25)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:24.187189Z","iopub.execute_input":"2022-07-31T17:30:24.187633Z","iopub.status.idle":"2022-07-31T17:30:24.657548Z","shell.execute_reply.started":"2022-07-31T17:30:24.187594Z","shell.execute_reply":"2022-07-31T17:30:24.656636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🔎 Visualization","metadata":{}},{"cell_type":"code","source":"positive_df = pd.DataFrame(os.listdir('../input/strip-ai-background-clot/positive'), columns=['filename'])\npositive_df['filepath'] = positive_df['filename'].apply(lambda x: os.path.join('../input/strip-ai-background-clot/positive', x))\npositive_df['label'] = 1\npositive_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:24.658972Z","iopub.execute_input":"2022-07-31T17:30:24.659738Z","iopub.status.idle":"2022-07-31T17:30:25.179314Z","shell.execute_reply.started":"2022-07-31T17:30:24.659698Z","shell.execute_reply":"2022-07-31T17:30:25.178251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"negative_df = pd.DataFrame(os.listdir('../input/strip-ai-background-clot/negative'), columns=['filename'])\nnegative_df['filepath'] = negative_df['filename'].apply(lambda x: os.path.join('../input/strip-ai-background-clot/negative', x))\nnegative_df['label'] = 0\nnegative_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:25.180864Z","iopub.execute_input":"2022-07-31T17:30:25.181726Z","iopub.status.idle":"2022-07-31T17:30:26.046827Z","shell.execute_reply.started":"2022-07-31T17:30:25.181686Z","shell.execute_reply":"2022-07-31T17:30:26.045874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.concat([positive_df, negative_df]).sample(frac=1)\nprint(data.shape)\ndisplay(data.head(10))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:26.048134Z","iopub.execute_input":"2022-07-31T17:30:26.0489Z","iopub.status.idle":"2022-07-31T17:30:26.069179Z","shell.execute_reply.started":"2022-07-31T17:30:26.04886Z","shell.execute_reply":"2022-07-31T17:30:26.068191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_background(filepath, threshold=10, display=False, ax=None):\n    image = cv2.imread(filepath)\n    h, w, _ = image.shape\n    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n    thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_OTSU + cv2.THRESH_BINARY_INV)[1]\n\n    pixels = cv2.countNonZero(thresh)\n    ratio = (pixels/(h * w)) * 100\n    #print('Pixel ratio: {:.2f}%'.format(ratio))\n    roi = 0\n    if ratio >= threshold:\n        roi = 1\n    \n    if display and ax is not None:\n        ax.imshow(thresh)\n        ax.set_title('Mostly Background' if not roi else 'Contains region of interest', fontsize=14)\n        ax.axis('off')\n    \n    return roi","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:26.070694Z","iopub.execute_input":"2022-07-31T17:30:26.071304Z","iopub.status.idle":"2022-07-31T17:30:26.079401Z","shell.execute_reply.started":"2022-07-31T17:30:26.071267Z","shell.execute_reply":"2022-07-31T17:30:26.078429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=2, ncols=6, figsize=(20,6))\nplt.suptitle(\"Samples\", fontsize = 16)\n\nfor i in range(0, 2*6):\n    image = cv2.imread(data['filepath'].values[i])\n    \n    x = i // 6\n    y = i % 6\n    axes[x, y].imshow(image, cmap=plt.cm.bone)\n    #axes[x, y].axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:26.080983Z","iopub.execute_input":"2022-07-31T17:30:26.081647Z","iopub.status.idle":"2022-07-31T17:30:28.814027Z","shell.execute_reply.started":"2022-07-31T17:30:26.081608Z","shell.execute_reply":"2022-07-31T17:30:28.813127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=5, ncols=8, figsize=(40,25))\n\nfor i in range(0, 5*8):\n    x = i // 8\n    y = i % 8\n    check_background(data['filepath'].values[np.random.randint(0, len(data))], 10, True, axes[x, y])","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:28.815031Z","iopub.execute_input":"2022-07-31T17:30:28.815358Z","iopub.status.idle":"2022-07-31T17:30:36.843708Z","shell.execute_reply.started":"2022-07-31T17:30:28.815322Z","shell.execute_reply":"2022-07-31T17:30:36.842867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = Parallel(n_jobs=-1)(delayed(check_background)(data['filepath'].values[i], 5) for i in tqdm(range(len(data))))\nprint(f'Accuracy: {accuracy_score(preds, data[\"label\"].values)}')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:36.845049Z","iopub.execute_input":"2022-07-31T17:30:36.845906Z","iopub.status.idle":"2022-07-31T17:30:36.849985Z","shell.execute_reply.started":"2022-07-31T17:30:36.845868Z","shell.execute_reply":"2022-07-31T17:30:36.848796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame(glob('../input/mayo-clinic-1024*[!test]/train/*.jpg'), columns=['filepath'])\ndf['image_id'] = df['filepath'].apply(lambda x: x.split('/')[-1].split('-')[0])\ndf = df.merge(train_data, on='image_id', how='left')\ndf = df.drop(['image_id', 'center_id', 'image_num', 'patient_id'], axis=1)\n\nprint(f'Number of Training Samples: {df.shape[0]}')\ndisplay(df.head(10))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:36.851206Z","iopub.execute_input":"2022-07-31T17:30:36.851958Z","iopub.status.idle":"2022-07-31T17:30:36.864837Z","shell.execute_reply.started":"2022-07-31T17:30:36.851863Z","shell.execute_reply":"2022-07-31T17:30:36.860308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"is_background = Parallel(n_jobs=-1)(delayed(check_background)(df['filepath'].values[i], 5) for i in tqdm(range(len(df))))\ndf['is_background'] = is_background\ndf.to_csv('dataset.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:36.86765Z","iopub.execute_input":"2022-07-31T17:30:36.867974Z","iopub.status.idle":"2022-07-31T17:30:36.871807Z","shell.execute_reply.started":"2022-07-31T17:30:36.867944Z","shell.execute_reply":"2022-07-31T17:30:36.871029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data = pd.read_csv('../input/stroke-blood-clot-classification/dataset.csv')\n# display(data.head())\ndata = df","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:36.872994Z","iopub.execute_input":"2022-07-31T17:30:36.873805Z","iopub.status.idle":"2022-07-31T17:30:37.187849Z","shell.execute_reply.started":"2022-07-31T17:30:36.873767Z","shell.execute_reply":"2022-07-31T17:30:37.18547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Percentage of Background Images in sliced dataset: {:.3f}'.format(data[data[\"is_background\"] == 0].count()[0] / len(data) * 100))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.188979Z","iopub.status.idle":"2022-07-31T17:30:37.190096Z","shell.execute_reply.started":"2022-07-31T17:30:37.189842Z","shell.execute_reply":"2022-07-31T17:30:37.189867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Lets visualize some images for sanity check if is_background are correct or not","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=3, ncols=6, figsize=(40,25))\nplt.suptitle(\"Background patches\", fontsize = 16)\n\nbackground_images = data[data[\"is_background\"] == 0]\n\nfor i in range(0, 3*6):\n    x = i // 6\n    y = i % 6\n    image = cv2.imread(background_images.sample(1)['filepath'].values[0])\n    axes[x, y].imshow(image, cmap=plt.cm.bone)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.191655Z","iopub.status.idle":"2022-07-31T17:30:37.192159Z","shell.execute_reply.started":"2022-07-31T17:30:37.191914Z","shell.execute_reply":"2022-07-31T17:30:37.191937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=3, ncols=6, figsize=(40,25))\nplt.suptitle(\"Non-Background patches\", fontsize = 16)\n\nroi_images = data[data[\"is_background\"] == 1]\n\nfor i in range(0, 3*6):\n    x = i // 6\n    y = i % 6\n    image = cv2.imread(roi_images.sample(1)['filepath'].values[0])\n    axes[x, y].imshow(image, cmap=plt.cm.bone)\n    axes[x, y].axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.193973Z","iopub.status.idle":"2022-07-31T17:30:37.194444Z","shell.execute_reply.started":"2022-07-31T17:30:37.194195Z","shell.execute_reply":"2022-07-31T17:30:37.194217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = data.loc[data['is_background'] != 0]\nprint(data.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.195869Z","iopub.status.idle":"2022-07-31T17:30:37.196682Z","shell.execute_reply.started":"2022-07-31T17:30:37.196428Z","shell.execute_reply":"2022-07-31T17:30:37.196452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"type_distribution = data['label'].value_counts()\nplt.figure(figsize=(20, 5))\nsns.barplot(x=type_distribution.values, y=list(id_to_label.values()), palette=colors)\nplt.title('Class Distribution', fontsize=25)\nplt.xlabel('Frequency', fontsize=25)\nplt.ylabel('Label', fontsize=25)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.19843Z","iopub.status.idle":"2022-07-31T17:30:37.198908Z","shell.execute_reply.started":"2022-07-31T17:30:37.19865Z","shell.execute_reply":"2022-07-31T17:30:37.198671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClinicDataset(Dataset):\n    def __init__(self, df, transforms=None, is_test=False):\n        self.df = df\n        self.transforms = transforms\n        self.is_test = is_test\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        image = Image.open(self.df['filepath'].values[idx])\n\n        if transforms is not None:\n            image = self.transforms(image)\n        \n        if self.is_test:\n            return image\n        \n        label = torch.tensor(self.df['label'].values[idx], dtype=torch.float)\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.200675Z","iopub.status.idle":"2022-07-31T17:30:37.201159Z","shell.execute_reply.started":"2022-07-31T17:30:37.200919Z","shell.execute_reply":"2022-07-31T17:30:37.200942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model('densenet121', pretrained=True, num_classes=1)\ntorch.nn.functional.sigmoid(model(torch.randn(4, 3, 512, 512)))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.202952Z","iopub.status.idle":"2022-07-31T17:30:37.203424Z","shell.execute_reply.started":"2022-07-31T17:30:37.203173Z","shell.execute_reply":"2022-07-31T17:30:37.203196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.classifier = nn.Sequential(\n    nn.Linear(1024, 256),\n    nn.ReLU(),\n    nn.Dropout(0.2),\n    nn.Linear(256, 1),\n    nn.Sigmoid()\n)\nmodel(torch.randn(4, 3, 512, 512))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.205206Z","iopub.status.idle":"2022-07-31T17:30:37.205676Z","shell.execute_reply.started":"2022-07-31T17:30:37.205437Z","shell.execute_reply":"2022-07-31T17:30:37.205459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The train function for every Epoch\ndef fit(model, dataset, dataloader, optim, criterion, mode='train'):\n    # Choice of training and testing mode\n    if mode == 'train':\n        model.train()\n    else:\n        model.eval()\n\n    running_loss = 0.0\n    running_corrects = 0.0\n\n    tqdm_loop = tqdm(\n        dataloader,\n        total=len(dataset) // dataloader.batch_size,\n        desc=mode, leave=True\n    )\n\n    # Loop over the dataloader and train over every batch of images\n    for i, data in enumerate(tqdm_loop):\n        # Copy data to the gpu\n        images, labels = data\n        images, labels = images.to(config['device']), labels.to(config['device'])\n        \n        # Zero the parameter gradients during training\n        if mode=='train':\n            optim.zero_grad()\n\n        # Predict classes using images from the training set\n        outputs = model(images)\n\n        # Compute the loss based on model output and real labels\n        loss = criterion(outputs.squeeze(1), labels)\n\n        # Calculate statistics\n        running_loss += loss.item()\n        running_corrects += (outputs.max(1)[1] == labels).sum().item()\n        \n        # Perform model updates according to the loss function (criterion)\n        if mode=='train':\n            # Backpropagate the loss\n            loss.backward()\n            # Adjust parameters based on the calculated gradients\n            optim.step()\n    \n    # Record the average statistics\n    epoch_loss = running_loss / dataset.__len__()\n    epoch_acc = running_corrects / dataset.__len__()\n    \n    tqdm_loop.set_postfix(loss=epoch_loss, acc=epoch_acc)\n\n    return epoch_loss, epoch_acc","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.207484Z","iopub.status.idle":"2022-07-31T17:30:37.207957Z","shell.execute_reply.started":"2022-07-31T17:30:37.207706Z","shell.execute_reply":"2022-07-31T17:30:37.207727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n0 = data.loc[data['label'] == 1].shape[0]\nn1 = data.loc[data['label'] == 1].shape[0]\n\nw0 = 1 - n0/(n0+n1)\nw1 = 1 - n1/(n0+n1)\n\nclass_weights=torch.FloatTensor([w0, w1]).to(config['device'])\n\noptim = torch.optim.AdamW(model.parameters(), lr=config['lr'])\nloss_fn = torch.nn.BCELoss()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.209715Z","iopub.status.idle":"2022-07-31T17:30:37.210194Z","shell.execute_reply.started":"2022-07-31T17:30:37.209955Z","shell.execute_reply":"2022-07-31T17:30:37.209977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, valid_df = train_test_split(data[:100], test_size=0.2)\n\nprint(train_df.shape, valid_df.shape)\n\ndata_transforms = {\n    'train': transforms.Compose([\n        transforms.Resize((512, 512)),\n#         transforms.RandomHorizontalFlip(),\n#         transforms.CenterCrop(10),\n#         transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n        transforms.ToTensor(),\n        \n    ]),\n    'valid': transforms.Compose([\n        transforms.Resize((512, 512)),\n        transforms.ToTensor(),\n    ])\n}\n\ntrain_dataset = ClinicDataset(train_df, transforms=data_transforms['train'])\nvalid_dataset = ClinicDataset(valid_df, transforms=data_transforms['valid'])\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=16, shuffle=True, drop_last=True)\nvalid_dataloader = DataLoader(valid_dataset, batch_size=16, shuffle=True, drop_last=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.211982Z","iopub.status.idle":"2022-07-31T17:30:37.212456Z","shell.execute_reply.started":"2022-07-31T17:30:37.212206Z","shell.execute_reply":"2022-07-31T17:30:37.212228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss, valid_loss = [], []\ntrain_acc, valid_acc = [], []\nbest_acc = 0\n\n# Run train loop for given number of epochs\nfor epoch in range(config['num_epochs']):\n    print(f'\\nEpoch: {epoch+1} / {config[\"num_epochs\"]}')\n    print('-' * 10)\n\n    # Train model one epoch and display statistics\n    train_epoch_loss, train_epoch_acc = fit(model.to(config['device']), train_dataset, train_dataloader, optim, loss_fn, mode='train')\n    train_loss.append(train_epoch_loss)\n    train_acc.append(train_epoch_acc)\n    print(f'Train Loss: {train_epoch_loss:.4f} | Train Acc: {train_epoch_acc:.4f}')\n\n    # Run validation for one epoch and display statistics\n    with torch.no_grad():\n        valid_epoch_loss, valid_epoch_acc = fit(model.to(config['device']), valid_dataset, valid_dataloader, optim, loss_fn, mode='valid')\n        valid_loss.append(valid_epoch_loss)\n        valid_acc.append(valid_epoch_acc)\n        print(f'Valid Loss: {valid_epoch_loss:.4f} | Valid Acc: {valid_epoch_acc:.4f}')\n    \n    if config['use_wandb']:\n        wandb.log({\n            \"train_loss\": train_epoch_loss,\n            \"valid_loss\": valid_epoch_loss,\n            \"train_acc\": train_epoch_acc,\n            \"valid_acc\": valid_epoch_acc,\n\n        })\n    \n    if valid_epoch_acc >= best_acc:\n        print(f'Model improved from {best_acc} to {valid_epoch_acc}, Saving best model...')\n        torch.save(model.state_dict(), f'efficientnet_b0_{valid_epoch_acc:.4f}.pt')\n        best_acc = valid_epoch_acc\n    \n\n# Save the model and all the metrics in a log file\n\ncheckpoint = {\n            'total_epochs'      : config['num_epochs'],\n            'state_dict'        : model.state_dict(),\n            'optimizer'         : optim.state_dict(),\n            'train_loss'        : train_loss,\n            'train_acc'         : train_acc,\n            'val_loss'          : valid_loss,\n            'val_acc'           : valid_acc,\n            }\n\ntorch.save(checkpoint, 'last_checkpoint.pt')\nprint(\"Model Saved\")","metadata":{"execution":{"iopub.status.busy":"2022-07-31T17:30:37.21538Z","iopub.status.idle":"2022-07-31T17:30:37.21594Z","shell.execute_reply.started":"2022-07-31T17:30:37.215637Z","shell.execute_reply":"2022-07-31T17:30:37.215662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}