{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":97984,"databundleVersionId":14096757},{"sourceType":"datasetVersion","sourceId":15977916,"datasetId":10246688,"databundleVersionId":16939090},{"sourceType":"datasetVersion","sourceId":15987333,"datasetId":10252848,"databundleVersionId":16949294},{"sourceType":"datasetVersion","sourceId":15991332,"datasetId":10255762,"databundleVersionId":16953682},{"sourceType":"datasetVersion","sourceId":13731160,"datasetId":8733970,"databundleVersionId":14479231},{"sourceType":"datasetVersion","sourceId":15934274,"datasetId":9997958,"databundleVersionId":16891884},{"sourceType":"datasetVersion","sourceId":14695416,"datasetId":9387663,"databundleVersionId":15539655},{"sourceType":"kernelVersion","sourceId":313720038}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n!pip install --no-deps segmentation-models-pytorch==0.5.0\nimport sys\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport multiprocessing as mp\nimport pickle\nfrom timeit import default_timer as timer\n\n# YOUR modified file takes priority\nsys.path.insert(0, '/kaggle/input/datasets/zahouaniyacine/my-stage2-lead-model') #mine\nsys.path.append('/kaggle/input/datasets/takashisomeya/physionet-final-submission-models')\nsys.path.append('/kaggle/input/datasets/hengck23/hengck23-demo-submit-physionet')\n\nprint(sys.path)\nDEVICE = 'cuda'\nFLOAT_TYPE = torch.float16   # or torch.float16 for T4\n\nKAGGLE_DIR   = '/kaggle/input/competitions/physionet-ecg-image-digitization'\nWEIGHT_DIR   = '/kaggle/input/datasets/takashisomeya/physionet-final-submission-models'\nOUT_DIR      = '/kaggle/working/outputs'\nSAVE_DIR     = '/kaggle/working/checkpoints'\n\nos.makedirs(OUT_DIR,  exist_ok=True)\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# Training config\nTARGET_TOTAL_EPOCHS = 4      # total target over all sessions\nEPOCHS_THIS_SESSION = 1      # exactly 1 epoch per Kaggle run\nLR = 3e-5\nBATCH_SIZE = 1\nSAVE_EVERY_STEPS = 1000      # emergency checkpoint inside epoch\nRESUME_CKPT = '/kaggle/input/datasets/zahouaniyacine/attention-checkpoints-3/cross_attn_b6_last_full_2.pth'           # set later to /kaggle/input/your-checkpoint-dataset/last_full.pth\n\nWINDOW_SIZE = 240 \nOFFSET      = 416\n#IGNORE_EDGE = 8\nxscale = 5000 / (2080 - 118)\naddx = 1\nyscale = 1\nIMGH, IMGW = int(1700 * yscale), int(2200 * xscale + addx)\n\nx0, x1 = 0, 5600\ny0, y1 = 0, 1696\nzero_mv = [703.5, 987.5, 1271.5, 1531.5]\n\n\n\nprint('constants ok')\n#Step 4 — Import your modified model\n\nimport stage2_lead_model\nprint('Loading from:', stage2_lead_model.__file__)\n# Must print your dataset path, not the original\n\nfrom stage2_lead_model import Net as LeadModel\nfrom stage2_smp_model  import Net as WholeModel   # from original dataset\nfrom stage2_common     import *\nfrom stage2_model      import prob_to_series_by_max\nprint('imports ok')\n\n\nimport os\n\nMASK_DIR = '/kaggle/input/notebooks/m1h4wk22/generate-pseudo-masks/output/masks'\nRECT_DIR  = '/kaggle/input/notebooks/m1h4wk22/generate-pseudo-masks/output/rectified'\n\nvalid_id = sorted([\n    f.replace('.mask-coo.npz', '')\n    for f in os.listdir(MASK_DIR)\n    if f.endswith('.mask-coo.npz')\n])\n\nprint('nb training samples =', len(valid_id))\nprint(valid_id[:5])\n\nFAIL_ID = []","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:12.878999Z","iopub.execute_input":"2026-04-28T16:34:12.880017Z","iopub.status.idle":"2026-04-28T16:34:29.061298Z","shell.execute_reply.started":"2026-04-28T16:34:12.879962Z","shell.execute_reply":"2026-04-28T16:34:29.060301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  Hors de la classe — fonction globale\ndef read_images(path):\n    image = cv2.imread(path, cv2.IMREAD_COLOR)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    image = cv2.resize(image, (IMGW, IMGH), interpolation=cv2.INTER_LINEAR)\n    trim_image = image.copy()[OFFSET:y1, x0:x1]\n    image = image[y0:y1, x0:x1]\n    H, W, _ = image.shape\n    lead_images = []\n    for zmv in zero_mv:\n        h0, h1 = int(zmv - WINDOW_SIZE), int(zmv + WINDOW_SIZE)\n        src_h0, src_h1 = max(0, h0), min(H, h1)\n        dst_h0 = src_h0 - h0\n        dst_h1 = dst_h0 + (src_h1 - src_h0)\n        lead_img = np.zeros((WINDOW_SIZE * 2, W, 3), np.uint8)\n        lead_img[dst_h0:dst_h1] = image[src_h0:src_h1]\n        lead_images.append(lead_img)\n    lead_images = np.stack(lead_images)\n    return trim_image, lead_images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.063168Z","iopub.execute_input":"2026-04-28T16:34:29.063690Z","iopub.status.idle":"2026-04-28T16:34:29.071089Z","shell.execute_reply.started":"2026-04-28T16:34:29.063648Z","shell.execute_reply":"2026-04-28T16:34:29.070122Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Step 6 — Dataset class**\nThe key thing here: you need ground-truth pixel masks. The 2nd solution trained on PTB-XL synthetic ECG images where the exact waveform is known, so the mask is computed from the known signal position. For your case, since you don't have those masks ready, the simplest approach is to use pseudo-labels: run Stage 2 inference once with the existing conv2d model to generate pixel maps, save them as .npy, then use those as training targets for fine-tuning the attention layers.","metadata":{}},{"cell_type":"code","source":"\n\n\n\ndef load_sparse_mask(path, shape=(4, 1700, 5600)):\n    \"\"\"\n    Charge un masque sparse COO .npz → dense float32 (4, H, W).\n    Compatible avec le format du repo Someya.\n    \"\"\"\n    d = np.load(path)\n    # Si le shape est sauvegardé dans le fichier, on l'utilise\n    H, W = int(d['shape'][1]), int(d['shape'][2])\n    mask = np.zeros((4, H, W), dtype=np.float32)\n    for ch in range(4):\n        key_y = f'ch{ch}_y'\n        key_x = f'ch{ch}_x'\n        key_v = f'ch{ch}_v'\n        if key_y in d and len(d[key_y]) > 0:\n            mask[ch, d[key_y], d[key_x]] = d[key_v]\n    return mask  # (4, H, W)\n\nclass Stage2AttentionDataset(Dataset):\n    def __init__(self, sample_ids, fail_ids=None):\n        if fail_ids is None:\n            fail_ids = []\n        self.ids = [s for s in sample_ids if s not in fail_ids]\n\n    def __len__(self):\n        return len(self.ids)\n\n  \n    def __getitem__(self, idx):\n        sample_id = self.ids[idx]\n\n        img_path = f'{RECT_DIR}/{sample_id}.rect.jpg'\n        trim_image, lead_images = read_images(img_path)\n\n        image_tensor = torch.from_numpy(\n            lead_images.transpose(0, 3, 1, 2)\n        ).byte()   # (4,3,H,W)\n\n        # ← Charge le sparse COO .npz au lieu du .mask.npy dense\n        mask = load_sparse_mask(f'{MASK_DIR}/{sample_id}.mask-coo.npz')\n        # mask: (4, H, W) → ajoute la dim canal attendue par le modèle\n        mask_tensor = torch.from_numpy(mask).float().unsqueeze(1)  # (4, 1, H, W) : .unsqueeze(1) ? \n\n        return {\n            'image': image_tensor,\n            'pixel': mask_tensor,\n            'sample_id': sample_id,\n        }\n\nprint('dataset class ok')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.072111Z","iopub.execute_input":"2026-04-28T16:34:29.072522Z","iopub.status.idle":"2026-04-28T16:34:29.100491Z","shell.execute_reply.started":"2026-04-28T16:34:29.072478Z","shell.execute_reply":"2026-04-28T16:34:29.099701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = Stage2AttentionDataset(valid_id[:4], FAIL_ID)\nsample = ds[0]\nprint(sample['image'].shape, sample['pixel'].shape)\nassert sample['image'].shape[-2:] == sample['pixel'].shape[-2:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.102577Z","iopub.execute_input":"2026-04-28T16:34:29.103173Z","iopub.status.idle":"2026-04-28T16:34:29.307258Z","shell.execute_reply.started":"2026-04-28T16:34:29.103144Z","shell.execute_reply":"2026-04-28T16:34:29.306287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test sur un seul sample\nsid = valid_id[0]\nimg_path = f'{RECT_DIR}/{sid}.rect.jpg'\nassert os.path.exists(img_path), f\"Image manquante: {img_path}\"\nmask_path = f'{MASK_DIR}/{sid}.mask-coo.npz'\nassert os.path.exists(mask_path), f\"Masque manquant: {mask_path}\"\n\ntrim, leads = read_images(img_path)\nmask = load_sparse_mask(mask_path)\nprint(f\"leads shape: {leads.shape}\")   # doit être (4, H, W, 3)\nprint(f\"mask shape:  {mask.shape}\")    # doit être (4, H, W)\nprint(f\"mask nonzero pixels: {np.count_nonzero(mask)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.308530Z","iopub.execute_input":"2026-04-28T16:34:29.308917Z","iopub.status.idle":"2026-04-28T16:34:29.449150Z","shell.execute_reply.started":"2026-04-28T16:34:29.308890Z","shell.execute_reply":"2026-04-28T16:34:29.448440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys, importlib\nimport torch.nn as nn\n\n# ── 1. Recharger le module original ──\nimport stage2_lead_model as _m\nimportlib.reload(_m)  # recharge depuis le .py read-only\n\n# ── 2. Redéfinir CrossLeadAttentionFusion proprement ──\nclass CrossLeadAttentionFusion(nn.Module):\n    def __init__(self, channels, num_leads=4, num_heads=4, dropout=0.0):\n        super().__init__()\n        self.channels  = channels\n        self.num_leads = num_leads\n        # Calcul local, ignore le num_heads passé en argument\n        self.num_heads = next(\n            (h for h in [8, 4, 2, 1] if channels % h == 0), 1\n        )\n        self.norm = nn.LayerNorm(channels)\n        self.attn = nn.MultiheadAttention(\n            embed_dim   = channels,\n            num_heads   = self.num_heads,\n            dropout     = dropout,\n            batch_first = True,\n        )\n        self.ffn = nn.Sequential(\n            nn.LayerNorm(channels),\n            nn.Linear(channels, channels * 2),\n            nn.GELU(),\n            nn.Linear(channels * 2, channels),\n        )\n\n    def forward(self, x, batch_size):\n        B = batch_size\n        _, C, H, W = x.shape\n        x_leads = x.view(B, self.num_leads, C, H, W)\n        x_hw    = x_leads.permute(0, 3, 4, 1, 2)\n        x_seq   = x_hw.reshape(B * H * W, self.num_leads, C)\n        x_norm  = self.norm(x_seq)\n        attn_out, _ = self.attn(x_norm, x_norm, x_norm)\n        x_seq   = x_seq + attn_out\n        x_seq   = x_seq + self.ffn(x_seq)\n        x_hw    = x_seq.reshape(B, H, W, self.num_leads, C)\n        x_out   = x_hw.permute(0, 3, 4, 1, 2)\n        return x_out.reshape(B * self.num_leads, C, H, W)\n\n# ── 3. Injecter dans le module chargé ──\n_m.CrossLeadAttentionFusion = CrossLeadAttentionFusion\n\n# ── 4. Rendre LeadModel accessible depuis ton notebook ──\nLeadModel = _m.Net\n\nprint(\"✅ CrossLeadAttentionFusion patché avec succès\")\n# Vérification rapide\nfor c in [200, 344, 576]:\n    nh = next((h for h in [8,4,2,1] if c % h == 0), 1)\n    print(f\"  channels={c} → num_heads={nh}, {c}%{nh}={c%nh} ✅\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.450107Z","iopub.execute_input":"2026-04-28T16:34:29.450582Z","iopub.status.idle":"2026-04-28T16:34:29.467322Z","shell.execute_reply.started":"2026-04-28T16:34:29.450542Z","shell.execute_reply":"2026-04-28T16:34:29.466524Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Step 7 — Training function (single GPU)**\n","metadata":{}},{"cell_type":"markdown","source":"def set_trainable(module, flag):\n    for p in module.parameters():\n        p.requires_grad = flag\n\ndef apply_progressive_unfreezing(model, epoch):\n    set_trainable(model.encoder, False)\n\n    set_trainable(model.decoder, True)\n    set_trainable(model.fusion_modules, True)\n    set_trainable(model.pixel_head, True)\n\n    enc = model.encoder.model if hasattr(model.encoder, 'model') else model.encoder\n\n    if epoch == 0:\n        stage = \"encoder frozen\"\n\n    elif epoch == 1:\n        if hasattr(enc, 'blocks'):\n            set_trainable(enc.blocks[-1], True)\n            if hasattr(enc, 'conv_head'):\n                set_trainable(enc.conv_head, True)\n            if hasattr(enc, 'bn2'):\n                set_trainable(enc.bn2, True)\n            stage = \"last encoder block\"\n        else:\n            stage = \"fallback frozen\"\n\n    else:\n        if hasattr(enc, 'blocks'):\n            set_trainable(enc.blocks[-1], True)\n            set_trainable(enc.blocks[-2], True)\n            if hasattr(enc, 'conv_head'):\n                set_trainable(enc.conv_head, True)\n            if hasattr(enc, 'bn2'):\n                set_trainable(enc.bn2, True)\n            stage = \"last 2 encoder blocks\"\n        else:\n            stage = \"fallback frozen\"\n\n    return stage\n\ndef train_one_gpu(gpu_id=0, assigned_ids=None, prev_fail_ids=None, result_file=None):\n    device = f'cuda:{gpu_id}'\n    if assigned_ids is None:\n        assigned_ids = valid_id\n    if prev_fail_ids is None:\n        prev_fail_ids = FAIL_ID\n\n    model = LeadModel(\n        encoder_name='tu-efficientnet_b6',\n        encoder_weights=None,\n        fusion_type='cross_attn',\n        fusion_levels=[3, 4],\n    )\n\n    ckpt_path = f'{WEIGHT_DIR}/series_b6_shared_conv2d_lb23.10.pth'\n    state = torch.load(ckpt_path, map_location='cpu')\n    missing, unexpected = model.load_state_dict(state, strict=False)\n    print(f'GPU{gpu_id} | missing keys: {len(missing)}')\n    print(f'GPU{gpu_id} | unexpected keys: {len(unexpected)}')\n\n    model.to(device)\n    model.output_type = ['loss', 'dice_loss', 'infer']\n\n    dataset = Stage2AttentionDataset(assigned_ids, prev_fail_ids)\n    loader = DataLoader(\n        dataset,\n        batch_size=1,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True\n    )\n\n    optimizer = AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(\n        optimizer,\n        T_max=EPOCHS * len(loader),\n        eta_min=3e-7\n    )\n    scaler = torch.amp.GradScaler('cuda')\n\n    best_loss = 1e9\n    best_ckpt = f'{SAVE_DIR}/cross_attn_b6_best.pth'\n\n    for epoch in range(EPOCHS):\n        stage = apply_progressive_unfreezing(model, epoch)\n        model.train()\n\n        epoch_loss = 0.0\n        start = timer()\n\n        for n, batch in enumerate(loader):\n            optimizer.zero_grad(set_to_none=True)\n\n            with torch.amp.autocast('cuda', dtype=torch.float16):\n                output = model(batch)\n                loss = output['pixel_loss'] + output['pixel_dice_loss']\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n\n            epoch_loss += loss.item()\n            print(\n                f'\\r GPU{gpu_id} epoch {epoch+1}/{EPOCHS} '\n                f'step {n+1}/{len(loader)} loss {loss.item():.4f}',\n                end='', flush=True\n            )\n\n        avg = epoch_loss / len(loader)\n        elapsed = timer() - start\n        print(f'\\n GPU{gpu_id} epoch {epoch+1}/{EPOCHS} avg_loss={avg:.4f} time={elapsed:.0f}s stage={stage}')\n\n        ep_ckpt = f'{SAVE_DIR}/cross_attn_b6_gpu{gpu_id}_ep{epoch+1}.pth'\n        torch.save(model.state_dict(), ep_ckpt)\n\n        if avg < best_loss:\n            best_loss = avg\n            torch.save(model.state_dict(), best_ckpt)\n            print(f' GPU{gpu_id} new best -> {best_ckpt}')\n\n    if result_file:\n        with open(result_file, 'wb') as f:\n            pickle.dump({'gpu': gpu_id, 'final_loss': best_loss}, f)","metadata":{"execution":{"iopub.status.busy":"2026-04-26T14:05:58.793675Z","iopub.execute_input":"2026-04-26T14:05:58.794263Z","iopub.status.idle":"2026-04-26T14:05:58.800426Z","shell.execute_reply.started":"2026-04-26T14:05:58.79424Z","shell.execute_reply":"2026-04-26T14:05:58.799836Z"}}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def optimizer_to(optimizer, device):\n    for state in optimizer.state.values():\n        for k, v in state.items():\n            if torch.is_tensor(v):\n                state[k] = v.to(device)\ndef set_trainable(module, flag):\n    for p in module.parameters():\n        p.requires_grad = flag\n\ndef apply_progressive_unfreezing(model, global_epoch):\n    set_trainable(model.encoder, False)\n\n    set_trainable(model.decoder, True)\n    set_trainable(model.fusion_modules, True)\n    set_trainable(model.pixel_head, True)\n\n    enc = model.encoder.model if hasattr(model.encoder, 'model') else model.encoder\n\n    if global_epoch == 0:\n        stage = \"encoder frozen\"\n\n    elif global_epoch == 1:\n        if hasattr(enc, 'blocks'):\n            set_trainable(enc.blocks[-1], True)\n            if hasattr(enc, 'conv_head'):\n                set_trainable(enc.conv_head, True)\n            if hasattr(enc, 'bn2'):\n                set_trainable(enc.bn2, True)\n            stage = \"last encoder block\"\n        else:\n            stage = \"fallback frozen\"\n\n    else:\n        if hasattr(enc, 'blocks'):\n            set_trainable(enc.blocks[-1], True)\n            set_trainable(enc.blocks[-2], True)\n            if hasattr(enc, 'conv_head'):\n                set_trainable(enc.conv_head, True)\n            if hasattr(enc, 'bn2'):\n                set_trainable(enc.bn2, True)\n            stage = \"last 2 encoder blocks\"\n        else:\n            stage = \"fallback frozen\"\n\n    return stage\n\n\ndef save_full_checkpoint(path, model, optimizer, scheduler, scaler,\n                         global_epoch, best_loss, last_avg_loss,\n                         completed_epoch, extra=None):\n    ckpt = {\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict(),\n        'scaler_state_dict': scaler.state_dict(),\n        'global_epoch': global_epoch,\n        'best_loss': best_loss,\n        'last_avg_loss': last_avg_loss,\n        'completed_epoch': completed_epoch,\n        'extra': extra if extra is not None else {},\n    }\n    tmp_path = path + '.tmp'\n    torch.save(ckpt, tmp_path)\n    os.replace(tmp_path, path)\n\ndef train_one_gpu(gpu_id=0, assigned_ids=None, prev_fail_ids=None,\n                  result_file=None, resume_ckpt=None):\n    device = torch.device(f'cuda:{gpu_id}')\n    if assigned_ids is None:\n        assigned_ids = valid_id\n    if prev_fail_ids is None:\n        prev_fail_ids = FAIL_ID\n\n    model = LeadModel(\n        encoder_name='tu-efficientnet_b6',\n        encoder_weights=None,\n        fusion_type='cross_attn',\n        fusion_levels=[3, 4],\n    ).to(device)\n\n    dataset = Stage2AttentionDataset(assigned_ids, prev_fail_ids)\n    loader = DataLoader(\n        dataset,\n        batch_size=1,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=False\n    )\n\n    optimizer = AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(\n        optimizer,\n        T_max=TARGET_TOTAL_EPOCHS * len(loader),\n        eta_min=3e-7\n    )\n    scaler = torch.amp.GradScaler('cuda')\n\n    start_global_epoch = 0\n    best_loss = 1e9\n\n    if resume_ckpt is not None and os.path.exists(resume_ckpt):\n        ckpt = torch.load(resume_ckpt, map_location='cpu')\n        model.load_state_dict(ckpt['model_state_dict'], strict=False)\n        optimizer.load_state_dict(ckpt['optimizer_state_dict'])\n        scheduler.load_state_dict(ckpt['scheduler_state_dict'])\n        scaler.load_state_dict(ckpt['scaler_state_dict'])\n\n        optimizer_to(optimizer, device)\n\n        completed_epoch = ckpt.get('completed_epoch', True)\n        if completed_epoch:\n            start_global_epoch = ckpt['global_epoch'] + 1\n        else:\n            start_global_epoch = ckpt['global_epoch']\n\n        best_loss = ckpt.get('best_loss', 1e9)\n        print(f'GPU{gpu_id} | resumed from {resume_ckpt}')\n        print(f'GPU{gpu_id} | start_global_epoch = {start_global_epoch}')\n        print(f'GPU{gpu_id} | best_loss = {best_loss:.6f}')\n        \n    model.to(device)\n    model.output_type = ['loss', 'dice_loss', 'infer']\n\n    last_ckpt = f'{SAVE_DIR}/cross_attn_b6_last_full.pth'\n    best_ckpt = f'{SAVE_DIR}/cross_attn_b6_best_full.pth'\n    final_out = '/kaggle/working/cross_attn_b6_final.pth'\n\n    end_global_epoch = min(start_global_epoch + EPOCHS_THIS_SESSION, TARGET_TOTAL_EPOCHS)\n\n    for global_epoch in range(start_global_epoch, end_global_epoch):\n        stage = apply_progressive_unfreezing(model, global_epoch)\n        model.train()\n\n        epoch_loss = 0.0\n        start = timer()\n\n        for n, batch in enumerate(loader):\n            optimizer.zero_grad(set_to_none=True)\n\n            with torch.amp.autocast('cuda', dtype=torch.float16):\n                output = model(batch)\n                loss = output['pixel_loss'] + output['pixel_dice_loss']\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n\n            epoch_loss += loss.item()\n\n            sid = batch['sample_id'][0] if isinstance(batch['sample_id'], list) else batch['sample_id']\n            print(\n                f'\\r GPU{gpu_id} global_epoch {global_epoch+1}/{TARGET_TOTAL_EPOCHS} '\n                f'step {n+1}/{len(loader)} loss {loss.item():.4f} id={sid}',\n                end='', flush=True\n            )\n\n            if (n + 1) % SAVE_EVERY_STEPS == 0:\n                save_full_checkpoint(\n                    last_ckpt,\n                    model, optimizer, scheduler, scaler,\n                    global_epoch=global_epoch,\n                    best_loss=best_loss,\n                    last_avg_loss=epoch_loss / (n + 1),\n                    completed_epoch=False,\n                    extra={'note': 'mid-epoch safety save'}\n                )\n                print(f'\\n GPU{gpu_id} emergency save -> {last_ckpt}')\n\n        avg = epoch_loss / len(loader)\n        elapsed = timer() - start\n        print(f'\\n GPU{gpu_id} global_epoch {global_epoch+1}/{TARGET_TOTAL_EPOCHS} avg_loss={avg:.4f} time={elapsed:.0f}s stage={stage}')\n\n        is_best = avg < best_loss\n        if is_best:\n            best_loss = avg\n\n        epoch_ckpt = f'{SAVE_DIR}/cross_attn_b6_ep{global_epoch+1}.pth'\n        save_full_checkpoint(\n            epoch_ckpt,\n            model, optimizer, scheduler, scaler,\n            global_epoch=global_epoch,\n            best_loss=best_loss,\n            last_avg_loss=avg,\n            completed_epoch=True,\n            extra={'stage': stage}\n        )\n\n        save_full_checkpoint(\n            last_ckpt,\n            model, optimizer, scheduler, scaler,\n            global_epoch=global_epoch,\n            best_loss=best_loss,\n            last_avg_loss=avg,\n            completed_epoch=True,\n            extra={'stage': stage}\n        )\n\n        print(f' GPU{gpu_id} saved epoch checkpoint -> {epoch_ckpt}')\n        print(f' GPU{gpu_id} updated last checkpoint -> {last_ckpt}')\n\n        if is_best:\n            save_full_checkpoint(\n                best_ckpt,\n                model, optimizer, scheduler, scaler,\n                global_epoch=global_epoch,\n                best_loss=best_loss,\n                last_avg_loss=avg,\n                completed_epoch=True,\n                extra={'stage': stage}\n            )\n            torch.save(model.state_dict(), final_out)\n            print(f' GPU{gpu_id} new best -> {best_ckpt}')\n            print(f' GPU{gpu_id} final model weights -> {final_out}')\n\n    if result_file:\n        with open(result_file, 'wb') as f:\n            pickle.dump({\n                'gpu': gpu_id,\n                'completed_until_global_epoch': end_global_epoch - 1,\n                'best_loss': best_loss,\n            }, f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.468848Z","iopub.execute_input":"2026-04-28T16:34:29.469126Z","iopub.status.idle":"2026-04-28T16:34:29.492540Z","shell.execute_reply.started":"2026-04-28T16:34:29.469102Z","shell.execute_reply":"2026-04-28T16:34:29.491801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob, os\n\nprint(\"working files:\")\nfor f in sorted(glob.glob('/kaggle/working/*')):\n    print(' ', f)\n\nprint(\"\\ncheckpoint files:\")\nfor f in sorted(glob.glob(f'{SAVE_DIR}/*')):\n    print(' ', f, os.path.getsize(f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:34:29.493656Z","iopub.execute_input":"2026-04-28T16:34:29.494049Z","iopub.status.idle":"2026-04-28T16:34:29.510485Z","shell.execute_reply.started":"2026-04-28T16:34:29.493995Z","shell.execute_reply":"2026-04-28T16:34:29.509876Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Step 8 — Single GPU launcher**\n","metadata":{}},{"cell_type":"code","source":"# Step 8 — Single GPU launcher (P100)\n# On entraîne UN seul modèle sur TOUT le dataset, sur cuda:0\n\nprint('Starting single-GPU training on cuda:0 ...')\ntrain_one_gpu(\n    gpu_id=0,\n    assigned_ids=valid_id,\n    prev_fail_ids=FAIL_ID,\n    result_file=f'{OUT_DIR}/train_result_gpu0.pkl',\n    resume_ckpt=RESUME_CKPT,\n)\nprint('Single-GPU training complete')\nprint(f'Checkpoints saved in {SAVE_DIR}')","metadata":{"execution":{"iopub.status.busy":"2026-04-28T16:34:29.511980Z","iopub.execute_input":"2026-04-28T16:34:29.512374Z","iopub.status.idle":"2026-04-28T16:35:13.561285Z","shell.execute_reply.started":"2026-04-28T16:34:29.512337Z","shell.execute_reply":"2026-04-28T16:35:13.560463Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Step 9 — After training: verify and save to output**\n","metadata":{}},{"cell_type":"code","source":"# Vérifie ce qui existe dans checkpoints\nimport glob\nprint(\"Checkpoints présents:\")\nfor f in sorted(glob.glob(f'{SAVE_DIR}/*.pth')):\n    print(' ', f)\n\nprint(\"\\nFichiers résultats:\")\nfor f in sorted(glob.glob(f'{OUT_DIR}/*.pkl')):\n    print(' ', f)\n    with open(f, 'rb') as fh:\n        print('  contenu:', pickle.load(fh))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:35:13.563508Z","iopub.execute_input":"2026-04-28T16:35:13.563779Z","iopub.status.idle":"2026-04-28T16:35:13.571277Z","shell.execute_reply.started":"2026-04-28T16:35:13.563754Z","shell.execute_reply":"2026-04-28T16:35:13.570473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Diagnostic : vérifie combien de GPUs sont disponibles ──\nimport torch\nprint(f\"GPUs disponibles: {torch.cuda.device_count()}\")\nfor i in range(torch.cuda.device_count()):\n    print(f\"  GPU{i}: {torch.cuda.get_device_name(i)} | \"\n          f\"VRAM: {torch.cuda.get_device_properties(i).total_memory / 1e9:.1f} GB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:35:13.572484Z","iopub.execute_input":"2026-04-28T16:35:13.572931Z","iopub.status.idle":"2026-04-28T16:35:13.585375Z","shell.execute_reply.started":"2026-04-28T16:35:13.572893Z","shell.execute_reply":"2026-04-28T16:35:13.584491Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 9 — After training: verify and save final checkpoint\n\nfinal_ckpt = f'{SAVE_DIR}/cross_attn_b6_best.pth'\nassert os.path.exists(final_ckpt), f\"Checkpoint introuvable: {final_ckpt}\"\n\n# Rebuild model exactly like in training\nmodel = LeadModel(\n    encoder_name='tu-efficientnet_b6',\n    encoder_weights=None,\n    fusion_type='cross_attn',\n    fusion_levels=[3, 4],\n)\n\nstate = torch.load(final_ckpt, map_location='cpu')\nmissing, unexpected = model.load_state_dict(state, strict=False)\n\nprint('Checkpoint reload ok')\nprint('missing keys:', len(missing))\nprint('unexpected keys:', len(unexpected))\n\nmodel.eval()\nprint('Model ready to use in inference notebook')\n\n# Copy to /kaggle/working so it appears in notebook output\nimport shutil\nfinal_out = '/kaggle/working/cross_attn_b6_final.pth'\nshutil.copy(final_ckpt, final_out)\n\nprint(f'Final checkpoint copied to: {final_out}')","metadata":{"execution":{"execution_failed":"2026-04-26T22:44:46.558Z"}}},{"cell_type":"code","source":"# Step 9 — Verify saved artifacts\n\nimport glob, os, pickle, torch\n\nprint(\"Checkpoint files:\")\nfor f in sorted(glob.glob(f'{SAVE_DIR}/*.pth')):\n    print(' ', os.path.basename(f), os.path.getsize(f))\n\nprint(\"\\nResult files:\")\nfor f in sorted(glob.glob(f'{OUT_DIR}/*.pkl')):\n    print(' ', os.path.basename(f))\n    with open(f, 'rb') as fh:\n        print('   ', pickle.load(fh))\n\nresume_ckpt = f'{SAVE_DIR}/cross_attn_b6_last_full.pth'\nbest_full   = f'{SAVE_DIR}/cross_attn_b6_best_full.pth'\nfinal_weights = '/kaggle/working/cross_attn_b6_final.pth'\n\nprint(\"\\nExists:\")\nprint(\" last_full   :\", os.path.exists(resume_ckpt))\nprint(\" best_full   :\", os.path.exists(best_full))\nprint(\" final_wts   :\", os.path.exists(final_weights))\n\nif os.path.exists(resume_ckpt):\n    ckpt = torch.load(resume_ckpt, map_location='cpu')\n    print(\"\\nResume checkpoint metadata:\")\n    print(\" global_epoch    :\", ckpt.get('global_epoch'))\n    print(\" completed_epoch :\", ckpt.get('completed_epoch'))\n    print(\" best_loss       :\", ckpt.get('best_loss'))\n    print(\" last_avg_loss   :\", ckpt.get('last_avg_loss'))\n    print(\" extra           :\", ckpt.get('extra', {}))\n\nif os.path.exists(final_weights):\n    model = LeadModel(\n        encoder_name='tu-efficientnet_b6',\n        encoder_weights=None,\n        fusion_type='cross_attn',\n        fusion_levels=[3, 4],\n    )\n    state = torch.load(final_weights, map_location='cpu')\n    missing, unexpected = model.load_state_dict(state, strict=False)\n    print(\"\\nFinal weights reload ok\")\n    print(\"missing keys:\", len(missing))\n    print(\"unexpected keys:\", len(unexpected))\n    model.eval()\n    print(\"Model ready for inference notebook\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:35:13.586487Z","iopub.execute_input":"2026-04-28T16:35:13.586846Z","iopub.status.idle":"2026-04-28T16:35:15.149460Z","shell.execute_reply.started":"2026-04-28T16:35:13.586820Z","shell.execute_reply":"2026-04-28T16:35:15.148600Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nprint(\"working files:\")\nfor f in sorted(glob.glob('/kaggle/working/*')):\n    print(\" \", f)\n\nprint(\"\\ncheckpoint files:\")\nfor f in sorted(glob.glob('/kaggle/working/checkpoints/*')):\n    print(\" \", f, os.path.getsize(f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T16:35:15.150586Z","iopub.execute_input":"2026-04-28T16:35:15.151043Z","iopub.status.idle":"2026-04-28T16:35:15.157468Z","shell.execute_reply.started":"2026-04-28T16:35:15.151011Z","shell.execute_reply":"2026-04-28T16:35:15.156735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport os\n\npaths = [\n    '/kaggle/working/checkpoints/cross_attn_b6_best.pth',\n    '/kaggle/working/checkpoints/cross_attn_b6_gpu0_ep1.pth',\n    '/kaggle/working/checkpoints/cross_attn_b6_gpu0_ep2.pth',\n    '/kaggle/working/cross_attn_b6_final.pth',\n]\n\nfor p in paths:\n    print(\"\\n\", p)\n    print(\"exists:\", os.path.exists(p))\n    if os.path.exists(p):\n        try:\n            state = torch.load(p, map_location='cpu')\n            print(\"load ok; type:\", type(state))\n            if isinstance(state, dict):\n                print(\"num keys:\", len(state))\n                print(\"first keys:\", list(state.keys())[:10])\n        except Exception as e:\n            print(\"load failed:\", e)","metadata":{"execution":{"iopub.status.busy":"2026-04-28T16:35:15.158501Z","iopub.execute_input":"2026-04-28T16:35:15.158796Z","iopub.status.idle":"2026-04-28T16:35:15.426481Z","shell.execute_reply.started":"2026-04-28T16:35:15.158770Z","shell.execute_reply":"2026-04-28T16:35:15.425492Z"},"trusted":true},"outputs":[],"execution_count":null}]}