{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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"},{"sourceId":7707802,"sourceType":"datasetVersion","datasetId":4500288}],"dockerImageVersionId":30665,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torchvision\nimport torchvision.transforms as transforms\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, TensorDataset\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils import resample\nimport numpy as np\nimport cv2\nimport pandas as pd\nimport time\nimport wandb\nfrom PIL import Image\nfrom sklearn.metrics import roc_curve, precision_recall_curve, precision_score, recall_score, f1_score, confusion_matrix\nimport torch.nn.functional as F","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-03T20:36:02.621310Z","iopub.execute_input":"2024-03-03T20:36:02.621588Z","iopub.status.idle":"2024-03-03T20:36:11.464532Z","shell.execute_reply.started":"2024-03-03T20:36:02.621564Z","shell.execute_reply":"2024-03-03T20:36:11.463578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define W&B parameters","metadata":{}},{"cell_type":"code","source":"PARAMS = {\n    'epochs': 10,\n    'learning_rate': 0.001,\n    'batch_size': 32,\n    'optimizer': 'sgd',  # 'adam', 'rmsprop', 'sgd', etc.\n    'loss_function': 'cross_entropy', # 'cross_entropy','BCE'\n    'momentum': 0.0,  # Add momentum for SGD optimizer\n    'model_architecture': 'resnet34',  # Change to any other model\n    'train_image_path' : '/kaggle/input/melanoma-resized-images-512512/train/train/',\n    'train_csv_path' : '/kaggle/input/siim-isic-melanoma-classification/train.csv',\n    'test_image_path' : '/kaggle/input/melanoma-resized-images-512512/test/test/',\n    'test_csv_path' : '/kaggle/input/siim-isic-melanoma-classification/test.csv'\n    \n}\nwandb.init(project=\"SIIM ISIC RESNET18 Model\", save_code=True, config=PARAMS)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:36:11.466326Z","iopub.execute_input":"2024-03-03T20:36:11.466627Z","iopub.status.idle":"2024-03-03T20:37:11.591806Z","shell.execute_reply.started":"2024-03-03T20:36:11.466603Z","shell.execute_reply":"2024-03-03T20:37:11.590943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_path = PARAMS['train_image_path']\ntest_image_path = PARAMS['test_image_path']","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:11.593427Z","iopub.execute_input":"2024-03-03T20:37:11.593735Z","iopub.status.idle":"2024-03-03T20:37:12.481722Z","shell.execute_reply.started":"2024-03-03T20:37:11.593710Z","shell.execute_reply":"2024-03-03T20:37:12.480871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = transforms.Compose([\n    transforms.Resize(256, interpolation=transforms.InterpolationMode.BILINEAR),\n    transforms.CenterCrop(224),\n    transforms.RandomHorizontalFlip(),  # Random horizontal flip\n    transforms.RandomVerticalFlip(),    # Random vertical flip\n    transforms.RandomRotation(45),      # Random rotation (-45 to +45 degrees)\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),  # Color jitter\n    transforms.ToTensor(),              # Convert to tensor\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:12.484139Z","iopub.execute_input":"2024-03-03T20:37:12.484558Z","iopub.status.idle":"2024-03-03T20:37:13.243653Z","shell.execute_reply.started":"2024-03-03T20:37:12.484519Z","shell.execute_reply":"2024-03-03T20:37:13.242777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load and Preprocess","metadata":{}},{"cell_type":"code","source":"def load_and_preprocess_images(image_paths, mode='train', transform=None):\n    images = []\n    for path in image_paths:\n        # Load image using OpenCV\n        if mode == \"test\":\n            img = cv2.imread(test_image_path + path + '.jpg')  \n        else:\n            img = cv2.imread(train_image_path + path + '.jpg')\n        pil_img = Image.fromarray(img)\n        # Apply data transformations if provided\n        if transform is not None:\n            augmented_img = transform(pil_img)\n            \n        images.append(np.array(augmented_img))\n    return np.array(images)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:13.244812Z","iopub.execute_input":"2024-03-03T20:37:13.245137Z","iopub.status.idle":"2024-03-03T20:37:13.984400Z","shell.execute_reply.started":"2024-03-03T20:37:13.245106Z","shell.execute_reply":"2024-03-03T20:37:13.983554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the function to load and preprocess images\ndef load_images(image_paths):\n    images = []\n    for path in image_paths:\n        # Load image using OpenCV\n        img = cv2.imread(train_image_path + path + '.jpg')\n        # Convert image to grayscale\n        #img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        images.append(img)\n    return np.array(images)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:13.985616Z","iopub.execute_input":"2024-03-03T20:37:13.985956Z","iopub.status.idle":"2024-03-03T20:37:14.754996Z","shell.execute_reply.started":"2024-03-03T20:37:13.985924Z","shell.execute_reply":"2024-03-03T20:37:14.754236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(PARAMS['train_csv_path'])\ntest_df = pd.read_csv(PARAMS['test_csv_path'])\ntrain_df.head()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:14.756485Z","iopub.execute_input":"2024-03-03T20:37:14.756806Z","iopub.status.idle":"2024-03-03T20:37:15.715870Z","shell.execute_reply.started":"2024-03-03T20:37:14.756776Z","shell.execute_reply":"2024-03-03T20:37:15.714990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_distribution = train_df['target'].value_counts()\nclass_distribution","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:15.717494Z","iopub.execute_input":"2024-03-03T20:37:15.717881Z","iopub.status.idle":"2024-03-03T20:37:16.479138Z","shell.execute_reply.started":"2024-03-03T20:37:15.717848Z","shell.execute_reply":"2024-03-03T20:37:16.478236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X = train_df['image_name'] #images\ny = train_df['target'] #target","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:16.480372Z","iopub.execute_input":"2024-03-03T20:37:16.480654Z","iopub.status.idle":"2024-03-03T20:37:17.246532Z","shell.execute_reply.started":"2024-03-03T20:37:16.480630Z","shell.execute_reply":"2024-03-03T20:37:17.245669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming X and y are your original data and labels\nX_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.2, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:17.251140Z","iopub.execute_input":"2024-03-03T20:37:17.251477Z","iopub.status.idle":"2024-03-03T20:37:18.071893Z","shell.execute_reply.started":"2024-03-03T20:37:17.251449Z","shell.execute_reply":"2024-03-03T20:37:18.070933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Separate majority and minority classes in training data","metadata":{}},{"cell_type":"code","source":"\nmajority_class = X_train[y_train == 0]\nminority_class = X_train[y_train == 1]","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:18.073754Z","iopub.execute_input":"2024-03-03T20:37:18.074026Z","iopub.status.idle":"2024-03-03T20:37:18.824888Z","shell.execute_reply.started":"2024-03-03T20:37:18.073997Z","shell.execute_reply":"2024-03-03T20:37:18.823862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Resampling","metadata":{}},{"cell_type":"code","source":"undersampled_majority_class = resample(majority_class,\n                                        replace=False,\n                                        n_samples=len(minority_class),\n                                        random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:18.825992Z","iopub.execute_input":"2024-03-03T20:37:18.826292Z","iopub.status.idle":"2024-03-03T20:37:19.548765Z","shell.execute_reply.started":"2024-03-03T20:37:18.826269Z","shell.execute_reply":"2024-03-03T20:37:19.547861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load and preprocess images for majority class\nundersampled_majority_images = load_and_preprocess_images(undersampled_majority_class,\"train\",transform=data_transforms)\n# Convert labels to float\nundersampled_majority_labels = np.zeros(len(undersampled_majority_images))\n\n# Load and preprocess images for minority class\nminority_images = load_and_preprocess_images(minority_class,\"train\",transform=data_transforms)\n# Convert labels to float\nminority_labels = np.ones(len(minority_images))","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:19.549970Z","iopub.execute_input":"2024-03-03T20:37:19.550277Z","iopub.status.idle":"2024-03-03T20:37:47.586568Z","shell.execute_reply.started":"2024-03-03T20:37:19.550253Z","shell.execute_reply":"2024-03-03T20:37:47.585675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Combine minority class with undersampled majority class\nundersampled_X_train = np.concatenate([undersampled_majority_images, minority_images])\nundersampled_y_train = np.concatenate([undersampled_majority_labels, minority_labels])\n\n\nundersampled_X_train_tensor = torch.tensor(undersampled_X_train, dtype=torch.float32)\nundersampled_y_train_tensor = torch.tensor(undersampled_y_train, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:47.588046Z","iopub.execute_input":"2024-03-03T20:37:47.588650Z","iopub.status.idle":"2024-03-03T20:37:48.755577Z","shell.execute_reply.started":"2024-03-03T20:37:47.588615Z","shell.execute_reply":"2024-03-03T20:37:48.754668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Shuffle the data\nshuffled_indices = np.random.permutation(len(undersampled_y_train))\nundersampled_X_train = undersampled_X_train_tensor[shuffled_indices]\nundersampled_y_train = undersampled_y_train_tensor[shuffled_indices]\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:48.756930Z","iopub.execute_input":"2024-03-03T20:37:48.757643Z","iopub.status.idle":"2024-03-03T20:37:50.010906Z","shell.execute_reply.started":"2024-03-03T20:37:48.757610Z","shell.execute_reply":"2024-03-03T20:37:50.009981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"undersampled_dataset = TensorDataset(undersampled_X_train_tensor, undersampled_y_train_tensor.long())","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:50.012084Z","iopub.execute_input":"2024-03-03T20:37:50.012429Z","iopub.status.idle":"2024-03-03T20:37:50.753243Z","shell.execute_reply.started":"2024-03-03T20:37:50.012405Z","shell.execute_reply":"2024-03-03T20:37:50.752243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define data loader\nbatch_size = PARAMS['batch_size']\n#undersampled_dataset.transform = data_transforms\nundersampled_dataloader = DataLoader(undersampled_dataset, batch_size=batch_size, shuffle=True)\nlen(undersampled_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:50.755129Z","iopub.execute_input":"2024-03-03T20:37:50.755816Z","iopub.status.idle":"2024-03-03T20:37:51.551997Z","shell.execute_reply.started":"2024-03-03T20:37:50.755776Z","shell.execute_reply":"2024-03-03T20:37:51.551090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:51.553427Z","iopub.execute_input":"2024-03-03T20:37:51.553738Z","iopub.status.idle":"2024-03-03T20:37:52.278764Z","shell.execute_reply.started":"2024-03-03T20:37:51.553709Z","shell.execute_reply":"2024-03-03T20:37:52.277912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nval_images = load_and_preprocess_images(X_val,\"train\",transform=data_transforms)\nval_labels = np.array(y_val)\n\nval_labels_tensor = torch.tensor(val_labels, dtype=torch.long)\n\n# Create validation dataset\nval_dataset = TensorDataset(torch.tensor(val_images, dtype=torch.float32), val_labels_tensor)\n\n# Define validation data loader\nval_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:37:52.280058Z","iopub.execute_input":"2024-03-03T20:37:52.280394Z","iopub.status.idle":"2024-03-03T20:40:47.584613Z","shell.execute_reply.started":"2024-03-03T20:37:52.280366Z","shell.execute_reply":"2024-03-03T20:40:47.583697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load pre-trained ResNet model","metadata":{}},{"cell_type":"code","source":"\nmodel = getattr(torchvision.models, PARAMS['model_architecture'])(pretrained=True).to(device)\nnum_ftrs = model.fc.in_features\nmodel.fc = nn.Linear(num_ftrs, 2).to(device) ","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:47.585722Z","iopub.execute_input":"2024-03-03T20:40:47.585982Z","iopub.status.idle":"2024-03-03T20:40:53.454156Z","shell.execute_reply.started":"2024-03-03T20:40:47.585959Z","shell.execute_reply":"2024-03-03T20:40:53.453236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer and Criterion Call Definition","metadata":{}},{"cell_type":"code","source":"def get_optimizer_and_criterion(PARAMS, model):\n    # Update the loss function based on the PARAMS\n    if PARAMS['loss_function'] == 'cross_entropy':\n        criterion = nn.CrossEntropyLoss()\n    elif PARAMS['loss_function'] == 'BCE':\n        criterion = nn.BCEWithLogitsLoss()\n    # Add more conditions for other loss functions as needed\n    else:\n        raise ValueError(f\"Unsupported loss function: {PARAMS['loss_function']}\")\n\n    # Update the optimizer based on the PARAMS\n    if PARAMS['optimizer'] == 'adam':\n        optimizer = optim.Adam(model.parameters(), lr=PARAMS['learning_rate'])\n    elif PARAMS['optimizer'] == 'rmsprop':\n        optimizer = optim.RMSprop(model.parameters(), lr=PARAMS['learning_rate'])\n    elif PARAMS['optimizer'] == 'sgd':\n        optimizer = optim.SGD(model.parameters(), lr=PARAMS['learning_rate'], momentum=PARAMS['momentum'])\n    else:\n        raise ValueError(f\"Unsupported optimizer: {PARAMS['optimizer']}\")\n\n    return optimizer, criterion","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:53.455436Z","iopub.execute_input":"2024-03-03T20:40:53.455759Z","iopub.status.idle":"2024-03-03T20:40:54.256303Z","shell.execute_reply.started":"2024-03-03T20:40:53.455730Z","shell.execute_reply":"2024-03-03T20:40:54.255316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Definition","metadata":{}},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, undersampled_dataloader, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    start_time = time.time()\n    \n    for inputs, labels in undersampled_dataloader:\n        optimizer.zero_grad()                \n        # Print the shape of the input tensor\n        #inputs = inputs.permute(0, 3, 1, 2)  # Rearrange dimensions\n        inputs, labels = inputs.to(device), labels.to(device)\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()    \n        \n        # Calculate accuracy\n        _, predicted = torch.max(outputs, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n    \n    end_time = time.time()\n    train_time = end_time - start_time\n    \n    training_loss = running_loss / len(undersampled_dataloader)\n    training_accuracy = 100 * correct / total\n    \n    return training_loss,training_accuracy,train_time","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:54.257781Z","iopub.execute_input":"2024-03-03T20:40:54.258218Z","iopub.status.idle":"2024-03-03T20:40:55.078725Z","shell.execute_reply.started":"2024-03-03T20:40:54.258164Z","shell.execute_reply":"2024-03-03T20:40:55.077683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:55.079953Z","iopub.execute_input":"2024-03-03T20:40:55.083497Z","iopub.status.idle":"2024-03-03T20:40:55.975944Z","shell.execute_reply.started":"2024-03-03T20:40:55.083469Z","shell.execute_reply":"2024-03-03T20:40:55.974918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Definition","metadata":{}},{"cell_type":"code","source":"def validate_model(model, criterion, val_loader, device):\n    model.to(device)\n\n        # Validation loop\n    model.eval()\n    val_running_loss = 0.0\n    val_correct = 0\n    val_total = 0\n    val_preds = []\n    val_labels = []\n    start_time = time.time()\n\n    with torch.no_grad():\n        for val_inputs, val_labels_batch in val_dataloader:\n        # Ensure correct data type and device for input data\n            val_inputs = val_inputs.to(device, dtype=torch.float32)\n            val_labels_batch = val_labels_batch.to(device, dtype=torch.long)\n\n            # Rearrange dimensions if necessary\n            #val_inputs = val_inputs.permute(0, 3, 1, 2)\n\n            # Forward pass\n            val_outputs = model(val_inputs)\n\n            # Calculate loss\n            val_loss = criterion(val_outputs, val_labels_batch)\n            val_running_loss += val_loss.item()\n\n                # Append predictions and true labels\n            val_preds.extend(torch.argmax(val_outputs, axis=1).cpu().numpy())\n            val_labels.extend(val_labels_batch.cpu().numpy())\n\n            # Calculate accuracy\n            val_total += val_labels_batch.size(0)\n            val_correct += (torch.argmax(val_outputs, axis=1) == val_labels_batch).sum().item()\n\n    # Calculate validation loss\n    val_loss = val_running_loss / len(val_dataloader)\n\n    # Calculate validation accuracy\n    val_accuracy = 100 * val_correct / val_total\n\n    end_time = time.time()\n    val_time = end_time - start_time\n\n    # Calculate precision, recall, and F1 score\n    precision = precision_score(val_labels, val_preds,average='weighted')\n    recall = recall_score(val_labels, val_preds,average='weighted')\n    f1 = f1_score(val_labels, val_preds,average='weighted')\n    \n    conf_matrix = confusion_matrix(val_labels, val_preds, normalize='true')  # Normalize confusion matrix\n    class_names = ['benign', 'malignant'] \n    # Plot confusion matrix with class names\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(conf_matrix, annot=True, cmap='Blues', fmt=\".2%\", xticklabels=class_names, yticklabels=class_names)\n    plt.title('Confusion Matrix')\n    plt.xlabel('Predicted Labels')\n    plt.ylabel('True Labels')\n    plt.show()\n\n\n    return val_loss, val_accuracy,val_time,precision,recall,f1,conf_matrix","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:55.977512Z","iopub.execute_input":"2024-03-03T20:40:55.977896Z","iopub.status.idle":"2024-03-03T20:40:56.849906Z","shell.execute_reply.started":"2024-03-03T20:40:55.977860Z","shell.execute_reply":"2024-03-03T20:40:56.848869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test Definition","metadata":{}},{"cell_type":"code","source":"def test_model(test_loader, model, device, test_df):\n    model.eval()\n    predictions = []\n    image_names = []\n\n    test_start_time = time.time()\n    with torch.no_grad():\n        for i, data in enumerate(test_loader):\n            data = data.to(device)\n            outputs = model(data)\n\n            probabilities = (torch.sigmoid(outputs) >= 0.5).cpu().numpy()\n            predictions.extend(probabilities.flatten())\n\n            image_names.extend(test_df['image_name'][i * test_loader.batch_size:(i + 1) * test_loader.batch_size])\n\n    test_end_time = time.time()\n    test_time = test_end_time - test_start_time\n\n    submission_df = pd.DataFrame({'image_name': image_names, 'target': predictions})\n    submission_df.to_csv('submission.csv', index=False)\n\n    print(f'Test Evaluation and Submission CSV is generated in: {test_time} seconds')","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:56.851317Z","iopub.execute_input":"2024-03-03T20:40:56.851797Z","iopub.status.idle":"2024-03-03T20:40:58.531916Z","shell.execute_reply.started":"2024-03-03T20:40:56.851756Z","shell.execute_reply":"2024-03-03T20:40:58.531006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training and Valaidation\nnum_epochs = PARAMS['epochs']\noptimizer,criterion = get_optimizer_and_criterion(PARAMS,model)  \nfor epoch in range(num_epochs):\n    train_loss, train_accuracy, train_time = train_model(model, criterion, optimizer, undersampled_dataloader, device)\n    print(f\"Epoch {epoch+1}, Train Loss: {train_loss}, Train Accuracy: {train_accuracy}%, Train Time: {train_time} seconds\")\n    wandb.log({\"Training Loss\": train_loss, \"Training Accuracy\": train_accuracy})\n\n    # Validation\n    val_loss, val_accuracy, val_time,precision,recall,f1,conf_matrix = validate_model(model, criterion, val_dataloader, device)\n    print(f\"Validation Loss: {val_loss}, Validation Accuracy: {val_accuracy}% Validation Time: {val_time} seconds\")\n    print(f\"Precision: {precision},Recall: {recall}, F1_Score: {f1}\")\n    wandb.log({\"Precision\": precision, \"Recall\": recall, \"F1 Score\": f1})\n    wandb.log({\"Validation Loss\": val_loss, \"Validation Accuracy\": val_accuracy})    ","metadata":{"execution":{"iopub.status.busy":"2024-03-03T20:40:58.533321Z","iopub.execute_input":"2024-03-03T20:40:58.533697Z","iopub.status.idle":"2024-03-03T20:42:31.965679Z","shell.execute_reply.started":"2024-03-03T20:40:58.533661Z","shell.execute_reply":"2024-03-03T20:42:31.964630Z"},"trusted":true},"execution_count":null,"outputs":[]}]}