{"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":13746387,"datasetId":8747012,"databundleVersionId":14496454},{"sourceType":"modelInstanceVersion","sourceId":623160,"databundleVersionId":14269261,"modelInstanceId":468838,"modelId":484689},{"sourceType":"modelInstanceVersion","sourceId":623170,"databundleVersionId":14269443,"modelInstanceId":468846,"modelId":484698}],"dockerImageVersionId":31154,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport os\nimport torch\nfrom tabulate import tabulate\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Install cc3d silently\ntry:\n    import cc3d\nexcept:\n    import subprocess\n    subprocess.run([\n        'pip', 'install', 'connected-components-3d', '--no-index',\n        '--find-links=file:///kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/setup',\n        '-q'\n    ], capture_output=True)\n    import cc3d\nfrom scipy.signal import savgol_filter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:47:38.483569Z","iopub.execute_input":"2026-04-26T14:47:38.483811Z","iopub.status.idle":"2026-04-26T14:47:49.458244Z","shell.execute_reply.started":"2026-04-26T14:47:38.483785Z","shell.execute_reply":"2026-04-26T14:47:49.457671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nProperly using hengck23's pretrained models with fixed imports\n\"\"\"\n\nprint(\"=\" * 80)\nprint(\"ECG INFERENCE\")\nprint(\"=\" * 80)\n\n# Setup paths\nimport sys\nbase_path = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet'\nsys.path.insert(0, base_path)\n\nKAGGLE_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nWEIGHT_DIR = f'{base_path}/weight'\nOUT_DIR = '/kaggle/working/output-submit'\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nFLOAT_TYPE = torch.float16\n\nos.makedirs(f'{OUT_DIR}/normalised', exist_ok=True)\nos.makedirs(f'{OUT_DIR}/rectified', exist_ok=True)\nos.makedirs(f'{OUT_DIR}/digitalised', exist_ok=True)\n\nprint(f\"\\n🔧 Device: {DEVICE}\")\nprint(f\"📁 Weights: {WEIGHT_DIR}\")\n\n# ================================================================================\n# DATA INSPECTION\n# ================================================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DATA INSPECTION\")\nprint(\"=\" * 80)\n\ntrain_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\ntest_df = pd.read_csv(f'{KAGGLE_DIR}/test.csv')\ntest_df['id'] = test_df['id'].astype(str)\nvalid_id = test_df['id'].unique().tolist()\n\nprint(\"\\n📊 DATASET OVERVIEW\")\nstats = [\n    [\"Training Samples\", len(train_df)],\n    [\"Test Images\", len(valid_id)],\n    [\"Total Test Rows\", len(test_df)],\n    [\"Sampling Frequency\", f\"{test_df['fs'].iloc[0]} Hz\"],\n]\nprint(tabulate(stats, headers=[\"Metric\", \"Value\"], tablefmt=\"fancy_grid\"))\n\nprint(\"\\n📊 TEST DATA STRUCTURE\")\nprint(tabulate(test_df.head(24), headers='keys', tablefmt='fancy_grid', showindex=False))\n\n# ================================================================================\n# LOAD MODELS\n# ================================================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOADING MODELS\")\nprint(\"=\" * 80)\n\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import load_net, image_to_batch, output_to_predict as s0_output, normalise_by_homography\n\nprint(\"\\n🔧 Loading Stage 0...\")\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, f'{WEIGHT_DIR}/stage0-last.checkpoint.pth')\nstage0_net.to(DEVICE).eval()\nprint(\"✅ Stage 0 loaded\")\n\nfrom stage1_model import Net as Stage1Net\nfrom stage1_common import load_net as s1_load, output_to_predict as s1_output, rectify_image\n\nprint(\"\\n🔧 Loading Stage 1...\")\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = s1_load(stage1_net, f'{WEIGHT_DIR}/stage1-last.checkpoint.pth')\nstage1_net.to(DEVICE).eval()\nprint(\"✅ Stage 1 loaded\")\n\nfrom stage2_model import Net as Stage2Net\nfrom stage2_common import load_net as s2_load, pixel_to_series, filter_series_by_limits\n\nprint(\"\\n🔧 Loading Stage 2...\")\nstage2_net = Stage2Net(pretrained=False)\nstage2_net = s2_load(stage2_net, f'{WEIGHT_DIR}/stage2-00005810.checkpoint.pth')\nstage2_net.to(DEVICE).eval()\nprint(\"✅ Stage 2 loaded\")\n\n# ================================================================================\n# PROCESSING PIPELINE\n# ================================================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"PROCESSING TEST IMAGES\")\nprint(\"=\" * 80)\n\nFAIL_ID = []\n\n# STAGE 0\nprint(\"\\n🔄 Stage 0: Normalization...\")\nfor n, sample_id in enumerate(tqdm(valid_id, desc=\"Stage 0\")):\n    try:\n        image = cv2.imread(f'{KAGGLE_DIR}/test/{sample_id}.png', cv2.IMREAD_COLOR_RGB)\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                rotated, keypoint = s0_output(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        \n    except Exception as e:\n        print(f\"\\n⚠️  {sample_id}: {e}\")\n        FAIL_ID.append(sample_id)\n    \n    torch.cuda.empty_cache()\n\nprint(f\"✅ Stage 0 complete: {len(valid_id) - len(FAIL_ID)}/{len(valid_id)} success\")\n\n# STAGE 1\nprint(\"\\n🔄 Stage 1: Rectification...\")\nfor n, sample_id in enumerate(tqdm(valid_id, desc=\"Stage 1\")):\n    if sample_id in FAIL_ID:\n        continue\n    \n    try:\n        image = cv2.imread(f'{OUT_DIR}/normalised/{sample_id}.norm.png', cv2.IMREAD_COLOR_RGB)\n        batch = {'image': torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0)}\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage1_net(batch)\n                gridpoint_xy, more = s1_output(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        \n    except Exception as e:\n        print(f\"\\n⚠️  {sample_id}: {e}\")\n        FAIL_ID.append(sample_id)\n    \n    torch.cuda.empty_cache()\n\nprint(f\"✅ Stage 1 complete: {len(valid_id) - len(FAIL_ID)}/{len(valid_id)} success\")\n\n# STAGE 2\nprint(\"\\n🔄 Stage 2: Signal Extraction...\")\nfor n, sample_id in enumerate(tqdm(valid_id, desc=\"Stage 2\")):\n    if sample_id in FAIL_ID:\n        continue\n    \n    try:\n        image = cv2.imread(f'{OUT_DIR}/rectified/{sample_id}.rect.png', cv2.IMREAD_COLOR_RGB)\n        d = test_df[(test_df['id']==sample_id) & (test_df['lead']=='II')].iloc[0]\n        length = d.number_of_rows\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 = 80.0\n        t0, t1 = 118, 2080\n        \n        crop = image[y0:y1, x0:x1]\n        batch = {'image': torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0)}\n        \n        with torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            with torch.no_grad():\n                output = stage2_net(batch)\n        \n        pixel = output['pixel'].float().data.cpu().numpy()[0]\n        series_in_pixel = pixel_to_series(pixel[..., t0:t1], zero_mv, length)\n        series = (np.array(zero_mv).reshape(4, 1) - series_in_pixel) / mv_to_pixel\n        series = filter_series_by_limits(series)\n\n        for i in range(series.shape[0]):\n            # 2. Add a Median Filter BEFORE Savgol. \n            # This instantly kills single-pixel anomalies/spikes without blurring the sharp QRS peaks.\n            from scipy.signal import medfilt\n            series[i] = medfilt(series[i], kernel_size=5)\n            \n            # Then apply Savgol to smooth the remaining curve\n            # series[i] = savgol_filter(series[i], window_length=11, polyorder=3)\n\n        if n < 3: \n            plt.figure(figsize=(20, 10))\n            # Display the crop used for inference\n            plt.imshow(crop) \n            \n            # We need to reverse the math: series -> pixels\n            for lead_idx in range(4): # The model outputs 4 rows of leads\n                # Reverse the normalization: (zero_mv - (series * 79.0))\n                y_coords = zero_mv[lead_idx] - (series[lead_idx] * mv_to_pixel)\n                \n                # Plot only the valid range\n                x_coords = np.arange(t0, t1)\n                # Ensure lengths match\n                limit = min(len(x_coords), len(y_coords))\n                plt.plot(x_coords[:limit], y_coords[:limit], color='red', linewidth=1)\n                \n            plt.title(f\"Prediction Overlay: {sample_id}\")\n            plt.show()\n\n        \n        np.save(f'{OUT_DIR}/digitalised/{sample_id}.series.npy', series)\n        \n    except Exception as e:\n        print(f\"\\n⚠️  {sample_id}: {e}\")\n        FAIL_ID.append(sample_id)\n    \n    torch.cuda.empty_cache()\n\nprint(f\"✅ Stage 2 complete: {len(valid_id) - len(FAIL_ID)}/{len(valid_id)} success\")\n\n# ================================================================================\n# CREATE SUBMISSION\n# ================================================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CREATING SUBMISSION\")\nprint(\"=\" * 80)\n\nsubmission_data = []\ngb = test_df.groupby('id')\n\nfor rec_idx, (sample_id, df) in enumerate(tqdm(gb, desc=\"Building submission\")):\n    \n    try:\n        series = np.load(f'{OUT_DIR}/digitalised/{sample_id}.series.npy')\n        _4_, L = series.shape\n        \n        series_by_lead = {}\n        for l in range(3):\n            lead_names = [\n                ['I', 'aVR', 'V1', 'V4'],\n                ['II', 'aVL', 'V2', 'V5'],\n                ['III', 'aVF', 'V3', 'V6'],\n            ][l]\n            \n            index = [int(round(1*L/4)), int(round(2*L/4)), int(round(3*L/4))]\n            split = np.split(series[l], index)\n            for k, s in zip(lead_names, split):\n                series_by_lead[k] = s\n        \n        series_by_lead['II'] = series[3]\n        \n    except:\n        series_by_lead = {}\n        for _, d in df.iterrows():\n            series_by_lead[d.lead] = np.zeros(d.number_of_rows, dtype=np.float32)\n    \n    for _, d in df.iterrows():\n        s = series_by_lead[d.lead]\n        target_len = int(d.number_of_rows)\n        \n        if len(s) != target_len:\n            x_old = np.linspace(0.0, 1.0, len(s), endpoint=False)\n            x_new = np.linspace(0.0, 1.0, target_len, endpoint=False)\n            s = np.interp(x_new, x_old, s)\n        \n        s = s.astype(np.float32)\n\n        # s = s - np.nanmedian(s)\n        \n        # # 2. Clip impossible voltage spikes caused by ink smudges/text\n        s = np.clip(s, -3.0, 3.0)\n        \n        for t in range(target_len):\n            submission_data.append({\n                'id': f'{sample_id}_{t}_{d.lead}',\n                'value': float(s[t])\n            })\n\nsubmission = pd.DataFrame(submission_data)\nsubmission.to_csv('submission.csv', index=False)\n\nprint(f\"\\n✅ Submission created: {len(submission):,} rows\")\n\n# ================================================================================\n# ANALYSIS\n# ================================================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"SUBMISSION ANALYSIS\")\nprint(\"=\" * 80)\n\nsubmission['lead'] = submission['id'].str.split('_').str[2]\n\nprint(\"\\n📊 PER-LEAD STATISTICS\")\nlead_stats = []\nfor lead in ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']:\n    data = submission[submission['lead'] == lead]['value']\n    lead_stats.append([\n        lead, len(data), f\"{data.mean():.6f}\", f\"{data.std():.6f}\",\n        f\"{data.min():.6f}\", f\"{data.max():.6f}\", (data != 0).sum()\n    ])\n\nprint(tabulate(lead_stats, headers=['Lead', 'Count', 'Mean', 'Std', 'Min', 'Max', 'Non-Zero'],\n              tablefmt='fancy_grid'))\n\nprint(f\"\\n✅ Failed IDs: {len(FAIL_ID)}\")\nif FAIL_ID:\n    print(f\"   {FAIL_ID}\")\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"✅ COMPLETE\")\nprint(\"=\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:47:49.459417Z","iopub.execute_input":"2026-04-26T14:47:49.460298Z","iopub.status.idle":"2026-04-26T14:48:12.139076Z","shell.execute_reply.started":"2026-04-26T14:47:49.460271Z","shell.execute_reply":"2026-04-26T14:48:12.138297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndef visualize_pipeline(sample_id):\n    \"\"\"\n    Visualizes the 3 stages of the pipeline for a single ID.\n    \"\"\"\n    # Paths to the saved intermediate files\n    norm_path = f'{OUT_DIR}/normalised/{sample_id}.norm.png'\n    rect_path = f'{OUT_DIR}/rectified/{sample_id}.rect.png'\n    signal_path = f'{OUT_DIR}/digitalised/{sample_id}.series.npy'\n    \n    # 1. Load Images\n    img_norm = cv2.imread(norm_path)\n    if img_norm is None:\n        print(f\"❌ Could not find normalized image for {sample_id}\")\n        return\n    img_norm = cv2.cvtColor(img_norm, cv2.COLOR_BGR2RGB)\n    \n    img_rect = cv2.imread(rect_path)\n    if img_rect is None:\n        print(f\"❌ Could not find rectified image for {sample_id}\")\n        return\n    img_rect = cv2.cvtColor(img_rect, cv2.COLOR_BGR2RGB)\n    \n    # 2. Load Signal Data\n    try:\n        series = np.load(signal_path)\n    except:\n        print(f\"❌ Could not find signal data for {sample_id}\")\n        return\n\n    # 3. Setup the Plot (1 Row, 3 Columns)\n    fig, axes = plt.subplots(1, 3, figsize=(24, 8))\n    \n    # --- PANEL 1: Stage 0 (Normalization) ---\n    axes[0].imshow(img_norm)\n    axes[0].set_title(\"Stage 0: Normalization\\n(Perspective Corrected)\", fontsize=14, fontweight='bold')\n    axes[0].axis('off')\n    \n    # --- PANEL 2: Stage 1 (Rectification) ---\n    axes[1].imshow(img_rect)\n    axes[1].set_title(\"Stage 1: Rectification\\n(Grid Unwarping)\", fontsize=14, fontweight='bold')\n    axes[1].axis('off')\n    \n    # --- PANEL 3: Stage 2 (Signal Extraction) ---\n    # We plot the rectified image again, but overlay the red signal line\n    axes[2].imshow(img_rect)\n    axes[2].set_title(\"Stage 2: Signal Extraction\\n(Red Line = Prediction)\", fontsize=14, fontweight='bold')\n    axes[2].axis('off')\n    \n    # Parameters for reconstruction (Must match your loop!)\n    zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n    mv_to_pixel = 80.0 # <--- Your new optimized value\n    t0, t1 = 118, 2080\n    \n    # Overlay the red lines\n    for lead_idx in range(4):\n        # Invert the math: Voltage -> Pixel Y-coordinate\n        y_coords = zero_mv[lead_idx] - (series[lead_idx] * mv_to_pixel)\n        x_coords = np.arange(t0, t1)\n        \n        # Clip to ensure we don't plot outside bounds\n        valid_len = min(len(x_coords), len(y_coords))\n        axes[2].plot(x_coords[:valid_len], y_coords[:valid_len], color='red', linewidth=1.5, alpha=0.9)\n\n    plt.tight_layout()\n    plt.show()\n\n# --- RUN IT ---\n# Pick the first 3 IDs from your validation list to inspect\nprint(f\"Visualizing the first 3 samples...\")\nfor sample in valid_id[:3]:\n    visualize_pipeline(sample)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:48:12.139926Z","iopub.execute_input":"2026-04-26T14:48:12.140188Z","iopub.status.idle":"2026-04-26T14:48:16.034675Z","shell.execute_reply.started":"2026-04-26T14:48:12.140149Z","shell.execute_reply":"2026-04-26T14:48:16.033558Z"}},"outputs":[],"execution_count":null}]}