{"cells": [{"cell_type": "markdown", "metadata": {}, "source": "# PhysioNet ECG Image Digitization\nClassical CV approach using grid removal and trace extraction"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "import cv2\nimport numpy as np\nfrom scipy import interpolate\nfrom scipy.ndimage import gaussian_filter1d\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "class ECGDigitizer:\n    \"\"\"Extract ECG signals from printout images using classical CV.\"\"\"\n\n    LEAD_ORDER = [\n        ['I', 'aVR', 'V1', 'V4'],\n        ['II', 'aVL', 'V2', 'V5'],\n        ['III', 'aVF', 'V3', 'V6'],\n    ]\n\n    def remove_grid(self, img):\n        \"\"\"Remove red/pink grid from ECG image.\"\"\"\n        if len(img.shape) == 2:\n            return img\n\n        hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV)\n        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n\n        lower_red1 = np.array([0, 20, 50])\n        upper_red1 = np.array([20, 255, 255])\n        lower_red2 = np.array([160, 20, 50])\n        upper_red2 = np.array([180, 255, 255])\n\n        mask1 = cv2.inRange(hsv, lower_red1, upper_red1)\n        mask2 = cv2.inRange(hsv, lower_red2, upper_red2)\n        grid_mask = mask1 | mask2\n\n        kernel = np.ones((3, 3), np.uint8)\n        grid_mask = cv2.dilate(grid_mask, kernel, iterations=1)\n        gray[grid_mask > 0] = 255\n\n        return gray\n\n    def get_trace_mask(self, gray_img):\n        \"\"\"Get binary mask of ECG traces.\"\"\"\n        _, trace_mask = cv2.threshold(gray_img, 80, 255, cv2.THRESH_BINARY_INV)\n        kernel = np.ones((2, 2), np.uint8)\n        trace_mask = cv2.morphologyEx(trace_mask, cv2.MORPH_OPEN, kernel)\n        return trace_mask\n\n    def detect_regions(self, img_shape):\n        \"\"\"Detect ECG regions based on standard 3x4+rhythm layout.\"\"\"\n        h, w = img_shape[:2]\n\n        trace_start = int(h * 0.32)\n        row_height = int(h * 0.18)\n        margin_x = int(w * 0.03)\n        col_width = (w - 2 * margin_x) // 4\n\n        regions = {}\n\n        for row_idx, leads in enumerate(self.LEAD_ORDER):\n            y_start = trace_start + row_idx * row_height\n            y_end = y_start + row_height\n\n            for col_idx, lead in enumerate(leads):\n                x_start = margin_x + col_idx * col_width\n                x_end = x_start + col_width\n                regions[lead] = (y_start, y_end, x_start, x_end)\n\n        rhythm_y_start = trace_start + 3 * row_height\n        rhythm_y_end = min(rhythm_y_start + row_height, int(h * 0.97))\n        regions['rhythm'] = (rhythm_y_start, rhythm_y_end, margin_x, w - margin_x)\n\n        return regions\n\n    def extract_signal(self, trace_region, num_samples):\n        \"\"\"Extract time series signal from a trace region.\"\"\"\n        h, w = trace_region.shape\n\n        crop_top = int(h * 0.08)\n        crop_bottom = int(h * 0.88)\n        cropped = trace_region[crop_top:crop_bottom, :]\n        h_crop = cropped.shape[0]\n\n        signal_y = np.full(w, np.nan)\n\n        for x in range(w):\n            col = cropped[:, x]\n            white_pixels = np.where(col > 0)[0]\n            if len(white_pixels) > 0:\n                signal_y[x] = np.mean(white_pixels)\n\n        valid = ~np.isnan(signal_y)\n        if np.sum(valid) < w * 0.1:\n            return np.zeros(num_samples)\n\n        x_valid = np.where(valid)[0]\n        y_valid = signal_y[valid]\n\n        try:\n            interp_func = interpolate.interp1d(\n                x_valid, y_valid, kind='linear',\n                bounds_error=False, fill_value=(y_valid[0], y_valid[-1])\n            )\n            signal_y = interp_func(np.arange(w))\n        except:\n            return np.zeros(num_samples)\n\n        signal_y = gaussian_filter1d(signal_y, sigma=3)\n\n        mid_section = signal_y[w//4:3*w//4]\n        baseline = np.median(mid_section)\n\n        signal_centered = baseline - signal_y\n        scale = 4.0 / h_crop\n        signal_mv = signal_centered * scale\n\n        if len(signal_mv) != num_samples:\n            x_orig = np.linspace(0, 1, len(signal_mv))\n            x_new = np.linspace(0, 1, num_samples)\n            interp_func = interpolate.interp1d(x_orig, signal_mv, kind='linear')\n            signal_mv = interp_func(x_new)\n\n        return np.clip(signal_mv, -5, 5)\n\n    def process_image(self, img_path, lead_samples, rhythm_samples):\n        \"\"\"Process ECG image and extract all leads.\"\"\"\n        img = cv2.imread(str(img_path))\n        if img is None:\n            return None\n\n        gray = self.remove_grid(img)\n        trace_mask = self.get_trace_mask(gray)\n        regions = self.detect_regions(img.shape)\n\n        signals = {}\n\n        for row_leads in self.LEAD_ORDER:\n            for lead in row_leads:\n                y1, y2, x1, x2 = regions[lead]\n                lead_region = trace_mask[y1:y2, x1:x2]\n                signals[lead] = self.extract_signal(lead_region, lead_samples)\n\n        y1, y2, x1, x2 = regions['rhythm']\n        rhythm_region = trace_mask[y1:y2, x1:x2]\n        signals['II_rhythm'] = self.extract_signal(rhythm_region, rhythm_samples)\n\n        return signals"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Paths\nINPUT_DIR = Path('/kaggle/input/physionet-ecg-image-digitization')\nTEST_CSV = INPUT_DIR / 'test.csv'\nTEST_IMG_DIR = INPUT_DIR / 'test'\n\nprint(f'Test CSV exists: {TEST_CSV.exists()}')\nprint(f'Test images dir exists: {TEST_IMG_DIR.exists()}')"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Load test data\ntest_df = pd.read_csv(TEST_CSV)\nprint(f'Test samples: {len(test_df)}')\nprint(f'Unique ECGs: {test_df[\"id\"].nunique()}')\ntest_df.head()"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Process all ECGs\ndigitizer = ECGDigitizer()\nresults = []\n\necg_ids = test_df['id'].unique()\n\nfor ecg_id in tqdm(ecg_ids, desc='Processing ECGs'):\n    ecg_rows = test_df[test_df['id'] == ecg_id]\n    img_path = TEST_IMG_DIR / f'{ecg_id}.png'\n\n    if not img_path.exists():\n        print(f'Warning: {img_path} not found')\n        for _, row in ecg_rows.iterrows():\n            for i in range(row['number_of_rows']):\n                results.append({'id': f\"{ecg_id}_{i}_{row['lead']}\", 'value': 0})\n        continue\n\n    try:\n        lead_samples = ecg_rows[ecg_rows['lead'] == 'I']['number_of_rows'].iloc[0]\n        rhythm_samples = ecg_rows[ecg_rows['lead'] == 'II']['number_of_rows'].iloc[0]\n\n        signals = digitizer.process_image(img_path, lead_samples, rhythm_samples)\n\n        if signals is None:\n            raise ValueError('Failed to process image')\n\n        for _, row in ecg_rows.iterrows():\n            lead = row['lead']\n            num_rows = row['number_of_rows']\n\n            if lead == 'II' and num_rows > lead_samples:\n                signal = signals.get('II_rhythm', np.zeros(num_rows))\n            else:\n                signal = signals.get(lead, np.zeros(num_rows))\n\n            if len(signal) != num_rows:\n                if len(signal) > num_rows:\n                    signal = signal[:num_rows]\n                else:\n                    signal = np.pad(signal, (0, num_rows - len(signal)), mode='edge')\n\n            for i, val in enumerate(signal):\n                results.append({\n                    'id': f\"{ecg_id}_{i}_{lead}\",\n                    'value': int(round(val * 1000))\n                })\n\n    except Exception as e:\n        print(f'Error processing {ecg_id}: {e}')\n        for _, row in ecg_rows.iterrows():\n            for i in range(row['number_of_rows']):\n                results.append({'id': f\"{ecg_id}_{i}_{row['lead']}\", 'value': 0})\n\nprint(f'Total results: {len(results)}')"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Create submission\nsubmission_df = pd.DataFrame(results)\nsubmission_df.to_csv('/kaggle/working/submission.csv', index=False)\n\nprint(f'Submission shape: {submission_df.shape}')\nprint(f'\\nValue statistics:')\nprint(submission_df['value'].describe())\nprint(f'\\nPreview:')\nsubmission_df.head(20)"}], "metadata": {"kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"}, "language_info": {"name": "python", "version": "3.10.0"}}, "nbformat": 4, "nbformat_minor": 4}