{"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":"import os, math, random, warnings, cv2\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision.transforms as T\nimport timm\nfrom tqdm import tqdm\nimport torch.nn.functional as F\n\n# -------------------------\n# Config\n# -------------------------\nSEED = 42\nIMG_SIZE = 512\nBATCH_SIZE = 8\nEPOCHS = 4\nWARMUP_EPOCHS = 1\nBASE_LR = 2e-5\nHEAD_LR = 8e-5\nWEIGHT_DECAY = 1e-2\nGRAD_CLIP_NORM = 1.0\nEMA_DECAY = 0.999\nFOCAL_GAMMA = 2.0\nNUM_WORKERS = 4\n\nTRAIN_CSV = \"/kaggle/input/grand-xray-slam-division-a/train1.csv\"\nTRAIN_DIR = \"/kaggle/input/grand-xray-slam-division-a/train1\"\nTEST_DIR  = \"/kaggle/input/grand-xray-slam-division-a/test1\"\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\nSAVE_PATH = \"convnextv2_gem_heavyattn_no_clahe.pth\"\n\n# -------------------------\n# Repro\n# -------------------------\ndef set_seed(seed=SEED):\n    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\nset_seed()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# -------------------------\n# Dataset without CLAHE\n# -------------------------\nclass XRayDataset(Dataset):\n    def __init__(self, df, image_dir, img_size=IMG_SIZE):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.img_size = img_size\n        self.tf = T.Compose([\n            T.ToTensor(),\n            T.RandomHorizontalFlip(p=0.5),\n            T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n        ])\n    def __len__(self): return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        path = os.path.join(self.image_dir, row['Image_name'])\n        img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n        if img is None:\n            img = np.zeros((self.img_size, self.img_size), dtype=np.uint8)\n        img = cv2.resize(img, (self.img_size, self.img_size), interpolation=cv2.INTER_CUBIC)\n        img = cv2.merge([img,img,img])\n        img = self.tf(img)\n        y = torch.tensor(row[LABEL_COLS].values.astype(np.float32), dtype=torch.float32)\n        return img, y\n\n# -------------------------\n# Focal Loss\n# -------------------------\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction=\"mean\"):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n    def forward(self, logits, targets):\n        bce = nn.functional.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n        pt = torch.exp(-bce)\n        loss = (1 - pt)**self.gamma * bce\n        if self.alpha is not None:\n            loss = loss * self.alpha.to(logits.device)\n        if self.reduction == \"mean\":\n            return loss.mean()\n        return loss.sum()\n\n# -------------------------\n# EMA\n# -------------------------\nclass EMA:\n    def __init__(self, model, decay=EMA_DECAY):\n        self.decay = decay\n        self.shadow = {n: p.detach().clone() for n, p in model.named_parameters() if p.requires_grad}\n    @torch.no_grad()\n    def update(self, model):\n        for n, p in model.named_parameters():\n            if n in self.shadow:\n                self.shadow[n].mul_(self.decay).add_(p.detach(), alpha=1-self.decay)\n    @torch.no_grad()\n    def apply_to(self, model):\n        for n, p in model.named_parameters():\n            if n in self.shadow:\n                p.copy_(self.shadow[n])\n\n# -------------------------\n# Data prep (use ALL data)\n# -------------------------\ndf = pd.read_csv(TRAIN_CSV)\ndf[LABEL_COLS] = df[LABEL_COLS].apply(pd.to_numeric, errors=\"coerce\").fillna(0)\n\nif 'No Finding' in LABEL_COLS:\n    others = [c for c in LABEL_COLS if c != 'No Finding']\n    df['No Finding'] = (df[others].sum(axis=1) == 0).astype(int)\n\n# alpha weights\npos_counts = df[LABEL_COLS].sum()\nneg_counts = len(df) - pos_counts\nalpha = torch.tensor((neg_counts/(pos_counts+1e-6)).values, dtype=torch.float32)\n\ntrain_ds = XRayDataset(df, TRAIN_DIR)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\n\n# -------------------------\n# GeM Pooling + Heavy Attention Head\n# -------------------------\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n    def forward(self, x):\n        return F.adaptive_avg_pool2d(x.clamp(min=self.eps).pow(self.p), (1,1)).pow(1./self.p)\n\nclass HeavyAttentionHead(nn.Module):\n    def __init__(self, in_ch, num_classes, num_heads=8, ff_mult=4):\n        super().__init__()\n        self.attn = nn.MultiheadAttention(embed_dim=in_ch, num_heads=num_heads, batch_first=True)\n        self.ff = nn.Sequential(\n            nn.LayerNorm(in_ch),\n            nn.Linear(in_ch, in_ch*ff_mult),\n            nn.GELU(),\n            nn.Linear(in_ch*ff_mult, in_ch),\n        )\n        self.ln = nn.LayerNorm(in_ch)\n        self.fc = nn.Linear(in_ch, num_classes)\n    def forward(self, x):\n        x = x.unsqueeze(1)\n        attn_out,_ = self.attn(x,x,x)\n        x = x + attn_out\n        x = x + self.ff(x)\n        x = self.ln(x)\n        x = x.squeeze(1)\n        return self.fc(x)\n\n# -------------------------\n# Full Model Module\n# -------------------------\nclass ConvNeXtV2_GeM_HeavyAttn(nn.Module):\n    def __init__(self, num_classes=len(LABEL_COLS)):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"convnextv2_base.fcmae_ft_in22k_in1k\",\n            pretrained=True,\n            num_classes=0\n        )\n        in_ch = self.backbone.num_features\n        self.gem = GeM()\n        self.dropout = nn.Dropout(0.3)\n        self.head = HeavyAttentionHead(in_ch, num_classes, num_heads=8, ff_mult=4)\n    def forward(self, x):\n        x = self.backbone.forward_features(x)\n        x = self.gem(x)\n        x = torch.flatten(x,1)\n        x = self.dropout(x)\n        x = self.head(x)\n        return x\n\ndef make_model():\n    return ConvNeXtV2_GeM_HeavyAttn(num_classes=len(LABEL_COLS))\n\ndef param_groups(m):\n    module = m.module if isinstance(m, nn.DataParallel) else m\n    base_params = [p for p in module.backbone.parameters()]\n    head_params = list(module.gem.parameters()) + list(module.dropout.parameters()) + list(module.head.parameters())\n    return [\n        {\"params\": base_params, \"lr\": BASE_LR, \"weight_decay\": WEIGHT_DECAY},\n        {\"params\": head_params, \"lr\": HEAD_LR, \"weight_decay\": WEIGHT_DECAY},\n    ]\n\ncriterion = FocalLoss(alpha=alpha, gamma=FOCAL_GAMMA)\n\n# -------------------------\n# Train on ALL data\n# -------------------------\nmodel = make_model().to(device)\n\nif torch.cuda.device_count() > 1:\n    print(f\"✅ Using {torch.cuda.device_count()} GPUs for DataParallel\")\n    model = nn.DataParallel(model)\n\noptimizer = optim.AdamW(param_groups(model))\nscheduler = optim.lr_scheduler.LambdaLR(\n    optimizer,\n    lr_lambda=lambda e: (e+1)/WARMUP_EPOCHS if e<WARMUP_EPOCHS\n    else 0.5*(1+math.cos(math.pi*(e-WARMUP_EPOCHS)/max(1,EPOCHS-WARMUP_EPOCHS)))\n)\nscaler = torch.cuda.amp.GradScaler(enabled=device.type==\"cuda\")\nema = EMA(model)\n\nfor epoch in range(EPOCHS):\n    model.train()\n    running=0\n    pbar=tqdm(train_loader,desc=f\"ConvNeXtV2+HeavyAttn Epoch {epoch+1}/{EPOCHS}\")\n    for imgs,y in pbar:\n        imgs,y=imgs.to(device),y.to(device)\n        optimizer.zero_grad(set_to_none=True)\n        with torch.cuda.amp.autocast(enabled=device.type==\"cuda\"):\n            out=model(imgs)\n            loss=criterion(out,y)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(),GRAD_CLIP_NORM)\n        scaler.step(optimizer); scaler.update()\n        ema.update(model)\n        running+=loss.item()\n        pbar.set_postfix(loss=running/max(1,pbar.n))\n    scheduler.step()\n\nema.apply_to(model)\nstate_dict = model.module.state_dict() if isinstance(model, nn.DataParallel) else model.state_dict()\ntorch.save(state_dict, SAVE_PATH)\nprint(f\"✅ Finished training on ALL data and saved to {SAVE_PATH}\")\n\n# -------------------------\n# Inference Dataset without CLAHE\n# -------------------------\nclass TestDataset(Dataset):\n    def __init__(self, df, image_dir):\n        self.df=df.reset_index(drop=True)\n        self.image_dir=image_dir\n        self.tf=T.Compose([\n            T.ToTensor(),\n            T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n        ])\n    def __len__(self):return len(self.df)\n    def __getitem__(self,idx):\n        row=self.df.iloc[idx]\n        path=os.path.join(self.image_dir,row['Image_name'])\n        img=cv2.imread(path,cv2.IMREAD_GRAYSCALE)\n        if img is None:\n            img=np.zeros((IMG_SIZE,IMG_SIZE),dtype=np.uint8)\n        img=cv2.resize(img,(IMG_SIZE,IMG_SIZE))\n        img=cv2.merge([img,img,img])\n        img=self.tf(img)\n        return img,row['Image_name']\n\n# Load best weights\nmodel_infer = make_model().to(device)\nif torch.cuda.device_count() > 1:\n    model_infer = nn.DataParallel(model_infer)\n\nstate_dict = torch.load(SAVE_PATH,map_location=device)\nif isinstance(model_infer, nn.DataParallel):\n    model_infer.module.load_state_dict(state_dict)\nelse:\n    model_infer.load_state_dict(state_dict)\nmodel_infer.eval()\n\ntest_names=sorted(os.listdir(TEST_DIR))\ntest_df=pd.DataFrame({'Image_name':test_names})\ntest_ds=TestDataset(test_df,TEST_DIR)\ntest_loader=DataLoader(test_ds,batch_size=BATCH_SIZE,shuffle=False,num_workers=NUM_WORKERS,pin_memory=True)\n\n@torch.no_grad()\ndef predict_tta(model,loader):\n    preds_all=[];names_all=[]\n    for imgs,names in tqdm(loader,desc=\"Predicting ConvNeXtV2+HeavyAttn\"):\n        imgs=imgs.to(device)\n        logits1=model(imgs)\n        imgs_flipped=torch.flip(imgs,dims=[3])\n        logits2=model(imgs_flipped)\n        logits=0.5*(logits1+logits2)\n        probs=torch.sigmoid(logits).cpu().numpy()\n        preds_all.append(probs)\n        names_all.extend(names)\n    return np.concatenate(preds_all,axis=0),names_all\n\npreds,names=predict_tta(model_infer,test_loader)\nsub=pd.DataFrame(preds,columns=LABEL_COLS)\nsub.insert(0,\"Image_name\",names)\nsub.to_csv(\"submission.csv\",index=False)\nprint(\"✅ Created submission.csv\")\nprint(sub.head())\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-24T04:56:55.684698Z","iopub.execute_input":"2025-09-24T04:56:55.684896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}