{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":623160,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":468838,"modelId":484689},{"sourceId":623170,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":468846,"modelId":484698}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/input/open-ecg-digitizer/pytorch/default/1\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pandas as pd\nimport numpy as np\nfrom scipy.signal import medfilt, savgol_filter, butter, filtfilt, resample\nimport cv2\nimport os\nfrom tqdm import tqdm\nfrom PIL import Image\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Import from Open-ECG-Digitizer\nfrom torchvision.io import read_image\nfrom src.model.unet import UNet\nfrom src.model.perspective_detector import PerspectiveDetector\nfrom src.model.cropper import Cropper\nfrom src.model.pixel_size_finder import PixelSizeFinder\nfrom src.model.signal_extractor import SignalExtractor\nfrom src.model.lead_identifier import LeadIdentifier\nimport yaml\n\nprint(\"Initializing...\")\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nLEADS = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n\n# ============================================================================\n# SIMPLIFIED IMAGE PREPROCESSING\n# ============================================================================\nclass ImagePreprocessor:\n    \"\"\"Minimal preprocessing to avoid signal degradation\"\"\"\n    \n    @staticmethod\n    def load_image(image_path):\n        \"\"\"Load image with minimal processing\"\"\"\n        try:\n            # Try with PIL first\n            image = Image.open(image_path)\n            image = np.array(image)\n        except:\n            # Fallback to OpenCV\n            image = cv2.imread(image_path)\n            if image is None:\n                raise ValueError(f\"Cannot load image: {image_path}\")\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Ensure RGB\n        if len(image.shape) == 2:\n            image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n        elif image.shape[2] == 4:\n            image = cv2.cvtColor(image, cv2.COLOR_RGBA2RGB)\n        \n        return image\n    \n    @staticmethod\n    def preprocess_for_model(image_path, resample_size=3200):\n        \"\"\"Convert image to tensor with minimal processing\"\"\"\n        # Load image\n        image = ImagePreprocessor.load_image(image_path)\n        \n        # Convert to tensor: HWC -> CHW, normalize to [0, 1]\n        img_tensor = torch.from_numpy(image).permute(2, 0, 1).float() / 255.0\n        img_tensor = img_tensor.unsqueeze(0)\n        \n        # Resize if too large\n        height, width = img_tensor.shape[2], img_tensor.shape[3]\n        max_dim = max(height, width)\n        \n        if max_dim > resample_size:\n            scale = resample_size / max_dim\n            new_size = (int(height * scale), int(width * scale))\n            img_tensor = F.interpolate(\n                img_tensor, \n                size=new_size, \n                mode='bilinear', \n                align_corners=False, \n                antialias=True\n            )\n        \n        return img_tensor\n\n\n# ============================================================================\n# MINIMAL SIGNAL POST-PROCESSING\n# ============================================================================\nclass SignalPostProcessor:\n    \"\"\"Minimal post-processing to preserve signal quality\"\"\"\n    \n    @staticmethod\n    def clean_nans(signal):\n        \"\"\"Remove NaN values\"\"\"\n        signal = np.array(signal, dtype=np.float64)\n        \n        if not np.any(np.isnan(signal)):\n            return signal\n        \n        # If more than 50% NaN, return zeros\n        if np.sum(np.isnan(signal)) > len(signal) * 0.5:\n            return np.zeros_like(signal)\n        \n        # Interpolate NaN values\n        nan_mask = np.isnan(signal)\n        if np.any(~nan_mask):\n            valid_indices = np.where(~nan_mask)[0]\n            nan_indices = np.where(nan_mask)[0]\n            signal[nan_indices] = np.interp(nan_indices, valid_indices, signal[valid_indices])\n        else:\n            signal[nan_mask] = 0\n        \n        return signal\n    \n    @staticmethod\n    def is_valid_signal(signal):\n        \"\"\"Check if signal is valid\"\"\"\n        if signal is None or len(signal) == 0:\n            return False\n        signal = np.array(signal)\n        if np.all(np.isnan(signal)) or np.all(signal == 0):\n            return False\n        # Check if signal has some variation\n        if np.std(signal) < 1e-6:\n            return False\n        return True\n    \n    @staticmethod\n    def minimal_filter(signal, fs=500):\n        \"\"\"Apply minimal filtering - only baseline removal\"\"\"\n        signal = SignalPostProcessor.clean_nans(signal)\n        \n        if len(signal) < 10:\n            return signal\n        \n        # Simple baseline removal using median\n        window = min(int(fs * 0.6), len(signal))\n        if window % 2 == 0:\n            window += 1\n        window = max(3, window)\n        \n        if len(signal) >= window:\n            try:\n                baseline = medfilt(signal, kernel_size=window)\n                signal = signal - baseline\n            except:\n                signal = signal - np.median(signal)\n        else:\n            signal = signal - np.median(signal)\n        \n        return signal\n\n\n# ============================================================================\n# MAIN PIPELINE - SIMPLIFIED\n# ============================================================================\nclass SimplifiedECGDigitizer:\n    def __init__(self, weights_path, layout_weights_path, device=DEVICE):\n        self.device = device\n        \n        # Load models\n        self.segmentation_model = self._load_segmentation_model(weights_path)\n        self.layout_model = self._load_layout_model(layout_weights_path)\n        \n        # Initialize components\n        self.perspective_detector = PerspectiveDetector(num_thetas=200)\n        self.cropper = Cropper(percentiles=(0.03, 0.97), alpha=0.99)\n        self.pixel_size_finder = PixelSizeFinder(\n            min_number_of_grid_lines=30,\n            max_number_of_grid_lines=100,\n            lower_grid_line_factor=0.1,\n        )\n        self.signal_extractor = SignalExtractor()\n        \n        # Load layouts\n        self.layouts = yaml.safe_load(\n            open(\"src/config/lead_layouts_george-moody-2024.yml\", \"r\")\n        )\n        \n        # Preprocessor and post-processor\n        self.preprocessor = ImagePreprocessor()\n        self.post_processor = SignalPostProcessor()\n        \n    def _load_segmentation_model(self, weights_path):\n        model = UNet(\n            num_in_channels=3,\n            num_out_channels=4,\n            dims=[32, 64, 128, 256, 320, 320, 320, 320],\n            depth=2,\n        )\n        state_dict = torch.load(weights_path, map_location=self.device)\n        state_dict = {k.replace(\"_orig_mod.\", \"\"): v for k, v in state_dict.items()}\n        model.load_state_dict(state_dict)\n        model.eval().to(self.device)\n        return model\n    \n    def _load_layout_model(self, weights_path):\n        model = UNet(\n            num_in_channels=1,\n            num_out_channels=13,\n            dims=[32, 64, 128, 256, 256],\n            depth=2,\n        )\n        state_dict = torch.load(weights_path, map_location=self.device)\n        state_dict = {k.replace(\"_orig_mod.\", \"\"): v for k, v in state_dict.items()}\n        model.load_state_dict(state_dict)\n        model.eval().to(self.device)\n        return model\n    \n    def extract_raw_signals(self, image_path, target_num_samples):\n        \"\"\"Extract signals using Open-ECG-Digitizer\"\"\"\n        try:\n            # Minimal preprocessing\n            input_img = self.preprocessor.preprocess_for_model(image_path, resample_size=3200)\n            \n            with torch.no_grad():\n                logits = self.segmentation_model(input_img.to(self.device))\n                output_probs = torch.softmax(logits, dim=1)\n            \n            # Crop and align\n            aligned_image, aligned_signal, aligned_grid, aligned_text = self._crop_image(\n                input_img, output_probs\n            )\n            \n            # Get pixel calibration\n            mm_per_pixel_x, mm_per_pixel_y = self.pixel_size_finder(aligned_grid)\n            avg_pixel_per_mm = (1 / mm_per_pixel_x + 1 / mm_per_pixel_y) / 2\n            \n            # Extract signals\n            signals = self.signal_extractor(aligned_signal.squeeze())\n            \n            # Identify leads - CRITICAL: Increase required_valid_samples\n            identifier = LeadIdentifier(\n                layouts=self.layouts,\n                unet=self.layout_model,\n                device=self.device,\n                possibly_flipped=True,  # Allow flipped detection\n                target_num_samples=target_num_samples,\n                required_valid_samples=1,  # Lower threshold to get more leads\n            )\n            \n            signals = identifier(signals, aligned_text, avg_pixel_per_mm=avg_pixel_per_mm)\n            \n            return signals\n            \n        except Exception as e:\n            print(f\"Error in signal extraction: {e}\")\n            import traceback\n            traceback.print_exc()\n            return {\"canonical_lines\": torch.zeros(12, target_num_samples)}\n    \n    def _crop_image(self, image, probs):\n        \"\"\"Crop and align image\"\"\"\n        alignment_params = self.perspective_detector(probs[0, 0])\n        source_points = self.cropper(probs[0, 1], alignment_params)\n        \n        signal_prob = probs[:, [2]]\n        grid_prob = probs[:, [0]]\n        text_prob = probs[:, [1]]\n        \n        aligned_image, aligned_signal, aligned_grid, aligned_text = self._align_feature_maps(\n            image, signal_prob, grid_prob, text_prob, source_points\n        )\n        \n        return aligned_image, aligned_signal, aligned_grid, aligned_text\n    \n    def _align_feature_maps(self, image, signal_prob, grid_prob, text_prob, source_points):\n        \"\"\"Align feature maps using perspective transform\"\"\"\n        aligned_signal = self.cropper.apply_perspective(signal_prob, source_points, fill_value=0)\n        aligned_image = self.cropper.apply_perspective(image, source_points, fill_value=0)\n        aligned_grid = self.cropper.apply_perspective(grid_prob, source_points, fill_value=0)\n        aligned_text = self.cropper.apply_perspective(text_prob, source_points, fill_value=0)\n        \n        aligned_image, aligned_signal, aligned_grid, aligned_text = self._crop_y(\n            aligned_image, aligned_signal, aligned_grid, aligned_text\n        )\n        \n        return aligned_image, aligned_signal, aligned_grid, aligned_text\n    \n    def _crop_y(self, image, signal_prob, grid_prob, text_prob):\n        \"\"\"Crop Y dimension\"\"\"\n        def get_bounds(tensor):\n            prob = torch.clamp(\n                tensor.squeeze().sum(dim=tensor.dim() - 3) - \n                tensor.squeeze().sum(dim=tensor.dim() - 3).mean(),\n                min=0,\n            )\n            non_zero = (prob > 0).nonzero(as_tuple=True)[0]\n            if non_zero.numel() == 0:\n                return 0, tensor.shape[2] - 1\n            return int(non_zero[0].item()), int(non_zero[-1].item())\n        \n        y1, y2 = get_bounds(signal_prob + grid_prob)\n        slices = (slice(None), slice(None), slice(y1, y2 + 1), slice(None))\n        \n        return image[slices], signal_prob[slices], grid_prob[slices], text_prob[slices]\n    \n    def post_process_signals(self, signals_dict, fs=500):\n        \"\"\"Minimal post-processing\"\"\"\n        processed = {}\n        \n        for lead, signal in signals_dict.items():\n            # Convert to numpy and clean\n            signal = np.array(signal, dtype=np.float64)\n            signal = self.post_processor.clean_nans(signal)\n            \n            if not self.post_processor.is_valid_signal(signal):\n                processed[lead] = signal\n                continue\n            \n            # Minimal filtering - only baseline removal\n            signal = self.post_processor.minimal_filter(signal, fs)\n            \n            processed[lead] = signal\n        \n        return processed\n    \n    def process_single_ecg(self, image_path, target_num_samples, fs=500):\n        \"\"\"Complete pipeline for single ECG\"\"\"\n        # Extract raw signals\n        raw_result = self.extract_raw_signals(image_path, target_num_samples)\n        \n        # Convert to dict\n        raw_signals = {}\n        valid_count = 0\n        \n        for i, lead in enumerate(LEADS):\n            signal_tensor = raw_result[\"canonical_lines\"][i].cpu().numpy()\n            \n            # Clean NaNs\n            signal_clean = self.post_processor.clean_nans(signal_tensor)\n            \n            # Check scaling - ECG signals should be in mV range (-5 to +5 typically)\n            # If signal is too large (>100), it might be in µV\n            signal_magnitude = np.abs(signal_clean).max()\n            if signal_magnitude > 100:\n                signal_clean = signal_clean * 1e-3  # Convert µV to mV\n            elif signal_magnitude > 10:\n                signal_clean = signal_clean * 0.1  # Scale down\n            \n            raw_signals[lead] = signal_clean\n            \n            if self.post_processor.is_valid_signal(signal_clean):\n                valid_count += 1\n        \n        print(f\"  Valid leads: {valid_count}/12\")\n        \n        # Minimal post-processing\n        final_signals = self.post_process_signals(raw_signals, fs)\n        \n        return final_signals\n\n\n# ============================================================================\n# SUBMISSION GENERATION\n# ============================================================================\ndef create_submission(digitizer, test_df, output_path):\n    \"\"\"Generate submission file\"\"\"\n    if os.path.exists(output_path):\n        os.remove(output_path)\n    \n    # Create header\n    pd.DataFrame(columns=[\"id\", \"value\"]).to_csv(output_path, index=False)\n    \n    old_id = None\n    cached_signals = None\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\" \" * 15 + \"SIMPLIFIED ECG DIGITIZER\")\n    print(\"=\"*70)\n    print(\"\\nApproach:\")\n    print(\"  ✓ Minimal preprocessing to preserve signal\")\n    print(\"  ✓ Lower lead identification threshold\")\n    print(\"  ✓ Robust NaN handling\")\n    print(\"  ✓ Proper signal scaling\")\n    print(\"=\"*70 + \"\\n\")\n    \n    for index, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Processing\"):\n        current_id = row[\"id\"]\n        \n        # Process new image\n        if current_id != old_id:\n            old_id = current_id\n            image_path = f\"/kaggle/input/physionet-ecg-image-digitization/test/{current_id}.png\"\n            target_num_samples = int(row[\"fs\"] * 10)  # 10 seconds\n            fs = int(row[\"fs\"])\n            \n            try:\n                print(f\"\\n[{current_id}] fs={fs}Hz\")\n                cached_signals = digitizer.process_single_ecg(image_path, target_num_samples, fs)\n                \n                # Report valid leads\n                valid_leads = [lead for lead in LEADS if lead in cached_signals and \n                             SignalPostProcessor.is_valid_signal(cached_signals[lead])]\n                print(f\"  Extracted: {', '.join(valid_leads)}\")\n                \n                # Show signal ranges\n                for lead in valid_leads[:3]:  # Show first 3\n                    sig = cached_signals[lead]\n                    print(f\"  {lead}: [{sig.min():.3f}, {sig.max():.3f}] mV, std={sig.std():.3f}\")\n                \n            except Exception as e:\n                print(f\"  ERROR: {e}\")\n                cached_signals = {lead: np.zeros(target_num_samples) for lead in LEADS}\n        \n        # Extract requested lead and rows\n        lead_name = row['lead']\n        number_of_rows = row[\"number_of_rows\"]\n        \n        if lead_name not in cached_signals:\n            signal_resampled = np.zeros(number_of_rows)\n        else:\n            signal = cached_signals[lead_name]\n            \n            # Resample to requested length\n            if len(signal) == 0 or not SignalPostProcessor.is_valid_signal(signal):\n                signal_resampled = np.zeros(number_of_rows)\n            else:\n                if len(signal) != number_of_rows:\n                    signal_resampled = resample(signal, number_of_rows)\n                else:\n                    signal_resampled = signal\n        \n        # Final cleanup - replace any remaining NaNs with 0\n        signal_resampled = np.nan_to_num(signal_resampled, nan=0.0, posinf=0.0, neginf=0.0)\n        \n        # Create submission chunk\n        chunk = []\n        for t in range(number_of_rows):\n            chunk.append({\n                \"id\": f\"{current_id}_{t}_{lead_name}\",\n                \"value\": float(signal_resampled[t])\n            })\n        \n        if chunk:\n            pd.DataFrame(chunk).to_csv(output_path, mode='a', index=False, header=False)\n    \n    print(\"\\n\" + \"=\"*70)\n    print(f\"✓ Submission saved to {output_path}\")\n    print(\"=\"*70)\n\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*70)\n    print(\" \" * 10 + \"PHYSIONET ECG DIGITIZATION - SIMPLIFIED\")\n    print(\"=\"*70)\n    print(\"\\nKey Changes:\")\n    print(\"  • Minimal image preprocessing\")\n    print(\"  • Lower lead identification threshold\")\n    print(\"  • Only baseline filtering (no bandpass)\")\n    print(\"  • Proper signal scaling checks\")\n    print(\"  • Robust NaN handling throughout\")\n    print(\"=\"*70 + \"\\n\")\n    \n    digitizer = SimplifiedECGDigitizer(\n        weights_path=\"/kaggle/input/open-ecg-digitizer-weights/pytorch/default/1/unet_weights_07072025.pt\",\n        layout_weights_path=\"/kaggle/input/open-ecg-digitizer-weights/pytorch/default/1/lead_name_unet_weights_07072025.pt\",\n        device=DEVICE\n    )\n    \n    test_df = pd.read_csv('/kaggle/input/physionet-ecg-image-digitization/test.csv')\n    \n    print(f\"Dataset Info:\")\n    print(f\"  Total rows: {len(test_df):,}\")\n    print(f\"  Unique images: {test_df['id'].nunique()}\")\n    print(f\"  Sampling rates: {sorted(test_df['fs'].unique())}\")\n    print(f\"  Leads: {sorted(test_df['lead'].unique())}\\n\")\n    \n    create_submission(digitizer, test_df, '/kaggle/working/submission.csv')\n    \n    print(\"\\n✓ Processing complete!\")\n    print(\"Check submission.csv for results.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}