{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\nThis is a significantly refactored version of @hengck23's original notebook: [https://www.kaggle.com/code/hengck23/demo-submission](https://www.kaggle.com/code/hengck23/demo-submission).  \nIt includes fixes from @kami1976 and a small TTA enhancement added by me.  However, the biggest change is the inclusion of a validation framework to score images from the training set all the way down to the lead level.  An obvious hint for improvement is included in the first set of images.  Enjoy!\n\n**Update (Version 5):** Included updates from Sesha Raju's notebook: [https://www.kaggle.com/code/seshurajup/henkgck-submission-v4-credits-to-hengck](https://www.kaggle.com/code/seshurajup/henkgck-submission-v4-credits-to-hengck) and additionally averages the short and full lead II signals to predict the first fourth of lead II.\n","metadata":{}},{"cell_type":"markdown","source":"# Imports, Metrics, and Utility Functions","metadata":{}},{"cell_type":"code","source":"try:\n    import cc3d\nexcept:\n    !pip install connected-components-3d --no-index --find-links=file:///kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/setup/\n\nimport cc3d\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport os\nimport sys\n\n# Add hengck23's modified library (by kami1976) to path\nsys.path.append('/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet')\n\nfrom stage0_common import image_to_batch, output_to_predict, normalise_by_homography, load_net\nfrom stage0_model import Net as Stage0Net\nfrom stage1_common import output_to_predict as stage1_output_to_predict, rectify_image\nfrom stage1_model import Net as Stage1Net\nfrom stage2_common import pixel_to_series, filter_series_by_limits\nfrom stage2_model import Net as Stage2Net","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:35.72077Z","iopub.execute_input":"2025-11-18T12:23:35.72096Z","iopub.status.idle":"2025-11-18T12:23:51.821382Z","shell.execute_reply.started":"2025-11-18T12:23:35.720943Z","shell.execute_reply":"2025-11-18T12:23:51.820746Z"},"_kg_hide-input":true,"_kg_hide-output":true},"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    merged_df['image_id'] = merged_df[row_id_column_name].str.split('_').str[0]\n    merged_df['row_id'] = merged_df[row_id_column_name].str.split('_').str[1].astype('int64')\n    merged_df['lead'] = merged_df[row_id_column_name].str.split('_').str[2]\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:51.822825Z","iopub.execute_input":"2025-11-18T12:23:51.823226Z","iopub.status.idle":"2025-11-18T12:23:52.344196Z","shell.execute_reply.started":"2025-11-18T12:23:51.823204Z","shell.execute_reply":"2025-11-18T12:23:52.343667Z"},"_kg_hide-input":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ECG lead structure for 4-row format\n# Rows 0-2 each contain 4 concatenated short leads (~2.5s each)\n# Row 3 contains the full 10s Lead II rhythm strip\nLEAD_NAMES_BY_ROW = [\n    ['I', 'aVR', 'V1', 'V4'],        # Row 0\n    ['II_short', 'aVL', 'V2', 'V5'], # Row 1 (II_short is the short segment)\n    ['III', 'aVF', 'V3', 'V6'],      # Row 2\n]\n\n# Type IDs for validation\nTYPE_IDS = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']\n\n\ndef interpolate_signal_to_length(signal, target_length):\n    \"\"\"\n    Interpolate signal to match target length.\n\n    Length correction by interpolation (from kami1976)\n    https://www.kaggle.com/code/kami1976/physionet-digitization-of-ecg-images-v22\n\n    Args:\n        signal: Input signal array\n        target_length: Desired output length\n\n    Returns:\n        Interpolated signal of length target_length\n    \"\"\"\n    if len(signal) == target_length:\n        return signal\n\n    x_old = np.linspace(0.0, 1.0, len(signal), endpoint=False)\n    x_new = np.linspace(0.0, 1.0, target_length, endpoint=False)\n    return np.interp(x_new, x_old, signal)\n\n\ndef extract_leads_from_series(series):\n    \"\"\"\n    Extract individual lead signals from 4-row series format.\n\n    Args:\n        series: (4, L) array from stage2 output\n            Row 0: I, aVR, V1, V4 concatenated\n            Row 1: II, aVL, V2, V5 concatenated\n            Row 2: III, aVF, V3, V6 concatenated\n            Row 3: Full Lead II rhythm strip\n\n    Returns:\n        dict: Mapping of lead_name -> signal array for all 12 leads\n    \"\"\"\n    L = series.shape[1]\n    predicted_leads = {}\n\n    # Extract leads using array_split for rows 0-2\n    for row_idx in range(3):\n        # array_split handles uneven divisions automatically\n        split = np.array_split(series[row_idx], 4)\n\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    predicted_leads['II'] = series[3]\n\n    return predicted_leads\n\n\n# [DL] This function copied from https://www.kaggle.com/code/hengck23/demo-submission\n# (make_test_fake_df()). Simulates test.csv metadata from train.csv by reading\n# ground truth CSVs.\ndef test_from_train_df(data_dir='/kaggle/input/physionet-ecg-image-digitization'):\n    valid_df = pd.read_csv(f'{data_dir}/train.csv')\n    valid_df['id'] = valid_df['id'].astype(str)\n    fake_test_df=[]\n    for i,d in valid_df.iterrows():\n        #if i==4: break\n        image_id = d['id']\n\n        truth_df = pd.read_csv(f'{data_dir}/train/{image_id}/{image_id}.csv')\n        non_nan_count = truth_df.count()\n        #print(i,image_id,non_nan_count)\n        #print(non_nan_count.index)\n\n        #lead\tfs\tnumber_of_rows\n        this_df = pd.DataFrame({\n            'id':image_id ,\n            'lead':non_nan_count.index,\n            'fs': d['fs'],\n            'number_of_rows':non_nan_count.values\n        })\n        fake_test_df.append(this_df)\n        if i==0: print(this_df)\n    fake_test_df = pd.concat(fake_test_df)\n    return fake_test_df\n\n\ndef validation_subset(image_ids, type_ids,\n                      data_dir='/kaggle/input/physionet-ecg-image-digitization'):\n    \"\"\"\n    Iterator that yields image filenames and lead records for validation.\n\n    Args:\n        image_ids: List of image IDs to process (e.g., ['7663343', '7663344'])\n        type_ids: List of type IDs to process (e.g., ['0001', '0003'])\n        data_dir: Path to PhysioNet data directory\n\n    Yields:\n        tuple: (filename, lead_records)\n            filename: Path to image file\n            lead_records: DataFrame with lead specifications for this image\n                (columns: id, lead, fs, number_of_rows)\n    \"\"\"\n    # Load train.csv to get valid image IDs\n    train_df = pd.read_csv(f'{data_dir}/train.csv')\n    train_df['id'] = train_df['id'].astype(str)\n    valid_image_ids = set(train_df['id'].values)\n\n    # Get lead records using test_from_train_df\n    test_df = test_from_train_df(data_dir)\n\n    # Filter requested image_ids to only those in train.csv\n    filtered_image_ids = [img_id for img_id in image_ids if img_id in valid_image_ids]\n\n    # Iterate through combinations\n    for image_id in filtered_image_ids:\n        # Get lead records for this image_id\n        lead_records = test_df[test_df['id'] == image_id].copy()\n\n        for type_id in type_ids:\n            filename = f'{data_dir}/train/{image_id}/{image_id}-{type_id}.png'\n            yield filename, lead_records\n\n\ndef plot_stage0_result(original, normalised, title='Stage 0 Result', show=True):\n    \"\"\"\n    Plot stage0 input and output side by side.\n\n    Args:\n        original: Original input image (RGB)\n        normalised: Normalized output image (RGB)\n        title: Plot title\n        show: If True, display plot on screen. If False, only prepare for saving.\n    \"\"\"\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n    axes[0].imshow(original)\n    axes[0].set_title('Original Image', fontsize=14)\n    axes[0].axis('off')\n\n    axes[1].imshow(normalised)\n    axes[1].set_title('Normalized Image', fontsize=14)\n    axes[1].axis('off')\n\n    fig.suptitle(title, fontsize=16)\n    plt.tight_layout()\n    if show:\n        plt.show()\n\n\ndef plot_stage1_result(normalised, rectified, title='Stage 1 Result', show=True):\n    \"\"\"\n    Plot stage1 input and output side by side.\n\n    Args:\n        normalised: Normalized input image from stage0 (RGB)\n        rectified: Rectified output image (RGB)\n        title: Plot title\n        show: If True, display plot on screen. If False, only prepare for saving.\n    \"\"\"\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n    axes[0].imshow(normalised)\n    axes[0].set_title('Normalized Image (from Stage 0)', fontsize=14)\n    axes[0].axis('off')\n\n    axes[1].imshow(rectified)\n    axes[1].set_title('Rectified Image', fontsize=14)\n    axes[1].axis('off')\n\n    fig.suptitle(title, fontsize=16)\n    plt.tight_layout()\n    if show:\n        plt.show()\n\n\ndef plot_stage2_result(predicted_leads, title='Stage 2 Result', ground_truth=None, overall_snr_db=None, lead_snrs=None, show=True):\n    \"\"\"\n    Plot stage2 extracted signals as 12 leads (3x4) plus full Lead II.\n\n    Args:\n        predicted_leads: dict mapping lead_name -> signal array for all 12 leads\n        title: Plot title\n        ground_truth: Optional pd.DataFrame with ground truth signals\n        overall_snr_db: Optional overall image SNR in dB to display in main title\n        lead_snrs: Optional dict mapping lead_name -> SNR in dB (pre-calculated)\n        show: If True, display plot on screen. If False, only prepare for saving.\n    \"\"\"\n\n    # 3x4 grid for 12 leads\n    fig1, axes = plt.subplots(3, 4, figsize=(16, 9))\n    title_text = f'{title}'\n    if overall_snr_db is not None:\n        title_text += f' (Overall SNR: {overall_snr_db:.2f} dB)'\n    fig1.suptitle(title_text, fontsize=16)\n\n    # Calculate y limits for each row\n    row_ylims = []\n    for row_idx in range(3):\n        row_signals = [predicted_leads[LEAD_NAMES_BY_ROW[row_idx][col_idx]] for col_idx in range(4)]\n        ymin = min(np.min(s) for s in row_signals)\n        ymax = max(np.max(s) for s in row_signals)\n        row_ylims.append((ymin, ymax))\n\n    for row_idx in range(3):\n        for col_idx in range(4):\n            lead_name = LEAD_NAMES_BY_ROW[row_idx][col_idx]\n            signal = predicted_leads[lead_name]\n            t = np.arange(len(signal))\n\n            ax = axes[row_idx, col_idx]\n\n            # Set title with SNR if available\n            lead_title = lead_name if lead_name != 'II_short' else 'II'\n            if lead_snrs is not None and lead_name in lead_snrs and lead_name != 'II_short':\n                lead_title = f'{lead_title} ({lead_snrs[lead_name]:.1f} dB)'\n\n            # Map II_short to II for ground truth lookup\n            gt_lead_name = 'II' if lead_name == 'II_short' else lead_name\n\n            if ground_truth is not None and gt_lead_name in ground_truth.columns:\n                # For small lead plots, only use as many ground truth samples as we have in prediction\n                # (e.g., for II_short, use first 1/4 of the full Lead II ground truth)\n                truth_signal = ground_truth[gt_lead_name].dropna().values[:len(signal)]\n                t_truth = np.arange(len(truth_signal))\n                ax.plot(t_truth, truth_signal, linewidth=0.8, color='red', alpha=0.7, label='Ground Truth')\n\n                # Align prediction to ground truth\n                aligned_signal = align_signals(truth_signal, signal, max_shift=len(truth_signal) * 0.2)\n                ax.plot(t_truth, aligned_signal, linewidth=0.8, color='blue', label='Prediction (aligned)')\n            else:\n                # Plot prediction without alignment\n                ax.plot(t, signal, linewidth=0.8, color='blue', label='Prediction')\n\n            ax.set_title(lead_title, fontsize=14)\n            ax.set_ylim(row_ylims[row_idx])\n            ax.grid(True, alpha=0.3)\n            if ground_truth is not None:\n                ax.legend(fontsize=10)\n\n    plt.tight_layout()\n    if show:\n        plt.show()\n\n    # Full Lead II\n    fig2, ax = plt.subplots(1, 1, figsize=(16, 3))\n    full_lead_ii = predicted_leads['II']\n    t = np.arange(len(full_lead_ii))\n\n    title_text = f'Full Lead II'\n    if lead_snrs is not None and 'II' in lead_snrs:\n        title_text += f' ({lead_snrs[\"II\"]:.1f} dB)'\n\n    # Plot ground truth if available\n    if ground_truth is not None and 'II' in ground_truth.columns:\n        truth_signal = ground_truth['II'].dropna().values\n        t_truth = np.arange(len(truth_signal))\n        ax.plot(t_truth, truth_signal, linewidth=0.8, color='red', alpha=0.7, label='Ground Truth')\n\n        # Align prediction to ground truth\n        aligned_signal = align_signals(truth_signal, full_lead_ii, max_shift=len(truth_signal) * 0.2)\n        ax.plot(t_truth, aligned_signal, linewidth=0.8, color='blue', label='Prediction (aligned)')\n    else:\n        # Plot prediction without alignment\n        ax.plot(t, full_lead_ii, linewidth=0.8, color='blue', label='Prediction')\n    ax.set_title(title_text, fontsize=14)\n    ax.grid(True, alpha=0.3)\n    if ground_truth is not None:\n        ax.legend(fontsize=10)\n    plt.tight_layout()\n    if show:\n        plt.show()\n\n\ndef plot_all_stages(image, normalised, rectified, predicted_leads, image_id, type_id,\n                    ground_truth=None, overall_snr_db=None, lead_snrs=None,\n                    output_dir='validation_output/plots', show=True):\n    \"\"\"\n    Plot all three stages, save to files, and optionally display.\n\n    Args:\n        image: Original input image (RGB)\n        normalised: Stage 0 normalized image (RGB)\n        rectified: Stage 1 rectified image (RGB)\n        predicted_leads: dict mapping lead_name -> signal array\n        image_id: Image ID string\n        type_id: Type ID string\n        ground_truth: Optional ground truth DataFrame\n        overall_snr_db: Optional overall SNR in dB\n        lead_snrs: Optional dict mapping lead_name -> SNR in dB (pre-calculated)\n        output_dir: Directory to save plots\n        show: If True, display plots on screen\n    \"\"\"\n    import os\n    os.makedirs(output_dir, exist_ok=True)\n    image_name = f'{image_id}-{type_id}'\n\n    # Stage 0: Normalization\n    plot_stage0_result(image, normalised, f'Stage 0: {image_name}', show=show)\n    plt.savefig(f'{output_dir}/stage0-{image_id}-{type_id}.png', dpi=100, bbox_inches='tight')\n    plt.close('all')\n\n    # Stage 1: Rectification\n    plot_stage1_result(normalised, rectified, f'Stage 1: {image_name}', show=show)\n    plt.savefig(f'{output_dir}/stage1-{image_id}-{type_id}.png', dpi=100, bbox_inches='tight')\n    plt.close('all')\n\n    # Stage 2: Signal extraction\n    plot_stage2_result(predicted_leads, f'Stage 2: {image_name}', ground_truth, overall_snr_db, lead_snrs, show=show)\n    plt.savefig(f'{output_dir}/stage2-{image_id}-{type_id}.png', dpi=100, bbox_inches='tight')\n    plt.close('all')\n\n\ndef calculate_lead_snr(truth_signal, pred_signal, fs):\n    \"\"\"\n    Calculate SNR for a single lead.\n\n    Args:\n        truth_signal: np.ndarray - ground truth signal\n        pred_signal: np.ndarray - predicted signal\n        fs: Sampling frequency\n\n    Returns:\n        float: SNR in dB for this lead\n    \"\"\"\n    # Align prediction to ground truth\n    aligned_pred = align_signals(truth_signal, pred_signal, int(fs * MAX_TIME_SHIFT))\n\n    # Compute power\n    p_signal, p_noise = compute_power(truth_signal, aligned_pred)\n\n    # Compute SNR\n    snr = compute_snr(p_signal, p_noise)\n    return 10 * np.log10(snr)\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        # Get prediction\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        # Create rows in competition format\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\ndef build_solution_rows(ground_truth, image_id, lead_records):\n    \"\"\"\n    Build solution rows from ground truth in competition format.\n\n    Args:\n        ground_truth: pd.DataFrame with ground truth signals\n        image_id: Image ID\n        lead_records: DataFrame with lead metadata (id, lead, fs, number_of_rows)\n\n    Returns:\n        list: Solution rows in format [{'id': ..., 'fs': ..., 'value': ...}, ...]\n    \"\"\"\n    solution_rows = []\n\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\n        # Get ground truth\n        if lead_name not in ground_truth.columns:\n            continue\n\n        truth_signal = ground_truth[lead_name].dropna().values[:num_rows]\n\n        # Create rows in competition format\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\n    return solution_rows\n\n\ndef load_ground_truth(image_id, data_dir):\n    \"\"\"\n    Load ground truth CSV for a training image.\n\n    Args:\n        image_id: Image ID\n        data_dir: Path to PhysioNet data directory\n\n    Returns:\n        pd.DataFrame: Ground truth signals with columns for each lead\n    \"\"\"\n    truth_path = f'{data_dir}/train/{image_id}/{image_id}.csv'\n    return pd.read_csv(truth_path)\n\n\ndef save_validation_csvs(solution_df, submission_df, image_id, type_id, output_dir):\n    \"\"\"\n    Save solution and submission CSVs for debugging/analysis.\n\n    Args:\n        solution_df: Solution DataFrame with ground truth\n        submission_df: Submission DataFrame with predictions\n        image_id: Image ID\n        type_id: Type ID\n        output_dir: Directory to save CSV files\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    solution_df.to_csv(f'{output_dir}/sol-{image_id}-{type_id}.csv', index=False)\n    submission_df.to_csv(f'{output_dir}/pred-{image_id}-{type_id}.csv', index=False)\n\n\ndef calculate_image_lead_snrs(predicted_leads, ground_truth, fs):\n    \"\"\"\n    Calculate SNR for each lead in an image.\n\n    Args:\n        predicted_leads: dict mapping lead_name -> signal array\n        ground_truth: pd.DataFrame with ground truth signals\n        fs: Sampling frequency\n\n    Returns:\n        dict: Mapping of lead_name -> SNR in dB for this image\n    \"\"\"\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\n        pred_signal = predicted_leads[lead_name]\n        truth_signal = ground_truth[lead_name].dropna().values[:len(pred_signal)]\n\n        lead_snr_db = calculate_lead_snr(truth_signal, pred_signal, fs)\n        image_lead_snrs[lead_name] = lead_snr_db\n\n    return image_lead_snrs\n\n\ndef build_image_stats_row(image_id, type_id, overall_snr_db, image_lead_snrs):\n    \"\"\"\n    Build statistics row for a single image.\n\n    Args:\n        image_id: Image ID\n        type_id: Type ID\n        overall_snr_db: Overall image SNR in dB\n        image_lead_snrs: dict mapping lead_name -> SNR in dB for this image\n\n    Returns:\n        dict: Statistics row with image_id, type_id, overall, and per-lead SNRs\n    \"\"\"\n    stats_row = {\n        'image_id': image_id,\n        'type_id': type_id,\n        'overall': overall_snr_db\n    }\n    # Add per-lead SNRs\n    for lead_name in LEADS:\n        stats_row[lead_name] = image_lead_snrs.get(lead_name, None)\n    return stats_row\n\n\ndef print_validation_summary(type_snrs, lead_snrs):\n    \"\"\"\n    Print summary statistics for validation run.\n\n    Args:\n        type_snrs: dict mapping type_id -> list of SNR values in dB\n        lead_snrs: dict mapping lead_name -> list of SNR values in dB\n    \"\"\"\n    print('\\n' + '='*80)\n    print('SUMMARY')\n    print('='*80)\n\n    # Overall average across all images\n    all_image_snrs = []\n    for type_id in type_snrs:\n        all_image_snrs.extend(type_snrs[type_id])\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        # Convert dB to linear, average, convert back to dB\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:\n            # Convert dB to linear, average, convert back to dB\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\ndef plot_snr_by_image_type(type_snrs, output_dir):\n    \"\"\"\n    Generate and save box plot for SNR by image type.\n\n    Args:\n        type_snrs: dict mapping type_id -> list of SNR values in dB\n        output_dir: Directory to save plot\n    \"\"\"\n    fig, ax = plt.subplots(1, 1, figsize=(10, 6))\n    type_data = [type_snrs[type_id] for type_id in sorted(type_snrs.keys())]\n    type_labels = sorted(type_snrs.keys())\n    ax.boxplot(type_data, labels=type_labels)\n    ax.set_xlabel('Image Type')\n    ax.set_ylabel('SNR (dB)')\n    ax.set_title('SNR Distribution by Image Type')\n    ax.grid(True, alpha=0.3, axis='y')\n    plt.tight_layout()\n    plt.savefig(f'{output_dir}/summary-snr-by-type.png', dpi=100, bbox_inches='tight')\n    plt.show()\n    plt.close()\n\n\ndef plot_snr_by_lead(lead_snrs, output_dir):\n    \"\"\"\n    Generate and save box plot for SNR by lead.\n\n    Args:\n        lead_snrs: dict mapping lead_name -> list of SNR values in dB\n        output_dir: Directory to save plot\n    \"\"\"\n    fig, ax = plt.subplots(1, 1, figsize=(14, 6))\n    lead_data = [lead_snrs[lead] for lead in LEADS if lead in lead_snrs]\n    lead_labels = [lead for lead in LEADS if lead in lead_snrs]\n    ax.boxplot(lead_data, labels=lead_labels)\n    ax.set_xlabel('Lead', fontsize=12)\n    ax.set_ylabel('SNR (dB)', fontsize=12)\n    ax.set_title('SNR Distribution by Lead', fontsize=16)\n    ax.grid(True, alpha=0.3, axis='y')\n    plt.tight_layout()\n    plt.savefig(f'{output_dir}/summary-snr-by-lead.png', dpi=100, bbox_inches='tight')\n    plt.show()\n    plt.close()\n\n\ndef save_image_stats(image_stats, output_dir):\n    \"\"\"\n    Save per-image statistics to CSV.\n\n    Args:\n        image_stats: List of statistics rows (dicts)\n        output_dir: Directory to save CSV file\n    \"\"\"\n    stats_df = pd.DataFrame(image_stats)\n    stats_df.to_csv(f'{output_dir}/stats.csv', index=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:52.34513Z","iopub.execute_input":"2025-11-18T12:23:52.345457Z","iopub.status.idle":"2025-11-18T12:23:52.384396Z","shell.execute_reply.started":"2025-11-18T12:23:52.34541Z","shell.execute_reply":"2025-11-18T12:23:52.383687Z"},"_kg_hide-input":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stages","metadata":{}},{"cell_type":"code","source":"def process_stage0(image, stage0_net, device='cuda', float_type=None):\n    \"\"\"\n    Run Stage 0: Normalize image by detecting keypoints and applying homography.\n\n    Args:\n        image: np.ndarray of input ECG image (RGB)\n        stage0_net: Loaded Stage 0 model\n        device: 'cuda' or 'cpu'\n        float_type: torch dtype for mixed precision\n\n    Returns:\n        tuple: (normalised, keypoint, homography)\n            normalised: np.ndarray of normalized image (RGB)\n            keypoint: detected keypoints\n            homography: homography matrix\n\n    Raises:\n        Exception: If stage 0 processing fails\n    \"\"\"\n    batch = image_to_batch(image)\n\n    with torch.amp.autocast('cuda', dtype=float_type):\n        with torch.no_grad():\n            output = stage0_net(batch)\n            rotated, keypoint = output_to_predict(image, batch, output)\n            normalised, keypoint, homography = normalise_by_homography(rotated, keypoint)\n\n    torch.cuda.empty_cache()\n\n    return normalised, keypoint, homography","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:52.385784Z","iopub.execute_input":"2025-11-18T12:23:52.385994Z","iopub.status.idle":"2025-11-18T12:23:52.404953Z","shell.execute_reply.started":"2025-11-18T12:23:52.385977Z","shell.execute_reply":"2025-11-18T12:23:52.404157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_stage1(normalised, stage1_net, device='cuda', float_type=None, num_tta=4):\n    \"\"\"\n    Run Stage 1: Grid detection and rectification with TTA.\n\n    Args:\n        normalised: np.ndarray of normalized image from stage0 (RGB)\n        stage1_net: Loaded Stage 1 model\n        device: 'cuda' or 'cpu'\n        float_type: torch dtype for mixed precision\n        num_tta: Number of TTA variations (default 4)\n\n    Returns:\n        tuple: (rectified, gridpoint_xy)\n            rectified: np.ndarray of rectified image (RGB)\n            gridpoint_xy: detected grid points\n\n    Raises:\n        Exception: If stage 1 processing fails\n    \"\"\"\n    # Prepare batch\n    batch = {\n        'image': torch.from_numpy(np.ascontiguousarray(normalised.transpose(2, 0, 1))).unsqueeze(0),\n    }\n\n    # TTA loop - brightness/contrast variations (compatible with training augmentations)\n    marker_logit_sum = 0\n    gridpoint_logit_sum = 0\n    gridhline_logit_sum = 0\n    gridvline_logit_sum = 0\n\n    for trial in range(num_tta):\n        crop = batch['image'].clone().float()\n\n        # Apply brightness/contrast adjustments (preserves spatial structure)\n        if trial == 0:\n            pass  # Original\n        elif trial == 1:\n            crop = (crop * 1.1).clamp(0, 255)  # Brighter\n        elif trial == 2:\n            crop = (crop * 0.9).clamp(0, 255)  # Darker\n        elif trial == 3:\n            # Higher contrast (scale around mean)\n            mean = crop.mean()\n            crop = ((crop - mean) * 1.2 + mean).clamp(0, 255)\n\n        crop = crop.byte()\n\n        with torch.amp.autocast('cuda', dtype=float_type):\n            with torch.no_grad():\n                output = stage1_net({'image': crop})\n\n                # Convert probabilities to logits (inverse softmax/sigmoid)\n                eps = 1e-7\n                marker_logit = torch.log(output['marker'].clamp(min=eps))\n                gridpoint_logit = torch.log(output['gridpoint'].clamp(min=eps, max=1-eps) /\n                                           (1 - output['gridpoint'].clamp(min=eps, max=1-eps)))\n                gridhline_logit = torch.log(output['gridhline'].clamp(min=eps))\n                gridvline_logit = torch.log(output['gridvline'].clamp(min=eps))\n\n                # Accumulate logits\n                marker_logit_sum = marker_logit_sum + marker_logit.float()\n                gridpoint_logit_sum = gridpoint_logit_sum + gridpoint_logit.float()\n                gridhline_logit_sum = gridhline_logit_sum + gridhline_logit.float()\n                gridvline_logit_sum = gridvline_logit_sum + gridvline_logit.float()\n\n    # Average logits and convert back to probabilities\n    output = {\n        'marker': torch.softmax(marker_logit_sum / num_tta, dim=1),\n        'gridpoint': torch.sigmoid(gridpoint_logit_sum / num_tta),\n        'gridhline': torch.softmax(gridhline_logit_sum / num_tta, dim=1),\n        'gridvline': torch.softmax(gridvline_logit_sum / num_tta, dim=1),\n    }\n\n    # Get grid points and rectify image\n    with torch.amp.autocast('cuda', dtype=float_type):\n        with torch.no_grad():\n            gridpoint_xy, more = stage1_output_to_predict(normalised, batch, output)\n            rectified = rectify_image(normalised, gridpoint_xy)\n\n    torch.cuda.empty_cache()\n\n    return rectified, gridpoint_xy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:52.405672Z","iopub.execute_input":"2025-11-18T12:23:52.405896Z","iopub.status.idle":"2025-11-18T12:23:52.424968Z","shell.execute_reply.started":"2025-11-18T12:23:52.405871Z","shell.execute_reply":"2025-11-18T12:23:52.424432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_stage2(rectified, signal_length, stage2_net, device='cuda', float_type=None):\n    \"\"\"\n    Run Stage 2: Extract ECG signal from rectified image.\n\n    Args:\n        rectified: np.ndarray of rectified image from stage1 (RGB)\n        signal_length: Expected signal length (from metadata)\n        stage2_net: Loaded Stage 2 model\n        device: 'cuda' or 'cpu'\n        float_type: torch dtype for mixed precision\n\n    Returns:\n        series: np.ndarray of shape (4, signal_length) containing extracted signals\n            Row 0: I, aVR, V1, V4 (concatenated, split into 4 quarters)\n            Row 1: II, aVL, V2, V5 (concatenated, split into 4 quarters)\n            Row 2: III, aVF, V3, V6 (concatenated, split into 4 quarters)\n            Row 3: II (full length)\n\n    Raises:\n        Exception: If stage 2 processing fails\n    \"\"\"\n    # Rectified coordinate frame parameters\n    x0, x1 = 0, 2176\n    y0, y1 = 0, 1696\n    zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n    mv_to_pixel = 79.0\n    t0, t1 = 118, 2080\n\n    # Crop rectified image\n    crop = rectified[y0:y1, x0:x1]\n    batch = {\n        'image': torch.from_numpy(np.ascontiguousarray(crop.transpose(2, 0, 1))).unsqueeze(0),\n    }\n\n    # Run stage2 model\n    with torch.amp.autocast('cuda', dtype=float_type):\n        with torch.no_grad():\n            output = stage2_net(batch)\n            pixel = output['pixel'].data.cpu().numpy()[0]\n\n    # Convert pixel predictions to series\n    series_in_pixel = pixel_to_series(pixel[..., t0:t1], zero_mv, signal_length)\n    series = (np.array(zero_mv).reshape(4, 1) - series_in_pixel) / mv_to_pixel\n    series = filter_series_by_limits(series)\n\n    torch.cuda.empty_cache()\n\n    return series","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:52.425615Z","iopub.execute_input":"2025-11-18T12:23:52.42582Z","iopub.status.idle":"2025-11-18T12:23:52.443137Z","shell.execute_reply.started":"2025-11-18T12:23:52.425798Z","shell.execute_reply":"2025-11-18T12:23:52.44256Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Full Pipeline","metadata":{}},{"cell_type":"code","source":"# Unified 3-stage pipeline for ECG image processing\n# Consolidates processing logic used by both inference and validation\n\n# Configuration\nKAGGLE_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nWEIGHT_DIR = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight'\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nFLOAT_TYPE = torch.float32\n\n# Load models once at module level\nprint('Loading models...')\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, f'{WEIGHT_DIR}/stage0-last.checkpoint.pth')\nstage0_net.to(DEVICE)\n\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(stage1_net, f'{WEIGHT_DIR}/stage1-last.checkpoint.pth')\nstage1_net.to(DEVICE)\n\nstage2_net = Stage2Net(pretrained=False)\nstage2_net = load_net(stage2_net, f'{WEIGHT_DIR}/stage2-00005810.checkpoint.pth')\nstage2_net.to(DEVICE)\nprint('Models loaded.')\n\n\ndef process_image_pipeline(image, lead_records):\n\t\"\"\"\n\tRun full 3-stage pipeline on an ECG image and extract individual leads.\n\n\tArgs:\n\t\timage: Input ECG image (RGB numpy array)\n\t\tlead_records: DataFrame with metadata for this image\n\n\tReturns:\n\t\ttuple: (predicted_leads, normalised, rectified)\n\t\t\tpredicted_leads: dict mapping lead_name -> signal array for all 12 leads\n\t\t\tnormalised: Stage 0 normalized image (for visualization)\n\t\t\trectified: Stage 1 rectified image (for visualization)\n\n\tRaises:\n\t\tException: If any stage fails\n\t\"\"\"\n\t# Stage 0: Normalize image\n\tnormalised, keypoint, homography = process_stage0(image, stage0_net, DEVICE, FLOAT_TYPE)\n\n\t# Stage 1: Rectify image using detected grid\n\trectified, gridpoint_xy = process_stage1(normalised, stage1_net, DEVICE, FLOAT_TYPE)\n\n\t# Stage 2: Extract ECG signal from rectified image\n\t# Get expected signal length from metadata\n\tsignal_length = lead_records[lead_records['lead'] == 'II'].iloc[0]['number_of_rows']\n\tseries = process_stage2(rectified, signal_length, stage2_net, DEVICE, FLOAT_TYPE)\n\n\t# Extract individual leads from 4-row series format\n\tpredicted_leads = extract_leads_from_series(series)\n\n\treturn predicted_leads, normalised, rectified","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:52.443782Z","iopub.execute_input":"2025-11-18T12:23:52.444012Z","iopub.status.idle":"2025-11-18T12:23:56.024554Z","shell.execute_reply.started":"2025-11-18T12:23:52.443992Z","shell.execute_reply":"2025-11-18T12:23:56.023787Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# Inference on test images and submission creation\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 from test set\nimage_ids = test_df['id'].unique().tolist()\n\n# Build submission\nsubmission_rows = []\ntotal_images = len(image_ids)\n\nfor n, image_id in enumerate(image_ids):\n\t# Get lead records for this image\n\tlead_records = test_df[test_df['id'] == image_id].copy()\n\n\tfilename = f'{KAGGLE_DIR}/test/{image_id}.png'\n\tprint(f'\\r Processing {n+1}/{total_images}: {image_id}', end='', flush=True)\n\n\ttry:\n\t\t# Read image\n\t\timage = cv2.imread(filename, cv2.IMREAD_COLOR_RGB)\n\n\t\t# Run 3-stage pipeline\n\t\tpredicted_leads, normalised, rectified = process_image_pipeline(image, lead_records)\n\n\t\t# Build submission rows for this image\n\t\tsubmission_rows.extend(build_submission_rows(predicted_leads, image_id, lead_records))\n\n\texcept Exception as e:\n\t\tprint(f'\\nFailed on {image_id}: {e}')\n\t\tcontinue\n\nprint('\\n\\nBuilding submission DataFrame...')\nsubmission_df = pd.DataFrame(submission_rows)\n\n# Save submission\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)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:23:56.025252Z","iopub.execute_input":"2025-11-18T12:23:56.025542Z","iopub.status.idle":"2025-11-18T12:24:05.421563Z","shell.execute_reply.started":"2025-11-18T12:23:56.025523Z","shell.execute_reply":"2025-11-18T12:24:05.420726Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Validation Against Training Set","metadata":{}},{"cell_type":"code","source":"# Validation on training images with ground truth scoring\n\n# Skip validation during competition rerun\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n\tprint('Skipping validation during competition rerun')\nelse:\n\t# Define which images to validate\n\ttrain_df = pd.read_csv(f'{KAGGLE_DIR}/train.csv')\n\ttrain_df['id'] = train_df['id'].astype(str)\n\timage_ids = train_df['id'].unique()[:10].tolist()  # 90 images\n\n\t# Initialize tracking for summary statistics\n\ttype_snrs = {type_id: [] for type_id in TYPE_IDS}\n\tlead_snrs = {lead: [] for lead in LEADS}\n\timage_stats = []  # Per-image statistics for stats.csv\n\n\t# Process validation images\n\tfor n, (filename, lead_records) in enumerate(validation_subset(image_ids, TYPE_IDS, KAGGLE_DIR)):\n\t\tprint(f'\\nProcessing {n+1}: {filename}', flush=True)\n\n\t\t# Extract image_id and type_id from filename\n\t\timage_name = os.path.basename(filename)\n\t\ttype_id = image_name.split('-')[1].replace('.png', '')\n\t\timage_id = lead_records['id'].iloc[0]\n\n\t\t# Read image\n\t\timage = cv2.imread(filename, cv2.IMREAD_COLOR_RGB)\n\n\t\t# Run 3-stage pipeline\n\t\ttry:\n\t\t\tpredicted_leads, normalised, rectified = process_image_pipeline(image, lead_records)\n\n\t\t\t# Load ground truth\n\t\t\tground_truth = load_ground_truth(image_id, KAGGLE_DIR)\n\n\t\t\t# Build submission and solution rows for this image\n\t\t\tsubmission_rows = build_submission_rows(predicted_leads, image_id, lead_records)\n\t\t\tsolution_rows = build_solution_rows(ground_truth, image_id, lead_records)\n\n\t\t\t# Convert to dataframes\n\t\t\tsubmission_df = pd.DataFrame(submission_rows)\n\t\t\tsolution_df = pd.DataFrame(solution_rows)\n\n\t\t\t# Calculate overall image SNR using competition scoring method\n\t\t\tfs = lead_records['fs'].iloc[0]\n\t\t\toverall_snr_db = score(solution_df, submission_df, 'id')\n\t\t\tprint(f'Overall SNR: {overall_snr_db:.2f} dB')\n\n\t\t\t# Save solution and submission CSVs for this image\n\t\t\tsave_validation_csvs(solution_df, submission_df, image_id, type_id, 'validation_output')\n\n\t\t\t# Calculate SNR for each lead\n\t\t\timage_lead_snrs = calculate_image_lead_snrs(predicted_leads, ground_truth, fs)\n\t\t\tfor lead_name, lead_snr_db in image_lead_snrs.items():\n\t\t\t\tlead_snrs[lead_name].append(lead_snr_db)\n\n\t\t\t# Plot all stages (show first 10 images only)\n\t\t\tshow_plots = (n < 10)\n\t\t\tplot_all_stages(image, normalised, rectified, predicted_leads, image_id, type_id,\n\t\t\t                ground_truth, overall_snr_db, image_lead_snrs, show=show_plots)\n\n\t\t\t# Track SNR by type (store dB values)\n\t\t\ttype_snrs[type_id].append(overall_snr_db)\n\n\t\t\t# Collect stats for this image\n\t\t\tstats_row = build_image_stats_row(image_id, type_id, overall_snr_db, image_lead_snrs)\n\t\t\timage_stats.append(stats_row)\n\n\t\texcept Exception as e:\n\t\t\tprint(f'Failed on {filename}: {e}')\n\t\t\tcontinue\n\n\t# Print summary statistics\n\tprint_validation_summary(type_snrs, lead_snrs)\n\n\t# Box plots\n\tplot_snr_by_image_type(type_snrs, 'validation_output/plots')\n\tplot_snr_by_lead(lead_snrs, 'validation_output/plots')\n\n\t# Save per-image statistics to CSV\n\tsave_image_stats(image_stats, 'validation_output')\n\tprint(f'Saved per-image statistics to validation_output/stats.csv ({len(image_stats)} images)')\n\n\tprint('Validation complete!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-18T12:24:05.422555Z","iopub.execute_input":"2025-11-18T12:24:05.422843Z"}},"outputs":[],"execution_count":null}]}