{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:24:37.162044Z","iopub.execute_input":"2025-03-11T10:24:37.162460Z","iopub.status.idle":"2025-03-11T10:25:31.633928Z","shell.execute_reply.started":"2025-03-11T10:24:37.162429Z","shell.execute_reply":"2025-03-11T10:25:31.633035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install efficientnet_pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:25:31.634954Z","iopub.execute_input":"2025-03-11T10:25:31.635331Z","iopub.status.idle":"2025-03-11T10:25:37.818523Z","shell.execute_reply.started":"2025-03-11T10:25:31.635310Z","shell.execute_reply":"2025-03-11T10:25:37.817697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Dataset, random_split\nfrom torchvision import datasets, models, transforms\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport warnings\nimport time\nimport cv2\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:25:37.820605Z","iopub.execute_input":"2025-03-11T10:25:37.820915Z","iopub.status.idle":"2025-03-11T10:25:44.415694Z","shell.execute_reply.started":"2025-03-11T10:25:37.820892Z","shell.execute_reply":"2025-03-11T10:25:44.414964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_time = time.time()\nBASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:25:44.417107Z","iopub.execute_input":"2025-03-11T10:25:44.417589Z","iopub.status.idle":"2025-03-11T10:25:44.420964Z","shell.execute_reply.started":"2025-03-11T10:25:44.417556Z","shell.execute_reply":"2025-03-11T10:25:44.420097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"brain_activities = ['Seizure', 'GPD', 'LRDA', 'Other', 'GRDA', 'LPD']\nactivity_mapping = {activity: idx for idx, activity in enumerate(brain_activities)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:25:44.421724Z","iopub.execute_input":"2025-03-11T10:25:44.421947Z","iopub.status.idle":"2025-03-11T10:25:44.441558Z","shell.execute_reply.started":"2025-03-11T10:25:44.421928Z","shell.execute_reply":"2025-03-11T10:25:44.440705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f\"{BASE_DIR}train.csv\")\n\ndf_toy = df.sample(frac=0.3, random_state=42)\n# Split 80% Train, 20% Temp (Validation + Test)\ntrain_df, temp_df = train_test_split(df_toy, test_size=0.4, random_state=42)\n\n# Split 10% Validation, 10% Test from Temp\nval_df, test_df = train_test_split(temp_df, test_size=0.5, random_state=42)\n\n# Save to CSV\ntrain_df.to_csv(\"train.csv\", index=False)\nval_df.to_csv(\"validation.csv\", index=False)\ntest_df.to_csv(\"test.csv\", index=False)\n\nprint(\"Splitting done! Train:\", len(train_df), \"Val:\", len(val_df), \"Test:\", len(test_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:39:55.917333Z","iopub.execute_input":"2025-03-11T10:39:55.917647Z","iopub.status.idle":"2025-03-11T10:39:56.190030Z","shell.execute_reply.started":"2025-03-11T10:39:55.917622Z","shell.execute_reply":"2025-03-11T10:39:56.189142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ChunkedBrainActivityDataset(Dataset):\n    def __init__(self, csv_file, base_dir, activity_mapping):\n        self.df = csv_file\n        self.base_dir = base_dir\n        self.activity_mapping = activity_mapping\n        self.resize_transform = transforms.Resize((224, 224))\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        spect_id, label, offset = self.df.iloc[idx][[\"spectrogram_id\", \"expert_consensus\", \"spectrogram_label_offset_seconds\"]]\n\n        temp_df = pd.read_parquet(f'{self.base_dir}/train_spectrograms/{spect_id}.parquet')\n        temp_df.drop(['time'], axis=1, inplace=True)\n\n        start = int(offset) // 2\n        temp_df = temp_df[start:start+300]\n        temp_df = np.log1p(temp_df)\n        temp_df /= temp_df.max()\n        temp_arr = np.nan_to_num(temp_df.to_numpy(), nan=1e-4)\n\n        # Use OpenCV to apply a colormap and convert to RGB\n        temp_arr_uint8 = np.uint8(255 * temp_arr)\n        rgb_image = cv2.applyColorMap(temp_arr_uint8, cv2.COLORMAP_JET)\n\n        # Normalize to [0, 1] and convert to tensor\n        rgb_image = rgb_image.astype(np.float32) / 255.0\n        rgb_image_tensor = torch.tensor(rgb_image).permute(2, 0, 1)  # (C, H, W)\n        rgb_image_tensor = self.resize_transform(rgb_image_tensor)\n            \n        y = self.activity_mapping[label]\n        y_tensor = torch.nn.functional.one_hot(torch.tensor(y, dtype=torch.long), num_classes=6).float()\n        \n        return rgb_image_tensor, y_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:30:49.392823Z","iopub.execute_input":"2025-03-11T10:30:49.393146Z","iopub.status.idle":"2025-03-11T10:30:49.400714Z","shell.execute_reply.started":"2025-03-11T10:30:49.393123Z","shell.execute_reply":"2025-03-11T10:30:49.399720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Now create DataLoader with the chunked dataset\n# chunk_size = 1000  # Adjust chunk size according to memory constraints\n\ntrain_dataset = ChunkedBrainActivityDataset(csv_file=train_df, base_dir=BASE_DIR, activity_mapping=activity_mapping)\nval_dataset = ChunkedBrainActivityDataset(csv_file=val_df, base_dir=BASE_DIR, activity_mapping=activity_mapping)\ntest_dataset = ChunkedBrainActivityDataset(csv_file=test_df, base_dir=BASE_DIR, activity_mapping=activity_mapping)\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers= 2, pin_memory=True, prefetch_factor=2)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False, num_workers= 2, pin_memory=True, prefetch_factor=2)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers= 2, pin_memory=True, prefetch_factor=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:30:50.401998Z","iopub.execute_input":"2025-03-11T10:30:50.402329Z","iopub.status.idle":"2025-03-11T10:30:50.407765Z","shell.execute_reply.started":"2025-03-11T10:30:50.402305Z","shell.execute_reply":"2025-03-11T10:30:50.406875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# import torch.optim as optim\n# from torch.utils.data import DataLoader\n\n# # Assuming your ChunkedBrainActivityDataset class and the DataLoader code\n# # for train_loader, val_loader, and test_loader are already defined.\n\n# # Define a logistic regression model for multi-class classification.\n# class LogisticRegressionModel(nn.Module):\n#     def __init__(self, input_dim, num_classes):\n#         super(LogisticRegressionModel, self).__init__()\n#         self.linear = nn.Linear(input_dim, num_classes)\n        \n#     def forward(self, x):\n#         # Flatten the input tensor: (batch, 3, 224, 224) => (batch, 3*224*224)\n#         x = x.view(x.size(0), -1)\n#         # Return the raw logits (CrossEntropyLoss applies softmax internally)\n#         return self.linear(x)\n\n# # Set device to GPU if available\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# # Calculate the size of the flattened image.\n# # Your images are of shape: (3, 224, 224)\n# input_dim = 3 * 224 * 224\n# num_classes = 6\n\n# # Instantiate the model, loss function, and optimizer.\n# model = LogisticRegressionModel(input_dim, num_classes).to(device)\n\n# # CrossEntropyLoss expects integer labels, not one-hot vectors.\n# criterion = nn.CrossEntropyLoss()\n# optimizer = optim.SGD(model.parameters(), lr=0.01)\n\n# # Number of training epochs\n# num_epochs = 10\n\n# for epoch in range(num_epochs):\n#     model.train()\n#     running_loss = 0.0\n#     correct = 0\n#     total = 0\n#     for images, targets in train_loader:\n#         # Move the batch to the device\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         # Convert one-hot targets to integer labels.\n#         labels = torch.argmax(targets, dim=1)\n\n#         optimizer.zero_grad()\n#         logits = model(images)\n#         loss = criterion(logits, labels)\n#         loss.backward()\n#         optimizer.step()\n\n#         running_loss += loss.item() * images.size(0)\n#         # Calculate accuracy for the batch\n#         _, preds = torch.max(logits, 1)\n#         total += labels.size(0)\n#         correct += (preds == labels).sum().item()\n        \n#     epoch_loss = running_loss / total\n#     epoch_acc = 100 * correct / total\n#     print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%\")\n\n\n# model.eval() # Set the model to evaluation mode\n# correct_test = 0\n# total_test = 0\n\n# with torch.no_grad(): # Disables gradient calculation\n#     for images, targets in test_loader:\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         # Convert one-hot target vectors to scalar class labels\n#         labels = torch.argmax(targets, dim=1)\n#         # Forward pass to get predictions\n#         logits = model(images)\n#         _, preds = torch.max(logits, dim=1)\n#         total_test += labels.size(0)\n#         correct_test += (preds == labels).sum().item()\n#     test_accuracy = 100 * correct_test / total_test\n#     print(f\"Test Accuracy: {test_accuracy:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:30:52.224554Z","iopub.execute_input":"2025-03-11T10:30:52.224841Z","iopub.status.idle":"2025-03-11T10:30:52.228999Z","shell.execute_reply.started":"2025-03-11T10:30:52.224820Z","shell.execute_reply":"2025-03-11T10:30:52.228195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# import torch.optim as optim\n# import torchvision.models as models  # ✅ Correct import\n\n# # Define the model that uses EfficientNet as an encoder with logistic regression as the classifier\n# class EfficientNetV2EncoderLogisticRegression(nn.Module):\n#     def __init__(self, num_classes=6):\n#         super(EfficientNetV2EncoderLogisticRegression, self).__init__()\n#         # ✅ Load pretrained EfficientNetV2-S correctly\n#         self.encoder = models.efficientnet_v2_s(weights=models.EfficientNet_V2_S_Weights.DEFAULT)\n        \n#         # Get the feature size from the last layer\n#         n_features = self.encoder.classifier[1].in_features\n        \n#         # Remove the classifier head\n#         self.encoder.classifier = nn.Identity()\n        \n#         # Add a logistic regression layer for classification\n#         self.logistic_regression = nn.Linear(n_features, num_classes)\n\n#     def forward(self, x):\n#         features = self.encoder(x)  # Extract features\n#         logits = self.logistic_regression(features)  # Apply classifier\n#         return logits\n\n# # Set device\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# # Instantiate the model and move it to the appropriate device\n# num_classes = 6\n# model = EfficientNetV2EncoderLogisticRegression(num_classes=num_classes).to(device)\n\n# # Load pretrained model weights\n# pretrained_path = \"/kaggle/input/cnn_efficientnet_v1/pytorch/default/1/HMS_model_v1_efficientnet_v2_s.pth\"\n# pretrained_dict = torch.load(pretrained_path, map_location=device)\n\n# # Get model's current state dict\n# model_dict = model.state_dict()\n\n# # ✅ Load only matching layers (ignore classifier mismatch)\n# pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict and v.shape == model_dict[k].shape}\n# model_dict.update(pretrained_dict)\n# model.load_state_dict(model_dict)\n\n# # ✅ Do NOT redefine classifier again (already done inside the model class)\n# # Move model to device\n# model.to(device)\n\n# print(\"Model loaded and classifier updated successfully!\")\n\n# # Define the loss function and optimizer\n# criterion = nn.CrossEntropyLoss()\n# optimizer = optim.Adam(model.parameters(), lr=0.001)\n\n# # Training loop\n# num_epochs = 10\n# for epoch in range(num_epochs):\n#     model.train()\n#     running_loss = 0.0\n#     correct = 0\n#     total = 0\n#     for images, targets in train_loader:\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         labels = torch.argmax(targets, dim=1)  # Convert one-hot encoded to class index\n        \n#         optimizer.zero_grad()\n#         logits = model(images)\n#         loss = criterion(logits, labels)\n#         loss.backward()\n#         optimizer.step()\n        \n#         running_loss += loss.item() * images.size(0)\n#         _, preds = torch.max(logits, dim=1)\n#         total += labels.size(0)\n#         correct += (preds == labels).sum().item()\n    \n#     epoch_loss = running_loss / total\n#     epoch_acc = 100 * correct / total\n#     print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%\")\n\n# # Evaluate on the test dataset\n# model.eval()\n# correct_test = 0\n# total_test = 0\n# with torch.no_grad():\n#     for images, targets in test_loader:\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         labels = torch.argmax(targets, dim=1)\n#         logits = model(images)\n#         _, preds = torch.max(logits, dim=1)\n#         total_test += labels.size(0)\n#         correct_test += (preds == labels).sum().item()\n\n# test_accuracy = 100 * correct_test / total_test\n# print(f\"Test Accuracy: {test_accuracy:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:30:55.496191Z","iopub.execute_input":"2025-03-11T10:30:55.496484Z","iopub.status.idle":"2025-03-11T10:30:55.500868Z","shell.execute_reply.started":"2025-03-11T10:30:55.496465Z","shell.execute_reply":"2025-03-11T10:30:55.499844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Define the model that uses a custom CNN encoder\n# class RandomCNNEncoderLogisticRegression(nn.Module):\n#     def __init__(self, num_classes=6):\n#         super(RandomCNNEncoderLogisticRegression, self).__init__()\n#         # Define a simple CNN encoder\n#         self.encoder = nn.Sequential(\n#             nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1),  # (C, H, W) -> (64, H, W)\n#             nn.ReLU(),\n#             nn.MaxPool2d(kernel_size=2, stride=2),  # Downsample (64, H/2, W/2)\n#             nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),  # (128, H/2, W/2)\n#             nn.ReLU(),\n#             nn.MaxPool2d(kernel_size=2, stride=2),  # Downsample (128, H/4, W/4)\n#             nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1),  # (256, H/4, W/4)\n#             nn.ReLU(),\n#             nn.AdaptiveAvgPool2d((1, 1))  # Global average pooling (256, 1, 1)\n#         )\n        \n#         # Logistic regression layer\n#         self.logistic_regression = nn.Linear(256, num_classes)\n\n#     def forward(self, x):\n#         features = self.encoder(x)  # Extract features\n#         features = features.view(features.size(0), -1)  # Flatten (B, 256)\n#         logits = self.logistic_regression(features)  # Apply classifier\n#         return logits\n\n# # Set device\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# # Instantiate the model\n# num_classes = 6\n# model = RandomCNNEncoderLogisticRegression(num_classes=num_classes).to(device)\n\n# # Define the loss function and optimizer\n# criterion = nn.CrossEntropyLoss()\n# optimizer = optim.Adam(model.parameters(), lr=0.001)\n\n# print(\"Model with random CNN encoder initialized successfully!\")\n\n# # Training loop (same as before)\n# num_epochs = 10\n# for epoch in range(num_epochs):\n#     model.train()\n#     running_loss = 0.0\n#     correct = 0\n#     total = 0\n#     for images, targets in train_loader:\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         labels = torch.argmax(targets, dim=1)  # Convert one-hot encoded to class index\n        \n#         optimizer.zero_grad()\n#         logits = model(images)\n#         loss = criterion(logits, labels)\n#         loss.backward()\n#         optimizer.step()\n        \n#         running_loss += loss.item() * images.size(0)\n#         _, preds = torch.max(logits, dim=1)\n#         total += labels.size(0)\n#         correct += (preds == labels).sum().item()\n    \n#     epoch_loss = running_loss / total\n#     epoch_acc = 100 * correct / total\n#     print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%\")\n\n# # Evaluate on the test dataset (same as before)\n# model.eval()\n# correct_test = 0\n# total_test = 0\n# with torch.no_grad():\n#     for images, targets in test_loader:\n#         images = images.to(device)\n#         targets = targets.to(device)\n#         labels = torch.argmax(targets, dim=1)\n#         logits = model(images)\n#         _, preds = torch.max(logits, dim=1)\n#         total_test += labels.size(0)\n#         correct_test += (preds == labels).sum().item()\n\n# test_accuracy = 100 * correct_test / total_test\n# print(f\"Test Accuracy: {test_accuracy:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:30:56.975909Z","iopub.execute_input":"2025-03-11T10:30:56.976263Z","iopub.status.idle":"2025-03-11T10:30:56.980524Z","shell.execute_reply.started":"2025-03-11T10:30:56.976235Z","shell.execute_reply":"2025-03-11T10:30:56.979631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass SimCLR(nn.Module):\n    def __init__(self, base_model, projection_dim=128):\n        super(SimCLR, self).__init__()\n        self.encoder = base_model\n        self.encoder.fc = nn.Identity()  # Remove the final fully connected layer\n        \n        # Projection head\n        self.projection = nn.Sequential(\n            nn.Linear(2048, 2048),\n            nn.ReLU(),\n            nn.Linear(2048, projection_dim)\n        )\n    \n    def forward(self, x):\n        h = self.encoder(x)\n        z = self.projection(h)\n        return h, z\n\nclass NTXentLoss(nn.Module):\n    def __init__(self, temperature=0.1):\n        super(NTXentLoss, self).__init__()\n        self.temperature = temperature\n    \n    def forward(self, out1, out2):\n        # Normalize the outputs\n        out1 = torch.nn.functional.normalize(out1, dim=1)\n        out2 = torch.nn.functional.normalize(out2, dim=1)\n        \n        # Compute similarity matrix\n        sim_matrix = torch.exp(torch.mm(out1, out2.T) / self.temperature)\n        \n        # Positive pairs are on the diagonal\n        pos_pairs = torch.diag(sim_matrix)\n        \n        # Negative pairs are all non-diagonal elements\n        neg_pairs = sim_matrix.sum(dim=1) - pos_pairs\n        \n        # Compute loss\n        loss = -torch.log(pos_pairs / neg_pairs).mean()\n        \n        return loss\n\n# Set device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Instantiate the model and move it to the appropriate device\nbase_model = models.resnet50(pretrained=False)\nmodel = SimCLR(base_model).to(device)\n\n# Define the loss function and optimizer\ncriterion = NTXentLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:31:06.664274Z","iopub.execute_input":"2025-03-11T10:31:06.664616Z","iopub.status.idle":"2025-03-11T10:31:07.093221Z","shell.execute_reply.started":"2025-03-11T10:31:06.664588Z","shell.execute_reply":"2025-03-11T10:31:07.092547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom torchvision.transforms.functional import to_pil_image\n\nclass RandomAugmentation(nn.Module):\n    def __init__(self):\n        super(RandomAugmentation, self).__init__()\n        self.augmentations = transforms.Compose([\n            transforms.RandomResizedCrop(224),\n            transforms.RandomHorizontalFlip(),\n            transforms.RandomApply([transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4, hue=0.2)], p=0.8),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n    \n    def forward(self, x):\n        print(\"RandomAugmentation forward called\")\n        # Handle input batch of shape (N, C, H, W)\n        if len(x.shape) == 4:\n            print(\"Batch input detected\")\n            return torch.stack([self.augmentations(to_pil_image(img)) for img in x])\n        # Handle single image of shape (C, H, W)\n        elif len(x.shape) == 3:\n            print(\"Single image input detected\")\n            return self.augmentations(to_pil_image(x))\n        else:\n            raise ValueError(f\"Invalid input shape {x.shape}. Expected (C, H, W) or (N, C, H, W).\")\n\n# Create two instances of the augmentation pipeline\naug1 = RandomAugmentation()\naug2 = RandomAugmentation()\n\n# Training loop\nnum_epochs = 10\nfor epoch in range(num_epochs):\n    print(f\"Starting epoch {epoch+1}/{num_epochs}\")\n    model.train()\n    running_loss = 0.0\n    for batch_idx, (images, _) in enumerate(train_loader):\n        print(f\"Processing batch {batch_idx+1}\")\n        # Ensure images are on CPU before converting to PIL\n        images = images.cpu()\n        \n        print(\"Applying random augmentations (aug1)\")\n        aug_images1 = aug1(images)  # Shape: (N, C, H, W)\n        \n        print(\"Applying random augmentations (aug2)\")\n        aug_images2 = aug2(images)  # Shape: (N, C, H, W)\n        \n        aug_images1 = aug_images1.to(device)\n        aug_images2 = aug_images2.to(device)\n        \n        print(\"Forward pass through the model\")\n        _, z1 = model(aug_images1)\n        _, z2 = model(aug_images2)\n        \n        print(\"Computing contrastive loss\")\n        loss = criterion(z1, z2)\n        \n        print(\"Performing backpropagation\")\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * images.size(0)\n    \n    epoch_loss = running_loss / len(train_loader.dataset)\n    print(f\"Epoch [{epoch+1}/{num_epochs}] completed, Loss: {epoch_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:31:08.108383Z","iopub.execute_input":"2025-03-11T10:31:08.108700Z","iopub.status.idle":"2025-03-11T10:35:59.011538Z","shell.execute_reply.started":"2025-03-11T10:31:08.108673Z","shell.execute_reply":"2025-03-11T10:35:59.010437Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fine-tuning with logistic regression\nclass FineTuneModel(nn.Module):\n    def __init__(self, encoder, num_classes=6):\n        super(FineTuneModel, self).__init__()\n        self.encoder = encoder\n        self.logistic_regression = nn.Linear(2048, num_classes)\n    \n    def forward(self, x):\n        h = self.encoder(x)\n        logits = self.logistic_regression(h)\n        return logits\n\n# Freeze the encoder and add a logistic regression layer\nfine_tune_model = FineTuneModel(model.encoder).to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(fine_tune_model.logistic_regression.parameters(), lr=0.001)\n\n# Fine-tuning loop\nnum_epochs = 5\nfor epoch in range(num_epochs):\n    fine_tune_model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    for images, targets in train_loader:\n        images = images.to(device)\n        targets = targets.to(device)\n        labels = torch.argmax(targets, dim=1)  # Convert one-hot encoded to class index\n        \n        optimizer.zero_grad()\n        logits = fine_tune_model(images)\n        loss = criterion(logits, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(logits, dim=1)\n        total += labels.size(0)\n        correct += (preds == labels).sum().item()\n    \n    epoch_loss = running_loss / total\n    epoch_acc = 100 * correct / total\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:36:08.407532Z","iopub.execute_input":"2025-03-11T10:36:08.407896Z","iopub.status.idle":"2025-03-11T10:37:25.549620Z","shell.execute_reply.started":"2025-03-11T10:36:08.407862Z","shell.execute_reply":"2025-03-11T10:37:25.548525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set the model to evaluation mode\nfine_tune_model.eval()\n\n# Initialize variables to track accuracy\ntest_correct = 0\ntest_total = 0\n\n# Disable gradient computation for evaluation\nwith torch.no_grad():\n    for images, targets in test_loader:\n        images = images.to(device)\n        targets = targets.to(device)\n        labels = torch.argmax(targets, dim=1)  # Convert one-hot encoded to class index\n        \n        # Forward pass\n        logits = fine_tune_model(images)\n        \n        # Get predictions\n        _, preds = torch.max(logits, dim=1)\n        \n        # Update accuracy counters\n        test_total += labels.size(0)\n        test_correct += (preds == labels).sum().item()\n\n# Calculate test accuracy\ntest_accuracy = 100 * test_correct / test_total\n\nprint(f\"Test Accuracy: {test_accuracy:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-11T10:37:36.857417Z","iopub.execute_input":"2025-03-11T10:37:36.857790Z","iopub.status.idle":"2025-03-11T10:37:42.953034Z","shell.execute_reply.started":"2025-03-11T10:37:36.857745Z","shell.execute_reply":"2025-03-11T10:37:42.952166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}