{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":319883136,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport ast\nimport json\nimport time\nimport copy\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport albumentations as A\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.metrics import roc_auc_score, average_precision_score\nfrom timm.scheduler import CosineLRScheduler\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ntorch.backends.cudnn.benchmark = True\n\nCACHE_ROOT = \"/kaggle/input/notebooks/vhthai/cache\"\n\nif not (\n    os.path.exists(os.path.join(CACHE_ROOT, \"train_split.csv\")) and\n    os.path.exists(os.path.join(CACHE_ROOT, \"val_split.csv\")) and\n    os.path.exists(os.path.join(CACHE_ROOT, \"cache\", \"train\")) and\n    os.path.exists(os.path.join(CACHE_ROOT, \"cache\", \"val\"))\n):\n    found_root = None\n\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        if (\n            \"train_split.csv\" in files and\n            \"val_split.csv\" in files and\n            os.path.isdir(os.path.join(root, \"cache\", \"train\")) and\n            os.path.isdir(os.path.join(root, \"cache\", \"val\"))\n        ):\n            found_root = root\n            break\n\n    if found_root is not None:\n        CACHE_ROOT = found_root\n\nTRAIN_CSV = os.path.join(CACHE_ROOT, \"train_split.csv\")\nVAL_CSV = os.path.join(CACHE_ROOT, \"val_split.csv\")\nTRAIN_DIR = os.path.join(CACHE_ROOT, \"cache\", \"train\")\nVAL_DIR = os.path.join(CACHE_ROOT, \"cache\", \"val\")\n\nOUTPUT_DIR = \"/kaggle/working/maxvit_final_outputs_v3\"\nFIGURE_DIR = os.path.join(OUTPUT_DIR, \"figures\")\n\nos.makedirs(OUTPUT_DIR, exist_ok=True)\nos.makedirs(FIGURE_DIR, exist_ok=True)\n\nprint(\"DEVICE:\", DEVICE)\nprint(\"CACHE_ROOT:\", CACHE_ROOT)\nprint(\"TRAIN_CSV:\", TRAIN_CSV)\nprint(\"VAL_CSV:\", VAL_CSV)\nprint(\"TRAIN_DIR:\", TRAIN_DIR)\nprint(\"VAL_DIR:\", VAL_DIR)\nprint(\"OUTPUT_DIR:\", OUTPUT_DIR)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CLASS_NAMES = [\n    \"Aortic enlargement\",\n    \"Atelectasis\",\n    \"Calcification\",\n    \"Cardiomegaly\",\n    \"Consolidation\",\n    \"ILD\",\n    \"Infiltration\",\n    \"Lung Opacity\",\n    \"Nodule/Mass\",\n    \"Other lesion\",\n    \"Pleural effusion\",\n    \"Pleural thickening\",\n    \"Pneumothorax\",\n    \"Pulmonary fibrosis\",\n    \"No finding\"\n]\n\nNUM_CLASSES = 15\nNO_FINDING_ID = 14\n\nMODEL_NAME = \"maxvit_tiny_tf_384\"\nIMG_SIZE = 384\n\nBATCH_SIZE = 8\nACCUMULATION_STEPS = 4\nEPOCHS = 40\nEARLY_STOP_PATIENCE = 8\nVAL_LOSS_DELTA = 1e-4\n\nLR_HEAD = 8e-5\nLR_BACKBONE = 5e-6\nWEIGHT_DECAY = 1e-4\nNUM_WORKERS = 4\n\nDROP_PATH_RATE = 0.10\nMAX_POS_WEIGHT = 8.0\nLABEL_SMOOTHING = 0.03\nNO_FINDING_LOSS_WEIGHT = 0.65\n\nEMA_DECAY = 0.995\nUSE_TTA_VALIDATION = True\nUSE_VAL_LOSS_PATIENCE_RELIEF = True\nMAX_GRAD_NORM = 1.0\n\nIMAGENET_STATS = {\n    \"mean\": [0.485, 0.456, 0.406],\n    \"std\": [0.229, 0.224, 0.225]\n}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_target(x):\n    if isinstance(x, str):\n        return np.array(ast.literal_eval(x), dtype=np.float32)\n    return np.array(x, dtype=np.float32)\n\ntrain_df = pd.read_csv(TRAIN_CSV)\nval_df = pd.read_csv(VAL_CSV)\n\ntrain_df[\"target\"] = train_df[\"target\"].apply(parse_target)\nval_df[\"target\"] = val_df[\"target\"].apply(parse_target)\n\ntrain_targets = np.vstack(train_df[\"target\"].values).astype(np.float32)\nval_targets = np.vstack(val_df[\"target\"].values).astype(np.float32)\n\npos_counts = train_targets.sum(axis=0)\nneg_counts = len(train_targets) - pos_counts\npos_ratio = pos_counts / len(train_targets)\n\nimbalance_df = pd.DataFrame({\n    \"class_id\": np.arange(NUM_CLASSES),\n    \"class_name\": CLASS_NAMES,\n    \"positive_count\": pos_counts.astype(int),\n    \"negative_count\": neg_counts.astype(int),\n    \"positive_ratio\": pos_ratio\n})\n\nimbalance_df.to_csv(os.path.join(OUTPUT_DIR, \"maxvit_class_imbalance_report.csv\"), index=False)\n\ndisplay(imbalance_df)\n\nprint(\"Train:\", len(train_df))\nprint(\"Val:\", len(val_df))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_pos_weight(targets, no_finding_id=14, max_pos_weight=8.0):\n    targets = targets.astype(np.float32)\n\n    pos_counts = targets.sum(axis=0) + 1e-6\n    neg_counts = len(targets) - pos_counts + 1e-6\n\n    pos_weight = neg_counts / pos_counts\n    pos_weight = np.power(pos_weight, 0.45)\n    pos_weight = np.clip(pos_weight, 1.0, max_pos_weight)\n\n    pos_weight[no_finding_id] = 0.75\n\n    return torch.tensor(pos_weight, dtype=torch.float32)\n\npos_weight = build_pos_weight(\n    train_targets,\n    no_finding_id=NO_FINDING_ID,\n    max_pos_weight=MAX_POS_WEIGHT\n).to(DEVICE)\n\npos_weight_df = pd.DataFrame({\n    \"class_id\": np.arange(NUM_CLASSES),\n    \"class_name\": CLASS_NAMES,\n    \"pos_weight\": pos_weight.detach().cpu().numpy()\n})\n\npos_weight_df.to_csv(os.path.join(OUTPUT_DIR, \"maxvit_pos_weight.csv\"), index=False)\n\ndisplay(pos_weight_df)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class XrayDataset(Dataset):\n    def __init__(self, df, img_dir, img_size=384, is_train=True):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.img_size = img_size\n        self.is_train = is_train\n\n        if is_train:\n            self.transform = A.Compose([\n                A.Resize(img_size, img_size),\n                A.HorizontalFlip(p=0.35),\n                A.ShiftScaleRotate(\n                    shift_limit=0.035,\n                    scale_limit=0.04,\n                    rotate_limit=5,\n                    border_mode=cv2.BORDER_CONSTANT,\n                    value=0,\n                    p=0.30\n                ),\n                A.RandomBrightnessContrast(\n                    brightness_limit=0.05,\n                    contrast_limit=0.05,\n                    p=0.30\n                ),\n                A.GaussNoise(var_limit=(5.0, 15.0), p=0.12),\n                A.Normalize(\n                    mean=IMAGENET_STATS[\"mean\"],\n                    std=IMAGENET_STATS[\"std\"]\n                ),\n                ToTensorV2()\n            ])\n        else:\n            self.transform = A.Compose([\n                A.Resize(img_size, img_size),\n                A.Normalize(\n                    mean=IMAGENET_STATS[\"mean\"],\n                    std=IMAGENET_STATS[\"std\"]\n                ),\n                ToTensorV2()\n            ])\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row[\"image_id\"]\n        path = os.path.join(self.img_dir, image_id + \".npy\")\n\n        img = np.load(path)\n\n        if img.dtype != np.uint8:\n            if img.max() <= 1.0:\n                img = (img * 255).astype(np.uint8)\n            else:\n                img = img.astype(np.uint8)\n\n        if img.ndim == 2:\n            img = np.stack([img] * 3, axis=-1)\n\n        if img.shape[-1] != 3:\n            img = img[..., :3]\n\n        image = self.transform(image=img)[\"image\"]\n        target = torch.tensor(row[\"target\"], dtype=torch.float32)\n\n        return image, target\n\ntrain_loader = DataLoader(\n    XrayDataset(train_df, TRAIN_DIR, IMG_SIZE, True),\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    drop_last=True\n)\n\nval_loader = DataLoader(\n    XrayDataset(val_df, VAL_DIR, IMG_SIZE, False),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StableSmoothBCELoss(nn.Module):\n    def __init__(\n        self,\n        pos_weight=None,\n        smoothing=0.03,\n        no_finding_id=14,\n        no_finding_loss_weight=0.65\n    ):\n        super().__init__()\n        self.pos_weight = pos_weight\n        self.smoothing = smoothing\n        self.no_finding_id = no_finding_id\n        self.no_finding_loss_weight = no_finding_loss_weight\n\n    def forward(self, logits, targets):\n        if self.smoothing > 0:\n            targets = targets * (1.0 - self.smoothing) + 0.5 * self.smoothing\n\n        loss = F.binary_cross_entropy_with_logits(\n            logits,\n            targets,\n            pos_weight=self.pos_weight,\n            reduction=\"none\"\n        )\n\n        class_weight = torch.ones(NUM_CLASSES, device=logits.device)\n        class_weight[self.no_finding_id] = self.no_finding_loss_weight\n\n        loss = loss * class_weight.view(1, -1)\n\n        return loss.mean()\n\ncriterion = StableSmoothBCELoss(\n    pos_weight=pos_weight,\n    smoothing=LABEL_SMOOTHING,\n    no_finding_id=NO_FINDING_ID,\n    no_finding_loss_weight=NO_FINDING_LOSS_WEIGHT\n)\n\nmonitor_criterion = nn.BCEWithLogitsLoss()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = timm.create_model(\n    MODEL_NAME,\n    pretrained=True,\n    num_classes=NUM_CLASSES,\n    drop_path_rate=DROP_PATH_RATE\n)\n\nmodel = model.to(DEVICE)\n\nema_model = copy.deepcopy(model)\nema_model.eval()\n\nfor p in ema_model.parameters():\n    p.requires_grad_(False)\n\ntry:\n    head_params = list(model.get_classifier().parameters())\nexcept:\n    head_params = []\n\nhead_param_ids = set(id(p) for p in head_params)\n\nif len(head_params) > 0:\n    backbone_params = [\n        p for p in model.parameters()\n        if id(p) not in head_param_ids\n    ]\nelse:\n    head_params = []\n    backbone_params = []\n\n    for name, param in model.named_parameters():\n        if (\"head\" in name) or (\"classifier\" in name) or (\"fc\" in name):\n            head_params.append(param)\n        else:\n            backbone_params.append(param)\n\noptimizer = torch.optim.AdamW(\n    [\n        {\"params\": backbone_params, \"lr\": LR_BACKBONE},\n        {\"params\": head_params, \"lr\": LR_HEAD}\n    ],\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler = CosineLRScheduler(\n    optimizer,\n    t_initial=EPOCHS,\n    warmup_t=4,\n    warmup_lr_init=1e-7,\n    lr_min=5e-7\n)\n\ntry:\n    scaler = torch.amp.GradScaler(device=DEVICE, enabled=(DEVICE == \"cuda\"))\nexcept:\n    scaler = torch.cuda.amp.GradScaler(enabled=(DEVICE == \"cuda\"))\n\ndef autocast_context():\n    return torch.amp.autocast(device_type=DEVICE, enabled=(DEVICE == \"cuda\"))\n\n@torch.no_grad()\ndef update_ema(model, ema_model, decay=0.995):\n    model_state = model.state_dict()\n    ema_state = ema_model.state_dict()\n\n    for key in ema_state.keys():\n        if key not in model_state:\n            continue\n\n        if torch.is_floating_point(ema_state[key]):\n            ema_state[key].mul_(decay).add_(model_state[key], alpha=1.0 - decay)\n        else:\n            ema_state[key].copy_(model_state[key])\n\nprint(\"Model ready:\", MODEL_NAME)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_multilabel_metrics(y_true, y_pred, class_names, no_finding_id=14):\n    per_class = []\n    auc_list = []\n    abnormal_auc_list = []\n    ap_list = []\n    abnormal_ap_list = []\n\n    for c, name in enumerate(class_names):\n        if len(np.unique(y_true[:, c])) < 2:\n            auc = np.nan\n            ap = np.nan\n        else:\n            auc = roc_auc_score(y_true[:, c], y_pred[:, c])\n            ap = average_precision_score(y_true[:, c], y_pred[:, c])\n\n        per_class.append({\n            \"class_id\": c,\n            \"class_name\": name,\n            \"auc\": auc,\n            \"ap\": ap,\n            \"positive_samples\": int(y_true[:, c].sum()),\n            \"positive_ratio\": float(y_true[:, c].mean())\n        })\n\n        if not np.isnan(auc):\n            auc_list.append(auc)\n\n            if c != no_finding_id:\n                abnormal_auc_list.append(auc)\n\n        if not np.isnan(ap):\n            ap_list.append(ap)\n\n            if c != no_finding_id:\n                abnormal_ap_list.append(ap)\n\n    metrics = {\n        \"macro_auc\": float(np.nanmean(auc_list)),\n        \"abnormal_macro_auc\": float(np.nanmean(abnormal_auc_list)),\n        \"macro_ap\": float(np.nanmean(ap_list)),\n        \"abnormal_macro_ap\": float(np.nanmean(abnormal_ap_list))\n    }\n\n    try:\n        metrics[\"micro_auc\"] = float(roc_auc_score(y_true.ravel(), y_pred.ravel()))\n    except:\n        metrics[\"micro_auc\"] = np.nan\n\n    return metrics, pd.DataFrame(per_class)\n\n@torch.no_grad()\ndef evaluate_model(eval_model, loader, use_tta=True):\n    eval_model.eval()\n\n    val_bce_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    for imgs, labels in tqdm(loader, desc=\"Valid\"):\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        with autocast_context():\n            logits = eval_model(imgs)\n            loss = monitor_criterion(logits, labels)\n\n            if use_tta:\n                flipped_imgs = torch.flip(imgs, dims=[3])\n                logits_flip = eval_model(flipped_imgs)\n                logits = (logits + logits_flip) / 2.0\n\n        probs = torch.sigmoid(logits).detach().cpu().numpy()\n\n        val_bce_loss += loss.item()\n        all_preds.append(probs)\n        all_labels.append(labels.detach().cpu().numpy())\n\n    val_bce_loss = val_bce_loss / len(loader)\n    all_preds = np.vstack(all_preds)\n    all_labels = np.vstack(all_labels)\n\n    metrics, per_class_df = compute_multilabel_metrics(\n        all_labels,\n        all_preds,\n        CLASS_NAMES,\n        no_finding_id=NO_FINDING_ID\n    )\n\n    return val_bce_loss, metrics, per_class_df, all_preds, all_labels","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_abnormal_macro_auc = -1.0\nbest_val_bce_loss_for_patience = float(\"inf\")\nepochs_without_improvement = 0\n\nhistory = {\n    \"epoch\": [],\n    \"train_loss\": [],\n    \"train_bce_loss\": [],\n    \"val_bce_loss\": [],\n    \"macro_auc\": [],\n    \"abnormal_macro_auc\": [],\n    \"micro_auc\": [],\n    \"macro_ap\": [],\n    \"abnormal_macro_ap\": [],\n    \"lr_backbone\": [],\n    \"lr_head\": [],\n    \"epoch_time_seconds\": []\n}\n\ntotal_start_time = time.time()\n\nfor epoch in range(EPOCHS):\n    epoch_start_time = time.time()\n\n    model.train()\n    running_loss = 0.0\n    running_bce_loss = 0.0\n    optimizer.zero_grad(set_to_none=True)\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch + 1}/{EPOCHS} - Train\")\n\n    for step, (imgs, labels) in enumerate(pbar):\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        with autocast_context():\n            outputs = model(imgs)\n            train_loss_value = criterion(outputs, labels)\n            bce_loss_value = monitor_criterion(outputs, labels)\n            loss = train_loss_value / ACCUMULATION_STEPS\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % ACCUMULATION_STEPS == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=MAX_GRAD_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n            update_ema(model, ema_model, decay=EMA_DECAY)\n\n        running_loss += train_loss_value.item()\n        running_bce_loss += bce_loss_value.item()\n\n        pbar.set_postfix(\n            loss=f\"{train_loss_value.item():.4f}\",\n            bce=f\"{bce_loss_value.item():.4f}\"\n        )\n\n    if (len(train_loader) % ACCUMULATION_STEPS) != 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=MAX_GRAD_NORM)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        update_ema(model, ema_model, decay=EMA_DECAY)\n\n    train_loss = running_loss / len(train_loader)\n    train_bce_loss = running_bce_loss / len(train_loader)\n\n    val_bce_loss, metrics, per_class_metrics_df, all_preds, all_labels = evaluate_model(\n        ema_model,\n        val_loader,\n        use_tta=USE_TTA_VALIDATION\n    )\n\n    macro_auc = metrics[\"macro_auc\"]\n    abnormal_macro_auc = metrics[\"abnormal_macro_auc\"]\n    micro_auc = metrics[\"micro_auc\"]\n    macro_ap = metrics[\"macro_ap\"]\n    abnormal_macro_ap = metrics[\"abnormal_macro_ap\"]\n\n    scheduler.step(epoch + 1)\n\n    current_lrs = [group[\"lr\"] for group in optimizer.param_groups]\n    lr_backbone = current_lrs[0]\n    lr_head = current_lrs[1] if len(current_lrs) > 1 else current_lrs[0]\n\n    epoch_time = time.time() - epoch_start_time\n\n    history[\"epoch\"].append(epoch + 1)\n    history[\"train_loss\"].append(train_loss)\n    history[\"train_bce_loss\"].append(train_bce_loss)\n    history[\"val_bce_loss\"].append(val_bce_loss)\n    history[\"macro_auc\"].append(macro_auc)\n    history[\"abnormal_macro_auc\"].append(abnormal_macro_auc)\n    history[\"micro_auc\"].append(micro_auc)\n    history[\"macro_ap\"].append(macro_ap)\n    history[\"abnormal_macro_ap\"].append(abnormal_macro_ap)\n    history[\"lr_backbone\"].append(lr_backbone)\n    history[\"lr_head\"].append(lr_head)\n    history[\"epoch_time_seconds\"].append(epoch_time)\n\n    history_df = pd.DataFrame(history)\n    history_df.to_csv(\n        os.path.join(OUTPUT_DIR, \"maxvit_final_training_history.csv\"),\n        index=False\n    )\n\n    print(f\"\\nEpoch {epoch + 1}/{EPOCHS}\")\n    print(f\"Train Loss:          {train_loss:.4f}\")\n    print(f\"Train BCE Loss:      {train_bce_loss:.4f}\")\n    print(f\"Val BCE Loss:        {val_bce_loss:.4f}\")\n    print(f\"Macro AUC:           {macro_auc:.6f}\")\n    print(f\"Abnormal Macro AUC:  {abnormal_macro_auc:.6f}\")\n    print(f\"Micro AUC:           {micro_auc:.6f}\")\n    print(f\"Macro AP:            {macro_ap:.6f}\")\n    print(f\"Abnormal Macro AP:   {abnormal_macro_ap:.6f}\")\n    print(f\"LR Backbone:         {lr_backbone:.8f}\")\n    print(f\"LR Head:             {lr_head:.8f}\")\n    print(f\"Time:                {int(epoch_time // 60)}m {int(epoch_time % 60)}s\")\n\n    auc_improved = abnormal_macro_auc > best_abnormal_macro_auc\n    val_loss_improved = val_bce_loss < best_val_bce_loss_for_patience - VAL_LOSS_DELTA\n\n    if val_loss_improved:\n        best_val_bce_loss_for_patience = val_bce_loss\n\n    if auc_improved:\n        best_abnormal_macro_auc = abnormal_macro_auc\n        epochs_without_improvement = 0\n\n        checkpoint = {\n            \"epoch\": epoch + 1,\n            \"model_name\": MODEL_NAME,\n            \"model_state_dict\": ema_model.state_dict(),\n            \"raw_model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"best_abnormal_macro_auc\": best_abnormal_macro_auc,\n            \"macro_auc\": macro_auc,\n            \"micro_auc\": micro_auc,\n            \"macro_ap\": macro_ap,\n            \"abnormal_macro_ap\": abnormal_macro_ap,\n            \"train_loss\": train_loss,\n            \"train_bce_loss\": train_bce_loss,\n            \"val_bce_loss\": val_bce_loss,\n            \"img_size\": IMG_SIZE,\n            \"num_classes\": NUM_CLASSES,\n            \"class_names\": CLASS_NAMES,\n            \"no_finding_id\": NO_FINDING_ID,\n            \"imagenet_stats\": IMAGENET_STATS,\n            \"seed\": SEED,\n            \"ema_decay\": EMA_DECAY,\n            \"use_tta_validation\": USE_TTA_VALIDATION,\n            \"early_stop_patience\": EARLY_STOP_PATIENCE,\n            \"val_loss_delta\": VAL_LOSS_DELTA,\n            \"drop_path_rate\": DROP_PATH_RATE,\n            \"lr_head\": LR_HEAD,\n            \"lr_backbone\": LR_BACKBONE,\n            \"label_smoothing\": LABEL_SMOOTHING,\n            \"no_finding_loss_weight\": NO_FINDING_LOSS_WEIGHT\n        }\n\n        torch.save(\n            checkpoint,\n            os.path.join(OUTPUT_DIR, \"best_maxvit_final_checkpoint.pth\")\n        )\n\n        torch.save(\n            ema_model.state_dict(),\n            os.path.join(OUTPUT_DIR, \"best_maxvit_final_model.pth\")\n        )\n\n        per_class_metrics_df.to_csv(\n            os.path.join(OUTPUT_DIR, \"best_maxvit_final_per_class_metrics.csv\"),\n            index=False\n        )\n\n        np.savez_compressed(\n            os.path.join(OUTPUT_DIR, \"best_maxvit_final_val_predictions.npz\"),\n            preds=all_preds,\n            labels=all_labels\n        )\n\n        print(f\"--> Saved New Best Model | Abnormal Macro AUC: {best_abnormal_macro_auc:.6f}\")\n\n    else:\n        if USE_VAL_LOSS_PATIENCE_RELIEF and val_loss_improved:\n            epochs_without_improvement = max(0, epochs_without_improvement - 1)\n            print(f\"AUC not improved, but Val BCE Loss improved. Patience: {epochs_without_improvement}/{EARLY_STOP_PATIENCE}\")\n        else:\n            epochs_without_improvement += 1\n            print(f\"No improvement for {epochs_without_improvement}/{EARLY_STOP_PATIENCE} epochs\")\n\n    if epochs_without_improvement >= EARLY_STOP_PATIENCE:\n        print(f\"Early stopping triggered at epoch {epoch + 1}\")\n        break\n\n    torch.cuda.empty_cache()\n\ntotal_time = time.time() - total_start_time\n\nprint(\"\\nTraining loop finished\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_df = pd.DataFrame(history)\nhistory_df.to_csv(\n    os.path.join(OUTPUT_DIR, \"maxvit_final_training_history.csv\"),\n    index=False\n)\n\nbest_idx = int(np.argmax(history_df[\"abnormal_macro_auc\"].values))\nbest_epoch = int(history_df.loc[best_idx, \"epoch\"])\n\nsummary = {\n    \"model\": \"MaxViT Tiny Final v3\",\n    \"model_name\": MODEL_NAME,\n    \"best_epoch\": best_epoch,\n    \"best_abnormal_macro_auc\": float(history_df.loc[best_idx, \"abnormal_macro_auc\"]),\n    \"best_macro_auc\": float(history_df.loc[best_idx, \"macro_auc\"]),\n    \"best_micro_auc\": float(history_df.loc[best_idx, \"micro_auc\"]),\n    \"best_macro_ap\": float(history_df.loc[best_idx, \"macro_ap\"]),\n    \"best_abnormal_macro_ap\": float(history_df.loc[best_idx, \"abnormal_macro_ap\"]),\n    \"best_val_bce_loss\": float(history_df.loc[best_idx, \"val_bce_loss\"]),\n    \"epochs_planned\": EPOCHS,\n    \"epochs_trained\": int(len(history_df)),\n    \"early_stopping_used\": True,\n    \"early_stop_patience\": EARLY_STOP_PATIENCE,\n    \"val_loss_delta\": VAL_LOSS_DELTA,\n    \"val_loss_patience_relief\": USE_VAL_LOSS_PATIENCE_RELIEF,\n    \"ema_decay\": EMA_DECAY,\n    \"tta_validation\": USE_TTA_VALIDATION,\n    \"img_size\": IMG_SIZE,\n    \"batch_size\": BATCH_SIZE,\n    \"accumulation_steps\": ACCUMULATION_STEPS,\n    \"effective_batch_size\": BATCH_SIZE * ACCUMULATION_STEPS,\n    \"lr_head\": LR_HEAD,\n    \"lr_backbone\": LR_BACKBONE,\n    \"weight_decay\": WEIGHT_DECAY,\n    \"drop_path_rate\": DROP_PATH_RATE,\n    \"label_smoothing\": LABEL_SMOOTHING,\n    \"no_finding_loss_weight\": NO_FINDING_LOSS_WEIGHT,\n    \"total_time_seconds\": total_time,\n    \"total_time_hours\": total_time / 3600,\n    \"output_folder\": OUTPUT_DIR\n}\n\nwith open(os.path.join(OUTPUT_DIR, \"maxvit_final_training_summary.json\"), \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nsummary_df = pd.DataFrame([\n    [\"Model\", \"MaxViT Tiny Final v3\"],\n    [\"Model Name\", MODEL_NAME],\n    [\"Best Epoch\", best_epoch],\n    [\"Epochs Trained\", f\"{len(history_df)}/{EPOCHS}\"],\n    [\"Best Abnormal Macro AUC\", round(float(history_df.loc[best_idx, \"abnormal_macro_auc\"]), 6)],\n    [\"Best Macro AUC\", round(float(history_df.loc[best_idx, \"macro_auc\"]), 6)],\n    [\"Best Micro AUC\", round(float(history_df.loc[best_idx, \"micro_auc\"]), 6)],\n    [\"Best Macro AP\", round(float(history_df.loc[best_idx, \"macro_ap\"]), 6)],\n    [\"Best Abnormal Macro AP\", round(float(history_df.loc[best_idx, \"abnormal_macro_ap\"]), 6)],\n    [\"Best Validation BCE Loss\", round(float(history_df.loc[best_idx, \"val_bce_loss\"]), 6)],\n    [\"Early Stopping Patience\", EARLY_STOP_PATIENCE],\n    [\"Val Loss Patience Relief\", str(USE_VAL_LOSS_PATIENCE_RELIEF)],\n    [\"EMA\", f\"Yes, decay={EMA_DECAY}\"],\n    [\"TTA Validation\", str(USE_TTA_VALIDATION)],\n    [\"Image Size\", IMG_SIZE],\n    [\"Batch Size\", BATCH_SIZE],\n    [\"Accumulation Steps\", ACCUMULATION_STEPS],\n    [\"Effective Batch Size\", BATCH_SIZE * ACCUMULATION_STEPS],\n    [\"Total Training Time\", f\"{total_time / 3600:.2f} hours\"],\n    [\"Output Folder\", OUTPUT_DIR]\n], columns=[\"Metric\", \"Value\"])\n\ndisplay(summary_df)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.rcParams.update({\n    \"font.size\": 11,\n    \"axes.titlesize\": 14,\n    \"axes.labelsize\": 12,\n    \"legend.fontsize\": 10,\n    \"figure.dpi\": 120\n})\n\nmax_epoch = int(history_df[\"epoch\"].max())\nxticks = list(range(5, max_epoch + 1, 5))\n\nif len(xticks) == 0 or xticks[-1] != max_epoch:\n    if max_epoch not in xticks:\n        xticks.append(max_epoch)\n\nfig, axes = plt.subplots(1, 2, figsize=(15, 5))\n\naxes[0].plot(\n    history_df[\"epoch\"],\n    history_df[\"train_bce_loss\"],\n    marker=\"o\",\n    markersize=4,\n    linewidth=2,\n    label=\"Train BCE Loss\"\n)\n\naxes[0].plot(\n    history_df[\"epoch\"],\n    history_df[\"val_bce_loss\"],\n    marker=\"s\",\n    markersize=4,\n    linewidth=2,\n    label=\"Validation BCE Loss\"\n)\n\naxes[0].scatter(\n    best_epoch,\n    history_df.loc[best_idx, \"val_bce_loss\"],\n    s=120,\n    zorder=5,\n    label=f\"Best AUC Epoch {best_epoch}\"\n)\n\naxes[0].set_title(\"MaxViT Training and Validation Loss\")\naxes[0].set_xlabel(\"Epoch\")\naxes[0].set_ylabel(\"BCE Loss\")\naxes[0].set_xticks(xticks)\naxes[0].grid(True, linestyle=\"--\", alpha=0.4)\naxes[0].legend()\n\naxes[1].plot(\n    history_df[\"epoch\"],\n    history_df[\"macro_auc\"],\n    marker=\"o\",\n    markersize=4,\n    linewidth=2,\n    label=\"Macro AUC\"\n)\n\naxes[1].plot(\n    history_df[\"epoch\"],\n    history_df[\"abnormal_macro_auc\"],\n    marker=\"s\",\n    markersize=4,\n    linewidth=2,\n    label=\"Abnormal Macro AUC\"\n)\n\naxes[1].scatter(\n    best_epoch,\n    history_df.loc[best_idx, \"abnormal_macro_auc\"],\n    s=120,\n    zorder=5,\n    label=f\"Best = {float(history_df.loc[best_idx, 'abnormal_macro_auc']):.6f}\"\n)\n\naxes[1].axhline(\n    float(history_df.loc[best_idx, \"abnormal_macro_auc\"]),\n    linestyle=\"--\",\n    linewidth=1.5,\n    alpha=0.7\n)\n\naxes[1].set_title(\"MaxViT Validation AUC Progression\")\naxes[1].set_xlabel(\"Epoch\")\naxes[1].set_ylabel(\"AUC\")\naxes[1].set_xticks(xticks)\naxes[1].set_ylim(\n    max(\n        0.0,\n        min(\n            history_df[\"abnormal_macro_auc\"].min(),\n            history_df[\"macro_auc\"].min()\n        ) - 0.05\n    ),\n    min(\n        1.0,\n        max(\n            history_df[\"abnormal_macro_auc\"].max(),\n            history_df[\"macro_auc\"].max()\n        ) + 0.03\n    )\n)\naxes[1].grid(True, linestyle=\"--\", alpha=0.4)\naxes[1].legend()\n\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.png\"),\n    bbox_inches=\"tight\",\n    dpi=200\n)\n\nplt.savefig(\n    os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.pdf\"),\n    bbox_inches=\"tight\"\n)\n\nplt.show()\n\nprint(\"Saved:\")\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.png\"))\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.pdf\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_metrics_path = os.path.join(OUTPUT_DIR, \"best_maxvit_final_per_class_metrics.csv\")\n\nif os.path.exists(best_metrics_path):\n    best_per_class_df = pd.read_csv(best_metrics_path)\n\n    plt.figure(figsize=(10, 6))\n\n    plot_df = best_per_class_df.sort_values(\"auc\", ascending=True)\n\n    plt.barh(plot_df[\"class_name\"], plot_df[\"auc\"])\n\n    plt.axvline(\n        plot_df[\"auc\"].mean(),\n        linestyle=\"--\",\n        linewidth=1.5,\n        label=f\"Mean AUC = {plot_df['auc'].mean():.4f}\"\n    )\n\n    plt.title(\"Per-Class AUC - MaxViT\")\n    plt.xlabel(\"AUC\")\n    plt.ylabel(\"Class\")\n    plt.xlim(0.5, 1.0)\n    plt.grid(axis=\"x\", linestyle=\"--\", alpha=0.4)\n    plt.legend()\n    plt.tight_layout()\n\n    plt.savefig(\n        os.path.join(FIGURE_DIR, \"maxvit_final_per_class_auc.png\"),\n        bbox_inches=\"tight\",\n        dpi=200\n    )\n\n    plt.savefig(\n        os.path.join(FIGURE_DIR, \"maxvit_final_per_class_auc.pdf\"),\n        bbox_inches=\"tight\"\n    )\n\n    plt.show()\n\n    display(best_per_class_df[[\n        \"class_id\",\n        \"class_name\",\n        \"auc\",\n        \"ap\",\n        \"positive_samples\",\n        \"positive_ratio\"\n    ]])\nelse:\n    print(\"Per-class metrics file not found.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nTraining Completed\")\nprint(f\"Best Epoch: {best_epoch}\")\nprint(f\"Best Abnormal Macro AUC: {float(history_df.loc[best_idx, 'abnormal_macro_auc']):.6f}\")\nprint(f\"Best Macro AUC: {float(history_df.loc[best_idx, 'macro_auc']):.6f}\")\nprint(f\"Best Micro AUC: {float(history_df.loc[best_idx, 'micro_auc']):.6f}\")\nprint(f\"Epochs Trained: {len(history_df)}/{EPOCHS}\")\nprint(f\"Total Time: {int(total_time // 3600)}h {int((total_time % 3600) // 60)}m {int(total_time % 60)}s\")\nprint(\"Saved outputs to:\", OUTPUT_DIR)\n\nprint(\"\\nImportant files:\")\nprint(os.path.join(OUTPUT_DIR, \"best_maxvit_final_checkpoint.pth\"))\nprint(os.path.join(OUTPUT_DIR, \"best_maxvit_final_model.pth\"))\nprint(os.path.join(OUTPUT_DIR, \"best_maxvit_final_per_class_metrics.csv\"))\nprint(os.path.join(OUTPUT_DIR, \"best_maxvit_final_val_predictions.npz\"))\nprint(os.path.join(OUTPUT_DIR, \"maxvit_final_training_history.csv\"))\nprint(os.path.join(OUTPUT_DIR, \"maxvit_final_training_summary.json\"))\nprint(os.path.join(OUTPUT_DIR, \"maxvit_class_imbalance_report.csv\"))\nprint(os.path.join(OUTPUT_DIR, \"maxvit_pos_weight.csv\"))\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.png\"))\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_loss_auc.pdf\"))\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_per_class_auc.png\"))\nprint(os.path.join(FIGURE_DIR, \"maxvit_final_per_class_auc.pdf\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}