{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":97984,"databundleVersionId":14096757,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":15884537,"datasetId":10184951,"databundleVersionId":16838352},{"sourceType":"datasetVersion","sourceId":15773400,"datasetId":10109573,"databundleVersionId":16718484},{"sourceType":"datasetVersion","sourceId":15870045,"datasetId":10174563,"databundleVersionId":16822712},{"sourceType":"datasetVersion","sourceId":15884760,"datasetId":10185101,"databundleVersionId":16838585},{"sourceType":"datasetVersion","sourceId":15878646,"datasetId":10180675,"databundleVersionId":16831987},{"sourceType":"datasetVersion","sourceId":15884654,"datasetId":10185027,"databundleVersionId":16838473},{"sourceType":"kernelVersion","sourceId":312109075,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":312121255,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\nfrom glob import glob\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict\nfrom tqdm import tqdm\nfrom scipy.signal import medfilt\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import r2_score\n    \nfrom tensorflow.keras.layers import Input, Dense, Activation, Reshape, GaussianNoise\nfrom tensorflow.keras.initializers import Constant\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.losses import MeanSquaredError\nfrom tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, TerminateOnNaN","metadata":{"_uuid":"65d50c78-742b-4f8b-a643-910efcf80720","_cell_guid":"fac93c2f-1b3b-4f63-abd9-20fd00368c02","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:29:46.172088Z","iopub.execute_input":"2026-04-22T21:29:46.172365Z","iopub.status.idle":"2026-04-22T21:29:54.137368Z","shell.execute_reply.started":"2026-04-22T21:29:46.172328Z","shell.execute_reply":"2026-04-22T21:29:54.136413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# dataloader","metadata":{"_uuid":"7401e42a-c596-407a-aa2d-85b734037bb5","_cell_guid":"ae6797c5-586a-433e-83ef-3253fa9ff56a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## marker finder","metadata":{"_uuid":"fcf3d224-2480-47b3-a6c0-deebe679805d","_cell_guid":"664ea88d-165e-4ec2-a259-14709c21f4ae","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nclass MarkerFinder:\n    \"\"\"This class finds the 13 markers in scanned ecg images and guesses the 4 line ends.\"\"\"\n    \n    def __init__(self, show_templates=False):\n        # Derive the templates from type 1 images\n        # np.max keeps the gridlines and markers and removes the ecg lines\n        ima = np.max([\n            cv2.imread('/kaggle/input/competitions/physionet-ecg-image-digitization/train/1006867983/1006867983-0001.png'),\n            cv2.imread('/kaggle/input/competitions/physionet-ecg-image-digitization/train/102150619/102150619-0001.png'),\n            cv2.imread('/kaggle/input/competitions/physionet-ecg-image-digitization/train/1079294623/1079294623-0001.png'),\n        ], axis=0)\n\n        # Template points in global coordinates of type 1 images\n        absolute_points = np.zeros((17, 2), dtype=int)\n        for i in range(3):\n            absolute_points[5 * i] = np.array([707 + 284 * i, 118]) # y, x\n            for j in range(1, 5):\n                absolute_points[5 * i + j] = np.array([707 + 284 * i, 118 + 492 * j])\n        absolute_points[5 * 3] = np.array([1535, 118])\n        absolute_points[5 * 3 + 1] = np.array([1535, 118 + 492 * 4])\n\n        # Top left corner of template rectangle\n        template_positions = [None] * 17\n        for i in range(len(absolute_points)):\n            if absolute_points[i][1] < 118 + 492 * 4:\n                if i % 5 == 0:\n                    template_positions[i] = (absolute_points[i][0] - 87, absolute_points[i][1] - 50) # y, x\n                else:\n                    template_positions[i] = (absolute_points[i][0] - 37, absolute_points[i][1] - 13)\n\n        # Height and width of the templates\n        template_sizes = np.array([(105, 60)] * 17) # height, width\n\n        # Transform the points to relative coordinates (inside the template)\n        template_points = [np.array([absolute_points[i][0] - template_positions[i][0],\n                                     absolute_points[i][1] - template_positions[i][1]])\n                           if template_positions[i] is not None\n                           else None\n                           for i in range(len(absolute_points))]\n\n        # Save the template matrices\n        templates = [None] * 17\n        for i in range(len(template_positions)):\n            if template_points[i] is not None:\n                template = (ima[template_positions[i][0]:template_positions[i][0]+template_sizes[i][0],\n                            template_positions[i][1]:template_positions[i][1]+template_sizes[i][1]])\n                templates[i] = template\n\n        if show_templates:\n            _, axs = plt.subplots(4, 4, figsize=(5, 7))\n            for i in range(len(template_positions)):\n                if template_points[i] is not None:\n                    template = templates[i].copy()\n                    cv2.rectangle(template,\n                                  (template_points[i][1]-1, template_points[i][0]-1),\n                                  (template_points[i][1]+1, template_points[i][0]+1), \n                                  [255, 0, 0], 2)\n                    axs[i // 5, i % 5].imshow(template)\n            for i in range(13, len(axs.ravel())):\n                axs.ravel()[i].axis('off')\n            plt.tight_layout()\n            plt.suptitle('The templates for the 13 markers', y=1.01)\n            plt.show()\n\n        self._absolute_points = absolute_points\n        self._template_positions = template_positions\n        self._template_sizes = template_sizes\n        self._template_points = template_points\n        self._templates = templates\n        \n    def find_markers(self, ima, warn=False, plot=False, title=''):\n        \"\"\"Return 17 markers as list of size-2 integer arrays (row, column)\"\"\"\n        if ima.shape[0] != 1652:\n            # For this pipeline, we bypass the error since we might process different sizes, \n            # but we scale it to expected size for the template matching to work\n            pass\n\n        markers = np.full((17, 2), -1)\n\n        # Find 13 template-based markers\n        for j in range(len(self._templates)):\n            if self._template_points[j] is not None:\n                t = self._template_positions[j][0]-100\n                l = max(self._template_positions[j][1]-100, 0)\n                \n                # Safety check to ensure search range doesn't go out of bounds\n                bottom_bound = min(ima.shape[0], self._template_positions[j][0]+100+self._template_sizes[j][0])\n                right_bound = min(ima.shape[1], self._template_positions[j][1]+250+self._template_sizes[j][0])\n                \n                search_range = ima[t:bottom_bound, l:right_bound]\n                \n                if search_range.shape[0] < self._templates[j].shape[0] or search_range.shape[1] < self._templates[j].shape[1]:\n                    continue # Skip if the image is too small for the template\n                \n                res = cv2.matchTemplate(search_range, self._templates[j], cv2.TM_CCOEFF)\n                min_val, max_val, min_loc, max_loc = cv2.minMaxLoc(res)\n    \n                top_left = max_loc\n                markers[j] = np.array((t + top_left[1] + self._template_points[j][0],\n                                       l + top_left[0] + self._template_points[j][1]))\n\n        # Guess the ends of the first three lines\n        for i in range(3):\n            if markers[5 * i + 3][0] != -1 and markers[5 * i + 2][0] != -1:\n                m = markers[5 * i + 3] * 2 - markers[5 * i + 2]\n                markers[5 * i + 4] = m\n\n        # Guess the end of the fourth line\n        if markers[14][0] != -1 and markers[9][0] != -1:\n            markers[16] = ((markers[14] * (284 + 260) - markers[9] * 260) / 284).astype(int)\n\n        return markers\n        \n    @staticmethod\n    def lead_info(lead):\n        \"\"\"Specify which markers mark the begin and the end of a lead.\"\"\"\n        begin, end = {\n            'I': (0, 1),\n            'II-subset': (5, 6),\n            'III': (10, 11),\n            'aVR': (1, 2),\n            'aVL': (6, 7),\n            'aVF': (11, 12),\n            'V1': (2, 3),\n            'V2': (7, 8),\n            'V3': (12, 13),\n            'V4': (3, 4),\n            'V5': (8, 9),\n            'V6': (13, 14),\n            'II': (15, 16), # The long rhythm strip at the bottom\n        }[lead]\n        \n        # We add a fallback for the long II lead to map it to the 3x4 grid for our UNet\n        row_idx = begin // 5\n        if lead == 'II':\n            row_idx = 3 # 4th row for rhythm strip\n            \n        return row_idx, begin, end\n\n# --- INSTANTIATE THE OBJECT HERE ---\nprint(\"Initializing MarkerFinder...\")\nmf = MarkerFinder(show_templates=False)\nprint(\"MarkerFinder ready!\")","metadata":{"_uuid":"876ef435-2076-45c0-ae89-4f8858acb7de","_cell_guid":"4e2fb01f-6940-425d-b003-810cb13645fa","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:29:54.138574Z","iopub.execute_input":"2026-04-22T21:29:54.139435Z","iopub.status.idle":"2026-04-22T21:29:54.415584Z","shell.execute_reply.started":"2026-04-22T21:29:54.139396Z","shell.execute_reply":"2026-04-22T21:29:54.414470Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ecg lead dataset class","metadata":{"_uuid":"07f0a717-7082-448b-8aec-81c0aeec19a3","_cell_guid":"7df20239-0059-4703-9bb6-566b21982edc","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"\n\n\n\n\nimport os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nclass ECGLeadDataset(Dataset):\n    def __init__(self, metadata_df, img_dir, label_dir, marker_finder, transforms=None):\n        self.metadata = metadata_df\n        self.img_dir = img_dir\n        self.label_dir = label_dir\n        self.mf = marker_finder # Pass the instantiated MarkerFinder object!\n        self.transforms = transforms\n        \n        self.lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n\n    def __len__(self):\n        return len(self.metadata) * 12\n\n    def __getitem__(self, idx):\n        img_idx = idx // 12\n        lead_idx = idx % 12\n        lead_name = self.lead_names[lead_idx]\n        \n        row = self.metadata.iloc[img_idx]\n        record_id = str(row['id'])\n        \n        # 1. Load Full Image\n        img_path = os.path.join(self.img_dir, record_id, f\"{record_id}-0001.png\")\n        image = cv2.imread(img_path)\n        image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        # 2. Use MarkerFinder for PERFECT Cropping\n        markers = self.mf.find_markers(image)\n        \n        # --- THE FIX: Unpack row_idx correctly ---\n        row_idx, begin_idx, end_idx = self.mf.lead_info(lead_name)\n        \n        # Get exact pixel coordinates for the horizontal bounds\n        left_x = markers[begin_idx][1]\n        right_x = markers[end_idx][1]\n\n        # Calculate vertical bounds dynamically\n        padding_top = 130\n        top_y = markers[begin_idx][0] - padding_top \n        \n        if row_idx < 2:\n            # For rows 0 and 1, the bottom is bounded by the marker of the NEXT row down\n            next_row_marker_idx = begin_idx + 5 \n            bottom_y = markers[next_row_marker_idx][0] - padding_top\n            \n        elif row_idx == 2:\n            # For row 2, the next row is the rhythm strip. \n            # The top of the rhythm strip is universally defined by marker 15.\n            bottom_y = markers[15][0] - padding_top\n            \n        else:\n            # For the bottom rhythm strip (row 3), we use a fixed height from the top marker\n            bottom_y = top_y + 350\n            \n        # Ensure we don't crop outside image boundaries\n        top_y = max(0, top_y)\n        bottom_y = min(image_rgb.shape[0], bottom_y)\n        left_x = max(0, left_x)\n        right_x = min(image_rgb.shape[1], right_x)\n\n        # Crop! No white margins, no overlapping leads.\n        crop = image_rgb[top_y:bottom_y, left_x:right_x]\n\n        \n        \n        # Fallback if the crop fails (e.g., markers not found or weird coordinates)\n        if crop.size == 0 or crop.shape[0] == 0 or crop.shape[1] == 0:\n             crop = np.zeros((300, 500, 3), dtype=np.uint8) \n\n        col_w = crop.shape[1]\n        row_h = crop.shape[0]\n\n        # 3. Generate Mask from CSV Ground Truth\n        csv_path = os.path.join(self.label_dir, record_id, f\"{record_id}.csv\")\n        labels = pd.read_csv(csv_path)\n        \n        # Strict typing to catch weird strings/NaNs\n        signal = pd.to_numeric(labels[lead_name], errors='coerce').values\n\n        # THE HORIZONTAL SQUISH FIX Strip away Kaggle's NaN padding to get the actual 2.5s of recorded data\n        signal_clean = signal[~np.isnan(signal)]\n\n        if len(signal_clean) < 2:\n            signal_clean = np.zeros(col_w)\n\n\n        x_coords = np.linspace(0, len(signal_clean)-1, col_w).astype(int)\n        signal_resampled = signal_clean[x_coords]\n        # ---------------------------------\n        \n        marker_y_in_crop = markers[begin_idx][0] - top_y\n        baseline_y = marker_y_in_crop \n        \n        pixels_per_mv = 80 \n        \n        # Now, regardless of how you pad or crop the image, the red line is glued \n        # 120 pixels below the marker, right where the true black ink sits!\n        y_coords = baseline_y - (signal_resampled * pixels_per_mv)\n        \n        # 3. Create an array of (x, y) points\n        x_coords = np.arange(col_w)\n        points = np.column_stack((x_coords, y_coords)).astype(np.int32)\n        points = points.reshape((-1, 1, 2))\n        \n        # 4. Draw a continuous, connected line!\n        mask = np.zeros((row_h, col_w), dtype=np.float32)\n        \n        line_thickness = 3 # A thickness of 3 is perfectly sufficient if aligned correctly.\n        cv2.polylines(mask, [points], isClosed=False, color=1.0, thickness=line_thickness)\n       \n        # -----------------------------------\n        \n        # 4. Augment \n        if self.transforms:\n            augmented = self.transforms(image=crop, mask=mask)\n            img_tensor = augmented['image']\n            mask_tensor = augmented['mask'].unsqueeze(0)\n        else:\n            img_tensor = torch.from_numpy(crop).permute(2,0,1).float()\n            mask_tensor = torch.from_numpy(mask).unsqueeze(0).float()\n\n        return img_tensor, mask_tensor","metadata":{"_uuid":"9476221e-fa51-4e72-9d82-e8bff4b4cd04","_cell_guid":"bc55bf94-d3f5-4ce3-a2ab-bf02ba160613","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:29:54.416703Z","iopub.execute_input":"2026-04-22T21:29:54.417050Z","iopub.status.idle":"2026-04-22T21:30:29.128359Z","shell.execute_reply.started":"2026-04-22T21:29:54.417019Z","shell.execute_reply":"2026-04-22T21:30:29.127127Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### data loader","metadata":{"_uuid":"a1071e0d-8aa0-46de-b328-5f694e981f83","_cell_guid":"756d719d-1cb2-4b49-a40e-0477b1fd5314","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import pandas as pd\n\n\ntrain_df = pd.read_csv('/kaggle/input/competitions/physionet-ecg-image-digitization/train.csv')\n\n\n\n\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\n# 1. Split your raw dataframe (80% train, 20% val)\ntrain_split_df, val_split_df = train_test_split(train_df, test_size=0.2, random_state=42)\ntrain_split_df = train_split_df.reset_index(drop=True)\nval_split_df = val_split_df.reset_index(drop=True)\n\n# 2. Define transforms for the FULL PAGE\ntrain_transforms = A.Compose([\n    A.Resize(256, 1024), # Resize the perfect crop to model input size\n    A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n    A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\nval_transforms = A.Compose([\n    A.Resize(256, 1024),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n# 3. Initialize the Datasets WITH the MarkerFinder\ntrain_dataset = ECGLeadDataset(\n    metadata_df=train_split_df, \n    img_dir='/kaggle/input/competitions/physionet-ecg-image-digitization/train', \n    label_dir='/kaggle/input/competitions/physionet-ecg-image-digitization/train', \n    marker_finder=mf, # <--- The new requirement!\n    transforms=train_transforms\n)\n\nval_dataset = ECGLeadDataset(\n    metadata_df=val_split_df, \n    img_dir='/kaggle/input/competitions/physionet-ecg-image-digitization/train', \n    label_dir='/kaggle/input/competitions/physionet-ecg-image-digitization/train',\n    marker_finder=mf, # <--- The new requirement!\n    transforms=val_transforms\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True) \nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False)","metadata":{"_uuid":"e761a2b8-4251-4d96-a6c5-55625eedeb9a","_cell_guid":"c25a0898-d12d-4465-ad7c-446fb91795af","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:29.130611Z","iopub.execute_input":"2026-04-22T21:30:29.130957Z","iopub.status.idle":"2026-04-22T21:30:29.159820Z","shell.execute_reply.started":"2026-04-22T21:30:29.130925Z","shell.execute_reply":"2026-04-22T21:30:29.158714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# val_transforms = A.Compose([\n#     A.Resize(256, 512),\n#     A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n#     ToTensorV2()\n# ])\n\n# val_dataset = ECGCropDataset(\n#     metadata_df=val_split_df, \n#     img_dir='/kaggle/input/competitions/physionet-ecg-image-digitization/train', \n#     mask_dir='/kaggle/working/masks/', \n#     transforms=val_transforms # Clean transforms\n# )\n# val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False) # No need to shuffle validation data","metadata":{"_uuid":"a290bca1-1173-46a6-a378-d4827f7f6410","_cell_guid":"c9f7c2ff-5b03-4a91-9e0a-7c69b63aafd7","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:29.163384Z","iopub.execute_input":"2026-04-22T21:30:29.163674Z","iopub.status.idle":"2026-04-22T21:30:29.168506Z","shell.execute_reply.started":"2026-04-22T21:30:29.163648Z","shell.execute_reply":"2026-04-22T21:30:29.167180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_df.head())\nprint(train_df.columns)","metadata":{"_uuid":"e2da6ede-3eeb-496c-bf72-35a91cf7f7ed","_cell_guid":"4d0bcadf-81cb-4e22-a170-229c5ab137f3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:29.169718Z","iopub.execute_input":"2026-04-22T21:30:29.170136Z","iopub.status.idle":"2026-04-22T21:30:29.193036Z","shell.execute_reply.started":"2026-04-22T21:30:29.170105Z","shell.execute_reply":"2026-04-22T21:30:29.191976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install segmentation-models-pytorch\n\n!pip install --no-index --find-links /kaggle/input/notebooks/ravnoorsingh101/smp-offline-wheels/smp-wheels segmentation-models-pytorch","metadata":{"_uuid":"953d1236-231a-4c97-80cb-8c8dab9c0bfd","_cell_guid":"85c7f373-f927-433a-8343-174ffc2ae226","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:29.194316Z","iopub.execute_input":"2026-04-22T21:30:29.194686Z","iopub.status.idle":"2026-04-22T21:30:33.383288Z","shell.execute_reply.started":"2026-04-22T21:30:29.194656Z","shell.execute_reply":"2026-04-22T21:30:33.382258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport torch\nimport segmentation_models_pytorch as smp\n\n# Initialize the Attention U-Net\nmodel = smp.Unet(\n    encoder_name=\"efficientnet-b0\", # Lightweight, highly accurate backbone\n    encoder_weights=None,     # Transfer learning from general vision\n    in_channels=3,                  # RGB augmented crops\n    classes=1,                      # Binary mask (0=Background, 1=Signal)\n    activation=None,                # We will apply sigmoid in the loss function\n    decoder_attention_type=\"scse\"   # The crucial Attention mechanism\n)\n\nformatted_weights_path = '/kaggle/input/notebooks/ravnoorsingh101/pre-trained-model-efficientnet-b0/smp_weights/smp_efficientnet_b0_imagenet.pth'\n\nmodel.encoder.load_state_dict(torch.load(formatted_weights_path), strict=False)\n\n# Move model to GPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)","metadata":{"_uuid":"0dc7ab85-3ae8-4919-88dc-6676ed7ceca2","_cell_guid":"6628e945-dc59-4d24-a208-5cc0683c954f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:33.384795Z","iopub.execute_input":"2026-04-22T21:30:33.385260Z","iopub.status.idle":"2026-04-22T21:30:39.505686Z","shell.execute_reply.started":"2026-04-22T21:30:33.385220Z","shell.execute_reply":"2026-04-22T21:30:39.504669Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### train model function","metadata":{"_uuid":"52ccf899-c049-4d64-87d0-96151cdaad34","_cell_guid":"be7dd583-2ee5-4b62-bdae-ac1ebe671842","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nclass ECGContinuityLoss(nn.Module):\n    def __init__(self, bce_weight=0.5, dice_weight=0.5, tv_weight=0.01):\n        super(ECGContinuityLoss, self).__init__()\n        self.bce_weight = bce_weight\n        self.dice_weight = dice_weight\n        self.tv_weight = tv_weight\n        self.bce = nn.BCEWithLogitsLoss()\n\n    def forward(self, inputs, targets):\n        # 1. Binary Cross Entropy (Pixel-wise accuracy)\n        bce_loss = self.bce(inputs, targets)\n        \n        # Apply sigmoid to get probabilities for Dice and TV\n        probs = torch.sigmoid(inputs)\n        \n        # 2. Dice Loss (Overlap accuracy)\n        smooth = 1e-6\n        intersection = (probs * targets).sum(dim=(2, 3))\n        union = probs.sum(dim=(2, 3)) + targets.sum(dim=(2, 3))\n        dice_loss = 1 - ((2. * intersection + smooth) / (union + smooth)).mean()\n        \n        # 3. Total Variation (Continuity Penalty)\n        # Penalizes differences between adjacent pixels in the x-direction (time)\n        tv_loss = torch.mean(torch.abs(probs[:, :, :, :-1] - probs[:, :, :, 1:]))\n        \n        # Combine them\n        total_loss = (self.bce_weight * bce_loss) + \\\n                     (self.dice_weight * dice_loss) + \\\n                     (self.tv_weight * tv_loss)\n                     \n        return total_loss\n\n# Initialize your custom loss and optimizer\ncriterion = ECGContinuityLoss(bce_weight=0.4, dice_weight=0.4, tv_weight=0.00)\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)","metadata":{"_uuid":"60e08e37-cd5e-4d47-9fc5-f045c7d693e6","_cell_guid":"8a9fa0dc-6393-48e9-93a2-802e17fa2dcf","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:39.506725Z","iopub.execute_input":"2026-04-22T21:30:39.507320Z","iopub.status.idle":"2026-04-22T21:30:39.519552Z","shell.execute_reply.started":"2026-04-22T21:30:39.507283Z","shell.execute_reply":"2026-04-22T21:30:39.518136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom tqdm import tqdm # For progress bars in Kaggle\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, num_epochs=15, patience=2):\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Scheduler: Reduces learning rate if validation loss stops improving\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)\n    \n    best_val_loss = float('inf')\n    epochs_without_improvement = 0\n    \n    # History tracking for plotting later\n    history = {'train_loss': [], 'val_loss': []}\n\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(\"-\" * 20)\n        \n        ### 1. TRAINING PHASE ###\n        model.train()\n        train_loss = 0.0\n        \n        # tqdm adds a nice progress bar in the Kaggle console\n        train_bar = tqdm(train_loader, desc=\"Training\")\n        \n        for images, masks in train_bar:\n            images = images.to(device)\n            masks = masks.to(device)\n            \n            # Forward pass\n            optimizer.zero_grad()\n            outputs = model(images)\n            \n            # Calculate custom Continuity Loss\n            loss = criterion(outputs, masks)\n            \n            # Backward pass & optimize\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item() * images.size(0)\n            train_bar.set_postfix({'loss': loss.item()})\n            \n        epoch_train_loss = train_loss / len(train_loader.dataset)\n        history['train_loss'].append(epoch_train_loss)\n        \n        ### 2. VALIDATION PHASE ###\n        model.eval()\n        val_loss = 0.0\n        \n        val_bar = tqdm(val_loader, desc=\"Validation\")\n        \n        with torch.no_grad(): # Disable gradient calculation for speed and memory\n            for images, masks in val_bar:\n                images = images.to(device)\n                masks = masks.to(device)\n                \n                outputs = model(images)\n                loss = criterion(outputs, masks)\n                \n                val_loss += loss.item() * images.size(0)\n                val_bar.set_postfix({'loss': loss.item()})\n                \n        epoch_val_loss = val_loss / len(val_loader.dataset)\n        history['val_loss'].append(epoch_val_loss)\n        \n        print(f\"Train Loss: {epoch_train_loss:.4f} | Val Loss: {epoch_val_loss:.4f}\")\n        \n        # Update learning rate scheduler\n        scheduler.step(epoch_val_loss)\n        \n        ### 3. EARLY STOPPING & CHECKPOINTING ###\n        if epoch_val_loss < best_val_loss:\n            best_val_loss = epoch_val_loss\n            epochs_without_improvement = 0\n            # Save the best model weights\n            torch.save(model.state_dict(), '/kaggle/working/best_attention_unet.pth')\n            print(\">>> Saved new best model!\")\n        else:\n            epochs_without_improvement += 1\n            print(f\"No improvement in validation loss for {epochs_without_improvement} epoch(s).\")\n            \n            if epochs_without_improvement >= patience:\n                print(\"Early stopping triggered. Halting training.\")\n                break\n                \n    return history\n\n# --- Execution ---\n# Assuming you have split your data and created train_loader and val_loader\n#history = train_model(model, train_loader, val_loader, criterion, optimizer, num_epochs=8)","metadata":{"_uuid":"fa76773a-62cc-435f-80e7-df74e8ffb33a","_cell_guid":"51121ebf-294a-4990-ae8d-1cfc2d4f25ac","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:39.520963Z","iopub.execute_input":"2026-04-22T21:30:39.521269Z","iopub.status.idle":"2026-04-22T21:30:39.552669Z","shell.execute_reply.started":"2026-04-22T21:30:39.521239Z","shell.execute_reply":"2026-04-22T21:30:39.550976Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{"_uuid":"8def8b79-96e9-4560-9719-6e9668381c5d","_cell_guid":"50fed29b-2fab-439c-92c1-afb561dec0a5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport segmentation_models_pytorch as smp\n\nmodel = smp.Unet(\n    encoder_name=\"efficientnet-b0\", \n    encoder_weights=None,      \n    in_channels=3,                  \n    classes=1,                      \n    activation=None,                \n    decoder_attention_type=\"scse\"   \n)\n\n# 2. Load YOUR trained weights (update the path to match your dataset name)\nweights_path = '/kaggle/input/datasets/ravnoor000/physionet-unet-weights-2/best_attention_unet (4).pth'\nmodel.load_state_dict(torch.load(weights_path, map_location=torch.device('cpu')))\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)","metadata":{"_uuid":"9eb98e30-140a-4d8a-ab58-5fba44d34b93","_cell_guid":"45f60a9f-78ed-4dbf-9c4e-03d1b9e42306","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:39.554932Z","iopub.execute_input":"2026-04-22T21:30:39.555538Z","iopub.status.idle":"2026-04-22T21:30:40.374820Z","shell.execute_reply.started":"2026-04-22T21:30:39.555501Z","shell.execute_reply":"2026-04-22T21:30:40.373724Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Mask and resutls","metadata":{"_uuid":"da60d378-0a06-4e51-bc35-4eae233ac2e5","_cell_guid":"6d20342c-1961-43ca-ba6a-b0a756783394","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\n\ndef get_visual_mask(model_output, threshold=0.5):\n    # Convert logits to probabilities (0 to 1)\n    probs = torch.sigmoid(model_output)\n    \n    # Binarize the output: pixels > 50% sure are 1 (white), else 0 (black)\n    binary_mask = (probs > threshold).float()\n    \n    # Move to CPU and convert to numpy for visualization\n    mask_np = binary_mask.squeeze().cpu().detach().numpy()\n    return mask_np\n\n# Example usage during evaluation:\n# plt.imshow(mask_np, cmap='gray')\n# plt.title(\"Attention U-Net Predicted Mask\")\n# plt.show()\n\n\nimport numpy as np\nfrom scipy.signal import savgol_filter\n\ndef extract_1d_signal_from_mask(mask_np):\n    \"\"\"\n    Extracts a 1D signal array from a 2D binary mask.\n    mask_np shape: (Height, Width)\n    \"\"\"\n    height, width = mask_np.shape\n    signal_1d = np.zeros(width)\n    prev_val = height // 2\n    \n    for x in range(width):\n        column = mask_np[:, x]\n        \n        # Find the indices (y-coordinates) where the mask is 1\n        y_indices = np.where(column > 0)[0]\n        \n        if len(y_indices) > 0:\n            # Center of mass: average the y-coordinates of the white pixels\n            # signal_1d[x] = np.mean(y_indices)\n\n            signal_1d[x] = y_indices[0]\n            prev_val = signal_1d[x]\n        else:\n            # If the model predicted a gap (no white pixels in this column),\n            # we will temporarily set it to NaN to interpolate later.\n            signal_1d[x] = np.nan\n            \n    # Interpolate to fill any NaN gaps (where the U-Net missed the line)\n    nan_mask = np.isnan(signal_1d)\n    if np.all(nan_mask):\n        # Edge Case: The model predicted a completely black mask.\n        # We cannot interpolate. Return a flat line in the exact center of the crop.\n        return np.full(width, height / 2.0)\n        \n    elif np.any(nan_mask):\n        # Standard Case: There are some gaps, but we have enough points to interpolate.\n        signal_1d[nan_mask] = np.interp(\n            np.flatnonzero(nan_mask), \n            np.flatnonzero(~nan_mask), \n            signal_1d[~nan_mask]\n        )\n        \n    return signal_1d\n\n\ndef pixels_to_millivolts(signal_1d, pixels_per_mV, actual_baseline_y):\n    \"\"\"\n    Converts pixel coordinates to voltage.\n    Note: In images, y=0 is the TOP of the image. \n    In signals, positive voltage goes UP. We must invert the y-axis.\n    \"\"\"\n    # 1. Find the baseline (assume the median of the signal is 0 mV)\n    baseline_pixel = actual_baseline_y\n    \n    # 2. Shift the signal so the baseline is at 0, and INVERT the axis\n    # (so a peak pointing to the top of the image is a positive voltage)\n    signal_shifted = baseline_pixel - signal_1d \n    \n    # 3. Convert to millivolts\n    signal_mv = signal_shifted / pixels_per_mV\n    \n    return signal_mv\n\ndef smooth_ecg_signal(signal_mv):\n    \"\"\"\n    Applies a Savitzky-Golay filter to remove jagged pixelation noise \n    while preserving the sharp peaks of the QRS complex.\n    \"\"\"\n    # window_length must be odd, polyorder is the polynomial degree\n    smoothed_signal = savgol_filter(signal_mv, window_length=11, polyorder=3)\n    return smoothed_signal\n\n\ndef get_dynamic_baseline(crop_gray):\n    \"\"\"\n    Uses morphological operations to isolate horizontal lines (baseline/grid).\n    \"\"\"\n    # 1. Threshold to get black ink\n    _, binary = cv2.threshold(crop_gray, 150, 255, cv2.THRESH_BINARY_INV)\n    \n    # 2. Morphological Kernel: Long horizontal line\n    kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (25, 1))\n    \n    # 3. Erode then Dilate (Open)\n    # This removes the sharp vertical QRS complexes and leaves the baseline\n    baseline_map = cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel)\n    \n    # 4. Find the Y-coordinate of the densest horizontal line\n    horizontal_density = np.sum(baseline_map, axis=1)\n    detected_baseline_y = np.argmax(horizontal_density)\n    \n    return detected_baseline_y\n\n\ndef estimate_pixels_per_mv(crop_gray):\n    \"\"\"\n    Estimate pixels per mV using vertical grid spacing.\n    \"\"\"\n    # Detect edges\n    edges = cv2.Canny(crop_gray, 50, 150)\n\n    # Hough lines to detect horizontal grid lines\n    lines = cv2.HoughLinesP(edges, 1, np.pi/180, threshold=100,\n                            minLineLength=50, maxLineGap=10)\n\n    if lines is None:\n        return 80.0  # fallback\n\n    y_coords = []\n\n    for line in lines:\n        x1, y1, x2, y2 = line[0]\n        # Keep only horizontal lines\n        if abs(y1 - y2) < 3:\n            y_coords.append(y1)\n\n    if len(y_coords) < 2:\n        return 80.0\n\n    y_coords = np.sort(y_coords)\n\n    # Compute distances between grid lines\n    diffs = np.diff(y_coords)\n\n    # Remove noise\n    diffs = diffs[(diffs > 5) & (diffs < 50)]\n\n    if len(diffs) == 0:\n        return 80.0\n\n    median_spacing = np.median(diffs)\n\n    # Small box = 1 mm → 10 mm = 1 mV\n    pixels_per_mv = median_spacing * 10\n\n    return pixels_per_mv","metadata":{"_uuid":"aad1a754-fd66-416c-baf3-e28a0d63611f","_cell_guid":"f19160ca-762e-427c-b8f8-6a9e0ea98326","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:40.376038Z","iopub.execute_input":"2026-04-22T21:30:40.376387Z","iopub.status.idle":"2026-04-22T21:30:40.393704Z","shell.execute_reply.started":"2026-04-22T21:30:40.376352Z","shell.execute_reply":"2026-04-22T21:30:40.392447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visual","metadata":{"_uuid":"54c2af7a-5b21-4451-900a-d407704f8588","_cell_guid":"232651d8-822a-4926-b58d-1fd679ac5fea","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# 1. Set model to evaluation mode (turns off dropout, batchnorm updates, etc.)\nmodel.eval()\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 2. Grab a single batch of images from your DataLoader\n# (Using train_loader here just to test that the pipeline runs)\nimages, true_masks = next(iter(train_loader))\nimages = images.to(device)\n\n# 3. Pass the images through the model to get the predictions\nwith torch.no_grad():\n    batch_outputs = model(images)\n\n# 4. Isolate the FIRST image in the batch to process\n# THIS is the missing variable!\nmodel_output = batch_outputs[0] \noriginal_image_tensor = images[0]\n\n# --- NOW WE RUN THE POST-PROCESSING PIPELINE ---\n\n# Step A: Get the 2D binary mask\nraw_mask = get_visual_mask(model_output, threshold=0.5)\n\n# Step B: Extract the 1D pixel heights (Center of Mass)\npixel_signal = extract_1d_signal_from_mask(raw_mask)\n\n# Step C: Convert pixels to Voltage \n# (Note: 50.0 is a placeholder. You'll need to check the Kaggle data to see exactly how many pixels = 1mV)\nactual_baseline_y_vis = 130.0\nvoltage_signal = pixels_to_millivolts(pixel_signal, pixels_per_mV=80.0, actual_baseline_y= actual_baseline_y_vis)\n\n# Step D: Smooth the signal with Savitzky-Golay\n\nfinal_signal = smooth_ecg_signal(voltage_signal)\n\n\n# --- VISUALIZATION: Let's see if it worked! ---\n\nfig, axs = plt.subplots(3, 1, figsize=(12, 10))\n\n# Plot 1: The Original Cropped Image\n# Convert tensor back to numpy and un-normalize for viewing\nimg_view = original_image_tensor.cpu().permute(1, 2, 0).numpy()\n# Undo the ImageNet normalization so it looks normal\nmean = np.array([0.485, 0.456, 0.406])\nstd = np.array([0.229, 0.224, 0.225])\nimg_view = std * img_view + mean\nimg_view = np.clip(img_view, 0, 1)\n\naxs[0].imshow(img_view)\naxs[0].set_title(\"1. Original Cropped Lead (Input)\")\naxs[0].axis('off')\n\n# Plot 2: The U-Net Binary Mask\naxs[1].imshow(raw_mask, cmap='gray')\naxs[1].set_title(\"2. Attention U-Net Prediction (2D Mask)\")\naxs[1].axis('off')\n\n# Plot 3: The Final Digitized 1D Signal\naxs[2].plot(final_signal, color='red')\naxs[2].set_title(\"3. Final Digitized Signal (1D Voltage)\")\naxs[2].set_ylabel(\"Millivolts (mV)\")\naxs[2].set_xlabel(\"Time (pixels)\")\naxs[2].grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"09b1a977-7c8c-430e-bace-7b478f360aee","_cell_guid":"a11871aa-0bfd-4e48-ae5e-4ea7456d7e58","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:40.395421Z","iopub.execute_input":"2026-04-22T21:30:40.396176Z","iopub.status.idle":"2026-04-22T21:30:50.612637Z","shell.execute_reply.started":"2026-04-22T21:30:40.396112Z","shell.execute_reply":"2026-04-22T21:30:50.611527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Grab one batch of TRAINING data\nimages, true_masks = next(iter(train_loader))\n\n# Un-normalize the image for viewing\nimg_view = images[0].permute(1, 2, 0).numpy()\nmean = np.array([0.485, 0.456, 0.406])\nstd = np.array([0.229, 0.224, 0.225])\nimg_view = std * img_view + mean\nimg_view = np.clip(img_view, 0, 1)\n\n# Get the true mask\nt_mask = true_masks[0].squeeze().numpy()\n\nplt.figure(figsize=(12, 6))\nplt.imshow(img_view)\n# Overlay the ground truth mask in RED\nplt.imshow(t_mask, cmap='Reds', alpha=0.5)\nplt.title(\"SANITY CHECK: Does the red mask exactly cover the black ink?\")\nplt.show()","metadata":{"_uuid":"a6f08c31-be7e-4d56-af5e-4d46d5823486","_cell_guid":"7ed4fc6a-fd37-47ac-a4ca-f8b7d6eedb83","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:50.614036Z","iopub.execute_input":"2026-04-22T21:30:50.614733Z","iopub.status.idle":"2026-04-22T21:30:53.929835Z","shell.execute_reply.started":"2026-04-22T21:30:50.614700Z","shell.execute_reply":"2026-04-22T21:30:53.928849Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Mask quality","metadata":{"_uuid":"14c2a8ec-c361-4811-acf6-1932c5333493","_cell_guid":"6b1ca92e-0232-4ed1-b50a-975fe857572c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"plt.figure(figsize=(10,4))\nplt.imshow(raw_mask, cmap='gray')\nplt.title(\"Predicted Mask\")\nplt.show()\n\nprint(\"Mask density:\", np.mean(raw_mask))","metadata":{"_uuid":"74997f5b-5910-41f2-9e58-0158129458b5","_cell_guid":"119393f0-6f7b-4c93-9d2a-5e62033e578e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:53.931057Z","iopub.execute_input":"2026-04-22T21:30:53.931399Z","iopub.status.idle":"2026-04-22T21:30:54.114820Z","shell.execute_reply.started":"2026-04-22T21:30:53.931360Z","shell.execute_reply":"2026-04-22T21:30:54.113952Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"plt.figure(figsize=(10,3))\nplt.plot(pixel_signal)\nplt.title(\"Pixel Signal (before voltage)\")\nplt.show()","metadata":{"_uuid":"8ec2cb31-8457-4174-9fdd-5dae3679d25b","_cell_guid":"2c31ddb7-aece-4040-b9c4-98b73d0466c6","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# 1. Get raw array from model\nraw_mask = get_visual_mask(model_output)\n\n# 2. Extract pixel heights\npixel_signal = extract_1d_signal_from_mask(raw_mask)\n\npixel_signal = np.nan_to_num(pixel_signal, nan=np.median(pixel_signal))\n\nimg_np = original_image_tensor.cpu().permute(1, 2, 0).numpy()\nmean = np.array([0.485, 0.456, 0.406])\nstd = np.array([0.229, 0.224, 0.225])\nimg_np = std * img_np + mean\nimg_np = np.clip(img_np, 0, 1)\n\n# Second, convert from [0.0, 1.0] float to [0, 255] uint8 (which OpenCV requires)\ncrop_uint8 = (img_np * 255).astype(np.uint8)\n\n# Third, convert to grayscale\ncrop_gray = cv2.cvtColor(crop_uint8, cv2.COLOR_RGB2GRAY)\n# 🔥 Estimate dynamically\npixels_per_mv = estimate_pixels_per_mv(crop_gray)\n\n# Safety fallback (VERY IMPORTANT)\nif pixels_per_mv < 30 or pixels_per_mv > 150:\n    pixels_per_mv = 100.0\n\nactual_baseline_y_vis = 130.0\n\n# 3. Convert to Voltage (You'll need to check the Kaggle metadata for the exact pixels_per_mV)\n# For example, if the grid resolution implies 50 pixels = 1 mV:\nvoltage_signal = pixels_to_millivolts(pixel_signal, pixels_per_mV=pixels_per_mv, actual_baseline_y = actual_baseline_y_vis)\n\n# 4. Smooth for final submission\nfinal_submission_signal = smooth_ecg_signal(voltage_signal)","metadata":{"_uuid":"d6cae417-35a3-4271-8f41-09d22404883d","_cell_guid":"3c17fa7d-8bbd-4dee-a011-1ef6756204a4","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:36:55.422060Z","iopub.execute_input":"2026-04-22T21:36:55.422455Z","iopub.status.idle":"2026-04-22T21:36:55.468599Z","shell.execute_reply.started":"2026-04-22T21:36:55.422420Z","shell.execute_reply":"2026-04-22T21:36:55.467353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- GENERATE FINAL KAGGLE SUBMISSION ---\nimport os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nprint(\"Generating Kaggle submission...\")\n\n# 1. Load the hidden test set metadata\ntest_df = pd.read_csv('/kaggle/input/competitions/physionet-ecg-image-digitization/test.csv')\n\n# 2. Setup model for pure inference\nmodel.eval()\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 3. Use the exact same transforms we used in validation\ntest_transforms = A.Compose([\n    A.Resize(256, 1024),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n# --- IMPROVED SUBMISSION LOOP ---\nsubmission_data = []\nold_id = None\n\n# Standard 3x4 Layout lookup\nlayout_map = {\n    'I': (0,0), 'aVR': (0,1), 'V1': (0,2), 'V4': (0,3),\n    'II': (1,0), 'aVL': (1,1), 'V2': (1,2), 'V5': (1,3),\n    'III': (2,0), 'aVF': (2,1), 'V3': (2,2), 'V6': (2,3)\n}\n\nfor idx, row in tqdm(test_df.iterrows(), total=len(test_df)):\n    record_id = str(row['id'])\n    lead_name = row['lead']\n    \n    if record_id != old_id:\n        img_path = os.path.join('/kaggle/input/competitions/physionet-ecg-image-digitization/test', f\"{record_id}.png\")\n        full_img = cv2.imread(img_path)\n\n        if full_img is None:\n            print(f\"Warning: Could not find image at {img_path}\")\n            full_img = np.zeros((1652, 2200, 3), dtype=np.uint8)\n        \n        full_img_rgb = cv2.cvtColor(full_img, cv2.COLOR_BGR2RGB)\n        h, w, _ = full_img.shape\n        old_id = record_id\n\n    # 1. GET LEAD CROP (Step 1)\n    markers = mf.find_markers(full_img)\n\n    row_idx, begin_idx, end_idx = mf.lead_info(lead_name)\n    \n    left_x = markers[begin_idx][1]\n    right_x = markers[end_idx][1]\n    \n    padding_top = 130\n    top_y = markers[begin_idx][0] - padding_top \n    \n    if row_idx < 2:\n        bottom_y = markers[begin_idx + 5][0] - padding_top\n    elif row_idx == 2:\n        bottom_y = markers[15][0] - padding_top\n    else:\n        bottom_y = top_y + 350\n    \n    top_y = max(0, top_y)\n    bottom_y = min(full_img.shape[0], bottom_y)\n    left_x = max(0, left_x)\n    right_x = min(full_img.shape[1], right_x)\n    \n    crop = full_img_rgb[top_y:bottom_y, left_x:right_x]\n    \n    # 2. DETECT BASELINE (Step 2)\n    #crop_gray = cv2.cvtColor(crop, cv2.COLOR_RGB2GRAY)\n    #actual_baseline_y = get_dynamic_baseline(crop_gray)\n    \n    actual_baseline_y = markers[begin_idx][0] - top_y\n    \n    # 3. U-NET INFERENCE (On High-Res Patch)\n    # Resize to model input size (e.g., 256x512) but keep aspect ratio\n    input_tensor = test_transforms(image=crop)['image'].unsqueeze(0).to(device)\n    with torch.no_grad():\n        mask_logits = model(input_tensor)[0]\n\n    row_h = crop.shape[0]\n    col_w = crop.shape[1]\n    # 4. SIGNAL EXTRACTION\n    mask = get_visual_mask(mask_logits) # Your existing function\n    pixel_signal = extract_1d_signal_from_mask(mask) # Your existing function\n    \n    # Re-scale back to crop height for voltage conversion\n    pixel_signal_scaled = pixel_signal * (row_h / mask.shape[0])\n    \n    # Use detected baseline instead of np.median\n    # Voltage = (Baseline - PixelY) / Scale\n    # 80 is a better estimate for these scans than 50\n    # signal_mv = (actual_baseline_y - pixel_signal_scaled) / 80.0\n\n    pixels_per_mv = estimate_pixels_per_mv(crop)\n    #pixels_per_mv = 150\n\n    signal_mv = (actual_baseline_y - pixel_signal_scaled) / pixels_per_mv\n    # signal_mv = signal_mv * 1.5\n    # signal_mv = np.clip(signal_mv, -3, 3)\n    # Final cleanup\n    #final_signal = smooth_ecg_signal(signal_mv)\n\n    final_signal = signal_mv\n    \n    # Interpolate to Kaggle's required length\n    final_interpolated = np.interp(\n        np.linspace(0, 1, int(row['number_of_rows'])),\n        np.linspace(0, 1, len(final_signal)),\n        final_signal\n    )\n\n    # Append to submission list (Standard formatting)\n    for t, val in enumerate(final_interpolated):\n        submission_data.append({'id': f\"{record_id}_{t}_{lead_name}\", 'value': val})\n\nprint(\"NaNs in final:\", np.isnan(final_signal).sum())\nprint(\"Min/Max:\", final_signal.min(), final_signal.max())\n\n# Save to CSV exactly as Kaggle expects\nsubmission_df = pd.DataFrame(submission_data)\nprint(f\"\\nCreated submission with {len(submission_df)} rows.\")\nsubmission_df.to_csv('submission.csv', index=False)\nprint(\"Saved to submission.csv! Ready to submit to the leaderboard.\")","metadata":{"_uuid":"5d2b6cdf-3f84-4bea-9f7a-4a015751d2f5","_cell_guid":"e2d35a62-2d7e-4fea-9901-7d8f50c1d529","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:38:38.109743Z","iopub.execute_input":"2026-04-22T21:38:38.110718Z","iopub.status.idle":"2026-04-22T21:38:58.100340Z","shell.execute_reply.started":"2026-04-22T21:38:38.110680Z","shell.execute_reply":"2026-04-22T21:38:58.099255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(crop)\nplt.title(f\"Lead: {lead_name}\")\nplt.show()","metadata":{"_uuid":"aa54b7be-ea5e-4acd-a9da-fde4cc8dab8a","_cell_guid":"331f6d3b-5cf9-4403-8dd4-1a127a907833","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-04-22T21:30:54.137436Z","iopub.status.idle":"2026-04-22T21:30:54.137786Z","shell.execute_reply.started":"2026-04-22T21:30:54.137603Z","shell.execute_reply":"2026-04-22T21:30:54.137623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import mean_squared_error, mean_absolute_error\n\ndef compare_submissions(ref_csv_path, pred_csv_path, plot_id=None):\n    print(\"Loading submissions for comparison...\")\n    \n    # 1. Load the CSVs\n    df_ref = pd.read_csv(ref_csv_path)\n    df_pred = pd.read_csv(pred_csv_path)\n    \n    # 2. Merge them on the 'id' column to ensure exact alignment\n    df_merged = pd.merge(df_ref, df_pred, on='id', suffixes=('_ref', '_pred'))\n    \n    # Check if merge was successful\n    if len(df_merged) == 0:\n        print(\"Error: No matching IDs found between the two CSVs.\")\n        return\n        \n    print(f\"Successfully aligned {len(df_merged)} rows.\\n\")\n    \n    # 3. Calculate Overall Metrics\n    y_true = df_merged['value_ref'].values\n    y_pred = df_merged['value_pred'].values\n    \n    mse = mean_squared_error(y_true, y_pred)\n    mae = mean_absolute_error(y_true, y_pred)\n    \n    # Correlation checks if the *shape* is right, even if the scale/offset is wrong\n    correlation = np.corrcoef(y_true, y_pred)[0, 1] \n    \n    print(\"--- GLOBAL METRICS ---\")\n    print(f\"Mean Squared Error (MSE):  {mse:.5f}  <-- Lower is better\")\n    print(f\"Mean Absolute Error (MAE): {mae:.5f}  <-- Lower is better\")\n    print(f\"Pearson Correlation:       {correlation:.5f}  <-- Closer to 1.0 is better\")\n    \n    # 4. Optional: Visual Comparison of a specific lead\n    if plot_id:\n        # Filter for a specific record and lead, e.g., '1053922973_I'\n        # The 'id' in the CSV looks like '1053922973_0_I'\n        target_prefix = f\"{plot_id}_\"\n        \n        # We need a regex or string match to pull out all time steps for this lead\n        df_plot = df_merged[df_merged['id'].str.contains(target_prefix, regex=False)]\n        \n        if len(df_plot) > 0:\n            plt.figure(figsize=(15, 4))\n            plt.plot(df_plot['value_ref'].values, label='Reference (Score: 15)', color='green', alpha=0.7, linewidth=2)\n            plt.plot(df_plot['value_pred'].values, label='Your Model', color='red', alpha=0.7, linewidth=2)\n            \n            plt.title(f\"Comparison for Lead: {plot_id}\")\n            plt.xlabel(\"Time Steps\")\n            plt.ylabel(\"Voltage (mV)\")\n            plt.grid(True)\n            plt.legend()\n            plt.show()\n        else:\n            print(f\"Could not find data to plot for ID matching: {plot_id}\")\n\n# --- HOW TO USE IT ---\n# Change these paths to wherever your files are located\nbaseline_path = '/kaggle/input/datasets/ravnoor000/score-23/submission_23.csv'  \n#baseline_path = '/kaggle/input/datasets/ravnoor000/score-23/submission_23.csv'  \nyour_path = '/kaggle/working/submission.csv'\n\n# Run the comparison!\n# To plot a specific lead, pass the record_id and lead name, e.g., '1053922973_.*_I' \n# (The easiest way is just to pass '1053922973' to see the first lead it finds, or '1053922973_.*_I' with regex if you tweak the function)\ncompare_submissions(baseline_path, your_path, plot_id=\"2352854581\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-22T21:39:46.299808Z","iopub.execute_input":"2026-04-22T21:39:46.300165Z","iopub.status.idle":"2026-04-22T21:39:46.768024Z","shell.execute_reply.started":"2026-04-22T21:39:46.300136Z","shell.execute_reply":"2026-04-22T21:39:46.766829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# --- LOAD FILES ---\nsub_good = pd.read_csv('/kaggle/input/datasets/ravnoor000/score-23/submission_23.csv')   # better one\nsub_yours = pd.read_csv('/kaggle/working/submission.csv')  # your file\n\n# --- SORT (VERY IMPORTANT) ---\nsub_good = sub_good.sort_values('id').reset_index(drop=True)\nsub_yours = sub_yours.sort_values('id').reset_index(drop=True)\n\n# --- CHECK ALIGNMENT ---\nassert (sub_good['id'] == sub_yours['id']).all(), \"IDs do not match!\"\n\n# --- EXTRACT VALUES ---\ny_true = sub_good['value'].values\ny_pred = sub_yours['value'].values\n\n# --- METRICS ---\nmae = np.mean(np.abs(y_true - y_pred))\nrmse = np.sqrt(np.mean((y_true - y_pred) ** 2))\ncorr = np.corrcoef(y_true, y_pred)[0, 1]\n\nprint(f\"MAE:  {mae:.5f}\")\nprint(f\"RMSE: {rmse:.5f}\")\nprint(f\"Correlation: {corr:.5f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-22T21:40:01.327623Z","iopub.execute_input":"2026-04-22T21:40:01.327980Z","iopub.status.idle":"2026-04-22T21:40:01.582290Z","shell.execute_reply.started":"2026-04-22T21:40:01.327948Z","shell.execute_reply":"2026-04-22T21:40:01.581188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_random_signals(sub_good, sub_yours, num_samples=5):\n    ids = sub_good['id'].values\n    \n    # group by record (each signal)\n    unique_records = list(set([i.split('_')[0] + '_' + i.split('_')[2] for i in ids]))\n    \n    samples = np.random.choice(unique_records, num_samples, replace=False)\n    \n    for sample in samples:\n        idxs = [i for i, x in enumerate(ids) if sample in x]\n        \n        true_signal = sub_good.iloc[idxs]['value'].values\n        pred_signal = sub_yours.iloc[idxs]['value'].values\n        \n        plt.figure(figsize=(10,3))\n        plt.plot(true_signal, label='Score 23 (Better)', linewidth=2)\n        plt.plot(pred_signal, label='Yours', alpha=0.7)\n        plt.title(f\"Signal Comparison: {sample}\")\n        plt.legend()\n        plt.grid()\n        plt.show()\n\nplot_random_signals(sub_good, sub_yours)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-22T21:27:53.904411Z","iopub.execute_input":"2026-04-22T21:27:53.904763Z","iopub.status.idle":"2026-04-22T21:27:54.644154Z","shell.execute_reply.started":"2026-04-22T21:27:53.904729Z","shell.execute_reply":"2026-04-22T21:27:54.643112Z"}},"outputs":[],"execution_count":null}]}