{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.18","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":1195048,"sourceType":"datasetVersion","datasetId":680469}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{},"version_major":2,"version_minor":0}},"colab":{"provenance":[]}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<div align='center'><font size=\"5\" color='##353B47'>SIIM ISIC</font></div>\n<div align='center'><font size=\"4\" color=\"##353B47\">Skin Cancer Classification Using PyTorch</font></div>\n<br>\n<hr>\n","metadata":{"id":"W2QIR6JZpZrC"}},{"cell_type":"markdown","source":"# Introduction  \n\nThis notebook aims to develop a deep learning model for an API that enhances a **💬 Messenger chatbot** on a Facebook page. The chatbot assists users in **🔍 detecting skin cancer, 🩺 providing medical advice, and ❓ answering frequently asked questions (FAQs)**. The full project is available on GitHub: [**🌐 DermAid - AI Copilot for Skin Cancer Detection**](https://github.com/Chaimaaorg/DermAid---AI-Copilot-for-Skin-Cancer-Detection).  \n\n<div style=\"text-align: center;\">\n  <img src=\"https://github.com/Chaimaaorg/DermAid---AI-Copilot-for-Skin-Cancer-Detection/blob/master/assets/demo-pic.png?raw=true\" alt=\"Melanoma challenge\" width=\"800\"/>\n</div>\n","metadata":{"id":"-3r4T3xDdWug"}},{"cell_type":"markdown","source":"* ### **📚 Background on Skin Cancer and Melanoma**\n\n> Skin cancer is one of the most common cancers worldwide and occurs when **abnormal skin cells grow uncontrollably**, often forming **malignant tumors**. It mainly falls into two categories:\n> * **🦠 Non-melanomas (Carcinomas)**: Include **basal cell carcinoma (BCC)**, **squamous cell carcinoma (SCC)**, and **Merkel cell carcinoma (MCC)**.\n* **🎨 Melanomas**: The **most dangerous** type, developing from **melanocytes**, the pigment-producing cells responsible for skin color.\n\n> 🧬 Cancer starts when the normal process of **cell growth and death breaks down**, causing **damaged or unnecessary cells** to survive and divide uncontrollably.\n\n* ### **🔍 What is Melanoma?**\n\n> Melanoma is a **serious and aggressive form of skin cancer** that often resembles a mole and can sometimes develop from existing moles. Unlike other types of skin cancer, melanomas can appear **anywhere on the body**, including areas **not typically exposed to the sun**, such as the back, scalp, or under nails. It remains relatively **less common** than other skin cancers, yet significantly more dangerous. In **2019**, approximately **192,000 new cases** were reported in the U.S., nearly **96,000 of which were invasive**, and the disease was projected to cause around **7,200 deaths**. Fortunately, melanoma is **highly treatable when detected early**, making timely and accurate diagnosis crucial.\n\n\n  <img src=\"https://melanomaresearch.com.au/wp-content/uploads/2020/04/Hero-image-1200x860-8-1024x734.jpg\" alt=\"Melanoma challenge\" width=\"800\"/>\n\n* ### **🌞 Causes and Detection**\n\n> * The main cause is **exposure to UV radiation** from the sun or **tanning machines**, which can damage DNA and trigger mutations.\n> * Traditional diagnosis via **dermoscopy and biopsy** is effective but can be invasive and slow.\n> * Recently, **deep learning models** offer promising solutions for **fast, automated, and accurate melanoma detection**.\n","metadata":{"id":"ghqK4BtxdWuk"}},{"cell_type":"markdown","source":"* ### **🎯 Objective of This Notebook**\n\n> Skin cancer is the most prevalent cancer globally, and while **melanoma** is **less common**, it accounts for **\\~75% of skin cancer-related deaths**. Early and accurate detection significantly improves patient survival, yet diagnosis remains challenging due to the subtle visual differences between **benign and malignant** lesions.\n\n> In this notebook, we focus on developing a **🤖 CNN-based model** to support this critical task, aligning with the goals of the **SIIM-ISIC melanoma classification competition**. Specifically, we implement the **🧠 EfficientNetB1** architecture, known for its powerful feature extraction capabilities, to classify skin lesions using a dataset of **expert-annotated dermoscopic images**. The model is trained to distinguish between **✅ benign and ❌ malignant** skin lesions, helping pave the way for **AI-assisted early diagnosis**.\n\n* ### **📏 Evaluation Metric**\n    \n> The evaluation metric is the ROC-AUC score as expected in medical imaging competition, and in healthcare competition broadly speaking. Since the dataset is likely to be highly imbalanced (more non-cancerous moles than cancerous ones) we want a metric that takes into account this imbalance.\n\n* ### **🚀 Next Steps**\n\n> This initial model **📷 uses image data only**. In a future notebook, we will extend our work by building a **🌐 multimodal model** that incorporates both **🖼️ image data** and **📊 patient metadata** (e.g., age, sex, lesion location). This approach aligns with the competition's design, where **contextual patient information** can play a vital role in improving melanoma detection and **better supporting clinical dermatologists**.\n","metadata":{"id":"XOffERzzrLno"}},{"cell_type":"markdown","source":"# Requirements","metadata":{"id":"KrGE-XFIvsFl"}},{"cell_type":"code","source":"!pip install -q efficientnet_pytorch","metadata":{"id":"u0i_RVPEvyBa","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install plotly","metadata":{"id":"2xbks8uWvtL8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch_xla.core.xla_model as xm\nprint(xm.xla_device())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load libraries","metadata":{"id":"LFYEplvOvop2"}},{"cell_type":"code","source":"# 📊 Data manipulation\nimport os\nimport re\nimport glob\nimport gc\nimport random\n\nimport numpy as np\nimport pandas as pd\n\n# 📈 Visualization\nfrom sklearn.utils import resample\nimport matplotlib.pyplot as plt\nimport plotly.graph_objects as go\n\nimport seaborn as sns\n\n# 🖼️ Image processing\nimport cv2\nfrom PIL import Image\n\n# 🧠 Deep Learning - PyTorch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, lr_scheduler\nfrom torch.utils.data.sampler import WeightedRandomSampler\n\n# 🧠 Deep Learning - Models & Augmentation\nfrom efficientnet_pytorch import EfficientNet\nimport albumentations\n\n# 🧪 Evaluation & Cross-validation\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import roc_auc_score\n\n# ⚡ TPU support (PyTorch/XLA)\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\n# 📦 Progress bar\nfrom tqdm import tqdm","metadata":{"id":"zO9CmHqsvzYs","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{"id":"efm1iJubv3dW"}},{"cell_type":"code","source":"# 📁 Paths\nDATA_ROOT_PATH = \"../input/siim-isic-melanoma-classification/\"\nTRAIN_CSV = DATA_ROOT_PATH + \"train.csv\"\nTEST_CSV = DATA_ROOT_PATH + \"test.csv\"\nSUBMISSION_CSV = DATA_ROOT_PATH + \"sample_submission.csv\"\n\n# 📄 Load DataFrames\ntrain_df = pd.read_csv(TRAIN_CSV, na_values=['unknown'])\ntest_df = pd.read_csv(TEST_CSV)\n\n# 🖼️ Image settings\nWIDTH = 224\nHEIGHT = 224\nMEAN = (0.485, 0.456, 0.406)\nSTD = (0.229, 0.224, 0.225)\n\n# ⚙️ Training hyperparameters\nTRAIN_BATCH_SIZE = 128\nVALID_BATCH_SIZE = 128\nEPOCHS = 10\nLR = 1e-3\nFOLDS = 5\nSEED = 0\nVERBOSE_STEP = 1","metadata":{"id":"K-Z9Hw7Hv6nd","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Exploratory Data Analysis","metadata":{"id":"xtUdazWhdWum"}},{"cell_type":"markdown","source":"## EDA - Dataset Overview","metadata":{"id":"utXV-bBGyGs7"}},{"cell_type":"code","source":"trn_len_df = len(train_df)\ntst_len_df = len(test_df)\nprint(f\"There are {trn_len_df} images in the training set\")\nprint(f\"There are {tst_len_df} images in the test set\")","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_kg_hide-input":true,"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","id":"qyqHJ-A4dWuo","outputId":"357834c5-2d9b-4e51-96bc-6fc35700f74b","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.head()","metadata":{"_kg_hide-input":true,"id":"J1c0jZhodWup","outputId":"d8e8a218-9b8c-48eb-94b0-208c018cfda9","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df.head()","metadata":{"_kg_hide-input":true,"id":"ef1AZcAUdWup","outputId":"b828a462-fdd4-465d-807a-3eb52c6f9284","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df['target'].value_counts()","metadata":{"id":"vP7vpmhazGiy","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<ul>\n  <li>The dataset contains 33,126 entries.</li>\n  <li>The <b>target</b> variable is highly imbalanced, with only <b>1.76%</b> of cases labeled as malignant.</li>\n  <li>Patients can have multiple samples, so it is crucial to <b>avoid data leakage</b> by ensuring that the same patient does not appear in both the training and validation sets; otherwise, the model would unfairly benefit from prior knowledge during evaluation.</li>\n  <li>The dataset provides <b>meta-features</b> that can be used to enrich our model.</li>\n</ul>\n<p>Let’s now explore these meta-features and examine their distribution.</p>","metadata":{"id":"KXZQYfsxdWup"}},{"cell_type":"markdown","source":"## EDA - Contextual features\n\nBefore examining the distributions of the meta-features, let’s first understand what each of these features represents:\n\n* **image\\_name**: the filename associated with each image or TFRecord.\n* **patient\\_id**: a unique identifier assigned to each patient.\n* **sex**: the biological sex of the patient (left blank if unknown).\n* **age\\_approx**: an estimated age of the patient.\n* **anatom\\_site\\_general\\_challenge**: the anatomical location where the image was taken.\n* **diagnosis**: information indicating whether the lesion is malignant.\n* **benign\\_malignant**: a binary variable representing the malignancy status, corresponding to the target label.\n","metadata":{"id":"jR8nPNqFdWup"}},{"cell_type":"code","source":"plt.figure(figsize=(10,6))\nsns.heatmap(train_df.isnull(), cbar=False, cmap='viridis', yticklabels=False)\nplt.title('Missing Data Heatmap', fontsize=16)\nplt.xlabel('Columns')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# NaN values\nnan_stats = train_df.isna().sum() / len(train_df) * 100\n\nstats = pd.DataFrame({\n    'columns': train_df.columns,\n    'NaN statistics (in %)': nan_stats\n})\n\nstats = stats.sort_values(by=['NaN statistics (in %)'], ascending=False)\n\nstats = stats.reset_index(drop=True)\nstats.head(5)","metadata":{"_kg_hide-input":true,"id":"BZF06GP3dWup","outputId":"33c31671-3265-42de-ac66-053af9093fb0","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>Key observations:</p>\n<ul>\n    <li>The dataset contains very few missing (NaN) values overall.</li>\n    <li>The <code>diagnosis</code> column has a high proportion of missing entries, likely because the diagnoses are not known for most cases.</li>\n    <li>For the <code>anatom_site_general_challenge</code> column, we can introduce a new category labeled <i>unknown</i> to handle missing values.</li>\n</ul>\n","metadata":{"id":"ydxsaT2zdWup"}},{"cell_type":"code","source":"# Patient distribution\nprint(f\"There are {train_df['patient_id'].nunique()} unique patients for {len(train_df)} images in the training set.\")\nprint(f\"There are {test_df['patient_id'].nunique()} unique patients for {len(test_df)} images in the training set.\")\n\nfig, ax = plt.subplots(1, 2, figsize=(15, 10))\n\nsns.countplot(x='patient_id',data=train_df, ax=ax[0],palette='viridis')\nax[0].set_title('Patient distribution in the training set')\n\nsns.countplot(x='patient_id',data=test_df, ax=ax[1],palette='viridis')\nax[1].set_title('Patient distribution in the test set')\n\nplt.show()","metadata":{"_kg_hide-input":true,"id":"H0V0Z8kddWuq","outputId":"d772f1df-e04d-43f7-b963-83bcf2156246","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>Key observations:</p>\n<ul>\n    <li>The dataset contains a relatively small number of unique patients.</li>\n    <li>Image distribution is imbalanced — some patients have over 40 images, which increases the risk of overfitting to those individuals.</li>\n    <li>Both the training and test sets exhibit similar patterns, with certain patients contributing a large number of images.</li>\n</ul>\n\n<p>Next, we'll check whether any patients appear in both the training and test sets simultaneously.</p>","metadata":{"id":"03h5FvURdWuq"}},{"cell_type":"code","source":"trn_patients = set(train_df['patient_id'])\ntst_patients = set(test_df['patient_id'])\n\ninter_patients = len(trn_patients.intersection(tst_patients))\n\nprint(f'There are {inter_patients} common patients in the training and test sets.')","metadata":{"_kg_hide-input":true,"id":"78BSaY9sdWuq","outputId":"59363538-af3b-46f6-d4d7-beea21d237f3","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(10, 5))\n\nsns.countplot(x='sex',data=train_df, ax=ax[0],palette='viridis')\nax[0].set_title(\"Sex distribution in the training set\")\n\nsns.countplot(x='sex',data=test_df, ax=ax[1],palette='viridis')\nax[1].set_title(\"Sex distribution in the test set\")\n\nplt.show()","metadata":{"_kg_hide-input":true,"id":"iL5BUEUYdWuq","outputId":"c0b0d287-99d3-42c1-a9b2-96fea1be4fa6","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>Observations:</p>\n<ul>\n    <li>Both datasets contain more images of male patients than female. We’ll later explore whether gender has any correlation with the target variable.</li>\n</ul>\n","metadata":{"id":"sd8uikTkdWuq"}},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(10, 5))\n\nsns.histplot(data=train_df, x='age_approx', ax=ax[0], kde=True, color='skyblue')\nax[0].set_title(\"Age distribution in the training set\")\n\nsns.histplot(data=test_df, x='age_approx', ax=ax[1], kde=True, color='salmon')\nax[1].set_title(\"Age distribution in the test set\")\n\nplt.tight_layout()\nplt.show()","metadata":{"_kg_hide-input":true,"id":"hJk3h8g_dWur","outputId":"0de1f9ab-fb50-40f1-c67f-a7110347a1e4","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>We observe the following:</p>\n<ul> \n    <li>The age distribution is similar in both the training and test sets.</li> \n    <li>Interestingly, despite the training and test sets containing images from 2056 and 690 patients respectively, the age values are quite limited — with only 18 distinct ages represented in the test set.</li>\n</ul>","metadata":{"id":"w8LN5PmfdWur"}},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(10, 5))\n\nchart = sns.countplot(x='anatom_site_general_challenge',data=train_df, ax=ax[0],palette='viridis')\nax[0].set_title(\"Anatomical site in the training set\")\nchart.set_xticklabels(chart.get_xticklabels(), rotation=45)\n\nchart2 = sns.countplot(x='anatom_site_general_challenge',data=test_df, ax=ax[1],palette='viridis')\nax[1].set_title(\"Anatomical site in the test set\")\nchart2.set_xticklabels(chart2.get_xticklabels(), rotation=45)\n\nplt.show()","metadata":{"_kg_hide-input":true,"id":"o2NtrSmmdWur","outputId":"bbcfd15c-edfa-4312-b324-a98a2630f4b4","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>It is observed that:</p>\n<ul> \n    <li>The majority of images are taken from the torso and limbs (such as feet, hands, and arms)</li>\n    <li>The test set shares a similar distribution with the training set.</li> \n</ul>\n","metadata":{"id":"PCxX_8LedWur"}},{"cell_type":"code","source":"chart3 = sns.countplot(train_df['diagnosis'])\nchart3.set_xticklabels(chart3.get_xticklabels(), rotation=45)\n\nplt.show()","metadata":{"_kg_hide-input":true,"id":"PnGpR9t2dWur","outputId":"0400d1b6-3284-4164-f98e-7c26c408ae2f","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>Observations:</p>\n<ul>\n    <li><b>The majority of values are NaN, although these are not shown in the plot.</b></li>\n    <li>The distribution is once again dominated by 'unknown' and 'nevus' labels.</li>\n    <li>This data is not provided in the test set.</li>\n</ul>\n","metadata":{"id":"Oia1VeTtdWut"}},{"cell_type":"markdown","source":"## EDA - Target distribution","metadata":{"id":"PJm3rINNdWut"}},{"cell_type":"code","source":"print(f\"There are {len(train_df[train_df['target'] == 0])} negative labels.\")\nprint(f\"There are {len(train_df[train_df['target'] == 1])} positive labels.\")\n\nsns.countplot(x='target', data=train_df, palette='viridis') \nplt.show()","metadata":{"_kg_hide-input":true,"id":"N0ATZDvcdWut","outputId":"a2de5a8b-5b5c-4abf-e076-5b7e7a78db66","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<p>As expected the data is highly imbalanced like in most medical datasets. We will need to come up with some forms of tricks to prevent a model from always predicting negative. Example: weighted loss, ...</p>","metadata":{"id":"7DzzCzQLdWut"}},{"cell_type":"markdown","source":"## EDA - Correlation\n\n<p>With a clearer view of the feature distributions, we can now explore how both continuous and categorical variables relate to the target variable.</p>\n<b>Continuous feature:</b>\n\n<ul> <li>Age</li> </ul>\n<b>Categorical features:</b>\n\n<ul> <li>Sex</li> <li>Anatomical site</li> </ul>","metadata":{"id":"v7vKs2OldWuu"}},{"cell_type":"code","source":"fig = plt.figure(figsize=(7,5))\nax = sns.countplot(x=\"target\", hue=\"sex\", data=train_df)\n\nfor p in ax.patches:\n    height = p.get_height()\n    ax.text(p.get_x() + p.get_width()/2, height+10, '{:1.2f}%'.format(100*height/len(train_df)), ha=\"center\")","metadata":{"_kg_hide-input":true,"id":"tvg5Z_FWdWuu","outputId":"4884410d-03db-4d61-ee3a-12c077b2a74b","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = plt.figure(figsize=(7,5))\nax = sns.countplot(x=\"target\", hue=\"anatom_site_general_challenge\", data=train_df)\n\nfor p in ax.patches:\n    height = p.get_height()\n    ax.text(p.get_x() + p.get_width()/2, height+15, '{:1.2f}%'.format(100*height/len(train_df)), ha=\"center\")","metadata":{"_kg_hide-input":true,"id":"IH03kFC_dWuu","outputId":"96cd1610-e43b-4e4e-eca4-90236ab2ef49","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Categorizing age\ntrain_df['age_cat'] = '0 / 20 years'\ntrain_df.loc[(train_df['age_approx'] > 20) & (train_df['age_approx'] <= 40), 'age_cat'] = '20 / 40 years'\ntrain_df.loc[(train_df['age_approx'] > 40) & (train_df['age_approx'] <= 60), 'age_cat'] = '40 / 60 years'\ntrain_df.loc[(train_df['age_approx'] > 60) & (train_df['age_approx'] <= 80), 'age_cat'] = '60 / 80 years'\ntrain_df.loc[(train_df['age_approx'] > 80), 'age_approx'] = '80+ years'","metadata":{"id":"QyZBiPYTLZNN","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sex_dim = go.parcats.Dimension(values=train_df[\"sex\"], label=\"Sex\")\nage_dim = go.parcats.Dimension(values=train_df[\"age_cat\"], label=\"Age Category\")\nsite_dim = go.parcats.Dimension(values=train_df[\"anatom_site_general_challenge\"], label=\"Anatomical Site\")\ndiagnosis_dim = go.parcats.Dimension(values=train_df[\"diagnosis\"], label=\"Diagnosis\")\ntarget_dim = go.parcats.Dimension(values=train_df[\"target\"], label=\"Target\")\n\ncolor = train_df[\"target\"]\ncolorscale = [[0, '#83a79a'], [1, 'goldenrod']]\n\n# Créer la figure avec go.Parcats\nfig = go.Figure(data=[\n    go.Parcats(\n        dimensions=[sex_dim, age_dim, site_dim, diagnosis_dim,target_dim],\n        line={'color': color, 'colorscale': colorscale},\n        hoveron='color',\n        hoverinfo='count+probability',\n        labelfont={'size': 14},\n        tickfont={'size': 12},\n        arrangement='freeform'\n    )\n])\n\nfig.update_layout(title='Parallel Categories Diagram for Trainset')\nfig.show(renderer='iframe')","metadata":{"id":"5fm-SrmoLZ6-","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The parallel category diagram provides a comprehensive visualization of the relationships between demographic, anatomical, and diagnostic features in the dataset, along with their association with malignancy (target variable). The majority of cases are found in males and individuals aged 40 to 80 years, with the torso, lower extremity, and upper extremity being the most common anatomical sites. The dataset is predominantly composed of benign diagnoses, particularly nevus, while malignant cases (target = 1), shown in red, appear less frequently but are notably associated with melanoma. Additionally, malignancies tend to be more prevalent in older age groups and specific anatomical sites, suggesting a correlation between age, lesion location, and cancer risk. This visualization highlights the imbalanced nature of the dataset, where malignant cases are underrepresented, which is crucial for model training and classification performance.","metadata":{"id":"XVLFdRtDLoBP"}},{"cell_type":"code","source":"## Separate minority and majority class\ntrain_no_target = train_df[train_df['target']==0]\ntrain_target = train_df[train_df['target']==1]","metadata":{"id":"1d7U39koMQRM","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Downsampling majority class\ndf_majority_downsampled = resample(train_no_target,\n                                   replace=False, ## sample without replacement\n                                   n_samples=584, ## to match minority class\n                                   random_state=42)\n\n## Combine minority class with downsampled majority class\ntrain_downsampled = pd.concat([df_majority_downsampled, train_target])","metadata":{"id":"ZMqRuLFeMR0i","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sex_dim = go.parcats.Dimension(values=train_downsampled[\"sex\"], label=\"Sex\")\nage_dim = go.parcats.Dimension(values=train_downsampled[\"age_cat\"], label=\"Age Category\")\nsite_dim = go.parcats.Dimension(values=train_downsampled[\"anatom_site_general_challenge\"], label=\"Anatomical Site\")\ndiagnosis_dim = go.parcats.Dimension(values=train_downsampled[\"diagnosis\"], label=\"Diagnosis\")\ntarget_dim = go.parcats.Dimension(values=train_downsampled[\"target\"], label=\"Target\")\n\ncolor = train_downsampled[\"target\"]\ncolorscale = [[0, '#83a79a'], [1, 'goldenrod']]\n\n# Créer la figure avec go.Parcats\nfig = go.Figure(data=[\n    go.Parcats(\n        dimensions=[sex_dim, age_dim, site_dim, diagnosis_dim,target_dim],\n        line={'color': color, 'colorscale': colorscale},\n        hoveron='color',\n        hoverinfo='count+probability',\n        labelfont={'size': 14},\n        tickfont={'size': 12},\n        arrangement='freeform'\n    )\n])\n\nfig.update_layout(title='Parallel Categories Diagram for Trainset')\nfig.show(renderer='iframe')","metadata":{"id":"qggz0vqjMUe1","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The diagram reveals that both sexes and a wide range of age groups are represented among both positive and negative cases, though there appears to be a higher prevalence of positive cases in the 40–60 and 60–80 age categories. Melanoma cases are strongly associated with the positive target, while diagnoses such as nevus and lentigo NOS are more often associated with negative cases. Anatomical site also shows variation: cases involving the torso and lower extremities appear across both target values.","metadata":{"id":"AwgOECsjMWR1"}},{"cell_type":"markdown","source":"## EDA - Image visualization\n\n<p>Now, let's visualize some images from the training and the test sets to see any noticeable differences betweeen both sets.</p>","metadata":{"id":"HVoc6WmHdWuu"}},{"cell_type":"code","source":"# Training set\n\nimg_names = glob.glob('../input/siim-isic-melanoma-classification/jpeg/train/*.jpg')\n\nfig, ax = plt.subplots(4, 4, figsize=(20, 20))\n\nfor i in range(16):\n    x = i // 4\n    y = i % 4\n\n    path = img_names[i]\n    image_id = path.split(\"/\")[5][:-4]\n\n    target = train_df.loc[train_df['image_name'] == image_id, 'target'].tolist()[0]\n\n    img = Image.open(path)\n\n    ax[x, y].imshow(img)\n    ax[x, y].axis('off')\n    ax[x, y].set_title(f'ID: {image_id}, Target: {target}')\n\nfig.suptitle(\"Training set samples\", fontsize=15)","metadata":{"_kg_hide-input":true,"id":"3-d55hPNdWuu","outputId":"0acd3ff7-852c-4316-d2d8-c360e6f159b8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test set\n\nimg_names = glob.glob('../input/siim-isic-melanoma-classification/jpeg/test/*.jpg')\n\nfig, ax = plt.subplots(4, 4, figsize=(20, 20))\n\nfor i in range(16):\n    x = i // 4\n    y = i % 4\n\n    path = img_names[i]\n    image_id = path.split(\"/\")[5][:-4]\n\n    img = Image.open(path)\n\n    ax[x, y].imshow(img)\n    ax[x, y].axis('off')\n    ax[x, y].set_title(f'ID: {image_id}')\n\nfig.suptitle(\"Test set samples\", fontsize=15)","metadata":{"_kg_hide-input":true,"id":"5HscSb5LdWuv","outputId":"2b0f6294-38c3-4584-92e2-b28b8264b9af","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We note that the pictures in the training and test sets appear to be high quality, although some may contain obstacles like hairs or artifacts such as \"mm\" markings or graduations.\n","metadata":{"id":"tivn8Iy6dWuv"}},{"cell_type":"markdown","source":"# Dataset Preparation\n\n<p>Based on the insights from the EDA, we can outline some important considerations for modeling.</p>\n<ul>\n  <li>First, it is essential to ensure that patients do not appear in both the training and validation sets simultaneously.</li>\n  <li>Second, due to the strong class imbalance in the data, we should apply techniques such as weighted loss functions or focal loss to address this issue.</li>\n</ul>\n","metadata":{"id":"xftYz_sKdWuz"}},{"cell_type":"markdown","source":"### Utility functions and classes","metadata":{}},{"cell_type":"markdown","source":"* This function ensures reproducibility by setting the random seed across Python, NumPy, and PyTorch environments.","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed):\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    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)","metadata":{"id":"i43xownudWu1","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n* Now let's define a utility class to keep track of and update loss values during training. ","metadata":{}},{"cell_type":"code","source":"class AverageMeter:\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"id":"u2WqmTx4dWu1","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Custom PyTorch Dataset class","metadata":{}},{"cell_type":"markdown","source":"Let's define a custom PyTorch Dataset class for loading melanoma images, applying resizing and augmentations, and preparing the data in the format required for model training.","metadata":{}},{"cell_type":"code","source":"class BaseMelanomaDataset:\n    def __init__(self, image_paths, resize=True, augmentations=None):\n        self.image_paths = image_paths\n        self.augmentations = augmentations\n        self.resize = resize\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def load_image(self, idx):\n        image = Image.open(self.image_paths[idx])\n        if self.resize:\n            image = image.resize((WIDTH, HEIGHT), resample=Image.BILINEAR)\n        image = np.array(image)\n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image)\n            image = augmented['image']\n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        return torch.tensor(image, dtype=torch.float)","metadata":{"id":"va6_GLyXdWu1","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MelanomaTrainDataset(BaseMelanomaDataset):\n    def __init__(self, image_paths, targets, resize=True, augmentations=None):\n        super().__init__(image_paths, resize, augmentations)\n        self.targets = targets\n\n    def __getitem__(self, idx):\n        image = self.load_image(idx)\n        target = torch.tensor(self.targets[idx], dtype=torch.long)\n        return {'image': image, 'targets': target}\n\nclass MelanomaTestDataset(BaseMelanomaDataset):\n    def __getitem__(self, idx):\n        image = self.load_image(idx)\n        return {'image': image}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We need to the image_name column to full image paths so they can be used by the MelanomaDataset class.","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(TRAIN_CSV)\ndf['image_name'] = df['image_name'].apply(lambda x: f'../input/siic-isic-224x224-images/train/{x}.png')\ndf.head()","metadata":{"id":"CGYoFyYqdWu1","outputId":"21c7c76b-ecb9-4d6c-dc16-65e6219721cd","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generating folds\nkf = KFold(FOLDS,shuffle=True, random_state=SEED)\ndf = df.sample(frac=1).reset_index(drop=True)\n\nfor f, (_, val_index) in enumerate(kf.split(df, df)):\n    df.loc[val_index, 'kfold'] = f\n\nprint(df['kfold'].value_counts())","metadata":{"id":"ZBVb8OmpdWu1","outputId":"d4c064ff-ecb5-4bf6-cf0c-a3b875d0bba2","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset Transformation","metadata":{"id":"UmK5IqI5NPOg"}},{"cell_type":"markdown","source":"We will need to apply some transformations on the train dataset in order to improve the model's **generalization** and **robustness** by simulating real-world variations in medical images, especially skin lesions, without altering the underlying pathology:\n* **ShiftScaleRotate** (`p=0.9`) simulates changes in image position, scale, and rotation — useful since lesion position/orientation varies naturally.\n* **CLAHE** (`p=0.5`) enhances local contrast, which can help highlight lesion boundaries and details in dermoscopic images.\n* **HorizontalFlip & VerticalFlip** (`p=0.5`) mirror images to account for symmetry in lesions and to reduce positional bias.\n* **RandomBrightnessContrast** (`p=0.9`) simulates lighting variation — helpful for dealing with images taken under different conditions.\n* **Normalize** applies mean and standard deviation normalization to standardize pixel values, which helps training convergence and stability.","metadata":{}},{"cell_type":"code","source":"train_transforms = albumentations.Compose([\n    albumentations.ShiftScaleRotate(p=0.9),\n    albumentations.CLAHE(p=0.5),\n    albumentations.HorizontalFlip(p=0.5),\n    albumentations.VerticalFlip(p=0.5),\n    albumentations.RandomBrightnessContrast(p=0.9),\n    albumentations.Normalize(mean=MEAN, std=STD, always_apply=True)\n])\n\ntest_transforms = albumentations.Compose([\n    albumentations.Normalize(mean=MEAN, std=STD, always_apply=True)\n])","metadata":{"id":"atKm7oU-dWu2","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Modeling","metadata":{"id":"vUoNdlKzdWu2"}},{"cell_type":"markdown","source":"After experimenting with **SEResNeXt50\\_32x4d** and **EfficientNet-B0** in older versions of this notebook, both models showed signs of **overfitting** or **underperformance** due to either excessive complexity or limited capacity. In contrast, the current **MelanomaModel**, based on **EfficientNet-B1**, provided a better balance between depth, resolution, and regularization. Its compound scaling strategy, combined with dropout and a lightweight custom head, allowed it to generalize better on the melanoma classification task. This architecture captured subtle image features effectively while avoiding overfitting, leading to improved performance.","metadata":{}},{"cell_type":"code","source":"class MelanomaModel(nn.Module):\n    def __init__(self):\n        super(MelanomaModel, self).__init__()\n\n        self.encoder = EfficientNet.from_pretrained(\"efficientnet-b1\")\n        self.dropout = nn.Dropout(0.3)\n        self.head = nn.Linear(1280, 1)\n\n    def forward(self, image):\n        batch_size, _, _, _ = image.shape\n\n        x = self.encoder.extract_features(image)\n        x = F.adaptive_avg_pool2d(x, 1).reshape(batch_size, -1)\n\n        x = self.dropout(x)\n        logit = self.head(x)\n\n        return logit","metadata":{"id":"eHkoEyQsdWu2","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss Function","metadata":{}},{"cell_type":"markdown","source":"The **Focal Loss** is a modified version of binary cross-entropy loss that focuses more on *hard examples* (misclassified or uncertain ones) and down-weights *easy examples*.\n* In **imbalanced classification tasks**, such as melanoma detection (where positive cases are rare), standard **binary cross-entropy** tends to be dominated by easy/majority class examples.\n* Focal Loss **down-weights easy examples** and **focuses learning on hard, misclassified samples**.\n* The **focusing parameter $\\gamma$** reduces the loss contribution from well-classified examples (where $p_t$ is high), allowing the model to learn more from difficult cases.\n\nFor **binary classification**, given:\n\n* $y \\in \\{0, 1\\}$: the true label\n* $p \\in [0, 1]$: the predicted probability for the class $y = 1$\n* $\\alpha \\in [0,1]$: balancing factor (optional, usually set to 1)\n* $\\gamma \\geq 0$: focusing parameter (controls how much to focus on hard examples)\n\nThe **Focal Loss** is defined as:\n\n$$\n\\text{FL}(p, y) = -\\alpha \\cdot (1 - p_t)^\\gamma \\cdot \\log(p_t)\n$$\n\nWhere:\n\n$$\np_t = \\begin{cases}\np & \\text{if } y = 1 \\\\\n1 - p & \\text{if } y = 0\n\\end{cases}\n$$","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=1, gamma=2, logits=False, reduce=True):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.logits = logits\n        self.reduce = reduce\n\n    def forward(self, inputs, targets):\n        if self.logits:\n            BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduce=None)\n        else:\n            BCE_loss = F.binary_cross_entropy(inputs, targets, reduce=None)\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n\n        if self.reduce:\n            return torch.mean(F_loss)\n        else:\n            return F_loss","metadata":{"id":"YScYlwILdWu3","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def loss_fn(outputs, targets):\n    return FocalLoss(logits=True)(outputs, targets)","metadata":{"id":"JT_MZxIddWu3","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Weighted Sampler\n\n* Assigning **higher weights** to samples from the **minority class**.\n* Ensuring that during **training with a `WeightedRandomSampler`**, each class is sampled more equally. This helps the model **see more of the rare class during training**, improving recall/sensitivity for minority classes like melanoma.\n\nGiven:\n\n* $B \\in [0, 1]$: a balancing factor (here $B = 0.5$, meaning equal importance to both classes)\n* $N_1$: number of positive samples ($y = 1$)\n* $N_0$: number of negative samples ($y = 0$)\n* $C_1 = 2B$, $C_0 = 2(1 - B)$: total class weights (sums to 2)\n\nWe define a **weight function** $w(y)$ for each sample with label $y \\in \\{0, 1\\}$:\n\n$$\nw(y) = \n\\begin{cases}\n\\frac{C_0}{N_0} & \\text{if } y = 0 \\\\\n\\frac{C_1}{N_1} & \\text{if } y = 1\n\\end{cases}\n$$\n\nThen, the final **sample weight vector** is:\n\n$$\n\\text{weights} = [w(y_i) \\text{ for } y_i \\in \\text{df.target}]\n$$\nThe **parameter $B$** controls the desired balance:\n\n* $B = 0.5$: equal class importance\n* $B > 0.5$: prioritize positive class more\n* $B < 0.5$: prioritize negative class more","metadata":{}},{"cell_type":"code","source":"# Deriving weights for sampler\ndef generate_weights(df):\n    B = 0.5\n\n    C = np.array([B, (1 - B)])*2\n    ones = len(df.query('target == 1'))\n    zeros = len(df.query('target == 0'))\n\n    weightage_fn = {0: C[1]/zeros, 1: C[0]/ones}\n    return [weightage_fn[target] for target in df.target]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training strategy","metadata":{}},{"cell_type":"markdown","source":"A standard training strategy commonly used for classification tasks, involving mini-batch processing, forward and backward propagation, optimizer updates, and optional learning rate scheduling.","metadata":{}},{"cell_type":"code","source":"def train_fn(data_loader, model, optimizer, device, scheduler=None):\n    model.train()\n\n    losses = AverageMeter()\n\n    tk0 = tqdm(data_loader, total=len(data_loader))\n\n    for bi, d in enumerate(tk0):\n        images = d['image']\n        targets = d['targets']\n\n        images = images.to(device, dtype=torch.float)\n        targets = targets.to(device, dtype=torch.long)\n\n        model.zero_grad()\n        outputs = model(images)\n        targets = targets.view(-1, 1).type_as(outputs)\n\n        loss = loss_fn(outputs, targets)\n        loss.backward()\n        xm.optimizer_step(optimizer, barrier=True)\n\n        if scheduler:\n            scheduler.step()\n\n        losses.update(loss.item(), images.size(0))\n\n        tk0.set_postfix(loss=losses.avg)\n    return losses.avg  ","metadata":{"id":"Cbe8IJFjdWu3","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation Strategy","metadata":{}},{"cell_type":"markdown","source":"Standard evaluation loop that computes the average loss and collects predictions without updating model weights.","metadata":{}},{"cell_type":"code","source":"def eval_fn(data_loader, model, device):\n    model.eval()\n\n    losses = AverageMeter()\n    final_preds = []\n\n    with torch.no_grad():\n        tk0 = tqdm(data_loader, total=len(data_loader))\n\n        for bi, d in enumerate(tk0):\n            images = d['image']\n            targets = d['targets']\n\n            images = images.to(device, dtype=torch.float)\n            targets = targets.to(device, dtype=torch.long)\n\n            outputs = model(images)\n            targets = targets.view(-1, 1).type_as(outputs)\n\n            loss = loss_fn(outputs, targets)\n            losses.update(loss.item(), images.size(0))\n\n            final_preds.extend(outputs.cpu().detach().numpy().tolist())\n\n    return losses.avg, final_preds","metadata":{"id":"R4Hvwhc4dWu4","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Plotting learning curves ","metadata":{}},{"cell_type":"code","source":"def plot_learning_curves(history, fold):\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n\n    fig, ax1 = plt.subplots(figsize=(10, 5))\n\n    # Plot loss\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Loss', color='tab:red')\n    ax1.plot(epochs, history['train_loss'], label='Train Loss', color='tab:red', linestyle='-')\n    ax1.plot(epochs, history['val_loss'], label='Validation Loss', color='tab:red', linestyle='--')\n    ax1.tick_params(axis='y', labelcolor='tab:red')\n\n    # Plot AUC on secondary axis\n    ax2 = ax1.twinx()\n    ax2.set_ylabel('Validation ROC AUC', color='tab:blue')\n    ax2.plot(epochs, history['val_auc'], label='Validation ROC AUC', color='tab:blue', linestyle='-')\n    ax2.tick_params(axis='y', labelcolor='tab:blue')\n\n    # Combine legends\n    lines_1, labels_1 = ax1.get_legend_handles_labels()\n    lines_2, labels_2 = ax2.get_legend_handles_labels()\n    ax1.legend(lines_1 + lines_2, labels_1 + labels_2, loc='center right')\n\n    plt.title(f'Learning Curves for Fold {fold}')\n    plt.show()","metadata":{"id":"j5Hra4PvQgFa","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"Now, let's bring everything together","metadata":{}},{"cell_type":"code","source":"def run_fold(fold):\n    device = xm.xla_device()\n    model = MelanomaModel().to(device)\n    best_auc = 0\n\n    # Selecting fold\n    train_df = df[df['kfold'] != fold].reset_index(drop=True)\n    valid_df = df[df['kfold'] == fold].reset_index(drop=True)\n\n    weights = generate_weights(train_df)\n\n    # Loading data\n\n    train_dataset = MelanomaTrainDataset(\n        image_paths=train_df['image_name'],\n        targets=train_df['target'],\n        resize=True,\n        augmentations=train_transforms,\n    )\n\n    valid_dataset = MelanomaTrainDataset(\n        image_paths=valid_df['image_name'],\n        targets=valid_df['target'],\n        resize=True,\n        augmentations=test_transforms,\n     )\n\n    train_sampler = WeightedRandomSampler(weights, len(train_df))\n\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=TRAIN_BATCH_SIZE,\n        sampler=train_sampler,\n        num_workers=8\n    )\n\n    valid_loader = torch.utils.data.DataLoader(\n        valid_dataset,\n        batch_size=VALID_BATCH_SIZE,\n        shuffle=False,\n        num_workers=8\n    )\n\n    # Optimizer and scheduler\n\n    num_train_steps = int(len(train_df) / TRAIN_BATCH_SIZE * EPOCHS)\n    optimizer = torch.optim.Adam(model.parameters(), lr=LR)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=num_train_steps,\n        eta_min=1e-6\n    )\n\n    # Training loop\n    history = {\n        'train_loss': [],\n        'val_loss': [],\n        'val_auc': [],\n    }\n    for epoch in range(EPOCHS):\n        train_loss = train_fn(train_loader, model, optimizer, device=device, scheduler=scheduler)\n        val_loss, y_pred = eval_fn(valid_loader, model, device=device)\n\n        y_pred = np.array(y_pred)\n        val_auc = roc_auc_score(valid_df['target'].values, y_pred)\n\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        history['val_auc'].append(val_auc)\n        xm.master_print(f\"Epoch = {epoch}, val_loss = {val_loss}, val_auc = {val_auc}\")\n\n        if val_auc > best_auc:\n            xm.save(model.state_dict(), f\"model_{fold}.bin\")\n            xm.master_print('Validation score improved ({} --> {}). Saving model!'.format(best_auc, val_auc))\n            best_auc = val_auc\n    plot_learning_curves(history, fold)","metadata":{"id":"e8KqVaSPdWu4","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_fold(0)","metadata":{"id":"uYwj3O-ddWu5","outputId":"6d5b01fa-2e81-4e16-91fd-b183a94fa636","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_fold(1)","metadata":{"id":"T1l4z0U0dWu5","outputId":"8a6344af-75b3-412d-bbbf-4c0ab44e1cd2","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_fold(2)","metadata":{"id":"8pEe3OMqdWu5","outputId":"4adefd92-ba2f-43c6-cff1-8ceea9fa6d21","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_fold(3)","metadata":{"id":"9YM2P6h0dWu5","outputId":"4b660734-13ac-42e4-ea56-5a5d31fee575","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_fold(4)","metadata":{"id":"j40FoWKEdWu5","outputId":"b6a34b3d-4ba0-48ef-efcd-3b168ad08993","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{"id":"F5_JOQ5sdWu5"}},{"cell_type":"code","source":"test_df = pd.read_csv(TEST_CSV)\ntest_df['image_name'] = test_df['image_name'].apply(lambda x: f'../input/siic-isic-224x224-images/test/{x}.png')\ntest_df.head()","metadata":{"id":"GB4nDCvodWu6","outputId":"d7b4406f-7fe1-4c07-c946-0e8bd2fd6202","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Test-Time Augmentation (TTA)\nAlthough augmentation is typically used during training, applying augmentations at inference time—known as test-time augmentation—helps improve model robustness and accuracy. By generating multiple transformed versions of each test image and averaging their predictions, TTA reduces model bias and better handles variations in the data, leading to improved performance in competition settings.\n\n","metadata":{}},{"cell_type":"code","source":"test_aug_transforms = albumentations.Compose([\n\n    albumentations.ShiftScaleRotate(p=0.9),\n\n    albumentations.OneOf([\n\n        albumentations.CLAHE(p=0.5),\n        albumentations.HueSaturationValue(p=0.5),\n\n    ]),\n\n    albumentations.OneOf([\n\n        albumentations.HorizontalFlip(p=0.5),\n        albumentations.VerticalFlip(p=0.5),\n\n    ]),\n\n    albumentations.RandomBrightnessContrast(p=0.9),\n\n    albumentations.Normalize(mean=MEAN, std=STD, always_apply=True)\n\n])","metadata":{"id":"aNpRI4a9dWu6","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Simple Transformation","metadata":{}},{"cell_type":"code","source":"test_transforms = albumentations.Compose([\n    albumentations.Normalize(mean=MEAN, std=STD, always_apply=True)\n])","metadata":{"id":"N6yV05LUdWu6","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We will average the predictions obtained from the original images and their augmented versions.\n\n","metadata":{}},{"cell_type":"markdown","source":"### Prediction function","metadata":{}},{"cell_type":"code","source":"def predict(data_loader, model, device):\n    model.eval()\n\n    final_preds = []\n\n    with torch.no_grad():\n        tk0 = tqdm(data_loader, total=len(data_loader))\n\n        for bi, d in enumerate(tk0):\n            images = d['image']\n\n            images = images.to(device, dtype=torch.float)\n\n            outputs = model(images)\n\n            final_preds.extend(outputs.cpu().detach().numpy().tolist())\n\n    return final_preds","metadata":{"id":"TTZMy6jhdWu7","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset Preparation","metadata":{}},{"cell_type":"code","source":"test_dataset = MelanomaTestDataset(\n    image_paths=test_df['image_name'],\n    resize=True,\n    augmentations=test_transforms,\n)\n\ntest_loader = torch.utils.data.DataLoader(\n    test_dataset,\n    batch_size=VALID_BATCH_SIZE,\n    shuffle=False,\n    num_workers=8\n)\n\n\ntest_aug_dataset = MelanomaTestDataset(\n    image_paths=test_df['image_name'],\n    resize=True,\n    augmentations=test_aug_transforms,\n)\n\ntest_aug_loader = torch.utils.data.DataLoader(\n    test_aug_dataset,\n    batch_size=VALID_BATCH_SIZE,\n    shuffle=False,\n    num_workers=8\n)","metadata":{"id":"buQqLxPhdWu7","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_PATHS = [f'model_{fold}.bin' for fold in range(4)]","metadata":{"id":"nyAbTtVOdWu7","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = xm.xla_device()\npredictions = []\n\nfor path in MODEL_PATHS:\n    model = MelanomaModel().to(device)\n    model.load_state_dict(torch.load(path))\n\n    preds = predict(test_loader, model, device)\n    preds_aug = predict(test_aug_loader, model, device)\n\n    preds = np.array(preds)\n    preds_aug = np.array(preds_aug)\n\n    final_preds = np.mean([preds, preds_aug], axis=0)\n\n    predictions.append(final_preds)","metadata":{"id":"L-f8UkSqdWu8","outputId":"7e9eb622-eb2d-4ce9-e9b4-d2e49e414fec","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission","metadata":{"id":"4xE5ACFKdWu8"}},{"cell_type":"code","source":"def sigmoid(x):\n    return 1 / (1 + np.exp(-x))","metadata":{"id":"KrAc66w8dWu8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = np.array(predictions)\npredictions = np.mean(predictions, axis=0)","metadata":{"id":"VlWzpQDWdWu8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = sigmoid(predictions)","metadata":{"id":"WL64LhTwdWu8","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(predictions)","metadata":{"id":"HJkLPJfRdWu8","outputId":"84ab0d4e-2440-4b15-d6b6-72f38781d960","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_df = pd.read_csv(SUBMISSION_CSV)\nsub_df['target'] = predictions","metadata":{"id":"oSn2wqjZdWu9","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)","metadata":{"id":"NOivHuuSdWu9","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<hr> <br> <div align='justify'><font color=\"##353B47\" size=\"4\">Thank you for exploring this notebook! I hope it provided clear insights and addressed your questions or curiosity. <u>Your constructive feedback is highly appreciated</u>—it fuels my growth and inspires me to create even better content. As a lifelong learner and enthusiast, I aim to deepen my understanding while sharing knowledge with others. If you found this helpful, I’d be grateful for an <u>upvote or share</u> to support my work.</font></div> <br> <div align='center'><font color=\"##353B47\" size=\"3\">Wishing you endless curiosity and passion.</font></div>","metadata":{}}]}