{"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":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14736935,"sourceType":"datasetVersion","datasetId":9417319}],"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    !ls /kaggle/input/gavin-submit-physionet/setup\n    !pip install connected-components-3d --no-index --find-links=file:///kaggle/input/gavin-submit-physionet/setup/\n\nimport os\nimport cc3d\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport matplotlib\nimport shutil\nfrom scipy.signal import resample_poly, resample\nimport torchaudio.functional as AF\nimport sys\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\nprint('import ok!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:33:37.884356Z","iopub.execute_input":"2026-02-05T04:33:37.884593Z","iopub.status.idle":"2026-02-05T04:33:40.784366Z","shell.execute_reply.started":"2026-02-05T04:33:37.884566Z","shell.execute_reply":"2026-02-05T04:33:40.782519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nMODE   = 'submit'  # submit  local fake\nDEVICE = 'cuda'\nFLOAT_TYPE = torch.float16\nFAIL_ID = []\n\nKAGGLE_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nLIBS_DIR = '/kaggle/input/gavin-submit-physionet'\nWEIGHT_DIR = '/kaggle/input/gavin-submit-physionet/weight'\nOUT_DIR = f'/kaggle/working/output-{MODE}'\n\nsys.path.append(LIBS_DIR)\n\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    \n    for i,d in valid_df.iterrows():\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\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            \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    from sample_list import ERROR_ID\n    valid_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n    valid_df['id']=valid_df['id'].astype(str)\n\n    valid_id = [\n        '640106434-0012', '1752267607-0011', '2299138053-0005', '2352363003-0011', '3286023290-0009', \n        '3394549140-0005', '3394549140-0006', '3394549140-0010', '3446730127-0010', '3487334519-0006', \n        '3451522192-0005', '2352363003-0011', '3938971616-0012', '3286023290-0009',\n        '3523783811-0011', '3607942876-0011', '3657702924-0006', '3938971616-0012'\n    ]\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\nif MODE == 'fake':\n    valid_df = make_test_fake_df()\n    valid_df['id']=valid_df['id'].astype(str) \n    valid_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    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        return length\n    if MODE == 'fake':\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        return 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-02-05T04:33:40.785536Z","iopub.execute_input":"2026-02-05T04:33:40.786165Z","iopub.status.idle":"2026-02-05T04:33:40.817616Z","shell.execute_reply.started":"2026-02-05T04:33:40.786127Z","shell.execute_reply":"2026-02-05T04:33:40.816914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage0\nprint('*** STARTING STAGE0 ***')\n\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import *\n\nos.makedirs(f'{OUT_DIR}/normalised', exist_ok=True)\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 = output_to_predict(image, batch, output)\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                except:\n                    FAIL_ID.append(sample_id)\n\n        torch.cuda.empty_cache()\n        if 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-02-05T04:33:40.820335Z","iopub.execute_input":"2026-02-05T04:33:40.820621Z","iopub.status.idle":"2026-02-05T04:33:49.922822Z","shell.execute_reply.started":"2026-02-05T04:33:40.820602Z","shell.execute_reply":"2026-02-05T04:33:49.922135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage1\nprint('*** STARTING STAGE1 ***')\n\nfrom stage1_common import *\nfrom stage1_unet_points import UNet\n\nos.makedirs(f'{OUT_DIR}/rectified', exist_ok=True)\n\ndef remove_module_prefix(state_dict):\n    new_state_dict = {}\n    for k, v in state_dict.items():\n        if k.startswith(\"module.\"):\n            new_state_dict[k[7:]] = v   # 去掉 module.\n        else:\n            new_state_dict[k] = v\n    return new_state_dict\n\ndef load_model(**kwargs):\n    weights_path = kwargs.pop(\"weights_path\", None)\n    model = UNet(**kwargs)\n    state = torch.load(weights_path, map_location='cpu', weights_only=False) # ['state_dict']\n    state = remove_module_prefix(state)\n    state = {k.replace(\"_orig_mod.\", \"\"): v for k, v in state.items()}\n    verbose = model.load_state_dict(state, strict=False)\n    print(verbose)\n    model.eval()\n    return model\n\ndef run_stage1():\n    \n    stage1_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/stage1-points-e60.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage1_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        if sample_id in FAIL_ID: continue\n\n        image = cv2.imread(f'{OUT_DIR}/normalised/{sample_id}.norm.png', cv2.IMREAD_COLOR_RGB)\n        batch = torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0).to(DEVICE)\n        num_tta = 1\n\n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE): #torch.bfloat16\n            with torch.no_grad():\n                output = stage1_net(batch)\n\n                try:\n                    # gridpoint_xy, more = output_to_predict(image, batch, output)\n                    gridpoint_xy, more = output_to_predict_optimized(image, batch, output)\n                    rectified = rectify_image(image, gridpoint_xy)\n                    # ---\n                    cv2.imwrite(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR))\n                    np.save(f'{OUT_DIR}/rectified/{sample_id}.gridpoint_xy.npy',gridpoint_xy)\n                except:\n                    FAIL_ID.append(sample_id)\n\n        torch.cuda.empty_cache()\n        if n<10: # optional: show results\n            overlay = draw_mapping(image, gridpoint_xy) #\n            ghfiltered, gvfiltered = draw_results_stage1(more)\n            \n            \n            print('')\n            print('demo results for stage1--------------')\n            print(sample_id)\n            plt.imshow(overlay);plt.show()\n            plt.imshow(gvfiltered);plt.show()\n            plt.imshow(ghfiltered);plt.show()\n            plt.imshow(rectified);plt.show()\n             \n    print('')\n\nrun_stage1()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage1() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:33:49.923748Z","iopub.execute_input":"2026-02-05T04:33:49.924036Z","iopub.status.idle":"2026-02-05T04:33:57.903418Z","shell.execute_reply.started":"2026-02-05T04:33:49.924015Z","shell.execute_reply":"2026-02-05T04:33:57.902553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom stage2_common import *\nfrom stage2_unet_query import UNet\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n    \ndef run_stage2():\n    stage2_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/output.5120.hflip.e150.query.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage2_net.to(DEVICE)\n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n        \n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n        \n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, target_len=5120)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, target_len=5120)\n\n        L = output['prob'].shape[-1]\n        series = output[f'y_mv_{L}']\n   \n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n        \n        \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.npy', series)\n\n\n        if n<20: \n            \n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:33:57.904296Z","iopub.execute_input":"2026-02-05T04:33:57.904510Z","iopub.status.idle":"2026-02-05T04:34:02.321113Z","shell.execute_reply.started":"2026-02-05T04:33:57.904495Z","shell.execute_reply":"2026-02-05T04:34:02.320300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# stage2\nprint('*** STARTING STAGE2 ***')\nfrom stage2_common import *\nfrom stage2_unet import UNet\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n    \ndef run_stage2():\n    stage2_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/ckpt-5120-full-hflip-e150.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage2_net.to(DEVICE)\n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n\n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, resample_length=5120)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, resample_length=5120)\n\n        L = output['prob'].shape[-1]\n        series = output[f'y_mv_{L}']\n\n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n    \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.v2.npy', series)\n\n\n        if n<20: # optional: show results\n            # overlay = draw_lead_pixel(crop, pixel)\n            # plt.imshow(overlay); plt.show()\n     \n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:02.321919Z","iopub.execute_input":"2026-02-05T04:34:02.322184Z","iopub.status.idle":"2026-02-05T04:34:07.454454Z","shell.execute_reply.started":"2026-02-05T04:34:02.322165Z","shell.execute_reply":"2026-02-05T04:34:07.453696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom stage2_common import *\nfrom stage2_unet import UNet\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n    \ndef run_stage2():\n    stage2_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/ckpt-5120-full-extra-e80-hflip.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage2_net.to(DEVICE)\n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n        \n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, resample_length=5120)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, resample_length=5120)\n\n        L = output['prob'].shape[-1]\n        series = output[f'y_mv_{L}']\n\n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n        \n        \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.v3.npy', series)\n\n\n        if n<20: # optional: show results\n\n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:07.455501Z","iopub.execute_input":"2026-02-05T04:34:07.455758Z","iopub.status.idle":"2026-02-05T04:34:11.854540Z","shell.execute_reply.started":"2026-02-05T04:34:07.455738Z","shell.execute_reply":"2026-02-05T04:34:11.853615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom stage2_common import *\nfrom stage2_unet_split import UNet\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n\n    \ndef run_stage2():\n    stage2_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/ckpt.full.10250.split.e70.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage2_net.to(DEVICE)\n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n    \n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n        \n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n\n        L = 10250\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, resample_length=L)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, resample_length=L)\n\n        series = output[f'y_mv_{L}']\n\n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n        \n        \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.v4.npy', series)\n\n\n        if n<20: # optional: show results\n\n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:11.855567Z","iopub.execute_input":"2026-02-05T04:34:11.855885Z","iopub.status.idle":"2026-02-05T04:34:19.666543Z","shell.execute_reply.started":"2026-02-05T04:34:11.855858Z","shell.execute_reply":"2026-02-05T04:34:19.665827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom stage2_common import *\nfrom stage2_unet import UNet\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n    \ndef run_stage2():\n    stage2_net = load_model(\n        weights_path=f\"{WEIGHT_DIR}/ckpt-5120-full-pretrain-e100-hflip.pth\",\n        num_in_channels=3,\n        num_out_channels=4,\n        dims=[32, 64, 128, 256, 320, 320, 320, 320],\n        depth=2,\n    )\n    stage2_net.to(DEVICE)\n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n\n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, resample_length=5120)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, resample_length=5120)\n\n        L = output['prob'].shape[-1]\n        series = output[f'y_mv_{L}']\n\n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n        \n        \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.v5.npy', series)\n\n\n        if n<20: # optional: show results\n\n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:19.667376Z","iopub.execute_input":"2026-02-05T04:34:19.667653Z","iopub.status.idle":"2026-02-05T04:34:23.969739Z","shell.execute_reply.started":"2026-02-05T04:34:19.667628Z","shell.execute_reply":"2026-02-05T04:34:23.969008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# stage2\nprint('*** STARTING STAGE2 ***')\n\nfrom stage2_common import *\nfrom stage2_unet_heng import Net\n\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n\n    \ndef run_stage2():\n\n    stage2_net = Net(pretrained=False)\n\n    state = torch.load(\n        f\"{WEIGHT_DIR}/ckpt.heng.5120.full.hflip.e140.pth\", \n        map_location='cpu'\n    )\n    state = {k.replace(\"_orig_mod.\", \"\"): v for k, v in state.items()}\n    verbose = stage2_net.load_state_dict(state, strict=False)\n    print(verbose)\n    stage2_net.to(DEVICE)\n    \n\n    \n    snr_list = []\n    start_timer = timer()\n    for n, sample_id in enumerate(valid_id):\n\n        timestamp = time_to_str(timer() - start_timer, 'sec')\n        print(f'\\r\\t {n:4d} {sample_id}', timestamp, end='', flush=True)\n        if sample_id in FAIL_ID: continue\n\n        x0, x1 = 0, 2176\n        y0, y1 = 0, 1696\n        zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n        mv_to_pixel = 79\n        t0, t1 = timespan = 118, 2080\n\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n    \n        length = read_sampling_length(sample_id) #5120\n\n        crop = image[y0:y1, x0:x1][:, t0:t1]\n        \n        H, W, C = crop.shape\n        input_tensor = torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0) / 255.0\n        input_tensor = input_tensor.to(DEVICE)\n\n        L = 5120 # 10250\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(input_tensor, resample_length=L)\n\n                # hflip tta\n                input_hflip = torch.flip(input_tensor, dims=[3])\n                out_hflip = stage2_net(input_hflip, resample_length=L)\n\n        series = output[f'y_mv_{L}']\n\n        # hflip tta\n        series_hflip = torch.flip(out_hflip[f'y_mv_{L}'], dims=[2])\n\n        # tta average\n        series = (series + series_hflip) / 2.0\n        \n        \n        if L != length:\n            series = stage2_net.signal_head.resample_torch(x=series, num=length)\n\n        series = series[0].cpu().numpy()\n            \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.v6.npy', series)\n\n\n        if n<20: # optional: show results\n\n            if MODE=='local':\n                truth_df = read_truth_series(sample_id,KAGGLE_DIR)\n                truth_series = truth_df[['series0','series1','series2','series3',]].values.T\n\n            t = np.arange(len(series[0]))\n            fig, axes = plt.subplots(4, 1, figsize=(12, 10))\n            for j in range(4):\n                snr=0\n                axes[j].plot(t, series[j], alpha=1.0, color='blue', linewidth=1, label='predict')\n                if MODE=='local':\n                    axes[j].plot(t, truth_series[j], alpha=0.5, color='red', linewidth=1,label='truth')\n                    snr = -np_snr(series[j], truth_series[j])\n                    snr_list.append(snr)\n\n                axes[j].set_title(f'snr {snr:8.3f} {length}')\n                axes[j].legend()\n            plt.show()\n    print('')\n    print(np.mean(snr_list))\n    \nrun_stage2()\nprint('FAIL_ID:', FAIL_ID)\nprint('run_stage2() ok!!!\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:23.970693Z","iopub.execute_input":"2026-02-05T04:34:23.971269Z","iopub.status.idle":"2026-02-05T04:34:29.005385Z","shell.execute_reply.started":"2026-02-05T04:34:23.971241Z","shell.execute_reply":"2026-02-05T04:34:29.004616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#make sbmission csv\ndef make_submission():\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        try:\n            series_v1 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.npy')\n            series_v2 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.v2.npy')\n            series_v3 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.v3.npy')\n            series_v4 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.v4.npy')\n            series_v5 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.v5.npy')\n            series_v6 = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.v6.npy')\n\n            series = (series_v1 * 0.17 + series_v2 * 0.18 + series_v3 * 0.17 + series_v4 * 0.15 + series_v5 * 0.18 + series_v6 * 0.15)\n\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                length=[\n                    df[df['lead']==lead[j]].iloc[0].number_of_rows\n                    for j in range(4)\n                ]\n                if lead[0]=='II':\n                    length[0] = length[0]-sum(length[1:])\n\n                index = np.cumsum(length)[:-1]\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            series_by_lead['II'] = series[3]\n            #print(series_by_lead)\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        for j,d in df.iterrows():\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            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\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:34:29.006231Z","iopub.execute_input":"2026-02-05T04:34:29.006580Z","iopub.status.idle":"2026-02-05T04:34:29.723783Z","shell.execute_reply.started":"2026-02-05T04:34:29.006548Z","shell.execute_reply":"2026-02-05T04:34:29.722924Z"}},"outputs":[],"execution_count":null}]}