{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012},{"sourceId":14499498,"sourceType":"datasetVersion","datasetId":9102698},{"sourceId":14564099,"sourceType":"datasetVersion","datasetId":9260130},{"sourceId":14583741,"sourceType":"datasetVersion","datasetId":9210722},{"sourceId":14661585,"sourceType":"datasetVersion","datasetId":9156624}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport os\nimport sys\nimport time\nfrom scipy.interpolate import griddata\n\nsys.path.append('/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet')\nsys.path.append('/kaggle/input/ecg-lib')\n\nfrom stage0_common import image_to_batch, load_net\nfrom stage0_model import Net as Stage0Net\nfrom stage1_model import Net as Stage1Net\nfrom ecg_utils import (\n    LABEL_POSITIONS, TEMPLATE_SIZE, PRECISION_LEADS, TYPE_IDS, LEAD_NAMES_BY_ROW,\n    refine_keypoint, validation_subset\n)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:26.128084Z","iopub.execute_input":"2026-01-29T09:32:26.128946Z","iopub.status.idle":"2026-01-29T09:32:40.406974Z","shell.execute_reply.started":"2026-01-29T09:32:26.128908Z","shell.execute_reply":"2026-01-29T09:32:40.406277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nConfiguration for inference3 pipeline.\n\"\"\"\n\n# Paths\nKAGGLE_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nECGLIB_DIR = '/kaggle/input/ecg-lib'\n\n# Model paths\nSTAGE0_MODEL_PATH = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage0-last.checkpoint.pth'\nSTAGE1_MODEL_PATH = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage1-last.checkpoint.pth'\nSTAGE2_MODEL_PATH = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage2-00005810.checkpoint.pth'\n# Single model (legacy):\n# STAGED_MODEL_PATH = '/kaggle/input/stage-d-models/best_model.pth'\n# Ensemble (list of paths):\nSTAGED_MODEL_PATHS = [\n    '/kaggle/input/stage-d-models/best_model-aug-20w-nn-convnext-cdr0.01-2x1-f1.pth',\n]\nREFINEMENT_MODEL_PATHS = [\n    '/kaggle/input/stage-c-refinement-models/best_model-stage1-refinement-v2-phase2a-31x31.pth',\n    '/kaggle/input/stage-c-refinement-models/best_model-stage1-refinement-v2-phase2b-31x31.pth',\n]\nREFINEMENT_PATCH_SIZE = 31  # Must match trained model's patch_size\nUSE_LEGACY_REFINEMENT = False  # True for old 4-layer model, False for dynamic RF model\n\n# Stage C options\n# USE_ADV_KP_ALGOS: Use Advanced Keypoint Algorithms (outlier detection, recovery)\n# USE_GP_REFINEMENT: Gridpoint refinement method\n#   None = no refinement\n#   'nn' = neural network refinement (trained on type 0001 only)\n#   'corr' = template correlation refinement\n# USE_CANONICAL_HR: Use high-res canonical (2400x1920) for rectification\n#   True = rectify 2400x1920 canonical, upscale gridpoints\n#   False = rectify 1440x1152 stage1 input (hengck23's original)\nUSE_ADV_KP_ALGOS = True\nUSE_GP_REFINEMENT = 'nn'\nUSE_CANONICAL_HR = True  # False to reproduce hengck23's ~16 dB result\n\n# Correlation refinement rotation angles (degrees)\n# None or [] = no rotation, just use template as-is\n# List of angles = try each, pick best correlation (GPU batched)\nCORR_ROTATION_ANGLES = None #[-3, -2, -1, 0, 1, 2, 3]\n\n# Stage D model selection\nUSE_STAGEH = False  # True = hengck23's model, False = custom Stage D\n\n# Stage D encoder backbone\nENCODER = 'convnextv2_base'  # 'resnet34', 'convnextv2_base', 'convnextv2_tiny', 'efficientnetv2_m'\n\n# Stage D resolution scaling (base is 2200x1700)\n# Must match the model being used\nX_SCALE = 2  # 2, 4, or 8\nY_SCALE = 1  # 1 or 2\n\n# Signal extraction method: 'softmax', 'adjacent', 'parabolic', 'centroid'\nEXTRACTION_METHOD = 'centroid'\n\n# Centroid extraction window (pixels). Should scale with Y_SCALE to maintain consistent mV coverage.\nCENTROID_RANGE = 100 * Y_SCALE\n\n# Lead II ensembling: average II_short with Lead II first quarter\n# Set False for Stage R training data generation (need raw predictions)\nUSE_ENSEMBLE = True\n\n# Stage E: Post-processing (Savgol filter + Einthoven correction)\n# If USE_STAGE_E = True, USE_STAGE_R is ignored\nUSE_STAGE_E = False\n\n# Stage R: Signal refinement network\nUSE_STAGE_R = False\nSTAGE_R_MODEL_PATH = '/kaggle/input/stage-r-models/best_model-stage_r_v5.pth'\n\n# Validation options\n#TYPE_IDS_TO_RUN = ['0001']  # Which types to validate\n#IMAGE_IDS_TO_RUN = None  # List of specific image IDs, or None for all\nTYPE_IDS_TO_RUN = None  # Which types to validate\nIMAGE_IDS_TO_RUN = ['2368837031',  '1084993373',  '452845711',  '2101175952',  '91831581',\n                    '49746380',  '2729579623',  '565587895',  '314575601',  '3753632017']\n# IMAGE_IDS_TO_RUN = ['1135737846',  '1512936796',  '1590643291',  '170158778',  '2063154718',  '2221052980', \n#                     '2933652299',  '3434428361',  '3511927436',  '3944034484', '4050367691', '937503588']\nFOLD_TO_RUN = None  # Fold number (1-10) to validate, or None for all\nSPLIT_CSV = '/kaggle/input/ecg-lib/split.csv'  # Path to split.csv with fold column\n\n# Output options\nSAVE_RECTIFIED = True\nRECTIFIED_DIR = 'rectified'\nSAVE_CANONICAL = True\nCANONICAL_DIR = 'canonical'\nSAVE_GRIDPOINTS = True\nGRIDPOINTS_DIR = 'gridpoints'\n\n# Stage R training data output\n# Saves pred/solution CSVs for Stage R training\n# Set USE_ENSEMBLE=False when generating this data\nSAVE_STAGED_PREDS = False\nSTAGED_PREDS_DIR = 'staged_preds'\n\n# Stage S training data output\n# Saves raw 4x3925 predictions as .npz for Stage S (learned resampling)\nSAVE_RAW_PREDS = False\nRAW_PREDS_DIR = 'raw_preds'\n\n# Synthetic image processing\n# When True, processes synthetic images instead of competition images\nSYNTHETIC = False\nSYNTHETIC_DIR = '/home/ubuntu/PhysioNet/synthetic'\nSYNTHETIC_START_ID = '000000'  # Inclusive\nSYNTHETIC_END_ID = '010000'    # Inclusive\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.408592Z","iopub.execute_input":"2026-01-29T09:32:40.409198Z","iopub.status.idle":"2026-01-29T09:32:40.417964Z","shell.execute_reply.started":"2026-01-29T09:32:40.409169Z","shell.execute_reply":"2026-01-29T09:32:40.417157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Scoring functions from PhysioNet ECG Digitization competition\n\nfrom typing import Tuple\n\nimport numpy as np\nimport pandas as pd\n\nimport scipy.optimize\nimport scipy.signal\n\n\nLEADS = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\nMAX_TIME_SHIFT = 0.2\nPERFECT_SCORE = 384\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef compute_power(label: np.ndarray, prediction: np.ndarray) -> Tuple[float, float]:\n    if label.ndim != 1 or prediction.ndim != 1:\n        raise ParticipantVisibleError('Inputs must be 1-dimensional arrays.')\n    finite_mask = np.isfinite(prediction)\n    if not np.any(finite_mask):\n        raise ParticipantVisibleError(\"The 'prediction' array contains no finite values (all NaN or inf).\")\n\n    prediction[~np.isfinite(prediction)] = 0\n    noise = label - prediction\n    p_signal = np.sum(label**2)\n    p_noise = np.sum(noise**2)\n    return p_signal, p_noise\n\n\ndef compute_snr(signal: float, noise: float) -> float:\n    if noise == 0:\n        # Perfect reconstruction\n        snr = PERFECT_SCORE\n    elif signal == 0:\n        snr = 0\n    else:\n        snr = min((signal / noise), PERFECT_SCORE)\n    return snr\n\n\ndef align_signals(label: np.ndarray, pred: np.ndarray, max_shift: float = float('inf')) -> np.ndarray:\n    if np.any(~np.isfinite(label)):\n        raise ParticipantVisibleError('values in label should all be finite')\n    if np.sum(np.isfinite(pred)) == 0:\n        raise ParticipantVisibleError('prediction can not all be infinite')\n\n    # Initialize the reference and digitized signals.\n    label_arr = np.asarray(label, dtype=np.float64)\n    pred_arr = np.asarray(pred, dtype=np.float64)\n\n    label_mean = np.mean(label_arr)\n    pred_mean = np.mean(pred_arr)\n\n    label_arr_centered = label_arr - label_mean\n    pred_arr_centered = pred_arr - pred_mean\n\n    # Compute the correlation between the reference and digitized signals and locate the maximum correlation.\n    correlation = scipy.signal.correlate(label_arr_centered, pred_arr_centered, mode='full')\n\n    n_label = np.size(label_arr)\n    n_pred = np.size(pred_arr)\n\n    lags = scipy.signal.correlation_lags(n_label, n_pred, mode='full')\n    valid_lags_mask = (lags >= -max_shift) & (lags <= max_shift)\n\n    max_correlation = np.nanmax(correlation[valid_lags_mask])\n    all_max_indices = np.flatnonzero(correlation == max_correlation)\n    best_idx = min(all_max_indices, key=lambda i: abs(lags[i]))\n    time_shift = lags[best_idx]\n    start_padding_len = max(time_shift, 0)\n    pred_slice_start = max(-time_shift, 0)\n    pred_slice_end = min(n_label - time_shift, n_pred)\n    end_padding_len = max(n_label - n_pred - time_shift, 0)\n    aligned_pred = np.concatenate((np.full(start_padding_len, np.nan), pred_arr[pred_slice_start:pred_slice_end], np.full(end_padding_len, np.nan)))\n\n    def objective_func(v_shift):\n        return np.nansum((label_arr - (aligned_pred - v_shift)) ** 2)\n\n    if np.any(np.isfinite(label_arr) & np.isfinite(aligned_pred)):\n        results = scipy.optimize.minimize_scalar(objective_func, method='Brent')\n        vertical_shift = results.x\n        aligned_pred -= vertical_shift\n    return aligned_pred\n\n\ndef _calculate_image_score(group: pd.DataFrame) -> float:\n    \"\"\"Helper function to calculate the total SNR score for a single image group.\"\"\"\n\n    unique_fs_values = group['fs'].unique()\n    if len(unique_fs_values) != 1:\n        raise ParticipantVisibleError('Sampling frequency should be consistent across each ecg')\n    sampling_frequency = unique_fs_values[0]\n    if sampling_frequency != int(len(group[group['lead'] == 'II']) / 10):\n        raise ParticipantVisibleError('The sequence_length should be sampling frequency * 10s')\n    sum_signal = 0\n    sum_noise = 0\n    for lead in LEADS:\n        sub = group[group['lead'] == lead]\n        label = sub['value_true'].values\n        pred = sub['value_pred'].values\n\n        aligned_pred = align_signals(label, pred, int(sampling_frequency * MAX_TIME_SHIFT))\n        p_signal, p_noise = compute_power(label, aligned_pred)\n        sum_signal += p_signal\n        sum_noise += p_noise\n    return compute_snr(sum_signal, sum_noise)\n\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str) -> float:\n    \"\"\"\n    Compute the mean Signal-to-Noise Ratio (SNR) across multiple ECG leads and images for the PhysioNet 2025 competition.\n    The final score is the average of the sum of SNRs over different lines, averaged over all unique images.\n    Args:\n        solution: DataFrame with ground truth values. Expected columns: 'id' and one for each lead.\n        submission: DataFrame with predicted values. Expected columns: 'id' and one for each lead.\n        row_id_column_name: The name of the unique identifier column, typically 'id'.\n    Returns:\n        The final competition score.\n\n    Examples\n    --------\n    >>> import pandas as pd\n    >>> import numpy as np\n    >>> row_id_column_name = \"id\"\n    >>> solution = pd.DataFrame({'id': ['343_0_I', '343_1_I', '343_2_I', '343_0_III', '343_1_III','343_2_III','343_0_aVR', '343_1_aVR','343_2_aVR',\\\n    '343_0_aVL', '343_1_aVL', '343_2_aVL', '343_0_aVF', '343_1_aVF','343_2_aVF','343_0_V1', '343_1_V1', '343_2_V1','343_0_V2', '343_1_V2','343_2_V2',\\\n    '343_0_V3', '343_1_V3', '343_2_V3','343_0_V4', '343_1_V4', '343_2_V4', '343_0_V5', '343_1_V5','343_2_V5','343_0_V6', '343_1_V6','343_2_V6',\\\n    '343_0_II', '343_1_II','343_2_II', '343_3_II', '343_4_II', '343_5_II','343_6_II', '343_7_II','343_8_II','343_9_II','343_10_II','343_11_II'],\\\n    'fs': [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1],\\\n    'value':[0.1,0.3,0.4,0.6,0.6,0.4,0.2,0.3,0.4,0.5,0.2,0.7,0.2,0.3,0.4,0.8,0.6,0.7, 0.2,0.3,-0.1,0.5,0.6,0.7,0.2,0.9,0.4,0.5,0.6,0.7,0.1,0.3,0.4,\\\n    0.6,0.6,0.4,0.2,0.3,0.4,0.5,0.2,0.7,0.2,0.3,0.4]})\n    >>> submission = solution.copy()\n    >>> round(score(solution, submission, row_id_column_name), 4)\n    25.8433\n    >>> submission.loc[0, 'value'] = 0.9 # Introduce some noise\n    >>> round(score(solution, submission, row_id_column_name), 4)\n    13.6291\n    >>> submission.loc[4, 'value'] = 0.3 # Introduce some noise\n    >>> round(score(solution, submission, row_id_column_name), 4)\n    13.0576\n\n    >>> solution = pd.DataFrame({'id': ['343_0_I', '343_1_I', '343_2_I', '343_0_III', '343_1_III','343_2_III','343_0_aVR', '343_1_aVR','343_2_aVR',\\\n    '343_0_aVL', '343_1_aVL', '343_2_aVL', '343_0_aVF', '343_1_aVF','343_2_aVF','343_0_V1', '343_1_V1', '343_2_V1','343_0_V2', '343_1_V2','343_2_V2',\\\n    '343_0_V3', '343_1_V3', '343_2_V3','343_0_V4', '343_1_V4', '343_2_V4', '343_0_V5', '343_1_V5','343_2_V5','343_0_V6', '343_1_V6','343_2_V6',\\\n    '343_0_II', '343_1_II','343_2_II', '343_3_II', '343_4_II', '343_5_II','343_6_II', '343_7_II','343_8_II','343_9_II','343_10_II','343_11_II'],\\\n    'fs': [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1],\\\n    'value':[0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0]})\n    >>> round(score(solution, submission, row_id_column_name), 4)\n    -384\n    >>> submission = solution.copy()\n    >>> round(score(solution, submission, row_id_column_name), 4)\n    25.8433\n\n    >>> # test alignment\n    >>> label = np.array([0, 1, 2, 1, 0])\n    >>> pred = np.array([0, 1, 2, 1, 0])\n    >>> aligned = align_signals(label, pred)\n    >>> expected_array = np.array([0, 1, 2, 1, 0])\n    >>> np.allclose(aligned, expected_array, equal_nan=True)\n    True\n\n    >>> # Test 2: Vertical shift (DC offset) should be removed\n    >>> label = np.array([0, 1, 2, 1, 0])\n    >>> pred = np.array([10, 11, 12, 11, 10])\n    >>> aligned = align_signals(label, pred)\n    >>> expected_array = np.array([0, 1, 2, 1, 0])\n    >>> np.allclose(aligned, expected_array, equal_nan=True)\n    True\n\n    >>> # Test 3: Time shift should be corrected\n    >>> label = np.array([0, 0, 1, 2, 1, 0., 0.])\n    >>> pred = np.array([1, 2, 1, 0, 0, 0, 0])\n    >>> aligned = align_signals(label, pred)\n    >>> expected_array = np.array([np.nan, np.nan, 1, 2, 1, 0, 0])\n    >>> np.allclose(aligned, expected_array, equal_nan=True)\n    True\n\n    >>> # Test 4: max_shift constraint prevents optimal alignment\n    >>> label = np.array([0, 0, 0, 0, 1, 2, 1]) # Peak is far\n    >>> pred = np.array([1, 2, 1, 0, 0, 0, 0])\n    >>> aligned = align_signals(label, pred, max_shift=10)\n    >>> expected_array = np.array([ np.nan, np.nan, np.nan, np.nan, 1, 2, 1])\n    >>> np.allclose(aligned, expected_array, equal_nan=True)\n    True\n\n    \"\"\"\n    for df in [solution, submission]:\n        if row_id_column_name not in df.columns:\n            raise ParticipantVisibleError(f\"'{row_id_column_name}' column not found in DataFrame.\")\n        if df['value'].isna().any():\n            raise ParticipantVisibleError('NaN exists in solution/submission')\n        if not np.isfinite(df['value']).all():\n            raise ParticipantVisibleError('Infinity exists in solution/submission')\n\n    submission = submission[['id', 'value']]\n    merged_df = pd.merge(solution, submission, on=row_id_column_name, suffixes=('_true', '_pred'))\n    # Parse id format: {image_id}_{timestep}_{lead} - image_id may contain underscores\n    # Split from right: lead is last, timestep is second-to-last, rest is image_id\n    merged_df['lead'] = merged_df[row_id_column_name].str.rsplit('_', n=1).str[-1]\n    merged_df['_temp'] = merged_df[row_id_column_name].str.rsplit('_', n=1).str[0]\n    merged_df['row_id'] = merged_df['_temp'].str.rsplit('_', n=1).str[-1].astype('int64')\n    merged_df['image_id'] = merged_df['_temp'].str.rsplit('_', n=1).str[0]\n    merged_df.drop(columns=['_temp'], inplace=True)\n    merged_df.sort_values(by=['image_id', 'row_id', 'lead'], inplace=True)\n    image_scores = merged_df.groupby('image_id').apply(_calculate_image_score, include_groups=False)\n    return max(float(10 * np.log10(image_scores.mean())), -PERFECT_SCORE)\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.419280Z","iopub.execute_input":"2026-01-29T09:32:40.419657Z","iopub.status.idle":"2026-01-29T09:32:40.843627Z","shell.execute_reply.started":"2026-01-29T09:32:40.419613Z","shell.execute_reply":"2026-01-29T09:32:40.843068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage A: Orientation correction and keypoint detection.\n\nInput: Original image (any orientation)\nOutput: Upright image, keypoints\n\nUses hengck23's stage0 model and helper functions.\n\"\"\"\n\nimport torch\nfrom stage0_common import image_to_batch, output_to_predict\n\n\ndef process_stage_a(image, stage0_net):\n    \"\"\"\n    Stage A: Correct image orientation and extract keypoints.\n\n    Args:\n        image: RGB image (H, W, 3) numpy array\n        stage0_net: Loaded Stage 0 model\n\n    Returns:\n        rotated: Upright image (H, W, 3) numpy array\n        keypoint: List of [x, y, label, leadname] for each detected keypoint\n    \"\"\"\n    batch = image_to_batch(image)\n\n    with torch.amp.autocast('cuda'):\n        with torch.no_grad():\n            output = stage0_net(batch)\n\n    rotated, keypoint = output_to_predict(image, batch, output)\n\n    torch.cuda.empty_cache()\n\n    return rotated, keypoint\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.845342Z","iopub.execute_input":"2026-01-29T09:32:40.845784Z","iopub.status.idle":"2026-01-29T09:32:40.851227Z","shell.execute_reply.started":"2026-01-29T09:32:40.845756Z","shell.execute_reply":"2026-01-29T09:32:40.850454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage B: Homography to canonical size.\n\nInput: Upright image from Stage A, keypoints from Stage A\nOutput: Canonical images warped via homography\n    - canonical_lr (1440x1152): Always created using hengck23's REF_PT9 method\n    - canonical_hr (2400x1920): Only if USE_CANONICAL_HR=True, using our LABEL_POSITIONS\n\nNo correlation refinement - uses keypoints directly from stage0.\n\"\"\"\n\nimport os\nimport cv2\nimport numpy as np\nfrom ecg_utils import LABEL_POSITIONS, PRECISION_LEADS\nfrom stage0_common import normalise_by_homography\n# Note: USE_CANONICAL_HR from config (defined in config cell)\n\n# Canonical output sizes\nCANONICAL_SIZE_HR = (2400, 1920)  # (width, height) - high-res\nCANONICAL_SIZE_LR = (1440, 1152)  # (width, height) - hengck23's original\n\n# Padding for high-res (to center 2200x1700 content in 2400x1920)\nPADDING_X_HR = (2400 - 2200) // 2  # 100px on each side\nPADDING_Y_HR = (1920 - 1700) // 2  # 110px on each side\n\n\n# -----------------------------------------------------------------------------\n# hengck23's REF_PT9 for canonical_lr (from stage0_common.py)\n# -----------------------------------------------------------------------------\n\n# Load reference gridpoints from hengck23's data\n# Note: __file__ not available in notebook cells, so use Kaggle path directly\nGRIDPOINT_PATH_KAGGLE = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/640106434-0001.gridpoint_xy.npy'\n\nif os.path.exists(GRIDPOINT_PATH_KAGGLE):\n    gridpoint0001_xy = np.load(GRIDPOINT_PATH_KAGGLE)\nelse:\n    # Local development fallback - adjust path as needed\n    gridpoint0001_xy = np.load('archive2/hengck23-submit-physionet/640106434-0001.gridpoint_xy.npy')\n\n\ndef make_ref_pt9():\n    \"\"\"\n    Create hengck23's 9-point reference for homography (from stage0_common.py).\n\n    Uses specific gridpoints from type 0001 reference, scaled and shifted\n    to center content in 1440x1152 output.\n    \"\"\"\n    h0001, w0001 = 1700, 2200\n    ref_pt = []\n    for j, i in [\n        [19, 3],\n        [26, 3],\n        [33, 3],\n    ]:\n        x, y = gridpoint0001_xy[j, i + 13]\n        ref_pt.append([x, y])\n        x, y = gridpoint0001_xy[j, i + 25]\n        ref_pt.append([x, y])\n        x, y = gridpoint0001_xy[j, i + 38]\n        ref_pt.append([x, y])\n\n    ref_pt = np.array(ref_pt, np.float32)\n    scale = 1280 / w0001\n    ref_pt = ref_pt * [[scale, scale]]\n    shift = (1440 - 1280) / 2\n    ref_pt = ref_pt + [[shift, shift]] + [[-6, +10]]\n    return ref_pt\n\n\nREF_PT9 = make_ref_pt9()\n\n# Mapping from label to index in REF_PT9 (order matches hengck23's loop)\nLABEL_TO_REF_IDX = {\n    2: 0,   # aVR\n    3: 1,   # V1\n    4: 2,   # V4\n    6: 3,   # aVL\n    7: 4,   # V2\n    8: 5,   # V5\n    10: 6,  # aVF\n    11: 7,  # V3\n    12: 8,  # V6\n}\n\n\n# -----------------------------------------------------------------------------\n# Keypoint conversion\n# -----------------------------------------------------------------------------\n\ndef keypoints_to_dict(keypoint_list):\n    \"\"\"\n    Convert hengck23's keypoint format to dict.\n\n    Args:\n        keypoint_list: List of [x, y, label, leadname]\n\n    Returns:\n        keypoints: dict mapping leadname to (x, y, label)\n    \"\"\"\n    keypoints = {}\n    for kp in keypoint_list:\n        x, y, label, leadname = kp\n        keypoints[leadname] = (x, y, label)\n    return keypoints\n\n\n# -----------------------------------------------------------------------------\n# Homography for canonical_lr (hengck23's method)\n# -----------------------------------------------------------------------------\n\ndef compute_homography_lr(keypoint_list):\n    \"\"\"\n    Compute homography using hengck23's 9-point method for canonical_lr.\n\n    Args:\n        keypoint_list: List of [x, y, label, leadname] from Stage A\n\n    Returns:\n        H: 3x3 homography matrix\n        match: array indicating RANSAC inliers\n    \"\"\"\n    # Extract 9 keypoints sorted by REF_PT9 index order\n    pt9 = [None] * 9\n    for kp in keypoint_list:\n        x, y, label, leadname = kp\n        if label in LABEL_TO_REF_IDX:\n            idx = LABEL_TO_REF_IDX[label]\n            pt9[idx] = [x, y]\n\n    # Filter out missing keypoints and get corresponding REF_PT9 entries\n    src_pts = []\n    dst_pts = []\n    for idx, pt in enumerate(pt9):\n        if pt is not None:\n            src_pts.append(pt)\n            dst_pts.append(REF_PT9[idx])\n\n    src_pts = np.array(src_pts, np.float32)\n    dst_pts = np.array(dst_pts, np.float32)\n\n    if len(src_pts) < 4:\n        raise ValueError(f\"Not enough keypoints for homography: {len(src_pts)}\")\n\n    H, match = cv2.findHomography(src_pts, dst_pts, method=cv2.RANSAC)\n    return H, match\n\n\ndef warp_to_canonical_lr(image, H):\n    \"\"\"Warp image to canonical_lr size (1440x1152) using homography.\"\"\"\n    canonical_w, canonical_h = CANONICAL_SIZE_LR\n    return cv2.warpPerspective(image, H, (canonical_w, canonical_h))\n\n\n# -----------------------------------------------------------------------------\n# Homography for canonical_hr (our method)\n# -----------------------------------------------------------------------------\n\ndef compute_homography_hr(keypoints):\n    \"\"\"\n    Compute homography using our LABEL_POSITIONS for canonical_hr.\n\n    Args:\n        keypoints: dict mapping lead name to (x, y, label)\n\n    Returns:\n        H: 3x3 homography matrix\n        inliers: number of RANSAC inliers\n    \"\"\"\n    src_points = []\n    dst_points = []\n\n    for lead_name, pos in keypoints.items():\n        if lead_name not in PRECISION_LEADS:\n            continue\n\n        x, y = pos[0], pos[1]\n        if x == 0 and y == 0:\n            continue\n\n        src_points.append([x, y])\n\n        template_pos = LABEL_POSITIONS[lead_name]\n        dst_points.append([template_pos['col'] + PADDING_X_HR, template_pos['row'] + PADDING_Y_HR])\n\n    if len(src_points) < 4:\n        raise ValueError(f\"Not enough keypoints for homography: {len(src_points)}\")\n\n    src_points = np.array(src_points, dtype=np.float32)\n    dst_points = np.array(dst_points, dtype=np.float32)\n\n    H, mask = cv2.findHomography(src_points, dst_points, cv2.RANSAC, 5.0)\n    inliers = mask.sum() if mask is not None else 0\n\n    return H, inliers\n\n\ndef warp_to_canonical_hr(image, H):\n    \"\"\"Warp image to canonical_hr size (2400x1920) using homography.\"\"\"\n    canonical_w, canonical_h = CANONICAL_SIZE_HR\n    return cv2.warpPerspective(image, H, (canonical_w, canonical_h))\n\n\n# -----------------------------------------------------------------------------\n# Main processing function\n# -----------------------------------------------------------------------------\n\ndef process_stage_b(image, keypoint_list):\n    \"\"\"\n    Stage B: Warp to canonical size using keypoints from Stage A.\n\n    Args:\n        image: RGB image (H, W, 3) numpy array (upright from Stage A)\n        keypoint_list: List of [x, y, label, leadname] from Stage A\n\n    Returns:\n        canonical_lr: RGB image (1152, 1440, 3) - always created (hengck23's method)\n        canonical_hr: RGB image (1920, 2400, 3) or None - only if USE_CANONICAL_HR\n        keypoints: dict mapping lead name to (x, y, label)\n        H_lr: 3x3 homography matrix for canonical_lr\n        H_hr: 3x3 homography matrix for canonical_hr (or None)\n    \"\"\"\n    # Convert keypoint format\n    keypoints = keypoints_to_dict(keypoint_list)\n\n    # Always compute canonical_lr using hengck23's exact method\n    canonical_lr, _, H_lr = normalise_by_homography(image, keypoint_list)\n\n    # Optionally compute canonical_hr using our method\n    if USE_CANONICAL_HR:\n        H_hr, inliers_hr = compute_homography_hr(keypoints)\n        canonical_hr = warp_to_canonical_hr(image, H_hr)\n    else:\n        H_hr = None\n        canonical_hr = None\n\n    return canonical_lr, canonical_hr, keypoints, H_lr, H_hr\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.852370Z","iopub.execute_input":"2026-01-29T09:32:40.852652Z","iopub.status.idle":"2026-01-29T09:32:40.877171Z","shell.execute_reply.started":"2026-01-29T09:32:40.852628Z","shell.execute_reply":"2026-01-29T09:32:40.876462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage C: Gridpoint detection and image rectification.\n\nInput:\n    - canonical_lr (1440x1152): Always provided, used for gridpoint detection\n    - canonical_hr (2400x1920): Optional, used for rectification if USE_CANONICAL_HR=True\nOutput: Rectified image (2200x1700), gridpoint_xy arrays\n\nModes (controlled by config):\n- USE_CANONICAL_HR=False: Rectify canonical_lr (hengck23's original approach)\n- USE_CANONICAL_HR=True: Rectify canonical_hr (preserves resolution)\n- USE_ADV_KP_ALGOS=True: Add outlier detection and recovery\n- USE_GP_REFINEMENT: Gridpoint refinement method\n    None = no refinement\n    'nn' = neural network refinement\n    'corr' = template correlation refinement\n\"\"\"\n\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom scipy.interpolate import griddata, RBFInterpolator\nfrom stage1_common import output_to_predict as stage1_output_to_predict\n# Note: USE_ADV_KP_ALGOS, USE_GP_REFINEMENT, USE_CANONICAL_HR, CORR_ROTATION_ANGLES, ECGLIB_DIR from config (defined in config cell)\n\n# Stage 1 model input size\nSTAGE1_WIDTH = 1440\nSTAGE1_HEIGHT = 1152\n\n# Output size after rectification\nOUTPUT_WIDTH = 2200\nOUTPUT_HEIGHT = 1700\n\n# Template for correlation refinement (type 0001 reference)\n# Note: ECGLIB_DIR defined in config cell\n# TEMPLATE_GRIDPOINT_XY loaded in pipeline.py, passed to process_stage_c\n# TEMPLATE_IMAGE loaded in pipeline.py, passed to process_stage_c\n\n\n# -----------------------------------------------------------------------------\n# Core functions (always used)\n# -----------------------------------------------------------------------------\n\ndef downscale_for_stage1(canonical):\n    \"\"\"Downscale canonical image to stage1 input size.\"\"\"\n    return cv2.resize(canonical, (STAGE1_WIDTH, STAGE1_HEIGHT), interpolation=cv2.INTER_LINEAR)\n\n\ndef upscale_gridpoints(gridpoint_xy, src_size, dst_size):\n    \"\"\"Upscale gridpoint coordinates from stage1 size to canonical size.\"\"\"\n    src_w, src_h = src_size\n    dst_w, dst_h = dst_size\n\n    scale_x = dst_w / src_w\n    scale_y = dst_h / src_h\n\n    scaled_xy = gridpoint_xy.copy()\n    scaled_xy[..., 0] *= scale_x\n    scaled_xy[..., 1] *= scale_y\n\n    return scaled_xy\n\n\ndef transform_gridpoints(gridpoint_xy, H):\n    \"\"\"Transform gridpoint coordinates using homography matrix.\"\"\"\n    rows, cols = gridpoint_xy.shape[:2]\n    transformed_xy = gridpoint_xy.copy()\n\n    for j in range(rows):\n        for i in range(cols):\n            x, y = gridpoint_xy[j, i]\n            if x == 0 and y == 0:\n                continue  # Skip missing gridpoints\n\n            # Homogeneous coordinates\n            p = np.array([x, y, 1.0])\n            p_transformed = H @ p\n\n            # Convert back from homogeneous\n            if p_transformed[2] != 0:\n                transformed_xy[j, i, 0] = p_transformed[0] / p_transformed[2]\n                transformed_xy[j, i, 1] = p_transformed[1] / p_transformed[2]\n\n    return transformed_xy\n\n\ndef interpolate_mapping(gridpoint_xy):\n    \"\"\"Interpolate missing gridpoints using cubic interpolation.\"\"\"\n    mx, my = np.meshgrid(np.arange(0, 57), np.arange(0, 44))\n    coord = np.stack([mx, my], axis=-1).reshape(-1, 2)\n    value = gridpoint_xy.copy().reshape(-1, 2)\n    missing = np.all(value == [0, 0], axis=1)\n\n    if missing.sum() == len(missing):\n        return gridpoint_xy\n\n    interpolate_xy = griddata(coord[~missing], value[~missing], (mx, my), method='cubic')\n    interpolate_xy[np.isnan(interpolate_xy)] = 0\n    return interpolate_xy\n\n\ndef save_gridpoints_csv(gridpoint_xy, csv_path):\n    \"\"\"Save gridpoint coordinates to CSV.\"\"\"\n    import pandas as pd\n    rows = []\n    for j in range(44):\n        for i in range(57):\n            x, y = gridpoint_xy[j, i]\n            if x == 0 and y == 0:\n                continue  # Skip missing gridpoints\n            rows.append({'j': j, 'i': i, 'x': x, 'y': y})\n    df = pd.DataFrame(rows)\n    df.to_csv(csv_path, index=False)\n\n\ndef rectify_image(image, gridpoint_xy):\n    \"\"\"Rectify image using gridpoint mapping.\"\"\"\n    H, W = OUTPUT_HEIGHT, OUTPUT_WIDTH\n    H1, W1 = image.shape[:2]\n\n    sparse_map = gridpoint_xy / [[[W1-1, H1-1]]] * 2 - 1\n    sparse_map = torch.from_numpy(np.ascontiguousarray(sparse_map.transpose(2, 0, 1))).unsqueeze(0).float()\n    dense_map = F.interpolate(sparse_map, size=(H, W), mode='bilinear', align_corners=True)\n    distort = torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0).float()\n    rectified = F.grid_sample(\n        distort, dense_map.permute(0, 2, 3, 1), mode='bilinear', padding_mode='border', align_corners=False\n    )\n    rectified = rectified.data.cpu().numpy()\n    rectified = rectified[0].transpose(1, 2, 0).astype(np.uint8)\n    return rectified\n\n\n# -----------------------------------------------------------------------------\n# Advanced keypoint algorithms (when USE_ADV_KP_ALGOS=True)\n# -----------------------------------------------------------------------------\n\ndef detect_outliers(gridpoint_xy, threshold=20.0):\n    \"\"\"Detect gridpoint outliers by checking if each point fits local grid spacing.\"\"\"\n    cleaned_xy = gridpoint_xy.copy()\n    rows, cols = 44, 57\n    num_outliers = 0\n\n    def is_valid(j, i):\n        if j < 0 or j >= rows or i < 0 or i >= cols:\n            return False\n        return not (cleaned_xy[j, i, 0] == 0 and cleaned_xy[j, i, 1] == 0)\n\n    def estimate_from_neighbors(j, i):\n        estimates = []\n        # From left neighbor\n        if is_valid(j, i-1) and is_valid(j, i-2):\n            spacing = cleaned_xy[j, i-1] - cleaned_xy[j, i-2]\n            estimates.append(cleaned_xy[j, i-1] + spacing)\n        # From right neighbor\n        if is_valid(j, i+1) and is_valid(j, i+2):\n            spacing = cleaned_xy[j, i+1] - cleaned_xy[j, i+2]\n            estimates.append(cleaned_xy[j, i+1] + spacing)\n        # From top neighbor\n        if is_valid(j-1, i) and is_valid(j-2, i):\n            spacing = cleaned_xy[j-1, i] - cleaned_xy[j-2, i]\n            estimates.append(cleaned_xy[j-1, i] + spacing)\n        # From bottom neighbor\n        if is_valid(j+1, i) and is_valid(j+2, i):\n            spacing = cleaned_xy[j+1, i] - cleaned_xy[j+2, i]\n            estimates.append(cleaned_xy[j+1, i] + spacing)\n        if estimates:\n            return np.mean(estimates, axis=0)\n        return None\n\n    for j in range(rows):\n        for i in range(cols):\n            # Skip edges - only check interior gridpoints\n            if j == 0 or j == 43 or i == 0 or i == 56:\n                continue\n            if not is_valid(j, i):\n                continue\n            actual = cleaned_xy[j, i]\n            expected = estimate_from_neighbors(j, i)\n            if expected is not None:\n                dist = np.linalg.norm(actual - expected)\n                if dist > threshold:\n                    cleaned_xy[j, i] = [0, 0]\n                    num_outliers += 1\n\n    return cleaned_xy, num_outliers\n\n\ndef recover_missing_gridpoints(gridpoint_xy):\n    \"\"\"Recover missing gridpoints using cascaded interpolation/extrapolation.\"\"\"\n    recovered_xy = gridpoint_xy.copy()\n    rows, cols = 44, 57\n\n    def is_valid(j, i):\n        if j < 0 or j >= rows or i < 0 or i >= cols:\n            return False\n        return not (recovered_xy[j, i, 0] == 0 and recovered_xy[j, i, 1] == 0)\n\n    def get_missing_indices():\n        missing = []\n        for j in range(rows):\n            for i in range(cols):\n                # Skip edges - don't recover edge gridpoints\n                if j == 0 or j == 43 or i == 0 or i == 56:\n                    continue\n                if not is_valid(j, i):\n                    missing.append((j, i))\n        return missing\n\n    # Pass 1: Interior interpolation (both axes)\n    for j, i in get_missing_indices():\n        if is_valid(j, i-1) and is_valid(j, i+1) and is_valid(j-1, i) and is_valid(j+1, i):\n            row_interp = (recovered_xy[j, i-1] + recovered_xy[j, i+1]) / 2\n            col_interp = (recovered_xy[j-1, i] + recovered_xy[j+1, i]) / 2\n            recovered_xy[j, i] = (row_interp + col_interp) / 2\n\n    # Pass 2: Single-axis interpolation\n    for j, i in get_missing_indices():\n        if is_valid(j, i-1) and is_valid(j, i+1):\n            recovered_xy[j, i] = (recovered_xy[j, i-1] + recovered_xy[j, i+1]) / 2\n        elif is_valid(j-1, i) and is_valid(j+1, i):\n            recovered_xy[j, i] = (recovered_xy[j-1, i] + recovered_xy[j+1, i]) / 2\n\n    # Pass 3: Local extrapolation\n    for j, i in get_missing_indices():\n        if is_valid(j, i+1) and is_valid(j, i+2):\n            spacing = recovered_xy[j, i+1] - recovered_xy[j, i+2]\n            recovered_xy[j, i] = recovered_xy[j, i+1] + spacing\n        elif is_valid(j, i-1) and is_valid(j, i-2):\n            spacing = recovered_xy[j, i-1] - recovered_xy[j, i-2]\n            recovered_xy[j, i] = recovered_xy[j, i-1] + spacing\n        elif is_valid(j+1, i) and is_valid(j+2, i):\n            spacing = recovered_xy[j+1, i] - recovered_xy[j+2, i]\n            recovered_xy[j, i] = recovered_xy[j+1, i] + spacing\n        elif is_valid(j-1, i) and is_valid(j-2, i):\n            spacing = recovered_xy[j-1, i] - recovered_xy[j-2, i]\n            recovered_xy[j, i] = recovered_xy[j-1, i] + spacing\n\n    # Pass 4: Diagonal interpolation/extrapolation\n    for j, i in get_missing_indices():\n        diag_points = []\n        for dj, di in [(-1, -1), (-1, 1), (1, -1), (1, 1)]:\n            if is_valid(j+dj, i+di):\n                if is_valid(j-dj, i-di):\n                    est = (recovered_xy[j+dj, i+di] + recovered_xy[j-dj, i-di]) / 2\n                    diag_points.append(est)\n                elif is_valid(j+2*dj, i+2*di):\n                    spacing = recovered_xy[j+dj, i+di] - recovered_xy[j+2*dj, i+2*di]\n                    diag_points.append(recovered_xy[j+dj, i+di] + spacing)\n        if diag_points:\n            recovered_xy[j, i] = np.mean(diag_points, axis=0)\n\n    # Pass 5: Global cubic griddata (last resort)\n    if get_missing_indices():\n        recovered_xy = interpolate_mapping(recovered_xy)\n\n    return recovered_xy\n\n\n# -----------------------------------------------------------------------------\n# Gridpoint refinement (when USE_GP_REFINEMENT is set)\n# -----------------------------------------------------------------------------\n\ndef soft_argmax(heatmap):\n    \"\"\"\n    Compute sub-pixel coordinates from heatmap using weighted average.\n\n    Args:\n        heatmap: (B, 1, H, W) predicted heatmaps (sigmoid output)\n\n    Returns:\n        coords: (B, 2) tensor of (x, y) coordinates\n    \"\"\"\n    B, _, H, W = heatmap.shape\n    device = heatmap.device\n\n    x_coords = torch.arange(W, dtype=torch.float32, device=device)\n    y_coords = torch.arange(H, dtype=torch.float32, device=device)\n\n    # Normalize heatmap to sum to 1\n    heatmap_sum = heatmap.sum(dim=(2, 3), keepdim=True) + 1e-8\n    weights = heatmap / heatmap_sum\n\n    # Compute expected x and y\n    x_map = x_coords.view(1, 1, 1, W).expand(B, 1, H, W)\n    y_map = y_coords.view(1, 1, H, 1).expand(B, 1, H, W)\n\n    x = (weights * x_map).sum(dim=(2, 3))\n    y = (weights * y_map).sum(dim=(2, 3))\n\n    return torch.cat([x, y], dim=1)\n\n\n# Note: REFINEMENT_PATCH_SIZE defined in config cell\nREFINEMENT_PATCH_CENTER = (REFINEMENT_PATCH_SIZE - 1) // 2\n\n\ndef refine_gridpoints_network(canonical, gridpoint_xy, refinement_net):\n    \"\"\"\n    Refine gridpoint positions using trained refinement network.\n\n    Args:\n        canonical: RGB image used for rectification\n        gridpoint_xy: (44, 57, 2) array of gridpoint coordinates\n        refinement_net: Loaded refinement network\n\n    Returns:\n        refined_xy: (44, 57, 2) array with refined gridpoint coordinates\n    \"\"\"\n    device = next(refinement_net.parameters()).device\n    H, W = canonical.shape[:2]\n    rows, cols = 44, 57\n    half = REFINEMENT_PATCH_CENTER\n\n    refined_xy = gridpoint_xy.copy()\n\n    # Collect valid gridpoints and their patches\n    patches = []\n    indices = []\n\n    for j in range(rows):\n        for i in range(cols):\n            # Skip edges and corners - only refine interior gridpoints\n            if j == 0 or j == rows - 1 or i == 0 or i == cols - 1:\n                continue\n\n            # Skip upper right 5x5 corner (QR code region)\n            if j < 5 and i >= 52:\n                continue\n\n            x, y = gridpoint_xy[j, i]\n            if x == 0 and y == 0:\n                continue  # Skip missing gridpoints\n\n            # Check bounds for patch extraction\n            x_int, y_int = int(round(x)), int(round(y))\n            if x_int - half < 0 or x_int + half + 1 > W:\n                continue\n            if y_int - half < 0 or y_int + half + 1 > H:\n                continue\n\n            # Extract patch\n            patch = canonical[y_int - half:y_int + half + 1, x_int - half:x_int + half + 1]\n            patches.append(patch)\n            indices.append((j, i, x, y, x_int, y_int))\n\n    if not patches:\n        return refined_xy\n\n    # Stack patches and run through network\n    patches_np = np.stack(patches, axis=0)  # (N, H, W, 3)\n    patches_tensor = torch.from_numpy(patches_np).permute(0, 3, 1, 2).float().to(device) / 255.0\n\n    with torch.no_grad():\n        heatmaps = refinement_net(patches_tensor)  # (N, 1, H, W)\n        coords = soft_argmax(heatmaps)  # (N, 2) - predicted (x, y) in patch coordinates\n\n    coords_np = coords.cpu().numpy()\n\n    # Update gridpoint positions with refined offsets\n    for idx, (j, i, x_orig, y_orig, x_int, y_int) in enumerate(indices):\n        pred_x, pred_y = coords_np[idx]\n        # Offset from patch center\n        offset_x = pred_x - REFINEMENT_PATCH_CENTER\n        offset_y = pred_y - REFINEMENT_PATCH_CENTER\n        # Apply offset to integer center (not original float)\n        refined_xy[j, i, 0] = x_int + offset_x\n        refined_xy[j, i, 1] = y_int + offset_y\n\n    return refined_xy\n\n\ndef compute_interpolation_deviation(gridpoint_xy):\n    \"\"\"\n    Compute how much each gridpoint deviates from a smooth surface fit.\n\n    Fits a smooth surface (RBF) to map grid indices (i,j) -> pixel positions (x,y).\n    The deviation is the distance between actual position and surface prediction.\n    Outliers will have high deviation because they don't fit the smooth surface.\n\n    Args:\n        gridpoint_xy: (44, 57, 2) array of x,y coordinates\n\n    Returns:\n        deviations: (44, 57) array of distances (pixels)\n        interp_xy: (2508, 2) array of interpolated positions (for outlier replacement)\n    \"\"\"\n    mx, my = np.meshgrid(np.arange(0, 57), np.arange(0, 44))\n    coord = np.stack([mx, my], axis=-1).reshape(-1, 2).astype(np.float64)\n    value = gridpoint_xy.copy().reshape(-1, 2).astype(np.float64)\n\n    # Identify valid (non-zero) gridpoints\n    valid = ~np.all(value == [0, 0], axis=1)\n\n    if valid.sum() < 4:\n        return np.zeros((44, 57), dtype=np.float32), value\n\n    # Fit RBF interpolator (smooth surface) to all valid points\n    # Using thin_plate_spline kernel with smoothing to get a smooth fit\n    # that doesn't exactly pass through each point\n    rbf = RBFInterpolator(\n        coord[valid], value[valid],\n        kernel='thin_plate_spline',\n        smoothing=1.0  # Add smoothing so it doesn't interpolate exactly\n    )\n\n    # Predict positions for all grid points\n    interp_xy = rbf(coord)\n\n    # Compute distances between actual and smoothed positions\n    distances = np.sqrt(np.sum((value - interp_xy) ** 2, axis=1))\n    distances[~valid] = 0  # Invalid points get zero deviation\n\n    # Reshape to grid\n    deviations = distances.reshape(44, 57)\n\n    return deviations, interp_xy\n\n\ndef fix_outliers(gridpoint_xy, threshold=0.4, max_iterations=3):\n    \"\"\"\n    Detect and fix gridpoint outliers using iterative TPS fitting.\n\n    Args:\n        gridpoint_xy: (44, 57, 2) array of x,y coordinates\n        threshold: Distance threshold in pixels for outlier detection\n        max_iterations: Maximum iterations for fixing adjacent outliers\n\n    Returns:\n        fixed_xy: (44, 57, 2) array with outliers replaced by interpolated values\n        num_fixed: Total number of outliers fixed\n    \"\"\"\n    fixed_xy = gridpoint_xy.copy()\n    total_fixed = 0\n\n    for iteration in range(max_iterations):\n        deviations, interp_xy = compute_interpolation_deviation(fixed_xy)\n\n        # Find outliers\n        outlier_mask = deviations > threshold\n        num_outliers = outlier_mask.sum()\n\n        if num_outliers == 0:\n            break\n\n        # Replace outliers with interpolated values\n        outlier_indices = np.where(outlier_mask)\n        for j, i in zip(outlier_indices[0], outlier_indices[1]):\n            flat_idx = j * 57 + i\n            fixed_xy[j, i] = interp_xy[flat_idx]\n\n        total_fixed += num_outliers\n\n    return fixed_xy, total_fixed\n\n\ndef _rotate_template_gpu(patch_t, angle_deg, device):\n    \"\"\"Rotate a template patch on GPU using affine transform.\"\"\"\n    if angle_deg == 0:\n        return patch_t\n\n    angle_rad = np.deg2rad(angle_deg)\n    cos_a, sin_a = np.cos(angle_rad), np.sin(angle_rad)\n\n    # Rotation matrix (no translation, rotating around center)\n    theta = torch.tensor([\n        [cos_a, -sin_a, 0],\n        [sin_a, cos_a, 0]\n    ], dtype=torch.float32, device=device).unsqueeze(0)\n\n    grid = F.affine_grid(theta, patch_t.shape, align_corners=False)\n    rotated = F.grid_sample(patch_t, grid, mode='bilinear', padding_mode='zeros', align_corners=False)\n    return rotated\n\n\ndef refine_gridpoints_correlation(image, gridpoint_xy, template_image, template_gridpoint_xy,\n                                  patch_size=15, search_radius=5, rotation_angles=None):\n    \"\"\"\n    Refine gridpoint positions using template correlation (batched GPU version).\n\n    Uses a single clean gridpoint patch from position (10, 10) in the template\n    as the reference for all gridpoints. Processes all valid gridpoints in a\n    single batched GPU operation.\n\n    Args:\n        image: RGB image (H, W, 3) numpy array (canonical)\n        gridpoint_xy: (44, 57, 2) array of x,y coordinates from model\n        template_image: RGB image of type 0001 template\n        template_gridpoint_xy: (44, 57, 2) array of reference gridpoint positions\n        patch_size: Size of correlation patch (default: 15)\n        search_radius: Search radius around each point (default: 5)\n        rotation_angles: List of angles in degrees to try, or None for no rotation\n\n    Returns:\n        refined_xy: (44, 57, 2) array of refined x,y coordinates\n    \"\"\"\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    h, w = image.shape[:2]\n    half_patch = patch_size // 2\n    region_size = 2 * search_radius + patch_size\n    corr_size = 2 * search_radius + 1\n\n    # Use clean gridpoint at (10, 10) as template for all\n    tx, ty = template_gridpoint_xy[10, 10]\n    tcx, tcy = int(round(tx)), int(round(ty))\n    template_patch = template_image[tcy-half_patch:tcy+half_patch+1, tcx-half_patch:tcx+half_patch+1]\n\n    # Pre-compute template tensor\n    patch_t = torch.from_numpy(template_patch.astype(np.float32)).permute(2, 0, 1).unsqueeze(0).to(device)\n    n_elements = patch_size * patch_size * 3\n\n    # Prepare rotated templates if rotation is enabled\n    if rotation_angles and len(rotation_angles) > 0:\n        angles = rotation_angles\n    else:\n        angles = [0]  # No rotation, just use original\n\n    # Pre-compute all rotated templates: list of (patch_centered, patch_norm)\n    rotated_templates = []\n    for angle in angles:\n        rotated = _rotate_template_gpu(patch_t, angle, device)\n        patch_mean = rotated.mean()\n        patch_centered = rotated - patch_mean\n        patch_norm = torch.sqrt((patch_centered ** 2).sum())\n        rotated_templates.append((patch_centered, patch_norm))\n\n    refined_xy = gridpoint_xy.copy()\n\n    # Collect all valid gridpoints and their regions\n    valid_indices = []  # (j, i) pairs\n    valid_origins = []  # (x0, y0) pairs\n    regions_list = []\n\n    for j in range(44):\n        for i in range(57):\n            # Skip edges - only refine interior gridpoints\n            if j == 0 or j == 43 or i == 0 or i == 56:\n                continue\n\n            x, y = gridpoint_xy[j, i]\n\n            # Skip missing gridpoints\n            if x == 0 and y == 0:\n                continue\n\n            # Skip if search region would be out of bounds\n            cx, cy = int(round(x)), int(round(y))\n            x0 = cx - search_radius - half_patch\n            y0 = cy - search_radius - half_patch\n            x1 = x0 + region_size\n            y1 = y0 + region_size\n\n            if x0 < 0 or y0 < 0 or x1 > w or y1 > h:\n                continue\n\n            # Extract search region\n            region = image[y0:y1, x0:x1]\n            if region.shape[0] != region_size or region.shape[1] != region_size:\n                continue\n\n            valid_indices.append((j, i))\n            valid_origins.append((x0, y0))\n            regions_list.append(region)\n\n    if len(regions_list) == 0:\n        return refined_xy\n\n    # Stack all regions into batch tensor: (N, 3, region_size, region_size)\n    regions_np = np.stack(regions_list, axis=0).astype(np.float32)\n    regions_t = torch.from_numpy(regions_np).permute(0, 3, 1, 2).to(device)\n    N = regions_t.shape[0]\n\n    # ones_kernel for local sum: (1, 3, patch_size, patch_size)\n    ones_kernel = torch.ones(1, 3, patch_size, patch_size, device=device)\n\n    # Pre-compute local stats (same for all rotations)\n    local_sum = F.conv2d(regions_t, ones_kernel)  # (N, 1, corr_size, corr_size)\n    regions_sq = regions_t ** 2\n    local_sq_sum = F.conv2d(regions_sq, ones_kernel)\n    local_var = local_sq_sum - (local_sum ** 2) / n_elements\n    local_std = torch.sqrt(torch.clamp(local_var, min=1e-8))\n\n    # Compute correlations for all rotations and find best\n    best_corr = torch.full((N, corr_size, corr_size), -float('inf'), device=device)\n\n    for patch_centered, patch_norm in rotated_templates:\n        cross_corr = F.conv2d(regions_t, patch_centered)  # (N, 1, corr_size, corr_size)\n        corr = cross_corr / (patch_norm * local_std)\n        corr = corr.squeeze(1)  # (N, corr_size, corr_size)\n        best_corr = torch.maximum(best_corr, corr)\n\n    # Find peaks for best correlations\n    corr_flat = best_corr.view(N, -1)  # (N, corr_size*corr_size)\n    peak_indices = corr_flat.argmax(dim=1)  # (N,)\n    iy = (peak_indices // corr_size).cpu().numpy()\n    ix = (peak_indices % corr_size).cpu().numpy()\n\n    # Move correlation maps to CPU for sub-pixel refinement\n    corr_np = best_corr.cpu().numpy()\n\n    # Sub-pixel refinement and update positions\n    for idx, ((j, i), (x0, y0)) in enumerate(zip(valid_indices, valid_origins)):\n        c = corr_np[idx]\n        pix, piy = ix[idx], iy[idx]\n\n        # Parabolic sub-pixel refinement in x\n        if 0 < pix < corr_size - 1:\n            denom = c[piy, pix-1] - 2*c[piy, pix] + c[piy, pix+1]\n            if abs(denom) > 1e-8:\n                dx = (c[piy, pix-1] - c[piy, pix+1]) / (2 * denom)\n            else:\n                dx = 0\n        else:\n            dx = 0\n\n        # Parabolic sub-pixel refinement in y\n        if 0 < piy < corr_size - 1:\n            denom = c[piy-1, pix] - 2*c[piy, pix] + c[piy+1, pix]\n            if abs(denom) > 1e-8:\n                dy = (c[piy-1, pix] - c[piy+1, pix]) / (2 * denom)\n            else:\n                dy = 0\n        else:\n            dy = 0\n\n        # Update position\n        new_x = x0 + pix + dx + half_patch\n        new_y = y0 + piy + dy + half_patch\n        refined_xy[j, i] = [new_x, new_y]\n\n    return refined_xy\n\n\n# -----------------------------------------------------------------------------\n# Main processing function\n# -----------------------------------------------------------------------------\n\ndef process_stage_c(canonical_lr, canonical_hr, stage1_net, refinement_nets=None,\n                    template_image=None, template_gridpoint_xy=None,\n                    H_lr=None, H_hr=None):\n    \"\"\"\n    Stage C: Detect gridpoints and rectify image.\n\n    Args:\n        canonical_lr: RGB image (1152, 1440, 3) from Stage B (always provided)\n        canonical_hr: RGB image (1920, 2400, 3) from Stage B (or None if USE_CANONICAL_HR=False)\n        stage1_net: Loaded Stage 1 model\n        refinement_nets: List of refinement networks for cascade (when USE_GP_REFINEMENT='nn')\n        template_image: Template image for correlation refinement (when USE_GP_REFINEMENT='corr')\n        template_gridpoint_xy: Template gridpoints (44, 57, 2) for correlation refinement\n        H_lr: Homography matrix for canonical_lr (required if USE_CANONICAL_HR=True)\n        H_hr: Homography matrix for canonical_hr (required if USE_CANONICAL_HR=True)\n\n    Returns:\n        rectified: RGB image (1700, 2200, 3) rectified\n        canonical: The canonical image used for rectification (lr or hr)\n        gridpoint_xy_raw: (44, 57, 2) array - raw from stage1 (after transform if USE_CANONICAL_HR)\n        gridpoint_xy_recovered: (44, 57, 2) array - after outlier detection/recovery (or same as raw if disabled)\n        gridpoint_xy_refined: (44, 57, 2) array - after refinement (or same as recovered if disabled)\n    \"\"\"\n    # Run stage1 model on canonical_lr (always 1440x1152)\n    image_tensor = torch.from_numpy(canonical_lr.astype(np.uint8)).permute(2, 0, 1).unsqueeze(0)\n    device = next(stage1_net.parameters()).device\n    batch = {'image': image_tensor.to(device)}\n\n    with torch.amp.autocast('cuda'):\n        with torch.no_grad():\n            output = stage1_net(batch)\n\n    # Extract gridpoints using hengck23's full output_to_predict (includes segment linking)\n    gridpoint_xy, _ = stage1_output_to_predict(canonical_lr, batch, output)\n\n    # Choose image for rectification and transform gridpoints accordingly\n    if USE_CANONICAL_HR and canonical_hr is not None:\n        # Use high-res canonical (2400x1920) - transform gridpoints via homography\n        canonical = canonical_hr\n        # H_composite maps canonical_lr space -> canonical_hr space\n        H_composite = H_hr @ np.linalg.inv(H_lr)\n        gridpoint_xy = transform_gridpoints(gridpoint_xy, H_composite)\n    else:\n        # Use canonical_lr (1440x1152) - hengck23's original approach\n        canonical = canonical_lr\n\n    gridpoint_xy_raw = gridpoint_xy.copy()\n\n    # Advanced keypoint algorithms (outlier detection + recovery)\n    if USE_ADV_KP_ALGOS:\n        gridpoint_xy, _ = detect_outliers(gridpoint_xy, threshold=20.0)\n        gridpoint_xy = recover_missing_gridpoints(gridpoint_xy)\n    gridpoint_xy_recovered = gridpoint_xy.copy()\n\n    # Gridpoint refinement\n    if USE_GP_REFINEMENT == 'nn' and refinement_nets:\n        for refinement_net in refinement_nets:\n            gridpoint_xy = refine_gridpoints_network(canonical, gridpoint_xy, refinement_net)\n            gridpoint_xy, _ = fix_outliers(gridpoint_xy, threshold=0.4, max_iterations=3)\n    elif USE_GP_REFINEMENT == 'corr' and template_image is not None and template_gridpoint_xy is not None:\n        gridpoint_xy = refine_gridpoints_correlation(\n            canonical, gridpoint_xy, template_image, template_gridpoint_xy,\n            rotation_angles=CORR_ROTATION_ANGLES\n        )\n        gridpoint_xy, _ = fix_outliers(gridpoint_xy, threshold=0.4, max_iterations=3)\n    gridpoint_xy_refined = gridpoint_xy.copy()\n\n    # Interpolate any remaining missing gridpoints\n    gridpoint_xy = interpolate_mapping(gridpoint_xy)\n\n    # Rectify image\n    rectified = rectify_image(canonical, gridpoint_xy)\n\n    torch.cuda.empty_cache()\n\n    return rectified, canonical, gridpoint_xy_raw, gridpoint_xy_recovered, gridpoint_xy_refined\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.878517Z","iopub.execute_input":"2026-01-29T09:32:40.878915Z","iopub.status.idle":"2026-01-29T09:32:40.957110Z","shell.execute_reply.started":"2026-01-29T09:32:40.878876Z","shell.execute_reply":"2026-01-29T09:32:40.956226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage D: Signal extraction from rectified image.\n\nInput: Rectified image (2200x1700) from Stage C\nOutput: 4-row signal series (I/aVR/V1/V4, II/aVL/V2/V5, III/aVF/V3/V6, II-rhythm)\n\nUses trained Stage D model at configurable resolution (X_SCALE x Y_SCALE) with multi-patch inference.\n\"\"\"\n\nimport numpy as np\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nfrom scipy.interpolate import pchip_interpolate\n\n# Note: X_SCALE, Y_SCALE, EXTRACTION_METHOD, ENCODER defined in config cell\n\nENCODER_CONFIGS = {\n    'resnet34': {\n        'model_name': 'resnet34.a3_in1k',\n        'dims': [64, 128, 256, 512],\n    },\n    'convnextv2_base': {\n        'model_name': 'convnextv2_base.fcmae_ft_in22k_in1k',\n        'dims': [128, 256, 512, 1024],\n    },\n    'convnextv2_tiny': {\n        'model_name': 'convnextv2_tiny.fcmae_ft_in22k_in1k',\n        'dims': [96, 192, 384, 768],\n    },\n    'efficientnetv2_m': {\n        'model_name': 'tf_efficientnetv2_m.in21k_ft_in1k',\n        'dims': [48, 80, 176, 512],\n    },\n}\n\n# Base parameters at 1x resolution (reference values)\n_T0_1X = 118\n_T1_1X = 2080\n_ZERO_MV_1X = [701.5, 986.0, 1270.0, 1530.0]\n_MV_TO_PIXEL_1X = 79.0\n_T1_DRIFT_2X = 1  # Drift correction at 2x scale\n\n# Derived dimensions from scaling\nFULL_WIDTH = 2200 * X_SCALE\nFULL_HEIGHT = 1700 * Y_SCALE\nCROP_HEIGHT = (FULL_HEIGHT // 32) * 32  # Round down to multiple of 32 for U-Net\nY_OFFSET = (FULL_HEIGHT - CROP_HEIGHT) // 2\n\n# Signal extraction parameters (scaled from 1x reference)\nT0 = _T0_1X * X_SCALE\nT1 = _T1_1X * X_SCALE + _T1_DRIFT_2X * (X_SCALE // 2)  # Drift scales linearly\nZERO_MV = [z * Y_SCALE for z in _ZERO_MV_1X]\nMV_TO_PIXEL = _MV_TO_PIXEL_1X * Y_SCALE\n\n# Original crop coordinates from rectified image (1x resolution, for reference)\nX0_1X, X1_1X = 0, 2176\nY0_1X, Y1_1X = 0, 1696\n\n\n# -----------------------------------------------------------------------------\n# Model architecture (StageDNet - matches training-stage_d/model.py)\n# -----------------------------------------------------------------------------\n\nclass MyCoordDecoderBlock(nn.Module):\n    \"\"\"Decoder block with coordinate injection for spatial awareness.\"\"\"\n\n    def __init__(self, in_channel, skip_channel, out_channel, scale=2):\n        super().__init__()\n        self.scale = scale\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(in_channel + skip_channel + 2, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_channel, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x, skip=None):\n        x = F.interpolate(x, scale_factor=self.scale, mode='nearest')\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)\n\n        b, c, h, w = x.shape\n        coordx, coordy = torch.meshgrid(\n            torch.linspace(-2, 2, w, dtype=x.dtype, device=x.device),\n            torch.linspace(-2, 2, h, dtype=x.dtype, device=x.device),\n            indexing='xy'\n        )\n        coordxy = torch.stack([coordx, coordy], dim=1).reshape(1, 2, h, w).repeat(b, 1, 1, 1)\n        x = torch.cat([x, coordxy], dim=1)\n\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return x\n\n\nclass MyCoordUnetDecoder(nn.Module):\n    \"\"\"U-Net decoder with coordinate-aware blocks.\"\"\"\n\n    def __init__(self, in_channel, skip_channel, out_channel, scale=[2, 2, 2, 2]):\n        super().__init__()\n        i_channel = [in_channel] + out_channel[:-1]\n        block = [\n            MyCoordDecoderBlock(i, s, o, sc)\n            for i, s, o, sc in zip(i_channel, skip_channel, out_channel, scale)\n        ]\n        self.block = nn.ModuleList(block)\n\n    def forward(self, feature, skip):\n        d = feature\n        for i, block in enumerate(self.block):\n            d = block(d, skip[i])\n        return d\n\n\ndef encode_with_resnet(encoder, x):\n    \"\"\"Extract multi-scale features from ResNet encoder (skips maxpool for stride 16 bottleneck).\"\"\"\n    encode = []\n    x = encoder.conv1(x)\n    x = encoder.bn1(x)\n    x = encoder.act1(x)\n    x = encoder.layer1(x)\n    encode.append(x)\n    x = encoder.layer2(x)\n    encode.append(x)\n    x = encoder.layer3(x)\n    encode.append(x)\n    x = encoder.layer4(x)\n    encode.append(x)\n    return encode\n\n\nclass StageDNet(nn.Module):\n    \"\"\"\n    Stage D network for ECG signal extraction.\n\n    Input: RGB image patch (B, 3, H, W), uint8\n    Output: 4-channel probability mask (B, 4, H, W)\n    \"\"\"\n\n    def __init__(self, pretrained=True, encoder_name=None):\n        super().__init__()\n\n        # Use encoder from config if not specified\n        if encoder_name is None:\n            encoder_name = ENCODER\n\n        self.encoder_name = encoder_name\n        cfg = ENCODER_CONFIGS[encoder_name]\n        encoder_dim = cfg['dims']\n        decoder_dim = [256, 128, 64, 32]\n\n        # ResNet: manual extraction skipping maxpool → stride 16 bottleneck\n        # Others: features_only=True → stride 32 bottleneck (needs 5 decoder blocks)\n        if encoder_name == 'resnet34':\n            self.encoder = timm.create_model(\n                model_name=cfg['model_name'],\n                pretrained=pretrained,\n                in_chans=3,\n                num_classes=0,\n                global_pool=''\n            )\n            # 4 decoder blocks: 16x upsample for stride 16 bottleneck\n            self.decoder = MyCoordUnetDecoder(\n                in_channel=encoder_dim[-1],\n                skip_channel=encoder_dim[:-1][::-1] + [0],\n                out_channel=decoder_dim,\n                scale=[2, 2, 2, 2]\n            )\n            final_decoder_channels = 32\n        else:\n            self.encoder = timm.create_model(\n                model_name=cfg['model_name'],\n                pretrained=pretrained,\n                features_only=True,\n                in_chans=3,\n            )\n            # 5 decoder blocks: 32x upsample for stride 32 bottleneck\n            decoder_dim_5 = [256, 128, 64, 32, 16]\n            self.decoder = MyCoordUnetDecoder(\n                in_channel=encoder_dim[-1],\n                skip_channel=encoder_dim[:-1][::-1] + [0, 0],\n                out_channel=decoder_dim_5,\n                scale=[2, 2, 2, 2, 2]\n            )\n            final_decoder_channels = 16\n\n        # +2 for coordy and coordx channels before final prediction\n        self.pixel = nn.Conv2d(final_decoder_channels + 2, 4, 1)\n\n        # Normalization constants\n        self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n    def forward(self, image, x_offset=None):\n        \"\"\"\n        Args:\n            image: (B, 3, H, W) tensor, uint8 or float [0-255]\n            x_offset: (B,) tensor or scalar, horizontal offset of patch within full image\n\n        Returns:\n            pixel: (B, 4, H, W) sigmoid probabilities\n        \"\"\"\n        B, _, H, W = image.shape\n\n        # Normalize\n        x = image.float() / 255\n        x = (x - self.mean) / self.std\n\n        # Coordinate channel for y-position awareness\n        coordy = torch.arange(H, device=x.device).reshape(1, 1, H, 1).repeat(B, 1, 1, W)\n        coordy = coordy.float() / (H - 1) * 2 - 1\n\n        # Coordinate channel for x-position awareness (absolute within full image)\n        local_x = torch.arange(W, device=x.device, dtype=x.dtype)\n        if x_offset is not None:\n            if isinstance(x_offset, (int, float)):\n                x_offset = torch.full((B,), x_offset, device=x.device, dtype=x.dtype)\n            x_offset = x_offset.view(B, 1, 1, 1)\n            coordx = (x_offset + local_x.view(1, 1, 1, W)) / (FULL_WIDTH - 1) * 2 - 1\n            coordx = coordx.repeat(1, 1, H, 1)\n        else:\n            coordx = local_x.view(1, 1, 1, W) / (W - 1) * 2 - 1\n            coordx = coordx.repeat(B, 1, H, 1)\n\n        # Encode-decode\n        if self.encoder_name == 'resnet34':\n            encode = encode_with_resnet(self.encoder, x)\n            skip_suffix = [None]  # 4 decoder blocks\n        else:\n            encode = self.encoder(x)  # features_only=True returns list of feature maps\n            skip_suffix = [None, None]  # 5 decoder blocks\n        last = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + skip_suffix)\n\n        # Add coordy and coordx, then predict\n        last = torch.cat([last, coordy, coordx], dim=1)\n        pixel = self.pixel(last)\n\n        return torch.sigmoid(pixel)\n\n    def forward_logits(self, image, x_offset=None):\n        \"\"\"Forward pass returning logits (for BCEWithLogitsLoss or softmax extraction).\"\"\"\n        B, _, H, W = image.shape\n\n        x = image.float() / 255\n        x = (x - self.mean) / self.std\n\n        coordy = torch.arange(H, device=x.device).reshape(1, 1, H, 1).repeat(B, 1, 1, W)\n        coordy = coordy.float() / (H - 1) * 2 - 1\n\n        local_x = torch.arange(W, device=x.device, dtype=x.dtype)\n        if x_offset is not None:\n            if isinstance(x_offset, (int, float)):\n                x_offset = torch.full((B,), x_offset, device=x.device, dtype=x.dtype)\n            x_offset = x_offset.view(B, 1, 1, 1)\n            coordx = (x_offset + local_x.view(1, 1, 1, W)) / (FULL_WIDTH - 1) * 2 - 1\n            coordx = coordx.repeat(1, 1, H, 1)\n        else:\n            coordx = local_x.view(1, 1, 1, W) / (W - 1) * 2 - 1\n            coordx = coordx.repeat(B, 1, H, 1)\n\n        # Encode-decode\n        if self.encoder_name == 'resnet34':\n            encode = encode_with_resnet(self.encoder, x)\n            skip_suffix = [None]  # 4 decoder blocks\n        else:\n            encode = self.encoder(x)  # features_only=True returns list of feature maps\n            skip_suffix = [None, None]  # 5 decoder blocks\n        last = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + skip_suffix)\n\n        last = torch.cat([last, coordy, coordx], dim=1)\n        logits = self.pixel(last)\n\n        return logits\n\n\n# -----------------------------------------------------------------------------\n# Processing functions\n# -----------------------------------------------------------------------------\n\ndef process_stage_d(rectified, signal_length, stage2_net, threshold=0.001, staged_net_gpu1=None, return_logits_only=False):\n    \"\"\"\n    Stage D: Extract ECG signals from rectified image using overlapping patches at scaled resolution.\n\n    Args:\n        rectified: RGB image (1700, 2200, 3) numpy array from Stage C\n        signal_length: Expected signal length (from metadata)\n        stage2_net: Loaded Stage D model (StageDNet)\n        threshold: Softmax probability threshold (default 0.001 = 30dB SNR)\n        staged_net_gpu1: Optional pre-loaded model copy on GPU 1 for multi-GPU inference\n\n    Returns:\n        series: (4, signal_length) array of extracted signals\n        crop: (CROP_HEIGHT, FULL_WIDTH, 3) cropped input image at scaled resolution\n        stitched: (4, CROP_HEIGHT, FULL_WIDTH) stitched logits at scaled resolution\n        patches: list of NUM_PATCHES input patches, each (CROP_HEIGHT, 1280, 3)\n        patch_logits: (NUM_PATCHES, 4, CROP_HEIGHT, 1280) raw logits for each patch\n    \"\"\"\n    # Patching parameters (at scaled resolution)\n    PATCH_W = 1280          # Model input width\n    TRUSTED_W = 1120        # Trusted center region (80 margin each side)\n    MARGIN = 80             # Artifact-prone margin on each side\n    STRIDE = 1100           # Output stride per patch (PATCH_W - 180 overlap)\n    PAD = 90                # Padding for edge patches (80 margin + 10 discarded)\n\n    device = next(stage2_net.parameters()).device\n\n    # Step 1: Upscale rectified image to target resolution\n    upscaled = cv2.resize(rectified, (FULL_WIDTH, FULL_HEIGHT), interpolation=cv2.INTER_CUBIC)\n\n    # Step 2: Crop to model input dimensions\n    y_offset = (FULL_HEIGHT - CROP_HEIGHT) // 2\n    crop = upscaled[y_offset:y_offset + CROP_HEIGHT, :, :]\n    H, W, C = crop.shape\n\n    # Step 3: Pad left and right for edge patches\n    crop_padded = np.pad(crop, ((0, 0), (PAD, PAD), (0, 0)), mode='constant', constant_values=0)\n\n    # Step 4: Extract patches (number scales with X_SCALE: 2x→4, 4x→8, 8x→16)\n    NUM_PATCHES = 2 * X_SCALE\n    patch_starts = [i * STRIDE for i in range(NUM_PATCHES)]\n\n    patches = []\n    for start in patch_starts:\n        patch = crop_padded[:, start:start + PATCH_W, :]\n        patches.append(patch)\n\n    x_offsets = [start - PAD for start in patch_starts]\n\n    # Step 5: Run model on patches in batches of 2\n    if staged_net_gpu1 is not None:\n        # Multi-GPU: process batches in parallel on different GPUs\n        from concurrent.futures import ThreadPoolExecutor\n\n        def process_batch(net, dev, batch_indices):\n            batch_np = np.stack([patches[i].transpose(2, 0, 1) for i in batch_indices])\n            batch_tensor = torch.from_numpy(np.ascontiguousarray(batch_np)).to(dev)\n            offsets = torch.tensor([x_offsets[i] for i in batch_indices], device=dev, dtype=torch.float32)\n\n            with torch.amp.autocast('cuda'):\n                with torch.no_grad():\n                    logits = net.forward_logits(batch_tensor, x_offset=offsets)\n                    result = logits.float().cpu().numpy()\n\n            del batch_tensor, logits\n            torch.cuda.empty_cache()\n            return result\n\n        # Split patches between GPUs: first half on GPU0, second half on GPU1\n        half = NUM_PATCHES // 2\n        gpu0_indices = list(range(0, half))\n        gpu1_indices = list(range(half, NUM_PATCHES))\n\n        with ThreadPoolExecutor(max_workers=2) as executor:\n            future0 = executor.submit(process_batch, stage2_net, device, gpu0_indices)\n            future1 = executor.submit(process_batch, staged_net_gpu1, 'cuda:1', gpu1_indices)\n            logits_list = [future0.result(), future1.result()]\n    else:\n        # Single GPU: sequential processing in batches of 2\n        logits_list = []\n        with torch.amp.autocast('cuda'):\n            with torch.no_grad():\n                for i in range(0, NUM_PATCHES, 2):\n                    batch_np = np.stack([patches[i].transpose(2, 0, 1), patches[i+1].transpose(2, 0, 1)])\n                    batch_tensor = torch.from_numpy(np.ascontiguousarray(batch_np)).to(device)\n                    offsets = torch.tensor([x_offsets[i], x_offsets[i+1]], device=device, dtype=torch.float32)\n                    logits = stage2_net.forward_logits(batch_tensor, x_offset=offsets)\n                    logits_list.append(logits.float().cpu().numpy())\n                    del batch_tensor, logits\n                    torch.cuda.empty_cache()\n\n    logits_np = np.concatenate(logits_list, axis=0)\n\n    # Step 6: Extract trusted regions and stitch\n    full_W = W\n    stitched = np.zeros((4, H, full_W), dtype=np.float32)\n    counts = np.zeros((1, 1, full_W), dtype=np.float32)\n\n    for i, start in enumerate(patch_starts):\n        patch_out = logits_np[i]\n\n        trust_start = MARGIN\n        trust_end = MARGIN + TRUSTED_W\n\n        if i == 0:\n            trust_start = PAD\n        if i == NUM_PATCHES - 1:\n            trust_end = PATCH_W - PAD\n\n        trusted = patch_out[:, :, trust_start:trust_end]\n\n        out_start = start - PAD + trust_start\n        out_end = out_start + (trust_end - trust_start)\n\n        out_start = max(0, out_start)\n        out_end = min(full_W, out_end)\n\n        trusted_start = 0\n        trusted_end = out_end - out_start\n\n        stitched[:, :, out_start:out_end] += trusted[:, :, trusted_start:trusted_end]\n        counts[:, :, out_start:out_end] += 1\n\n    counts = np.maximum(counts, 1)\n    stitched = stitched / counts\n\n    # Early return for ensemble logit averaging\n    if return_logits_only:\n        torch.cuda.empty_cache()\n        return None, crop, stitched, patches, logits_np\n\n    # Step 7: Extract signal region and convert to series\n    signal_logits = stitched[:, :, T0:T1]\n\n    # Extract raw series (fixed length, no interpolation) for Stage S training\n    raw_series_in_pixel = pixel_to_series_raw(signal_logits, threshold, method=EXTRACTION_METHOD)\n    raw_series = (np.array(ZERO_MV).reshape(4, 1) - raw_series_in_pixel) / MV_TO_PIXEL\n    raw_series = filter_series(raw_series)\n\n    # Extract interpolated series for output\n    series_in_pixel = pixel_to_series(signal_logits, ZERO_MV, signal_length, threshold, method=EXTRACTION_METHOD)\n\n    # Convert pixel coordinates to mV\n    series = (np.array(ZERO_MV).reshape(4, 1) - series_in_pixel) / MV_TO_PIXEL\n\n    # Apply clipping filter\n    series = filter_series(series)\n\n    torch.cuda.empty_cache()\n\n    return series, raw_series, crop, stitched, patches, logits_np\n\n\ndef extract_series_from_stitched(stitched, signal_length, threshold=0.001):\n    \"\"\"Convert stitched logits to series (for ensemble averaging).\n\n    Args:\n        stitched: (4, H, W) stitched logits from process_stage_d\n        signal_length: Target output length\n\n    Returns:\n        series: (4, signal_length) array of extracted signals in mV\n    \"\"\"\n    signal_logits = stitched[:, :, T0:T1]\n    series_in_pixel = pixel_to_series(signal_logits, ZERO_MV, signal_length, threshold, method=EXTRACTION_METHOD)\n    series = (np.array(ZERO_MV).reshape(4, 1) - series_in_pixel) / MV_TO_PIXEL\n    series = filter_series(series)\n    return series\n\n\ndef pixel_to_series_raw(logits, threshold=0.001, method='adjacent'):\n    \"\"\"Convert pixel logits to signal series WITHOUT interpolation.\n\n    Returns series with same width as input logits.\n\n    Args:\n        logits: (C, H, W) array of raw logits\n        threshold: Probability threshold for softmax method\n        method: 'softmax', 'adjacent', 'parabolic', or 'centroid'\n\n    Returns:\n        series: (C, W) array of y-positions (fixed length)\n    \"\"\"\n    C, H, W = logits.shape\n    series = np.zeros((C, W), dtype=np.float32)\n\n    for c in range(C):\n        for x in range(W):\n            col = logits[c, :, x]\n\n            if method == 'softmax':\n                col_exp = np.exp(col - col.max())\n                col_softmax = col_exp / col_exp.sum()\n                y_coords = np.arange(H, dtype=np.float32)\n                mask = col_softmax > threshold\n                if mask.sum() > 0:\n                    masked_softmax = col_softmax * mask\n                    masked_softmax = masked_softmax / masked_softmax.sum()\n                    series[c, x] = np.sum(y_coords * masked_softmax)\n                else:\n                    series[c, x] = col.argmax()\n\n            elif method == 'adjacent':\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n                if peak == 0:\n                    adj = 1\n                elif peak == H - 1:\n                    adj = H - 2\n                else:\n                    adj = peak - 1 if col_probs[peak - 1] > col_probs[peak + 1] else peak + 1\n                p_peak = col_probs[peak]\n                p_adj = col_probs[adj]\n                series[c, x] = (peak * p_peak + adj * p_adj) / (p_peak + p_adj)\n\n            elif method == 'parabolic':\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n                if peak == 0 or peak == H - 1:\n                    series[c, x] = peak\n                else:\n                    y0 = col_probs[peak - 1]\n                    y1 = col_probs[peak]\n                    y2 = col_probs[peak + 1]\n                    denom = y0 - 2 * y1 + y2\n                    if abs(denom) > 1e-8:\n                        delta = 0.5 * (y0 - y2) / denom\n                        delta = np.clip(delta, -1, 1)\n                    else:\n                        delta = 0\n                    series[c, x] = peak + delta\n\n            elif method == 'centroid':\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n                window = int(round(CENTROID_RANGE))\n                start = max(0, peak - window)\n                end = min(H, peak + window + 1)\n                local_probs = col_probs[start:end]\n                local_y = np.arange(start, end, dtype=np.float32)\n                series[c, x] = np.sum(local_y * local_probs) / np.sum(local_probs)\n\n            else:\n                raise ValueError(f\"Unknown method: {method}\")\n\n    return series\n\n\ndef pixel_to_series(logits, zero_mv, signal_length, threshold=0.001, method='adjacent'):\n    \"\"\"Convert pixel logits to signal series.\n\n    Args:\n        logits: (C, H, W) array of raw logits\n        zero_mv: List of zero-mV y-positions (unused, kept for API compatibility)\n        signal_length: Target output length\n        threshold: Probability threshold for softmax method\n        method: 'softmax', 'adjacent', 'parabolic', or 'centroid'\n            - softmax: Weighted average of all pixels above threshold\n            - adjacent: Argmax + weighted average with highest adjacent neighbor\n            - parabolic: Parabolic fit of argmax and 2 neighbors\n            - centroid: Local centroid around argmax (±3 pixels)\n\n    Returns:\n        series: (C, signal_length) array of y-positions\n    \"\"\"\n    C, H, W = logits.shape\n    series = np.zeros((C, signal_length), dtype=np.float32)\n\n    for c in range(C):\n        col_series = np.zeros(W, dtype=np.float32)\n\n        for x in range(W):\n            col = logits[c, :, x]\n\n            if method == 'softmax':\n                # Softmax weighted average with threshold\n                col_exp = np.exp(col - col.max())\n                col_softmax = col_exp / col_exp.sum()\n                y_coords = np.arange(H, dtype=np.float32)\n\n                mask = col_softmax > threshold\n                if mask.sum() > 0:\n                    masked_softmax = col_softmax * mask\n                    masked_softmax = masked_softmax / masked_softmax.sum()\n                    col_series[x] = np.sum(y_coords * masked_softmax)\n                else:\n                    col_series[x] = col.argmax()\n\n            elif method == 'adjacent':\n                # Argmax + adjacent weighted average\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n\n                if peak == 0:\n                    adj = 1\n                elif peak == H - 1:\n                    adj = H - 2\n                else:\n                    adj = peak - 1 if col_probs[peak - 1] > col_probs[peak + 1] else peak + 1\n\n                p_peak = col_probs[peak]\n                p_adj = col_probs[adj]\n                col_series[x] = (peak * p_peak + adj * p_adj) / (p_peak + p_adj)\n\n            elif method == 'parabolic':\n                # Parabolic fit of argmax and 2 neighbors\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n\n                if peak == 0 or peak == H - 1:\n                    col_series[x] = peak\n                else:\n                    # Fit parabola to (peak-1, peak, peak+1)\n                    y0 = col_probs[peak - 1]\n                    y1 = col_probs[peak]\n                    y2 = col_probs[peak + 1]\n                    # Parabolic interpolation: delta = 0.5 * (y0 - y2) / (y0 - 2*y1 + y2)\n                    denom = y0 - 2 * y1 + y2\n                    if abs(denom) > 1e-8:\n                        delta = 0.5 * (y0 - y2) / denom\n                        delta = np.clip(delta, -1, 1)  # Clamp to adjacent pixels\n                    else:\n                        delta = 0\n                    col_series[x] = peak + delta\n\n            elif method == 'centroid':\n                # Local centroid around argmax (±CENTROID_RANGE pixels)\n                # Note: CENTROID_RANGE defined in config cell, should scale with Y_SCALE\n                col_probs = 1.0 / (1.0 + np.exp(-col))\n                peak = col_probs.argmax()\n\n                window = int(round(CENTROID_RANGE))\n                start = max(0, peak - window)\n                end = min(H, peak + window + 1)\n\n                local_probs = col_probs[start:end]\n                local_y = np.arange(start, end, dtype=np.float32)\n                col_series[x] = np.sum(local_y * local_probs) / np.sum(local_probs)\n\n            else:\n                raise ValueError(f\"Unknown method: {method}\")\n\n        x_old = np.linspace(0, 1, W)\n        x_new = np.linspace(0, 1, signal_length)\n        series[c] = pchip_interpolate(x_old, col_series, x_new)\n\n    return series\n\n\ndef filter_series(series):\n    \"\"\"Clip signal amplitudes per lead.\n\n    Per-lead thresholds from clipping analysis at 25 dB SNR target.\n    Row structure:\n      Row 0: I, aVR, V1, V4\n      Row 1: II_short, aVL, V2, V5\n      Row 2: III, aVF, V3, V6\n      Row 3: Lead II (full)\n    \"\"\"\n    # Per-lead clip thresholds (min, max) in mV from 25 dB analysis\n    LEAD_CLIPS = {\n        'I':    (-1.9,  4.4),\n        'II':   (-1.9,  5.9),\n        'III':  (-2.8,  4.0),\n        'aVR':  (-5.5,  1.2),\n        'aVL':  (-1.7,  2.2),\n        'aVF':  (-2.2,  3.7),\n        'V1':   (-4.7,  2.0),\n        'V2':   (-4.5,  2.5),\n        'V3':   (-4.3,  4.3),\n        'V4':   (-4.0,  4.9),\n        'V5':   (-3.7,  4.5),\n        'V6':   (-2.4,  4.5),  # V6 max from V5 (outlier in training data)\n    }\n\n    # Row-to-lead mapping (4 leads per row, in quarter order)\n    ROW_LEADS = [\n        ['I', 'aVR', 'V1', 'V4'],\n        ['II', 'aVL', 'V2', 'V5'],  # II_short uses II thresholds\n        ['III', 'aVF', 'V3', 'V6'],\n    ]\n\n    C, L = series.shape\n\n    # Clip rows 0-2 (short leads) by quarter\n    for j in range(3):\n        for i in range(4):\n            i0 = i * (L // 4)\n            i1 = (i + 1) * (L // 4) if i < 3 else L\n            lead = ROW_LEADS[j][i]\n            series[j, i0:i1] = np.clip(series[j, i0:i1], *LEAD_CLIPS[lead])\n\n    # Clip row 3 (Lead II rhythm strip)\n    series[3] = np.clip(series[3], *LEAD_CLIPS['II'])\n\n    return series\n\n\n# -----------------------------------------------------------------------------\n# Model loading\n# -----------------------------------------------------------------------------\n\ndef load_staged_net(checkpoint_path, device='cuda', encoder_name=None):\n    \"\"\"Load Stage D model from checkpoint.\n\n    Args:\n        checkpoint_path: Path to model checkpoint\n        device: Device to load model on\n        encoder_name: Encoder backbone name (default: use ENCODER from config)\n    \"\"\"\n    model = StageDNet(pretrained=False, encoder_name=encoder_name)\n    model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n    model = model.to(device)\n    model.eval()\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:40.958541Z","iopub.execute_input":"2026-01-29T09:32:40.958969Z","iopub.status.idle":"2026-01-29T09:32:41.026832Z","shell.execute_reply.started":"2026-01-29T09:32:40.958942Z","shell.execute_reply":"2026-01-29T09:32:41.026130Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage E: Post-processing (Savitzky-Golay filter + Einthoven correction)\n\nBased on public notebook techniques:\n1. Savitzky-Golay smoothing filter\n2. Einthoven's law correction (Lead II = Lead I + Lead III)\n\"\"\"\n\nimport numpy as np\nfrom scipy.signal import savgol_filter\n\n# Note: LEAD_NAMES_BY_ROW defined in imports cell\n\n\ndef apply_savgol(series, window_length=7, polyorder=2):\n    \"\"\"\n    Apply Savitzky-Golay filter to each row.\n\n    Args:\n        series: (4, L) array\n        window_length: Filter window (must be odd)\n        polyorder: Polynomial order\n\n    Returns:\n        Smoothed series (4, L)\n    \"\"\"\n    result = series.copy()\n    for i in range(series.shape[0]):\n        # Window must be <= signal length\n        wl = min(window_length, len(series[i]))\n        if wl % 2 == 0:\n            wl -= 1\n        if wl >= polyorder + 2:\n            result[i] = savgol_filter(series[i], window_length=wl, polyorder=polyorder)\n    return result\n\n\ndef apply_einthoven_correction(series, alpha=0.33):\n    \"\"\"\n    Apply Einthoven's law correction to leads I, II_short, III, and Lead II Q1.\n\n    Einthoven's law: Lead II = Lead I + Lead III\n    Any deviation is extraction error. Distribute it across all three leads.\n\n    Also applies correction to first quarter of Lead II (row 3) so that\n    USE_ENSEMBLE averaging doesn't undo the correction.\n\n    Args:\n        series: (4, L) array\n        alpha: Error distribution weight (0.33 = equal distribution)\n\n    Returns:\n        Corrected series (4, L)\n    \"\"\"\n    result = series.copy()\n\n    # Split rows 0-2 into quarters (each row has 4 leads)\n    row0_splits = np.array_split(result[0], 4)  # I, aVR, V1, V4\n    row1_splits = np.array_split(result[1], 4)  # II_short, aVL, V2, V5\n    row2_splits = np.array_split(result[2], 4)  # III, aVF, V3, V6\n\n    # Get the limb leads (first element of each row)\n    lead_I = row0_splits[0]\n    lead_II_short = row1_splits[0]\n    lead_III = row2_splits[0]\n\n    # Check lengths match (they should from array_split)\n    min_len = min(len(lead_I), len(lead_II_short), len(lead_III))\n\n    # Compute Einthoven error: II should equal I + III\n    error = lead_II_short[:min_len] - (lead_I[:min_len] + lead_III[:min_len])\n\n    # Distribute error to rows 0-2\n    lead_I[:min_len] = lead_I[:min_len] + (alpha * error)\n    lead_III[:min_len] = lead_III[:min_len] + (alpha * error)\n    lead_II_short[:min_len] = lead_II_short[:min_len] - (alpha * error)\n\n    # Also apply same correction to Lead II first quarter (row 3)\n    # This ensures USE_ENSEMBLE averaging doesn't undo the correction\n    lead_II_full = result[3]\n    quarter_len = min(min_len, len(lead_II_full) // 4)\n    lead_II_full[:quarter_len] = lead_II_full[:quarter_len] - (alpha * error[:quarter_len])\n\n    # Reconstruct rows\n    result[0] = np.concatenate(row0_splits)\n    result[1] = np.concatenate(row1_splits)\n    result[2] = np.concatenate(row2_splits)\n    result[3] = lead_II_full\n\n    return result\n\n\ndef process_stage_e(series, savgol_window=7, savgol_polyorder=2, einthoven_alpha=0.33):\n    \"\"\"\n    Apply Stage E post-processing.\n\n    Args:\n        series: (4, L) array from Stage D\n        savgol_window: Savitzky-Golay window length\n        savgol_polyorder: Savitzky-Golay polynomial order\n        einthoven_alpha: Einthoven error distribution weight\n\n    Returns:\n        Processed series (4, L)\n    \"\"\"\n    # 1. Savitzky-Golay smoothing\n    #series = apply_savgol(series, savgol_window, savgol_polyorder)\n\n    # 2. Einthoven correction\n    series = apply_einthoven_correction(series, einthoven_alpha)\n\n    return series\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:41.027989Z","iopub.execute_input":"2026-01-29T09:32:41.028310Z","iopub.status.idle":"2026-01-29T09:32:41.053109Z","shell.execute_reply.started":"2026-01-29T09:32:41.028283Z","shell.execute_reply":"2026-01-29T09:32:41.052359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage H: Signal extraction using hengck23's model.\n\nInput: Rectified image (2200x1700) from Stage C\nOutput: 4-row signal series (I/aVR/V1/V4, II/aVL/V2/V5, III/aVF/V3/V6, II-rhythm)\n\nUses hengck23's stage2 model (no resize, native resolution).\n\"\"\"\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nfrom stage2_common import pixel_to_series as pixel_to_series_h, filter_series_by_limits as filter_series_h\n\n# Crop coordinates from rectified image (Stage H specific - avoid collision with Stage D)\nX0_H, X1_H = 0, 2176\nY0_H, Y1_H = 0, 1696\n\n# Signal extraction parameters (hengck23's native resolution)\nT0_H, T1_H = 118, 2080  # X range for signal area\nZERO_MV_H = [703.5, 987.5, 1271.5, 1531.5]  # Y positions of zero lines for 4 rows\nMV_TO_PIXEL_H = 79.0  # pixels per mV\n\n\n# -----------------------------------------------------------------------------\n# Model architecture (hengck23's stage2_model.py)\n# -----------------------------------------------------------------------------\n\nclass MyCoordDecoderBlockH(nn.Module):\n    def __init__(self, in_channel, skip_channel, out_channel, scale=2):\n        super().__init__()\n        self.scale = scale\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(in_channel + skip_channel + 2, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n        self.attention1 = nn.Identity()\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_channel, out_channel, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channel),\n            nn.ReLU(inplace=True),\n        )\n        self.attention2 = nn.Identity()\n\n    def forward(self, x, skip=None):\n        x = F.interpolate(x, scale_factor=self.scale, mode='nearest')\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n\n        b, c, h, w = x.shape\n        coordx, coordy = torch.meshgrid(\n            torch.linspace(-2, 2, w, dtype=x.dtype, device=x.device),\n            torch.linspace(-2, 2, h, dtype=x.dtype, device=x.device),\n            indexing='xy'\n        )\n        coordxy = torch.stack([coordx, coordy], dim=1).reshape(1, 2, h, w).repeat(b, 1, 1, 1)\n        x = torch.cat([x, coordxy], dim=1)\n\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass MyCoordUnetDecoderH(nn.Module):\n    def __init__(self, in_channel, skip_channel, out_channel, scale=[2, 2, 2, 2]):\n        super().__init__()\n        self.center = nn.Identity()\n\n        i_channel = [in_channel] + out_channel[:-1]\n        s_channel = skip_channel\n        o_channel = out_channel\n        block = [\n            MyCoordDecoderBlockH(i, s, o, sc)\n            for i, s, o, sc in zip(i_channel, s_channel, o_channel, scale)\n        ]\n        self.block = nn.ModuleList(block)\n\n    def forward(self, feature, skip):\n        d = self.center(feature)\n        decode = []\n        for i, block in enumerate(self.block):\n            s = skip[i]\n            d = block(d, s)\n            decode.append(d)\n        last = d\n        return last, decode\n\n\ndef encode_with_resnet_h(e, x):\n    encode = []\n    x = e.conv1(x)\n    x = e.bn1(x)\n    x = e.act1(x)\n    x = e.layer1(x)\n    encode.append(x)\n    x = e.layer2(x)\n    encode.append(x)\n    x = e.layer3(x)\n    encode.append(x)\n    x = e.layer4(x)\n    encode.append(x)\n    return encode\n\n\nclass Stage2Net(nn.Module):\n    \"\"\"Stage 2 model from hengck23 (native resolution, larger decoder).\"\"\"\n\n    def __init__(self, pretrained=True):\n        super(Stage2Net, self).__init__()\n        encoder_dim = [64, 128, 256, 512]\n        decoder_dim = [256, 128, 64, 32]\n\n        self.output_type = ['infer', 'loss']\n        self.register_buffer('D', torch.tensor(0))\n        self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).reshape(1, 3, 1, 1))\n        self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).reshape(1, 3, 1, 1))\n\n        self.encoder = timm.create_model(\n            model_name='resnet34.a3_in1k', pretrained=pretrained, in_chans=3, num_classes=0, global_pool=''\n        )\n\n        self.decoder = MyCoordUnetDecoderH(\n            in_channel=encoder_dim[-1],\n            skip_channel=encoder_dim[:-1][::-1] + [0],\n            out_channel=decoder_dim,\n            scale=[2, 2, 2, 2]\n        )\n        # Note: hengck23 adds coordy (+1 channel) before pixel head\n        self.pixel = nn.Conv2d(decoder_dim[-1] + 1, 4, 1)\n\n    def forward(self, batch):\n        device = self.D.device\n        image = batch['image'].to(device)\n        B, _, H, W = image.shape\n        x = image.float() / 255\n        x = (x - self.mean) / self.std\n\n        # Coordinate maps\n        coordy = torch.arange(H, device=device).reshape(1, 1, H, 1).repeat(B, 1, 1, W)\n        coordy = coordy / (H - 1) * 2 - 1\n\n        # Encoder-decoder\n        encode = encode_with_resnet_h(self.encoder, x)\n        last, _ = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + [None])\n\n        # Add coordy before pixel head\n        last = torch.cat([last, coordy], dim=1)\n        pixel = self.pixel(last)\n\n        output = {}\n        if 'infer' in self.output_type:\n            output['pixel'] = torch.sigmoid(pixel)\n\n        return output\n\n\n# -----------------------------------------------------------------------------\n# Processing functions\n# -----------------------------------------------------------------------------\n\ndef process_stage_h(rectified, signal_length, stage2_net):\n    \"\"\"\n    Stage H: Extract ECG signals using hengck23's model.\n\n    Args:\n        rectified: RGB image (1700, 2200, 3) numpy array from Stage C\n        signal_length: Expected signal length (from metadata)\n        stage2_net: Loaded Stage 2 model (Stage2Net)\n\n    Returns:\n        series: (4, signal_length) array of extracted signals\n    \"\"\"\n    device = next(stage2_net.parameters()).device\n\n    # Crop rectified image\n    crop = rectified[Y0_H:Y1_H, X0_H:X1_H]\n\n    # Convert to batch format expected by hengck23's model (uint8, BCHW)\n    batch = {\n        'image': torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0),\n    }\n\n    # Run model\n    with torch.amp.autocast('cuda'):\n        with torch.no_grad():\n            output = stage2_net(batch)\n\n    # Convert pixel predictions to series using hengck23's functions\n    pixel = output['pixel'].float().cpu().numpy()[0]\n    series_in_pixel = pixel_to_series_h(pixel[..., T0_H:T1_H], ZERO_MV_H, signal_length)\n    series = (np.array(ZERO_MV_H).reshape(4, 1) - series_in_pixel) / MV_TO_PIXEL_H\n\n    # Apply clipping filter\n    series = filter_series_h(series)\n\n    torch.cuda.empty_cache()\n\n    return series\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:41.054263Z","iopub.execute_input":"2026-01-29T09:32:41.054914Z","iopub.status.idle":"2026-01-29T09:32:41.086184Z","shell.execute_reply.started":"2026-01-29T09:32:41.054888Z","shell.execute_reply":"2026-01-29T09:32:41.085386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nStage R: Signal Refinement Network\n\nInference code for Stage R. Refines Stage D output signals using a\nmulti-scale dilated 1D convnet that learns to correct systematic errors.\n\nNote: Model architecture must match training-stage_r/model.py\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\n\n# -----------------------------------------------------------------------------\n# Model Architecture (must match training)\n# -----------------------------------------------------------------------------\n\nclass ConvBlock(nn.Module):\n    \"\"\"1D Convolutional block with residual connection.\"\"\"\n\n    def __init__(self, channels, kernel_size=7, dilation=1, dropout=0.2):\n        super().__init__()\n        padding = (kernel_size - 1) * dilation // 2\n\n        self.conv1 = nn.Conv1d(channels, channels, kernel_size,\n                               padding=padding, dilation=dilation)\n        self.bn1 = nn.BatchNorm1d(channels)\n        self.dropout1 = nn.Dropout(dropout)\n\n        self.conv2 = nn.Conv1d(channels, channels, kernel_size,\n                               padding=padding, dilation=dilation)\n        self.bn2 = nn.BatchNorm1d(channels)\n        self.dropout2 = nn.Dropout(dropout)\n\n    def forward(self, x):\n        residual = x\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.dropout1(x)\n        x = self.bn2(self.conv2(x))\n        x = self.dropout2(x)\n        return F.relu(x + residual)\n\n\nclass ChannelMixing(nn.Module):\n    \"\"\"Cross-lead mixing using 1x1 convolutions.\"\"\"\n\n    def __init__(self, channels, dropout=0.2):\n        super().__init__()\n        self.mix = nn.Conv1d(channels, channels, kernel_size=1)\n        self.bn = nn.BatchNorm1d(channels)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        residual = x\n        x = self.bn(self.mix(x))\n        x = self.dropout(x)\n        return F.relu(x + residual)\n\n\nclass StageRNet(nn.Module):\n    \"\"\"Stage R Signal Refinement Network.\"\"\"\n\n    def __init__(self,\n                 in_channels=16,\n                 hidden_channels=96,\n                 kernel_size=7,\n                 dropout=0.2,\n                 use_xcoord=True):\n        super().__init__()\n\n        self.use_xcoord = use_xcoord\n        input_channels = in_channels + 1 if use_xcoord else in_channels\n\n        # Input projection\n        self.input_proj = nn.Conv1d(input_channels, hidden_channels, kernel_size=1)\n        self.input_bn = nn.BatchNorm1d(hidden_channels)\n\n        # Multi-scale dilated convolutional blocks\n        dilations = [1, 2, 4, 8, 16, 32]\n        self.blocks = nn.ModuleList()\n        self.channel_mixers = nn.ModuleList()\n\n        for i, dilation in enumerate(dilations):\n            self.blocks.append(ConvBlock(hidden_channels, kernel_size, dilation, dropout))\n\n            if i % 2 == 1:\n                self.channel_mixers.append(ChannelMixing(hidden_channels, dropout))\n\n        # Output projection\n        self.output_proj = nn.Conv1d(hidden_channels, in_channels, kernel_size=1)\n\n    def forward(self, x, mask=None):\n        batch_size, _, length = x.shape\n        residual = x\n\n        if self.use_xcoord:\n            xcoord = torch.linspace(0, 1, length, device=x.device, dtype=x.dtype)\n            xcoord = xcoord.view(1, 1, length).expand(batch_size, 1, length)\n            x = torch.cat([x, xcoord], dim=1)\n\n        x = F.relu(self.input_bn(self.input_proj(x)))\n\n        mixer_idx = 0\n        for i, block in enumerate(self.blocks):\n            x = block(x)\n            if i % 2 == 1:\n                x = self.channel_mixers[mixer_idx](x)\n                mixer_idx += 1\n\n        x = self.output_proj(x)\n        x = x + residual\n\n        if mask is not None:\n            x = x * mask\n\n        return x\n\n\n# -----------------------------------------------------------------------------\n# Series <-> 16-channel conversion\n# -----------------------------------------------------------------------------\n\n# Lead names by row (matches training-stage_r/dataset.py)\nLEAD_NAMES_BY_ROW = [\n    ['I', 'aVR', 'V1', 'V4'],      # Row 0 quarters\n    ['II_short', 'aVL', 'V2', 'V5'],  # Row 1 quarters\n    ['III', 'aVF', 'V3', 'V6'],    # Row 2 quarters\n]\n\n\ndef compute_quarter_boundaries(total_len):\n    \"\"\"\n    Compute quarter boundaries using floor(N * fraction).\n    Must match training-stage_r/dataset.py exactly.\n    \"\"\"\n    return [\n        0,\n        int(total_len * 0.25),\n        int(total_len * 0.50),\n        int(total_len * 0.75),\n        total_len,\n    ]\n\n\ndef series_to_16_channels(series):\n    \"\"\"\n    Convert 4-row series to 16-channel format for Stage R.\n\n    Args:\n        series: (4, L) array - Stage D output\n            Row 0: I, aVR, V1, V4 concatenated\n            Row 1: II_short, aVL, V2, V5 concatenated\n            Row 2: III, aVF, V3, V6 concatenated\n            Row 3: Full Lead II rhythm strip\n\n    Returns:\n        (16, quarter_length) array:\n            Channels 0-11: 12 short leads (11 standard + II_short), each quarter_length\n            Channels 12-15: 4 quarters of Lead II\n    \"\"\"\n    signal_length = series.shape[1]\n    bounds = compute_quarter_boundaries(signal_length)\n    quarter_lens = [bounds[i + 1] - bounds[i] for i in range(4)]\n    out_len = max(quarter_lens)  # Use max to avoid losing samples from longer quarters\n\n    channels = np.zeros((16, out_len), dtype=np.float32)\n\n    # Channels 0-11: Short leads from rows 0-2 (4 leads per row)\n    for row_idx in range(3):\n        for q_idx in range(4):\n            channel_idx = row_idx * 4 + q_idx\n            q_start = bounds[q_idx]\n            q_end = bounds[q_idx + 1]\n            quarter = series[row_idx, q_start:q_end]\n            q_len = len(quarter)\n            channels[channel_idx, :q_len] = quarter\n            # Pad with end value if shorter than out_len\n            if q_len < out_len:\n                channels[channel_idx, q_len:] = quarter[-1]\n\n    # Channels 12-15: Lead II quarters\n    for q_idx in range(4):\n        q_start = bounds[q_idx]\n        q_end = bounds[q_idx + 1]\n        quarter = series[3, q_start:q_end]\n        q_len = len(quarter)\n        channels[12 + q_idx, :q_len] = quarter\n        if q_len < out_len:\n            channels[12 + q_idx, q_len:] = quarter[-1]\n\n    return channels\n\n\ndef channels_16_to_series(channels, original_length):\n    \"\"\"\n    Convert 16-channel format back to 4-row series.\n\n    Args:\n        channels: (16, quarter_length) array\n        original_length: Original signal length (needed for proper reconstruction)\n\n    Returns:\n        (4, original_length) array - same format as Stage D output\n    \"\"\"\n    bounds = compute_quarter_boundaries(original_length)\n    quarter_lens = [bounds[i + 1] - bounds[i] for i in range(4)]\n\n    series = np.zeros((4, original_length), dtype=np.float32)\n\n    # Rows 0-2: Place each channel at proper quarter position\n    for row_idx in range(3):\n        for q_idx in range(4):\n            channel_idx = row_idx * 4 + q_idx\n            q_start = bounds[q_idx]\n            q_len = quarter_lens[q_idx]\n            series[row_idx, q_start:q_start + q_len] = channels[channel_idx, :q_len]\n\n    # Row 3: Lead II quarters\n    for q_idx in range(4):\n        q_start = bounds[q_idx]\n        q_len = quarter_lens[q_idx]\n        series[3, q_start:q_start + q_len] = channels[12 + q_idx, :q_len]\n\n    return series\n\n\n# -----------------------------------------------------------------------------\n# Model loading and inference\n# -----------------------------------------------------------------------------\n\n_stage_r_model = None\n\n\ndef load_stage_r_model(model_path, device='cuda'):\n    \"\"\"Load Stage R model from checkpoint.\"\"\"\n    global _stage_r_model\n\n    checkpoint = torch.load(model_path, map_location=device)\n\n    # Infer model config from checkpoint (or use defaults)\n    model = StageRNet(\n        in_channels=16,\n        hidden_channels=96,\n        kernel_size=7,\n        dropout=0.0,  # No dropout at inference\n        use_xcoord=True\n    )\n\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.to(device)\n    model.eval()\n\n    _stage_r_model = model\n    return model\n\n\ndef refine_series(series, device='cuda'):\n    \"\"\"\n    Apply Stage R refinement to series.\n\n    Args:\n        series: (4, L) array - Stage D output\n\n    Returns:\n        (4, L) array - Refined series\n    \"\"\"\n    global _stage_r_model\n\n    if _stage_r_model is None:\n        raise RuntimeError(\"Stage R model not loaded. Call load_stage_r_model() first.\")\n\n    original_length = series.shape[1]\n\n    # Convert to 16 channels\n    channels = series_to_16_channels(series)\n\n    # To tensor\n    x = torch.from_numpy(channels).unsqueeze(0).to(device)  # (1, 16, quarter_len)\n\n    # Inference\n    with torch.no_grad():\n        refined = _stage_r_model(x)\n\n    # Back to numpy\n    refined_channels = refined.squeeze(0).cpu().numpy()\n\n    # Convert back to series (needs original length for proper reconstruction)\n    refined_series = channels_16_to_series(refined_channels, original_length)\n\n    return refined_series\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:41.088646Z","iopub.execute_input":"2026-01-29T09:32:41.088955Z","iopub.status.idle":"2026-01-29T09:32:41.113316Z","shell.execute_reply.started":"2026-01-29T09:32:41.088930Z","shell.execute_reply":"2026-01-29T09:32:41.112295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nPipeline: Model loading and main processing function.\n\nLoads all models once at module level and provides process_image() function\nthat chains Stage A -> B -> C -> D/H.\n\"\"\"\n\nimport numpy as np\nimport torch\n# Note: All config variables defined in config cell\n# Note: Stage0Net, Stage1Net, load_net, LEAD_NAMES_BY_ROW from imports cell\n# Note: process_stage_a, process_stage_b, process_stage_c defined in stage cells\n# Note: StageDNet, process_stage_d, Stage2Net, process_stage_h defined in stage cells\n\n\n# -----------------------------------------------------------------------------\n# Refinement Network Definitions\n# -----------------------------------------------------------------------------\n\nclass RefinementNetLegacy(torch.nn.Module):\n    \"\"\"\n    Legacy gridpoint refinement network (4 fixed layers, RF=17).\n    Use for models trained before patch_size support was added.\n    \"\"\"\n\n    def __init__(self, in_channels=3, base_channels=64):\n        super().__init__()\n        import torch.nn as nn\n\n        self.conv1 = nn.Conv2d(in_channels, base_channels, 3, padding=1, dilation=1)\n        self.bn1 = nn.BatchNorm2d(base_channels)\n        self.conv2 = nn.Conv2d(base_channels, base_channels, 3, padding=2, dilation=2)\n        self.bn2 = nn.BatchNorm2d(base_channels)\n        self.conv3 = nn.Conv2d(base_channels, base_channels * 2, 3, padding=4, dilation=4)\n        self.bn3 = nn.BatchNorm2d(base_channels * 2)\n        self.conv4 = nn.Conv2d(base_channels * 2, base_channels * 2, 3, padding=1, dilation=1)\n        self.bn4 = nn.BatchNorm2d(base_channels * 2)\n        self.conv_out = nn.Conv2d(base_channels * 2, 1, 1)\n\n    def forward(self, x):\n        import torch.nn.functional as F\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = F.relu(self.bn2(self.conv2(x)))\n        x = F.relu(self.bn3(self.conv3(x)))\n        x = F.relu(self.bn4(self.conv4(x)))\n        return torch.sigmoid(self.conv_out(x))\n\n\ndef _compute_dilations_for_rf(target_rf):\n    \"\"\"Compute dilation sequence to achieve target receptive field.\"\"\"\n    dilations = []\n    rf = 1\n    d = 1\n    while rf < target_rf:\n        dilations.append(d)\n        rf += (3 - 1) * d\n        d *= 2\n    dilations.append(1)  # Final refinement layer\n    rf += 2\n    return dilations, rf\n\n\nclass RefinementNet(torch.nn.Module):\n    \"\"\"\n    Gridpoint refinement network with dynamic RF.\n    Input: (patch_size x patch_size) RGB patch centered on coarse gridpoint detection\n    Output: (patch_size x patch_size) heatmap with Gaussian peak at true gridpoint position\n    Automatically configures dilations to achieve RF >= patch_size.\n    \"\"\"\n\n    def __init__(self, in_channels=3, base_channels=64, patch_size=31):\n        super().__init__()\n        import torch.nn as nn\n\n        self.patch_size = patch_size\n        dilations, actual_rf = _compute_dilations_for_rf(patch_size)\n        self.dilations = dilations\n        self.actual_rf = actual_rf\n\n        self.convs = nn.ModuleList()\n        self.bns = nn.ModuleList()\n\n        in_ch = in_channels\n        for i, d in enumerate(dilations):\n            out_ch = base_channels if i < 2 else base_channels * 2\n            self.convs.append(nn.Conv2d(in_ch, out_ch, 3, padding=d, dilation=d))\n            self.bns.append(nn.BatchNorm2d(out_ch))\n            in_ch = out_ch\n\n        self.conv_out = nn.Conv2d(in_ch, 1, 1)\n\n    def forward(self, x):\n        import torch.nn.functional as F\n        for conv, bn in zip(self.convs, self.bns):\n            x = F.relu(bn(conv(x)))\n        return torch.sigmoid(self.conv_out(x))\n\n\n# -----------------------------------------------------------------------------\n# Model Loading\n# -----------------------------------------------------------------------------\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nprint('Loading models...')\n\n# Stage 0: Orientation + keypoint detection\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, STAGE0_MODEL_PATH)\nstage0_net.to(DEVICE)\n\n# Stage 1: Gridpoint detection\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(stage1_net, STAGE1_MODEL_PATH)\nstage1_net.to(DEVICE)\nstage1_net.eval()\n\n# Refinement networks (optional, for USE_GP_REFINEMENT='nn')\n# Cascade: models applied in sequence\nrefinement_nets = []\nif USE_GP_REFINEMENT == 'nn':\n    for path in REFINEMENT_MODEL_PATHS:\n        if USE_LEGACY_REFINEMENT:\n            net = RefinementNetLegacy(in_channels=3, base_channels=128)\n        else:\n            net = RefinementNet(in_channels=3, base_channels=128, patch_size=REFINEMENT_PATCH_SIZE)\n        checkpoint = torch.load(path, map_location=DEVICE)\n        net.load_state_dict(checkpoint['model_state_dict'])\n        net.to(DEVICE)\n        net.eval()\n        refinement_nets.append(net)\n    if USE_LEGACY_REFINEMENT:\n        print(f'Refinement networks: loaded {len(refinement_nets)} model(s) [legacy]')\n    else:\n        print(f'Refinement networks: loaded {len(refinement_nets)} model(s), patch_size={REFINEMENT_PATCH_SIZE}')\n\n# Template for correlation refinement (for USE_GP_REFINEMENT='corr')\ntemplate_image = None\ntemplate_gridpoint_xy = None\nif USE_GP_REFINEMENT == 'corr':\n    import cv2\n    template_image = cv2.imread(f'{ECGLIB_DIR}/1006427285-0001.png', cv2.IMREAD_COLOR_RGB)\n    template_gridpoint_xy = np.load(f'{ECGLIB_DIR}/640106434-0001.gridpoint_xy.npy')\n    print('Template for correlation refinement: loaded')\n\n# Stage D/H: Signal extraction\nNUM_GPUS = torch.cuda.device_count() if torch.cuda.is_available() else 0\nstaged_nets = []\nstaged_net_devices = []  # Track which device each model is on\nif USE_STAGEH:\n    staged_net = Stage2Net(pretrained=False)\n    staged_net = load_net(staged_net, STAGE2_MODEL_PATH)\n    staged_net.to(DEVICE)\n    staged_net.eval()\n    staged_nets.append(staged_net)\n    staged_net_devices.append(DEVICE)\n    print('Using hengck23 Stage 2 model (stage_h)')\nelse:\n    for i, path in enumerate(STAGED_MODEL_PATHS):\n        # Multi-GPU ensemble: distribute models across GPUs\n        if NUM_GPUS >= 2 and len(STAGED_MODEL_PATHS) > 1:\n            gpu_id = i % NUM_GPUS\n            device = f'cuda:{gpu_id}'\n        else:\n            device = DEVICE\n        net = StageDNet(pretrained=False)\n        net.load_state_dict(torch.load(path, map_location=device))\n        net.to(device)\n        net.eval()\n        staged_nets.append(net)\n        staged_net_devices.append(device)\n    if NUM_GPUS >= 2 and len(STAGED_MODEL_PATHS) > 1:\n        print(f'Loaded {len(staged_nets)} Stage D model(s) across {min(NUM_GPUS, len(staged_nets))} GPU(s)')\n    else:\n        print(f'Loaded {len(staged_nets)} Stage D model(s)')\n\n# Legacy single-model reference (for compatibility)\nstaged_net = staged_nets[0]\n\n# Multi-GPU: load second copy on GPU 1 if available (only for single model)\n# Note: NUM_GPUS defined above during model loading\nstaged_net_gpu1 = None\nif NUM_GPUS >= 2 and not USE_STAGEH and len(staged_nets) == 1:\n    import copy\n    staged_net_gpu1 = copy.deepcopy(staged_net).to('cuda:1')\n    staged_net_gpu1.eval()\n    print(f'Multi-GPU: Stage D model replicated to cuda:1')\n\n# Stage R: Signal refinement (optional)\n# Note: load_stage_r_model, refine_series defined in stage_r.py cell\nif USE_STAGE_R:\n    load_stage_r_model(STAGE_R_MODEL_PATH, DEVICE)\n    print(f'Stage R model loaded from {STAGE_R_MODEL_PATH}')\n\nprint('Models loaded.')\nprint(f'  USE_ADV_KP_ALGOS: {USE_ADV_KP_ALGOS}')\nprint(f'  USE_GP_REFINEMENT: {USE_GP_REFINEMENT}')\nprint(f'  USE_CANONICAL_HR: {USE_CANONICAL_HR}')\nprint(f'  USE_STAGEH: {USE_STAGEH}')\nprint(f'  USE_STAGE_E: {USE_STAGE_E}')\nprint(f'  USE_STAGE_R: {USE_STAGE_R}')\n\n\n# -----------------------------------------------------------------------------\n# Lead extraction utilities\n# -----------------------------------------------------------------------------\n\ndef extract_leads_from_series(series, ensemble=True):\n    \"\"\"\n    Extract individual lead signals from 4-row series format.\n\n    Args:\n        series: (4, L) array from stage D output\n            Row 0: I, aVR, V1, V4 concatenated\n            Row 1: II_short, aVL, V2, V5 concatenated\n            Row 2: III, aVF, V3, V6 concatenated\n            Row 3: Full Lead II rhythm strip\n        ensemble: If True, average II_short with Lead II first quarter\n\n    Returns:\n        dict: Mapping of lead_name -> signal array for all 12 leads + II_short\n    \"\"\"\n    predicted_leads = {}\n\n    # Extract leads using array_split for rows 0-2\n    for row_idx in range(3):\n        split = np.array_split(series[row_idx], 4)\n        for lead_name, signal in zip(LEAD_NAMES_BY_ROW[row_idx], split):\n            predicted_leads[lead_name] = signal\n\n    # Row 3: Full Lead II rhythm strip\n    lead_ii = series[3].copy()\n\n    if ensemble:\n        # Average II_short with Lead II first quarter\n        ii_short = predicted_leads['II_short']\n        quarter_len = len(ii_short)\n        lead_ii[:quarter_len] = (ii_short + lead_ii[:quarter_len]) / 2.0\n\n    predicted_leads['II'] = lead_ii\n\n    return predicted_leads\n\n\ndef interpolate_signal_to_length(signal, target_length):\n    \"\"\"Interpolate signal to match target length using PCHIP.\"\"\"\n    if len(signal) == target_length:\n        return signal\n\n    from scipy.interpolate import PchipInterpolator\n\n    x_old = np.arange(len(signal))\n    x_new = np.linspace(0, len(signal) - 1, target_length)\n\n    interpolator = PchipInterpolator(x_old, signal)\n    return interpolator(x_new)\n\n\n# -----------------------------------------------------------------------------\n# Main pipeline function\n# -----------------------------------------------------------------------------\n\ndef process_image(image, signal_length, collect_artifacts=False):\n    \"\"\"\n    Run full pipeline on an ECG image.\n\n    Args:\n        image: Input ECG image (RGB numpy array)\n        signal_length: Expected signal length from metadata\n        collect_artifacts: If True, collect intermediate results for debugging\n\n    Returns:\n        series: (4, signal_length) array of extracted signals\n        artifacts: dict of intermediate results (empty if collect_artifacts=False)\n    \"\"\"\n    artifacts = {}\n    timings = {}\n\n    # Stage A: Orientation correction and keypoint detection\n    t0 = time.time()\n    rotated, keypoints = process_stage_a(image, stage0_net)\n    timings['stage_a'] = time.time() - t0\n\n    # Stage B: Homography to canonical size\n    t0 = time.time()\n    canonical_lr, canonical_hr, keypoints_dict, H_lr, H_hr = process_stage_b(rotated, keypoints)\n    timings['stage_b'] = time.time() - t0\n\n    # Stage C: Gridpoint detection and image rectification\n    t0 = time.time()\n    rectified, canonical, gridpoint_xy_raw, gridpoint_xy_recovered, gridpoint_xy_refined = process_stage_c(\n        canonical_lr, canonical_hr, stage1_net, refinement_nets,\n        template_image, template_gridpoint_xy, H_lr, H_hr\n    )\n    timings['stage_c'] = time.time() - t0\n\n    # Stage D/H: Signal extraction from rectified image\n    t0 = time.time()\n    raw_series = None  # Raw series before interpolation (for Stage S training)\n    if USE_STAGEH:\n        series = process_stage_h(rectified, signal_length, staged_nets[0])\n        crop, stitched, patches, patch_logits = None, None, None, None\n    elif len(staged_nets) == 1:\n        # Single model (possibly with multi-GPU)\n        series, raw_series, crop, stitched, patches, patch_logits = process_stage_d(\n            rectified, signal_length, staged_nets[0], staged_net_gpu1=staged_net_gpu1\n        )\n    else:\n        # Ensemble: average logits, then extract series\n        # Multi-GPU: run models in parallel on their respective GPUs\n        def run_model(net):\n            _, crop, stitched, patches, patch_logits = process_stage_d(\n                rectified, signal_length, net, staged_net_gpu1=None, return_logits_only=True\n            )\n            return stitched\n\n        if NUM_GPUS >= 2 and len(staged_nets) > 1:\n            from concurrent.futures import ThreadPoolExecutor\n            with ThreadPoolExecutor(max_workers=NUM_GPUS) as executor:\n                stitched_list = list(executor.map(run_model, staged_nets))\n        else:\n            stitched_list = [run_model(net) for net in staged_nets]\n\n        avg_stitched = np.mean(stitched_list, axis=0)\n        series = extract_series_from_stitched(avg_stitched, signal_length)\n        # For artifacts (ensemble mode: individual model outputs not meaningful)\n        crop, stitched, patches, patch_logits = None, avg_stitched, None, None\n    timings['stage_d'] = time.time() - t0\n\n    # Stage E: Post-processing (Savgol + Einthoven) OR Stage R: Signal refinement\n    # Note: process_stage_e defined in stage_e.py cell\n    # Note: refine_series defined in stage_r.py cell\n    t0 = time.time()\n    if USE_STAGE_E:\n        series = process_stage_e(series)\n    elif USE_STAGE_R:\n        series = refine_series(series, device=DEVICE)\n    timings['stage_e_r'] = time.time() - t0\n\n    artifacts['timings'] = timings\n\n    if collect_artifacts:\n        artifacts['rotated'] = rotated\n        artifacts['keypoints'] = keypoints\n        artifacts['keypoints_dict'] = keypoints_dict\n        artifacts['canonical_lr'] = canonical_lr\n        artifacts['canonical_hr'] = canonical_hr\n        artifacts['canonical'] = canonical  # The one used for rectification\n        artifacts['H_lr'] = H_lr\n        artifacts['H_hr'] = H_hr\n        artifacts['rectified'] = rectified\n        artifacts['gridpoint_xy_raw'] = gridpoint_xy_raw\n        artifacts['gridpoint_xy_recovered'] = gridpoint_xy_recovered\n        artifacts['gridpoint_xy_refined'] = gridpoint_xy_refined\n        artifacts['raw_series'] = raw_series\n        artifacts['crop'] = crop\n        artifacts['stitched'] = stitched\n        artifacts['patches'] = patches\n        artifacts['patch_logits'] = patch_logits\n\n    return series, artifacts\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:41.114542Z","iopub.execute_input":"2026-01-29T09:32:41.115486Z","iopub.status.idle":"2026-01-29T09:32:50.227565Z","shell.execute_reply.started":"2026-01-29T09:32:41.115444Z","shell.execute_reply":"2026-01-29T09:32:50.226683Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nInference: Run on test images and create submission.csv\n\"\"\"\n\n# -----------------------------------------------------------------------------\n# Submission building utilities\n# -----------------------------------------------------------------------------\n\ndef build_submission_rows(predicted_leads, image_id, lead_records):\n    \"\"\"\n    Build submission rows from predicted leads in competition format.\n\n    Args:\n        predicted_leads: dict mapping lead_name -> signal array\n        image_id: Image ID\n        lead_records: DataFrame with lead metadata (id, lead, fs, number_of_rows)\n\n    Returns:\n        list: Submission rows in format [{'id': ..., 'value': ...}, ...]\n    \"\"\"\n    submission_rows = []\n\n    for _, record in lead_records.iterrows():\n        lead_name = record['lead']\n        num_rows = int(record['number_of_rows'])\n\n        if lead_name not in predicted_leads:\n            continue\n\n        pred_signal = predicted_leads[lead_name]\n        pred_signal = interpolate_signal_to_length(pred_signal, num_rows)\n\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_{lead_name}'\n            submission_rows.append({'id': row_id, 'value': float(pred_signal[t])})\n\n    return submission_rows\n\n\n# -----------------------------------------------------------------------------\n# Main inference loop\n# -----------------------------------------------------------------------------\n\n# Load test metadata\ntest_df = pd.read_csv(f'{KAGGLE_DIR}/test.csv')\ntest_df['id'] = test_df['id'].astype(str)\n\n# Get unique image IDs\nimage_ids = test_df['id'].unique().tolist()\n\n# Build submission\nsubmission_rows = []\ntotal_images = len(image_ids)\nstart_time = time.time()\nstage_times = {'stage_a': 0, 'stage_b': 0, 'stage_c': 0, 'stage_d': 0, 'stage_e_r': 0}\n\nfor n, image_id in enumerate(image_ids):\n    lead_records = test_df[test_df['id'] == image_id].copy()\n    filename = f'{KAGGLE_DIR}/test/{image_id}.png'\n    print(f'\\rProcessing {n+1}/{total_images}: {image_id}', end='', flush=True)\n\n    try:\n        image = cv2.imread(filename, cv2.IMREAD_COLOR_RGB)\n        signal_length = lead_records[lead_records['lead'] == 'II'].iloc[0]['number_of_rows']\n\n        series, artifacts = process_image(image, signal_length)\n        for stage, t in artifacts.get('timings', {}).items():\n            stage_times[stage] += t\n        predicted_leads = extract_leads_from_series(series, ensemble=USE_ENSEMBLE)\n        submission_rows.extend(build_submission_rows(predicted_leads, image_id, lead_records))\n\n    except Exception as e:\n        print(f'\\nFailed on {image_id}: {e}')\n        continue\n\nelapsed = time.time() - start_time\nper_image = elapsed / total_images\nprojected_1000 = per_image * 1000\nprint(f'\\n\\nElapsed: {elapsed:.1f}s ({per_image:.2f}s/image)')\nprint(f'Projected for 1000 images: {projected_1000/60:.1f} min')\nprint(f'\\nStage breakdown (avg per image):')\nfor stage, t in stage_times.items():\n    avg = t / total_images\n    pct = 100 * t / elapsed if elapsed > 0 else 0\n    print(f'  {stage}: {avg:.3f}s ({pct:.1f}%)')\n\nprint('\\nBuilding submission DataFrame...')\nsubmission_df = pd.DataFrame(submission_rows)\n\nsubmission_path = 'submission.csv'\nsubmission_df.to_csv(submission_path, index=False)\nprint(f'Submission saved to {submission_path}')\nprint(f'Total rows: {len(submission_df)}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:32:50.228697Z","iopub.execute_input":"2026-01-29T09:32:50.229034Z","iopub.status.idle":"2026-01-29T09:33:21.065001Z","shell.execute_reply.started":"2026-01-29T09:32:50.229002Z","shell.execute_reply":"2026-01-29T09:33:21.064135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nValidate: Run on training images with ground truth scoring.\n\"\"\"\n\n# Previous cells provide:\n#   - validation_subset from ecglib (via imports.py)\n#   - extract_leads_from_series from pipeline.py\n#   - save_gridpoints_csv from stage_c.py\n#   - scoring functions from metric.py\n\n\n# -----------------------------------------------------------------------------\n# Synthetic image support\n# -----------------------------------------------------------------------------\n\n# Variant digit -> sample rate (Hz)\nVARIANT_TO_FS = {\n    '1': 250,\n    '2': 256,\n    '3': 500,\n    '4': 512,\n    '5': 1000,\n    '6': 1025,\n}\n\n\ndef synthetic_subset(start_id, end_id, synthetic_dir):\n    \"\"\"\n    Iterator that yields synthetic image filenames and lead records.\n\n    Args:\n        start_id: Starting synth_id (inclusive), e.g., '000000'\n        end_id: Ending synth_id (inclusive), e.g., '010000'\n        synthetic_dir: Path to synthetic data directory\n\n    Yields:\n        tuple: (filename, lead_records)\n            filename: Path to image file\n            lead_records: DataFrame with lead specifications\n                (columns: id, lead, fs, number_of_rows)\n    \"\"\"\n    # List all synth_id directories\n    for synth_id in sorted(os.listdir(synthetic_dir)):\n        # Filter by ID range (string comparison works for zero-padded IDs)\n        if synth_id < start_id or synth_id > end_id:\n            continue\n\n        # Check if directory contains expected files\n        img_path = f'{synthetic_dir}/{synth_id}/{synth_id}.png'\n        if not os.path.exists(img_path):\n            continue\n\n        # Parse variant from last digit\n        variant = synth_id[-1]\n        if variant not in VARIANT_TO_FS:\n            continue\n\n        fs = VARIANT_TO_FS[variant]\n        signal_length = fs * 10  # PTB-XL is 10 seconds\n        quarter_length = signal_length // 4\n\n        # Build lead_records DataFrame (same format as competition)\n        # Lead II is full length, all others are quarter length\n        lead_rows = []\n        for lead_name in LEADS:\n            if lead_name == 'II':\n                num_rows = signal_length\n            else:\n                num_rows = quarter_length\n            lead_rows.append({\n                'id': synth_id,\n                'lead': lead_name,\n                'fs': fs,\n                'number_of_rows': num_rows,\n            })\n        lead_records = pd.DataFrame(lead_rows)\n\n        yield img_path, lead_records\n\n\ndef load_synthetic_ground_truth(synth_id, synthetic_dir):\n    \"\"\"Load ground truth CSV for a synthetic image.\"\"\"\n    truth_path = f'{synthetic_dir}/{synth_id}/{synth_id}.csv'\n    return pd.read_csv(truth_path)\n\n\n# -----------------------------------------------------------------------------\n# Validation utilities\n# -----------------------------------------------------------------------------\n\ndef load_ground_truth(image_id, data_dir=KAGGLE_DIR):\n    \"\"\"Load ground truth CSV for a training image.\"\"\"\n    truth_path = f'{data_dir}/train/{image_id}/{image_id}.csv'\n    return pd.read_csv(truth_path)\n\n\ndef build_submission_rows(predicted_leads, image_id, lead_records):\n    \"\"\"Build submission rows from predicted leads.\"\"\"\n    submission_rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record['lead']\n        num_rows = int(record['number_of_rows'])\n        if lead_name not in predicted_leads:\n            continue\n        pred_signal = predicted_leads[lead_name]\n        pred_signal = interpolate_signal_to_length(pred_signal, num_rows)\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_{lead_name}'\n            submission_rows.append({'id': row_id, 'value': float(pred_signal[t])})\n    return submission_rows\n\n\ndef build_solution_rows(ground_truth, image_id, lead_records):\n    \"\"\"Build solution rows from ground truth.\"\"\"\n    solution_rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record['lead']\n        num_rows = int(record['number_of_rows'])\n        fs = record['fs']\n        if lead_name not in ground_truth.columns:\n            continue\n        truth_signal = ground_truth[lead_name].dropna().values[:num_rows]\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_{lead_name}'\n            solution_rows.append({'id': row_id, 'fs': fs, 'value': float(truth_signal[t])})\n    return solution_rows\n\n\ndef calculate_lead_snr(truth_signal, pred_signal, fs):\n    \"\"\"Calculate SNR for a single lead.\"\"\"\n    aligned_pred = align_signals(truth_signal, pred_signal, int(fs * MAX_TIME_SHIFT))\n    p_signal, p_noise = compute_power(truth_signal, aligned_pred)\n    snr = compute_snr(p_signal, p_noise)\n    return 10 * np.log10(snr)\n\n\ndef calculate_image_lead_snrs(predicted_leads, ground_truth, fs):\n    \"\"\"Calculate SNR for each lead in an image.\"\"\"\n    image_lead_snrs = {}\n    for lead_name in LEADS:\n        if lead_name not in ground_truth.columns or lead_name not in predicted_leads:\n            continue\n        pred_signal = predicted_leads[lead_name]\n        truth_signal = ground_truth[lead_name].dropna().values\n        pred_signal_interp = interpolate_signal_to_length(pred_signal, len(truth_signal))\n        lead_snr_db = calculate_lead_snr(truth_signal, pred_signal_interp, fs)\n        image_lead_snrs[lead_name] = lead_snr_db\n    return image_lead_snrs\n\n\ndef save_staged_prediction(predicted_leads, image_id, type_id, lead_records, output_dir):\n    \"\"\"\n    Save Stage D predictions for Stage R training.\n\n    Saves two files:\n    - pred-{image_id}-{type_id}.csv: 12 standard leads (I, II, III, aVR, aVL, aVF, V1-V6)\n    - pred-{image_id}-{type_id}_II_short.csv: II_short separately (if present)\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n\n    # Build rows for 12 standard leads\n    pred_rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record['lead']\n        num_rows = int(record['number_of_rows'])\n        if lead_name not in predicted_leads:\n            continue\n        pred_signal = predicted_leads[lead_name]\n        pred_signal = interpolate_signal_to_length(pred_signal, num_rows)\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_{lead_name}'\n            pred_rows.append({'id': row_id, 'value': float(pred_signal[t])})\n\n    # Save main prediction file\n    pred_df = pd.DataFrame(pred_rows)\n    pred_path = f'{output_dir}/pred-{image_id}-{type_id}.csv'\n    pred_df.to_csv(pred_path, index=False)\n\n    # Save II_short separately if present\n    if 'II_short' in predicted_leads:\n        ii_record = lead_records[lead_records['lead'] == 'II'].iloc[0]\n        # II_short is 1/4 the length of full Lead II\n        num_rows = int(ii_record['number_of_rows']) // 4\n        pred_signal = predicted_leads['II_short']\n        pred_signal = interpolate_signal_to_length(pred_signal, num_rows)\n\n        ii_short_rows = []\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_II_short'\n            ii_short_rows.append({'id': row_id, 'value': float(pred_signal[t])})\n\n        ii_short_df = pd.DataFrame(ii_short_rows)\n        ii_short_path = f'{output_dir}/pred-{image_id}-{type_id}_II_short.csv'\n        ii_short_df.to_csv(ii_short_path, index=False)\n\n\ndef save_staged_solution(ground_truth, image_id, type_id, lead_records, output_dir):\n    \"\"\"\n    Save ground truth for Stage R training.\n\n    Saves one file:\n    - sol-{image_id}-{type_id}.csv: 12 standard leads\n\n    Note: No separate II_short file - ground truth II_short is the same as\n    the first quarter of the rhythm strip Lead II.\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n\n    sol_rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record['lead']\n        num_rows = int(record['number_of_rows'])\n        if lead_name not in ground_truth.columns:\n            continue\n        truth_signal = ground_truth[lead_name].dropna().values[:num_rows]\n        for t in range(num_rows):\n            row_id = f'{image_id}_{t}_{lead_name}'\n            sol_rows.append({'id': row_id, 'value': float(truth_signal[t])})\n\n    sol_df = pd.DataFrame(sol_rows)\n    sol_path = f'{output_dir}/sol-{image_id}-{type_id}.csv'\n    sol_df.to_csv(sol_path, index=False)\n\n\n# -----------------------------------------------------------------------------\n# Main validation loop\n# -----------------------------------------------------------------------------\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    print('Skipping validation during competition rerun')\nelse:\n    type_ids = TYPE_IDS_TO_RUN if TYPE_IDS_TO_RUN else TYPE_IDS\n\n    # Choose iterator based on mode\n    if SYNTHETIC:\n        print(f'Synthetic mode: {SYNTHETIC_START_ID} to {SYNTHETIC_END_ID}')\n        print(f'Synthetic dir: {SYNTHETIC_DIR}')\n        image_iterator = synthetic_subset(SYNTHETIC_START_ID, SYNTHETIC_END_ID, SYNTHETIC_DIR)\n        # Synthetic images are all type 0001\n        type_ids = ['0001']\n    else:\n        # Compute effective image IDs based on fold and explicit list\n        effective_image_ids = IMAGE_IDS_TO_RUN\n        if FOLD_TO_RUN is not None:\n            split_df = pd.read_csv(SPLIT_CSV)\n            fold_ids = set(split_df[split_df['fold'] == FOLD_TO_RUN]['image_id'].astype(str))\n            if IMAGE_IDS_TO_RUN is not None:\n                # Union: include images from fold OR explicit list\n                effective_image_ids = list(fold_ids | set(IMAGE_IDS_TO_RUN))\n            else:\n                # Just the fold\n                effective_image_ids = list(fold_ids)\n            print(f'Running fold {FOLD_TO_RUN}: {len(effective_image_ids)} images')\n        image_iterator = validation_subset(effective_image_ids, TYPE_IDS_TO_RUN, KAGGLE_DIR)\n\n    type_snrs = {type_id: [] for type_id in type_ids}\n    lead_snrs = {lead: [] for lead in LEADS}\n    image_stats = []\n    snr_results = []  # For CSV output\n\n    for n, (filename, lead_records) in enumerate(image_iterator):\n        print(f'\\nProcessing {n+1}: {filename}', flush=True)\n\n        image_name = os.path.basename(filename)\n        image_id = lead_records['id'].iloc[0]\n\n        # Parse type_id based on mode\n        if SYNTHETIC:\n            type_id = '0001'  # Synthetic images are always type 0001\n        else:\n            type_id = image_name.split('-')[1].replace('.png', '')\n\n        image = cv2.imread(filename, cv2.IMREAD_COLOR_RGB)\n        if image is None:\n            print(f'Could not read {filename}')\n            continue\n\n        try:\n            signal_length = lead_records[lead_records['lead'] == 'II'].iloc[0]['number_of_rows']\n\n            # Collect artifacts if saving debug outputs or raw predictions\n            collect_artifacts = SAVE_RECTIFIED or SAVE_CANONICAL or SAVE_GRIDPOINTS or SAVE_RAW_PREDS\n            series, artifacts = process_image(image, signal_length, collect_artifacts)\n\n            # Extract leads from series\n            predicted_leads = extract_leads_from_series(series, ensemble=USE_ENSEMBLE)\n\n            # Save rectified image\n            if SAVE_RECTIFIED:\n                os.makedirs(RECTIFIED_DIR, exist_ok=True)\n                rectified_path = f'{RECTIFIED_DIR}/{image_id}-{type_id}.png'\n                cv2.imwrite(rectified_path, cv2.cvtColor(artifacts['rectified'], cv2.COLOR_RGB2BGR))\n\n            # Save canonical image\n            if SAVE_CANONICAL:\n                os.makedirs(CANONICAL_DIR, exist_ok=True)\n                canonical_path = f'{CANONICAL_DIR}/{image_id}-{type_id}.png'\n                cv2.imwrite(canonical_path, cv2.cvtColor(artifacts['canonical'], cv2.COLOR_RGB2BGR))\n\n            # Save gridpoints\n            if SAVE_GRIDPOINTS:\n                os.makedirs(GRIDPOINTS_DIR, exist_ok=True)\n                save_gridpoints_csv(artifacts['gridpoint_xy_raw'], f'{GRIDPOINTS_DIR}/{image_id}-{type_id}-raw.csv')\n                save_gridpoints_csv(artifacts['gridpoint_xy_recovered'], f'{GRIDPOINTS_DIR}/{image_id}-{type_id}-recovered.csv')\n                save_gridpoints_csv(artifacts['gridpoint_xy_refined'], f'{GRIDPOINTS_DIR}/{image_id}-{type_id}-refined.csv')\n\n            # Load ground truth based on mode\n            if SYNTHETIC:\n                ground_truth = load_synthetic_ground_truth(image_id, SYNTHETIC_DIR)\n            else:\n                ground_truth = load_ground_truth(image_id, KAGGLE_DIR)\n\n            # Save Stage R training data\n            if SAVE_STAGED_PREDS:\n                save_staged_prediction(predicted_leads, image_id, type_id, lead_records, STAGED_PREDS_DIR)\n                save_staged_solution(ground_truth, image_id, type_id, lead_records, STAGED_PREDS_DIR)\n\n            # Save Stage S training data (raw 4xN predictions, fixed length before interpolation)\n            if SAVE_RAW_PREDS and artifacts.get('raw_series') is not None:\n                os.makedirs(RAW_PREDS_DIR, exist_ok=True)\n                np.savez_compressed(f'{RAW_PREDS_DIR}/{image_id}-{type_id}.npz', preds=artifacts['raw_series'])\n\n            submission_rows = build_submission_rows(predicted_leads, image_id, lead_records)\n            solution_rows = build_solution_rows(ground_truth, image_id, lead_records)\n\n            submission_df = pd.DataFrame(submission_rows)\n            solution_df = pd.DataFrame(solution_rows)\n\n            fs = lead_records['fs'].iloc[0]\n            overall_snr_db = score(solution_df, submission_df, 'id')\n            print(f'Overall SNR: {overall_snr_db:.2f} dB')\n\n            image_lead_snrs = calculate_image_lead_snrs(predicted_leads, ground_truth, fs)\n            for lead_name, lead_snr_db in image_lead_snrs.items():\n                lead_snrs[lead_name].append(lead_snr_db)\n\n            type_snrs[type_id].append(overall_snr_db)\n\n            stats_row = {\n                'image_id': image_id,\n                'type_id': type_id,\n                'signal_length': signal_length,\n                'overall': overall_snr_db\n            }\n            for lead_name in LEADS:\n                stats_row[lead_name] = image_lead_snrs.get(lead_name, None)\n            image_stats.append(stats_row)\n\n        except Exception as e:\n            print(f'Failed on {filename}: {e}')\n            continue\n\n    # Print summary\n    print('\\n' + '='*80)\n    print('SUMMARY')\n    print('='*80)\n\n    all_image_snrs = []\n    for type_id in type_snrs:\n        all_image_snrs.extend(type_snrs[type_id])\n\n    if all_image_snrs:\n        overall_avg_linear = np.mean([10 ** (db / 10) for db in all_image_snrs])\n        overall_avg_db = 10 * np.log10(overall_avg_linear)\n        print(f'\\nOverall Average SNR: {overall_avg_db:.2f} dB (n={len(all_image_snrs)} images)')\n\n        print('\\nAverage SNR by Image Type:')\n        for type_id in sorted(type_snrs.keys()):\n            if type_snrs[type_id]:\n                avg_linear = np.mean([10 ** (db / 10) for db in type_snrs[type_id]])\n                avg_db = 10 * np.log10(avg_linear)\n                print(f'  Type {type_id}: {avg_db:.2f} dB (n={len(type_snrs[type_id])})')\n\n        print('\\nAverage SNR by Lead:')\n        for lead_name in LEADS:\n            if lead_name in lead_snrs and lead_snrs[lead_name]:\n                avg_linear = np.mean([10 ** (db / 10) for db in lead_snrs[lead_name]])\n                avg_db = 10 * np.log10(avg_linear)\n                print(f'  {lead_name:>3}: {avg_db:.2f} dB (n={len(lead_snrs[lead_name])})')\n\n    # Save results to CSV\n    if image_stats:\n        stats_df = pd.DataFrame(image_stats)\n        if SYNTHETIC:\n            stats_csv_path = f'validation_snr_synth_{SYNTHETIC_START_ID}_{SYNTHETIC_END_ID}.csv'\n        elif FOLD_TO_RUN is not None:\n            stats_csv_path = f'validation_snr_{FOLD_TO_RUN}.csv'\n        else:\n            stats_csv_path = 'validation_snr.csv'\n        stats_df.to_csv(stats_csv_path, index=False)\n        print(f'\\nSaved SNR results to: {stats_csv_path}')\n\n    print('\\nValidation complete!')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T09:33:21.066216Z","iopub.execute_input":"2026-01-29T09:33:21.066553Z"}},"outputs":[],"execution_count":null}]}