{"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":"none","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:09.191966Z","iopub.execute_input":"2025-12-10T18:46:09.192326Z","iopub.status.idle":"2025-12-10T18:46:09.197386Z","shell.execute_reply.started":"2025-12-10T18:46:09.192302Z","shell.execute_reply":"2025-12-10T18:46:09.196198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\nimport seaborn as sns\n\n# Set style\nplt.style.use('seaborn-v0_8-darkgrid')\nsns.set_palette(\"husl\")\n\n# Data paths\nBASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\nTRAIN_PATH = BASE_PATH / 'train'\nTEST_PATH = BASE_PATH / 'test'\n\nprint(\"=\" * 80)\nprint(\"ECG IMAGE DIGITIZATION - DATA EXPLORATION\")\nprint(\"=\" * 80)\n\n# ============================================================================\n# 1. LOAD METADATA\n# ============================================================================\nprint(\"\\n📊 LOADING METADATA...\")\ntrain_df = pd.read_csv(BASE_PATH / 'train.csv')\ntest_df = pd.read_csv(BASE_PATH / 'test.csv')\n\nprint(f\"\\n✓ Training samples: {len(train_df)}\")\nprint(f\"✓ Test samples: {len(test_df)}\")\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING METADATA SAMPLE\")\nprint(\"=\" * 80)\nprint(train_df.head(10))\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TRAINING METADATA STATISTICS\")\nprint(\"=\" * 80)\nprint(train_df.describe())\n\nprint(f\"\\nSampling Frequencies (fs): {sorted(train_df['fs'].unique())}\")\nprint(f\"Signal Lengths (sig_len): {sorted(train_df['sig_len'].unique())}\")\n\n# ============================================================================\n# 2. EXPLORE A SINGLE SAMPLE IN DETAIL\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"DETAILED SAMPLE EXPLORATION\")\nprint(\"=\" * 80)\n\n# Pick first sample\nsample_id = train_df['id'].iloc[0]\nsample_fs = train_df[train_df['id'] == sample_id]['fs'].values[0]\nsample_sig_len = train_df[train_df['id'] == sample_id]['sig_len'].values[0]\n\nprint(f\"\\nSample ID: {sample_id}\")\nprint(f\"Sampling Frequency: {sample_fs} Hz\")\nprint(f\"Signal Length: {sample_sig_len} samples ({sample_sig_len/sample_fs:.1f} seconds)\")\n\n# Load ground truth signal\nsignal_path = TRAIN_PATH / str(sample_id) / f\"{sample_id}.csv\"\nsignal_df = pd.read_csv(signal_path)\n\nprint(f\"\\nGround Truth Shape: {signal_df.shape}\")\nprint(f\"Leads: {list(signal_df.columns)}\")\nprint(f\"\\nFirst 5 samples of each lead:\")\nprint(signal_df.head())\n\nprint(f\"\\nSignal Statistics (in mV):\")\nprint(signal_df.describe())\n\n# ============================================================================\n# 3. VISUALIZE DEGRADATION TYPES\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VISUALIZING DEGRADATION TYPES\")\nprint(\"=\" * 80)\n\ndegradation_types = {\n    '0001': 'Original (Clean)',\n    '0003': 'Color Scan',\n    '0004': 'B&W Scan',\n    '0005': 'Mobile Photo',\n    '0006': 'Screen Photo',\n    '0009': 'Stained',\n    '0010': 'Damaged',\n    '0011': 'Moldy Color',\n    '0012': 'Moldy B&W'\n}\n\nfig, axes = plt.subplots(3, 3, figsize=(20, 18))\naxes = axes.flatten()\n\nfor idx, (seg_id, label) in enumerate(degradation_types.items()):\n    img_path = TRAIN_PATH / str(sample_id) / f\"{sample_id}-{seg_id}.png\"\n    \n    if img_path.exists():\n        img = Image.open(img_path)\n        axes[idx].imshow(img)\n        axes[idx].set_title(f'{label}\\n(Segment {seg_id})', fontsize=12, fontweight='bold')\n        axes[idx].axis('off')\n        \n        # Print image properties\n        print(f\"\\n{label} ({seg_id}):\")\n        print(f\"  - Size: {img.size[0]}x{img.size[1]} pixels\")\n        print(f\"  - Mode: {img.mode}\")\n        print(f\"  - Format: {img.format}\")\n    else:\n        axes[idx].text(0.5, 0.5, 'Not Found', ha='center', va='center')\n        axes[idx].axis('off')\n\nplt.tight_layout()\nplt.savefig('degradation_types_overview.png', dpi=150, bbox_inches='tight')\nprint(\"\\n✓ Saved: degradation_types_overview.png\")\nplt.show()\n\n# ============================================================================\n# 4. COMPARE IMAGE vs GROUND TRUTH\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"IMAGE vs GROUND TRUTH COMPARISON\")\nprint(\"=\" * 80)\n\n# Load the cleanest image\nclean_img_path = TRAIN_PATH / str(sample_id) / f\"{sample_id}-0001.png\"\nclean_img = Image.open(clean_img_path)\n\n# Standard ECG leads\nleads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n\n# Create visualization\nfig = plt.figure(figsize=(20, 16))\n\n# Show the image\nax1 = plt.subplot(2, 1, 1)\nax1.imshow(clean_img)\nax1.set_title(f'ECG Image (Sample {sample_id})', fontsize=14, fontweight='bold')\nax1.axis('off')\n\n# Show ground truth signals\nax2 = plt.subplot(2, 1, 2)\ntime = np.arange(len(signal_df)) / sample_fs\n\nfor i, lead in enumerate(leads):\n    # Offset each lead for visualization\n    offset = i * 2\n    ax2.plot(time, signal_df[lead] + offset, label=lead, linewidth=0.5)\n\nax2.set_xlabel('Time (seconds)', fontsize=12)\nax2.set_ylabel('Lead (with offset)', fontsize=12)\nax2.set_title('Ground Truth Signals', fontsize=14, fontweight='bold')\nax2.legend(loc='right', bbox_to_anchor=(1.05, 0.5))\nax2.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig('image_vs_groundtruth.png', dpi=150, bbox_inches='tight')\nprint(\"✓ Saved: image_vs_groundtruth.png\")\nplt.show()\n\n# ============================================================================\n# 5. INDIVIDUAL LEAD ANALYSIS\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"INDIVIDUAL LEAD ANALYSIS\")\nprint(\"=\" * 80)\n\nfig, axes = plt.subplots(4, 3, figsize=(20, 16))\naxes = axes.flatten()\n\ntime = np.arange(len(signal_df)) / sample_fs\n\nfor i, lead in enumerate(leads):\n    axes[i].plot(time, signal_df[lead], linewidth=0.8, color=f'C{i}')\n    axes[i].set_title(f'Lead {lead}', fontweight='bold')\n    axes[i].set_xlabel('Time (s)')\n    axes[i].set_ylabel('Amplitude (mV)')\n    axes[i].grid(True, alpha=0.3)\n    \n    # Add statistics\n    mean_val = signal_df[lead].mean()\n    std_val = signal_df[lead].std()\n    min_val = signal_df[lead].min()\n    max_val = signal_df[lead].max()\n    \n    stats_text = f'μ={mean_val:.2f} σ={std_val:.2f}\\nmin={min_val:.2f} max={max_val:.2f}'\n    axes[i].text(0.02, 0.98, stats_text, transform=axes[i].transAxes,\n                fontsize=8, verticalalignment='top',\n                bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n\nplt.tight_layout()\nplt.savefig('individual_leads.png', dpi=150, bbox_inches='tight')\nprint(\"✓ Saved: individual_leads.png\")\nplt.show()\n\n# ============================================================================\n# 6. LEAD II SPECIAL ANALYSIS (10 seconds)\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LEAD II SPECIAL ANALYSIS (10 seconds)\")\nprint(\"=\" * 80)\n\nlead_ii_duration = len(signal_df) / sample_fs\nother_leads_duration = 2.5  # seconds\n\nprint(f\"Lead II duration: {lead_ii_duration:.1f} seconds\")\nprint(f\"Other leads duration: {other_leads_duration:.1f} seconds\")\nprint(f\"Lead II is {lead_ii_duration/other_leads_duration:.1f}x longer\")\n\n# Calculate expected points for test submission\nprint(f\"\\nExpected points per lead (for submission):\")\nprint(f\"  Lead II: floor({sample_fs} * 10) = {int(sample_fs * 10)} points\")\nprint(f\"  Others:  floor({sample_fs} * 2.5) = {int(sample_fs * 2.5)} points\")\n\n# ============================================================================\n# 7. DATASET STATISTICS\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OVERALL DATASET STATISTICS\")\nprint(\"=\" * 80)\n\nprint(f\"\\nTotal training samples: {len(train_df)}\")\nprint(f\"Total training images: {len(train_df) * 9} (9 per sample)\")\nprint(f\"\\nSampling frequency distribution:\")\nprint(train_df['fs'].value_counts().sort_index())\n\n# Check file existence\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FILE STRUCTURE VERIFICATION\")\nprint(\"=\" * 80)\n\nsample_count = min(3, len(train_df))\nfor i in range(sample_count):\n    sid = train_df['id'].iloc[i]\n    print(f\"\\nSample {i+1}: ID {sid}\")\n    \n    csv_path = TRAIN_PATH / str(sid) / f\"{sid}.csv\"\n    print(f\"  Ground truth CSV: {'✓ Found' if csv_path.exists() else '✗ Missing'}\")\n    \n    for seg in ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']:\n        img_path = TRAIN_PATH / str(sid) / f\"{sid}-{seg}.png\"\n        status = '✓' if img_path.exists() else '✗'\n        print(f\"  Image {seg}: {status}\")\n\n# ============================================================================\n# 8. TEST SET PREVIEW\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"TEST SET PREVIEW\")\nprint(\"=\" * 80)\n\nprint(test_df.head(15))\n\nprint(f\"\\nUnique test IDs: {test_df['id'].nunique()}\")\nprint(f\"Total test rows: {len(test_df)}\")\nprint(f\"\\nLeads per test image: {test_df.groupby('id')['lead'].count().value_counts()}\")\n\n# Show a test image\nif len(test_df) > 0:\n    test_id = test_df['id'].iloc[0]\n    test_img_path = TEST_PATH / f\"{test_id}.png\"\n    \n    if test_img_path.exists():\n        fig, ax = plt.subplots(1, 1, figsize=(16, 10))\n        test_img = Image.open(test_img_path)\n        ax.imshow(test_img)\n        ax.set_title(f'Test Image: {test_id}', fontsize=14, fontweight='bold')\n        ax.axis('off')\n        plt.tight_layout()\n        plt.savefig('test_image_sample.png', dpi=150, bbox_inches='tight')\n        print(\"\\n✓ Saved: test_image_sample.png\")\n        plt.show()\n\n# ============================================================================\n# 9. KEY INSIGHTS SUMMARY\n# ============================================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"KEY INSIGHTS FOR MODEL DEVELOPMENT\")\nprint(\"=\" * 80)\n\nprint(\"\"\"\n1. IMAGE PROCESSING CHALLENGES:\n   - Handle 9 different degradation types per ECG\n   - Grid lines need to be detected and removed\n   - Various noise patterns (stains, mold, damage)\n   - Different color spaces (color vs B&W)\n\n2. SIGNAL EXTRACTION:\n   - 12 leads arranged in standard grid layout\n   - Lead II: 10 seconds (4x longer than others)\n   - Other leads: 2.5 seconds each\n   - Need precise pixel-to-mV and pixel-to-time mapping\n\n3. SUBMISSION FORMAT:\n   - One row per (base_id, row_id, lead) combination\n   - Lead II: floor(fs * 10) rows\n   - Other leads: floor(fs * 2.5) rows\n   - Total: ~1000 test images × 12 leads each\n\n4. SUGGESTED APPROACH:\n   - Start with clean images (segment 0001)\n   - Detect ECG grid and lead boundaries\n   - Extract signal trace from each lead region\n   - Convert pixels to mV values\n   - Handle degraded images progressively\n\n5. EVALUATION:\n   - Likely MSE or MAE between predicted and ground truth\n   - Time alignment is critical (small shifts = large errors)\n   - Lead II carries more weight (4x more points)\n\"\"\")\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"EXPLORATION COMPLETE!\")\nprint(\"=\" * 80)\nprint(\"\\nGenerated files:\")\nprint(\"  - degradation_types_overview.png\")\nprint(\"  - image_vs_groundtruth.png\")\nprint(\"  - individual_leads.png\")\nprint(\"  - test_image_sample.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:09.374757Z","iopub.execute_input":"2025-12-10T18:46:09.375059Z","iopub.status.idle":"2025-12-10T18:46:46.729962Z","shell.execute_reply.started":"2025-12-10T18:46:09.375040Z","shell.execute_reply":"2025-12-10T18:46:46.728815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"\"\"\n# ECG Image Digitization - Complete Pipeline\n# Combines classical CV techniques with deep learning for robust signal extraction\n# \"\"\"\n\n# import numpy as np\n# import pandas as pd\n# import cv2\n# from PIL import Image\n# from pathlib import Path\n# from scipy import signal as scipy_signal, interpolate\n# from scipy.ndimage import gaussian_filter1d\n# from skimage import morphology, filters\n# import matplotlib.pyplot as plt\n# from tqdm import tqdm\n\n# # ============================================================================\n# # PART 1: CLASSICAL SIGNAL PROCESSING & PREPROCESSING\n# # ============================================================================\n\n# class ECGPreprocessor:\n#     \"\"\"Handles image preprocessing: grid removal, denoising, normalization\"\"\"\n    \n#     def __init__(self):\n#         self.grid_color_ranges = {\n#             'pink': ([150, 50, 100], [255, 150, 200]),  # Pink grid lines\n#             'red': ([0, 0, 150], [100, 100, 255]),      # Red grid lines\n#             'light_red': ([180, 180, 200], [255, 255, 255])  # Light background\n#         }\n    \n#     def load_image(self, img_path):\n#         \"\"\"Load and convert image to RGB\"\"\"\n#         img = Image.open(img_path)\n#         if img.mode == 'RGBA':\n#             img = img.convert('RGB')\n#         return np.array(img)\n    \n#     def remove_grid_lines(self, img, method='hough'):\n#         \"\"\"Remove grid lines using multiple strategies\"\"\"\n        \n#         if method == 'color_filter':\n#             return self._remove_grid_by_color(img)\n#         elif method == 'hough':\n#             return self._remove_grid_by_hough(img)\n#         elif method == 'morphology':\n#             return self._remove_grid_by_morphology(img)\n#         else:\n#             # Combine all methods\n#             img1 = self._remove_grid_by_color(img)\n#             img2 = self._remove_grid_by_morphology(img1)\n#             return img2\n    \n#     def _remove_grid_by_color(self, img):\n#         \"\"\"Remove grid lines by color filtering\"\"\"\n#         img_copy = img.copy()\n        \n#         # Convert to HSV for better color filtering\n#         hsv = cv2.cvtColor(img, cv2.COLOR_RGB2HSV)\n        \n#         # Create mask for grid colors\n#         mask = np.zeros(img.shape[:2], dtype=np.uint8)\n        \n#         # Detect pink/red grid lines\n#         for color_name, (lower, upper) in self.grid_color_ranges.items():\n#             lower = np.array(lower, dtype=np.uint8)\n#             upper = np.array(upper, dtype=np.uint8)\n#             color_mask = cv2.inRange(img, lower, upper)\n#             mask = cv2.bitwise_or(mask, color_mask)\n        \n#         # Inpaint the grid areas\n#         result = cv2.inpaint(img_copy, mask, 3, cv2.INPAINT_TELEA)\n#         return result\n    \n#     def _remove_grid_by_hough(self, img):\n#         \"\"\"Remove grid lines using Hough line detection\"\"\"\n#         gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n#         edges = cv2.Canny(gray, 50, 150, apertureSize=3)\n        \n#         # Detect lines\n#         lines = cv2.HoughLinesP(edges, 1, np.pi/180, threshold=100,\n#                                 minLineLength=100, maxLineGap=10)\n        \n#         # Create mask of detected lines\n#         mask = np.zeros(gray.shape, dtype=np.uint8)\n#         if lines is not None:\n#             for line in lines:\n#                 x1, y1, x2, y2 = line[0]\n#                 # Only remove horizontal and vertical lines (grid)\n#                 if abs(y1 - y2) < 5 or abs(x1 - x2) < 5:\n#                     cv2.line(mask, (x1, y1), (x2, y2), 255, 2)\n        \n#         # Inpaint\n#         result = cv2.inpaint(img, mask, 3, cv2.INPAINT_TELEA)\n#         return result\n    \n#     def _remove_grid_by_morphology(self, img):\n#         \"\"\"Remove thin grid lines using morphological operations\"\"\"\n#         gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        \n#         # Remove thin horizontal lines\n#         kernel_h = cv2.getStructuringElement(cv2.MORPH_RECT, (40, 1))\n#         detected_h = cv2.morphologyEx(gray, cv2.MORPH_OPEN, kernel_h)\n        \n#         # Remove thin vertical lines\n#         kernel_v = cv2.getStructuringElement(cv2.MORPH_RECT, (1, 40))\n#         detected_v = cv2.morphologyEx(gray, cv2.MORPH_OPEN, kernel_v)\n        \n#         # Combine detected grids\n#         grid_mask = cv2.bitwise_or(detected_h, detected_v)\n#         grid_mask = cv2.threshold(grid_mask, 0, 255, cv2.THRESH_BINARY)[1]\n        \n#         # Inpaint\n#         result = cv2.inpaint(img, grid_mask, 3, cv2.INPAINT_TELEA)\n#         return result\n    \n#     def denoise(self, img, method='bilateral'):\n#         \"\"\"Denoise image while preserving edges\"\"\"\n#         if method == 'bilateral':\n#             return cv2.bilateralFilter(img, 9, 75, 75)\n#         elif method == 'nlm':\n#             return cv2.fastNlMeansDenoisingColored(img, None, 10, 10, 7, 21)\n#         else:\n#             return cv2.GaussianBlur(img, (5, 5), 0)\n    \n#     def binarize(self, img, method='adaptive'):\n#         \"\"\"Convert to binary image for signal extraction\"\"\"\n#         gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        \n#         if method == 'adaptive':\n#             binary = cv2.adaptiveThreshold(gray, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,\n#                                           cv2.THRESH_BINARY_INV, 11, 2)\n#         elif method == 'otsu':\n#             _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)\n#         else:\n#             _, binary = cv2.threshold(gray, 127, 255, cv2.THRESH_BINARY_INV)\n        \n#         return binary\n    \n#     def full_preprocess(self, img_path):\n#         \"\"\"Complete preprocessing pipeline\"\"\"\n#         img = self.load_image(img_path)\n#         img = self.remove_grid_lines(img, method='combined')\n#         img = self.denoise(img, method='bilateral')\n#         binary = self.binarize(img, method='adaptive')\n#         return img, binary\n\n\n# # ============================================================================\n# # PART 2: LEAD DETECTION & SEGMENTATION\n# # ============================================================================\n\n# class LeadDetector:\n#     \"\"\"Detect and segment individual ECG leads from the image\"\"\"\n    \n#     def __init__(self):\n#         self.lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n#                           'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n#         self.layout = 'standard_4x3'  # 4 rows, 3 columns\n    \n#     def detect_lead_boundaries(self, binary_img):\n#         \"\"\"Detect boundaries of each lead strip\"\"\"\n#         # Project horizontally to find row separations\n#         h_projection = np.sum(binary_img, axis=1)\n        \n#         # Find valleys (lead separators) in projection\n#         valleys = self._find_valleys(h_projection, min_distance=100)\n        \n#         # Divide into rows\n#         if len(valleys) >= 3:  # Expecting 4 rows\n#             row_boundaries = sorted(valleys[:3])\n#         else:\n#             # Fallback: divide equally\n#             height = binary_img.shape[0]\n#             row_boundaries = [height//4, height//2, 3*height//4]\n        \n#         # Divide columns similarly\n#         v_projection = np.sum(binary_img, axis=0)\n#         h_valleys = self._find_valleys(v_projection, min_distance=100)\n        \n#         if len(h_valleys) >= 2:\n#             col_boundaries = sorted(h_valleys[:2])\n#         else:\n#             width = binary_img.shape[1]\n#             col_boundaries = [width//3, 2*width//3]\n        \n#         return self._create_lead_regions(binary_img.shape, row_boundaries, col_boundaries)\n    \n#     def _find_valleys(self, projection, min_distance=50):\n#         \"\"\"Find valleys in projection (low signal areas)\"\"\"\n#         # Smooth projection\n#         smoothed = gaussian_filter1d(projection, sigma=10)\n        \n#         # Find local minima\n#         valleys = []\n#         for i in range(min_distance, len(smoothed) - min_distance):\n#             if smoothed[i] < smoothed[i-min_distance:i].min() * 1.1 and \\\n#                smoothed[i] < smoothed[i+1:i+min_distance].min() * 1.1:\n#                 valleys.append(i)\n        \n#         return valleys\n    \n#     def _create_lead_regions(self, img_shape, row_boundaries, col_boundaries):\n#         \"\"\"Create bounding boxes for each lead\"\"\"\n#         height, width = img_shape\n#         rows = [0] + row_boundaries + [height]\n#         cols = [0] + col_boundaries + [width]\n        \n#         regions = {}\n#         lead_idx = 0\n        \n#         # Standard 4x3 layout with Lead II spanning full width in last row\n#         for r in range(len(rows) - 1):\n#             if r < 3:  # First 3 rows: 3 leads each\n#                 for c in range(len(cols) - 1):\n#                     if lead_idx < len(self.lead_names):\n#                         regions[self.lead_names[lead_idx]] = {\n#                             'y1': rows[r], 'y2': rows[r+1],\n#                             'x1': cols[c], 'x2': cols[c+1]\n#                         }\n#                         lead_idx += 1\n#             else:  # Last row: Lead II spans full width\n#                 if lead_idx < len(self.lead_names):\n#                     regions[self.lead_names[lead_idx]] = {\n#                         'y1': rows[r], 'y2': rows[r+1],\n#                         'x1': 0, 'x2': width\n#                     }\n#                     lead_idx += 1\n        \n#         return regions\n\n\n# # ============================================================================\n# # PART 3: SIGNAL TRACE EXTRACTION\n# # ============================================================================\n\n# class SignalExtractor:\n#     \"\"\"Extract signal traces from binary images\"\"\"\n    \n#     def extract_trace(self, binary_region, method='skeleton'):\n#         \"\"\"Extract the signal trace from a lead region\"\"\"\n#         if method == 'skeleton':\n#             return self._extract_via_skeleton(binary_region)\n#         elif method == 'column_scan':\n#             return self._extract_via_column_scan(binary_region)\n#         else:\n#             return self._extract_hybrid(binary_region)\n    \n#     def _extract_via_skeleton(self, binary_region):\n#         \"\"\"Extract trace using morphological skeletonization\"\"\"\n#         # Skeletonize the signal\n#         skeleton = morphology.skeletonize(binary_region > 0)\n        \n#         # Extract coordinates\n#         trace_points = []\n#         width = skeleton.shape[1]\n        \n#         for x in range(width):\n#             column = skeleton[:, x]\n#             y_coords = np.where(column)[0]\n#             if len(y_coords) > 0:\n#                 # Take median if multiple points\n#                 y = int(np.median(y_coords))\n#                 trace_points.append((x, y))\n        \n#         return np.array(trace_points)\n    \n#     def _extract_via_column_scan(self, binary_region):\n#         \"\"\"Extract trace by scanning each column\"\"\"\n#         height, width = binary_region.shape\n#         trace_points = []\n        \n#         for x in range(width):\n#             column = binary_region[:, x]\n            \n#             # Find center of mass in this column\n#             y_coords = np.where(column > 0)[0]\n            \n#             if len(y_coords) > 0:\n#                 # Weight by intensity\n#                 weights = column[y_coords]\n#                 y = int(np.average(y_coords, weights=weights))\n#                 trace_points.append((x, y))\n#             elif len(trace_points) > 0:\n#                 # Interpolate from previous point\n#                 trace_points.append((x, trace_points[-1][1]))\n        \n#         return np.array(trace_points)\n    \n#     def _extract_hybrid(self, binary_region):\n#         \"\"\"Combine both methods for robustness\"\"\"\n#         skel_trace = self._extract_via_skeleton(binary_region)\n#         scan_trace = self._extract_via_column_scan(binary_region)\n        \n#         # Average the two\n#         if len(skel_trace) > 0 and len(scan_trace) > 0:\n#             # Align and average\n#             min_len = min(len(skel_trace), len(scan_trace))\n#             combined = np.zeros((min_len, 2))\n#             combined[:, 0] = skel_trace[:min_len, 0]\n#             combined[:, 1] = (skel_trace[:min_len, 1] + scan_trace[:min_len, 1]) / 2\n#             return combined\n        \n#         return skel_trace if len(skel_trace) > 0 else scan_trace\n    \n#     def convert_to_signal(self, trace_points, target_length, fs, lead_duration):\n#         \"\"\"Convert pixel coordinates to signal values\"\"\"\n#         if len(trace_points) == 0:\n#             return np.zeros(target_length)\n        \n#         # Flip y-coordinates (image y increases downward)\n#         x_coords = trace_points[:, 0]\n#         y_coords = trace_points[:, 1]\n        \n#         # Normalize x to time domain\n#         x_normalized = np.linspace(0, lead_duration, len(x_coords))\n        \n#         # Normalize y to voltage (assuming grid calibration)\n#         y_mean = np.mean(y_coords)\n#         y_std = np.std(y_coords)\n        \n#         # Rough calibration: 1mm = 0.1mV, estimate pixels per mm\n#         # This needs calibration from known ECG standards\n#         voltage = -(y_coords - y_mean) / (y_std + 1e-8)  # Flip and normalize\n        \n#         # Interpolate to target sampling rate\n#         target_time = np.linspace(0, lead_duration, target_length)\n#         interpolator = interpolate.interp1d(x_normalized, voltage, \n#                                            kind='cubic', fill_value='extrapolate')\n#         signal = interpolator(target_time)\n        \n#         # Apply bandpass filter (typical ECG: 0.5-150 Hz)\n#         if fs > 2:\n#             sos = scipy_signal.butter(4, [0.5, min(150, fs/2-1)], \n#                                      btype='band', fs=fs, output='sos')\n#             signal = scipy_signal.sosfilt(sos, signal)\n        \n#         return signal\n\n\n# # ============================================================================\n# # PART 4: COMPLETE BASELINE PIPELINE\n# # ============================================================================\n\n# class ECGDigitizer:\n#     \"\"\"Complete pipeline for ECG digitization\"\"\"\n    \n#     def __init__(self):\n#         self.preprocessor = ECGPreprocessor()\n#         self.lead_detector = LeadDetector()\n#         self.signal_extractor = SignalExtractor()\n    \n#     def digitize(self, img_path, fs, lead_durations=None):\n#         \"\"\"\n#         Digitize an ECG image to signals\n        \n#         Args:\n#             img_path: Path to ECG image\n#             fs: Sampling frequency\n#             lead_durations: Dict mapping lead names to durations (seconds)\n#                           Default: Lead II = 10s, others = 2.5s\n        \n#         Returns:\n#             Dictionary of lead_name -> signal array\n#         \"\"\"\n#         if lead_durations is None:\n#             lead_durations = {\n#                 'II': 10.0,\n#                 **{lead: 2.5 for lead in self.lead_detector.lead_names if lead != 'II'}\n#             }\n        \n#         # Preprocess\n#         print(f\"Preprocessing {img_path}...\")\n#         img, binary = self.preprocessor.full_preprocess(img_path)\n        \n#         # Detect leads\n#         print(\"Detecting lead regions...\")\n#         regions = self.lead_detector.detect_lead_boundaries(binary)\n        \n#         # Extract signals\n#         print(\"Extracting signals...\")\n#         signals = {}\n        \n#         for lead_name in self.lead_detector.lead_names:\n#             if lead_name not in regions:\n#                 print(f\"Warning: Lead {lead_name} not found in regions\")\n#                 continue\n            \n#             region = regions[lead_name]\n#             binary_crop = binary[region['y1']:region['y2'], region['x1']:region['x2']]\n            \n#             # Extract trace\n#             trace = self.signal_extractor.extract_trace(binary_crop, method='hybrid')\n            \n#             # Convert to signal\n#             duration = lead_durations.get(lead_name, 2.5)\n#             target_length = int(fs * duration)\n#             signal = self.signal_extractor.convert_to_signal(\n#                 trace, target_length, fs, duration\n#             )\n            \n#             signals[lead_name] = signal\n#             print(f\"  {lead_name}: {len(signal)} samples\")\n        \n#         return signals\n\n\n# # ============================================================================\n# # PART 5: INFERENCE & SUBMISSION GENERATION\n# # ============================================================================\n\n# def generate_submission(test_csv_path, test_images_path, output_path):\n#     \"\"\"Generate submission file for Kaggle\"\"\"\n    \n#     test_df = pd.read_csv(test_csv_path)\n#     digitizer = ECGDigitizer()\n    \n#     submission_rows = []\n    \n#     # Group by image ID\n#     for img_id in tqdm(test_df['id'].unique(), desc=\"Processing test images\"):\n#         img_path = Path(test_images_path) / f\"{img_id}.png\"\n        \n#         if not img_path.exists():\n#             print(f\"Warning: Image {img_id} not found\")\n#             continue\n        \n#         # Get metadata for this image\n#         img_meta = test_df[test_df['id'] == img_id].iloc[0]\n#         fs = img_meta['fs']\n        \n#         # Digitize\n#         try:\n#             signals = digitizer.digitize(img_path, fs)\n            \n#             # Create submission rows\n#             for _, row in test_df[test_df['id'] == img_id].iterrows():\n#                 lead = row['lead']\n#                 num_rows = row['number_of_rows']\n                \n#                 if lead in signals:\n#                     signal = signals[lead][:num_rows]  # Truncate to expected length\n                    \n#                     # Pad if necessary\n#                     if len(signal) < num_rows:\n#                         signal = np.pad(signal, (0, num_rows - len(signal)), \n#                                       mode='edge')\n                    \n#                     # Create rows\n#                     for row_id in range(num_rows):\n#                         submission_rows.append({\n#                             'id': f\"{img_id}_{row_id}_{lead}\",\n#                             'value': signal[row_id]\n#                         })\n#                 else:\n#                     print(f\"Warning: Lead {lead} not extracted for {img_id}\")\n#                     # Fill with zeros\n#                     for row_id in range(num_rows):\n#                         submission_rows.append({\n#                             'id': f\"{img_id}_{row_id}_{lead}\",\n#                             'value': 0.0\n#                         })\n        \n#         except Exception as e:\n#             print(f\"Error processing {img_id}: {e}\")\n#             # Fill with zeros\n#             for _, row in test_df[test_df['id'] == img_id].iterrows():\n#                 lead = row['lead']\n#                 num_rows = row['number_of_rows']\n#                 for row_id in range(num_rows):\n#                     submission_rows.append({\n#                         'id': f\"{img_id}_{row_id}_{lead}\",\n#                         'value': 0.0\n#                     })\n    \n#     # Create submission DataFrame\n#     submission_df = pd.DataFrame(submission_rows)\n#     submission_df.to_parquet(output_path, index=False)\n#     print(f\"\\n✓ Submission saved to {output_path}\")\n#     print(f\"  Total rows: {len(submission_df)}\")\n    \n#     return submission_df\n\n\n# # ============================================================================\n# # EXAMPLE USAGE\n# # ============================================================================\n\n# if __name__ == \"__main__\":\n#     # Test on a single image\n#     BASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\n    \n#     # Example: digitize one training image\n#     sample_id = '7663343'\n#     img_path = BASE_PATH / 'train' / sample_id / f'{sample_id}-0001.png'\n    \n#     digitizer = ECGDigitizer()\n#     signals = digitizer.digitize(img_path, fs=500)\n    \n#     # Visualize results\n#     fig, axes = plt.subplots(4, 3, figsize=(20, 16))\n#     axes = axes.flatten()\n    \n#     for idx, (lead, signal) in enumerate(signals.items()):\n#         if idx < 12:\n#             time = np.arange(len(signal)) / 500\n#             axes[idx].plot(time, signal, linewidth=0.8)\n#             axes[idx].set_title(f'Lead {lead}', fontweight='bold')\n#             axes[idx].set_xlabel('Time (s)')\n#             axes[idx].set_ylabel('Amplitude (mV)')\n#             axes[idx].grid(True, alpha=0.3)\n    \n#     plt.tight_layout()\n#     plt.savefig('baseline_extraction_result.png', dpi=150)\n#     plt.show()\n    \n#     # Generate full submission\n#     print(\"\\n\" + \"=\"*80)\n#     print(\"GENERATING SUBMISSION\")\n#     print(\"=\"*80)\n    \n#     submission = generate_submission(\n#         test_csv_path=BASE_PATH / 'test.csv',\n#         test_images_path=BASE_PATH / 'test',\n#         output_path='submission_baseline.parquet'\n#     )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:46.731634Z","iopub.execute_input":"2025-12-10T18:46:46.731922Z","iopub.status.idle":"2025-12-10T18:46:57.387777Z","shell.execute_reply.started":"2025-12-10T18:46:46.731901Z","shell.execute_reply":"2025-12-10T18:46:57.386455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nECG Image Digitization - Complete Pipeline\nCombines classical CV techniques with deep learning for robust signal extraction\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom PIL import Image\nfrom pathlib import Path\nfrom scipy import signal as scipy_signal, interpolate\nfrom scipy.ndimage import gaussian_filter1d\nfrom skimage import morphology, filters\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\n# ============================================================================\n# PART 1: CLASSICAL SIGNAL PROCESSING & PREPROCESSING\n# ============================================================================\n\nclass ECGPreprocessor:\n    \"\"\"Handles image preprocessing: grid removal, denoising, normalization\"\"\"\n    \n    def __init__(self):\n        self.grid_color_ranges = {\n            'pink': ([150, 50, 100], [255, 150, 200]),  # Pink grid lines\n            'red': ([0, 0, 150], [100, 100, 255]),      # Red grid lines\n            'light_red': ([180, 180, 200], [255, 255, 255])  # Light background\n        }\n    \n    def load_image(self, img_path):\n        \"\"\"Load and convert image to RGB\"\"\"\n        img = Image.open(img_path)\n        if img.mode == 'RGBA':\n            img = img.convert('RGB')\n        return np.array(img)\n    \n    def remove_grid_lines(self, img, method='hough'):\n        \"\"\"Remove grid lines using multiple strategies\"\"\"\n        \n        if method == 'color_filter':\n            return self._remove_grid_by_color(img)\n        elif method == 'hough':\n            return self._remove_grid_by_hough(img)\n        elif method == 'morphology':\n            return self._remove_grid_by_morphology(img)\n        else:\n            # Combine all methods\n            img1 = self._remove_grid_by_color(img)\n            img2 = self._remove_grid_by_morphology(img1)\n            return img2\n    \n    def _remove_grid_by_color(self, img):\n        \"\"\"Remove grid lines by color filtering\"\"\"\n        img_copy = img.copy()\n        \n        # Convert to HSV for better color filtering\n        hsv = cv2.cvtColor(img, cv2.COLOR_RGB2HSV)\n        \n        # Create mask for grid colors\n        mask = np.zeros(img.shape[:2], dtype=np.uint8)\n        \n        # Detect pink/red grid lines\n        for color_name, (lower, upper) in self.grid_color_ranges.items():\n            lower = np.array(lower, dtype=np.uint8)\n            upper = np.array(upper, dtype=np.uint8)\n            color_mask = cv2.inRange(img, lower, upper)\n            mask = cv2.bitwise_or(mask, color_mask)\n        \n        # Inpaint the grid areas\n        result = cv2.inpaint(img_copy, mask, 3, cv2.INPAINT_TELEA)\n        return result\n    \n    def _remove_grid_by_hough(self, img):\n        \"\"\"Remove grid lines using Hough line detection\"\"\"\n        gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        edges = cv2.Canny(gray, 50, 150, apertureSize=3)\n        \n        # Detect lines\n        lines = cv2.HoughLinesP(edges, 1, np.pi/180, threshold=100,\n                                minLineLength=100, maxLineGap=10)\n        \n        # Create mask of detected lines\n        mask = np.zeros(gray.shape, dtype=np.uint8)\n        if lines is not None:\n            for line in lines:\n                x1, y1, x2, y2 = line[0]\n                # Only remove horizontal and vertical lines (grid)\n                if abs(y1 - y2) < 5 or abs(x1 - x2) < 5:\n                    cv2.line(mask, (x1, y1), (x2, y2), 255, 2)\n        \n        # Inpaint\n        result = cv2.inpaint(img, mask, 3, cv2.INPAINT_TELEA)\n        return result\n    \n    def _remove_grid_by_morphology(self, img):\n        \"\"\"Remove thin grid lines using morphological operations\"\"\"\n        gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        \n        # Remove thin horizontal lines\n        kernel_h = cv2.getStructuringElement(cv2.MORPH_RECT, (40, 1))\n        detected_h = cv2.morphologyEx(gray, cv2.MORPH_OPEN, kernel_h)\n        \n        # Remove thin vertical lines\n        kernel_v = cv2.getStructuringElement(cv2.MORPH_RECT, (1, 40))\n        detected_v = cv2.morphologyEx(gray, cv2.MORPH_OPEN, kernel_v)\n        \n        # Combine detected grids\n        grid_mask = cv2.bitwise_or(detected_h, detected_v)\n        grid_mask = cv2.threshold(grid_mask, 0, 255, cv2.THRESH_BINARY)[1]\n        \n        # Inpaint\n        result = cv2.inpaint(img, grid_mask, 3, cv2.INPAINT_TELEA)\n        return result\n    \n    def denoise(self, img, method='bilateral'):\n        \"\"\"Denoise image while preserving edges\"\"\"\n        if method == 'bilateral':\n            return cv2.bilateralFilter(img, 9, 75, 75)\n        elif method == 'nlm':\n            return cv2.fastNlMeansDenoisingColored(img, None, 10, 10, 7, 21)\n        else:\n            return cv2.GaussianBlur(img, (5, 5), 0)\n    \n    def binarize(self, img, method='adaptive'):\n        \"\"\"Convert to binary image for signal extraction\"\"\"\n        gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        \n        if method == 'adaptive':\n            binary = cv2.adaptiveThreshold(gray, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,\n                                           cv2.THRESH_BINARY_INV, 11, 2)\n        elif method == 'otsu':\n            _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)\n        else:\n            _, binary = cv2.threshold(gray, 127, 255, cv2.THRESH_BINARY_INV)\n        \n        return binary\n    \n    def full_preprocess(self, img_path):\n        \"\"\"Complete preprocessing pipeline\"\"\"\n        img = self.load_image(img_path)\n        img = self.remove_grid_lines(img, method='combined')\n        img = self.denoise(img, method='bilateral')\n        binary = self.binarize(img, method='adaptive')\n        return img, binary\n\n\n# ============================================================================\n# PART 2: LEAD DETECTION & SEGMENTATION\n# ============================================================================\n\nclass LeadDetector:\n    \"\"\"Detect and segment individual ECG leads from the image\"\"\"\n    \n    def __init__(self):\n        self.lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n                          'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n        self.layout = 'standard_4x3'  # 4 rows, 3 columns\n    \n    def detect_lead_boundaries(self, binary_img):\n        \"\"\"Detect boundaries of each lead strip\"\"\"\n        # Project horizontally to find row separations\n        h_projection = np.sum(binary_img, axis=1)\n        \n        # Find valleys (lead separators) in projection\n        valleys = self._find_valleys(h_projection, min_distance=100)\n        \n        # Divide into rows\n        if len(valleys) >= 3:  # Expecting 4 rows\n            row_boundaries = sorted(valleys[:3])\n        else:\n            # Fallback: divide equally\n            height = binary_img.shape[0]\n            row_boundaries = [height//4, height//2, 3*height//4]\n        \n        # Divide columns similarly\n        v_projection = np.sum(binary_img, axis=0)\n        h_valleys = self._find_valleys(v_projection, min_distance=100)\n        \n        if len(h_valleys) >= 2:\n            col_boundaries = sorted(h_valleys[:2])\n        else:\n            width = binary_img.shape[1]\n            col_boundaries = [width//3, 2*width//3]\n        \n        return self._create_lead_regions(binary_img.shape, row_boundaries, col_boundaries)\n    \n    def _find_valleys(self, projection, min_distance=50):\n        \"\"\"Find valleys in projection (low signal areas)\"\"\"\n        # Smooth projection\n        smoothed = gaussian_filter1d(projection, sigma=10)\n        \n        # Find local minima\n        valleys = []\n        for i in range(min_distance, len(smoothed) - min_distance):\n            if smoothed[i] < smoothed[i-min_distance:i].min() * 1.1 and \\\n               smoothed[i] < smoothed[i+1:i+min_distance].min() * 1.1:\n                valleys.append(i)\n        \n        return valleys\n    \n    def _create_lead_regions(self, img_shape, row_boundaries, col_boundaries):\n        \"\"\"Create bounding boxes for each lead\"\"\"\n        height, width = img_shape\n        rows = [0] + row_boundaries + [height]\n        cols = [0] + col_boundaries + [width]\n        \n        regions = {}\n        lead_idx = 0\n        \n        # Standard 4x3 layout with Lead II spanning full width in last row\n        for r in range(len(rows) - 1):\n            if r < 3:  # First 3 rows: 3 leads each\n                for c in range(len(cols) - 1):\n                    if lead_idx < len(self.lead_names):\n                        regions[self.lead_names[lead_idx]] = {\n                            'y1': rows[r], 'y2': rows[r+1],\n                            'x1': cols[c], 'x2': cols[c+1]\n                        }\n                        lead_idx += 1\n            else:  # Last row: Lead II spans full width\n                if lead_idx < len(self.lead_names):\n                    regions[self.lead_names[lead_idx]] = {\n                        'y1': rows[r], 'y2': rows[r+1],\n                        'x1': 0, 'x2': width\n                    }\n                    lead_idx += 1\n        \n        return regions\n\n\n# ============================================================================\n# PART 3: SIGNAL TRACE EXTRACTION (ROBUST VERSION)\n# ============================================================================\n\nclass SignalExtractor:\n    \"\"\"Extract signal traces from binary images\"\"\"\n    \n    def extract_trace(self, binary_region, method='skeleton'):\n        \"\"\"Extract the signal trace from a lead region\"\"\"\n        if method == 'skeleton':\n            return self._extract_via_skeleton(binary_region)\n        elif method == 'column_scan':\n            return self._extract_via_column_scan(binary_region)\n        else:\n            return self._extract_hybrid(binary_region)\n    \n    def _extract_via_skeleton(self, binary_region):\n        \"\"\"Extract trace using morphological skeletonization\"\"\"\n        # Skeletonize the signal\n        skeleton = morphology.skeletonize(binary_region > 0)\n        \n        # Extract coordinates\n        trace_points = []\n        width = skeleton.shape[1]\n        \n        for x in range(width):\n            column = skeleton[:, x]\n            y_coords = np.where(column)[0]\n            if len(y_coords) > 0:\n                # Take median if multiple points\n                y = int(np.median(y_coords))\n                trace_points.append((x, y))\n        \n        return np.array(trace_points)\n    \n    def _extract_via_column_scan(self, binary_region):\n        \"\"\"Extract trace by scanning each column\"\"\"\n        height, width = binary_region.shape\n        trace_points = []\n        \n        for x in range(width):\n            column = binary_region[:, x]\n            \n            # Find center of mass in this column\n            y_coords = np.where(column > 0)[0]\n            \n            if len(y_coords) > 0:\n                # Weight by intensity\n                weights = column[y_coords]\n                y = int(np.average(y_coords, weights=weights))\n                trace_points.append((x, y))\n            elif len(trace_points) > 0:\n                # Interpolate from previous point (simple hold)\n                trace_points.append((x, trace_points[-1][1]))\n        \n        return np.array(trace_points)\n    \n    def _extract_hybrid(self, binary_region):\n        \"\"\"Combine both methods for robustness\"\"\"\n        skel_trace = self._extract_via_skeleton(binary_region)\n        scan_trace = self._extract_via_column_scan(binary_region)\n        \n        # Average the two\n        if len(skel_trace) > 0 and len(scan_trace) > 0:\n            # Align and average\n            min_len = min(len(skel_trace), len(scan_trace))\n            combined = np.zeros((min_len, 2))\n            combined[:, 0] = skel_trace[:min_len, 0]\n            combined[:, 1] = (skel_trace[:min_len, 1] + scan_trace[:min_len, 1]) / 2\n            return combined\n        \n        return skel_trace if len(skel_trace) > 0 else scan_trace\n    \n    def convert_to_signal(self, trace_points, target_length, fs, lead_duration):\n        \"\"\"Convert pixel coordinates to signal values with robust interpolation\"\"\"\n        # 1. Handle Empty or Extremely Sparse Data\n        if len(trace_points) < 2:\n            return np.zeros(target_length)\n        \n        # Flip y-coordinates (image y increases downward)\n        x_coords = trace_points[:, 0]\n        y_coords = trace_points[:, 1]\n        \n        # 2. Sort by X to ensure monotonicity (crucial for interpolation)\n        sort_idxs = np.argsort(x_coords)\n        x_coords = x_coords[sort_idxs]\n        y_coords = y_coords[sort_idxs]\n\n        # 3. Remove duplicate X values if any (keep first)\n        _, unique_indices = np.unique(x_coords, return_index=True)\n        x_coords = x_coords[unique_indices]\n        y_coords = y_coords[unique_indices]\n\n        # Normalize x to time domain\n        x_normalized = np.linspace(0, lead_duration, len(x_coords))\n        \n        # Normalize y to voltage\n        y_mean = np.mean(y_coords)\n        y_std = np.std(y_coords)\n        voltage = -(y_coords - y_mean) / (y_std + 1e-8) \n        \n        # 4. Adaptive Interpolation Strategy\n        # Cubic spline requires at least 4 points. Fallback to linear if fewer.\n        kind = 'cubic' if len(x_coords) > 3 else 'linear'\n        \n        try:\n            target_time = np.linspace(0, lead_duration, target_length)\n            interpolator = interpolate.interp1d(x_normalized, voltage, \n                                              kind=kind, fill_value='extrapolate')\n            signal = interpolator(target_time)\n        except Exception as e:\n            # Absolute fallback if spline fails\n            print(f\"Interpolation fallback (linear) due to: {e}\")\n            interpolator = interpolate.interp1d(x_normalized, voltage, \n                                              kind='linear', fill_value='extrapolate')\n            signal = interpolator(target_time)\n        \n        # Apply bandpass filter (typical ECG: 0.5-150 Hz)\n        if fs > 2 and len(signal) > 10:\n            try:\n                sos = scipy_signal.butter(4, [0.5, min(150, fs/2-1)], \n                                     btype='band', fs=fs, output='sos')\n                signal = scipy_signal.sosfilt(sos, signal)\n            except ValueError:\n                pass\n        \n        return signal\n\n\n# ============================================================================\n# PART 4: COMPLETE BASELINE PIPELINE\n# ============================================================================\n\nclass ECGDigitizer:\n    \"\"\"Complete pipeline for ECG digitization\"\"\"\n    \n    def __init__(self):\n        self.preprocessor = ECGPreprocessor()\n        self.lead_detector = LeadDetector()\n        self.signal_extractor = SignalExtractor()\n    \n    def digitize(self, img_path, fs, lead_durations=None):\n        \"\"\"\n        Digitize an ECG image to signals\n        \n        Args:\n            img_path: Path to ECG image\n            fs: Sampling frequency\n            lead_durations: Dict mapping lead names to durations (seconds)\n                            Default: Lead II = 10s, others = 2.5s\n        \n        Returns:\n            Dictionary of lead_name -> signal array\n        \"\"\"\n        if lead_durations is None:\n            lead_durations = {\n                'II': 10.0,\n                **{lead: 2.5 for lead in self.lead_detector.lead_names if lead != 'II'}\n            }\n        \n        # Preprocess\n        print(f\"Preprocessing {img_path}...\")\n        img, binary = self.preprocessor.full_preprocess(img_path)\n        \n        # Detect leads\n        print(\"Detecting lead regions...\")\n        regions = self.lead_detector.detect_lead_boundaries(binary)\n        \n        # Extract signals\n        print(\"Extracting signals...\")\n        signals = {}\n        \n        for lead_name in self.lead_detector.lead_names:\n            if lead_name not in regions:\n                print(f\"Warning: Lead {lead_name} not found in regions\")\n                continue\n            \n            region = regions[lead_name]\n            binary_crop = binary[region['y1']:region['y2'], region['x1']:region['x2']]\n            \n            # Extract trace\n            trace = self.signal_extractor.extract_trace(binary_crop, method='hybrid')\n            \n            # Convert to signal\n            duration = lead_durations.get(lead_name, 2.5)\n            target_length = int(fs * duration)\n            signal = self.signal_extractor.convert_to_signal(\n                trace, target_length, fs, duration\n            )\n            \n            signals[lead_name] = signal\n            print(f\"  {lead_name}: {len(signal)} samples\")\n        \n        return signals\n\n\n# ============================================================================\n# PART 5: INFERENCE & SUBMISSION GENERATION\n# ============================================================================\n\ndef generate_submission(test_csv_path, test_images_path, output_path):\n    \"\"\"Generate submission file for Kaggle\"\"\"\n    \n    test_df = pd.read_csv(test_csv_path)\n    digitizer = ECGDigitizer()\n    \n    submission_rows = []\n    \n    # Group by image ID\n    for img_id in tqdm(test_df['id'].unique(), desc=\"Processing test images\"):\n        img_path = Path(test_images_path) / f\"{img_id}.png\"\n        \n        if not img_path.exists():\n            print(f\"Warning: Image {img_id} not found\")\n            continue\n        \n        # Get metadata for this image\n        img_meta = test_df[test_df['id'] == img_id].iloc[0]\n        fs = img_meta['fs']\n        \n        # Digitize\n        try:\n            signals = digitizer.digitize(img_path, fs)\n            \n            # Create submission rows\n            for _, row in test_df[test_df['id'] == img_id].iterrows():\n                lead = row['lead']\n                num_rows = row['number_of_rows']\n                \n                if lead in signals:\n                    signal = signals[lead][:num_rows]  # Truncate to expected length\n                    \n                    # Pad if necessary\n                    if len(signal) < num_rows:\n                        signal = np.pad(signal, (0, num_rows - len(signal)), \n                                      mode='edge')\n                    \n                    # Create rows\n                    for row_id in range(num_rows):\n                        submission_rows.append({\n                            'id': f\"{img_id}_{row_id}_{lead}\",\n                            'value': signal[row_id]\n                        })\n                else:\n                    print(f\"Warning: Lead {lead} not extracted for {img_id}\")\n                    # Fill with zeros\n                    for row_id in range(num_rows):\n                        submission_rows.append({\n                            'id': f\"{img_id}_{row_id}_{lead}\",\n                            'value': 0.0\n                        })\n        \n        except Exception as e:\n            print(f\"Error processing {img_id}: {e}\")\n            # Fill with zeros\n            for _, row in test_df[test_df['id'] == img_id].iterrows():\n                lead = row['lead']\n                num_rows = row['number_of_rows']\n                for row_id in range(num_rows):\n                    submission_rows.append({\n                        'id': f\"{img_id}_{row_id}_{lead}\",\n                        'value': 0.0\n                    })\n    \n    # Create submission DataFrame\n    submission_df = pd.DataFrame(submission_rows)\n    submission_df.to_csv(output_path, index=False)\n    print(f\"\\n✓ Submission saved to {output_path}\")\n    print(f\"  Total rows: {len(submission_df)}\")\n    \n    return submission_df\n\n\n# ============================================================================\n# EXAMPLE USAGE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    # Test on a single image\n    BASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\n    \n    # Example: digitize one training image (if available)\n    # Check if directory exists before running example\n    if BASE_PATH.exists():\n        sample_id = '7663343' # Replace with an ID that exists in your dataset\n        img_path = BASE_PATH / 'train' / sample_id / f'{sample_id}-0001.png'\n        \n        if img_path.exists():\n            digitizer = ECGDigitizer()\n            signals = digitizer.digitize(img_path, fs=500)\n            \n            # Visualize results\n            fig, axes = plt.subplots(4, 3, figsize=(20, 16))\n            axes = axes.flatten()\n            \n            for idx, (lead, signal) in enumerate(signals.items()):\n                if idx < 12:\n                    time = np.arange(len(signal)) / 500\n                    axes[idx].plot(time, signal, linewidth=0.8)\n                    axes[idx].set_title(f'Lead {lead}', fontweight='bold')\n                    axes[idx].set_xlabel('Time (s)')\n                    axes[idx].set_ylabel('Amplitude (mV)')\n                    axes[idx].grid(True, alpha=0.3)\n            \n            plt.tight_layout()\n            plt.savefig('baseline_extraction_result.png', dpi=150)\n            plt.show()\n        \n        # Generate full submission\n        print(\"\\n\" + \"=\"*80)\n        print(\"GENERATING SUBMISSION\")\n        print(\"=\"*80)\n        \n        if (BASE_PATH / 'test.csv').exists():\n            submission = generate_submission(\n                test_csv_path=BASE_PATH / 'test.csv',\n                test_images_path=BASE_PATH / 'test',\n                output_path='submission.csv'\n            )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:57.388395Z","iopub.status.idle":"2025-12-10T18:46:57.388665Z","shell.execute_reply.started":"2025-12-10T18:46:57.388536Z","shell.execute_reply":"2025-12-10T18:46:57.388548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"\"\"\n# ECG Image Digitization - Deep Learning Solution\n# Advanced models: UNet for segmentation, Vision Transformers, and ensemble methods\n# \"\"\"\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# import numpy as np\n# import pandas as pd\n# import cv2\n# from PIL import Image\n# from pathlib import Path\n# import albumentations as A\n# from albumentations.pytorch import ToTensorV2\n# from tqdm import tqdm\n# import segmentation_models_pytorch as smp\n\n# # ============================================================================\n# # PART 1: CUSTOM DATASET\n# # ============================================================================\n\n# class ECGImageDataset(Dataset):\n#     \"\"\"Dataset for ECG image-to-signal conversion\"\"\"\n    \n#     def __init__(self, df, base_path, segment_id='0001', transform=None, \n#                  mode='train', img_size=(512, 512)):\n#         self.df = df\n#         self.base_path = Path(base_path)\n#         self.segment_id = segment_id\n#         self.transform = transform\n#         self.mode = mode\n#         self.img_size = img_size\n#         self.lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n#                           'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    \n#     def __len__(self):\n#         return len(self.df)\n    \n#     def __getitem__(self, idx):\n#         row = self.df.iloc[idx]\n#         img_id = row['id']\n#         fs = row['fs']\n        \n#         # Load image\n#         if self.mode == 'train':\n#             img_path = self.base_path / 'train' / str(img_id) / f'{img_id}-{self.segment_id}.png'\n#         else:\n#             img_path = self.base_path / 'test' / f'{img_id}.png'\n        \n#         img = cv2.imread(str(img_path))\n#         img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n#         # Resize\n#         img = cv2.resize(img, self.img_size)\n        \n#         if self.mode == 'train':\n#             # Load ground truth signals\n#             signal_path = self.base_path / 'train' / str(img_id) / f'{img_id}.csv'\n#             signals_df = pd.read_csv(signal_path)\n            \n#             # Create target tensor (12 leads)\n#             targets = []\n#             for lead in self.lead_names:\n#                 signal = signals_df[lead].values\n#                 signal = signal[~np.isnan(signal)]  # Remove NaN\n#                 targets.append(signal)\n            \n#             # Apply augmentations\n#             if self.transform:\n#                 augmented = self.transform(image=img)\n#                 img = augmented['image']\n#             else:\n#                 img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n            \n#             return {\n#                 'image': img,\n#                 'targets': targets,\n#                 'fs': fs,\n#                 'img_id': img_id\n#             }\n        \n#         else:  # test mode\n#             if self.transform:\n#                 augmented = self.transform(image=img)\n#                 img = augmented['image']\n#             else:\n#                 img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n            \n#             return {\n#                 'image': img,\n#                 'fs': fs,\n#                 'img_id': img_id\n#             }\n\n\n# # ============================================================================\n# # PART 2: UNET-BASED SEGMENTATION MODEL\n# # ============================================================================\n\n# class ECGSegmentationUNet(nn.Module):\n#     \"\"\"\n#     UNet for ECG signal segmentation\n#     Segments the ECG traces from the image\n#     \"\"\"\n    \n#     def __init__(self, encoder_name='resnet34', encoder_weights='imagenet', \n#                  in_channels=3, classes=12):\n#         super().__init__()\n        \n#         # Use segmentation_models_pytorch for robust UNet\n#         self.model = smp.Unet(\n#             encoder_name=encoder_name,\n#             encoder_weights=encoder_weights,\n#             in_channels=in_channels,\n#             classes=classes,\n#             activation=None\n#         )\n    \n#     def forward(self, x):\n#         return self.model(x)\n\n\n# class DoubleConv(nn.Module):\n#     \"\"\"Double convolution block for UNet\"\"\"\n#     def __init__(self, in_ch, out_ch):\n#         super().__init__()\n#         self.conv = nn.Sequential(\n#             nn.Conv2d(in_ch, out_ch, 3, padding=1),\n#             nn.BatchNorm2d(out_ch),\n#             nn.ReLU(inplace=True),\n#             nn.Conv2d(out_ch, out_ch, 3, padding=1),\n#             nn.BatchNorm2d(out_ch),\n#             nn.ReLU(inplace=True)\n#         )\n    \n#     def forward(self, x):\n#         return self.conv(x)\n\n\n# class CustomUNet(nn.Module):\n#     \"\"\"Custom UNet for ECG segmentation with attention\"\"\"\n    \n#     def __init__(self, in_channels=3, num_leads=12):\n#         super().__init__()\n        \n#         # Encoder\n#         self.enc1 = DoubleConv(in_channels, 64)\n#         self.enc2 = DoubleConv(64, 128)\n#         self.enc3 = DoubleConv(128, 256)\n#         self.enc4 = DoubleConv(256, 512)\n        \n#         self.pool = nn.MaxPool2d(2)\n        \n#         # Bottleneck\n#         self.bottleneck = DoubleConv(512, 1024)\n        \n#         # Decoder\n#         self.up4 = nn.ConvTranspose2d(1024, 512, 2, stride=2)\n#         self.dec4 = DoubleConv(1024, 512)\n        \n#         self.up3 = nn.ConvTranspose2d(512, 256, 2, stride=2)\n#         self.dec3 = DoubleConv(512, 256)\n        \n#         self.up2 = nn.ConvTranspose2d(256, 128, 2, stride=2)\n#         self.dec2 = DoubleConv(256, 128)\n        \n#         self.up1 = nn.ConvTranspose2d(128, 64, 2, stride=2)\n#         self.dec1 = DoubleConv(128, 64)\n        \n#         # Output: 12 channels for 12 leads\n#         self.out = nn.Conv2d(64, num_leads, 1)\n    \n#     def forward(self, x):\n#         # Encoder\n#         e1 = self.enc1(x)\n#         e2 = self.enc2(self.pool(e1))\n#         e3 = self.enc3(self.pool(e2))\n#         e4 = self.enc4(self.pool(e3))\n        \n#         # Bottleneck\n#         b = self.bottleneck(self.pool(e4))\n        \n#         # Decoder\n#         d4 = self.dec4(torch.cat([self.up4(b), e4], dim=1))\n#         d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1))\n#         d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))\n#         d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))\n        \n#         return self.out(d1)\n\n\n# # ============================================================================\n# # PART 3: SIGNAL REGRESSION MODEL\n# # ============================================================================\n\n# class ECGSignalRegressor(nn.Module):\n#     \"\"\"\n#     End-to-end model that predicts signal values directly from images\n#     Uses CNN encoder + LSTM decoder\n#     \"\"\"\n    \n#     def __init__(self, img_size=(512, 512), num_leads=12, max_seq_len=10000):\n#         super().__init__()\n        \n#         # CNN Encoder (extract features from image)\n#         self.encoder = nn.Sequential(\n#             nn.Conv2d(3, 64, 7, stride=2, padding=3),\n#             nn.BatchNorm2d(64),\n#             nn.ReLU(inplace=True),\n#             nn.MaxPool2d(3, stride=2, padding=1),\n            \n#             # ResNet-style blocks\n#             self._make_layer(64, 128, 3),\n#             self._make_layer(128, 256, 4),\n#             self._make_layer(256, 512, 6),\n            \n#             nn.AdaptiveAvgPool2d((1, 1))\n#         )\n        \n#         # LSTM Decoder (generate sequences)\n#         self.lstm = nn.LSTM(512, 512, num_layers=2, batch_first=True, \n#                            dropout=0.3, bidirectional=True)\n        \n#         # Lead-specific heads\n#         self.lead_heads = nn.ModuleList([\n#             nn.Linear(1024, 1) for _ in range(num_leads)\n#         ])\n        \n#         self.max_seq_len = max_seq_len\n    \n#     def _make_layer(self, in_ch, out_ch, num_blocks):\n#         layers = []\n#         layers.append(nn.Conv2d(in_ch, out_ch, 3, stride=2, padding=1))\n#         layers.append(nn.BatchNorm2d(out_ch))\n#         layers.append(nn.ReLU(inplace=True))\n        \n#         for _ in range(num_blocks - 1):\n#             layers.append(nn.Conv2d(out_ch, out_ch, 3, padding=1))\n#             layers.append(nn.BatchNorm2d(out_ch))\n#             layers.append(nn.ReLU(inplace=True))\n        \n#         return nn.Sequential(*layers)\n    \n#     def forward(self, x, seq_lengths):\n#         \"\"\"\n#         Args:\n#             x: (B, 3, H, W) - batch of images\n#             seq_lengths: (B, 12) - sequence length for each lead\n        \n#         Returns:\n#             List of (B, seq_len) tensors for each lead\n#         \"\"\"\n#         batch_size = x.size(0)\n        \n#         # Extract features\n#         features = self.encoder(x)  # (B, 512, 1, 1)\n#         features = features.view(batch_size, -1)  # (B, 512)\n        \n#         # Generate sequences for each lead\n#         lead_outputs = []\n        \n#         for lead_idx in range(12):\n#             # Get max sequence length for this lead in batch\n#             max_len = int(seq_lengths[:, lead_idx].max().item())\n            \n#             # Repeat features for sequence\n#             seq_features = features.unsqueeze(1).repeat(1, max_len, 1)  # (B, max_len, 512)\n            \n#             # LSTM\n#             lstm_out, _ = self.lstm(seq_features)  # (B, max_len, 1024)\n            \n#             # Lead-specific prediction\n#             lead_pred = self.lead_heads[lead_idx](lstm_out).squeeze(-1)  # (B, max_len)\n            \n#             lead_outputs.append(lead_pred)\n        \n#         return lead_outputs\n\n\n# # ============================================================================\n# # PART 4: VISION TRANSFORMER APPROACH\n# # ============================================================================\n\n# class ECGViT(nn.Module):\n#     \"\"\"\n#     Vision Transformer for ECG image processing\n#     Can capture long-range dependencies in ECG traces\n#     \"\"\"\n    \n#     def __init__(self, img_size=512, patch_size=16, in_channels=3, \n#                  num_leads=12, embed_dim=768, depth=12, num_heads=12):\n#         super().__init__()\n        \n#         self.patch_size = patch_size\n#         self.num_patches = (img_size // patch_size) ** 2\n        \n#         # Patch embedding\n#         self.patch_embed = nn.Conv2d(in_channels, embed_dim, \n#                                      kernel_size=patch_size, stride=patch_size)\n        \n#         # Positional embedding\n#         self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, embed_dim))\n        \n#         # Transformer encoder\n#         encoder_layer = nn.TransformerEncoderLayer(\n#             d_model=embed_dim, \n#             nhead=num_heads, \n#             dim_feedforward=4*embed_dim,\n#             dropout=0.1,\n#             activation='gelu',\n#             batch_first=True\n#         )\n#         self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=depth)\n        \n#         # Signal generation heads\n#         self.lead_heads = nn.ModuleList([\n#             nn.Sequential(\n#                 nn.Linear(embed_dim, 512),\n#                 nn.ReLU(),\n#                 nn.Dropout(0.1),\n#                 nn.Linear(512, 1)\n#             ) for _ in range(num_leads)\n#         ])\n    \n#     def forward(self, x):\n#         B = x.shape[0]\n        \n#         # Patch embedding\n#         x = self.patch_embed(x)  # (B, embed_dim, H', W')\n#         x = x.flatten(2).transpose(1, 2)  # (B, num_patches, embed_dim)\n        \n#         # Add positional embedding\n#         x = x + self.pos_embed\n        \n#         # Transformer\n#         x = self.transformer(x)  # (B, num_patches, embed_dim)\n        \n#         # Global average pooling\n#         x = x.mean(dim=1)  # (B, embed_dim)\n        \n#         # Generate lead predictions (placeholder - needs proper sequence generation)\n#         lead_outputs = [head(x) for head in self.lead_heads]\n        \n#         return lead_outputs\n\n\n# # ============================================================================\n# # PART 5: TRAINING UTILITIES\n# # ============================================================================\n\n# class ECGLoss(nn.Module):\n#     \"\"\"Custom loss for ECG signal prediction\"\"\"\n    \n#     def __init__(self, mse_weight=1.0, dtw_weight=0.0):\n#         super().__init__()\n#         self.mse_weight = mse_weight\n#         self.dtw_weight = dtw_weight\n    \n#     def forward(self, predictions, targets):\n#         \"\"\"\n#         Args:\n#             predictions: List of (B, seq_len) tensors\n#             targets: List of (B, seq_len) tensors\n#         \"\"\"\n#         total_loss = 0.0\n        \n#         for pred, target in zip(predictions, targets):\n#             # Handle variable lengths\n#             min_len = min(pred.size(1), target.size(1))\n#             pred = pred[:, :min_len]\n#             target = target[:, :min_len]\n            \n#             # MSE loss\n#             mse_loss = F.mse_loss(pred, target)\n#             total_loss += self.mse_weight * mse_loss\n        \n#         return total_loss / len(predictions)\n\n\n# def train_model(model, train_loader, val_loader, epochs=50, device='cuda'):\n#     \"\"\"Train the ECG model\"\"\"\n    \n#     model = model.to(device)\n#     optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)\n#     scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs)\n#     criterion = ECGLoss()\n    \n#     best_val_loss = float('inf')\n    \n#     for epoch in range(epochs):\n#         # Training\n#         model.train()\n#         train_loss = 0.0\n        \n#         for batch in tqdm(train_loader, desc=f'Epoch {epoch+1}/{epochs}'):\n#             images = batch['image'].to(device)\n#             targets = batch['targets']\n            \n#             optimizer.zero_grad()\n            \n#             # Forward pass\n#             outputs = model(images)\n            \n#             # Compute loss\n#             loss = criterion(outputs, targets)\n            \n#             # Backward pass\n#             loss.backward()\n#             optimizer.step()\n            \n#             train_loss += loss.item()\n        \n#         train_loss /= len(train_loader)\n        \n#         # Validation\n#         model.eval()\n#         val_loss = 0.0\n        \n#         with torch.no_grad():\n#             for batch in val_loader:\n#                 images = batch['image'].to(device)\n#                 targets = batch['targets']\n                \n#                 outputs = model(images)\n#                 loss = criterion(outputs, targets)\n                \n#                 val_loss += loss.item()\n        \n#         val_loss /= len(val_loader)\n        \n#         print(f'Epoch {epoch+1}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}')\n        \n#         # Save best model\n#         if val_loss < best_val_loss:\n#             best_val_loss = val_loss\n#             torch.save(model.state_dict(), 'best_ecg_model.pth')\n#             print('  ✓ Saved best model')\n        \n#         scheduler.step()\n    \n#     return model\n\n\n# # ============================================================================\n# # PART 6: ENSEMBLE & POST-PROCESSING\n# # ============================================================================\n\n# class ECGEnsemble:\n#     \"\"\"Ensemble multiple models for robust predictions\"\"\"\n    \n#     def __init__(self, models, weights=None):\n#         self.models = models\n#         self.weights = weights if weights else [1.0 / len(models)] * len(models)\n    \n#     def predict(self, image):\n#         \"\"\"Ensemble prediction from multiple models\"\"\"\n#         predictions = []\n        \n#         for model, weight in zip(self.models, self.weights):\n#             model.eval()\n#             with torch.no_grad():\n#                 pred = model(image)\n#                 predictions.append([(p * weight).cpu().numpy() for p in pred])\n        \n#         # Average predictions\n#         ensemble_pred = []\n#         for lead_idx in range(len(predictions[0])):\n#             lead_preds = [pred[lead_idx] for pred in predictions]\n#             ensemble_pred.append(np.mean(lead_preds, axis=0))\n        \n#         return ensemble_pred\n\n\n# def post_process_signal(signal, fs):\n#     \"\"\"Post-process extracted signal\"\"\"\n#     from scipy import signal as scipy_signal\n    \n#     # Remove baseline wander (high-pass filter at 0.5 Hz)\n#     sos = scipy_signal.butter(4, 0.5, btype='high', fs=fs, output='sos')\n#     signal_filtered = scipy_signal.sosfilt(sos, signal)\n    \n#     # Remove high-frequency noise (low-pass filter at 150 Hz)\n#     sos = scipy_signal.butter(4, min(150, fs/2-1), btype='low', fs=fs, output='sos')\n#     signal_filtered = scipy_signal.sosfilt(sos, signal_filtered)\n    \n#     # Remove outliers (clip to 5 standard deviations)\n#     mean, std = np.mean(signal_filtered), np.std(signal_filtered)\n#     signal_filtered = np.clip(signal_filtered, mean - 5*std, mean + 5*std)\n    \n#     return signal_filtered\n\n\n# # ============================================================================\n# # EXAMPLE USAGE & TRAINING PIPELINE\n# # ============================================================================\n\n# def main():\n#     # Configuration\n#     BASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\n#     IMG_SIZE = (512, 512)\n#     BATCH_SIZE = 4\n#     EPOCHS = 50\n#     DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n#     print(f\"Using device: {DEVICE}\")\n    \n#     # Load metadata\n#     train_df = pd.read_csv(BASE_PATH / 'train.csv')\n    \n#     # Split train/val\n#     from sklearn.model_selection import train_test_split\n#     train_df, val_df = train_test_split(train_df, test_size=0.2, random_state=42)\n    \n#     # Data augmentation\n#     train_transform = A.Compose([\n#         A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n#         A.HorizontalFlip(p=0.0),  # Don't flip ECGs!\n#         A.RandomBrightnessContrast(p=0.5),\n#         A.GaussNoise(p=0.3),\n#         A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n#         ToTensorV2()\n#     ])\n    \n#     val_transform = A.Compose([\n#         A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n#         A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n#         ToTensorV2()\n#     ])\n    \n#     # Create datasets\n#     train_dataset = ECGImageDataset(train_df, BASE_PATH, segment_id='0001', \n#                                    transform=train_transform, mode='train')\n#     val_dataset = ECGImageDataset(val_df, BASE_PATH, segment_id='0001',\n#                                  transform=val_transform, mode='train')\n    \n#     train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, \n#                              shuffle=True, num_workers=2)\n#     val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, \n#                            shuffle=False, num_workers=2)\n    \n#     # Initialize model\n#     print(\"\\n\" + \"=\"*80)\n#     print(\"TRAINING UNET SEGMENTATION MODEL\")\n#     print(\"=\"*80)\n    \n#     model = ECGSegmentationUNet(encoder_name='resnet34', encoder_weights='imagenet')\n    \n#     # Train\n#     trained_model = train_model(model, train_loader, val_loader, \n#                                 epochs=EPOCHS, device=DEVICE)\n    \n#     print(\"\\n✓ Training complete!\")\n#     print(\"Best model saved to: best_ecg_model.pth\")\n\n\n# if __name__ == \"__main__\":\n#     main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:57.389760Z","iopub.status.idle":"2025-12-10T18:46:57.390019Z","shell.execute_reply.started":"2025-12-10T18:46:57.389898Z","shell.execute_reply":"2025-12-10T18:46:57.389910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nAdvanced ECG Digitization Techniques\n- SAM2 integration for robust segmentation\n- Handling extreme degradations (mold, stains, damage)\n- Multi-stage processing pipeline\n\"\"\"\n\nimport torch\nimport numpy as np\nimport cv2\nfrom pathlib import Path\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# ============================================================================\n# PART 1: SAM2 INTEGRATION FOR ROBUST SEGMENTATION\n# ============================================================================\n\nclass SAM2ECGSegmenter:\n    \"\"\"\n    Use SAM2 (Segment Anything Model 2) for robust ECG trace segmentation\n    Even works on heavily degraded images\n    \"\"\"\n    \n    def __init__(self, checkpoint_path=None):\n        \"\"\"\n        Initialize SAM2 model\n        \n        Installation:\n        pip install git+https://github.com/facebookresearch/segment-anything-2.git\n        \n        Download checkpoint:\n        wget https://dl.fbaipublicfiles.com/segment_anything_2/sam2_hiera_large.pt\n        \"\"\"\n        try:\n            from sam2.build_sam import build_sam2\n            from sam2.sam2_image_predictor import SAM2ImagePredictor\n            \n            if checkpoint_path is None:\n                checkpoint_path = \"sam2_hiera_large.pt\"\n            \n            # Build SAM2 model\n            sam2_model = build_sam2(\"sam2_hiera_l.yaml\", checkpoint_path)\n            self.predictor = SAM2ImagePredictor(sam2_model)\n            self.available = True\n            print(\"✓ SAM2 loaded successfully\")\n            \n        except ImportError:\n            print(\"⚠ SAM2 not available. Install with:\")\n            print(\"  pip install git+https://github.com/facebookresearch/segment-anything-2.git\")\n            self.predictor = None\n            self.available = False\n    \n    def segment_ecg_traces(self, image, prompt_points=None, prompt_labels=None):\n        \"\"\"\n        Segment ECG traces using SAM2\n        \n        Args:\n            image: RGB image as numpy array\n            prompt_points: Optional (N, 2) array of point prompts\n            prompt_labels: Optional (N,) array of labels (1=foreground, 0=background)\n        \n        Returns:\n            masks: Segmentation masks for ECG traces\n        \"\"\"\n        if not self.available:\n            return None\n        \n        # Set image\n        self.predictor.set_image(image)\n        \n        # Auto-generate prompts if not provided\n        if prompt_points is None:\n            prompt_points, prompt_labels = self._generate_auto_prompts(image)\n        \n        # Predict masks\n        masks, scores, logits = self.predictor.predict(\n            point_coords=prompt_points,\n            point_labels=prompt_labels,\n            multimask_output=True\n        )\n        \n        # Select best mask\n        best_mask_idx = np.argmax(scores)\n        return masks[best_mask_idx]\n    \n    def _generate_auto_prompts(self, image):\n        \"\"\"\n        Automatically generate prompt points for ECG traces\n        Uses color/intensity analysis to find likely trace locations\n        \"\"\"\n        # Convert to grayscale\n        gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        \n        # Find dark regions (ECG traces are typically dark)\n        _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)\n        \n        # Find contours\n        contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n        \n        # Sample points from largest contours\n        prompt_points = []\n        prompt_labels = []\n        \n        # Sort contours by area\n        contours = sorted(contours, key=cv2.contourArea, reverse=True)[:20]\n        \n        for contour in contours:\n            if cv2.contourArea(contour) > 100:\n                # Sample a point from this contour\n                M = cv2.moments(contour)\n                if M[\"m00\"] != 0:\n                    cx = int(M[\"m10\"] / M[\"m00\"])\n                    cy = int(M[\"m01\"] / M[\"m00\"])\n                    prompt_points.append([cx, cy])\n                    prompt_labels.append(1)  # Foreground\n        \n        # Add some background points\n        h, w = image.shape[:2]\n        background_points = [\n            [10, 10], [w-10, 10], [10, h-10], [w-10, h-10]\n        ]\n        for pt in background_points:\n            prompt_points.append(pt)\n            prompt_labels.append(0)  # Background\n        \n        return np.array(prompt_points), np.array(prompt_labels)\n    \n    def segment_all_leads(self, image, num_leads=12):\n        \"\"\"\n        Segment all 12 ECG leads separately\n        Returns dict of lead_name -> mask\n        \"\"\"\n        lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n                     'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n        \n        # Divide image into lead regions (4x3 grid)\n        h, w = image.shape[:2]\n        \n        lead_masks = {}\n        \n        # First 3 rows: 3 leads each\n        for row in range(3):\n            for col in range(3):\n                lead_idx = row * 3 + col\n                if lead_idx < len(lead_names):\n                    y1 = row * h // 4\n                    y2 = (row + 1) * h // 4\n                    x1 = col * w // 3\n                    x2 = (col + 1) * w // 3\n                    \n                    region = image[y1:y2, x1:x2]\n                    mask = self.segment_ecg_traces(region)\n                    \n                    if mask is not None:\n                        lead_masks[lead_names[lead_idx]] = {\n                            'mask': mask,\n                            'bbox': (y1, y2, x1, x2)\n                        }\n        \n        # Last row: Lead II spans full width\n        if len(lead_names) > 9:\n            y1 = 3 * h // 4\n            y2 = h\n            region = image[y1:y2, :]\n            mask = self.segment_ecg_traces(region)\n            \n            if mask is not None:\n                lead_masks['II'] = {\n                    'mask': mask,\n                    'bbox': (y1, y2, 0, w)\n                }\n        \n        return lead_masks\n\n\n# ============================================================================\n# PART 2: ADVANCED PREPROCESSING FOR DEGRADED IMAGES\n# ============================================================================\n\nclass RobustPreprocessor:\n    \"\"\"Handle extreme degradations: mold, stains, damage, etc.\"\"\"\n    \n    def __init__(self):\n        self.techniques = []\n    \n    def remove_mold_stains(self, image):\n        \"\"\"Remove mold and stain artifacts\"\"\"\n        # Convert to LAB color space\n        lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n        l, a, b = cv2.split(lab)\n        \n        # Detect mold (usually green/blue tint)\n        # In LAB: a = green-red, b = blue-yellow\n        mold_mask = ((a < 128) & (b < 128)).astype(np.uint8) * 255\n        \n        # Detect brown stains\n        stain_mask = ((a > 135) & (b > 135)).astype(np.uint8) * 255\n        \n        # Combine masks\n        damage_mask = cv2.bitwise_or(mold_mask, stain_mask)\n        \n        # Dilate mask to cover surrounding areas\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))\n        damage_mask = cv2.dilate(damage_mask, kernel, iterations=2)\n        \n        # Inpaint damaged areas\n        restored = cv2.inpaint(image, damage_mask, 7, cv2.INPAINT_TELEA)\n        \n        return restored, damage_mask\n    \n    def restore_damaged_regions(self, image):\n        \"\"\"Restore physically damaged/torn areas\"\"\"\n        gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        \n        # Detect very bright or very dark anomalous regions\n        _, bright_mask = cv2.threshold(gray, 240, 255, cv2.THRESH_BINARY)\n        _, dark_mask = cv2.threshold(gray, 15, 255, cv2.THRESH_BINARY_INV)\n        \n        damage_mask = cv2.bitwise_or(bright_mask, dark_mask)\n        \n        # Remove small noise\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))\n        damage_mask = cv2.morphologyEx(damage_mask, cv2.MORPH_OPEN, kernel)\n        damage_mask = cv2.morphologyEx(damage_mask, cv2.MORPH_CLOSE, kernel)\n        \n        # Inpaint\n        restored = cv2.inpaint(image, damage_mask, 10, cv2.INPAINT_NS)\n        \n        return restored, damage_mask\n    \n    def enhance_faded_scans(self, image):\n        \"\"\"Enhance contrast in faded scans\"\"\"\n        # Convert to LAB\n        lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n        l, a, b = cv2.split(lab)\n        \n        # Apply CLAHE (Contrast Limited Adaptive Histogram Equalization)\n        clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n        l_enhanced = clahe.apply(l)\n        \n        # Merge back\n        enhanced_lab = cv2.merge([l_enhanced, a, b])\n        enhanced = cv2.cvtColor(enhanced_lab, cv2.COLOR_LAB2RGB)\n        \n        return enhanced\n    \n    def correct_mobile_photo_distortion(self, image):\n        \"\"\"Correct perspective distortion in mobile photos\"\"\"\n        gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        \n        # Detect edges\n        edges = cv2.Canny(gray, 50, 150, apertureSize=3)\n        \n        # Find lines (page borders)\n        lines = cv2.HoughLinesP(edges, 1, np.pi/180, threshold=100,\n                               minLineLength=100, maxLineGap=10)\n        \n        if lines is not None and len(lines) >= 4:\n            # Find corners (intersection of lines)\n            # Simplified: just use detected page boundary\n            # In production, implement full perspective correction\n            pass\n        \n        return image  # Placeholder\n    \n    def denoise_heavily(self, image):\n        \"\"\"Aggressive denoising for very noisy images\"\"\"\n        # Multiple passes of Non-Local Means denoising\n        denoised = cv2.fastNlMeansDenoisingColored(image, None, 10, 10, 7, 21)\n        denoised = cv2.fastNlMeansDenoisingColored(denoised, None, 5, 5, 7, 21)\n        \n        return denoised\n    \n    def full_robust_preprocess(self, image, degradation_type='unknown'):\n        \"\"\"\n        Complete robust preprocessing pipeline\n        \n        Args:\n            image: Input image\n            degradation_type: One of 'mold', 'stained', 'damaged', 'faded', \n                            'mobile', 'screen', or 'unknown'\n        \"\"\"\n        processed = image.copy()\n        masks = {}\n        \n        if degradation_type in ['mold', 'stained', 'unknown']:\n            processed, mold_mask = self.remove_mold_stains(processed)\n            masks['mold'] = mold_mask\n        \n        if degradation_type in ['damaged', 'unknown']:\n            processed, damage_mask = self.restore_damaged_regions(processed)\n            masks['damage'] = damage_mask\n        \n        if degradation_type in ['faded', 'unknown']:\n            processed = self.enhance_faded_scans(processed)\n        \n        if degradation_type in ['mobile', 'screen', 'unknown']:\n            processed = self.correct_mobile_photo_distortion(processed)\n        \n        # Always denoise\n        processed = self.denoise_heavily(processed)\n        \n        return processed, masks\n\n\n# ============================================================================\n# PART 3: MULTI-STAGE PIPELINE\n# ============================================================================\n\nclass MultiStageECGPipeline:\n    \"\"\"\n    Multi-stage pipeline combining classical CV, DL, and SAM2\n    \"\"\"\n    \n    def __init__(self, use_sam2=False):\n        self.robust_preprocessor = RobustPreprocessor()\n        \n        if use_sam2:\n            self.sam2_segmenter = SAM2ECGSegmenter()\n        else:\n            self.sam2_segmenter = None\n    \n    def process_degraded_image(self, image_path, degradation_type='unknown'):\n        \"\"\"\n        Process a degraded ECG image through complete pipeline\n        \n        Returns:\n            - preprocessed image\n            - lead masks (if SAM2 available)\n            - extracted signals\n        \"\"\"\n        # Stage 1: Load image\n        image = cv2.imread(str(image_path))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Stage 2: Robust preprocessing\n        print(f\"Stage 1: Preprocessing ({degradation_type})...\")\n        preprocessed, damage_masks = self.robust_preprocessor.full_robust_preprocess(\n            image, degradation_type\n        )\n        \n        # Stage 3: SAM2 segmentation (if available)\n        lead_masks = None\n        if self.sam2_segmenter and self.sam2_segmenter.available:\n            print(\"Stage 2: SAM2 segmentation...\")\n            lead_masks = self.sam2_segmenter.segment_all_leads(preprocessed)\n        \n        # Stage 4: Signal extraction (placeholder - use previous methods)\n        print(\"Stage 3: Signal extraction...\")\n        signals = self._extract_signals_from_masks(preprocessed, lead_masks)\n        \n        return {\n            'original': image,\n            'preprocessed': preprocessed,\n            'damage_masks': damage_masks,\n            'lead_masks': lead_masks,\n            'signals': signals\n        }\n    \n    def _extract_signals_from_masks(self, image, lead_masks):\n        \"\"\"Extract signals from segmented masks\"\"\"\n        if lead_masks is None:\n            return None\n        \n        signals = {}\n        \n        for lead_name, mask_info in lead_masks.items():\n            mask = mask_info['mask']\n            bbox = mask_info['bbox']\n            \n            # Extract centerline of mask (skeleton)\n            from skimage.morphology import skeletonize\n            skeleton = skeletonize(mask > 0)\n            \n            # Convert to coordinates\n            y_coords, x_coords = np.where(skeleton)\n            \n            if len(x_coords) > 0:\n                # Sort by x coordinate\n                sorted_indices = np.argsort(x_coords)\n                x_sorted = x_coords[sorted_indices]\n                y_sorted = y_coords[sorted_indices]\n                \n                # Remove duplicates in x\n                unique_x, unique_indices = np.unique(x_sorted, return_index=True)\n                unique_y = y_sorted[unique_indices]\n                \n                signals[lead_name] = {\n                    'x': unique_x,\n                    'y': unique_y,\n                    'bbox': bbox\n                }\n        \n        return signals\n    \n    def visualize_results(self, results, save_path='pipeline_results.png'):\n        \"\"\"Visualize all stages of processing\"\"\"\n        fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n        \n        # Original\n        axes[0, 0].imshow(results['original'])\n        axes[0, 0].set_title('Original Image', fontsize=12, fontweight='bold')\n        axes[0, 0].axis('off')\n        \n        # Preprocessed\n        axes[0, 1].imshow(results['preprocessed'])\n        axes[0, 1].set_title('Preprocessed', fontsize=12, fontweight='bold')\n        axes[0, 1].axis('off')\n        \n        # Damage masks\n        if 'mold' in results['damage_masks']:\n            axes[0, 2].imshow(results['damage_masks']['mold'], cmap='hot')\n            axes[0, 2].set_title('Mold/Stain Detection', fontsize=12, fontweight='bold')\n            axes[0, 2].axis('off')\n        \n        # Lead masks (if available)\n        if results['lead_masks']:\n            # Show first few lead masks\n            lead_names = list(results['lead_masks'].keys())[:3]\n            for idx, lead_name in enumerate(lead_names):\n                if idx < 3:\n                    mask = results['lead_masks'][lead_name]['mask']\n                    axes[1, idx].imshow(mask, cmap='viridis')\n                    axes[1, idx].set_title(f'Lead {lead_name} Mask', \n                                         fontsize=12, fontweight='bold')\n                    axes[1, idx].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f\"✓ Saved visualization to {save_path}\")\n        \n        return fig\n\n\n# ============================================================================\n# PART 4: USAGE EXAMPLES\n# ============================================================================\n\ndef example_process_all_degradation_types():\n    \"\"\"Example: Process all 9 degradation types for a single ECG\"\"\"\n    \n    BASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\n    sample_id = '7663343'\n    \n    # Map segment IDs to degradation types\n    degradation_map = {\n        '0001': 'clean',\n        '0003': 'faded',\n        '0004': 'faded',\n        '0005': 'mobile',\n        '0006': 'screen',\n        '0009': 'stained',\n        '0010': 'damaged',\n        '0011': 'mold',\n        '0012': 'mold'\n    }\n    \n    # Initialize pipeline\n    pipeline = MultiStageECGPipeline(use_sam2=False)  # Set True if SAM2 available\n    \n    results_all = {}\n    \n    for segment_id, deg_type in degradation_map.items():\n        img_path = BASE_PATH / 'train' / sample_id / f'{sample_id}-{segment_id}.png'\n        \n        if img_path.exists():\n            print(f\"\\n{'='*80}\")\n            print(f\"Processing {segment_id} ({deg_type})\")\n            print(f\"{'='*80}\")\n            \n            results = pipeline.process_degraded_image(img_path, deg_type)\n            results_all[segment_id] = results\n            \n            # Visualize\n            pipeline.visualize_results(results, \n                                      save_path=f'results_{segment_id}_{deg_type}.png')\n    \n    return results_all\n\n\ndef example_sam2_only():\n    \"\"\"Example: Use only SAM2 for segmentation\"\"\"\n    \n    # Initialize SAM2\n    segmenter = SAM2ECGSegmenter()\n    \n    if not segmenter.available:\n        print(\"SAM2 not available. Please install:\")\n        print(\"pip install git+https://github.com/facebookresearch/segment-anything-2.git\")\n        return\n    \n    # Load image\n    img_path = '/kaggle/input/physionet-ecg-image-digitization/train/7663343/7663343-0001.png'\n    image = cv2.imread(img_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # Segment all leads\n    lead_masks = segmenter.segment_all_leads(image)\n    \n    # Visualize\n    fig, axes = plt.subplots(3, 4, figsize=(20, 15))\n    axes = axes.flatten()\n    \n    for idx, (lead_name, mask_info) in enumerate(lead_masks.items()):\n        if idx < 12:\n            mask = mask_info['mask']\n            axes[idx].imshow(mask, cmap='viridis')\n            axes[idx].set_title(f'Lead {lead_name}', fontweight='bold')\n            axes[idx].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig('sam2_segmentation_results.png', dpi=150)\n    print(\"✓ Saved SAM2 results\")\n\n\nif __name__ == \"__main__\":\n    print(\"=\"*80)\n    print(\"ADVANCED ECG PROCESSING PIPELINE\")\n    print(\"=\"*80)\n    \n    print(\"\\nAvailable examples:\")\n    print(\"1. example_process_all_degradation_types() - Process all 9 variants\")\n    print(\"2. example_sam2_only() - Use SAM2 for segmentation\")\n    \n    # Run example\n    example_process_all_degradation_types()\n    example_sam2_only()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:57.391367Z","iopub.status.idle":"2025-12-10T18:46:57.391690Z","shell.execute_reply.started":"2025-12-10T18:46:57.391525Z","shell.execute_reply":"2025-12-10T18:46:57.391537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nCOMPLETE ECG DIGITIZATION SOLUTION\nEnd-to-end pipeline integrating all techniques\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\n\nclass Config:\n    \"\"\"Centralized configuration\"\"\"\n    \n    # Paths\n    BASE_PATH = Path('/kaggle/input/physionet-ecg-image-digitization')\n    OUTPUT_PATH = Path('/kaggle/working')\n    \n    # Model settings\n    IMG_SIZE = (512, 512)\n    BATCH_SIZE = 8\n    NUM_WORKERS = 2\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    # Training\n    EPOCHS = 50\n    LEARNING_RATE = 1e-4\n    WEIGHT_DECAY = 1e-5\n    \n    # ECG specifics\n    LEAD_NAMES = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n                  'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    LEAD_DURATIONS = {'II': 10.0}  # Lead II is 10s, others default to 2.5s\n    \n    # Degradation types for training augmentation\n    SEGMENT_IDS = ['0001', '0003', '0004', '0005', '0006', \n                   '0009', '0010', '0011', '0012']\n    \n    # Model ensemble\n    USE_ENSEMBLE = True\n    ENSEMBLE_MODELS = ['unet_resnet34', 'unet_resnet50', 'efficientunet']\n\n\n# ============================================================================\n# COMPLETE PIPELINE CLASS\n# ============================================================================\n\nclass ECGCompletePipeline:\n    \"\"\"\n    Complete end-to-end pipeline\n    Integrates: preprocessing → segmentation → signal extraction → post-processing\n    \"\"\"\n    \n    def __init__(self, config=None, use_sam2=False):\n        self.config = config or Config()\n        self.use_sam2 = use_sam2\n        \n        # Initialize components\n        self._setup_preprocessing()\n        self._setup_models()\n        \n        print(f\"✓ Pipeline initialized on {self.config.DEVICE}\")\n    \n    def _setup_preprocessing(self):\n        \"\"\"Initialize preprocessing modules\"\"\"\n        from scipy.ndimage import gaussian_filter1d\n        self.gaussian_filter = gaussian_filter1d\n        \n    def _setup_models(self):\n        \"\"\"Load trained models\"\"\"\n        self.models = {}\n        \n        # Load ensemble of models if available\n        if self.config.USE_ENSEMBLE:\n            for model_name in self.config.ENSEMBLE_MODELS:\n                model_path = self.config.OUTPUT_PATH / f'{model_name}.pth'\n                if model_path.exists():\n                    try:\n                        model = self._load_model(model_name, model_path)\n                        self.models[model_name] = model\n                        print(f\"✓ Loaded {model_name}\")\n                    except Exception as e:\n                        print(f\"⚠ Failed to load {model_name}: {e}\")\n        \n        if not self.models:\n            print(\"⚠ No trained models found. Will use classical methods only.\")\n    \n    def _load_model(self, model_name, model_path):\n        \"\"\"Load a trained model\"\"\"\n        # Implementation depends on model architecture\n        # Placeholder for now\n        return None\n    \n    # ========================================================================\n    # PREPROCESSING\n    # ========================================================================\n    \n    def preprocess_image(self, image, degradation_level='medium'):\n        \"\"\"\n        Comprehensive preprocessing\n        \n        Args:\n            image: Input RGB image (numpy array)\n            degradation_level: 'low', 'medium', 'high', or 'extreme'\n        \n        Returns:\n            Preprocessed image\n        \"\"\"\n        processed = image.copy()\n        \n        # Step 1: Denoise\n        if degradation_level in ['high', 'extreme']:\n            processed = cv2.fastNlMeansDenoisingColored(processed, None, 10, 10, 7, 21)\n        else:\n            processed = cv2.bilateralFilter(processed, 9, 75, 75)\n        \n        # Step 2: Contrast enhancement\n        lab = cv2.cvtColor(processed, cv2.COLOR_RGB2LAB)\n        l, a, b = cv2.split(lab)\n        \n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        l = clahe.apply(l)\n        \n        processed = cv2.merge([l, a, b])\n        processed = cv2.cvtColor(processed, cv2.COLOR_LAB2RGB)\n        \n        # Step 3: Remove grid lines\n        processed = self._remove_grid_adaptive(processed)\n        \n        return processed\n    \n    def _remove_grid_adaptive(self, image):\n        \"\"\"Adaptive grid removal\"\"\"\n        # Detect grid color\n        hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n        \n        # Pink/red grid detection\n        lower_pink = np.array([140, 30, 100])\n        upper_pink = np.array([180, 255, 255])\n        pink_mask = cv2.inRange(hsv, lower_pink, upper_pink)\n        \n        # Light red/orange grid\n        lower_orange = np.array([0, 30, 150])\n        upper_orange = np.array([20, 255, 255])\n        orange_mask = cv2.inRange(hsv, lower_orange, upper_orange)\n        \n        grid_mask = cv2.bitwise_or(pink_mask, orange_mask)\n        \n        # Morphological cleanup\n        kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (3, 3))\n        grid_mask = cv2.morphologyEx(grid_mask, cv2.MORPH_CLOSE, kernel)\n        \n        # Inpaint\n        result = cv2.inpaint(image, grid_mask, 5, cv2.INPAINT_TELEA)\n        \n        return result\n    \n    # ========================================================================\n    # SIGNAL EXTRACTION\n    # ========================================================================\n    \n    def extract_signals(self, image, fs, use_dl=True):\n        \"\"\"\n        Extract signals from preprocessed image\n        \n        Args:\n            image: Preprocessed image\n            fs: Sampling frequency\n            use_dl: Use deep learning if available\n        \n        Returns:\n            Dictionary of lead_name -> signal array\n        \"\"\"\n        if use_dl and self.models:\n            return self._extract_signals_dl(image, fs)\n        else:\n            return self._extract_signals_classical(image, fs)\n    \n    def _extract_signals_dl(self, image, fs):\n        \"\"\"Deep learning signal extraction\"\"\"\n        signals = {}\n        \n        # Resize and normalize\n        img_tensor = self._prepare_image_tensor(image)\n        \n        # Ensemble prediction\n        predictions = []\n        for model_name, model in self.models.items():\n            model.eval()\n            with torch.no_grad():\n                pred = model(img_tensor)\n                predictions.append(pred)\n        \n        # Average predictions\n        if predictions:\n            avg_pred = torch.mean(torch.stack(predictions), dim=0)\n            \n            # Convert to signals for each lead\n            for lead_idx, lead_name in enumerate(self.config.LEAD_NAMES):\n                duration = self.config.LEAD_DURATIONS.get(lead_name, 2.5)\n                target_len = int(fs * duration)\n                \n                # Extract and resample\n                signal = avg_pred[0, lead_idx, :target_len].cpu().numpy()\n                signals[lead_name] = signal\n        \n        return signals\n    \n    def _extract_signals_classical(self, image, fs):\n        \"\"\"Classical CV signal extraction\"\"\"\n        signals = {}\n        \n        # Convert to grayscale\n        gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        \n        # Binarize\n        _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)\n        \n        # Detect lead regions\n        regions = self._detect_lead_regions(binary)\n        \n        # Extract each lead\n        for lead_name, region in regions.items():\n            y1, y2, x1, x2 = region\n            lead_binary = binary[y1:y2, x1:x2]\n            \n            # Extract trace\n            trace = self._extract_trace_from_binary(lead_binary)\n            \n            # Convert to signal\n            duration = self.config.LEAD_DURATIONS.get(lead_name, 2.5)\n            signal = self._trace_to_signal(trace, fs, duration)\n            \n            signals[lead_name] = signal\n        \n        return signals\n    \n    def _detect_lead_regions(self, binary_image):\n        \"\"\"Detect regions for each ECG lead\"\"\"\n        h, w = binary_image.shape\n        \n        # Standard 4x3 grid layout\n        regions = {}\n        lead_idx = 0\n        \n        # First 3 rows: 3 leads each\n        for row in range(3):\n            for col in range(3):\n                if lead_idx < 9:\n                    y1 = int(row * h / 4)\n                    y2 = int((row + 1) * h / 4)\n                    x1 = int(col * w / 3)\n                    x2 = int((col + 1) * w / 3)\n                    \n                    regions[self.config.LEAD_NAMES[lead_idx]] = (y1, y2, x1, x2)\n                    lead_idx += 1\n        \n        # Last row: Lead II (full width) and two more leads\n        y1 = int(3 * h / 4)\n        y2 = h\n        \n        # Lead II spans full width or partial - adjust based on layout\n        regions['II'] = (y1, y2, 0, w)\n        \n        return regions\n    \n    def _extract_trace_from_binary(self, binary_region):\n        \"\"\"Extract centerline trace from binary region\"\"\"\n        from skimage.morphology import skeletonize\n        \n        # Skeletonize\n        skeleton = skeletonize(binary_region > 0)\n        \n        # Extract coordinates\n        coords = []\n        width = skeleton.shape[1]\n        \n        for x in range(width):\n            col = skeleton[:, x]\n            y_vals = np.where(col)[0]\n            \n            if len(y_vals) > 0:\n                y = int(np.median(y_vals))\n                coords.append((x, y))\n        \n        return np.array(coords) if coords else np.array([])\n    \n    def _trace_to_signal(self, trace, fs, duration):\n        \"\"\"Convert pixel trace to signal values\"\"\"\n        from scipy import interpolate\n        \n        if len(trace) == 0:\n            return np.zeros(int(fs * duration))\n        \n        x_coords = trace[:, 0]\n        y_coords = trace[:, 1]\n        \n        # Normalize\n        x_norm = np.linspace(0, duration, len(x_coords))\n        y_mean = np.mean(y_coords)\n        y_std = np.std(y_coords) + 1e-8\n        \n        # Flip y (image coordinates) and normalize\n        y_signal = -(y_coords - y_mean) / y_std\n        \n        # Interpolate to target length\n        target_len = int(fs * duration)\n        target_time = np.linspace(0, duration, target_len)\n        \n        f = interpolate.interp1d(x_norm, y_signal, kind='cubic', \n                                fill_value='extrapolate')\n        signal = f(target_time)\n        \n        # Post-process\n        signal = self._post_process_signal(signal, fs)\n        \n        return signal\n    \n    def _prepare_image_tensor(self, image):\n        \"\"\"Convert image to model input tensor\"\"\"\n        # Resize\n        img = cv2.resize(image, self.config.IMG_SIZE)\n        \n        # Normalize\n        img = img.astype(np.float32) / 255.0\n        mean = np.array([0.485, 0.456, 0.406])\n        std = np.array([0.229, 0.224, 0.225])\n        img = (img - mean) / std\n        \n        # To tensor\n        img = torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0)\n        img = img.to(self.config.DEVICE)\n        \n        return img\n    \n    # ========================================================================\n    # POST-PROCESSING\n    # ========================================================================\n    \n    def _post_process_signal(self, signal, fs):\n        \"\"\"Post-process extracted signal\"\"\"\n        from scipy import signal as scipy_signal\n        \n        # Bandpass filter (0.5-150 Hz)\n        if fs > 2:\n            nyq = fs / 2\n            low = 0.5 / nyq\n            high = min(150, nyq - 1) / nyq\n            \n            b, a = scipy_signal.butter(4, [low, high], btype='band')\n            signal_filtered = scipy_signal.filtfilt(b, a, signal)\n        else:\n            signal_filtered = signal\n        \n        # Remove outliers\n        mean, std = np.mean(signal_filtered), np.std(signal_filtered)\n        signal_filtered = np.clip(signal_filtered, mean - 5*std, mean + 5*std)\n        \n        # Smooth slightly\n        signal_filtered = self.gaussian_filter(signal_filtered, sigma=1)\n        \n        return signal_filtered\n    \n    # ========================================================================\n    # INFERENCE\n    # ========================================================================\n    \n    def process_single_image(self, img_path, fs, degradation_level='medium'):\n        \"\"\"\n        Process a single ECG image\n        \n        Returns:\n            Dictionary of lead_name -> signal array\n        \"\"\"\n        # Load image\n        image = cv2.imread(str(img_path))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Preprocess\n        preprocessed = self.preprocess_image(image, degradation_level)\n        \n        # Extract signals\n        signals = self.extract_signals(preprocessed, fs, use_dl=bool(self.models))\n        \n        return signals\n    \n    def generate_submission(self, test_csv_path=None, test_dir=None, \n                          output_path=None):\n        \"\"\"\n        Generate Kaggle submission file\n        \"\"\"\n        if test_csv_path is None:\n            test_csv_path = self.config.BASE_PATH / 'test.csv'\n        if test_dir is None:\n            test_dir = self.config.BASE_PATH / 'test'\n        if output_path is None:\n            output_path = self.config.OUTPUT_PATH / 'submission.parquet'\n        \n        # Load test metadata\n        test_df = pd.read_csv(test_csv_path)\n        \n        submission_rows = []\n        \n        # Process each unique test image\n        for img_id in tqdm(test_df['id'].unique(), desc=\"Processing test images\"):\n            img_path = Path(test_dir) / f\"{img_id}.png\"\n            \n            if not img_path.exists():\n                print(f\"⚠ Image {img_id} not found\")\n                continue\n            \n            # Get metadata\n            img_rows = test_df[test_df['id'] == img_id]\n            fs = img_rows['fs'].iloc[0]\n            \n            try:\n                # Process image\n                signals = self.process_single_image(img_path, fs, \n                                                   degradation_level='medium')\n                \n                # Create submission rows\n                for _, row in img_rows.iterrows():\n                    lead = row['lead']\n                    num_rows = row['number_of_rows']\n                    \n                    if lead in signals:\n                        signal = signals[lead]\n                        \n                        # Truncate or pad to expected length\n                        if len(signal) > num_rows:\n                            signal = signal[:num_rows]\n                        elif len(signal) < num_rows:\n                            signal = np.pad(signal, (0, num_rows - len(signal)), \n                                          mode='edge')\n                        \n                        # Create rows\n                        for row_id in range(num_rows):\n                            submission_rows.append({\n                                'id': f\"{img_id}_{row_id}_{lead}\",\n                                'value': float(signal[row_id])\n                            })\n                    else:\n                        # Fill with zeros if lead not found\n                        for row_id in range(num_rows):\n                            submission_rows.append({\n                                'id': f\"{img_id}_{row_id}_{lead}\",\n                                'value': 0.0\n                            })\n            \n            except Exception as e:\n                print(f\"⚠ Error processing {img_id}: {e}\")\n                # Fill with zeros on error\n                for _, row in img_rows.iterrows():\n                    lead = row['lead']\n                    num_rows = row['number_of_rows']\n                    for row_id in range(num_rows):\n                        submission_rows.append({\n                            'id': f\"{img_id}_{row_id}_{lead}\",\n                            'value': 0.0\n                        })\n        \n        # Create submission DataFrame\n        submission_df = pd.DataFrame(submission_rows)\n        submission_df.to_parquet(output_path, index=False)\n        \n        print(f\"\\n✓ Submission saved to {output_path}\")\n        print(f\"  Total rows: {len(submission_df)}\")\n        print(f\"  Unique IDs: {submission_df['id'].nunique()}\")\n        \n        return submission_df\n\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\ndef main():\n    \"\"\"Main execution function\"\"\"\n    \n    print(\"=\"*80)\n    print(\"ECG IMAGE DIGITIZATION - COMPLETE SOLUTION\")\n    print(\"=\"*80)\n    \n    # Initialize pipeline\n    config = Config()\n    pipeline = ECGCompletePipeline(config=config, use_sam2=False)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"TESTING ON SAMPLE\")\n    print(\"=\"*80)\n    \n    # Test on one training sample\n    sample_id = '7663343'\n    img_path = config.BASE_PATH / 'train' / sample_id / f'{sample_id}-0001.png'\n    \n    if img_path.exists():\n        signals = pipeline.process_single_image(img_path, fs=500)\n        \n        print(f\"\\n✓ Extracted {len(signals)} leads:\")\n        for lead, signal in signals.items():\n            print(f\"  {lead}: {len(signal)} samples, \"\n                  f\"range [{signal.min():.3f}, {signal.max():.3f}]\")\n    \n    # Generate submission\n    print(\"\\n\" + \"=\"*80)\n    print(\"GENERATING SUBMISSION\")\n    print(\"=\"*80)\n    \n    submission = pipeline.generate_submission()\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"COMPLETE!\")\n    print(\"=\"*80)\n    print(\"\\nNext steps:\")\n    print(\"1. Train deep learning models on full dataset\")\n    print(\"2. Fine-tune on each degradation type\")\n    print(\"3. Implement model ensemble\")\n    print(\"4. Optimize hyperparameters\")\n    print(\"5. Add SAM2 integration for robust segmentation\")\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T18:46:57.393005Z","iopub.status.idle":"2025-12-10T18:46:57.393405Z","shell.execute_reply.started":"2025-12-10T18:46:57.393167Z","shell.execute_reply":"2025-12-10T18:46:57.393186Z"}},"outputs":[],"execution_count":null}]}