{"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":[{"sourceType":"competition","sourceId":97984,"databundleVersionId":14096757},{"sourceType":"datasetVersion","sourceId":13731160,"datasetId":8733970,"databundleVersionId":14479231},{"sourceType":"datasetVersion","sourceId":15621323,"datasetId":9997958,"databundleVersionId":16555629},{"sourceType":"datasetVersion","sourceId":14695416,"datasetId":9387663,"databundleVersionId":15539655}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Reference\n- https://www.kaggle.com/code/hengck23/demo-submission","metadata":{}},{"cell_type":"code","source":"!ls /kaggle/input/hengck23-demo-submit-physionet/setup\n!pip install connected-components-3d --no-index --find-links=file:///kaggle/input/hengck23-demo-submit-physionet/setup/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:31.593428Z","iopub.execute_input":"2026-04-20T14:09:31.594236Z","iopub.status.idle":"2026-04-20T14:09:35.058359Z","shell.execute_reply.started":"2026-04-20T14:09:31.594197Z","shell.execute_reply":"2026-04-20T14:09:35.057605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cc3d\nimport cv2\nimport pandas as pd\nimport numpy as np\nfrom scipy import signal\nimport torch\nimport matplotlib.pyplot as plt\nimport matplotlib\n#matplotlib.use('TkAgg')\nimport shutil\nimport copy\nimport multiprocessing as mp\nimport pickle\nimport os\nimport sys\nimport multiprocessing as mp\n#mp.set_start_method('spawn', force=True)\nsys.path.insert(0, '/kaggle/input/my-stage2-lead-model') \nsys.path.append('/kaggle/input/hengck23-demo-submit-physionet')\nsys.path.append('/kaggle/input/physionet-final-submission-models')\nfrom timeit import default_timer as timer\nimport pickle\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n\nfrom stage2_smp_model import Net as WholeModel\nfrom stage2_lead_model import Net as LeadModel\nfrom stage2_common import *\nfrom stage2_model import prob_to_series_by_max\nprint('import ok!!!')\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.059783Z","iopub.execute_input":"2026-04-20T14:09:35.060289Z","iopub.status.idle":"2026-04-20T14:09:35.067432Z","shell.execute_reply.started":"2026-04-20T14:09:35.060262Z","shell.execute_reply":"2026-04-20T14:09:35.066587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODE   = 'local'  # submit  local fake\nDEVICE = 'cuda'\nFLOAT_TYPE = torch.float16 #torch.bfloat16\nFAIL_ID = []\n# training config\nEPOCHS        = 8\nLR            = 1e-4\nBATCH_SIZE    = 1\nFREEZE_EPOCHS = 2\nWINDOW_SIZE = 240\nOFFSET      = 416\n\nSAVE_DIR      = '/kaggle/working/checkpoints'\nos.makedirs(SAVE_DIR, exist_ok=True)\nKAGGLE_DIR = \\\n\t'/kaggle/input/physionet-ecg-image-digitization'\nWEIGHT_DIR = \\\n\t'/kaggle/input/hengck23-demo-submit-physionet/weight'\nOUT_DIR = '/kaggle/working/output' \nos.makedirs(f'{OUT_DIR}/rectified', exist_ok=True)\nos.makedirs(f'{OUT_DIR}/masks',     exist_ok=True)\ndef make_test_fake_df(): \n    valid_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n    valid_df.loc[:,'id']=valid_df['id'].astype(str) \n    fake_test_df=[]\n    for i,d in valid_df.iterrows():\n        #if i==4: break\n        image_id = d['id']\n    \n        truth_df = pd.read_csv(f'{KAGGLE_DIR}/train/{image_id}/{image_id}.csv')\n        non_nan_count = truth_df.count()\n        #print(i,image_id,non_nan_count)\n        #print(non_nan_count.index)\n    \n        #lead\tfs\tnumber_of_rows \n        this_df = pd.DataFrame({\n            'id':image_id ,\n            'lead':non_nan_count.index,\n            'fs': d['fs'],\n            'number_of_rows':non_nan_count.values \n        })\n        fake_test_df.append(this_df)\n        if i==0: print(this_df)\n    fake_test_df = pd.concat(fake_test_df)\n    return fake_test_df\n\n\n# set valid/test data\nif MODE == 'local':\n\tfrom sample_list import ERROR_ID\n\tvalid_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n\tvalid_df['id']=valid_df['id'].astype(str)\n    \n\tvalid_id = [\n\t\t#f'{image_id}-{type_id}' for image_id in ERROR_ID\n\t\tf'{image_id}-{type_id}' for image_id in valid_df['id'].values\n\t\tfor type_id in ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']\n\t]\n\t#valid_id = [\n      #  '11842146-0012','144746082-0009','225208096-0006', '2289894144-0012','1617515072-0006',\n     #   '2289894144-0010','2566168201-0009', '2659677149-0011'\n   # ]\n    \nif MODE == 'submit':\n\tvalid_df = pd.read_csv(f'{KAGGLE_DIR}/test.csv')\n\tvalid_df['id']=valid_df['id'].astype(str) \n\tvalid_id = valid_df['id'].unique().tolist()\n\nif MODE == 'fake':\n\tvalid_df = make_test_fake_df()\n\tvalid_df['id']=valid_df['id'].astype(str) \n\tvalid_id = valid_df['id'].unique().tolist()\n\n#--------------------------------------\n\ndef read_image(sample_id):\n    if MODE == 'local':\n        image_id, type_id = sample_id.split('-')\n        image = cv2.imread(f'{KAGGLE_DIR}/train/{image_id}/{image_id}-{type_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n    if MODE == 'submit':\n        image_id = sample_id\n        image = cv2.imread(f'{KAGGLE_DIR}/test/{image_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n    if MODE == 'fake':\n        image_id = sample_id \n        type_id = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012'][\n            int(image_id)%9\n        ] \n        image = cv2.imread(f'{KAGGLE_DIR}/train/{image_id}/{image_id}-{type_id}.png', cv2.IMREAD_COLOR_RGB)\n        return image\n\ndef read_sampling_length(sample_id):\n\tif MODE == 'local':\n\t\timage_id, type_id = sample_id.split('-')\n\t\td = valid_df[valid_df['id']==image_id].iloc[0]\n\t\tlength = d.sig_len\n\t\treturn length\n\tif MODE == 'submit':\n\t\timage_id = sample_id\n\t\td = valid_df[\n\t\t\t(valid_df['id']==image_id) & (valid_df['lead']=='II')\n\t\t].iloc[0]\n\t\tlength = d.number_of_rows\n\t\treturn length\n\tif MODE == 'fake':\n\t\timage_id = sample_id\n\t\td = valid_df[\n\t\t\t(valid_df['id']==image_id) & (valid_df['lead']=='II')\n\t\t].iloc[0]\n\t\tlength = d.number_of_rows\n\t\treturn length\n\n#valid_id = valid_id[:300]\nprint('valid_id:', len(valid_id))\nprint('\\t', valid_id[:3], '...')\nprint('setting ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.068353Z","iopub.execute_input":"2026-04-20T14:09:35.068558Z","iopub.status.idle":"2026-04-20T14:09:35.095013Z","shell.execute_reply.started":"2026-04-20T14:09:35.068542Z","shell.execute_reply":"2026-04-20T14:09:35.094213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── ADD: sparse COO mask save/load (replaces dense .npy) ──\ndef save_sparse_mask_coo(mask_dense, save_path):\n    \"\"\"\n    mask_dense: (4, H, W) float32  — the dense prob map from the conv2d model\n    Saves only non-zero entries in COO format → tiny file\n    \"\"\"\n    shape = np.array(mask_dense.shape)\n    arrays = {'shape': shape}\n    for i in range(mask_dense.shape[0]):\n        ys, xs = np.where(mask_dense[i] > 0.3)   # threshold: keep confident pixels only\n        vs = mask_dense[i, ys, xs]\n        arrays[f'ch{i}_y'] = ys.astype(np.int32)\n        arrays[f'ch{i}_x'] = xs.astype(np.int32)\n        arrays[f'ch{i}_v'] = vs.astype(np.float32)\n    np.savez_compressed(save_path, **arrays)\n\ndef load_sparse_mask_coo(filepath):\n    \"\"\"Returns (4, H, W) float32 — identical interface to the old np.load\"\"\"\n    data = np.load(filepath)\n    shape = tuple(data['shape'])\n    mask = np.zeros(shape, dtype=np.float32)\n    for i in range(shape[0]):\n        y = data[f'ch{i}_y']\n        x = data[f'ch{i}_x']\n        v = data[f'ch{i}_v']\n        mask[i, y, x] = v\n    return mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.096602Z","iopub.execute_input":"2026-04-20T14:09:35.096886Z","iopub.status.idle":"2026-04-20T14:09:35.103185Z","shell.execute_reply.started":"2026-04-20T14:09:35.096871Z","shell.execute_reply":"2026-04-20T14:09:35.102213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# output_to_predict (different from stage0's version) , rectify_image,draw_mapping,draw_results_stage1\nos.makedirs(f'{OUT_DIR}/rectified', exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.104122Z","iopub.execute_input":"2026-04-20T14:09:35.104366Z","iopub.status.idle":"2026-04-20T14:09:35.118632Z","shell.execute_reply.started":"2026-04-20T14:09:35.104346Z","shell.execute_reply":"2026-04-20T14:09:35.117869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stage 2 only for pseudo-mask generation\n#print('*** STARTING STAGE2 MASK GENERATION ***')\n\nos.makedirs(f'{OUT_DIR}/masks', exist_ok=True)\n\n\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\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\n    trim_image = image.copy()[OFFSET:y1, x0:x1]\n    image = image[y0:y1, x0:x1]\n    H, W, _ = image.shape\n\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\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\n    lead_images = np.stack(lead_images)  # (4, H, W, 3)\n    return trim_image, lead_images\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.119353Z","iopub.execute_input":"2026-04-20T14:09:35.119854Z","iopub.status.idle":"2026-04-20T14:09:35.131449Z","shell.execute_reply.started":"2026-04-20T14:09:35.119829Z","shell.execute_reply":"2026-04-20T14:09:35.130840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('*** STARTING SINGLE PIPELINE LOOP: Stage0 → Stage1 → Masks ***')\n\nfrom stage0_common import time_to_str\nfrom stage0_model import Net as Stage0Net\nfrom stage1_model import Net as Stage1Net\nimport stage0_common as s0c\nimport stage1_common as s1c\n# Load all three models once before the loop\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, f'{WEIGHT_DIR}/stage0-last.checkpoint.pth')\nstage0_net.to(DEVICE).eval()\n\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(stage1_net, f'{WEIGHT_DIR}/stage1-last.checkpoint.pth')\nstage1_net.to(DEVICE).eval()\n\nmask_model = LeadModel(\n    encoder_name='tu-timm/tf_efficientnet_b6.ns_jft_in1k',\n    encoder_weights=None,\n    fusion_type='shared_conv2d',\n)\nstate = torch.load(\n    '/kaggle/input/physionet-final-submission-models/series_b6_shared_conv2d_lb23.10.pth',\n    map_location='cpu'\n)\nmask_model.load_state_dict(state, strict=False)\nmask_model.to(DEVICE).eval()\nmask_model.output_type = ['infer']\n\nstart_timer = timer()\nfor n, sample_id in enumerate(valid_id):\n    timestamp = time_to_str(timer() - start_timer, 'sec')\n    print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n\n    # ── Stage 0 ──\n    try:\n        image = read_image(sample_id)\n        batch = s0c.image_to_batch(image)\n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output0 = stage0_net(batch)\n        rotated, keypoint = s0c.output_to_predict(image, batch, output0)\n        normalised, keypoint, homo = s0c.normalise_by_homography(rotated, keypoint)\n        # normalised stays in RAM — NOT saved to disk\n    except Exception as e:\n        print(f'\\nStage0 failed {sample_id}: {e}')\n        FAIL_ID.append(sample_id)\n        torch.cuda.empty_cache()\n        continue\n\n    # ── Stage 1 ──\n    try:\n        batch1 = {\n            'image': torch.from_numpy(\n                np.ascontiguousarray(normalised.transpose(2, 0, 1))\n            ).unsqueeze(0)\n        }\n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output1 = stage1_net(batch1)\n        gridpoint_xy, more = s1c.output_to_predict(normalised, batch1, output1)\n        rectified = s1c.rectify_image(normalised, gridpoint_xy)\n        # Save only rectified — this is what training needs\n        cv2.imwrite(\n            f'{OUT_DIR}/rectified/{sample_id}.rect.jpg',\n            cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR),\n            [int(cv2.IMWRITE_JPEG_QUALITY), 70]\n        )\n    except Exception as e:\n        print(f'\\nStage1 failed {sample_id}: {e}')\n        FAIL_ID.append(sample_id)\n        torch.cuda.empty_cache()\n        continue\n\n    # ── Stage 2 mask (pseudo-label) ──\n    try:\n        _, lead_images = read_images(f'{OUT_DIR}/rectified/{sample_id}.rect.jpg')\n        lead_tensor = torch.from_numpy(\n            lead_images.transpose(0, 3, 1, 2)\n        ).contiguous().unsqueeze(0).to(DEVICE)   # (1, 4, 3, H, W)\n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output2 = mask_model({'image': lead_tensor})\n        mask_dense = output2['pixel'].squeeze(0).squeeze(1).float().cpu().numpy()  # (4, H, W)\n        save_sparse_mask_coo(mask_dense, f'{OUT_DIR}/masks/{sample_id}.mask-coo.npz')\n        \n    except Exception as e:\n        print(f'\\nMask failed {sample_id}: {e}')\n        # Don't add to FAIL_ID — rectified image was saved, mask can be retried\n\n    torch.cuda.empty_cache()\n\nprint('')\nprint('Pipeline done. FAIL_ID:', FAIL_ID)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-20T14:09:35.132357Z","iopub.execute_input":"2026-04-20T14:09:35.132581Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# training","metadata":{}},{"cell_type":"code","source":"\n    # ---- Actual training Dataset ----\nos.makedirs(f'{OUT_DIR}/masks', exist_ok=True)\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    def __getitem__(self, idx):\n        sample_id = self.ids[idx]\n\n        _, lead_images = read_images(\n            f'{OUT_DIR}/rectified/{sample_id}.rect.jpg'\n        )\n\n        image_tensor = torch.from_numpy(\n            lead_images.transpose(0, 3, 1, 2)\n        ).byte()\n\n        mask = load_sparse_mask_coo(\n            f'{OUT_DIR}/masks/{sample_id}.mask-coo.npz'\n        )\n        mask_tensor = torch.from_numpy(mask).float()\n\n        return {\n            'image': image_tensor,\n            'pixel': mask_tensor,\n            'sample_id': sample_id,\n        }\n        \nprint('dataset class ok')","metadata":{"execution":{"iopub.status.busy":"2026-04-20T14:03:01.395817Z","iopub.status.idle":"2026-04-20T14:03:01.396137Z","shell.execute_reply.started":"2026-04-20T14:03:01.395979Z","shell.execute_reply":"2026-04-20T14:03:01.395992Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_gpu(gpu_id=0, assigned_ids=None, prev_fail_ids=None,\n                  result_file=None):\n    device = f'cuda:{gpu_id}'\n    if assigned_ids  is None: assigned_ids  = valid_id\n    if prev_fail_ids is None: prev_fail_ids = FAIL_ID\n\n    # ── Build model with YOUR new cross_attn fusion ──\n    model = LeadModel(\n        encoder_name='tu-timm-efficientnet_b6.ns_jft_in1k',\n        encoder_weights=None,\n        fusion_type='cross_attn',\n        fusion_levels=[3, 4],      # deep levels only — safe for memory\n    )\n\n    # ── Load conv2d checkpoint as starting point ──\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 (new attn layers): {len(missing)}')\n    print(f'GPU{gpu_id} | unexpected keys (old fusion): {len(unexpected)}')\n    # missing = your new CrossLeadAttentionFusion params (random init — expected)\n    # unexpected = old CrossLeadFusion params (discarded — expected)\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(dataset, batch_size=BATCH_SIZE, shuffle=True,\n                         num_workers=2, pin_memory=True)\n\n    optimizer = AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS * len(loader),\n                                  eta_min=LR * 0.01)\n\n    scaler = torch.cuda.amp.GradScaler()\n\n    for epoch in range(EPOCHS):\n\n        # ── Freeze encoder for first FREEZE_EPOCHS epochs ──\n        # Only the new attention fusion layers train freely at first\n        freeze = (epoch < FREEZE_EPOCHS)\n        for p in model.encoder.parameters():\n            p.requires_grad = not freeze\n        for p in model.decoder.parameters():\n            p.requires_grad = not freeze\n        # Fusion modules always train\n        for p in model.fusion_modules.parameters():\n            p.requires_grad = True\n\n        model.train()\n        epoch_loss = 0.0\n        start = timer()\n\n        for n, batch in enumerate(loader):\n            with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n                output = model(batch)\n                loss = output['pixel_loss'] + output['pixel_dice_loss']\n\n            optimizer.zero_grad()\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(f'\\r GPU{gpu_id} epoch {epoch+1}/{EPOCHS} '\n                  f'step {n+1}/{len(loader)} '\n                  f'loss {loss.item():.4f}',\n                  end='', flush=True)\n\n        avg = epoch_loss / len(loader)\n        elapsed = timer() - start\n        print(f'\\n GPU{gpu_id} epoch {epoch+1} avg_loss={avg:.4f} '\n              f'time={elapsed:.0f}s  encoder_frozen={freeze}')\n\n        # Save checkpoint every epoch\n        ckpt_out = f'{SAVE_DIR}/cross_attn_b6_gpu{gpu_id}_ep{epoch+1}.pth'\n        torch.save(model.state_dict(), ckpt_out)\n        print(f' GPU{gpu_id} saved → {ckpt_out}')\n\n    if result_file:\n        with open(result_file, 'wb') as f:\n            pickle.dump({'gpu': gpu_id, 'final_loss': avg}, f)\n\nprint('train_one_gpu defined ok')","metadata":{"execution":{"iopub.status.busy":"2026-04-20T14:03:01.397236Z","iopub.status.idle":"2026-04-20T14:03:01.397444Z","shell.execute_reply.started":"2026-04-20T14:03:01.397346Z","shell.execute_reply":"2026-04-20T14:03:01.397354Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_parallel():\n    mid = len(valid_id) // 2\n    ids_gpu0 = valid_id[:mid]\n    ids_gpu1 = valid_id[mid:]\n\n    res0 = f'{OUT_DIR}/train_result_gpu0.pkl'\n    res1 = f'{OUT_DIR}/train_result_gpu1.pkl'\n\n    p0 = mp.Process(target=train_one_gpu, args=(0, ids_gpu0, FAIL_ID, res0))\n    p1 = mp.Process(target=train_one_gpu, args=(1, ids_gpu1, FAIL_ID, res1))\n\n    p0.start(); p1.start()\n    p0.join();  p1.join()\n\n    print('Training complete')\n    print(f'Checkpoints saved in {SAVE_DIR}')\n\ntrain_parallel()","metadata":{"execution":{"iopub.status.busy":"2026-04-20T14:03:01.398731Z","iopub.status.idle":"2026-04-20T14:03:01.398976Z","shell.execute_reply.started":"2026-04-20T14:03:01.398851Z","shell.execute_reply":"2026-04-20T14:03:01.398860Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best checkpoint and verify forward pass\nmodel = LeadModel(\n    encoder_name='tu-timm-efficientnet_b6.ns_jft_in1k',\n    encoder_weights=None,\n    fusion_type='cross_attn',\n    fusion_levels=[3, 4],\n)\nstate = torch.load(f'{SAVE_DIR}/cross_attn_b6_gpu0_ep{EPOCHS}.pth',\n                   map_location='cpu')\nmodel.load_state_dict(state)\nmodel.eval()\nprint('Checkpoint loads ok — ready to use in inference notebook')\n\n# Copy to /kaggle/working so it appears as a notebook output\nimport shutil\nshutil.copy(\n    f'{SAVE_DIR}/cross_attn_b6_gpu0_ep{EPOCHS}.pth',\n    '/kaggle/working/cross_attn_b6_final.pth'\n)","metadata":{"execution":{"iopub.status.busy":"2026-04-20T14:03:01.399965Z","iopub.status.idle":"2026-04-20T14:03:01.400268Z","shell.execute_reply.started":"2026-04-20T14:03:01.400107Z","shell.execute_reply":"2026-04-20T14:03:01.400120Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage1","metadata":{}},{"cell_type":"markdown","source":"# Stage2","metadata":{}}]}