{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14068806,"sourceType":"datasetVersion","datasetId":8955071},{"sourceId":14077258,"sourceType":"datasetVersion","datasetId":8961306},{"sourceId":678236,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":514384,"modelId":529028},{"sourceId":678435,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":514544,"modelId":529191}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2 as cv\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\nDATA_PATH = \"/kaggle/input/physionet-ecg-image-digitization\"\ntrain_df = pd.read_csv(DATA_PATH  + '/train.csv')\ntest_df = pd.read_csv(DATA_PATH + '/test.csv')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:56.499366Z","iopub.execute_input":"2025-12-10T05:52:56.500086Z","iopub.status.idle":"2025-12-10T05:52:57.136641Z","shell.execute_reply.started":"2025-12-10T05:52:56.500064Z","shell.execute_reply":"2025-12-10T05:52:57.135946Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Function Definitions","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\ndef extract_patches(strip_img, patch_size=64, stride=32):\n    \"\"\"\n    Extract non-overlapping or optionally strided patches from an image.\n\n    Parameters\n    ----------\n    strip_img : np.ndarray\n        2D (grayscale) or 3D (color) image array.\n    patch_size : int\n        Size of the square patches (default: 64).\n    stride : int or None\n        Step between patches. If None, defaults to patch_size (non-overlapping).\n\n    Returns\n    -------\n    patches : list of np.ndarray\n        List of extracted image patches.\n    coords : list of tuples\n        List of (y, x) coordinates for the top-left corner of each patch.\n    \"\"\"\n    if stride is None:\n        stride = patch_size\n\n    H, W = strip_img.shape[:2]\n    patches = []\n    coords = []\n\n    for y in range(0, H - patch_size + 1, stride):\n        for x in range(0, W - patch_size + 1, stride):\n            patch = strip_img[y:y+patch_size, x:x+patch_size]\n            patches.append(patch)\n            coords.append((y, x))\n\n    return patches, coords\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.137872Z","iopub.execute_input":"2025-12-10T05:52:57.138092Z","iopub.status.idle":"2025-12-10T05:52:57.143502Z","shell.execute_reply.started":"2025-12-10T05:52:57.138075Z","shell.execute_reply":"2025-12-10T05:52:57.142610Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n\ndef crop_and_deskew_ecg(image):\n    \"\"\"\n    Crop the ECG to the bounding rectangle of the detected grid, then deskew that crop.\n    Returns the deskewed crop, original cropped image, and the rotation angle.\n    \"\"\"\n    # --- 1. Find grid and crop ---\n    grid_contour = find_grid_robust(image)\n    x, y, w, h = cv.boundingRect(grid_contour)\n    cropped_img = image[y:y+h, x:x+w]\n\n    # --- 2. Compute rotation angle for deskew ---\n    rect = cv.minAreaRect(grid_contour)\n    angle = rect[-1]\n    width, height = rect[1]\n    if width < height:\n        angle += 90\n\n    # --- 3. Rotate cropped image ---\n    if abs(angle) < 0.5:\n        (ch, cw) = cropped_img.shape[:2]\n        center = (cw // 2, ch // 2)\n        M = cv.getRotationMatrix2D(center, angle, 1.0)\n        rotated = cv.warpAffine(cropped_img, M, (cw, ch),\n                                flags=cv.INTER_LINEAR,\n                                borderMode=cv.BORDER_REPLICATE)\n    else:\n        rotated = cropped_img.copy()\n    rotated = rotated[rotated.shape[0] // 3:, :]\n    return rotated, cropped_img, angle\n\n\ndef visualize_top_waveform_patches(strip_img, mask, patches, labels, n_patches=6):\n    \"\"\"\n    Visualize the topmost waveform strip with mask overlay and sample patches.\n\n    Parameters:\n    - strip_img: cropped waveform strip (grayscale)\n    - mask: corresponding binary waveform mask\n    - patches: extracted patches from the strip\n    - labels: 0/1 labels corresponding to patches\n    - n_patches: number of random patches to display\n    \"\"\"\n\n    # --- Show strip with mask overlay ---\n    plt.figure(figsize=(12,4))\n\n    plt.subplot(1,2,1)\n    plt.imshow(strip_img, cmap='gray')\n    plt.title(\"Cropped Top Waveform\")\n    plt.axis('off')\n\n    plt.subplot(1,2,2)\n    plt.imshow(strip_img, cmap='gray')\n    plt.imshow(mask, cmap='Reds', alpha=0.5)  # overlay mask in red\n    plt.title(\"Waveform Mask Overlay\")\n    plt.axis('off')\n\n    plt.show()\n\n    # --- Display a few example patches ---\n    if patches is not None and labels is not None:\n        idxs = np.random.choice(len(patches), min(n_patches, len(patches)), replace=False)\n        plt.figure(figsize=(12,2))\n        for i, idx in enumerate(idxs):\n            plt.subplot(1, n_patches, i+1)\n            plt.imshow(patches[idx], cmap='gray')\n            plt.title(\"Waveform\" if labels[idx]==1 else \"Background\")\n            plt.axis('off')\n        plt.show()\n\n\ndef make_waveform_mask(img):\n    \"\"\"\n    Generates a pixel-accurate mask for ONE ECG lead waveform.\n    \"\"\"\n    mask = (img > 200)\n    return mask\n\n\ndef find_grid_robust(img):\n    \"\"\"\n    Robust grid detection with dynamic kernel\n    \"\"\"\n    # Step 1: Convert to grayscale.\n    gray = cv.cvtColor(img, cv.COLOR_BGR2GRAY)\n    # Step 2: Gaussian blur (reduce noise).\n    blur_kernel = 5  # Try: 3, 5, 7, 9 (must be odd).\n    blurred = cv.GaussianBlur(gray, (blur_kernel, blur_kernel), 0)\n    \n    # Step 3: Closing on grayscale (fills dark gaps/valleys).\n    kernel_size = 21\n    \n    kernel = cv.getStructuringElement(cv.MORPH_ELLIPSE, (kernel_size, kernel_size))\n    closed = cv.morphologyEx(blurred, cv.MORPH_CLOSE, kernel)\n    \n    # Step 4: Normalize to handle lighting variations.\n    div = np.float32(blurred) / (closed + 1e-8)\n    normalized = np.uint8(cv.normalize(div, div, 0, 255, cv.NORM_MINMAX))\n    \n    # Step 5: Adaptive thresholding (local).\n    block_size = 51  # Try: 11, 21, 31, 51, 101 (must be odd, larger = smoother).\n    C = 10  # Try: 0-10 (constant subtracted from mean).\n    \n    thresh = cv.adaptiveThreshold(\n        normalized, \n        255, \n        cv.ADAPTIVE_THRESH_GAUSSIAN_C,  # or ADAPTIVE_THRESH_MEAN_C\n        cv.THRESH_BINARY_INV, \n        block_size, \n        C\n    )\n\n    contours, hierarchy = cv.findContours(thresh, cv.RETR_TREE, cv.CHAIN_APPROX_SIMPLE)\n\n    largest_contour = max(contours, key=cv.contourArea)\n    largest_area = cv.contourArea(largest_contour)\n\n    return largest_contour\n\n\ndef crop_img(img):\n    \"\"\"\n    Crops provided image for waveform detection\n    \"\"\"\n    grid_contour = find_grid_robust(img)\n    x, y, w, h = cv.boundingRect(grid_contour)\n    cropped = img[y:y+h, x:x+w]\n    return cropped, grid_contour\n\n\ndef remove_ecg_grid(img, min_grid_length_ratio=0.5, min_speck_area=75):\n    \"\"\"\n    Removes grid lines from an ECG image while preserving the waveform.\n\n    Parameters:\n    - img: BGR image\n    - min_grid_length_ratio: fraction of image width/height to classify long straight components as grid\n    - min_speck_area: minimum area of waveform components to keep\n\n    Returns:\n    - clean_waveform: binary mask of waveform\n    \"\"\"\n    import cv2 as cv\n    import numpy as np\n\n    \n    blurred = cv.GaussianBlur(img, (5, 5), 0)\n\n    # --- Background normalization ---\n    kernel = cv.getStructuringElement(cv.MORPH_ELLIPSE, (21, 21))\n    closed = cv.morphologyEx(blurred, cv.MORPH_CLOSE, kernel)\n    div = np.float32(blurred) / (closed + 1e-8)\n    normalized = np.uint8(cv.normalize(div, div, 0, 255, cv.NORM_MINMAX))\n\n    # --- Adaptive threshold to get waveform candidates ---\n    thresh = cv.adaptiveThreshold(\n        normalized, 255,\n        cv.ADAPTIVE_THRESH_GAUSSIAN_C,\n        cv.THRESH_BINARY_INV,\n        201, 70\n    )\n\n    H, W = thresh.shape\n\n    # --- Step 1: Detect horizontal and vertical lines (likely grid) ---\n    h_kernel = cv.getStructuringElement(cv.MORPH_RECT, (max(30, W//30), 1))\n    horizontal = cv.morphologyEx(thresh, cv.MORPH_OPEN, h_kernel, iterations=1)\n\n    v_kernel = cv.getStructuringElement(cv.MORPH_RECT, (1, max(30, H//30)))\n    vertical = cv.morphologyEx(thresh, cv.MORPH_OPEN, v_kernel, iterations=1)\n\n    grid = cv.add(horizontal, vertical)\n\n    # --- Step 2: Remove long straight components based on geometry, preserving steep waveform peaks ---\n    num_labels, labels, stats, _ = cv.connectedComponentsWithStats(grid, connectivity=8)\n    grid_clean = np.zeros_like(grid)\n    for i in range(1, num_labels):\n        w = stats[i, cv.CC_STAT_WIDTH]\n        h = stats[i, cv.CC_STAT_HEIGHT]\n        aspect_ratio = w / h if h > 0 else 0\n\n        # Horizontal lines: long width relative to height\n        if w >= min_grid_length_ratio * W:\n            grid_clean[labels == i] = 255\n        # Vertical lines: long height relative to width, but ignore very steep waveform peaks\n        elif h >= min_grid_length_ratio * H and aspect_ratio < 0.05:  # adjust 0.05 if needed\n            grid_clean[labels == i] = 255\n\n    # Subtract grid from thresholded image\n    waveform = cv.subtract(thresh, grid_clean)\n\n    # --- Step 3: Enhance waveform edges ---\n    sx = cv.Sobel(normalized, cv.CV_32F, 1, 0, ksize=5)\n    sy = cv.Sobel(normalized, cv.CV_32F, 0, 1, ksize=5)\n    abs_grad = cv.normalize(np.abs(sx) + np.abs(sy), None, 0, 255, cv.NORM_MINMAX).astype(np.uint8)\n    boosted = cv.addWeighted(waveform, 1.0, abs_grad, 0.25, 0)\n\n    # Re-binarize softly\n    _, boosted = cv.threshold(boosted, 70, 255, cv.THRESH_BINARY)\n\n    # --- Step 4: Remove tiny specks and keep waveform ---\n    num_labels, labels, stats, _ = cv.connectedComponentsWithStats(boosted, connectivity=8)\n    clean_waveform = np.zeros_like(boosted)\n    for i in range(1, num_labels):\n        if stats[i, cv.CC_STAT_AREA] >= min_speck_area:\n            clean_waveform[labels == i] = 255\n\n    return clean_waveform\n\n\ndef partition_into_four_strips(img):\n    \"\"\"\n    Takes an ECG image with 4 horizontal waveform bands\n    and returns a list of 4 vertically cropped strips.\n\n    Input:\n        img : numpy array (H x W x C)\n\n    Output:\n        strips : list of 4 numpy arrays (each is a cropped vertical strip)\n    \"\"\"\n    H = img.shape[0]\n    W = img.shape[1]\n\n    # Compute equally spaced boundaries\n    # (Assumes all 4 bands are same height)\n    strip_height = H // 4\n\n    strips = []\n    for i in range(4):\n        top = i * strip_height\n        bottom = (i + 1) * strip_height\n        strips.append(img[top:bottom, :, :])\n\n    return strips\n\n\ndef generate_patches_with_mask(\n    img,                    # grayscale or grayscale-converted image\n    patch_size=64,\n    stride=32,\n    y_baseline=160,        # if None → auto center\n    px_per_mV=400,\n    leads_total=4\n):\n    \"\"\"\n    Generates training patches for a single ECG lead using the new mask generation logic.\n    \"\"\"\n   \n    # Ensure grayscale\n    if len(img.shape) == 3:\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n\n    H, W = img.shape\n\n    if y_baseline is None:\n        y_baseline = H // 2\n\n    # --- 1. Generate accurate waveform mask ---\n    mask = make_waveform_mask(img=img)\n\n    # --- 2. Extract patches ---\n    patches = []\n    labelsClf = []\n    labelsReg = []\n    coords = []   # (y0, x0) — used for visualization later\n\n    for y0 in range(0, H - patch_size + 1, stride):\n        for x0 in range(0, W - patch_size + 1, stride):\n\n            patch = img[y0:y0+patch_size, x0:x0+patch_size]\n            patch_mask = mask[y0:y0+patch_size, x0:x0+patch_size]\n\n            # Label is 1 if any mask pixels inside patch are white\n            labelClf = 1 if np.mean(patch_mask) > 0.03 else 0\n            labelReg = np.mean(patch_mask) / 255\n\n            patches.append(patch)\n            labelsClf.append(labelClf)\n            labelsReg.append(labelReg)\n            coords.append((y0, x0))\n\n    return (\n        np.array(patches),\n        np.array(labelsClf),\n        np.array(labelsReg),\n        np.array(coords),   # helpful for visualization\n        mask                # return mask for debugging\n    )\n\n\ndef crop_side_bars(strip, threshold=100, min_fraction=0.95):\n    \"\"\"\n    Crops vertical side bars from an ECG strip.\n    \n    Parameters:\n    - strip: 2D NumPy array (grayscale) of the ECG strip.\n    - threshold: pixel intensity below which a pixel is considered background.\n    - min_fraction: minimum fraction of pixels in a column below threshold\n                    to consider it a bar.\n                    \n    Returns:\n    - cropped_strip: 2D array without the side bars.\n    - x_start: leftmost pixel kept\n    - x_end: rightmost pixel kept\n    \"\"\"\n    H, W = strip.shape\n    mask = strip < threshold  # True for background pixels\n    \n    # Count fraction of background pixels in each column\n    col_fraction = mask.sum(axis=0) / H\n    \n    # Identify columns to keep\n    cols_to_keep = np.where(col_fraction < min_fraction)[0]\n    \n    if len(cols_to_keep) == 0:\n        # fallback: keep everything\n        return strip, 0, W\n    \n    x_start, x_end = cols_to_keep[0], cols_to_keep[-1] + 1\n    cropped_strip = strip[:, x_start:x_end]\n\n    if cropped_strip.shape[1] < 1500:\n        print(f\"Overcrop prevented! Returning original strip because width={cropped_strip.shape[1]}\")\n        return strip, 0, W\n    \n    return cropped_strip, x_start, x_end\n\n\ndef detect_grid_lines(strip, thresh_val=50):\n    \"\"\"\n    Detect horizontal and vertical grid lines in an ECG strip.\n    Returns binary image with only the grid lines.\n    \"\"\"\n    # Threshold to isolate bright grid lines\n    binary = cv.adaptiveThreshold(strip, 255, cv.ADAPTIVE_THRESH_MEAN_C, cv.THRESH_BINARY_INV, blockSize=15, C=10)\n\n    # Optional: morphological operations to enhance lines\n    horizontal_kernel = cv.getStructuringElement(cv.MORPH_RECT, (80, 1))\n    vertical_kernel   = cv.getStructuringElement(cv.MORPH_RECT, (1, 25))\n\n    horizontal_lines = cv.morphologyEx(binary, cv.MORPH_OPEN, horizontal_kernel)\n    vertical_lines   = cv.morphologyEx(binary, cv.MORPH_OPEN, vertical_kernel)\n\n    grid_lines = cv.bitwise_or(horizontal_lines, vertical_lines)\n    return grid_lines\n\n\ndef estimate_skew_angle(grid_lines):\n    \"\"\"\n    Estimates rotation angle to deskew the strip.\n    Returns angle in degrees.\n    \"\"\"\n    lines = cv.HoughLines(grid_lines, 1, np.pi/180, 100)\n    angles = []\n\n    if lines is not None:\n        print(\"hello\")\n        for rho, theta in lines[:,0]:\n            angle_deg = (theta * 180 / np.pi) - 90  # convert rad to deg\n            if -45 < angle_deg < 45:\n                angles.append(angle_deg)\n\n    if len(angles) == 0:\n        return 0  # fallback\n    return np.median(angles)\n\n\ndef deskew_strip_using_grid(strip):\n    \"\"\"\n    Deskew a single ECG strip using the gridlines present in the strip.\n    Assumes gridlines are black on a white background.\n    \n    Parameters:\n    - strip: grayscale image of a single ECG strip (numpy array)\n    \n    Returns:\n    - deskewed_strip: rotated image to align gridlines horizontally\n    - angle: angle in degrees that was used to rotate the strip\n    \"\"\"\n    # --- 1. Ensure grayscale ---\n    if len(strip.shape) == 3:\n        strip_gray = cv.cvtColor(strip, cv.COLOR_BGR2GRAY)\n    else:\n        strip_gray = strip.copy()\n\n    h, w = strip_gray.shape\n\n    # --- 2. Binarize (black grid on white background) ---\n    binary = cv.adaptiveThreshold(\n        strip_gray,\n        255,\n        cv.ADAPTIVE_THRESH_MEAN_C,\n        cv.THRESH_BINARY_INV,\n        blockSize=max(15, (w//200)|1),  # odd block size, scales with width\n        C=1\n    )\n\n    # --- 3. Morphological opening to extract horizontal lines ---\n    horizontal_kernel = cv.getStructuringElement(cv.MORPH_RECT, (w // 80, 1))\n    horizontal_lines = cv.morphologyEx(binary, cv.MORPH_OPEN, horizontal_kernel)\n\n    # --- 4. Probabilistic Hough Transform ---\n    lines = cv.HoughLinesP(\n        horizontal_lines,\n        rho=1,\n        theta=np.pi / 180,\n        threshold=max(50, w//60),\n        minLineLength=w // 10,\n        maxLineGap=max(5, w//200)\n    )\n\n    if lines is None:\n        # No lines found, return original\n        print(\"no liens\")\n        return strip.copy(), 0.0\n\n    # --- 5. Compute median angle ---\n    angles = []\n    for line in lines:\n        x1, y1, x2, y2 = line[0]\n        angle = np.arctan2(y2 - y1, x2 - x1) * 180 / np.pi\n        angles.append(angle)\n    median_angle = np.median(angles)\n\n    # --- 6. Rotate strip ---\n    center = (w // 2, h // 2)\n    M = cv.getRotationMatrix2D(center, median_angle, 1.0)\n    deskewed_strip = cv.warpAffine(\n        strip,\n        M,\n        (w, h),\n        flags=cv.INTER_LINEAR,\n        borderMode=cv.BORDER_REPLICATE\n    )\n\n    return deskewed_strip, median_angle\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.144348Z","iopub.execute_input":"2025-12-10T05:52:57.144579Z","iopub.status.idle":"2025-12-10T05:52:57.179434Z","shell.execute_reply.started":"2025-12-10T05:52:57.144564Z","shell.execute_reply":"2025-12-10T05:52:57.178696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport pytesseract\n\nimport cv2\nimport numpy as np\nimport pytesseract\n\ndef remove_text_and_get_coords(img, threshold_val=250, text_thresh=245, dilate_kernel_size=2, min_confidence=50):\n    \"\"\"\n    Removes text from an ECG image using Tesseract OCR with preprocessing to detect small/thin text.\n    Returns inpainted image with text regions set to black and coordinates of removed text.\n\n    Args:\n        img: Input image (color or grayscale)\n        threshold_val: Any pixel < threshold_val will be set to 0\n        text_thresh: Threshold for detecting bright text (default 245)\n        dilate_kernel_size: Kernel size for dilation (small thickening of thin text)\n        min_confidence: Minimum OCR confidence to consider a detection\n\n    Returns:\n        result_img: Image with text removed and threshold applied\n        removed_coords: numpy array of [x, y, w, h] for removed regions\n    \"\"\"\n    img_copy = img.copy()\n\n    # Convert to grayscale\n    if len(img_copy.shape) == 3:\n        gray = cv2.cvtColor(img_copy, cv2.COLOR_BGR2GRAY)\n    else:\n        gray = img_copy.copy()\n\n    # Threshold bright text\n    _, thresh = cv2.threshold(gray, text_thresh, 255, cv2.THRESH_BINARY_INV)\n\n    # Slightly dilate thin text\n    kernel = np.ones((dilate_kernel_size, dilate_kernel_size), np.uint8)\n    thresh_dilated = cv2.dilate(thresh, kernel, iterations=1)\n\n    # Tesseract OCR with config\n    custom_config = r'--oem 3 --psm 6'\n    data = pytesseract.image_to_data(thresh_dilated, config=custom_config, output_type=pytesseract.Output.DICT)\n\n    # Build mask and store bounding boxes\n    mask = np.zeros_like(gray, dtype=np.uint8)\n    removed_coords = []\n\n    for i in range(len(data['text'])):\n        conf = int(data['conf'][i])\n        text = data['text'][i].strip()\n        if conf >= min_confidence and text != '':\n            x, y, w, h = data['left'][i], data['top'][i], data['width'][i], data['height'][i]\n            cv2.rectangle(mask, (x, y), (x + w, y + h), 255, -1)\n            removed_coords.append([x, y, w, h])\n\n    removed_coords = np.array(removed_coords)\n\n    # Inpaint text regions\n    result_img = cv2.inpaint(img_copy, mask, 3, cv2.INPAINT_TELEA)\n\n    # Explicitly set removed regions to black\n    result_img[mask == 255] = 0\n\n    # Threshold remaining pixels\n    result_img[result_img < threshold_val] = 0\n\n    return result_img, removed_coords\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.181116Z","iopub.execute_input":"2025-12-10T05:52:57.181329Z","iopub.status.idle":"2025-12-10T05:52:57.199311Z","shell.execute_reply.started":"2025-12-10T05:52:57.181311Z","shell.execute_reply":"2025-12-10T05:52:57.198720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef get_strip_patches(df, image_id, image_type, strip_num, patch_size=64):\n    \"\"\"\n    Returns all patches corresponding to a specific strip_num for a given image.\n    \n    Args:\n        df: DataFrame with columns ['image_id', 'image_type', 'coords', 'image_patch', 'strip_num']\n        image_id: ID of the image\n        image_type: Type of the image\n        strip_num: Which strip to extract (1,2,3,4)\n        patch_size: Size of each patch (assumes square patches)\n    \n    Returns:\n        df_strip: DataFrame containing only patches in the specified strip\n    \"\"\"\n    # Filter by image_id, image_type, and strip_num\n    df_strip = df[\n        (df['image_id'] == image_id) &\n        (df['image_type'] == image_type) &\n        (df['strip_num'] == strip_num)\n    ].copy()\n    \n    return df_strip\n\n\ndef reconstruct_strip(df_strip, patch_size=64):\n    \"\"\"\n    Reconstructs a strip image from patches in df_strip, placing only patches\n    with labelClf == 1. No modifications to patch pixel values are performed.\n\n    Args:\n        df_strip: DataFrame with columns ['coords', 'image_patch', 'labelClf']\n                  - 'coords' is a tuple (y, x)\n                  - 'image_patch' is a flattened array/list of length patch_size*patch_size\n                  - 'labelClf' indicates whether this patch contains a waveform (1) or not\n        patch_size: Size of each patch (assumes square)\n\n    Returns:\n        canvas: 2D numpy array representing the reconstructed strip\n    \"\"\"\n    if df_strip.empty:\n        return np.zeros((patch_size, patch_size), dtype=np.uint8)\n\n    ys = df_strip[\"coords\"].apply(lambda c: int(c[0]))\n    xs = df_strip[\"coords\"].apply(lambda c: int(c[1]))\n\n    max_y = ys.max() + patch_size\n    max_x = xs.max() + patch_size\n\n    canvas = np.zeros((max_y, max_x), dtype=np.uint8)\n\n    for _, row in df_strip.iterrows():\n        if row.get(\"labelClf\", 0) != 1:\n            continue  # skip non-waveform patches\n\n        y, x = row[\"coords\"]\n        patch_flat = np.array(row[\"image_patch\"], dtype=np.uint8)\n        if patch_flat.size != patch_size * patch_size:\n            raise ValueError(f\"Patch at index {_} has incorrect size {patch_flat.size}\")\n\n        patch_2d = patch_flat.reshape(patch_size, patch_size)\n        canvas[y:y+patch_size, x:x+patch_size] = patch_2d\n\n    return canvas\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.200041Z","iopub.execute_input":"2025-12-10T05:52:57.200300Z","iopub.status.idle":"2025-12-10T05:52:57.219526Z","shell.execute_reply.started":"2025-12-10T05:52:57.200284Z","shell.execute_reply":"2025-12-10T05:52:57.218873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def reconstruct_strip_debug(df_strip, patch_size=64):\n    if df_strip.empty:\n        return np.zeros((patch_size, patch_size), dtype=np.uint8)\n\n    ys = [c[0] for c in df_strip[\"coords\"]]\n    xs = [c[1] for c in df_strip[\"coords\"]]\n\n    max_y = max(ys) + patch_size\n    max_x = max(xs) + patch_size\n\n    canvas = np.zeros((max_y, max_x), dtype=np.uint8)\n\n    placed = 0\n    for i, row in df_strip.iterrows():\n        y, x = row[\"coords\"]\n        if row[\"labelClf\"] != 1:\n            continue\n        patch = np.array(row[\"image_patch\"], dtype=np.uint8)\n        if patch.size != patch_size**2:\n            print(f\"Skipping patch {i}: wrong size {patch.size}\")\n            continue\n        patch = patch.reshape(patch_size, patch_size)\n\n        # Safety: check bounds\n        if y + patch_size > canvas.shape[0] or x + patch_size > canvas.shape[1]:\n            print(f\"Skipping patch {i}: coords out of bounds ({y},{x})\")\n            continue\n\n        canvas[y:y+patch_size, x:x+patch_size] = patch\n        placed += 1\n\n    print(f\"Patches placed: {placed}/{len(df_strip)}\")\n    return canvas\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.220214Z","iopub.execute_input":"2025-12-10T05:52:57.220435Z","iopub.status.idle":"2025-12-10T05:52:57.235646Z","shell.execute_reply.started":"2025-12-10T05:52:57.220414Z","shell.execute_reply":"2025-12-10T05:52:57.234907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def scale_ppm(base_size, base_ppm, new_size):\n    \"\"\"\n    base_size: (H, W) of reference image\n    base_ppm: known px per mm for reference image\n    new_size: (H, W) of new image\n\n    Returns scaled pixels per mm (float)\n    \"\"\"\n    _, base_w = base_size\n    _, new_w = new_size\n\n    # Scale linearly by width\n    return base_ppm * (new_w / base_w)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.236337Z","iopub.execute_input":"2025-12-10T05:52:57.236540Z","iopub.status.idle":"2025-12-10T05:52:57.248413Z","shell.execute_reply.started":"2025-12-10T05:52:57.236526Z","shell.execute_reply":"2025-12-10T05:52:57.247859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_lead_gt_trace(groundTruth_df, lead_name, color=\"red\", name=None):\n    \"\"\"\n    Extracts the non-NaN portion of a single ECG lead and returns\n    a Plotly Scatter trace for plotting.\n\n    Parameters\n    ----------\n    groundTruth_df : pd.DataFrame\n        ECG CSV containing multiple leads.\n    lead_name : str\n        Name of the lead to extract (e.g., 'aVR', 'V1').\n    color : str\n        Line color for the plot.\n    name : str\n        Optional trace name. If None, defaults to lead_name.\n\n    Returns\n    -------\n    trace : plotly.graph_objects.Scatter\n        Plotly trace for the ground truth lead.\n    \"\"\"\n    y_vals = groundTruth_df[lead_name].to_numpy(dtype=float)\n    # Get indices of non-NaN values\n    x_vals = np.flatnonzero(~np.isnan(y_vals))\n    y_vals = y_vals[x_vals]\n\n    trace = go.Scatter(\n        x=x_vals,\n        y=y_vals,\n        mode=\"lines\",\n        name=name if name else lead_name,\n        line=dict(width=1, dash=\"dot\", color=color)\n    )\n    return trace","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.249268Z","iopub.execute_input":"2025-12-10T05:52:57.249489Z","iopub.status.idle":"2025-12-10T05:52:57.265354Z","shell.execute_reply.started":"2025-12-10T05:52:57.249472Z","shell.execute_reply":"2025-12-10T05:52:57.264568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pixels_to_mV(waveform_pixels, baseline_px, px_per_mV=80):\n    \"\"\"\n    Converts y-pixels relative to baseline to mV.\n    \n    waveform_pixels: array of y-pixels in image coordinates\n    baseline_px: pixel coordinate of the baseline\n    \"\"\"\n    # Distance from baseline (positive = above baseline, negative = below)\n    dist_from_baseline = baseline_px - waveform_pixels\n    return dist_from_baseline / px_per_mV\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.266190Z","iopub.execute_input":"2025-12-10T05:52:57.266752Z","iopub.status.idle":"2025-12-10T05:52:57.279053Z","shell.execute_reply.started":"2025-12-10T05:52:57.266731Z","shell.execute_reply":"2025-12-10T05:52:57.278243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\nimport numpy as np\n\ndef extract_waveform_pixels(strip_binary, grid_px=40, debug=False):\n    \"\"\"\n    Extract a 1-pixel-per-column waveform from a binary strip, ignoring\n    leading/trailing 2-grid-square margins.\n\n    Heuristics preserved:\n      - pairwise diff < 20 → median\n      - pairwise diff > 90 → prev_y\n      - otherwise → weighted avg based on distance from prev_y\n    \"\"\"\n\n    H, W = strip_binary.shape\n    margin = 2 * grid_px\n    waveform = np.full(W, np.nan, dtype=np.float32)\n\n    # Collect y indices per column\n    cols_ys = [None] * W\n    nonzero_counts = np.zeros(W, dtype=int)\n    for x in range(W):\n        ys = np.where(strip_binary[:, x] > 0)[0]\n        if ys.size:\n            cols_ys[x] = ys\n            nonzero_counts[x] = ys.size\n\n    # Set margins as NaN\n    for x in range(min(margin, W)):\n        waveform[x] = np.nan\n    for x in range(max(0, W - margin), W):\n        waveform[x] = np.nan\n\n    # Identify usable columns\n    nz_cols = np.nonzero(nonzero_counts)[0]\n    nz_cols = nz_cols[(nz_cols >= margin) & (nz_cols < W - margin)]\n    if nz_cols.size == 0:\n        return waveform\n\n    # Seed first column\n    first_col = nz_cols[0]\n    ys0 = cols_ys[first_col]\n    prev_y = int(np.median(ys0))\n    waveform[first_col] = prev_y\n\n    # Walk forward\n    for x in range(first_col + 1, nz_cols[-1] + 1):\n\n        ys = cols_ys[x]\n        if ys is None:\n            waveform[x] = np.nan\n            continue\n\n        pairwise_diff = np.max(ys) - np.min(ys)\n\n        if pairwise_diff < 20:\n            # Tight cluster → strong confidence\n            y = int(np.median(ys))\n\n        else:\n            # ---- NEW WEIGHTED AVERAGE CASE (20 ≤ diff ≤ 90) ----\n            diffs = np.abs(ys - prev_y) + 1e-3  # avoid div by zero\n            weights = 1.0 / diffs               # closer → higher weight\n            y = np.sum(weights * ys) / np.sum(weights)\n            y = int(y)\n        # ----------------------------------\n\n        waveform[x] = y\n        prev_y = y\n\n    # Interpolate internal NaNs\n    nan_mask = np.isnan(waveform)\n    good_idx = np.flatnonzero(~nan_mask)\n    bad_idx = np.flatnonzero(nan_mask)\n\n    if good_idx.size > 1:\n        left = good_idx[0]\n        right = good_idx[-1]\n        internal_bad = [x for x in bad_idx if left < x < right]\n        if internal_bad:\n            internal_bad = np.array(internal_bad)\n            waveform[internal_bad] = np.interp(\n                internal_bad,\n                good_idx,\n                waveform[good_idx]\n            )\n\n    if debug:\n        print(\"NaNs remaining:\", np.sum(np.isnan(waveform)))\n        print(\"Waveform y-range:\", (np.nanmin(waveform), np.nanmax(waveform)))\n\n    return waveform\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.281292Z","iopub.execute_input":"2025-12-10T05:52:57.281504Z","iopub.status.idle":"2025-12-10T05:52:57.295688Z","shell.execute_reply.started":"2025-12-10T05:52:57.281489Z","shell.execute_reply":"2025-12-10T05:52:57.295043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.interpolate import interp1d\nimport numpy as np\n\nfrom scipy.interpolate import interp1d\nimport numpy as np\n\ndef resample_no_padding(signal, target_len):\n    \"\"\"\n    Resample extracted waveform to exactly `target_len`\n    using interpolation. Ignores NaN-only boundaries.\n    \"\"\"\n    signal = np.asarray(signal, dtype=float)\n\n    # ------------------\n    # 1. Remove leading/trailing NaNs\n    # ------------------\n    valid = ~np.isnan(signal)\n    if valid.sum() < 2:\n        return np.full(target_len, np.nan)\n\n    first = np.argmax(valid)\n    last = len(signal) - np.argmax(valid[::-1]) - 1\n    trimmed = signal[first:last+1]\n\n    # ------------------\n    # 2. If still too few points, return NaNs\n    # ------------------\n    if len(trimmed) < 2:\n        return np.full(target_len, np.nan)\n\n    # ------------------\n    # 3. Create interpolation grid\n    # ------------------\n    old_x = np.linspace(0, 1, len(trimmed))\n    new_x = np.linspace(0, 1, target_len)\n\n    # ------------------\n    # 4. Interpolate\n    # ------------------\n    resampled = np.interp(new_x, old_x, trimmed)\n\n    return resampled\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.296447Z","iopub.execute_input":"2025-12-10T05:52:57.296728Z","iopub.status.idle":"2025-12-10T05:52:57.312865Z","shell.execute_reply.started":"2025-12-10T05:52:57.296713Z","shell.execute_reply":"2025-12-10T05:52:57.312188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef extract_ground_truth_strip(df, strip_num, fs10=5000):\n    \"\"\"\n    Extract the ground-truth 10-second waveform for the chosen strip number.\n    \n    Parameters\n    ----------\n    df : pandas.DataFrame\n        The dataframe for a single ECG recording (one image).\n        Columns include the 12 standard leads (I, II, III, aVR, aVL, aVF, V1-V6)\n        Each column is length fs*10sec with internal NaNs where the lead is not shown.\n        \n    strip_num : int\n        Which strip to extract (1, 2, 3, or 4).\n        \n    fs10 : int\n        Total length of a 10-second strip (default = 5000 samples).\n        \n    Returns\n    -------\n    full_wave : np.ndarray\n        5000-sample 1D array representing the correct 10-second ground truth\n        waveform for that strip.\n    \"\"\"\n    \n    # ----------------------------\n    # Strip → leads mapping\n    # ----------------------------\n    strip_map = {\n        1: [\"I\", \"aVR\", \"V1\", \"V4\"],\n        2: [\"II\", \"aVL\", \"V2\", \"V5\"],\n        3: [\"III\", \"aVF\", \"V3\", \"V6\"],\n        4: [\"II\", \"II\", \"II\", \"II\"],   # continuous 10-sec II\n    }\n    \n    if strip_num not in strip_map:\n        raise ValueError(f\"Invalid strip_num {strip_num}. Must be 1–4.\")\n\n    leads = strip_map[strip_num]\n\n    # Length per quadrant (5000 / 4)\n    seg_len = fs10 // 4\n\n    segments = []\n\n    for i, lead_name in enumerate(leads):\n        col = df[lead_name].values  # numpy array\n        \n        # Expected boundaries of this quadrant\n        start = i * seg_len\n        end = (i + 1) * seg_len\n\n        # Slice the expected region\n        seg = col[start:end]\n\n        # In the dataset, unused regions are filled with NaN.\n        # We want only the actual non-NaN signal.\n        valid = seg[~np.isnan(seg)]\n\n        if len(valid) == 0:\n            # If for some reason the segment is empty, fill with NaN\n            valid = np.full(seg_len, np.nan)\n        else:\n            # If the valid region isn’t exactly seg_len long, resample it\n            if len(valid) != seg_len:\n                valid = np.interp(\n                    np.linspace(0, 1, seg_len),\n                    np.linspace(0, 1, len(valid)),\n                    valid\n                )\n        \n        segments.append(valid)\n\n    # Concatenate the 4 × 1250 signals\n    full_wave = np.concatenate(segments)\n\n    # Ensure final length is exactly fs10\n    if len(full_wave) != fs10:\n        full_wave = full_wave[:fs10]\n\n    return full_wave\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.313468Z","iopub.execute_input":"2025-12-10T05:52:57.313650Z","iopub.status.idle":"2025-12-10T05:52:57.323270Z","shell.execute_reply.started":"2025-12-10T05:52:57.313636Z","shell.execute_reply":"2025-12-10T05:52:57.322747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom scipy.interpolate import interp1d\n\ndef resample_no_padding(wave, signal_len):\n    \"\"\"\n    Extracts only the non-NaN interior of a waveform and resamples it\n    to exactly `signal_len` samples. No padding, no shifting.\n\n    Parameters\n    ----------\n    wave : 1D array\n        Extracted waveform with possible NaNs at edges.\n    signal_len : int\n        Desired output length.\n\n    Returns\n    -------\n    out : 1D array, shape (signal_len,)\n        Resampled waveform with no padding.\n        Returns all NaNs if the waveform has no valid samples.\n    \"\"\"\n\n    wave = wave.astype(float)\n    \n    # Find non-NaN region\n    good = ~np.isnan(wave)\n    good_idx = np.flatnonzero(good)\n\n    if good_idx.size == 0:\n        return np.full(signal_len, np.nan)\n\n    # Extract valid interior segment\n    left = good_idx[0]\n    right = good_idx[-1]\n    interior = wave[left : right + 1]\n\n    # If interior is constant or length 1 → just repeat the value\n    if len(interior) == 1:\n        return np.full(signal_len, interior[0])\n\n    # Resample interior to desired length\n    x_old = np.linspace(0, 1, len(interior))\n    x_new = np.linspace(0, 1, signal_len)\n\n    f = interp1d(x_old, interior, kind=\"linear\")\n    return f(x_new)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.324065Z","iopub.execute_input":"2025-12-10T05:52:57.324327Z","iopub.status.idle":"2025-12-10T05:52:57.341156Z","shell.execute_reply.started":"2025-12-10T05:52:57.324307Z","shell.execute_reply":"2025-12-10T05:52:57.340434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef snr_waveforms(ground_truth, extracted):\n    \"\"\"\n    Compute SNR between a ground-truth waveform and an extracted waveform.\n    Handles unequal lengths and removes NaNs from extracted before comparison.\n    \"\"\"\n\n    # Convert to arrays\n    gt = np.asarray(ground_truth)\n    ext = np.asarray(extracted)\n\n    # Step 1 — Trim to equal length\n    L = min(len(gt), len(ext))\n    gt = gt[:L]\n    ext = ext[:L]\n\n    # Step 2 — Remove NaNs from extracted AND corresponding GT samples\n    valid = ~np.isnan(ext)\n    gt = gt[valid]\n    ext = ext[valid]\n\n    if gt.size == 0:\n        raise ValueError(\"No valid extracted samples after removing NaNs.\")\n\n    # Step 3 — Compute noise\n    noise = gt - ext\n\n    signal_power = np.sum(gt**2)\n    noise_power = np.sum(noise**2)\n\n    if noise_power == 0:\n        return np.inf\n\n    return 10 * np.log10(signal_power / noise_power)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.341853Z","iopub.execute_input":"2025-12-10T05:52:57.342139Z","iopub.status.idle":"2025-12-10T05:52:57.354357Z","shell.execute_reply.started":"2025-12-10T05:52:57.342123Z","shell.execute_reply":"2025-12-10T05:52:57.353675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_horizontal_lines(img, min_length=75):\n    \"\"\"\n    Removes horizontal lines longer than min_length pixels from a 2D image.\n\n    Args:\n        img: 2D numpy array (grayscale or binary)\n        min_length: Minimum length in pixels to consider a horizontal line\n\n    Returns:\n        img_clean: 2D numpy array with long horizontal lines removed\n    \"\"\"\n    # Make a copy so we don't modify original\n    # Ensure the image is uint8\n    if img.dtype != np.uint8:\n        img_uint8 = np.clip(img, 0, 255).astype(np.uint8)\n    else:\n        img_uint8 = img.copy()\n    \n    # Define horizontal structure\n    horizontal_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (min_length, 1))\n    \n    # Detect horizontal lines\n    detected_lines = cv2.morphologyEx(img_uint8, cv2.MORPH_OPEN, horizontal_kernel)\n    \n    # Invert detected lines mask\n    mask = cv2.bitwise_not(detected_lines)\n    \n    # Remove lines by masking\n    img_clean = cv2.bitwise_and(img_uint8, img_uint8, mask=mask)\n    \n    return img_clean","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.355043Z","iopub.execute_input":"2025-12-10T05:52:57.355233Z","iopub.status.idle":"2025-12-10T05:52:57.369504Z","shell.execute_reply.started":"2025-12-10T05:52:57.355213Z","shell.execute_reply":"2025-12-10T05:52:57.368814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport plotly.graph_objects as go\n\ndef plot_waveforms_snr(ground_truth, extracted, title=\"Waveform Comparison\"):\n    \"\"\"\n    Plots ground truth vs extracted waveform and prints SNR.\n\n    Args:\n        ground_truth: 1D numpy array of ground truth waveform\n        extracted: 1D numpy array of extracted waveform\n        title: Plot title\n\n    Returns:\n        fig: Plotly figure\n        snr: Signal-to-noise ratio in dB\n    \"\"\"\n    # Ensure same length\n    min_len = min(len(ground_truth), len(extracted))\n    gt = ground_truth[:min_len]\n    ex = extracted[:min_len]\n\n    # Compute SNR\n    noise = gt - ex\n    snr = 10 * np.log10(np.sum(gt**2) / (np.sum(noise**2) + 1e-12))  # avoid div0\n\n    # Create Plotly figure\n    fig = go.Figure()\n    fig.add_trace(go.Scatter(y=gt, mode='lines', name='Ground Truth', line=dict(color='blue')))\n    fig.add_trace(go.Scatter(y=ex, mode='lines', name='Extracted', line=dict(color='red')))\n    fig.update_layout(title=f\"{title} | SNR: {snr:.2f} dB\",\n                      xaxis_title=\"Sample Index\",\n                      yaxis_title=\"Amplitude\",\n                      legend=dict(x=0.02, y=0.98))\n    fig.show()\n\n    return fig, snr\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.370298Z","iopub.execute_input":"2025-12-10T05:52:57.370536Z","iopub.status.idle":"2025-12-10T05:52:57.382717Z","shell.execute_reply.started":"2025-12-10T05:52:57.370512Z","shell.execute_reply":"2025-12-10T05:52:57.381953Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Time Series Data Work","metadata":{}},{"cell_type":"code","source":"pip install opencv-python pytesseract","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:52:57.383437Z","iopub.execute_input":"2025-12-10T05:52:57.383657Z","iopub.status.idle":"2025-12-10T05:53:05.076728Z","shell.execute_reply.started":"2025-12-10T05:52:57.383633Z","shell.execute_reply":"2025-12-10T05:53:05.075683Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare Models","metadata":{}},{"cell_type":"code","source":"import pickle\n\nwith open(\"/kaggle/input/physionet-waveform-clf-rf/scikitlearn/default/1/rf_model.pkl\", \"rb\") as f:\n    rf = pickle.load(f)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:53:05.077934Z","iopub.execute_input":"2025-12-10T05:53:05.078348Z","iopub.status.idle":"2025-12-10T05:53:06.826323Z","shell.execute_reply.started":"2025-12-10T05:53:05.078319Z","shell.execute_reply":"2025-12-10T05:53:06.825569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# -----------------------------\n# SAME MODEL DEFINITION\n# -----------------------------\nH, W = 64, 64  \n\nclass CNN(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.conv1 = nn.Conv2d(1, 16, 3, padding=1)\n        self.conv2 = nn.Conv2d(16, 32, 3, padding=1)\n        self.fc1   = nn.Linear(32 * (H//4) * (W//4), 64)\n        self.fc2   = nn.Linear(64, num_classes)\n\n    def forward(self, x):\n        x = F.relu(self.conv1(x))\n        x = F.max_pool2d(x, 2)\n        x = F.relu(self.conv2(x))\n        x = F.max_pool2d(x, 2)\n        x = x.view(x.size(0), -1)\n        x = F.relu(self.fc1(x))\n        return self.fc2(x)\n\n# --------------------------------------\n# LOAD MODEL FROM KAGGLE DATASET\n# --------------------------------------\n\n\nmodel_path = \"/kaggle/input/physionet-waveform-classifier-cnn-final/pytorch/default/1/cnn_final.pth\"\n\n\nnum_classes = 2\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = CNN(num_classes).to(device)\n\n\nstate_dict = torch.load(model_path, map_location=torch.device(device))\n\nmodel.load_state_dict(state_dict)\nmodel.eval()\n\nprint(\"Model loaded and ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:53:06.827361Z","iopub.execute_input":"2025-12-10T05:53:06.827925Z","iopub.status.idle":"2025-12-10T05:53:07.099888Z","shell.execute_reply.started":"2025-12-10T05:53:06.827902Z","shell.execute_reply":"2025-12-10T05:53:07.099135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_patches(model, patches, batch_size=256):\n    \"\"\"\n    patches: list or numpy array of patches, flat or 2D\n    \"\"\"\n    model.eval()\n\n    # Detect model device automatically\n    device = next(model.parameters()).device\n\n    all_preds = []\n    all_probs = []\n\n    patches = np.array(patches)\n\n    # If patches are flat (4096,) → reshape\n    if patches.ndim == 2 and patches.shape[1] == 4096:\n        patches = patches.reshape(-1, 64, 64)\n\n    patches = patches.astype(np.float32)\n\n    with torch.no_grad():\n        for i in range(0, len(patches), batch_size):\n            batch = patches[i:i+batch_size]\n\n            batch = torch.tensor(batch).unsqueeze(1).to(device)  # (B,1,64,64)\n\n            logits = model(batch)\n            probs  = F.softmax(logits, dim=1)\n            preds  = torch.argmax(probs, dim=1)\n\n            all_preds.append(preds.cpu().numpy())\n            all_probs.append(probs.cpu().numpy())\n\n    return np.concatenate(all_preds), np.concatenate(all_probs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:53:07.100816Z","iopub.execute_input":"2025-12-10T05:53:07.101376Z","iopub.status.idle":"2025-12-10T05:53:07.107047Z","shell.execute_reply.started":"2025-12-10T05:53:07.101357Z","shell.execute_reply":"2025-12-10T05:53:07.106432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Perform Preprocessing on images, extract waveform without using ML, for Ablation Study","metadata":{}},{"cell_type":"code","source":"from plotly.subplots import make_subplots\nimport plotly.graph_objects as go\nimport statistics\nimport cv2 as cv\nimport numpy as np\n\n# Random set of ECG IDs\nid_set = train_df['id'].values[42:48]\ntotal_snr = 0;\n\nfor curid in id_set:\n    metadata = train_df[train_df[\"id\"] == curid]\n    signal_length = int(metadata[\"fs\"].values[0] * 10)  # 10 seconds\n\n    # Load ECG image\n    img_path = f\"/kaggle/input/physionet-ecg-image-digitization/train/{curid}/{curid}-0001.png\"\n    img = cv.imread(img_path)\n    ppm = scale_ppm((1700, 2200), 8, (img.shape[0], img.shape[1]))\n    px_per_bigSquare = ppm * 5\n\n    # Preprocess\n    result, _, _ = crop_and_deskew_ecg(img)\n    bands = partition_into_four_strips(result)\n    if len(bands) < 2:\n        continue  # skip if less than 2 strips\n\n    b = bands[0]  # strip 1\n\n    gray = cv.cvtColor(b, cv.COLOR_BGR2GRAY)\n    noBars, _, _ = crop_side_bars(gray)\n    if noBars.shape[1] < 40:\n        continue\n\n    deskewed, _ = deskew_strip_using_grid(noBars)\n    clean = remove_ecg_grid(deskewed)\n    noLines = remove_horizontal_lines(clean)\n    img_no_text, text_coords = remove_text_and_get_coords(noLines)\n\n    # Extract waveform\n    waveform_ys = extract_waveform_pixels(img_no_text)\n    resampled_no_header = waveform_ys[int(3 * px_per_bigSquare):]  # skip left 3 big squares\n    extracted_waveform = resample_no_padding(resampled_no_header, signal_length)\n    extracted_timeseries = pixels_to_mV(\n        extracted_waveform,\n        statistics.mode(extracted_waveform) - 3\n    )\n\n    # Load GT CSV\n    groundTruth_df = pd.read_csv(\n        f\"/kaggle/input/physionet-ecg-image-digitization/train/{curid}/{curid}.csv\"\n    )\n\n    # Extract strip 1 GT\n    stripGT = extract_ground_truth_strip(groundTruth_df, 1)  # strip 1\n\n    # Ensure it is a 1D float array\n    stripGT = np.asarray(stripGT, dtype=float).flatten()\n\n    ext_len = len(extracted_timeseries)\n    gt_len = len(stripGT)\n\n    # Pad with NaN if GT shorter than extracted waveform\n    gt_aligned = np.full(ext_len, np.nan)\n    assign_len = min(gt_len, ext_len)  # don't exceed extracted length\n    gt_aligned[:assign_len] = stripGT[:assign_len]\n\n\n    # Compute SNR using only valid GT points\n    valid_mask = ~np.isnan(gt_aligned)\n    noise = extracted_timeseries[valid_mask] - gt_aligned[valid_mask]\n    snr = np.mean(gt_aligned[valid_mask]**2) / (np.mean(noise**2) + 1e-9)\n    snr_db = 10 * np.log10(snr)\n    total_snr += snr_db\n    # Plot extracted waveform\n    fig = go.Figure()\n    fig.add_trace(\n        go.Scatter(\n            y=extracted_timeseries,\n            name=\"Extracted waveform\",\n            mode=\"lines\",\n            line=dict(width=1)\n        )\n    )\n\n    # Plot individual GT leads overlaid\n    lead_colors = {'I': 'red', 'aVR': 'blue', 'V1': 'green', 'V4': 'orange'}\n    for lead_name, color in lead_colors.items():\n        gt_trace = get_lead_gt_trace(groundTruth_df, lead_name, color=color)\n        fig.add_trace(gt_trace)\n\n    fig.update_layout(\n        height=400,\n        width=900,\n        title=f\"ECG {curid} - Strip 1 | SNR={snr_db:.2f} dB\",\n        xaxis_title=\"Sample index\",\n        yaxis_title=\"mV\"\n    )\n\n    fig.show()\nprint(\"Average SNR: \", total_snr/6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:53:07.107814Z","iopub.execute_input":"2025-12-10T05:53:07.108042Z","iopub.status.idle":"2025-12-10T05:53:12.635583Z","shell.execute_reply.started":"2025-12-10T05:53:07.108003Z","shell.execute_reply":"2025-12-10T05:53:12.634681Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Use the trained ML model to classify small patches of image\n# Extract Waveform from image made of reconstructed, classified patches","metadata":{}},{"cell_type":"code","source":"import statistics\nimport plotly.graph_objects as go\n\n# Random set of ECG IDs\nid_set = train_df['id'].values[42:48]\ntotal_snr= 0\nfor curid in id_set:\n    metadata = train_df[train_df[\"id\"] == curid]\n    signal_length = int(metadata[\"fs\"].values[0] * 10)  # 10 seconds\n\n    # Load ECG image\n    img_path = f\"/kaggle/input/physionet-ecg-image-digitization/train/{curid}/{curid}-0001.png\"\n    img = cv.imread(img_path)\n    ppm = scale_ppm((1700, 2200), 8, (img.shape[0], img.shape[1]))\n    px_per_bigSquare = ppm * 5\n\n    # Preprocess\n    result, _, _ = crop_and_deskew_ecg(img)\n    bands = partition_into_four_strips(result)\n    if len(bands) < 2:\n        continue  # skip if less than 2 strips\n\n    b = bands[1]  # strip 1\n\n    gray = cv.cvtColor(b, cv.COLOR_BGR2GRAY)\n    noBars, _, _ = crop_side_bars(gray)\n    if noBars.shape[0] < 40:\n        continue\n\n    deskewed, _ = deskew_strip_using_grid(noBars)\n    clean = remove_ecg_grid(deskewed)\n    noLines = remove_horizontal_lines(clean)\n    img_no_text, text_coords = remove_text_and_get_coords(noLines)\n    \n    # Extract patches and predict with RF\n    patches, coords = extract_patches(img_no_text)\n    X = np.array([p.flatten() for p in patches])\n    labels, probs = predict_patches(model, X)\n\n    patchDF_TEMP = pd.DataFrame({\n        'coords': coords,\n        'image_patch': [p.flatten() for p in patches],\n        'labelClf': labels\n    })\n\n    waveform_patches = patchDF_TEMP[patchDF_TEMP['labelClf'] == 1]\n    \n    # Reconstruct strip using only positive patches\n    strip_img = reconstruct_strip(waveform_patches)\n    img_no_text, text_coords = remove_text_and_get_coords(strip_img)\n    noLines = remove_horizontal_lines(img_no_text)\n    waveform_ys = extract_waveform_pixels(noLines)\n    \n    # Resample to match signal length\n    resampled_no_header = waveform_ys[int(px_per_bigSquare * 3):]\n   \n    resampled_waveform_ys = resample_no_padding(resampled_no_header, signal_length)\n    extracted_timeseries = pixels_to_mV(resampled_waveform_ys, statistics.mode(resampled_waveform_ys)-3)\n\n    # Load GT CSV\n    groundTruth_df = pd.read_csv(\n        f\"/kaggle/input/physionet-ecg-image-digitization/train/{curid}/{curid}.csv\"\n    )\n\n    # Extract strip 1 GT\n    stripGT = extract_ground_truth_strip(groundTruth_df, 1)  # strip 1\n\n    # Ensure it is a 1D float array\n    stripGT = np.asarray(stripGT, dtype=float).flatten()\n\n    ext_len = len(extracted_timeseries)\n    gt_len = len(stripGT)\n\n    # Pad with NaN if GT shorter than extracted waveform\n    gt_aligned = np.full(ext_len, np.nan)\n    assign_len = min(gt_len, ext_len)  # don't exceed extracted length\n    gt_aligned[:assign_len] = stripGT[:assign_len]\n\n\n    # Compute SNR using only valid GT points\n    valid_mask = ~np.isnan(gt_aligned)\n    noise = extracted_timeseries[valid_mask] - gt_aligned[valid_mask]\n    snr = np.mean(gt_aligned[valid_mask]**2) / (np.mean(noise**2) + 1e-9)\n    snr_db = 10 * np.log10(snr)\n    total_snr += snr_db\n    # Plot extracted waveform\n    fig = go.Figure()\n    fig.add_trace(\n        go.Scatter(\n            y=extracted_timeseries,\n            name=\"Extracted waveform\",\n            mode=\"lines\",\n            line=dict(width=1)\n        )\n    )\n\n    # Plot individual GT leads overlaid\n    lead_colors = {'I': 'red', 'aVR': 'blue', 'V1': 'green', 'V4': 'orange'}\n    for lead_name, color in lead_colors.items():\n        gt_trace = get_lead_gt_trace(groundTruth_df, lead_name, color=color)\n        fig.add_trace(gt_trace)\n\n    fig.update_layout(\n        height=400,\n        width=900,\n        title=f\"ECG {curid} - Strip 1 | SNR={snr_db:.2f} dB\",\n        xaxis_title=\"Sample index\",\n        yaxis_title=\"mV\"\n    )\n\n    fig.show()\nprint(\"Average SNR: \", total_snr/6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T05:53:21.580450Z","iopub.execute_input":"2025-12-10T05:53:21.581205Z","iopub.status.idle":"2025-12-10T05:53:26.266694Z","shell.execute_reply.started":"2025-12-10T05:53:21.581185Z","shell.execute_reply":"2025-12-10T05:53:26.265843Z"}},"outputs":[],"execution_count":null}]}