{"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":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":14560686,"sourceType":"datasetVersion","datasetId":9026412}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"try:\n    import cc3d\nexcept:\n    #https://pypi.org/project/connected-components-3d/\n    #!pip install connected-components-3d\n\n    !ls /kaggle/input/physionet-data/hengck23-demo-submit-physionet/setup\n    !pip install connected-components-3d --no-index --find-links=file:///kaggle/input/physionet-data/hengck23-demo-submit-physionet/setup\n\nimport cc3d\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport matplotlib\n#matplotlib.use('TkAgg')\nimport shutil\nfrom scipy.signal import resample\nfrom scipy.interpolate import PchipInterpolator\n\n\nimport sys, os\nfrom timeit import default_timer as timer\nsys.path.append('/kaggle/input/physionet-data/hengck23-demo-submit-physionet')\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:19:23.763087Z","iopub.execute_input":"2026-01-22T00:19:23.763325Z","iopub.status.idle":"2026-01-22T00:19:38.378404Z","shell.execute_reply.started":"2026-01-22T00:19:23.76329Z","shell.execute_reply":"2026-01-22T00:19:38.377579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IS_KAGGLE = True\nFOLD = 0\nMODE   = 'submit'  # submit  local fake\nDEVICE = 'cuda'\nFLOAT_TYPE = torch.float16 #torch.bfloat16\nFAIL_ID = []\n\nif IS_KAGGLE:\n    sys.path.append('/kaggle/input/hengck23-demo-submit-physionet')\n    KAGGLE_DIR = \\\n        '/kaggle/input/physionet-ecg-image-digitization'\n    WEIGHT_DIR = \\\n        '/kaggle/input/physionet-data/hengck23-demo-submit-physionet/weight'\n    OUT_DIR = \\\n        f'/kaggle/working/output-{MODE}'\nelse:\n    sys.path.append('data/hengck23-demo-submit-physionet')\n    KAGGLE_DIR = \\\n        'data'\n    WEIGHT_DIR = \\\n        'data/hengck23-demo-submit-physionet/weight'\n    OUT_DIR = \\\n        f'kaggle/working/output-{MODE}'\n\n\nclass CFG:\n    CSV_PATH = 'data/train.csv'\n    DATA_DIR = 'data/train'\n    STAGE0_DIR = f'{OUT_DIR}/normalised'\n    STAGE1_DIR = f'{OUT_DIR}/rectified'\n\n    OUT_IMG_DIR = f'{OUT_DIR}/rectified_resized'\n    OUT_SCALE_DIR = os.path.join(OUT_IMG_DIR, 'scales')\n\n    BASE_WIDTH = 1440\n    BASE_HEIGHT = 1152\n    BASE_RECT_W = 2200\n    BASE_RECT_H = 1700\n\n    FIXED_RECT_WIDTH = 4400\n    FIXED_RECT_HEIGHT = 1700 #1700 or 3400\n\n\nif MODE == 'local':\n    import os\n    from glob import glob\n    # from sample_list import ERROR_ID\n    valid_df = pd.read_csv(f'/kaggle/input/physionet-data/train_folds.csv')\n    valid_df = valid_df[valid_df[\"fold\"]==FOLD]\n    valid_df['id']=valid_df['id'].astype(str)\n    type_ids = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']\n\n    valid_id = []\n    for sample_id in valid_df[\"id\"].astype(str).tolist():  # fold==0 的 sample_id\n        for tid in type_ids:\n            png_path = os.path.join(\"data/train_flat\", f\"{sample_id}-{tid}.png\")\n            if os.path.exists(png_path):\n                valid_id.append(f\"{sample_id}-{tid}\")\n\n    valid_id = valid_id[:10]\n    \nif MODE == 'submit':\n    valid_df = pd.read_csv(f'{KAGGLE_DIR}/test.csv')\n    valid_df['id']=valid_df['id'].astype(str) \n    valid_id = valid_df['id'].unique().tolist()\n\n\n#--------------------------------------\n\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t/60)\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min)\n\n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\n\telse:\n\t\traise NotImplementedError\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\n\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-01-22T00:19:38.380326Z","iopub.execute_input":"2026-01-22T00:19:38.380646Z","iopub.status.idle":"2026-01-22T00:19:38.426077Z","shell.execute_reply.started":"2026-01-22T00:19:38.380619Z","shell.execute_reply":"2026-01-22T00:19:38.425396Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage0\nprint('*** STARTING STAGE0 ***')\n\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import * #load_net, image_to_batch\n\nos.makedirs(f'{OUT_DIR}/normalised', exist_ok=True)\n\ndef normalise_by_homography(image, keypoint, ref_pt9):\n    pt9 = [[k[0], k[1]] for k in keypoint]\n    normalised, homo, match = normalise_image(image, pt9, ref_pt9)\n    for i in range(len(keypoint)):\n        keypoint[i].append(match[i])\n    return normalised, keypoint, homo\n\nlead_name_to_label = {\n    'None': 0,\n    'I': 1,\n    'aVR': 2,\n    'V1': 3,\n    'V4': 4,\n    'II': 5,\n    'aVL': 6,\n    'V2': 7,\n    'V5': 8,\n    'III': 9,\n    'aVF': 10,\n    'V3': 11,\n    'V6': 12,\n    'II-rhythm': 13,\n}\n\nlabel_to_lead_name = {v: k for k, v in lead_name_to_label.items()}\n\n\ndef marker_to_keypoint(image, orientation, marker, scale, label_to_lead_name):\n    orientation = orientation.data.cpu().numpy().reshape(-1)\n    marker = marker.permute(0, 2, 3, 1).float().data.cpu().numpy()[0]\n\n    k = orientation.argmax()\n    if k != 0:\n        if k <= 3:\n            k = -k\n        else:\n            print(f'k={k} rotation unknown????')\n\n    marker = np.rot90(marker, k, axes=(0, 1))\n    keypoint = []\n    thresh = marker.argmax(-1)\n    for label in [2, 3, 4, 6, 7, 8, 10, 11, 12]:\n        cc = cc3d.connected_components(thresh == label)\n        stats = cc3d.statistics(cc)\n        center = stats['centroids'][1:]\n        area = stats['voxel_counts'][1:]\n        argsort = np.argsort(area)[::-1]\n        center = center[argsort]\n        area = area[argsort]\n\n        center = np.append(center, [[0, 0]], axis=0)\n        area = np.append(area, [1], axis=0)\n\n        for (y, x), a in zip(center[:1], area[:1]):\n            leadname = label_to_lead_name[label]\n            x, y = x / scale, y / scale\n            keypoint.append([x, y, label, leadname])\n\n    return keypoint, k\n    \ndef output_to_predict(image, batch, output, label_to_lead_name):\n    marker = 0\n    orientation = 0\n\n    num_tta = len(batch['tta'])\n    sH, sW = batch['sH'], batch['sW']\n    scale = batch['scale']\n\n    for b in range(num_tta):\n        tta = batch['tta'][b]\n\n        mk = output['marker'][[b]]\n        on = output['orientation'][b]\n\n        if tta == 1:\n            mk = torch.flip(mk, [2]).contiguous()\n            on = on[[4, 5, 6, 7, 0, 1, 2, 3]]\n        elif tta == 2:\n            mk = torch.flip(mk, [3]).contiguous()\n            on = on[[6, 7, 4, 5, 2, 3, 0, 1]]\n        elif tta == 3:\n            mk = torch.flip(mk, [2, 3]).contiguous()\n            on = on[[2, 3, 0, 1, 6, 7, 4, 5]]\n        else:\n            pass\n\n        orientation += on\n        marker += mk[..., :sH, :sW]\n\n    marker = marker / num_tta\n    orientation = orientation / num_tta\n    keypoint, k = marker_to_keypoint(image, orientation, marker, scale, label_to_lead_name)\n    rotated = np.ascontiguousarray(np.rot90(image, k, axes=(0, 1)))\n    print(\"rotated:!!!!!!!\",k)\n    return rotated, keypoint, k\n\ndef normalise_by_homography(image, keypoint):\n\t#[ k[-1] for k in keypoint]\n\tpt9 = [[ k[0],k[1]] for k in keypoint]\n\tnormalised, homo, match = normalise_image(image, pt9)\n\tfor i in range(len(keypoint)):\n\t\tkeypoint[i].append(match[i])\n\treturn normalised, keypoint, homo\n\ndef run_stage0():\n    stage0_net = Stage0Net(pretrained=False)\n    stage0_net = load_net(stage0_net, f'{WEIGHT_DIR}/stage0-last.checkpoint.pth')\n    stage0_net.to(DEVICE)\n\n    start_timer = timer()\n    for 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        image = read_image(sample_id)\n        batch = image_to_batch(image)\n\n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage0_net(batch)\n\n                try:\n                    rotated, keypoint, k = output_to_predict(image, batch, output, label_to_lead_name)\n                    # print(\"rotated:\",rotated)\n                    normalised, keypoint, homo = normalise_by_homography(rotated, keypoint)\n                    # ---\n                    cv2.imwrite(f'{OUT_DIR}/normalised/{sample_id}.norm.png', cv2.cvtColor(normalised, cv2.COLOR_RGB2BGR))\n                    np.save(f'{OUT_DIR}/normalised/{sample_id}.homo.npy', homo)\n                    np.save(f'{OUT_DIR}/normalised/{sample_id}.rotation.npy', np.array([k]))\n                except:\n                    print(\"sample_id:\",sample_id)\n                    FAIL_ID.append(sample_id)\n\n        torch.cuda.empty_cache()\n        if 0: # n<10: # optional: show results\n            overlay = draw_results_stage0(rotated, keypoint)\n            print('')\n            print('demo results for stage0--------------')\n            print(sample_id)\n            plt.imshow(image);plt.show()\n            plt.imshow(overlay);plt.show()\n            plt.imshow(normalised);plt.show()\n            \n    print('')\n\nrun_stage0()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage0() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:19:38.42668Z","iopub.execute_input":"2026-01-22T00:19:38.426863Z","iopub.status.idle":"2026-01-22T00:19:52.430954Z","shell.execute_reply.started":"2026-01-22T00:19:38.426848Z","shell.execute_reply":"2026-01-22T00:19:52.430294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage1\nprint('*** STARTING STAGE1 ***')\n\nfrom stage1_model import Net as Stage1Net\nfrom stage1_common import *\n\nos.makedirs(f'{OUT_DIR}/rectified', exist_ok=True)\n\ndef run_stage1():\n\tstage1_net = Stage1Net(pretrained=False)\n\tstage1_net = load_net(stage1_net, f'{WEIGHT_DIR}/stage1-last.checkpoint.pth')\n\tstage1_net.to(DEVICE)\n\n\tstart_timer = timer()\n\tfor n, sample_id in enumerate(valid_id):\n\t\ttimestamp = time_to_str(timer() - start_timer, 'sec')\n\t\tprint(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n\t\tif sample_id in FAIL_ID: continue\n\n\t\timage = cv2.imread(f'{OUT_DIR}/normalised/{sample_id}.norm.png', cv2.IMREAD_COLOR_RGB)\n\t\tbatch = {\n\t\t\t'image': torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0),\n\t\t}\n\t\tnum_tta = 1\n\n\t\twith torch.amp.autocast('cuda', dtype=FLOAT_TYPE): #torch.bfloat16\n\t\t\twith torch.no_grad():\n\t\t\t\toutput = stage1_net(batch)\n\n\t\t\t\ttry:\n\t\t\t\t\tgridpoint_xy, more = output_to_predict(image, batch, output)\n\t\t\t\t\trectified = rectify_image(image, gridpoint_xy)\n\t\t\t\t\t# ---\n\t\t\t\t\t# cv2.imwrite(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR))\n\t\t\t\t\tnp.save(f'{OUT_DIR}/rectified/{sample_id}.gridpoint_xy.npy',gridpoint_xy)\n\t\t\t\texcept:\n\t\t\t\t\tFAIL_ID.append(sample_id)\n\n\t\ttorch.cuda.empty_cache()\n\t\tif 0: #n<10: # optional: show results\n\t\t\toverlay = draw_mapping(image, gridpoint_xy) #\n\t\t\tghfiltered, gvfiltered = draw_results_stage1(more)\n            \n\t\t\t\n\t\t\tprint('')\n\t\t\tprint('demo results for stage1--------------')\n\t\t\tprint(sample_id)\n\t\t\tplt.imshow(overlay);plt.show()\n\t\t\tplt.imshow(gvfiltered);plt.show()\n\t\t\tplt.imshow(ghfiltered);plt.show()\n\t\t\tplt.imshow(rectified);plt.show()\n             \n\tprint('')\n\nrun_stage1()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage1() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:19:52.431657Z","iopub.execute_input":"2026-01-22T00:19:52.431909Z","iopub.status.idle":"2026-01-22T00:19:58.808168Z","shell.execute_reply.started":"2026-01-22T00:19:52.431885Z","shell.execute_reply":"2026-01-22T00:19:58.807492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom multiprocessing import Pool, cpu_count\nfrom functools import partial\nimport torch\nimport torch.nn.functional as F\nimport gc\n\n\n    \ndef apply_rotation(image, k):\n    if k == 0:\n        return image\n    return np.ascontiguousarray(np.rot90(image, k, axes=(0, 1)))\n    \ndef get_highres_homography(homo_low, scale):\n    S = np.array([\n        [scale, 0, 0],\n        [0, scale, 0],\n        [0, 0, 1]\n    ])\n    return S @ homo_low\n\n\ndef rectify_highres(normalized_img, grid_small, scale, target_shape):\n    target_h, target_w = target_shape\n    H_norm, W_norm = normalized_img.shape[:2]\n\n    grid_large = grid_small * scale\n\n    grid_norm = grid_large / np.array([[[W_norm - 1, H_norm - 1]]]) * 2 - 1\n\n    sparse_map = torch.from_numpy(\n        np.ascontiguousarray(grid_norm.transpose(2, 0, 1))\n    ).unsqueeze(0).float()\n\n    dense_map = F.interpolate(\n        sparse_map,\n        size=(target_h, target_w),\n        mode='bilinear',\n        align_corners=True\n    )\n\n    distort = torch.from_numpy(\n        np.ascontiguousarray(normalized_img.transpose(2, 0, 1))\n    ).unsqueeze(0).float()\n\n    rectified = F.grid_sample(\n        distort,\n        dense_map.permute(0, 2, 3, 1),\n        mode='bilinear',\n        padding_mode='border',\n        align_corners=False\n    )\n\n    rectified = rectified[0].permute(1, 2, 0).byte().cpu().numpy()\n\n    return rectified\n\n\ndef compute_fixed_size_rectified_img(sample_id, img_path, stage0_dir, stage1_dir, out_img_dir, out_scale_dir):\n    save_img_path = os.path.join(out_img_dir, f'{sample_id}.rect.png')\n    save_scale_path = os.path.join(out_scale_dir, f'{sample_id}.npy')\n    homo_path = os.path.join(stage0_dir, f'{sample_id}.homo.npy')\n    grid_path = os.path.join(stage1_dir, f'{sample_id}.gridpoint_xy.npy')\n    rotation_path = os.path.join(stage0_dir, f'{sample_id}.rotation.npy')\n\n\n    if not (os.path.exists(homo_path) and os.path.exists(grid_path) and os.path.exists(rotation_path)):  \n        raise FileNotFoundError(f\"Missing dependency files for {sample_id}\")\n        \n    original_img = cv2.imread(img_path)\n    ##########\n    # h, w = original_img.shape[:2]\n    # scale = CFG.BASE_WIDTH*15 / w\n    # original_img = cv2.resize(original_img, (int(w*scale),int(h * scale) ), interpolation=cv2.INTER_AREA)\n\n    #########\n    \n    if original_img is None:\n        raise ValueError(f\"Failed to read image: {img_path}\")\n\n    k = int(np.load(rotation_path)[0])\n    original_img = apply_rotation(original_img, k)\n\n    H_orig, W_orig = original_img.shape[:2]\n\n    scale = W_orig / CFG.BASE_WIDTH\n\n\n    homo = np.load(homo_path)\n    homo_highres = get_highres_homography(homo, scale)\n\n    norm_w = int(CFG.BASE_WIDTH * scale)\n    norm_h = int(CFG.BASE_HEIGHT * scale)\n\n    if 1:\n\n        normalized_highres = cv2.warpPerspective(\n            original_img,\n            homo_highres,\n            (norm_w, norm_h),\n            flags=cv2.INTER_CUBIC\n        )\n    else:\n        # print(f\"Scale for {sample_id}: {scale}\")\n        normalized_highres = cv2.warpPerspective(\n            original_img,\n            homo_highres,\n            (norm_w, norm_h),\n            flags=cv2.INTER_LINEAR   # save memory. if not good change it to INTER_CUBIC\n        )\n\n\n    del original_img\n    gc.collect()\n\n    grid_small = np.load(grid_path)\n    target_rect_w_natural = int(CFG.BASE_RECT_W * scale)\n    target_rect_h_natural = int(CFG.BASE_RECT_H * scale)\n\n\n    rectified_natural = rectify_highres(\n        normalized_highres,\n        grid_small,\n        scale,\n        (target_rect_h_natural, target_rect_w_natural)\n    )\n    del normalized_highres\n    del grid_small\n    gc.collect()\n    rectified_final = cv2.resize(\n        rectified_natural,\n        (CFG.FIXED_RECT_WIDTH, CFG.FIXED_RECT_HEIGHT),\n        interpolation=cv2.INTER_CUBIC\n    )\n    del rectified_natural\n    gc.collect()\n        \n    cv2.imwrite(save_img_path, rectified_final)\n    np.save(save_scale_path, np.array([scale]))\n    return True\n    \n\n    \nos.makedirs(CFG.OUT_IMG_DIR, exist_ok=True)\nos.makedirs(CFG.OUT_SCALE_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:19:58.808883Z","iopub.execute_input":"2026-01-22T00:19:58.809183Z","iopub.status.idle":"2026-01-22T00:19:58.821473Z","shell.execute_reply.started":"2026-01-22T00:19:58.809147Z","shell.execute_reply":"2026-01-22T00:19:58.82097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/input/physionet-data')\n\nfrom models.stage2_net import Stage2Net, Stage2ConvNeXtV2, Stage2HRNet, Stage2EfficientNetV2\nfrom configs.stage2 import Stage2Config_4352x1696, Stage2ConvNeXtV2Config_4352x1696, Stage2HRNetConfig_4352x1696, Stage2EfficientNetV2Config_4352x1696\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:19:58.822217Z","iopub.execute_input":"2026-01-22T00:19:58.822444Z","iopub.status.idle":"2026-01-22T00:20:35.992073Z","shell.execute_reply.started":"2026-01-22T00:19:58.822429Z","shell.execute_reply":"2026-01-22T00:20:35.991249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom torch.utils.data import Dataset, DataLoader\n# from frog.stage0_common import load_net\nfrom stage2_model import prob_to_series_by_max\nfrom stage2_common import *\n# from utils.postprocess import pixel_to_series  #FIXME\n\ndef pixel_to_series(pixel, zero_mv, length):\n    if isinstance(pixel, np.ndarray):\n        pixel = torch.from_numpy(pixel)\n    \n    C, H, W = pixel.shape\n    \n    vals_max, indices_max = torch.max(pixel, dim=1)\n    \n    indices_left = torch.clamp(indices_max - 1, min=0)\n    indices_right = torch.clamp(indices_max + 1, max=H - 1)\n    \n    vals_left = torch.gather(pixel, 1, indices_left.unsqueeze(1)).squeeze(1)\n    vals_right = torch.gather(pixel, 1, indices_right.unsqueeze(1)).squeeze(1)\n    vals_center = vals_max\n    \n    denom = vals_left - 2 * vals_center + vals_right + 1e-8\n    delta = 0.5 * (vals_left - vals_right) / denom\n    \n    delta = torch.clamp(delta, -0.5, 0.5)\n    \n    y_coords = indices_max.float() + delta\n    \n    confidence_probs = vals_center\n    confidence_threshold = 0.15\n    \n    series = y_coords.cpu().numpy()\n    conf_np = confidence_probs.cpu().numpy()\n    \n    for j in range(C):\n        mask_miss = conf_np[j] < confidence_threshold\n        if np.all(mask_miss):\n            series[j][:] = zero_mv[j]\n        elif np.any(mask_miss):\n            valid_x = np.where(~mask_miss)[0]\n            valid_y = series[j][valid_x]\n            miss_x = np.where(mask_miss)[0]\n            series[j][miss_x] = np.interp(miss_x, valid_x, valid_y)\n            # if len(valid_x) > 2:\n            #     itp = PchipInterpolator(valid_x, valid_y)\n            #     series[j][miss_x] = itp(miss_x)\n            # else:\n            #     series[j][miss_x] = np.interp(miss_x, valid_x, valid_y)\n\n    if length is not None and length != W:\n        # series_t = torch.from_numpy(series).unsqueeze(0)\n        # series_t = F.interpolate(series_t, size=length, mode='linear', align_corners=False)\n        # series = series_t.squeeze(0).numpy()\n        series = resample(series, int(length), axis=-1)\n\n    return series.astype(np.float32)\n\n\n\ndef read_sampling_length(sample_id):\n    if MODE == 'local':\n        image_id, type_id = sample_id.split('-')\n        d = valid_df[valid_df['id']==image_id].iloc[0]\n        length = d.sig_len\n        return length\n    if MODE == 'submit':\n        image_id = sample_id\n        d = valid_df[\n            (valid_df['id']==image_id) & (valid_df['lead']=='II')\n        ].iloc[0]\n        length = d.number_of_rows\n        # length = d.fs*10 \n        return length\n        \n\n# Create directories\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n\n# Define Model Configurations\nSTAGE2_CHECKPOINTS = [\n    {\n        'model_class': Stage2ConvNeXtV2,\n        'config': Stage2ConvNeXtV2Config_4352x1696(),\n        'weight': 0.5,\n        'tta': True,\n        'checkpoints': [\n            {'fold': 0, 'path': '/kaggle/input/physionet-data/run15/fold_0/epoch_0024.pth'},\n            {'fold': 1, 'path': '/kaggle/input/physionet-data/run15/fold_1/epoch_0023.pth'},\n            {'fold': 2, 'path': '/kaggle/input/physionet-data/run15/fold_2/epoch_0028.pth'},\n            {'fold': 3, 'path': '/kaggle/input/physionet-data/run15/fold_3/epoch_0025.pth'},\n            {'fold': 4, 'path': '/kaggle/input/physionet-data/run15/fold_4/epoch_0026.pth'},\n        ],\n    },\n    {\n        'model_class': Stage2ConvNeXtV2,\n        'config': Stage2ConvNeXtV2Config_4352x1696(),\n        'weight': 0.3,\n        'tta': False,\n        'checkpoints': [\n            {'fold': 0, 'path': '/kaggle/input/physionet-data/run13/fold_0/epoch_0029.pth'},\n            {'fold': 1, 'path': '/kaggle/input/physionet-data/run13/fold_1/epoch_0032.pth'},\n            {'fold': 2, 'path': '/kaggle/input/physionet-data/run13/fold_2/epoch_0030.pth'},\n            {'fold': 3, 'path': '/kaggle/input/physionet-data/run13/fold_3/epoch_0030.pth'},\n            {'fold': 4, 'path': '/kaggle/input/physionet-data/run13/fold_4/epoch_0032.pth'},\n        ],\n    },\n    {\n        'model_class': Stage2HRNet,\n        'config': Stage2HRNetConfig_4352x1696(),\n        'weight': 0.2,\n        'tta': False,\n        'checkpoints': [\n            {'fold': 0, 'path': '/kaggle/input/physionet-data/run10/fold_0/epoch_0026.pth'},\n            {'fold': 1, 'path': '/kaggle/input/physionet-data/run10/fold_1/epoch_0029.pth'},\n            {'fold': 2, 'path': '/kaggle/input/physionet-data/run10/fold_2/epoch_0028.pth'},\n            {'fold': 3, 'path': '/kaggle/input/physionet-data/run10/fold_3/epoch_0019.pth'}, # FIXME\n            {'fold': 4, 'path': '/kaggle/input/physionet-data/run10/fold_4/epoch_0028.pth'},\n        ],\n    },\n    # {\n    #     'model_class': Stage2EfficientNetV2,\n    #     'config': Stage2EfficientNetV2Config_4352x1696(),\n    #     'weight': 0.1,\n    #     'tta': False,\n    #     'checkpoints': [\n    #         {'fold': 0, 'path': '/kaggle/input/physionet-data/run12/fold_0/epoch_0028.pth'},\n    #         {'fold': 1, 'path': '/kaggle/input/physionet-data/run12/fold_1/epoch_0028.pth'},\n    #         {'fold': 2, 'path': '/kaggle/input/physionet-data/run12/fold_2/epoch_0027.pth'},\n    #         {'fold': 3, 'path': '/kaggle/input/physionet-data/run12/fold_3/epoch_0027.pth'},\n    #         {'fold': 4, 'path': '/kaggle/input/physionet-data/run12/fold_4/epoch_0028.pth'},\n    #     ],\n    # },\n]\n\nclass Stage2Dataset(Dataset):\n    def __init__(self, sample_ids, mode, kaggle_dir, cfg):\n        self.sample_ids = sample_ids\n        self.mode = mode\n        self.kaggle_dir = kaggle_dir\n        self.cfg = cfg\n\n    def __len__(self):\n        return len(self.sample_ids)\n\n    def __getitem__(self, idx):\n        sample_id = self.sample_ids[idx]\n\n\n        length = read_sampling_length(sample_id)\n\n\n        if self.mode == \"submit\":\n            image_path = f'{self.kaggle_dir}/test/{sample_id}.png'\n        else:\n            image_path = f'data/train_flat/{sample_id}.png'\n            \n        try:\n            success = compute_fixed_size_rectified_img(\n                sample_id, image_path,\n                self.cfg.STAGE0_DIR, self.cfg.STAGE1_DIR,\n                self.cfg.OUT_IMG_DIR, self.cfg.OUT_SCALE_DIR\n            )\n            \n            if not success:\n                raise ValueError(f\"Preprocessing returned False for {sample_id}\")\n\n            rect_path = f'{self.cfg.OUT_IMG_DIR}/{sample_id}.rect.png'\n            image = cv2.imread(rect_path, cv2.IMREAD_COLOR)\n            if image is None:\n                raise ValueError(f\"Failed to read rectified image: {rect_path}\")\n\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n            image_tensor = torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).float()\n            \n            return {\n                'valid': True,\n                'sample_id': sample_id,\n                'image': image_tensor,\n                'length': length\n            }\n\n        except Exception as e:\n            \n            C = 3\n            H = self.cfg.FIXED_RECT_HEIGHT\n            W = self.cfg.FIXED_RECT_WIDTH\n            dummy_image = torch.zeros((C, H, W), dtype=torch.float32)\n            \n            return {\n                'valid': False,       \n                'sample_id': sample_id,\n                'image': dummy_image, \n                'length': length     \n            }\n\ndef run_stage2_parallel():\n    # --- 1. Initialize Models ---\n    models_info = []\n    total_weight = 0\n    \n    for model_cfg in STAGE2_CHECKPOINTS:\n        cfg = model_cfg['config']\n        cfg.pretrained = False\n        net = model_cfg['model_class'](cfg)\n        net.to(DEVICE)\n        \n        if torch.cuda.device_count() > 1:\n            net = torch.nn.DataParallel(net)\n            \n        net.eval()\n        \n        models_info.append({\n            'net': net,\n            'config': cfg,\n            'checkpoints': model_cfg['checkpoints'],\n            'weight': model_cfg['weight'],\n            'tta': model_cfg['tta'],\n            'name': model_cfg['model_class'].__name__,\n        })\n        total_weight += model_cfg['weight']\n        print(f\"Initialized: {model_cfg['model_class'].__name__} with {len(model_cfg['checkpoints'])} folds\")\n\n    # --- 2. Setup DataLoader ---\n    dataset = Stage2Dataset(valid_id, MODE, KAGGLE_DIR, CFG)\n    loader = DataLoader(dataset, batch_size=2, shuffle=False, num_workers=2, pin_memory=True, drop_last=False)\n    \n    print(f\"Processing {len(valid_id)} samples in {len(loader)} batches...\")\n    start_timer = timer()\n    \n    # --- 3. Inference Loop ---\n    for batch_idx, batch_data in enumerate(loader):\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\rBatch {batch_idx+1}/{len(loader)} done. Time: {timestamp}', end='', flush=True)\n        \n        valid_mask = batch_data['valid']\n        \n        current_bs = len(batch_data['sample_id'])\n        sample_ids = np.array(batch_data['sample_id'])\n        images = batch_data['image'].float()\n        lengths = batch_data['length'].numpy()\n        \n        sample_status = [True] * current_bs \n        batch_results = [ [] for _ in range(current_bs) ]\n        \n        # Iterate over architectures\n        for model_info in models_info:\n            net = model_info['net']\n            cfg = model_info['config']\n            w = model_info['weight']\n            use_tta = model_info[\"tta\"]\n            x0, x1 = cfg.crop_x_range\n            y0, y1 = cfg.crop_y_range\n            \n            try:\n                model_input = images[:, :, y0:y1, x0:x1].to(DEVICE)\n                \n                aug_dims_list = [None, [-1], [-2], [-1, -2]] if use_tta else [None]\n                # =================================================\n                \n                fold_pixel_preds = []\n                for ckpt in model_info['checkpoints']:\n                    if isinstance(net, torch.nn.DataParallel):\n                        load_net(net.module, ckpt['path'])\n                    else:\n                        load_net(net, ckpt['path'])\n                    \n                    sum_preds = None\n                    \n                    for dims in aug_dims_list:\n                        if dims is None:\n                            aug_input = model_input\n                        else:\n                            aug_input = torch.flip(model_input, dims=dims)\n                        \n                        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n                            with torch.no_grad():\n                                output = net({'image': aug_input})\n                                curr_pred = output['pixel'].float()\n                        \n                        if dims is not None:\n                            curr_pred = torch.flip(curr_pred, dims=dims)\n                        \n                        if sum_preds is None:\n                            sum_preds = curr_pred\n                        else:\n                            sum_preds += curr_pred\n                    \n                    avg_pred = sum_preds / len(aug_dims_list)\n                    fold_pixel_preds.append(avg_pred.cpu().numpy())\n            except Exception as e:\n                print(f\"\\n[Warn] Batch failed: {e}. Retrying per sample.\")\n                for k in range(current_bs):\n                    sample_status[k] = False\n\n            # --- Aggregate & Post-process (CPU) ---\n            t0, t1 = cfg.time_range\n            zero_mv = cfg.zero_mv_positions\n            mv_to_pixel = cfg.mv_to_pixel\n            \n            for i in range(current_bs):\n                if not valid_mask[i] or not sample_status[i]: continue\n                \n                try:\n                    sample_fold_series = []\n                    for f_idx in range(len(fold_pixel_preds)):\n                        pixel_map = fold_pixel_preds[f_idx][i]\n                        series_in_pixel = pixel_to_series(pixel_map[..., t0:t1], zero_mv, lengths[i])\n                        series = (np.array(zero_mv).reshape(4, 1) - series_in_pixel) / mv_to_pixel\n                        sample_fold_series.append(series)\n                    \n                    model_avg = np.mean(sample_fold_series, axis=0)\n                    batch_results[i].append(model_avg * w)\n                    \n                except Exception as e:\n                    print(f\"\\n[Error] Post-process failed for {sample_ids[i]}: {e}\")\n                    sample_status[i] = False\n\n        # --- 4. Final Average & Save ---\n        for i in range(current_bs):\n            output_path = f'{OUT_DIR}/digitalised/{sample_ids[i]}.series.npy'\n            \n            try:\n                if valid_mask[i] and sample_status[i] and len(batch_results[i]) == len(models_info):\n                    final_series = np.sum(batch_results[i], axis=0) / total_weight\n\n                    if np.isnan(final_series).any():\n                        raise ValueError(f\"NaN detected in prediction for {sample_ids[i]}\")\n                    if not np.isfinite(final_series).all():\n                        # print(f\"Non-finite detected, fallback to zero\")\n                        raise ValueError(f\"Non-finite values detected\")\n                        \n                    np.save(output_path, final_series)\n                    \n                    # Visualization (First sample only)\n                    if batch_idx == 0 and i == 0:\n                        print(f\"Visualizing sample: {sample_ids[i]}\")\n                        viz_cfg = models_info[0]['config']\n                        x0, x1 = viz_cfg.crop_x_range\n                        y0, y1 = viz_cfg.crop_y_range\n                        \n                        # Get numpy image from tensor (C,H,W) -> (H,W,C)\n                        raw_img = images[i].byte().numpy().transpose(1, 2, 0)\n                        crop_viz = raw_img[y0:y1, x0:x1]\n                        \n                        # Overlay from first model, first fold\n                        pixel_viz = fold_pixel_preds[0][i] # from last model loop\n                        overlay = draw_lead_pixel(crop_viz, pixel_viz)\n                        \n                        plt.figure(figsize=(10,5))\n                        plt.imshow(overlay)\n                        plt.title(f\"Overlay {sample_ids[i]}\")\n                        plt.show()\n                        \n                        # Plot series\n                        t = np.arange(len(final_series[0]))\n                        fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n                        for j in range(4):\n                            axes[j].plot(t, final_series[j], color='blue', linewidth=1, label='predict')\n                            axes[j].legend()\n                        plt.show()\n                else:\n                    raise ValueError(\"Marked as failed\")\n\n            except:\n                target_len = lengths[i] \n                series = np.zeros((4, int(target_len)))\n                np.save(output_path, series)\n                \n        torch.cuda.empty_cache()\n        if batch_idx % 50 == 0:\n            gc.collect()\n    print('\\nStage 2 Finished.')\n\nif __name__ == '__main__':\n    run_stage2_parallel()\n    print('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:20:35.994181Z","iopub.execute_input":"2026-01-22T00:20:35.994739Z","iopub.status.idle":"2026-01-22T00:22:10.444879Z","shell.execute_reply.started":"2026-01-22T00:20:35.994718Z","shell.execute_reply":"2026-01-22T00:22:10.444082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#make sbmission csv\n#FAIL_ID=[1053922973, ]\n\ndef dw_after_weighted(series_dict, weights={'I': 1.0, 'II': 2.0, 'III': 1.0}):\n    if all(k in series_dict for k in ['I', 'II', 'III']):\n        len_I = len(series_dict['I'])\n        len_II = len(series_dict['II'])\n        len_III = len(series_dict['III'])\n        n = min(len_I, len_II, len_III)\n\n        part_I = series_dict['I'][:n]\n        part_II = series_dict['II'][:n]\n        part_III = series_dict['III'][:n]\n\n        error = part_II - (part_I + part_III)\n\n        w_I = weights.get('I', 1.0)\n        w_II = weights.get('II', 1.0)\n        w_III = weights.get('III', 1.0)\n\n        inv_I = 1.0 / w_I\n        inv_II = 1.0 / w_II\n        inv_III = 1.0 / w_III\n        \n        total_inv_weight = inv_I + inv_II + inv_III\n\n        ratio_I = inv_I / total_inv_weight\n        ratio_II = inv_II / total_inv_weight\n        ratio_III = inv_III / total_inv_weight\n\n        corr_I = error * ratio_I\n        corr_III = error * ratio_III\n        corr_II = error * ratio_II\n\n        \n        series_dict['I'] = part_I + corr_I\n        series_dict['III'] = part_III + corr_III\n        \n        series_dict['II'][:n] = part_II - corr_II\n\n    return series_dict\n\n\ndef make_submission():\n        \n    print('===========================================')\n    print('making submission csv ...')\n    \n    submit_df=[]\n    gb = valid_df.groupby('id')\n    for i,(sample_id, df) in enumerate(gb):\n        \n        #if sample_id in FAIL_ID:\n        #\tseries_by_lead = {}\n        #\tfor j,d in df.iterrows():\n        #\t\tseries_by_lead[d.lead] = np.zeros(d.number_of_rows)\n        \n        try:\n            series = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.npy')\n            _4_,L = series.shape\n\n            #https://www.kaggle.com/competitions/physionet-ecg-image-digitization/discussion/613179#3306701\n            #may be even or odd????\n            series_by_lead={}\n            for l in range(3):\n                lead = [\n                    ['I',   'aVR', 'V1', 'V4'],\n                    ['II',  'aVL', 'V2', 'V5'],\n                    ['III', 'aVF', 'V3', 'V6'],\n                ][l]\n\n                 \n\n                index = [ \n                    int(round(1*L/4)),\n                    int(round(2*L/4)),\n                    int(round(3*L/4)),\n                ]\n                split = np.split(series[l], index)\n                #print(length)\n                for (k, s) in zip(lead, split):\n                    series_by_lead[k] = s\n                    #print(k,len(s))\n            \n            \n\n            if 1: # average II\n                short_ii = series_by_lead['II']\n                long_ii = series[3]\n                overlap_len = min(len(short_ii), len(long_ii))\n                averaged_head = (short_ii[:overlap_len] + long_ii[:overlap_len]) / 2.0\n                final_ii = np.concatenate([averaged_head, long_ii[overlap_len:]])\n                series_by_lead['II'] = final_ii\n            else:\n                # Row 3: Complete II (10s) - overwrite the short II from Row 1\n                series_by_lead['II'] = series[3]\n                \n            ############\n            series_by_lead = dw_after_weighted(series_by_lead, weights={'I': 1.0, 'II': 1.0, 'III': 1.0})\n            ############\n        \n        except: \n            series_by_lead = {}\n            for j,d in df.iterrows():\n                series_by_lead[d.lead] = np.zeros(d.number_of_rows)\n\n        #print('\\r\\t {sample_id}', end='', flush=True)\n        for j,d in df.iterrows():\n            #\n            \n            #assert(len(series_by_lead[d.lead])==d.number_of_rows)\n             \n            #probably error here ... ???\n            series_by_lead[d.lead] = np.concatenate([\n                series_by_lead[d.lead], np.zeros_like(series_by_lead[d.lead])\n            ])[:d.number_of_rows]\n\n\n            # series_by_lead[d.lead] = np.clip(series_by_lead[d.lead], -15.0, 15.0)\n            \n            #print(d.lead, len(series_by_lead[d.lead]),d.number_of_rows)\n            assert(len(series_by_lead[d.lead])==d.number_of_rows) \n            # print(f'\\r\\t {i} {sample_id} : {d.lead}', end='', flush=True)\n\n            row_id = [\n                f'{sample_id}_{i}_{d.lead}' for i in range(d.number_of_rows)\n            ]\n            this_df = pd.DataFrame({\n                'id':row_id,\n                'value': series_by_lead[d.lead].astype(np.float32),\n            })\n            submit_df.append(this_df)\n\n    print('')\n    submit_df = pd.concat(submit_df, axis=0, ignore_index=True, sort=False, copy=False)\n    print(submit_df)\n    submit_df.to_csv('submission.csv',index=False)\n\n\nif (MODE=='fake')|(MODE=='submit'):\n    make_submission()\n    print('make_submission() ok!!!\\n')\n    # if MODE=='submit':\n    #     shutil.rmtree(OUT_DIR)\n    !ls\n    #!rm -rf {OUT_DIR}\n\n'''\nfake:\n[21618231 rows x 2 columns]\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T00:23:31.388454Z","iopub.execute_input":"2026-01-22T00:23:31.38878Z","iopub.status.idle":"2026-01-22T00:23:32.031082Z","shell.execute_reply.started":"2026-01-22T00:23:31.388753Z","shell.execute_reply":"2026-01-22T00:23:32.030169Z"}},"outputs":[],"execution_count":null}]}