{"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},{"sourceType":"kernelVersion","sourceId":319883136}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, ast, time, json, random, copy, warnings\nimport numpy as np\nimport pandas as pd\nimport torch\nimport timm\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast\nfrom tqdm import tqdm\nfrom sklearn.metrics import roc_auc_score, average_precision_score\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nimport cv2\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CACHE_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    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/swin_transformer_final_outputs\"\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_DIR:\", TRAIN_DIR)\nprint(\"VAL_DIR:\", VAL_DIR)\nprint(\"OUTPUT_DIR:\", OUTPUT_DIR)","metadata":{"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 = \"swin_tiny_patch4_window7_224\"\nIMG_SIZE = 384\n\nBATCH_SIZE = 8\nACCUMULATION_STEPS = 4\nEPOCHS = 40\nPATIENCE = 8\n\nLR_HEAD = 8e-5\nLR_BACKBONE = 5e-6\nWEIGHT_DECAY = 1e-4\nNUM_WORKERS = 4\n\nDROP_RATE = 0.1\nDROP_PATH_RATE = 0.1\nMAX_GRAD_NORM = 1.0\nMAX_POS_WEIGHT = 5.0\nEMA_DECAY = 0.995\nVAL_LOSS_DELTA = 1e-4\nUSE_VAL_LOSS_PATIENCE_RELIEF = True\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]","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, \"swin_class_imbalance_report.csv\"), index=False)\ndisplay(imbalance_df)\n\npos_weight = neg_counts / (pos_counts + 1e-6)\npos_weight = np.sqrt(pos_weight)\npos_weight = np.clip(pos_weight, 1.0, MAX_POS_WEIGHT)\npos_weight[NO_FINDING_ID] = 0.75\n\npos_weight_tensor = torch.tensor(pos_weight, dtype=torch.float32).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\n})\n\npos_weight_df.to_csv(os.path.join(OUTPUT_DIR, \"swin_pos_weight.csv\"), index=False)\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.025,\n                    scale_limit=0.035,\n                    rotate_limit=5,\n                    border_mode=cv2.BORDER_CONSTANT,\n                    value=0,\n                    p=0.25\n                ),\n                A.RandomBrightnessContrast(\n                    brightness_limit=0.05,\n                    contrast_limit=0.05,\n                    p=0.25\n                ),\n                A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n                ToTensorV2()\n            ])\n        else:\n            self.transform = A.Compose([\n                A.Resize(img_size, img_size),\n                A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\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.1:\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] == 1:\n            img = np.repeat(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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_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":"model = timm.create_model(\n    MODEL_NAME,\n    pretrained=True,\n    num_classes=NUM_CLASSES,\n    img_size=IMG_SIZE,\n    drop_rate=DROP_RATE,\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\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)\nmonitor_criterion = nn.BCEWithLogitsLoss()\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 = [p for p in model.parameters() if id(p) not in head_param_ids]\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 = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=EPOCHS,\n    eta_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\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@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\ndef compute_metrics(y_true, y_pred):\n    per_class = []\n    aucs = []\n    abnormal_aucs = []\n    aps = []\n    abnormal_aps = []\n\n    for c, class_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\": class_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            aucs.append(auc)\n\n            if c != NO_FINDING_ID:\n                abnormal_aucs.append(auc)\n\n        if not np.isnan(ap):\n            aps.append(ap)\n\n            if c != NO_FINDING_ID:\n                abnormal_aps.append(ap)\n\n    try:\n        micro_auc = roc_auc_score(y_true.ravel(), y_pred.ravel())\n    except:\n        micro_auc = np.nan\n\n    metrics = {\n        \"macro_auc\": float(np.nanmean(aucs)),\n        \"abnormal_macro_auc\": float(np.nanmean(abnormal_aucs)),\n        \"micro_auc\": float(micro_auc),\n        \"macro_ap\": float(np.nanmean(aps)),\n        \"abnormal_macro_ap\": float(np.nanmean(abnormal_aps))\n    }\n\n    return metrics, pd.DataFrame(per_class)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate_model(eval_model, loader):\n    eval_model.eval()\n\n    val_weighted_loss = 0.0\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            weighted_loss = criterion(logits, labels)\n            bce_loss = monitor_criterion(logits, labels)\n\n        probs = torch.sigmoid(logits).detach().cpu().numpy()\n\n        val_weighted_loss += weighted_loss.item()\n        val_bce_loss += bce_loss.item()\n\n        all_preds.append(probs)\n        all_labels.append(labels.detach().cpu().numpy())\n\n    val_weighted_loss = val_weighted_loss / len(loader)\n    val_bce_loss = val_bce_loss / len(loader)\n\n    y_pred = np.vstack(all_preds)\n    y_true = np.vstack(all_labels)\n\n    metrics, per_class_df = compute_metrics(y_true, y_pred)\n\n    return val_weighted_loss, val_bce_loss, metrics, per_class_df, y_pred, y_true","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_abnormal_macro_auc = -1.0\nbest_val_bce_loss = float(\"inf\")\nepochs_without_improvement = 0\n\nhistory = {\n    \"epoch\": [],\n    \"train_loss\": [],\n    \"train_bce_loss\": [],\n    \"val_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            logits = model(imgs)\n            train_loss_value = criterion(logits, labels)\n            train_bce_loss_value = monitor_criterion(logits, 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_GRAD_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n            update_ema(model, ema_model, EMA_DECAY)\n\n        running_loss += train_loss_value.item()\n        running_bce_loss += train_bce_loss_value.item()\n\n        pbar.set_postfix(\n            loss=f\"{train_loss_value.item():.4f}\",\n            bce=f\"{train_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_GRAD_NORM)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        update_ema(model, ema_model, EMA_DECAY)\n\n    train_loss = running_loss / len(train_loader)\n    train_bce_loss = running_bce_loss / len(train_loader)\n\n    val_loss, val_bce_loss, metrics, per_class_metrics_df, all_preds, all_labels = validate_model(\n        ema_model,\n        val_loader\n    )\n\n    scheduler.step()\n\n    lr_backbone = optimizer.param_groups[0][\"lr\"]\n    lr_head = optimizer.param_groups[1][\"lr\"] if len(optimizer.param_groups) > 1 else lr_backbone\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    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_loss\"].append(val_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(os.path.join(OUTPUT_DIR, \"swin_transformer_training_history.csv\"), index=False)\n\n    print(f\"\\nEpoch {epoch + 1}/{EPOCHS}\")\n    print(f\"Train Weighted Loss: {train_loss:.4f}\")\n    print(f\"Train BCE Loss:      {train_bce_loss:.4f}\")\n    print(f\"Val Weighted Loss:   {val_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_bce_improved = val_bce_loss < best_val_bce_loss - VAL_LOSS_DELTA\n\n    if val_bce_improved:\n        best_val_bce_loss = 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            \"scheduler_state_dict\": scheduler.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_loss\": val_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            \"seed\": SEED,\n            \"ema_decay\": EMA_DECAY,\n            \"lr_head\": LR_HEAD,\n            \"lr_backbone\": LR_BACKBONE,\n            \"weight_decay\": WEIGHT_DECAY,\n            \"pos_weight\": pos_weight.tolist()\n        }\n\n        torch.save(checkpoint, os.path.join(OUTPUT_DIR, \"best_swin_transformer_checkpoint.pth\"))\n        torch.save(ema_model.state_dict(), os.path.join(OUTPUT_DIR, \"best_swin_transformer_model.pth\"))\n\n        per_class_metrics_df.to_csv(\n            os.path.join(OUTPUT_DIR, \"best_swin_transformer_per_class_metrics.csv\"),\n            index=False\n        )\n\n        np.savez_compressed(\n            os.path.join(OUTPUT_DIR, \"best_swin_transformer_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_bce_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}/{PATIENCE}\")\n        else:\n            epochs_without_improvement += 1\n            print(f\"No improvement for {epochs_without_improvement}/{PATIENCE} epochs\")\n\n    if epochs_without_improvement >= 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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_df = pd.DataFrame(history)\nhistory_df.to_csv(os.path.join(OUTPUT_DIR, \"swin_transformer_training_history.csv\"), index=False)\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\": \"Swin Transformer Tiny\",\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_patience\": PATIENCE,\n    \"ema_decay\": EMA_DECAY,\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_rate\": DROP_RATE,\n    \"drop_path_rate\": DROP_PATH_RATE,\n    \"total_time_seconds\": total_time,\n    \"total_time_hours\": total_time / 3600,\n    \"output_dir\": OUTPUT_DIR\n}\n\nwith open(os.path.join(OUTPUT_DIR, \"swin_transformer_training_summary.json\"), \"w\") as f:\n    json.dump(summary, f, indent=4)\n\ndisplay(history_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(\"Swin Transformer 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(\"Swin Transformer Validation AUC Progression\")\naxes[1].set_xlabel(\"Epoch\")\naxes[1].set_ylabel(\"AUC\")\naxes[1].set_xticks(xticks)\naxes[1].set_ylim(\n    max(0.0, min(history_df[\"abnormal_macro_auc\"].min(), history_df[\"macro_auc\"].min()) - 0.05),\n    min(1.0, max(history_df[\"abnormal_macro_auc\"].max(), history_df[\"macro_auc\"].max()) + 0.03)\n)\naxes[1].grid(True, linestyle=\"--\", alpha=0.4)\naxes[1].legend()\n\nplt.tight_layout()\nplt.savefig(os.path.join(FIGURE_DIR, \"swin_transformer_loss_auc.png\"), bbox_inches=\"tight\", dpi=200)\nplt.savefig(os.path.join(FIGURE_DIR, \"swin_transformer_loss_auc.pdf\"), bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_metrics_path = os.path.join(OUTPUT_DIR, \"best_swin_transformer_per_class_metrics.csv\")\nbest_per_class_df = pd.read_csv(best_metrics_path)\n\nplt.figure(figsize=(10, 6))\n\nplot_df = best_per_class_df.sort_values(\"auc\", ascending=True)\n\nplt.barh(plot_df[\"class_name\"], plot_df[\"auc\"])\n\nplt.axvline(\n    plot_df[\"auc\"].mean(),\n    linestyle=\"--\",\n    linewidth=1.5,\n    label=f\"Mean AUC = {plot_df['auc'].mean():.4f}\"\n)\n\nplt.title(\"Per-Class AUC - Swin Transformer\")\nplt.xlabel(\"AUC\")\nplt.ylabel(\"Class\")\nplt.xlim(0.5, 1.0)\nplt.grid(axis=\"x\", linestyle=\"--\", alpha=0.4)\nplt.legend()\nplt.tight_layout()\nplt.savefig(os.path.join(FIGURE_DIR, \"swin_transformer_per_class_auc.png\"), bbox_inches=\"tight\", dpi=200)\nplt.savefig(os.path.join(FIGURE_DIR, \"swin_transformer_per_class_auc.pdf\"), bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_metrics = pd.DataFrame([\n    [\"Model\", \"Swin Transformer Tiny\"],\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\", PATIENCE],\n    [\"EMA\", f\"Yes, decay={EMA_DECAY}\"],\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(final_metrics)\n\ndisplay(best_per_class_df[[\n    \"class_id\",\n    \"class_name\",\n    \"auc\",\n    \"ap\",\n    \"positive_samples\",\n    \"positive_ratio\"\n]])\n\nprint(\"\\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\"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_swin_transformer_checkpoint.pth\"))\nprint(os.path.join(OUTPUT_DIR, \"best_swin_transformer_model.pth\"))\nprint(os.path.join(OUTPUT_DIR, \"best_swin_transformer_per_class_metrics.csv\"))\nprint(os.path.join(OUTPUT_DIR, \"swin_transformer_training_history.csv\"))\nprint(os.path.join(OUTPUT_DIR, \"swin_transformer_training_summary.json\"))\nprint(os.path.join(FIGURE_DIR, \"swin_transformer_loss_auc.png\"))\nprint(os.path.join(FIGURE_DIR, \"swin_transformer_per_class_auc.png\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}