{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os \nimport torch\nimport random","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:05.234754Z","iopub.execute_input":"2024-12-16T01:51:05.235162Z","iopub.status.idle":"2024-12-16T01:51:09.261487Z","shell.execute_reply.started":"2024-12-16T01:51:05.235112Z","shell.execute_reply":"2024-12-16T01:51:09.260781Z"}},"outputs":[],"execution_count":null},{"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(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.262927Z","iopub.execute_input":"2024-12-16T01:51:09.263271Z","iopub.status.idle":"2024-12-16T01:51:09.271571Z","shell.execute_reply.started":"2024-12-16T01:51:09.263244Z","shell.execute_reply":"2024-12-16T01:51:09.270767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.272631Z","iopub.execute_input":"2024-12-16T01:51:09.272948Z","iopub.status.idle":"2024-12-16T01:51:09.341899Z","shell.execute_reply.started":"2024-12-16T01:51:09.272909Z","shell.execute_reply":"2024-12-16T01:51:09.341216Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Read Data","metadata":{}},{"cell_type":"code","source":"sample_submission = pd.read_csv(\"/kaggle/input/siim-isic-melanoma-classification/sample_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.343629Z","iopub.execute_input":"2024-12-16T01:51:09.343919Z","iopub.status.idle":"2024-12-16T01:51:09.375445Z","shell.execute_reply.started":"2024-12-16T01:51:09.343893Z","shell.execute_reply":"2024-12-16T01:51:09.374596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.376410Z","iopub.execute_input":"2024-12-16T01:51:09.376691Z","iopub.status.idle":"2024-12-16T01:51:09.389082Z","shell.execute_reply.started":"2024-12-16T01:51:09.376665Z","shell.execute_reply":"2024-12-16T01:51:09.388253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(\"/kaggle/input/siim-isic-melanoma-classification/train.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.390095Z","iopub.execute_input":"2024-12-16T01:51:09.390383Z","iopub.status.idle":"2024-12-16T01:51:09.469702Z","shell.execute_reply.started":"2024-12-16T01:51:09.390356Z","shell.execute_reply":"2024-12-16T01:51:09.468860Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test = pd.read_csv(\"/kaggle/input/siim-isic-melanoma-classification/test.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.470667Z","iopub.execute_input":"2024-12-16T01:51:09.470904Z","iopub.status.idle":"2024-12-16T01:51:09.496809Z","shell.execute_reply.started":"2024-12-16T01:51:09.470879Z","shell.execute_reply":"2024-12-16T01:51:09.496184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.497787Z","iopub.execute_input":"2024-12-16T01:51:09.498029Z","iopub.status.idle":"2024-12-16T01:51:09.510160Z","shell.execute_reply.started":"2024-12-16T01:51:09.498005Z","shell.execute_reply":"2024-12-16T01:51:09.509367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.511271Z","iopub.execute_input":"2024-12-16T01:51:09.511620Z","iopub.status.idle":"2024-12-16T01:51:09.524387Z","shell.execute_reply.started":"2024-12-16T01:51:09.511582Z","shell.execute_reply":"2024-12-16T01:51:09.523489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_dir = \"/kaggle/input/siim-isic-melanoma-classification/jpeg/train/\"\ntrain[\"image_path\"] = train[\"image_name\"].apply(lambda x: f\"{image_dir}{x}.jpg\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.527062Z","iopub.execute_input":"2024-12-16T01:51:09.527308Z","iopub.status.idle":"2024-12-16T01:51:09.543837Z","shell.execute_reply.started":"2024-12-16T01:51:09.527284Z","shell.execute_reply":"2024-12-16T01:51:09.543192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_dir = \"/kaggle/input/siim-isic-melanoma-classification/jpeg/test/\"\ntest[\"image_path\"] = test[\"image_name\"].apply(lambda x: f\"{image_dir}{x}.jpg\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.544906Z","iopub.execute_input":"2024-12-16T01:51:09.545580Z","iopub.status.idle":"2024-12-16T01:51:09.558105Z","shell.execute_reply.started":"2024-12-16T01:51:09.545530Z","shell.execute_reply":"2024-12-16T01:51:09.557466Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nfrom PIL import Image\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.559120Z","iopub.execute_input":"2024-12-16T01:51:09.559477Z","iopub.status.idle":"2024-12-16T01:51:09.738986Z","shell.execute_reply.started":"2024-12-16T01:51:09.559420Z","shell.execute_reply":"2024-12-16T01:51:09.738320Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# Assuming `train` is your DataFrame and it contains 'image_path' and 'benign_malignant' columns\n\n# Separate malignant and benign samples\nmalignant_samples = train[train['benign_malignant'] == 'malignant'].sample(10, replace=True)\nbenign_samples = train[train['benign_malignant'] == 'benign'].sample(15, replace=True)\n\n# Combine samples to form a grid\nrandom_samples = pd.concat([malignant_samples, benign_samples]).sample(frac=1).reset_index(drop=True)\n\n# Plot the 5x5 grid\nfig, axes = plt.subplots(5, 5, figsize=(15, 15))\naxes = axes.flatten()\n\nfor i, (index, row) in enumerate(random_samples.iterrows()):\n    img = Image.open(row['image_path'])\n    axes[i].imshow(img)\n    axes[i].set_title(row['benign_malignant'], fontsize=12, loc='center')\n    axes[i].axis('off')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:09.739943Z","iopub.execute_input":"2024-12-16T01:51:09.740202Z","iopub.status.idle":"2024-12-16T01:51:35.507237Z","shell.execute_reply.started":"2024-12-16T01:51:09.740175Z","shell.execute_reply":"2024-12-16T01:51:35.506154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset & Dataloader","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:35.509031Z","iopub.execute_input":"2024-12-16T01:51:35.509532Z","iopub.status.idle":"2024-12-16T01:51:35.513940Z","shell.execute_reply.started":"2024-12-16T01:51:35.509476Z","shell.execute_reply":"2024-12-16T01:51:35.513149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\n\nclass MelanomaDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, train: bool = True, transforms=None):\n        self.df = df.reset_index(drop=True)  # Reset index to avoid indexing issues\n        self.transforms = transforms\n        self.train = train\n        \n    def __getitem__(self, index):\n        # Load the image\n        image_path = self.df.loc[index, \"image_path\"]\n        image = Image.open(image_path)\n        if image is None:\n            raise ValueError(f\"Image not found at {image_path}\")\n        \n        # Apply transformations\n        if self.transforms:\n            image = self.transforms(image)\n        \n        if self.train:\n            label = self.df.iloc[index][\"target\"]\n            return image, torch.tensor(label, dtype=torch.long)\n        else:\n            return image\n    \n    def __len__(self):\n        return len(self.df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:35.515005Z","iopub.execute_input":"2024-12-16T01:51:35.515254Z","iopub.status.idle":"2024-12-16T01:51:35.527011Z","shell.execute_reply.started":"2024-12-16T01:51:35.515230Z","shell.execute_reply":"2024-12-16T01:51:35.526081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:35.528016Z","iopub.execute_input":"2024-12-16T01:51:35.528298Z","iopub.status.idle":"2024-12-16T01:51:35.547778Z","shell.execute_reply.started":"2024-12-16T01:51:35.528264Z","shell.execute_reply":"2024-12-16T01:51:35.546733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nX_train, X_val, y_train, y_val = train_test_split(train[[\"image_name\",\"image_path\",\"target\"]],train.target,random_state=42,test_size=0.2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:35.548792Z","iopub.execute_input":"2024-12-16T01:51:35.549361Z","iopub.status.idle":"2024-12-16T01:51:36.136015Z","shell.execute_reply.started":"2024-12-16T01:51:35.549330Z","shell.execute_reply":"2024-12-16T01:51:36.135314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X_test = test[[\"image_name\",\"image_path\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:36.137071Z","iopub.execute_input":"2024-12-16T01:51:36.137628Z","iopub.status.idle":"2024-12-16T01:51:36.142871Z","shell.execute_reply.started":"2024-12-16T01:51:36.137564Z","shell.execute_reply":"2024-12-16T01:51:36.141911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_train.value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:36.143914Z","iopub.execute_input":"2024-12-16T01:51:36.144157Z","iopub.status.idle":"2024-12-16T01:51:36.158782Z","shell.execute_reply.started":"2024-12-16T01:51:36.144132Z","shell.execute_reply":"2024-12-16T01:51:36.158059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_val.value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:36.159819Z","iopub.execute_input":"2024-12-16T01:51:36.160401Z","iopub.status.idle":"2024-12-16T01:51:36.169722Z","shell.execute_reply.started":"2024-12-16T01:51:36.160362Z","shell.execute_reply":"2024-12-16T01:51:36.168931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\ntest_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])]\n)\n\ntrain_transforms=transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.2),\n    transforms.RandomVerticalFlip(p=0.2),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),  # Apply color jitter \n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])] # This is from the original paper\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:36.170833Z","iopub.execute_input":"2024-12-16T01:51:36.171444Z","iopub.status.idle":"2024-12-16T01:51:37.270269Z","shell.execute_reply.started":"2024-12-16T01:51:36.171405Z","shell.execute_reply":"2024-12-16T01:51:37.269604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = MelanomaDataset(X_train,transforms=train_transforms)\nval_dataset = MelanomaDataset(X_val,train=True,transforms=test_transforms)\ntrain_dataloader = DataLoader(train_dataset,batch_size=64,shuffle=True,num_workers=4)\nval_dataloader = DataLoader(val_dataset,batch_size=32,shuffle=False,num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:37.271250Z","iopub.execute_input":"2024-12-16T01:51:37.271654Z","iopub.status.idle":"2024-12-16T01:51:37.277527Z","shell.execute_reply.started":"2024-12-16T01:51:37.271626Z","shell.execute_reply":"2024-12-16T01:51:37.276590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = MelanomaDataset(X_test,transforms=test_transforms,train=False)\ntest_dataloader = DataLoader(test_dataset,batch_size=32,shuffle=False,num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:37.278798Z","iopub.execute_input":"2024-12-16T01:51:37.279082Z","iopub.status.idle":"2024-12-16T01:51:37.290666Z","shell.execute_reply.started":"2024-12-16T01:51:37.279054Z","shell.execute_reply":"2024-12-16T01:51:37.289849Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Models","metadata":{}},{"cell_type":"code","source":"!pip install torch-summary\nfrom torchsummary import summary\nimport torchvision.models as models","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:37.291699Z","iopub.execute_input":"2024-12-16T01:51:37.291924Z","iopub.status.idle":"2024-12-16T01:51:47.007480Z","shell.execute_reply.started":"2024-12-16T01:51:37.291899Z","shell.execute_reply":"2024-12-16T01:51:47.006403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"efficientnet = models.efficientnet_b4(weights=models.EfficientNet_B4_Weights.IMAGENET1K_V1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.009070Z","iopub.execute_input":"2024-12-16T01:51:47.009496Z","iopub.status.idle":"2024-12-16T01:51:47.905710Z","shell.execute_reply.started":"2024-12-16T01:51:47.009435Z","shell.execute_reply":"2024-12-16T01:51:47.904918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = len(y_train.unique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.906849Z","iopub.execute_input":"2024-12-16T01:51:47.907249Z","iopub.status.idle":"2024-12-16T01:51:47.913704Z","shell.execute_reply.started":"2024-12-16T01:51:47.907198Z","shell.execute_reply":"2024-12-16T01:51:47.912669Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# summary(efficientnet,(3,224,224))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.914823Z","iopub.execute_input":"2024-12-16T01:51:47.915126Z","iopub.status.idle":"2024-12-16T01:51:47.924610Z","shell.execute_reply.started":"2024-12-16T01:51:47.915081Z","shell.execute_reply":"2024-12-16T01:51:47.923515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for name, layer in efficientnet.named_children():\n#     print(f\"{name}: {layer}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.930076Z","iopub.execute_input":"2024-12-16T01:51:47.930368Z","iopub.status.idle":"2024-12-16T01:51:47.938705Z","shell.execute_reply.started":"2024-12-16T01:51:47.930337Z","shell.execute_reply":"2024-12-16T01:51:47.937781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_features = efficientnet.classifier[1].in_features","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.939819Z","iopub.execute_input":"2024-12-16T01:51:47.940152Z","iopub.status.idle":"2024-12-16T01:51:47.949682Z","shell.execute_reply.started":"2024-12-16T01:51:47.940124Z","shell.execute_reply":"2024-12-16T01:51:47.948817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"efficientnet.classifier = nn.Sequential(\n    nn.Dropout(0.3),\n    nn.Linear(num_features, 256),  # Add an intermediate fully connected layer\n    nn.ReLU(),     \n    nn.Dropout(0.2), # Activation function\n    nn.Linear(256, 2)              # Final output layer (e.g., for binary classification)\n)\nefficientnet = efficientnet.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:47.950710Z","iopub.execute_input":"2024-12-16T01:51:47.950950Z","iopub.status.idle":"2024-12-16T01:51:48.153311Z","shell.execute_reply.started":"2024-12-16T01:51:47.950926Z","shell.execute_reply":"2024-12-16T01:51:48.152640Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.154275Z","iopub.execute_input":"2024-12-16T01:51:48.154551Z","iopub.status.idle":"2024-12-16T01:51:48.158533Z","shell.execute_reply.started":"2024-12-16T01:51:48.154525Z","shell.execute_reply":"2024-12-16T01:51:48.157647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Freeze all and unfreeze fc and layer 4\nfor param in efficientnet.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.159617Z","iopub.execute_input":"2024-12-16T01:51:48.159879Z","iopub.status.idle":"2024-12-16T01:51:48.169141Z","shell.execute_reply.started":"2024-12-16T01:51:48.159853Z","shell.execute_reply":"2024-12-16T01:51:48.168479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for name, param in efficientnet.named_parameters():\n#     print(f\"{name}: requires_grad={param.requires_grad}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.170010Z","iopub.execute_input":"2024-12-16T01:51:48.170257Z","iopub.status.idle":"2024-12-16T01:51:48.178189Z","shell.execute_reply.started":"2024-12-16T01:51:48.170228Z","shell.execute_reply":"2024-12-16T01:51:48.177338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\noptimizer = torch.optim.Adam([\n    {'params': efficientnet.features.parameters(), 'lr': 1e-5},  # Lower learning rate for pretrained layers\n    {'params': efficientnet.classifier.parameters(), 'lr': 1e-4}  # Higher learning rate for new layers\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.179109Z","iopub.execute_input":"2024-12-16T01:51:48.179387Z","iopub.status.idle":"2024-12-16T01:51:48.189361Z","shell.execute_reply.started":"2024-12-16T01:51:48.179355Z","shell.execute_reply":"2024-12-16T01:51:48.188495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.nn import CrossEntropyLoss\n\n# # Compute class weights (example: using pandas)\n# class_counts = X_train['target'].value_counts()\n# class_weights = 1.0 / class_counts\n# class_weights = class_weights / class_weights.sum()  # Normalize weights\n# class_weights = torch.tensor(class_weights.values, dtype=torch.float32).to(device)\n\n# # Define weighted cross-entropy loss\n# loss_fn = CrossEntropyLoss(weight=class_weights)\nloss_fn = CrossEntropyLoss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.190317Z","iopub.execute_input":"2024-12-16T01:51:48.190603Z","iopub.status.idle":"2024-12-16T01:51:48.198400Z","shell.execute_reply.started":"2024-12-16T01:51:48.190578Z","shell.execute_reply":"2024-12-16T01:51:48.197488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch.nn as nn\n# import torch\n\n# class FocalLoss(nn.Module):\n#     def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):\n#         super(FocalLoss, self).__init__()\n#         self.alpha = alpha\n#         self.gamma = gamma\n#         self.reduction = reduction\n\n#     def forward(self, inputs, targets):\n#         ce_loss = nn.CrossEntropyLoss(reduction='none')(inputs, targets)\n#         pt = torch.exp(-ce_loss)  # Probability of the correct class\n#         focal_loss = self.alpha * (1 - pt) ** self.gamma * ce_loss\n\n#         if self.reduction == 'mean':\n#             return focal_loss.mean()\n#         elif self.reduction == 'sum':\n#             return focal_loss.sum()\n#         else:\n#             return focal_loss\n\n# loss_fn = FocalLoss(alpha=0.25, gamma=2.0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.199357Z","iopub.execute_input":"2024-12-16T01:51:48.199652Z","iopub.status.idle":"2024-12-16T01:51:48.206299Z","shell.execute_reply.started":"2024-12-16T01:51:48.199627Z","shell.execute_reply":"2024-12-16T01:51:48.205677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchmetrics.classification.auroc import AUROC\n\n# AUC-based loss function\nauc_metric = AUROC(task='binary', num_classes=None)\n\n# Use this during validation to compute AUC\n# Training will use a weighted loss like CrossEntropyLoss or FocalLoss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:48.207263Z","iopub.execute_input":"2024-12-16T01:51:48.207554Z","iopub.status.idle":"2024-12-16T01:51:50.199328Z","shell.execute_reply.started":"2024-12-16T01:51:48.207529Z","shell.execute_reply":"2024-12-16T01:51:50.198424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\n\n# scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)\n\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\nscheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=1, verbose=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:50.200405Z","iopub.execute_input":"2024-12-16T01:51:50.200948Z","iopub.status.idle":"2024-12-16T01:51:50.205181Z","shell.execute_reply.started":"2024-12-16T01:51:50.200906Z","shell.execute_reply":"2024-12-16T01:51:50.204372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time  # Import time module\n\ndef train_step(dataloader, model, loss_fn, optimizer, scheduler, epoch, print_every=5000):\n    losses = []\n    model.train()\n\n    # Record the start time for the epoch\n    epoch_start_time = time.time()\n\n    for batch, (inputs, targets) in enumerate(dataloader):\n        inputs, targets = inputs.to(device), targets.to(device)\n        outputs = model(inputs)\n\n        optimizer.zero_grad()\n        loss = loss_fn(outputs, targets)  # Compute the loss\n        losses.append(loss.item())\n\n        loss.backward()  # Backpropagation\n        optimizer.step()  # Update weights\n\n        # Update scheduler if needed (e.g., per batch for warm restarts)\n        if isinstance(scheduler, torch.optim.lr_scheduler.CosineAnnealingWarmRestarts):\n            scheduler.step(epoch + batch / len(dataloader))  # Adjust learning rate\n\n        # Print results\n        if batch % print_every == 0:\n            current = (batch + 1) * len(inputs)\n            print(f'  Training: Loss = {loss.item():>7f} [{current:>5d}/{len(dataloader.dataset):>5d}]')\n\n    # Record the end time for the epoch\n    epoch_end_time = time.time()\n\n    # Calculate the duration of the epoch\n    epoch_duration = epoch_end_time - epoch_start_time\n\n    print(f'Epoch {epoch + 1} completed in {epoch_duration:.2f} seconds')\n\n    return losses","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:50.206097Z","iopub.execute_input":"2024-12-16T01:51:50.206315Z","iopub.status.idle":"2024-12-16T01:51:50.221207Z","shell.execute_reply.started":"2024-12-16T01:51:50.206292Z","shell.execute_reply":"2024-12-16T01:51:50.220403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchmetrics.classification import BinaryAUROC\n\ndef valid_step(dataloader, model, loss_fn, auc_metric):\n    model.eval()\n    total_loss = 0\n    all_preds, all_targets = [], []\n\n    with torch.no_grad():\n        for inputs, targets in dataloader:\n            inputs, targets = inputs.to(device), targets.to(device)\n            outputs = model(inputs)\n\n            # Compute loss\n            total_loss += loss_fn(outputs, targets).item()\n\n            # Collect predictions and targets for AUC calculation\n            probs = torch.softmax(outputs, dim=1)[:, 1]  # Probability for the positive class\n            all_preds.append(probs)\n            all_targets.append(targets)\n\n    # Concatenate all predictions and targets for AUC calculation\n    all_preds = torch.cat(all_preds)\n    all_targets = torch.cat(all_targets)\n    auc = auc_metric(all_preds, all_targets)\n\n    # Compute average loss\n    avg_loss = total_loss / len(dataloader)\n    print(f'  Validation: AUC={(auc * 100):>0.2f}%, Average_Loss={avg_loss:>8f}\\n')\n\n    return avg_loss, auc.item()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:50.222111Z","iopub.execute_input":"2024-12-16T01:51:50.222407Z","iopub.status.idle":"2024-12-16T01:51:50.232887Z","shell.execute_reply.started":"2024-12-16T01:51:50.222380Z","shell.execute_reply":"2024-12-16T01:51:50.232095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\n# Initialize AUC metric\nauc_metric = BinaryAUROC().to(device)\n\nepochs = 10\ntrain_losses = []\nvalid_losses = []\nvalid_aucs = []\n\nbest_auc = 0  # To track the best validation AUC\nearly_stopping_patience = 3  # Stop if no improvement after 3 epochs\nearly_stopping_counter = 0  # Count epochs without improvement\nbest_model_path = \"best_model.pth\"  # Path to save the best model weights\n\nfor epoch in range(epochs):\n    print(f'Epoch [{epoch+1:>2d}/{epochs}]\\n-------------------------------')\n\n    # Train step\n    epoch_train_losses = train_step(train_dataloader, efficientnet, loss_fn, optimizer, scheduler, epoch, print_every=25)\n    train_losses.extend(epoch_train_losses)\n\n    # Validation step\n    avg_valid_loss, valid_auc = valid_step(val_dataloader, efficientnet, loss_fn, auc_metric)\n    valid_losses.append(avg_valid_loss)\n    valid_aucs.append(valid_auc)\n\n    # Check if this is the best model so far\n    if valid_auc > best_auc:\n        print(f\"  Validation AUC improved from {best_auc * 100:.2f}% to {valid_auc * 100:.2f}%\")\n        best_auc = valid_auc\n        early_stopping_counter = 0  # Reset early stopping counter\n        # Save best model weights\n        torch.save(efficientnet.state_dict(), best_model_path)\n        print(f\"  Best model saved to {best_model_path}\")\n    else:\n        early_stopping_counter += 1\n        print(f\"  Validation AUC did not improve. Early stopping counter: {early_stopping_counter}/{early_stopping_patience}\")\n\n    # Scheduler step (if not per batch)\n    if isinstance(scheduler, torch.optim.lr_scheduler.ReduceLROnPlateau):\n        scheduler.step(valid_auc)  # Adjust based on validation AUC\n\n    print(f'  End of Epoch {epoch+1}: AUC={valid_auc * 100:.2f}%, LR={optimizer.param_groups[0][\"lr\"]:.6f}\\n')\n\n    # Check for early stopping\n    if early_stopping_counter >= early_stopping_patience:\n        print(\"Early stopping triggered. Stopping training.\")\n        break\n\nprint(\"Training completed.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:51:50.234087Z","iopub.execute_input":"2024-12-16T01:51:50.234372Z","iopub.status.idle":"2024-12-16T01:53:22.509863Z","shell.execute_reply.started":"2024-12-16T01:51:50.234345Z","shell.execute_reply":"2024-12-16T01:53:22.508683Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the model weights (state_dict)\ntorch.save(efficientnet.state_dict(), \"efficientnet_weights.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:22.511580Z","iopub.execute_input":"2024-12-16T01:53:22.512310Z","iopub.status.idle":"2024-12-16T01:53:22.671686Z","shell.execute_reply.started":"2024-12-16T01:53:22.512264Z","shell.execute_reply":"2024-12-16T01:53:22.670716Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 4))\n\nplt.plot(train_losses)\nplt.title('Cross Entropy Loss')\nplt.xlabel('Iteration')\nplt.ylabel('Loss')\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:22.672773Z","iopub.execute_input":"2024-12-16T01:53:22.673049Z","iopub.status.idle":"2024-12-16T01:53:22.895922Z","shell.execute_reply.started":"2024-12-16T01:53:22.673021Z","shell.execute_reply":"2024-12-16T01:53:22.895093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 4))\n\nplt.plot(valid_losses)\nplt.title('Cross Entropy Loss')\nplt.xlabel('Iteration')\nplt.ylabel('Loss')\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:22.896963Z","iopub.execute_input":"2024-12-16T01:53:22.897222Z","iopub.status.idle":"2024-12-16T01:53:23.130803Z","shell.execute_reply.started":"2024-12-16T01:53:22.897195Z","shell.execute_reply":"2024-12-16T01:53:23.129794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 4))\n\nplt.plot(valid_aucs)\nplt.title('Cross Entropy Loss')\nplt.xlabel('Iteration')\nplt.ylabel('Loss')\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:23.132091Z","iopub.execute_input":"2024-12-16T01:53:23.132515Z","iopub.status.idle":"2024-12-16T01:53:23.347305Z","shell.execute_reply.started":"2024-12-16T01:53:23.132438Z","shell.execute_reply":"2024-12-16T01:53:23.346418Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test ","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm \ndef test_step(dataloader, model):\n    model.eval()\n    all_preds= []\n\n    with torch.no_grad():\n        for inputs in tqdm(dataloader):\n            inputs = inputs.to(device)\n            outputs = model(inputs)\n\n            # Collect predictions and targets for AUC calculation\n            probs = torch.softmax(outputs, dim=1)[:, 1]  # Probability for the positive class\n            all_preds.append(probs)\n\n    # Concatenate all predictions and targets for AUC calculation\n    all_preds = torch.cat(all_preds)\n\n    return all_preds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:23.348424Z","iopub.execute_input":"2024-12-16T01:53:23.348741Z","iopub.status.idle":"2024-12-16T01:53:23.353948Z","shell.execute_reply.started":"2024-12-16T01:53:23.348712Z","shell.execute_reply":"2024-12-16T01:53:23.353063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = test_step(test_dataloader,efficientnet)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:23.354985Z","iopub.execute_input":"2024-12-16T01:53:23.355296Z","iopub.status.idle":"2024-12-16T01:53:30.062083Z","shell.execute_reply.started":"2024-12-16T01:53:23.355257Z","shell.execute_reply":"2024-12-16T01:53:30.060948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = predictions.detach().cpu().numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:30.063334Z","iopub.execute_input":"2024-12-16T01:53:30.063641Z","iopub.status.idle":"2024-12-16T01:53:30.069029Z","shell.execute_reply.started":"2024-12-16T01:53:30.063611Z","shell.execute_reply":"2024-12-16T01:53:30.068093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission[\"target\"] = predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:30.070111Z","iopub.execute_input":"2024-12-16T01:53:30.070399Z","iopub.status.idle":"2024-12-16T01:53:30.831019Z","shell.execute_reply.started":"2024-12-16T01:53:30.070372Z","shell.execute_reply":"2024-12-16T01:53:30.829785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-16T01:53:30.831943Z","iopub.status.idle":"2024-12-16T01:53:30.832291Z","shell.execute_reply.started":"2024-12-16T01:53:30.832132Z","shell.execute_reply":"2024-12-16T01:53:30.832150Z"}},"outputs":[],"execution_count":null}]}