{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install timm transformers -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T12:49:24.474229Z","iopub.execute_input":"2025-09-16T12:49:24.474857Z","iopub.status.idle":"2025-09-16T12:49:27.689708Z","shell.execute_reply.started":"2025-09-16T12:49:24.474802Z","shell.execute_reply":"2025-09-16T12:49:27.688683Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 1. Imports","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport os, time\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport torch, timm\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\nfrom torchvision import transforms\nfrom torch.optim import AdamW\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, precision_score, recall_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:22:14.770222Z","iopub.execute_input":"2025-09-16T13:22:14.770701Z","iopub.status.idle":"2025-09-16T13:22:14.775473Z","shell.execute_reply.started":"2025-09-16T13:22:14.770675Z","shell.execute_reply":"2025-09-16T13:22:14.774637Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Config","metadata":{}},{"cell_type":"code","source":"cfg = {\n    \"data_dir\": \"/kaggle/input/grand-xray-slam-division-a\",\n    \"train_csv\": \"train1.csv\",\n    \"test_csv\": \"sample_submission_1.csv\",\n    \"train_folder\": \"train1\",\n    \"test_folder\": \"test1\",\n\n    \"labels\": [\n        \"Atelectasis\", \"Cardiomegaly\", \"Consolidation\", \"Edema\",\n        \"Enlarged Cardiomediastinum\", \"Fracture\", \"Lung Lesion\",\n        \"Lung Opacity\", \"No Finding\", \"Pleural Effusion\", \"Pleural Other\",\n        \"Pneumonia\", \"Pneumothorax\", \"Support Devices\"\n    ],\n\n    \"img_col\": \"Image_name\",\n    \"img_size\": 224,\n    \"batch_size\": 32,\n    \"epochs\": 5,\n    \"lr\": 2e-5,\n    \"num_workers\": 4,\n    \"device\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    \"output_file\": \"submission.csv\",\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:24:17.470036Z","iopub.execute_input":"2025-09-16T13:24:17.470603Z","iopub.status.idle":"2025-09-16T13:24:17.474976Z","shell.execute_reply.started":"2025-09-16T13:24:17.47058Z","shell.execute_reply":"2025-09-16T13:24:17.474226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Dataset","metadata":{}},{"cell_type":"code","source":"LABEL_COLS = [\n    \"Atelectasis\", \"Cardiomegaly\", \"Consolidation\", \"Edema\",\n    \"Enlarged Cardiomediastinum\", \"Fracture\", \"Lung Lesion\",\n    \"Lung Opacity\", \"No Finding\", \"Pleural Effusion\", \"Pleural Other\",\n    \"Pneumonia\", \"Pneumothorax\", \"Support Devices\"\n]\n\nclass XRayDataset(Dataset):\n    def __init__(self, dataframe, img_dir, labels, img_col, transform=None, is_test=False):\n        self.df = dataframe.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.labels = labels\n        self.img_col = img_col\n        self.transform = transform\n        self.is_test = is_test\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.img_dir, row[self.img_col])\n        image = Image.open(img_path).convert(\"RGB\")\n\n        if self.transform:\n            image = self.transform(image)\n\n        if self.is_test:\n            return image, row[self.img_col]\n        else:\n            labels = torch.tensor(row[self.labels].values.astype(\"float32\"))\n            return image, labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:24:20.419992Z","iopub.execute_input":"2025-09-16T13:24:20.420669Z","iopub.status.idle":"2025-09-16T13:24:20.427852Z","shell.execute_reply.started":"2025-09-16T13:24:20.420644Z","shell.execute_reply":"2025-09-16T13:24:20.427123Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Transforms","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(os.path.join(cfg[\"data_dir\"], cfg[\"train_csv\"]))\ntest_df  = pd.read_csv(os.path.join(cfg[\"data_dir\"], cfg[\"test_csv\"]))\ntrain_split, valid_split = train_test_split(train_df, test_size=0.2, random_state=42)\n\ntrain_tfms = transforms.Compose([\n    transforms.Resize((cfg[\"img_size\"], cfg[\"img_size\"])),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(10),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.5,0.5,0.5), std=(0.5,0.5,0.5)),\n])\n\nvalid_tfms = transforms.Compose([\n    transforms.Resize((cfg[\"img_size\"], cfg[\"img_size\"])),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.5,0.5,0.5), std=(0.5,0.5,0.5)),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:24:23.32018Z","iopub.execute_input":"2025-09-16T13:24:23.320422Z","iopub.status.idle":"2025-09-16T13:24:23.720319Z","shell.execute_reply.started":"2025-09-16T13:24:23.320404Z","shell.execute_reply":"2025-09-16T13:24:23.719521Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Load Data","metadata":{}},{"cell_type":"code","source":"train_dataset = XRayDataset(train_split, os.path.join(cfg[\"data_dir\"], cfg[\"train_folder\"]),cfg[\"labels\"], cfg[\"img_col\"], transform=train_tfms)\nvalid_dataset = XRayDataset(valid_split, os.path.join(cfg[\"data_dir\"], cfg[\"train_folder\"]),cfg[\"labels\"], cfg[\"img_col\"], transform=valid_tfms)\ntest_dataset  = XRayDataset(test_df, os.path.join(cfg[\"data_dir\"], cfg[\"test_folder\"]),cfg[\"labels\"], cfg[\"img_col\"], transform=valid_tfms, is_test=True)\n\ntrain_loader = DataLoader(train_dataset, batch_size=cfg[\"batch_size\"], shuffle=True,num_workers=cfg[\"num_workers\"], pin_memory=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=cfg[\"batch_size\"], shuffle=False,num_workers=cfg[\"num_workers\"], pin_memory=True)\ntest_loader  = DataLoader(test_dataset, batch_size=cfg[\"batch_size\"], shuffle=False,num_workers=cfg[\"num_workers\"], pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:24:25.803132Z","iopub.execute_input":"2025-09-16T13:24:25.80379Z","iopub.status.idle":"2025-09-16T13:24:25.837351Z","shell.execute_reply.started":"2025-09-16T13:24:25.803766Z","shell.execute_reply":"2025-09-16T13:24:25.836626Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Model (ViT)","metadata":{}},{"cell_type":"code","source":"class ViTMultiLabel(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=True):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.backbone.head.in_features\n        self.backbone.reset_classifier(0)  # remove classification head\n        self.fc = nn.Sequential(\n            nn.Linear(in_features, in_features // 2),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(in_features // 2, num_classes)\n        )\n\n\n    def forward(self, x):\n        features = self.backbone(x)\n        out = self.fc(features)\n        return out\n\nmodel = ViTMultiLabel(\n    model_name=cfg.get(\"model_name\", \"vit_base_patch16_224\"),\n    num_classes=len(cfg[\"labels\"]),pretrained=True\n)\n\n# Multi-GPU support\nif torch.cuda.device_count() > 1:\n    print(\"Using\", torch.cuda.device_count(), \"GPUs\")\n    model = nn.DataParallel(model)\nmodel = model.to(cfg[\"device\"])\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = AdamW(model.parameters(), lr=cfg['lr'])\nscaler = torch.amp.GradScaler('cuda')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:25:41.91552Z","iopub.execute_input":"2025-09-16T13:25:41.91601Z","iopub.status.idle":"2025-09-16T13:25:43.48808Z","shell.execute_reply.started":"2025-09-16T13:25:41.915985Z","shell.execute_reply":"2025-09-16T13:25:43.487319Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Training & Validation","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(epoch):\n    model.train()\n    running_loss = 0.0\n    for images, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1} [Train]\"):\n        images, labels = images.to(cfg[\"device\"]), labels.to(cfg[\"device\"])\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n    avg_loss = running_loss / len(train_loader)\n    print(f\"Epoch {epoch+1} Train Loss: {avg_loss:.4f}\")\n\n\ndef validate(epoch, threshold=0.5):\n    model.eval()\n    running_loss = 0.0\n    all_labels, all_preds, all_probs = [], [], []\n\n    with torch.no_grad():\n        for images, labels in tqdm(valid_loader, desc=f\"Epoch {epoch+1} [Valid]\"):\n            images, labels = images.to(cfg[\"device\"]), labels.to(cfg[\"device\"])\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            running_loss += loss.item()\n\n            probs = torch.sigmoid(outputs).cpu().numpy()\n            preds = (probs > threshold).astype(int)\n\n            all_labels.append(labels.cpu().numpy())\n            all_probs.append(probs)\n            all_preds.append(preds)\n\n    avg_loss = running_loss / len(valid_loader)\n    all_labels = np.vstack(all_labels)\n    all_probs = np.vstack(all_probs)\n    all_preds = np.vstack(all_preds)\n\n    # Metrics\n    auc = roc_auc_score(all_labels, all_probs, average=\"macro\")\n    acc = accuracy_score(all_labels, all_preds)\n    f1_micro = f1_score(all_labels, all_preds, average=\"micro\")\n    f1_macro = f1_score(all_labels, all_preds, average=\"macro\")\n    precision = precision_score(all_labels, all_preds, average=\"macro\", zero_division=0)\n    recall = recall_score(all_labels, all_preds, average=\"macro\", zero_division=0)\n\n    print(\n        f\"Epoch {epoch+1} Valid Loss: {avg_loss:.4f}, \"\n        f\"AUC: {auc:.4f}, Acc: {acc:.4f}, \"\n        f\"F1(micro): {f1_micro:.4f}, F1(macro): {f1_macro:.4f}, \"\n        f\"Prec: {precision:.4f}, Recall: {recall:.4f}\"\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:26:35.390966Z","iopub.execute_input":"2025-09-16T13:26:35.391266Z","iopub.status.idle":"2025-09-16T13:26:35.398456Z","shell.execute_reply.started":"2025-09-16T13:26:35.391241Z","shell.execute_reply":"2025-09-16T13:26:35.397835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(cfg[\"epochs\"]):\n    train_one_epoch(epoch)\n    validate(epoch)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:26:37.640546Z","iopub.execute_input":"2025-09-16T13:26:37.640835Z","iopub.status.idle":"2025-09-16T13:45:55.499173Z","shell.execute_reply.started":"2025-09-16T13:26:37.640785Z","shell.execute_reply":"2025-09-16T13:45:55.495123Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 9. Inference on Test","metadata":{}},{"cell_type":"code","source":"model.eval()\nall_preds, all_ids = [], []\nwith torch.no_grad():\n    for images, ids in tqdm(test_loader, desc=\"Inference\"):\n        images = images.to(cfg[\"device\"])\n        outputs = model(images)\n        probs = torch.sigmoid(outputs).cpu().numpy()\n        preds = (probs > 0.5).astype(int)  # one-hot encoding\n        all_preds.append(preds)\n        all_ids.extend(ids)\n\nall_preds = np.vstack(all_preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:45:55.499974Z","iopub.status.idle":"2025-09-16T13:45:55.500402Z","shell.execute_reply.started":"2025-09-16T13:45:55.500189Z","shell.execute_reply":"2025-09-16T13:45:55.500224Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 10. Submission","metadata":{}},{"cell_type":"code","source":"submission = pd.DataFrame(all_preds, columns=cfg[\"labels\"])\nsubmission.insert(0, cfg[\"img_col\"], all_ids)\nsubmission.to_csv(cfg[\"output_file\"], index=False)\nprint(f\"✅ Submission saved to {cfg['output_file']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T13:45:55.501017Z","iopub.status.idle":"2025-09-16T13:45:55.501402Z","shell.execute_reply.started":"2025-09-16T13:45:55.501219Z","shell.execute_reply":"2025-09-16T13:45:55.501234Z"}},"outputs":[],"execution_count":null}]}