{"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":113002,"databundleVersionId":13471427,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, math, random, warnings, cv2\nimport time\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport timm\nfrom tqdm import tqdm\nimport torch.nn.functional as F\n\n# Suppress warnings for cleaner output\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"OPENCV_IO_MAX_IMAGE_PIXELS\"] = str(2**64)\n\n# -------------------------\n# Config \n# -------------------------\nSEED = 42\nIMG_SIZE = 512\nBATCH_SIZE = 8\n# --- Training Parameters ---\nEPOCHS = 12 \nWARMUP_EPOCHS = 1\nBASE_LR = 2e-5\nHEAD_LR = 8e-5\n# --------------------------\nWEIGHT_DECAY = 1e-2\nGRAD_CLIP_NORM = 1.0\nEMA_DECAY = 0.999\nFOCAL_GAMMA = 2.0 \nNUM_WORKERS = 4\nVAL_SPLIT_RATIO = 0.1 \n\n# --- Paths ---\nTRAIN_CSV = \"/kaggle/input/grand-xray-slam-division-b/train2.csv\"\nTRAIN_DIR = \"/kaggle/input/grand-xray-slam-division-b/train2\"\nTEST_DIR  = \"/kaggle/input/grand-xray-slam-division-b/test2\"\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# --- Checkpointing and Logging Paths ---\nMODEL_NAME = \"convnextv2_pe_heavyattn_ASL\" \nBEST_MODEL_PATH = f\"{MODEL_NAME}_best_auc.pth\" \nLAST_CHECKPOINT_PATH = f\"{MODEL_NAME}_checkpoint.pth\"\nLOG_FILE_PATH = f\"training_log_{MODEL_NAME}.txt\" \nEARLY_STOPPING_PATIENCE = 3 \nMIXUP_ALPHA = 1.0 \n\n# --- Resume Checkpoint Path (Optional) ---\nRESUME_CHECKPOINT_PATH = \"\" \n# -------------------------\n\n# -------------------------\n# Utilities: Repro, Logging, EarlyStop \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()\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef log_message(message, filepath=LOG_FILE_PATH):\n    \"\"\"Appends a message to the log file and prints it to console.\"\"\"\n    if not os.path.exists(filepath):\n        with open(filepath, \"w\") as f:\n            f.write(\"Log initialized.\\n\")\n    with open(filepath, \"a\") as f:\n        f.write(f\"{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())} - {message}\\n\")\n    print(message)\n\nclass EarlyStopper:\n    def __init__(self, patience=EARLY_STOPPING_PATIENCE, min_delta=0):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.counter = 0\n        self.best_metric = -np.inf\n        self.early_stop = False\n\n    def __call__(self, val_metric):\n        if val_metric > self.best_metric + self.min_delta:\n            self.best_metric = val_metric\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n\ndef load_checkpoint(filepath, model, optimizer, scheduler, scaler, ema, device):\n    \"\"\"Loads a full training checkpoint to resume training.\"\"\"\n    if not os.path.exists(filepath):\n        log_message(f\"⚠️ Checkpoint not found at {filepath}. Starting from epoch 0.\")\n        return 0, -np.inf\n\n    log_message(f\"⌛ Attempting to load checkpoint from {filepath} for resumption...\")\n    \n    try:\n        checkpoint = torch.load(filepath, map_location=device, weights_only=False)\n    except Exception as e:\n        log_message(f\"❌ Error loading checkpoint dictionary: {e}. Starting from epoch 0.\")\n        return 0, -np.inf\n\n    start_epoch = checkpoint.get('epoch', 0)\n    best_auc = checkpoint.get('best_auc', -np.inf)\n    \n    model_to_load = model.module if isinstance(model, nn.DataParallel) else model\n    try:\n        model_to_load.load_state_dict(checkpoint['model_state_dict'], strict=False) \n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        scheduler.load_state_dict(checkpoint['scheduler_state_dict']) \n        if 'scaler_state_dict' in checkpoint:\n            scaler.load_state_dict(checkpoint['scaler_state_dict'])\n        if 'ema_shadow' in checkpoint:\n            ema.shadow = checkpoint['ema_shadow']\n        log_message(f\"✅ Checkpoint partially loaded. Resuming from Epoch {start_epoch}, Best AUC: {best_auc:.5f}\")\n    except Exception as e:\n        log_message(f\"❌ Error during full state load: {e}. Starting from epoch 0.\")\n        start_epoch = 0 \n        best_auc = -np.inf\n\n    return start_epoch, best_auc\n\n# -------------------------\n# Dataset with Augmentations\n# -------------------------\nclass XRayDataset(Dataset):\n    def __init__(self, df, image_dir, img_size=IMG_SIZE, is_train=True):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.img_size = img_size\n        \n        if is_train:\n            self.tf = T.Compose([\n                T.ToPILImage(),\n                T.RandomAffine(degrees=10, translate=(0.05, 0.05), scale=(0.95, 1.05)), \n                T.RandomHorizontalFlip(p=0.5),\n                T.ToTensor(),\n                T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n            ])\n        else:\n            self.tf = T.Compose([\n                T.ToPILImage(),\n                T.ToTensor(),\n                T.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n            ])\n\n    def __len__(self): return len(self.df)\n    \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        \n        if img is None:\n            img = np.zeros((self.img_size, self.img_size), dtype=np.uint8)\n        \n        img = cv2.resize(img, (self.img_size, self.img_size), interpolation=cv2.INTER_CUBIC)\n        img = cv2.merge([img,img,img]) \n        \n        img_tensor = self.tf(img)\n        \n        y = torch.tensor(row[LABEL_COLS].values.astype(np.float32), dtype=torch.float32)\n        return img_tensor, y\n\n# -------------------------\n# Asymmetric Loss (ASL)\n# -------------------------\nclass AsymmetricLoss(nn.Module):\n    def __init__(self, gamma_neg=2.0, gamma_pos=1.0, clip=0.0, eps=1e-8, disable_torch_grad_focal=False):\n        super(AsymmetricLoss, self).__init__()\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.eps = eps\n        self.disable_torch_grad_focal = disable_torch_grad_focal\n        self.reduction = 'mean'\n\n    def forward(self, x, y):\n        \"\"\"\"\n        Args:\n            x: input logits (N, C)\n            y: targets (N, C)\n        \"\"\"\n        xs_pos = torch.sigmoid(x)\n        xs_neg = 1.0 - xs_pos\n\n        if self.clip > 0:\n            xs_neg = (xs_neg + self.clip).clamp(max=1)\n\n        # ----------------- Positive Loss -----------------\n        pt = xs_pos * y\n        \n        one_sided_gamma = self.gamma_pos\n        if self.disable_torch_grad_focal:\n            torch.set_grad_enabled(False)\n        asymmetric_weight = torch.pow(1 - pt, one_sided_gamma)\n        if self.disable_torch_grad_focal:\n            torch.set_grad_enabled(True)\n        \n        log_pos = torch.log(xs_pos.clamp(min=self.eps))\n        loss_pos = asymmetric_weight * log_pos\n\n        # ----------------- Negative Loss -----------------\n        pt = xs_neg * (1 - y)\n        \n        asymmetric_weight = torch.pow(1 - pt, self.gamma_neg)\n        \n        log_neg = torch.log(xs_neg.clamp(min=self.eps))\n        loss_neg = asymmetric_weight * log_neg\n\n        # Total Loss\n        loss = - (loss_pos * y + loss_neg * (1 - y))\n        \n        if self.reduction == 'mean':\n            return loss.mean()\n        else:\n            return loss.sum()\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# 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        self.dropout = nn.Dropout(0.3)\n\n    def forward(self, x):\n        x = x.unsqueeze(1) \n        attn_out,_ = self.attn(x,x,x)\n        x = x + attn_out \n        ff_out = self.ff(x)\n        x = x + ff_out \n        x = self.ln(x)\n        x = x.squeeze(1) \n        x = self.dropout(x)\n        return self.fc(x)\n\n# -------------------------\n# Full Model Module (With Positional Encoding)\n# -------------------------\nclass ConvNeXtV2_PE_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        \n        # --- Learnable Global Positional Embedding ---\n        self.global_pe = nn.Parameter(torch.randn(1, in_ch))\n        \n        self.head = HeavyAttentionHead(in_ch, num_classes, num_heads=8, ff_mult=4)\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        \n        # --- Apply Positional Encoding ---\n        x = x + self.global_pe \n        \n        x = self.head(x)\n        return x\n\ndef make_model():\n    return ConvNeXtV2_PE_HeavyAttn(num_classes=len(LABEL_COLS))\n\n# --- FIXED param_groups to correctly include Positional Embedding ---\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    \n    # FIX: module.global_pe is already a Parameter, so it is included directly in a list [].\n    head_params = (\n        list(module.gem.parameters()) + \n        [module.global_pe] +                   # <-- FIXED HERE\n        list(module.head.parameters())\n    )\n    \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\n# -------------------------\n# Evaluation Metric (Mean AUC - Unchanged)\n# -------------------------\n@torch.no_grad()\ndef evaluate_auc(model, loader, device):\n    model.eval()\n    all_preds = []; all_targets = []\n    for imgs, y in loader:\n        imgs = imgs.to(device)\n        with torch.cuda.amp.autocast(enabled=device.type == \"cuda\"):\n            logits = model(imgs)\n        probs = torch.sigmoid(logits).cpu().numpy()\n        all_preds.append(probs); all_targets.append(y.cpu().numpy())\n    all_preds = np.concatenate(all_preds, axis=0); all_targets = np.concatenate(all_targets, axis=0)\n    aucs = []\n    for i in range(all_targets.shape[1]):\n        if len(np.unique(all_targets[:, i])) > 1:\n            auc = roc_auc_score(all_targets[:, i], all_preds[:, i])\n            aucs.append(auc)\n    mean_auc = np.mean(aucs) if aucs else 0.0\n    return mean_auc\n\n# -------------------------\n# Data prep (with Validation Split - Unchanged)\n# -------------------------\ndf_full = pd.read_csv(TRAIN_CSV)\ndf_full[LABEL_COLS] = df_full[LABEL_COLS].apply(pd.to_numeric, errors=\"coerce\").fillna(0)\nif 'No Finding' in LABEL_COLS:\n    others = [c for c in LABEL_COLS if c != 'No Finding']\n    df_full['No Finding'] = (df_full[others].sum(axis=1) == 0).astype(int)\n\ndf_train, df_val = train_test_split(\n    df_full, test_size=VAL_SPLIT_RATIO, random_state=SEED, shuffle=True\n)\n\npos_counts = df_train[LABEL_COLS].sum()\nneg_counts = len(df_train) - pos_counts\nalpha_tensor = torch.tensor((neg_counts / (pos_counts + 1e-6)).values, dtype=torch.float32) \n\ntrain_ds = XRayDataset(df_train, TRAIN_DIR, is_train=True)\nval_ds = XRayDataset(df_val, TRAIN_DIR, is_train=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE * 2, shuffle=False, \n                        num_workers=NUM_WORKERS, pin_memory=True)\n\n# Instantiate the Asymmetric Loss\ncriterion = AsymmetricLoss(gamma_neg=FOCAL_GAMMA, gamma_pos=1.0)\n\n# -------------------------\n# Training Loop with Checkpointing and Early Stop\n# -------------------------\nmodel = make_model().to(device)\nif torch.cuda.device_count() > 1:\n    log_message(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))))\nscaler = torch.cuda.amp.GradScaler(enabled=device.type == \"cuda\")\nema = EMA(model)\nearly_stopper = EarlyStopper()\n\nSTART_EPOCH, best_auc = load_checkpoint(\n    RESUME_CHECKPOINT_PATH, model, optimizer, scheduler, scaler, ema, device\n)\n\nwith open(LOG_FILE_PATH, \"w\") as f:\n    f.write(f\"--- Grand X-Ray Slam Division A Training Log ({MODEL_NAME}) ---\\n\")\nif START_EPOCH > 0:\n    log_message(f\"--- RESUMING Training ({MODEL_NAME}) from Epoch {START_EPOCH} ---\")\nelse:\n    log_message(f\"--- Starting Training from Scratch ({MODEL_NAME}) ---\")\n    \nlog_message(f\"Hyperparams: Epochs={EPOCHS}, BS={BATCH_SIZE}, IMG_SIZE={IMG_SIZE}, LR(Base)={BASE_LR}, LR(Head)={HEAD_LR}, Loss=ASL(gamma_neg={FOCAL_GAMMA}, gamma_pos=1.0), Patience={EARLY_STOPPING_PATIENCE}\")\n\nfor epoch in range(START_EPOCH, EPOCHS): \n    start_time = time.time()\n    \n    # --- Training Phase ---\n    model.train()\n    running_loss = 0\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch + 1}/{EPOCHS} (Train)\")\n    \n    for batch_idx, (imgs, y) in enumerate(pbar):\n        # --- MixUp Augmentation ---\n        if random.random() < 0.5 and MIXUP_ALPHA > 0:\n            lam = np.random.beta(MIXUP_ALPHA, MIXUP_ALPHA) \n            index = torch.randperm(imgs.size(0)).to(imgs.device)\n            mixed_imgs = lam * imgs + (1 - lam) * imgs[index, :]\n            y_mix = lam * y + (1 - lam) * y[index, :]\n            imgs, y_target = mixed_imgs, y_mix\n        else:\n            y_target = y\n        # --- End MixUp ---\n        \n        imgs, y_target = imgs.to(device), y_target.to(device)\n        optimizer.zero_grad(set_to_none=True)\n        \n        with torch.cuda.amp.autocast(enabled=device.type == \"cuda\"):\n            out = model(imgs)\n        loss = criterion(out, y_target)\n        \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP_NORM)\n        scaler.step(optimizer)\n        scaler.update()\n        ema.update(model)\n        \n        running_loss += loss.item()\n        pbar.set_postfix(loss=running_loss / max(1, batch_idx + 1))\n        \n    scheduler.step()\n\n    # --- Validation Phase & Checkpointing ---\n    model_to_eval = model.module if isinstance(model, nn.DataParallel) else model\n    ema.apply_to(model_to_eval)\n    \n    current_auc = evaluate_auc(model, val_loader, device)\n    \n    # Logging\n    train_loss = running_loss / len(train_loader)\n    epoch_duration = time.time() - start_time\n    log_message(\n        f\"Epoch {epoch+1}/{EPOCHS} | Train Loss: {train_loss:.5f} | Val Mean AUC: {current_auc:.5f} | Time: {epoch_duration:.0f}s\"\n    )\n\n    # 1. Save Best Model\n    if current_auc > best_auc:\n        best_auc = current_auc\n        state_dict_to_save = model_to_eval.state_dict()\n        torch.save(state_dict_to_save, BEST_MODEL_PATH)\n        log_message(f\"⭐ NEW BEST Model saved with AUC: {best_auc:.5f} to {BEST_MODEL_PATH}\")\n\n    # 2. Save Last Checkpoint (for resume)\n    torch.save(\n        {\n            'epoch': epoch + 1,\n            'model_state_dict': model_to_eval.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'scaler_state_dict': scaler.state_dict(),\n            'ema_shadow': ema.shadow,\n            'best_auc': best_auc,\n        },\n        LAST_CHECKPOINT_PATH\n    )\n\n    # 3. Early Stopping Check\n    early_stopper(current_auc)\n    if early_stopper.early_stop:\n        log_message(f\"🚨 Early stopping triggered after {early_stopper.patience} epochs without improvement.\")\n        break\n\nlog_message(f\"✅ Finished training.\")\n\n# -------------------------\n# Inference on new test set (using the BEST saved model)\n# -------------------------\nclass TestDataset(Dataset):\n    def __init__(self, df, image_dir):\n        self.df=df.reset_index(drop=True); self.image_dir=image_dir\n        self.tf=T.Compose([\n            T.ToPILImage(),\n            T.ToTensor(), 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]; path=os.path.join(self.image_dir,row['Image_name'])\n        img=cv2.imread(path,cv2.IMREAD_GRAYSCALE)\n        if img is None: img=np.zeros((IMG_SIZE,IMG_SIZE),dtype=np.uint8)\n        img=cv2.resize(img,(IMG_SIZE,IMG_SIZE)); img=cv2.merge([img,img,img])\n        img=self.tf(img); return img,row['Image_name']\n\n# Load the BEST weights for inference\nmodel_infer = make_model().to(device)\nif torch.cuda.device_count() > 1:\n    model_infer = nn.DataParallel(model_infer)\n\nweights_path = BEST_MODEL_PATH\nif not os.path.exists(weights_path):\n    weights_path = LAST_CHECKPOINT_PATH if os.path.exists(LAST_CHECKPOINT_PATH) else RESUME_CHECKPOINT_PATH\n\nlog_message(f\"Loading weights from {weights_path} for inference...\")\n\ntry:\n    checkpoint = torch.load(weights_path, map_location=device, weights_only=False)\n    state_dict = checkpoint.get('model_state_dict', checkpoint)\nexcept Exception as e:\n    log_message(f\"❌ Error loading inference weights: {e}. Attempting simple load.\")\n    state_dict = torch.load(weights_path, map_location=device, weights_only=False)\n\ndef load_state(model, state_dict):\n    new_state_dict = {}\n    for k, v in state_dict.items():\n        if k.startswith('module.'):\n            new_state_dict[k[7:]] = v\n        else:\n            new_state_dict[k] = v\n    model.load_state_dict(new_state_dict, strict=False)\n\nif isinstance(model_infer, nn.DataParallel):\n    load_state(model_infer.module, state_dict)\nelse:\n    load_state(model_infer, state_dict)\n\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*2,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=f\"Predicting {MODEL_NAME}\"):\n        imgs=imgs.to(device)\n        with torch.cuda.amp.autocast(enabled=device.type == \"cuda\"):\n            logits1=model(imgs)\n        imgs_flipped=torch.flip(imgs,dims=[3])\n        with torch.cuda.amp.autocast(enabled=device.type == \"cuda\"):\n            logits2=model(imgs_flipped)\n            \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)\nlog_message(\"✅ Created submission.csv\")\nprint(sub.head())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-12T11:10:11.518571Z","iopub.execute_input":"2025-10-12T11:10:11.5189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}