{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":24800,"databundleVersionId":1831594,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install iterative-stratification\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:55:29.521115Z","iopub.execute_input":"2025-04-18T16:55:29.52148Z","iopub.status.idle":"2025-04-18T16:55:36.34453Z","shell.execute_reply.started":"2025-04-18T16:55:29.521454Z","shell.execute_reply":"2025-04-18T16:55:36.343309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\nfrom timm import create_model\nfrom collections import Counter\n\n# =============================================================================\n# 1. SET DEVICE\n# =============================================================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# =============================================================================\n# Use the new recommended import for apply_windowing (with fallback)\n# =============================================================================\ntry:\n    from pydicom.pixels import apply_windowing\nexcept ImportError:\n    from pydicom.pixel_data_handlers.util import apply_windowing\n\n# =============================================================================\n# 2. LOAD DATASET & AGGREGATE MULTIPLE LABELS PER IMAGE\n# =============================================================================\ndicom_folder = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train'\ncsv_path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv'\ndf_csv = pd.read_csv(csv_path)\n\nprint(\"CSV sample:\")\nprint(df_csv.sample(5))\nprint(\"Total CSV rows:\", len(df_csv))\nprint(\"Unique images in CSV:\", df_csv.image_id.nunique())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:55:36.346225Z","iopub.execute_input":"2025-04-18T16:55:36.346608Z","iopub.status.idle":"2025-04-18T16:55:42.439055Z","shell.execute_reply.started":"2025-04-18T16:55:36.346578Z","shell.execute_reply":"2025-04-18T16:55:42.43801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# --- Configuration ---\ndicom_folder = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train'  \ncsv_path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv'  \n\n# --- Step 1: Load the CSV and aggregate labels per image ---\ndf_csv = pd.read_csv(csv_path)\n\n# We will group by image_id to remove duplicates and aggregate the labels.\nimage_ids = []\nimage_paths = []\nall_labels = []\n\nfor image_id, group in df_csv.groupby('image_id'):\n    file_path = os.path.join(dicom_folder, f'{image_id}.dicom')\n    # Aggregate all class names for the image, remove duplicates and standardize the case.\n    labels = list({row['class_name'].strip().title() for _, row in group.iterrows()})\n    \n    image_ids.append(image_id)\n    image_paths.append(file_path)\n    all_labels.append(labels)\n\n# Create a DataFrame with one row per image id.\ndf_all = pd.DataFrame({\n    'image_id': image_ids,\n    'image_path': image_paths,\n    'labels': all_labels\n})\nprint(\"Total unique images loaded:\", len(df_all))\nprint(df_all.head())\n\n# --- Step 2: Select 10,000 images for training ---\n# Here, we randomly sample 10,000 unique images for your initial training set.\ntraining_df = df_all.sample(n=10000, random_state=42)\ntraining_image_ids = set(training_df['image_id'].values)\nprint(\"Training set images:\", len(training_df))\n\n# --- Step 3: From the remaining images, select only those with disease findings ---\nremaining_df = df_all[~df_all['image_id'].isin(training_image_ids)]\nprint(\"Remaining images:\", len(remaining_df))\n\n# Filter out images that have only \"No Finding\" in their labels.\ndiseased_df = remaining_df[remaining_df['labels'].apply(lambda x: x != [\"No Finding\"])]\nprint(\"Total diseased images loaded from the remaining set:\", len(diseased_df))\nprint(diseased_df.head())\n\n# Now we have:\n# - 'training_df' containing our initial 10,000 images.\n# - 'diseased_df' containing only the diseased images from the remaining 5,000 images.\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:55:50.991399Z","iopub.execute_input":"2025-04-18T16:55:50.992126Z","iopub.status.idle":"2025-04-18T16:55:55.582744Z","shell.execute_reply.started":"2025-04-18T16:55:50.992095Z","shell.execute_reply":"2025-04-18T16:55:55.58164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter\nimport itertools\n\n# Flatten the list of labels from the final training DataFrame\nall_labels = list(itertools.chain.from_iterable(training_df['labels'].tolist()))\n\n# Count occurrences for each label\nclass_counts = Counter(all_labels)\n\nprint(\"Total number for each class before:\")\nfor label, count in class_counts.items():\n    print(f\"{label}: {count}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:01.215308Z","iopub.execute_input":"2025-04-18T16:56:01.216533Z","iopub.status.idle":"2025-04-18T16:56:01.230355Z","shell.execute_reply.started":"2025-04-18T16:56:01.216484Z","shell.execute_reply":"2025-04-18T16:56:01.229108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_training_df = pd.concat([training_df, diseased_df], ignore_index=True)\nprint(\"Final training dataset size:\", len(final_training_df))\nfinal_training_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:03.832575Z","iopub.execute_input":"2025-04-18T16:56:03.832899Z","iopub.status.idle":"2025-04-18T16:56:03.854706Z","shell.execute_reply.started":"2025-04-18T16:56:03.832877Z","shell.execute_reply":"2025-04-18T16:56:03.853811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter\nimport itertools\n\n# Flatten the list of labels from the final training DataFrame\nall_labels = list(itertools.chain.from_iterable(final_training_df['labels'].tolist()))\n\n# Count occurrences for each label\nclass_counts = Counter(all_labels)\n\nprint(\"Total number for each class after:\")\nfor label, count in class_counts.items():\n    print(f\"{label}: {count}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:07.714377Z","iopub.execute_input":"2025-04-18T16:56:07.714728Z","iopub.status.idle":"2025-04-18T16:56:07.728961Z","shell.execute_reply.started":"2025-04-18T16:56:07.714692Z","shell.execute_reply":"2025-04-18T16:56:07.728116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = final_training_df.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:11.177487Z","iopub.execute_input":"2025-04-18T16:56:11.177911Z","iopub.status.idle":"2025-04-18T16:56:11.184394Z","shell.execute_reply.started":"2025-04-18T16:56:11.17788Z","shell.execute_reply":"2025-04-18T16:56:11.183316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\nall_classes = set()\nfor labs in df['labels']:\n    all_classes.update(labs)\nall_classes = sorted(all_classes)\nlabel_map = {c: i for i, c in enumerate(all_classes)}\nprint(\"Label mapping:\", label_map)\n\n# Compute class frequencies on the training set (using the df from above or your training_df)\npos_counts = np.zeros(len(label_map))\nfor labs in df['labels']:\n    multi_hot = np.zeros(len(label_map))\n    for lab in labs:\n        if lab in label_map:\n            multi_hot[label_map[lab]] = 1\n    pos_counts += multi_hot\nepsilon = 1e-6\nclass_weights = 1.0 / (pos_counts + epsilon)\nprint(\"Class weights:\", class_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:12.997925Z","iopub.execute_input":"2025-04-18T16:56:12.99823Z","iopub.status.idle":"2025-04-18T16:56:13.040716Z","shell.execute_reply.started":"2025-04-18T16:56:12.99821Z","shell.execute_reply":"2025-04-18T16:56:13.039696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# 0. IMPORTS\n# =============================================================================\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport pydicom\nfrom torch.cuda.amp import autocast, GradScaler\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\nimport matplotlib.pyplot as plt\n\n# =============================================================================\n# 1. CONFIGURATION\n# =============================================================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nIMG_SIZE = 256\nCROP_SIZE = 224\nBATCH_SIZE = 32\n\n# =============================================================================\n# 2. UTILITIES\n# =============================================================================\ndef build_label_map(df: pd.DataFrame) -> dict:\n    \"\"\"Creates a mapping from label to index.\"\"\"\n    all_labels = sorted(set(label for labs in df['labels'] for label in labs))\n    return {label: idx for idx, label in enumerate(all_labels)}\n\ndef create_label_matrix(df: pd.DataFrame, label_map: dict) -> np.ndarray:\n    \"\"\"Converts multi-labels into multi-hot encoded matrix.\"\"\"\n    matrix = np.zeros((len(df), len(label_map)), dtype=int)\n    for i, labs in enumerate(df['labels']):\n        for lab in labs:\n            matrix[i, label_map[lab]] = 1\n    return matrix\n\ndef compute_class_weights(df: pd.DataFrame, label_map: dict) -> torch.Tensor:\n    \"\"\"Computes inverse frequency weights for each class.\"\"\"\n    counts = np.zeros(len(label_map))\n    for labs in df['labels']:\n        for lab in labs:\n            counts[label_map[lab]] += 1\n    weights = 1.0 / (counts + 1e-6)\n    return torch.tensor(weights, dtype=torch.float32).to(device)\n\n# =============================================================================\n# 3. LABEL MAPPING\n# =============================================================================\nlabel_map = build_label_map(df)\nprint(\"Label Mapping:\", label_map)\n\n# =============================================================================\n# 4. STRATIFIED SPLIT\n# =============================================================================\ny_all = create_label_matrix(df, label_map)\n\nsplitter = MultilabelStratifiedShuffleSplit(n_splits=1, test_size=0.3, random_state=42)\ntrain_idx, temp_idx = next(splitter.split(np.zeros(len(df)), y_all))\ntrain_df, temp_df = df.iloc[train_idx], df.iloc[temp_idx]\n\ny_temp = create_label_matrix(temp_df, label_map)\nval_idx, test_idx = next(MultilabelStratifiedShuffleSplit(n_splits=1, test_size=0.5, random_state=42).split(\n                         np.zeros(len(temp_df)), y_temp))\nval_df, test_df = temp_df.iloc[val_idx], temp_df.iloc[test_idx]\n\nprint(f\"Train: {len(train_df)} | Val: {len(val_df)} | Test: {len(test_df)}\")\n\n# =============================================================================\n# 5. TRANSFORMS\n# =============================================================================\nbase_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.CenterCrop(CROP_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\nminority_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(0.2, 0.2, 0.2, 0.1),\n    transforms.CenterCrop(CROP_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\n# =============================================================================\n# 6. DATASET\n# =============================================================================\ndef apply_windowing(image: np.ndarray, dicom_data) -> np.ndarray:\n    \"\"\"Apply DICOM windowing using Window Center (WC) and Window Width (WW).\"\"\"\n    try:\n        wc = dicom_data.WindowCenter\n        ww = dicom_data.WindowWidth\n        if isinstance(wc, pydicom.multival.MultiValue):\n            wc = wc[0]\n        if isinstance(ww, pydicom.multival.MultiValue):\n            ww = ww[0]\n\n        img_min = wc - ww // 2\n        img_max = wc + ww // 2\n\n        windowed_image = np.clip(image, img_min, img_max)\n        windowed_image = ((windowed_image - img_min) / (img_max - img_min + 1e-8)) * 255.0\n        return windowed_image.astype(np.uint8)\n\n    except Exception as e:\n        # Fallback to default normalization if WC/WW not found\n        image = image - np.min(image)\n        image = image / (np.max(image) + 1e-8) * 255.0\n        return image.astype(np.uint8)\n\n\nclass XRayDataset(Dataset):\n    def __init__(self, df, transform=None, minority_transform=None, return_raw=False, label_map=None, minority_classes=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.minority_transform = minority_transform\n        self.return_raw = return_raw\n        self.label_map = label_map or build_label_map(df)\n        self.minority_classes = minority_classes or set()\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        path = self.df.iloc[idx]['image_path']\n        labels = self.df.iloc[idx]['labels']\n\n        dicom_data = pydicom.dcmread(path)\n        image = dicom_data.pixel_array.astype(np.float32)\n\n        if hasattr(dicom_data, 'RescaleIntercept') and hasattr(dicom_data, 'RescaleSlope'):\n            image = image * dicom_data.RescaleSlope + dicom_data.RescaleIntercept\n\n        if getattr(dicom_data, 'PhotometricInterpretation', '') == 'MONOCHROME1':\n            image = np.max(image) - image\n\n        image = apply_windowing(image, dicom_data)\n        low, high = np.percentile(image, [1, 99])\n        image = np.clip(image, low, high)\n        image = ((image - image.min()) / (image.max() - image.min() + 1e-8) * 255).astype(np.uint8)\n\n        image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n        image = cv2.resize(image, (IMG_SIZE, IMG_SIZE))\n\n        target = torch.zeros(len(self.label_map), dtype=torch.float32)\n        for lab in labels:\n            target[self.label_map[lab]] = 1.0\n\n        # Check if any minority class exists in target\n        use_minority = any(target[i] == 1 and i in self.minority_classes for i in range(len(target)))\n        image_aug = self.minority_transform(image) if use_minority and self.minority_transform else self.transform(image)\n\n        if self.return_raw:\n            return image, image_aug, target\n        return image_aug, target\n\n# =============================================================================\n# 7. DATA LOADERS\n# =============================================================================\nclass_weights = compute_class_weights(train_df, label_map)\nmedian_weight = torch.median(class_weights).item()\nminority_classes = {i for i, w in enumerate(class_weights.cpu().numpy()) if w > median_weight}\n\ntrain_dataset = XRayDataset(train_df, transform=base_transform, minority_transform=minority_transform,\n                            label_map=label_map, minority_classes=minority_classes)\nval_dataset   = XRayDataset(val_df, transform=base_transform, label_map=label_map)\ntest_dataset  = XRayDataset(test_df, transform=base_transform, label_map=label_map)\n\n# Weighted Sampler\nsample_weights = []\nfor _, target in train_dataset:\n    active_classes = [i for i, val in enumerate(target) if val == 1]\n    weight = np.mean([class_weights[i].item() for i in active_classes]) if active_classes else 1.0\n    sample_weights.append(weight)\n\nsampler = WeightedRandomSampler(weights=sample_weights, num_samples=len(sample_weights), replacement=True)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, sampler=sampler, num_workers=4, pin_memory=True)\nval_loader   = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\ntest_loader  = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\n\n\n# =============================================================================\n# 8. OPTIONAL VISUALIZATION FUNCTION (unchanged)\n# =============================================================================\ndef visualize_before_after():\n    vis_dataset = XRayDataset(train_df, transform=train_transform, minority_transform=minority_transform,\n                              return_raw=True, label_map=common_label_map, minority_classes=minority_classes)\n    seen = {}\n    for raw, transformed, target in vis_dataset:\n        for label_idx, present in enumerate(target):\n            if present == 1 and label_idx not in seen:\n                seen[label_idx] = (raw, transformed)\n        if len(seen) == len(common_label_map):\n            break\n\n    rev_label_map = {v: k for k, v in common_label_map.items()}\n    mean = np.array([0.485, 0.456, 0.406])\n    std  = np.array([0.229, 0.224, 0.225])\n\n    fig, axes = plt.subplots(len(seen), 2, figsize=(12, 6 * len(seen)))\n    if len(seen) == 1:\n        axes = np.expand_dims(axes, axis=0)\n    for label_idx, (raw_img, trans_tensor) in sorted(seen.items()):\n        trans_np = trans_tensor.cpu().numpy().transpose(1,2,0)\n        trans_np = std * trans_np + mean\n        trans_np = np.clip(trans_np, 0, 1)\n        class_name = rev_label_map[label_idx]\n        axes[label_idx, 0].imshow(raw_img)\n        axes[label_idx, 0].set_title(f\"{class_name} - Before\")\n        axes[label_idx, 0].axis(\"off\")\n        axes[label_idx, 1].imshow(trans_np)\n        axes[label_idx, 1].set_title(f\"{class_name} - After\")\n        axes[label_idx, 1].axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\nvisualize_before_after()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T16:56:16.894112Z","iopub.execute_input":"2025-04-18T16:56:16.89452Z","execution_failed":"2025-04-19T08:18:26.453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# DEFINE THE EFFICIENTNETV2-M MODEL (Better than B5)\n# =============================================================================\nclass EfficientNetV2MClassifier(nn.Module):\n    def __init__(self, num_classes):\n        super(EfficientNetV2MClassifier, self).__init__()\n        self.model = timm.create_model('tf_efficientnetv2_m', pretrained=True, num_classes=num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\nnum_classes = len(common_label_map)\nmodel = EfficientNetV2MClassifier(num_classes=num_classes).to(device)\n\n\n# =============================================================================\n# 9. DEFINE A COMBINED LOSS FUNCTION (Weighted BCE + Focal Loss)\n# =============================================================================\ndef combined_loss(outputs, targets, weight=None, alpha_focal=0.25, gamma=2.0, lambda_focal=1.0):\n    \"\"\"\n    Computes a combination of weighted BCEWithLogitsLoss and focal loss.\n    - weight: class weights (per-class)\n    - alpha_focal: balancing factor for focal loss\n    - gamma: focusing parameter for focal loss\n    - lambda_focal: coefficient to balance the focal component with the standard loss.\n    \"\"\"\n    # Standard weighted BCE loss\n    bce_loss = F.binary_cross_entropy_with_logits(outputs, targets, weight=weight, reduction='none')\n    \n    # Compute probabilities\n    probs = torch.sigmoid(outputs)\n    # For focal loss, compute p_t for each element\n    p_t = probs * targets + (1 - probs) * (1 - targets)\n    focal_factor = (1 - p_t) ** gamma\n    focal_loss = alpha_focal * focal_factor * bce_loss\n\n    # Combine the losses (you can adjust the balance as needed)\n    loss = bce_loss + lambda_focal * focal_loss\n    return loss.mean()\n\ncriterion = combined_loss  # Use our custom combined loss\noptimizer = optim.AdamW(model.parameters(), lr=5e-5, weight_decay=0.02)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10, eta_min=1e-6)\nscaler = GradScaler()\n\n# =============================================================================\n# 10. DEFINE THE MIXUP FUNCTION (adapted for multi-label targets)\n# =============================================================================\ndef mixup_data(x, y, alpha=0.2):\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n    batch_size = x.size(0)\n    index = torch.randperm(batch_size).to(device)\n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\n\n# =============================================================================\n# 11. TRAINING FUNCTION WITH MIXUP, EARLY STOPPING, AND 7 EPOCHS\n# =============================================================================\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, scaler,\n                epochs=7, patience=7, use_mixup=True, mixup_alpha=0.2):\n    best_val_loss = float(\"inf\")\n    best_val_acc = 0\n    patience_counter = 0\n\n    train_losses = []\n    val_losses = []\n    val_accuracies = []\n\n    for epoch in range(epochs):\n        model.train()\n        epoch_train_loss = 0.0\n        epoch_train_acc = 0.0\n        n_batches = 0\n\n        for images, targets in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\"):\n            images, targets = images.to(device), targets.to(device)\n            if use_mixup:\n                images, targets_a, targets_b, lam = mixup_data(images, targets, alpha=mixup_alpha)\n            optimizer.zero_grad()\n            with torch.amp.autocast(device_type='cuda'):\n                outputs = model(images)\n                if use_mixup:\n                    loss = lam * criterion(outputs, targets_a) + (1 - lam) * criterion(outputs, targets_b)\n                else:\n                    loss = criterion(outputs, targets)\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            epoch_train_loss += loss.item() * images.size(0)\n            preds = (torch.sigmoid(outputs) > 0.5).float()\n            batch_acc = (preds == targets).all(dim=1).float().mean().item()\n            epoch_train_acc += batch_acc\n            n_batches += 1\n\n        scheduler.step()\n        train_loss_epoch = epoch_train_loss / len(train_loader.dataset)\n        train_acc_epoch = epoch_train_acc / n_batches\n        train_losses.append(train_loss_epoch)\n\n        model.eval()\n        epoch_val_loss = 0.0\n        epoch_val_acc = 0.0\n        n_batches_val = 0\n        with torch.no_grad():\n            for images, targets in val_loader:\n                images, targets = images.to(device), targets.to(device)\n                with torch.amp.autocast(device_type='cuda'):\n                    outputs = model(images)\n                    loss = criterion(outputs, targets)\n                epoch_val_loss += loss.item() * images.size(0)\n                preds = (torch.sigmoid(outputs) > 0.5).float()\n                batch_acc = (preds == targets).all(dim=1).float().mean().item()\n                epoch_val_acc += batch_acc\n                n_batches_val += 1\n        val_loss_epoch = epoch_val_loss / len(val_loader.dataset)\n        val_acc_epoch = epoch_val_acc / n_batches_val\n        val_losses.append(val_loss_epoch)\n        val_accuracies.append(val_acc_epoch)\n\n        print(f\"Epoch {epoch+1}: Train Loss: {train_loss_epoch:.4f} | Val Loss: {val_loss_epoch:.4f} | Val Acc (exact match): {val_acc_epoch*100:.2f}%\")\n\n        if val_loss_epoch < best_val_loss:\n            best_val_loss = val_loss_epoch\n            best_val_acc = val_acc_epoch\n            patience_counter = 0\n            torch.save(model.state_dict(), \"best_model.pth\")\n        else:\n            patience_counter += 1\n            if patience_counter >= patience:\n                print(\"Early stopping triggered\")\n                break\n\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    return model, best_val_acc, best_val_loss, train_losses, val_losses, val_accuracies\n\n# =============================================================================\n# 12. TRAIN THE MODEL (7 Epochs)\n# =============================================================================\nmodel, best_val_acc, best_val_loss, train_losses, val_losses, val_accuracies = train_model(\n    model, train_loader, val_loader, criterion, optimizer, scheduler, scaler,\n    epochs=7, patience=7, use_mixup=True, mixup_alpha=0.2\n)\nprint(f\"\\nFinal Best Validation Exact Match Accuracy: {best_val_acc*100:.2f}%\")\nprint(f\"Final Best Validation Loss: {best_val_loss:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}