{"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":"none","dataSources":[{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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       # consider 6–8 if using more heads\nWARMUP_EPOCHS = 1\nBASE_LR = 2e-5\nHEAD_LR = 8e-5   # you may lower to 6e-5 for stability\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_multiheadattn_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 (no 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\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\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\n# -------------------------\n# Multi-Head Attention Block\n# -------------------------\nclass MultiHeadAttentionBlock(nn.Module):\n    def __init__(self, in_ch, num_classes, num_heads=4, ff_mult=4):\n        super().__init__()\n        self.num_heads = num_heads\n        self.attn_layers = nn.ModuleList()\n        self.ff_layers = nn.ModuleList()\n        self.ln_layers = nn.ModuleList()\n\n        for _ in range(num_heads):\n            attn = nn.MultiheadAttention(embed_dim=in_ch, num_heads=8, batch_first=True)\n            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            ln = nn.LayerNorm(in_ch)\n            self.attn_layers.append(attn)\n            self.ff_layers.append(ff)\n            self.ln_layers.append(ln)\n\n        self.fc_out = nn.Linear(in_ch * num_heads, num_classes)\n\n    def forward(self, x):\n        x = x.unsqueeze(1)  # [B,1,C]\n        head_outputs = []\n        for attn, ff, ln in zip(self.attn_layers, self.ff_layers, self.ln_layers):\n            attn_out,_ = attn(x,x,x)\n            h = x + attn_out\n            h = h + ff(h)\n            h = ln(h)\n            h = h.squeeze(1)  # [B,C]\n            head_outputs.append(h)\n        x_cat = torch.cat(head_outputs, dim=-1)  # [B, C*num_heads]\n        return self.fc_out(x_cat)\n\n# -------------------------\n# Full Model Module with Multi-Head Attention\n# -------------------------\nclass ConvNeXtV2_GeM_MultiHeadAttn(nn.Module):\n    def __init__(self, num_classes=len(LABEL_COLS), num_heads=4):\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 = MultiHeadAttentionBlock(in_ch, num_classes, num_heads=num_heads)\n\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_MultiHeadAttn(num_classes=len(LABEL_COLS), num_heads=4)\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+MultiHeadAttn 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\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+MultiHeadAttn\"):\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},"outputs":[],"execution_count":null}]}