{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaL4","dataSources":[{"sourceId":91496,"databundleVersionId":11802066,"sourceType":"competition"},{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"},{"sourceId":12875356,"sourceType":"datasetVersion","datasetId":8144903}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile /kaggle/working/train.py\nimport os\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.distributed as dist\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.distributed import DistributedSampler\nfrom torchvision import transforms\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import roc_auc_score\nfrom torchvision.models import densenet121, DenseNet121_Weights\n\n# ======================\n# Dataset\n# ======================\nclass ChestXRayDataset(Dataset):\n    def __init__(self, df, img_size=(1048, 1048), is_test=False, transforms=None):\n        self.df = df\n        self.img_size = img_size\n        self.is_test = is_test\n        self.label_columns = [\n            'Atelectasis', 'Cardiomegaly', 'Consolidation', 'Edema', 'Enlarged Cardiomediastinum',\n            'Fracture', 'Lung Lesion', 'Lung Opacity', 'No Finding', 'Pleural Effusion',\n            'Pleural Other', 'Pneumonia', 'Pneumothorax', 'Support Devices'\n        ]\n        self.image_dir = '/kaggle/input/grand-xray-slam-division-a/train1/' if not is_test else '/kaggle/input/grand-xray-slam-division-a/test1/'\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.image_dir, row['Image_name'])\n        \n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if img is None:\n            img = np.zeros((self.img_size[0], self.img_size[1], 3), dtype=np.uint8)\n\n        img = cv2.resize(img, self.img_size)\n\n        if self.transforms:\n            img = self.transforms(img)\n\n        if not self.is_test:\n            labels = row[self.label_columns].values.astype(np.float32)\n            return img, torch.tensor(labels)\n        return img\n\n# ======================\n# Model\n# ======================\ndef create_model(num_classes=14):\n    model = densenet121(weights=DenseNet121_Weights.IMAGENET1K_V1)\n    for param in model.parameters():\n        param.requires_grad = False\n    num_ftrs = model.classifier.in_features\n    model.classifier = nn.Linear(num_ftrs, num_classes)\n    for param in model.features.denseblock4.parameters():\n        param.requires_grad = True\n    return model\n\n\n# ======================\n# Main DDP Training\n# ======================\ndef main():\n    dist.init_process_group(backend=\"nccl\")  # NCCL = best for multi-GPU\n    local_rank = int(os.environ[\"LOCAL_RANK\"])\n    torch.cuda.set_device(local_rank)\n    device = torch.device(\"cuda\", local_rank)\n\n    # --- Load data ---\n    train_df = pd.read_csv('/kaggle/input/grand-xray-slam-division-a/train1.csv')\n    from sklearn.model_selection import train_test_split\n    train_data, val_data = train_test_split(\n        train_df, test_size=0.2, random_state=42, stratify=train_df['No Finding']\n    )\n\n    img_transforms = transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ])\n\n    train_dataset = ChestXRayDataset(train_data, img_size=(1048, 1048), transforms=img_transforms)\n    val_dataset = ChestXRayDataset(val_data, img_size=(1048, 1048), transforms=img_transforms)\n\n    train_sampler = DistributedSampler(train_dataset)\n    val_sampler = DistributedSampler(val_dataset, shuffle=False)\n\n    train_loader = DataLoader(train_dataset, batch_size=32, sampler=train_sampler, num_workers=4)\n    val_loader = DataLoader(val_dataset, batch_size=32, sampler=val_sampler, num_workers=4)\n\n    # --- Model, loss, optimizer ---\n    model = create_model().to(device)\n    model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank)\n\n    criterion = nn.BCEWithLogitsLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=0.0005)\n\n    num_epochs = 10\n\n    for epoch in range(num_epochs):\n        train_sampler.set_epoch(epoch)\n\n        # ---- Train ----\n        model.train()\n        train_preds, train_labels = [], []\n        running_train_loss = 0.0\n\n        for images, labels in tqdm(train_loader, disable=(dist.get_rank() != 0)):\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            running_train_loss += loss.item() * images.size(0)\n            train_preds.append(torch.sigmoid(outputs).detach().cpu().numpy())\n            train_labels.append(labels.detach().cpu().numpy())\n\n        # ---- Validation ----\n        model.eval()\n        val_preds, val_labels = [], []\n        running_val_loss = 0.0\n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                running_val_loss += loss.item() * images.size(0)\n                val_preds.append(torch.sigmoid(outputs).cpu().numpy())\n                val_labels.append(labels.cpu().numpy())\n\n        if dist.get_rank() == 0:  # Only rank 0 prints\n            train_preds = np.vstack(train_preds)\n            train_labels = np.vstack(train_labels)\n            val_preds = np.vstack(val_preds)\n            val_labels = np.vstack(val_labels)\n\n            from sklearn.metrics import roc_auc_score\n            train_auc = roc_auc_score(train_labels, train_preds, average='macro')\n            val_auc = roc_auc_score(val_labels, val_preds, average='macro')\n\n            print(f\"Epoch {epoch+1}/{num_epochs} | \"\n                  f\"Train Loss {running_train_loss/len(train_dataset):.4f} | Train AUC {train_auc:.4f} | \"\n                  f\"Val Loss {running_val_loss/len(val_dataset):.4f} | Val AUC {val_auc:.4f}\")\n\n    if dist.get_rank() == 0:\n        torch.save(model.module.state_dict(), \"best_model.pth\")\n\n    dist.destroy_process_group()\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !torchrun --nproc_per_node=4 /kaggle/working/train.py","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/inference.py\nimport os\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom tqdm.auto import tqdm\nfrom torchvision.models import densenet121, DenseNet121_Weights\n\n# ======================\n# Dataset\n# ======================\nclass ChestXRayTestDataset(Dataset):\n    def __init__(self, df, img_size=(512, 512), transforms=None):\n        self.df = df\n        self.img_size = img_size\n        self.image_dir = '/kaggle/input/grand-xray-slam-division-a/test1/'\n        self.transforms = transforms\n        self.label_columns = [\n            'Atelectasis', 'Cardiomegaly', 'Consolidation', 'Edema', 'Enlarged Cardiomediastinum',\n            'Fracture', 'Lung Lesion', 'Lung Opacity', 'No Finding', 'Pleural Effusion',\n            'Pleural Other', 'Pneumonia', 'Pneumothorax', 'Support Devices'\n        ]\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.image_dir, row['Image_name'])\n\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if img is None:\n            img = np.zeros((self.img_size[0], self.img_size[1], 3), dtype=np.uint8)\n        \n        img = cv2.resize(img, self.img_size)\n        img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n\n        if self.transforms:\n            img = self.transforms(img)\n\n        return img, row['Image_name']\n\n# ======================\n# Model\n# ======================\ndef create_model(num_classes=14):\n    model = densenet121(weights=DenseNet121_Weights.IMAGENET1K_V1)\n    num_ftrs = model.classifier.in_features\n    model.classifier = nn.Linear(num_ftrs, num_classes)\n    return model\n\n# ======================\n# Inference\n# ======================\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    # Load test data\n    test_files = os.listdir('/kaggle/input/grand-xray-slam-division-a/test1/')\n    test_df = pd.DataFrame({\"Image_name\": test_files})\n\n    img_transforms = transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ])\n\n    test_dataset = ChestXRayTestDataset(test_df, img_size=(512, 512), transforms=img_transforms)\n    test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=2)\n\n    # Load model\n    model = create_model()\n    model.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\", map_location=device))\n    model.to(device)\n    model.eval()\n\n    all_probs = []\n    all_names = []\n\n    with torch.no_grad():\n        for images, names in tqdm(test_loader, desc=\"Inference\", leave=True):\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.sigmoid(outputs).cpu().numpy()\n            all_probs.append(probs)\n            all_names.extend(names)\n\n    all_probs = np.vstack(all_probs)\n    submission = pd.DataFrame(all_probs, columns=test_dataset.label_columns)\n    submission.insert(0, \"Image_name\", all_names)\n\n    submission.to_csv(\"/kaggle/working/submission.csv\", index=False)\n    print(\"✅ Saved submission.csv\")\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python /kaggle/working/inference.py","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/train.py\n\"\"\"\nImproved train.py (image-only) for Grand X-Ray Slam\n- fixes warnings from logs (Albumentations, AMP, DDP grad bucket view)\n- uses modern torch.amp API, gradient_as_bucket_view for DDP\n- uses Affine (Albumentations) instead of ShiftScaleRotate\n- optimizer.zero_grad(set_to_none=True) for better perf\n- CosineAnnealingLR scheduler (per-epoch)\n\"\"\"\n\nimport os\nos.environ.setdefault(\"OMP_NUM_THREADS\", \"1\")  # recommended for DDP per-process thread control\n\nimport gc\nimport random\nfrom argparse import ArgumentParser\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.distributed as dist\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.distributed import DistributedSampler\n\nimport timm\nfrom sklearn.model_selection import GroupKFold, train_test_split\nfrom sklearn.metrics import roc_auc_score\n\n# ---------- helpers ----------\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\ndef free_memory():\n    gc.collect()\n    try:\n        torch.cuda.empty_cache()\n        torch.cuda.ipc_collect()\n    except Exception:\n        pass\n\n# ---------- labels ----------\nLABEL_COLS = [\n    'Atelectasis','Cardiomegaly','Consolidation','Edema','Enlarged Cardiomediastinum',\n    'Fracture','Lung Lesion','Lung Opacity','No Finding','Pleural Effusion',\n    'Pleural Other','Pneumonia','Pneumothorax','Support Devices'\n]\n\n# ---------- dataset ----------\nclass ChestXRayImageDataset(Dataset):\n    def __init__(self, df, image_dir, img_size=1024, transforms=None, image_col='Image_Name', label_cols=LABEL_COLS):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = Path(image_dir)\n        self.transforms = transforms\n        self.img_size = img_size\n        self.image_col = image_col\n        self.label_cols = label_cols\n\n    def __len__(self):\n        return len(self.df)\n\n    def _read_img(self, fname):\n        p = self.image_dir / fname\n        img = cv2.imread(str(p), cv2.IMREAD_UNCHANGED)\n        if img is None:\n            return np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n        if img.ndim == 2:\n            img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        elif img.shape[2] == 4:\n            img = cv2.cvtColor(img, cv2.COLOR_BGRA2RGB)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        return img\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        fname = row[self.image_col]\n        image = self._read_img(fname)\n        if self.transforms is not None:\n            image = self.transforms(image=image)['image']\n        labels = row[self.label_cols].values.astype(np.float32)\n        return image, torch.tensor(labels, dtype=torch.float32)\n\n# ---------- loss ----------\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, logits, targets):\n        bce = nn.functional.binary_cross_entropy_with_logits(logits, targets, reduction='none')\n        p_t = torch.exp(-bce)\n        focal = (1 - p_t) ** self.gamma * bce\n        if self.reduction == 'mean':\n            return focal.mean()\n        elif self.reduction == 'sum':\n            return focal.sum()\n        return focal\n\ndef hybrid_loss(logits, targets, pos_weight_tensor, alpha_bce=0.6, alpha_focal=0.4, focal_gamma=2.0):\n    bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)(logits, targets)\n    focal = FocalLoss(gamma=focal_gamma)(logits, targets)\n    return alpha_bce * bce + alpha_focal * focal\n\n# ---------- train/val ----------\ndef train_one_epoch(model, loader, optimizer, scaler, device, pos_weight_tensor, accumulation_steps=1):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    optimizer.zero_grad(set_to_none=True)\n    pbar = tqdm(loader, desc='Train', leave=False)\n    for step, (images, targets) in enumerate(pbar):\n        images = images.to(device, non_blocking=True)\n        targets = targets.to(device, non_blocking=True)\n\n        with torch.amp.autocast(device_type=device.type):\n            logits = model(images)\n            loss = hybrid_loss(logits, targets, pos_weight_tensor)\n\n        scaler.scale(loss / accumulation_steps).backward()\n\n        if (step + 1) % accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n\n        running_loss += float(loss.item()) * images.size(0)\n        all_preds.append(torch.sigmoid(logits).detach().cpu().numpy())\n        all_targets.append(targets.detach().cpu().numpy())\n\n        pbar.set_postfix({'loss': running_loss / ((step + 1) * loader.batch_size)})\n\n    epoch_loss = running_loss / len(loader.dataset)\n    preds = np.vstack(all_preds)\n    targets = np.vstack(all_targets)\n    return epoch_loss, preds, targets\n\ndef valid_one_epoch(model, loader, device):\n    model.eval()\n    all_preds, all_targets = [], []\n    with torch.no_grad():\n        pbar = tqdm(loader, desc='Valid', leave=False)\n        for images, targets in pbar:\n            images = images.to(device, non_blocking=True)\n            targets = targets.to(device, non_blocking=True)\n            logits = model(images)\n            all_preds.append(torch.sigmoid(logits).cpu().numpy())\n            all_targets.append(targets.cpu().numpy())\n    preds = np.vstack(all_preds)\n    targets = np.vstack(all_targets)\n    return preds, targets\n\ndef compute_per_class_auc(truths, preds, label_cols):\n    per_class = []\n    for i, name in enumerate(label_cols):\n        try:\n            a = roc_auc_score(truths[:, i], preds[:, i])\n        except Exception:\n            a = float('nan')\n        per_class.append(a)\n    macro = float(np.nanmean(per_class))\n    return macro, per_class\n\n# ---------- main ----------\ndef main():\n    parser = ArgumentParser()\n    parser.add_argument('--data-csv', type=str, default='/kaggle/input/grand-xray-slam-division-a/train1.csv')\n    parser.add_argument('--image-dir', type=str, default='/kaggle/input/grand-xray-slam-division-a/train1/')\n    parser.add_argument('--out-dir', type=str, default='/kaggle/working/')\n    parser.add_argument('--img-size', type=int, default=1024)\n    parser.add_argument('--batch-size', type=int, default=8)\n    parser.add_argument('--epochs', type=int, default=6)\n    parser.add_argument('--workers', type=int, default=4)\n    parser.add_argument('--lr', type=float, default=1e-4)\n    parser.add_argument('--accumulation', type=int, default=1)\n    parser.add_argument('--fold', type=int, default=0)\n    parser.add_argument('--n-folds', type=int, default=5)\n    parser.add_argument('--use-ddp', action='store_true')\n    parser.add_argument('--backbone', type=str, default='convnext_large')\n    args = parser.parse_args()\n\n    seed_everything(42)\n    os.makedirs(args.out_dir, exist_ok=True)\n\n    # performance flags\n    torch.backends.cudnn.benchmark = True\n\n    # DDP init\n    use_ddp = args.use_ddp\n    if use_ddp:\n        dist.init_process_group(backend='nccl')\n        local_rank = int(os.environ.get('LOCAL_RANK', '0'))\n        torch.cuda.set_device(local_rank)\n        device = torch.device(f'cuda:{local_rank}')\n    else:\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        local_rank = 0\n\n    # read csv\n    df = pd.read_csv(args.data_csv)\n    # detect image column\n    image_col_candidates = ['Image_Name','Image_name','Image']\n    image_col = None\n    for c in image_col_candidates:\n        if c in df.columns:\n            image_col = c\n            break\n    if image_col is None:\n        for c in df.columns:\n            if c.lower() == 'image_name' or c.lower() == 'image':\n                image_col = c\n                break\n    if image_col is None:\n        raise ValueError(\"Couldn't find image column in CSV. Expected 'Image_Name' or 'Image_name' or 'Image'.\")\n\n    # ensure labels exist (case-insensitive)\n    label_cols = []\n    for lc in LABEL_COLS:\n        if lc in df.columns:\n            label_cols.append(lc)\n        else:\n            matched = [c for c in df.columns if c.lower() == lc.lower()]\n            if matched:\n                label_cols.append(matched[0])\n            else:\n                raise ValueError(f\"Label column {lc} not found in CSV. Available: {df.columns.tolist()}\")\n\n    # rename image col\n    df = df.rename(columns={image_col: 'Image_Name'})\n\n    # compute pos_weight per class\n    total = len(df)\n    label_sums = df[label_cols].sum(axis=0).values.astype(np.float32)\n    negs = total - label_sums\n    pos_weight = (negs / (label_sums + 1e-6)).astype(np.float32)\n    pos_weight_tensor = torch.from_numpy(pos_weight).to(device)\n\n    # split: grouped by patient if available\n    if 'Patient_ID' in df.columns:\n        groups = df['Patient_ID'].values\n        gkf = GroupKFold(n_splits=args.n_folds)\n        splits = list(gkf.split(df, df[label_cols], groups))\n        train_idx, val_idx = splits[args.fold]\n    else:\n        stratify_col = df[label_cols[0]] if label_cols[0] in df.columns else None\n        train_idx, val_idx = train_test_split(np.arange(len(df)), test_size=0.2, random_state=42, stratify=stratify_col)\n\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    val_df = df.iloc[val_idx].reset_index(drop=True)\n\n    # transforms (Affine used instead of ShiftScaleRotate)\n    train_transforms = A.Compose([\n        A.Resize(args.img_size, args.img_size),\n        A.HorizontalFlip(p=0.5),\n        A.Affine(translate_percent=0.05, scale=(0.92, 1.08), rotate=(-10, 10), shear=0, p=0.5),\n        A.RandomBrightnessContrast(p=0.5),\n        A.CLAHE(p=0.3),\n        A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n        ToTensorV2(),\n    ])\n    val_transforms = A.Compose([\n        A.Resize(args.img_size, args.img_size),\n        A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n        ToTensorV2(),\n    ])\n\n    # datasets & loaders\n    train_ds = ChestXRayImageDataset(train_df, args.image_dir, img_size=args.img_size, transforms=train_transforms, image_col='Image_Name', label_cols=label_cols)\n    val_ds = ChestXRayImageDataset(val_df, args.image_dir, img_size=args.img_size, transforms=val_transforms, image_col='Image_Name', label_cols=label_cols)\n\n    if use_ddp:\n        train_sampler = DistributedSampler(train_ds)\n        val_sampler = DistributedSampler(val_ds, shuffle=False)\n    else:\n        train_sampler = None\n        val_sampler = None\n\n    train_loader = DataLoader(train_ds, batch_size=args.batch_size, sampler=train_sampler,\n                              shuffle=(train_sampler is None), num_workers=args.workers, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=args.batch_size, sampler=val_sampler,\n                            shuffle=False, num_workers=args.workers, pin_memory=True)\n\n    # model init + DDP with gradient_as_bucket_view to avoid grad/bucket warnings\n    model = timm.create_model(args.backbone, pretrained=True, num_classes=len(label_cols))\n    model.to(device)\n    if use_ddp:\n        model = torch.nn.parallel.DistributedDataParallel(\n            model,\n            device_ids=[local_rank],\n            output_device=local_rank,\n            find_unused_parameters=False,\n            gradient_as_bucket_view=True\n        )\n\n    optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max(1, args.epochs))\n    # scaler: use new API, specify device type when CUDA\n    try:\n        if device.type == 'cuda':\n            scaler = torch.amp.GradScaler(device_type='cuda')\n        else:\n            scaler = torch.amp.GradScaler()\n    except TypeError:\n        # fallback for older torch versions\n        scaler = torch.amp.GradScaler()\n\n    best_auc = 0.0\n\n    for epoch in range(args.epochs):\n        if use_ddp:\n            train_loader.sampler.set_epoch(epoch)\n\n        train_loss, train_preds, train_targets = train_one_epoch(\n            model, train_loader, optimizer, scaler, device, pos_weight_tensor, accumulation_steps=args.accumulation\n        )\n\n        val_preds, val_targets = valid_one_epoch(model, val_loader, device)\n\n        macro_auc, per_class = compute_per_class_auc(val_targets, val_preds, label_cols)\n\n        if local_rank == 0:\n            print(f\"Epoch {epoch+1}/{args.epochs} | TrainLoss: {train_loss:.4f} | Val AUC (macro): {macro_auc:.4f}\")\n            for ln, a in zip(label_cols, per_class):\n                print(f\"  {ln}: {a:.4f}\")\n            if macro_auc > best_auc:\n                best_auc = macro_auc\n                ckpt_path = os.path.join(args.out_dir, f\"best_{args.backbone}_fold{args.fold}_img{args.img_size}.pth\")\n                state = model.module.state_dict() if hasattr(model, 'module') else model.state_dict()\n                torch.save({'state_dict': state, 'auc': best_auc}, ckpt_path)\n                print(\"Saved checkpoint:\", ckpt_path)\n\n        # scheduler step per epoch\n        try:\n            scheduler.step()\n        except Exception:\n            pass\n\n        free_memory()\n\n    if use_ddp:\n        dist.destroy_process_group()\n\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # debug / quick: smaller image, fewer epochs\n# python /kaggle/working/train.py --data-csv /kaggle/input/grand-xray-slam-division-a/train1.csv \\\n#   --image-dir /kaggle/input/grand-xray-slam-division-a/train1/ \\\n#   --img-size 512 --batch-size 16 --epochs 2 --workers 4 --backbone convnext_small\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!torchrun --nproc_per_node=4 /kaggle/working/train.py --use-ddp --img-size 512 --batch-size 16 --epochs 10 --workers 8 --backbone convnext_large","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/inference.py\n\"\"\"\nInference (image-only) with optional TTA and checkpoint ensembling.\nUsage examples:\n  python inference.py --image-dir /kaggle/input/grand-xray-slam-division-a/test1/ --ckpt \"/kaggle/working/*best*.pth\" --img-size 1024 --batch-size 8 --tta-scales \"1024,800\" --tta-flip\n\"\"\"\n\nimport os\nimport glob\nimport argparse\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\n\nLABEL_COLS = [\n    'Atelectasis','Cardiomegaly','Consolidation','Edema','Enlarged Cardiomediastinum',\n    'Fracture','Lung Lesion','Lung Opacity','No Finding','Pleural Effusion',\n    'Pleural Other','Pneumonia','Pneumothorax','Support Devices'\n]\n\nclass TestImageDataset(Dataset):\n    def __init__(self, image_dir, files, transforms=None, img_size=1024):\n        self.image_dir = Path(image_dir)\n        self.files = files\n        self.transforms = transforms\n        self.img_size = img_size\n\n    def __len__(self):\n        return len(self.files)\n\n    def _read(self, fname):\n        p = self.image_dir / fname\n        img = cv2.imread(str(p), cv2.IMREAD_UNCHANGED)\n        if img is None:\n            img = np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n            return img\n        if img.ndim == 2:\n            img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        elif img.shape[2] == 4:\n            img = cv2.cvtColor(img, cv2.COLOR_BGRA2RGB)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        return img\n\n    def __getitem__(self, idx):\n        fname = self.files[idx]\n        img = self._read(fname)\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n        return img, fname\n\ndef load_state_stripped(path, map_location='cpu'):\n    ck = torch.load(path, map_location=map_location)\n    if isinstance(ck, dict) and 'state_dict' in ck:\n        state = ck['state_dict']\n    elif isinstance(ck, dict) and 'model' in ck:\n        state = ck['model']\n    elif isinstance(ck, dict):\n        state = ck\n    else:\n        state = ck\n    new_state = {}\n    for k, v in state.items():\n        new_key = k.replace('module.', '') if k.startswith('module.') else k\n        new_state[new_key] = v\n    return new_state\n\n@torch.no_grad()\ndef run_inference(model, loader, device):\n    model.eval()\n    sigmoid = nn.Sigmoid()\n    probs_list = []\n    names = []\n    for imgs, fnames in tqdm(loader, desc=\"Inf\", leave=False):\n        imgs = imgs.to(device, non_blocking=True)\n        logits = model(imgs)\n        probs = sigmoid(logits).cpu().numpy()\n        probs_list.append(probs)\n        names.extend(fnames)\n    probs = np.vstack(probs_list)\n    return probs, names\n\ndef build_transform(img_size, hflip=False):\n    tr = [A.Resize(img_size, img_size)]\n    if hflip:\n        tr.append(A.HorizontalFlip(p=1.0))\n    tr += [\n        A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n        ToTensorV2()\n    ]\n    return A.Compose(tr)\n\ndef main():\n    parser = argparse.ArgumentParser()\n    parser.add_argument('--image-dir', type=str, required=True)\n    parser.add_argument('--ckpt', type=str, default=None, help=\"glob or single path\")\n    parser.add_argument('--backbone', type=str, default='convnext_large')\n    parser.add_argument('--img-size', type=int, default=1024)\n    parser.add_argument('--batch-size', type=int, default=8)\n    parser.add_argument('--tta-scales', type=str, default=None, help=\"e.g. '1024,800'\")\n    parser.add_argument('--tta-flip', action='store_true')\n    parser.add_argument('--device', type=str, default='cuda')\n    parser.add_argument('--num-workers', type=int, default=4)\n    parser.add_argument('--out', type=str, default='/kaggle/working/submission.csv')\n    args = parser.parse_args()\n\n    device = torch.device(args.device if torch.cuda.is_available() else 'cpu')\n\n    # gather checkpoints\n    if args.ckpt:\n        ckpts = sorted(glob.glob(args.ckpt))\n    else:\n        ckpts = sorted(glob.glob('/kaggle/working/*best*.pth')) + sorted(glob.glob('/kaggle/working/*.pth'))\n    if len(ckpts) == 0:\n        raise FileNotFoundError(\"No checkpoints found. Provide --ckpt pattern or place pth files in /kaggle/working/\")\n\n    files = sorted([f for f in os.listdir(args.image_dir) if f.lower().endswith(('.jpg','.jpeg','.png'))])\n    if len(files) == 0:\n        raise FileNotFoundError(f\"No images found in {args.image_dir}\")\n\n    if args.tta_scales:\n        scales = [int(x.strip()) for x in args.tta_scales.split(',') if x.strip()]\n    else:\n        scales = [args.img_size]\n\n    ensemble_preds = None\n    for ckpt in ckpts:\n        print(\"Loading checkpoint:\", ckpt)\n        # load model\n        model = timm.create_model(args.backbone, pretrained=False, num_classes=len(LABEL_COLS))\n        sd = load_state_stripped(ckpt, map_location='cpu')\n        model.load_state_dict(sd, strict=False)\n        model.to(device)\n        model.eval()\n\n        model_accum = None\n        tta_count = 0\n\n        for scale in scales:\n            # normal\n            tr = build_transform(scale, hflip=False)\n            ds = TestImageDataset(args.image_dir, files, transforms=tr, img_size=scale)\n            loader = DataLoader(ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, pin_memory=True)\n            probs, names = run_inference(model, loader, device)\n            if model_accum is None:\n                model_accum = probs.copy()\n            else:\n                model_accum += probs\n            tta_count += 1\n\n            # flip\n            if args.tta_flip:\n                trf = build_transform(scale, hflip=True)\n                dsf = TestImageDataset(args.image_dir, files, transforms=trf, img_size=scale)\n                loaderf = DataLoader(dsf, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, pin_memory=True)\n                probs_f, _ = run_inference(model, loaderf, device)\n                model_accum += probs_f\n                tta_count += 1\n\n        model_accum = model_accum / float(tta_count)\n        if ensemble_preds is None:\n            ensemble_preds = model_accum.copy()\n        else:\n            ensemble_preds += model_accum\n\n        # free GPU\n        del model, sd\n        torch.cuda.empty_cache()\n\n    ensemble_preds = ensemble_preds / float(len(ckpts))\n    out_df = pd.DataFrame(ensemble_preds, columns=LABEL_COLS)\n    out_df.insert(0, 'Image_Name', names)\n    out_df.to_csv(args.out, index=False)\n    print(\"Saved submission:\", args.out)\n\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python /kaggle/working/inference.py \\\n  --image-dir /kaggle/input/grand-xray-slam-division-a/test1/ \\\n  --ckpt \"/kaggle/working/*best*.pth\" \\\n  --backbone convnext_large --img-size 1024 --batch-size 64 --tta-scales \"1024,800\" --tta-flip \\\n  --out /kaggle/working/submission.csv","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-02T08:09:53.894Z"}},"outputs":[],"execution_count":null}]}