{"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":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Happy new year 2026 my fellow kagglers! If you are a serious data scientist active on kaggle and would love to work with me on kaggle tasks or any other useful data, and model development, fine-tuning etc, please feel free to DM me, \n\nAttached is the Comprehensive Explanation of ECG Image Digitization Pipeline with respect to PhysioNet - Digitization of ECG Images. I hope the Organizers and fellow consistent kaggler find this notebook useful in their quest for predictive model development. \n\n\nOverview\n\nThis is a sophisticated pipeline for digitizing ECG images from the PhysioNet ECG Image Digitization Challenge. The system converts scanned ECG images back into time-series signal data with enhanced accuracy and provides extensive visualization capabilities.\n\nArchitecture Breakdown: \n\n# 1. Configuration Management\npython\n@dataclass\nclass Config:\nCentralizes all parameters (paths, processing settings, model parameters)\n\nUses Python dataclass for clean configuration management\n\nKey parameters:\n\nImage processing: DPI, grid removal, thresholding\n\nSignal processing: Bandpass filters, notch filters\n\nModel training: Validation split, epochs, learning rate\n\n# 2. Visualization & EDA Module (ECGVisualizer)\nPurpose: Comprehensive exploratory data analysis and visualization\n\nKey Features:\n\nDataset Statistics: Visualizes distribution of patients, signals per patient, lead distribution\n\nSignal Examples: Shows actual ECG waveforms from different leads\n\nProcessing Pipeline: Visualizes each stage of image processing\n\nTraining Metrics: Tracks loss, SNR, MAE over epochs\n\nPrediction Quality: Analyzes SNR, MAE, RMSE distributions\n\nSignal Comparison: Overlay of ground truth vs predicted signals\n\n# 3. Image Preprocessing Module (EnhancedECGFrame)\nPurpose: Extract ECG signals from scanned images\n\nProcessing Pipeline:\n\nRotation Detection: Uses Hough Transform to detect and correct image tilt\n\nNoise Reduction: Fast Non-Local Means Denoising\n\nContrast Enhancement: CLAHE (Contrast Limited Adaptive Histogram Equalization)\n\nGrid Removal: Morphological operations to remove graph paper lines\n\nHorizontal kernel (40x1) for horizontal lines\n\nVertical kernel (1x40) for vertical lines\n\nBinarization: Otsu's thresholding\n\nLead Extraction: Connected component analysis to identify individual leads\n\nKey Innovations:\n\nGrid spacing detection using autocorrelation\n\nCenter-of-mass tracking for signal extraction\n\nRobust component filtering by aspect ratio and size\n\n# 4. Signal Processing Module (EnhancedSignalProcessor)\nPurpose: Clean and process extracted ECG signals\n\nProcessing Chain:\n\nBandpass Filter (0.5-40 Hz): Removes high-frequency noise and baseline wander\n\nBaseline Removal: Median filter to remove slow drifts\n\nNotch Filter (60 Hz): Removes power line interference\n\nR-peak Detection: Pan-Tompkins algorithm for heart rate calculation\n\nSNR Computation: Signal-to-noise ratio estimation\n\nResampling: Cubic interpolation for uniform signal length\n\n# 5. Data Loading Module (DataLoader)\nPurpose: Handle training and test data in the PhysioNet format\n\nStructure:\n\nTraining Data: Patient folders with PNG images and CSV signal files\n\nTest Data: Individual PNG images with metadata in test.csv\n\nKey Correction: Fixed column name from 'patient_id' to 'id' to match competition format\n\n# 6. Evaluation Metric (ECGSNRMetric)\nPurpose: Implement competition-specific SNR metric\n\nFeatures:\n\nSignal Alignment: Cross-correlation for time shift compensation\n\nVertical Offset Removal: Removes DC bias\n\nCombined SNR: Aggregates across all 12 leads\n\nRobust Handling: Deals with varying signal lengths\n\n# 7. Submission Generator\nPurpose: Format predictions for competition submission\n\nFormat: Creates rows with ID pattern: {patient_id}_{timestep}_{lead_name}\n\nMain Pipeline Execution Flow\nPhase 1: Data Loading\nLoads 977 training samples and 2 test samples\n\nOrganizes data into structured dictionaries\n\nPhase 2: Exploratory Data Analysis\nGenerated comprehensive statistics visualizations\n\nShowed signal examples from different leads\n\nPhase 3: Training Image Processing\nIssue Encountered:\n\nProcessed 5 training images successfully\n\nHowever, no ground truth comparisons were made (processed_train remained empty)\n\nLikely because CSV ground truth files weren't found or had mismatched formats\n\nPhase 4: Test Image Processing\nSuccessfully Processed:\n\nPatient 1053922973: Extracted 4 leads (I, II, III, aVR)\n\nPatient 2352854581: Extracted 1 lead (I)\n\nEach signal: 2500 samples (5 seconds at 500 Hz)\n\nPhase 5: Submission Creation\nCreated submission with 12,500 rows (5 leads × 2500 timepoints)\n\nSaved as Parquet format\n\nSample submission shows non-zero values, indicating successful digitization\n\nKey Results Analysis\nSuccessful Aspects:\nImage Processing Works: Successfully extracted signals from test images\n\nPipeline Execution: All modules execute without critical errors\n\nVisualization: Comprehensive plots generated and saved\n\nSubmission Format: Correctly formatted for competition\n\nIssues Identified:\nTraining Data Integration: Ground truth CSVs not properly linked\n\nprocessed_train remained empty (line 1161 in results)\n\nNo SNR/MAE metrics computed for training\n\nPartial Lead Extraction:\n\nFirst test patient: Only 4/12 leads extracted\n\nSecond test patient: Only 1/12 leads extracted\n\nData Loading Assumptions: Code assumes specific file naming conventions that might not match actual data\n\nTechnical Highlights:\nAdvanced Image Processing:\n\npython\n# Grid removal using morphological operations\nhorizontal_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (40, 1))\nvertical_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (1, 40))\nRobust Signal Alignment:\n\npython\n# Cross-correlation for time shift detection\ncorrelation = np.correlate(predicted, ground_truth, mode='full')\noptimal_shift = lags[np.argmax(correlation)]\nComprehensive Visualization:\n\n# 6 different plot types covering EDA, processing, training, and evaluation\n\nProfessional formatting with consistent color schemes\n\nHigh-resolution output (300 DPI)\n\nPotential Improvements\nImmediate Fixes:\nFix Training Data Loading: Update paths and file naming to match actual data structure\n\nImprove Lead Detection: Adjust connected component parameters to catch all 12 leads\n\nAdd Error Handling: Better exception handling for missing ground truth files\n\nAdvanced Enhancements:\nDeep Learning Integration: Could add CNN for direct image-to-signal mapping\n\nEnsemble Methods: Combine multiple extraction algorithms\n\nReal-time Processing: Optimize for clinical applications\n\nQuality Metrics: Add signal quality indices for reliability assessment\n\n# Conclusion\n\nThis pipeline represents a sophisticated approach to ECG image digitization with:\n\n# Robust preprocessing (rotation correction, grid removal, denoising)\n\n# Advanced signal processing (filtering, alignment, SNR computation)\n\n# Comprehensive visualization (EDA, training monitoring, quality assessment)\n\n# Competition-ready output (properly formatted submission)\n\nThe main issue preventing full training evaluation appears to be data path/format mismatches, but the core digitization functionality works as demonstrated by the successful test processing and submission generation.","metadata":{}},{"cell_type":"code","source":"# =============================================================================\n# ENHANCED ECG IMAGE DIGITIZATION PIPELINE WITH EDA & VISUALIZATION\n# =============================================================================\nimport os\nimport glob\nimport json\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom typing import Tuple, List, Dict, Optional, Union\nfrom dataclasses import dataclass\nfrom scipy import signal, ndimage\nfrom scipy.interpolate import interp1d\nfrom scipy.signal import butter, filtfilt, find_peaks, correlation_lags, iirnotch\nfrom scipy.ndimage import median_filter, gaussian_filter\nfrom scipy.optimize import minimize\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.model_selection import train_test_split, KFold\nfrom sklearn.metrics import mean_squared_error, mean_absolute_error, r2_score\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Set visualization style\nsns.set_style(\"whitegrid\")\nplt.rcParams['figure.figsize'] = (12, 6)\nplt.rcParams['font.size'] = 10\n\n# =============================================================================\n# CONFIGURATION & PATHS\n# =============================================================================\n@dataclass\nclass Config:\n    \"\"\"Configuration for ECG digitization pipeline\"\"\"\n    # Paths\n    root_path: str = \"/kaggle/input/physionet-ecg-image-digitization\"\n    test_images: str = \"/kaggle/input/physionet-ecg-image-digitization/test\"\n    train_images: str = \"/kaggle/input/physionet-ecg-image-digitization/train\"\n    sample_submission: str = \"/kaggle/input/physionet-ecg-image-digitization/sample_submission.parquet\"\n    test_csv: str = \"/kaggle/input/physionet-ecg-image-digitization/test.csv\"\n    train_csv: str = \"/kaggle/input/physionet-ecg-image-digitization/train.csv\"\n    \n    # Processing parameters - OPTIMIZED\n    target_dpi: int = 300\n    grid_removal_sigma: float = 1.5\n    lead_extraction_threshold: float = 0.75\n    min_signal_length: int = 100\n    max_signal_length: int = 10000\n    default_sampling_rate: int = 500\n    interpolation_method: str = 'cubic'\n    \n    # Enhanced filtering\n    bandpass_low: float = 0.5\n    bandpass_high: float = 40.0\n    notch_freq: float = 60.0  # Power line interference\n    baseline_window: float = 0.2\n    \n    # Model parameters\n    validation_split: float = 0.2\n    n_folds: int = 5\n    batch_size: int = 32\n    epochs: int = 50\n    learning_rate: float = 0.001\n    patience: int = 10\n    \n    # Output\n    output_dir: str = \"/kaggle/working\"\n    plots_dir: str = \"/kaggle/working/plots\"\n    submission_file: str = \"/kaggle/working/submission.parquet\"\n\nconfig = Config()\n\n# =============================================================================\n# VISUALIZATION & EDA MODULE\n# =============================================================================\nclass ECGVisualizer:\n    \"\"\"Comprehensive visualization and EDA for ECG data\"\"\"\n    \n    def __init__(self, output_dir: str):\n        self.output_dir = output_dir\n        os.makedirs(output_dir, exist_ok=True)\n        self.color_palette = sns.color_palette(\"husl\", 12)\n        \n    def plot_dataset_statistics(self, train_data: Dict, filename: str = \"dataset_stats.png\"):\n        \"\"\"Plot comprehensive dataset statistics\"\"\"\n        fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n        fig.suptitle('ECG Dataset Statistics & Distribution Analysis', fontsize=16, fontweight='bold')\n        \n        # Extract statistics\n        num_patients = len(train_data)\n        signals_per_patient = []\n        signal_lengths = []\n        leads_distribution = {lead: 0 for lead in ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', \n                                                     'V1', 'V2', 'V3', 'V4', 'V5', 'V6']}\n        \n        for patient_id, data in train_data.items():\n            if 'csv_data' in data:\n                csv_data = data['csv_data']\n                for col in csv_data.columns:\n                    if col in leads_distribution:\n                        leads_distribution[col] += 1\n                        signal_lengths.append(len(csv_data[col].dropna()))\n                signals_per_patient.append(len(csv_data.columns))\n        \n        # 1. Number of patients\n        axes[0, 0].bar(['Training Patients'], [num_patients], color=self.color_palette[0])\n        axes[0, 0].set_ylabel('Count')\n        axes[0, 0].set_title('Dataset Size')\n        axes[0, 0].text(0, num_patients + 5, f'N = {num_patients}', ha='center', fontweight='bold')\n        \n        # 2. Signals per patient distribution\n        if signals_per_patient:\n            axes[0, 1].hist(signals_per_patient, bins=20, color=self.color_palette[1], edgecolor='black', alpha=0.7)\n            axes[0, 1].set_xlabel('Number of Signals')\n            axes[0, 1].set_ylabel('Frequency')\n            axes[0, 1].set_title(f'Signals per Patient (Mean: {np.mean(signals_per_patient):.1f})')\n            axes[0, 1].axvline(np.mean(signals_per_patient), color='red', linestyle='--', label='Mean')\n            axes[0, 1].legend()\n        \n        # 3. Signal length distribution\n        if signal_lengths:\n            axes[0, 2].hist(signal_lengths, bins=30, color=self.color_palette[2], edgecolor='black', alpha=0.7)\n            axes[0, 2].set_xlabel('Signal Length (samples)')\n            axes[0, 2].set_ylabel('Frequency')\n            axes[0, 2].set_title(f'Signal Length Distribution (Mean: {np.mean(signal_lengths):.0f})')\n            axes[0, 2].axvline(np.mean(signal_lengths), color='red', linestyle='--', label='Mean')\n            axes[0, 2].legend()\n        \n        # 4. Lead distribution\n        leads = list(leads_distribution.keys())\n        counts = list(leads_distribution.values())\n        axes[1, 0].bar(leads, counts, color=self.color_palette[:len(leads)])\n        axes[1, 0].set_xlabel('Lead Name')\n        axes[1, 0].set_ylabel('Count')\n        axes[1, 0].set_title('Distribution of ECG Leads')\n        axes[1, 0].tick_params(axis='x', rotation=45)\n        \n        # 5. Signal quality metrics (using first few signals)\n        snr_values = []\n        baseline_drift = []\n        sample_count = 0\n        for patient_id, data in train_data.items():\n            if sample_count >= 100:\n                break\n            if 'csv_data' in data:\n                for col in data['csv_data'].columns[:3]:\n                    sig = data['csv_data'][col].dropna().values\n                    if len(sig) > 100:\n                        # Estimate SNR\n                        sig_std = np.std(sig)\n                        noise_est = np.std(np.diff(sig))\n                        if noise_est > 0:\n                            snr_values.append(20 * np.log10(sig_std / noise_est))\n                        # Baseline drift\n                        baseline_drift.append(np.abs(np.mean(sig[:100]) - np.mean(sig[-100:])))\n                        sample_count += 1\n        \n        if snr_values:\n            axes[1, 1].hist(snr_values, bins=25, color=self.color_palette[3], edgecolor='black', alpha=0.7)\n            axes[1, 1].set_xlabel('Estimated SNR (dB)')\n            axes[1, 1].set_ylabel('Frequency')\n            axes[1, 1].set_title(f'Signal Quality (SNR) Distribution')\n            axes[1, 1].axvline(np.mean(snr_values), color='red', linestyle='--', label=f'Mean: {np.mean(snr_values):.1f} dB')\n            axes[1, 1].legend()\n        \n        # 6. Summary statistics table\n        axes[1, 2].axis('off')\n        summary_data = [\n            ['Total Patients', f'{num_patients}'],\n            ['Avg Signals/Patient', f'{np.mean(signals_per_patient):.1f}' if signals_per_patient else 'N/A'],\n            ['Avg Signal Length', f'{np.mean(signal_lengths):.0f}' if signal_lengths else 'N/A'],\n            ['Total Signals', f'{sum(counts)}'],\n            ['Avg SNR', f'{np.mean(snr_values):.1f} dB' if snr_values else 'N/A'],\n            ['Unique Leads', f'{sum(1 for c in counts if c > 0)}']\n        ]\n        table = axes[1, 2].table(cellText=summary_data, colLabels=['Metric', 'Value'],\n                                 loc='center', cellLoc='left')\n        table.auto_set_font_size(False)\n        table.set_fontsize(10)\n        table.scale(1, 2)\n        axes[1, 2].set_title('Summary Statistics')\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Dataset statistics saved to: {filepath}\")\n    \n    def plot_signal_examples(self, train_data: Dict, num_examples: int = 6, \n                            filename: str = \"signal_examples.png\"):\n        \"\"\"Plot example ECG signals from different leads\"\"\"\n        fig, axes = plt.subplots(3, 2, figsize=(16, 12))\n        fig.suptitle('Example ECG Signals Across Different Leads', fontsize=16, fontweight='bold')\n        axes = axes.flatten()\n        \n        leads_to_plot = ['I', 'II', 'V1', 'V2', 'V5', 'V6']\n        plotted = 0\n        \n        for patient_id, data in train_data.items():\n            if plotted >= num_examples:\n                break\n            if 'csv_data' in data:\n                csv_data = data['csv_data']\n                for i, lead in enumerate(leads_to_plot):\n                    if lead in csv_data.columns and plotted < num_examples:\n                        sig = csv_data[lead].dropna().values\n                        if len(sig) > 100:\n                            time = np.arange(len(sig)) / 500  # Assuming 500 Hz\n                            axes[i].plot(time[:2000], sig[:2000], color=self.color_palette[i], linewidth=0.8)\n                            axes[i].set_title(f'Lead {lead} - Patient {patient_id}')\n                            axes[i].set_xlabel('Time (s)')\n                            axes[i].set_ylabel('Amplitude (mV)')\n                            axes[i].grid(True, alpha=0.3)\n                            plotted += 1\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Signal examples saved to: {filepath}\")\n    \n    def plot_processing_pipeline(self, original: np.ndarray, processed_stages: Dict,\n                                filename: str = \"processing_pipeline.png\"):\n        \"\"\"Visualize image processing pipeline stages\"\"\"\n        num_stages = len(processed_stages) + 1\n        fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n        fig.suptitle('ECG Image Processing Pipeline Stages', fontsize=16, fontweight='bold')\n        axes = axes.flatten()\n        \n        # Original\n        axes[0].imshow(original, cmap='gray')\n        axes[0].set_title('1. Original Image')\n        axes[0].axis('off')\n        \n        # Processed stages\n        for i, (stage_name, stage_image) in enumerate(processed_stages.items(), 1):\n            if i < len(axes):\n                if len(stage_image.shape) == 2:\n                    axes[i].imshow(stage_image, cmap='gray')\n                else:\n                    axes[i].imshow(stage_image)\n                axes[i].set_title(f'{i+1}. {stage_name}')\n                axes[i].axis('off')\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Processing pipeline visualization saved to: {filepath}\")\n    \n    def plot_training_metrics(self, history: Dict, filename: str = \"training_metrics.png\"):\n        \"\"\"Plot training and validation metrics over epochs\"\"\"\n        fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n        fig.suptitle('Model Training Metrics & Performance', fontsize=16, fontweight='bold')\n        \n        epochs = range(1, len(history['train_loss']) + 1)\n        \n        # Loss\n        axes[0, 0].plot(epochs, history['train_loss'], 'b-o', label='Training Loss', markersize=4)\n        axes[0, 0].plot(epochs, history['val_loss'], 'r-s', label='Validation Loss', markersize=4)\n        axes[0, 0].set_xlabel('Epoch')\n        axes[0, 0].set_ylabel('Loss (MSE)')\n        axes[0, 0].set_title('Training vs Validation Loss')\n        axes[0, 0].legend()\n        axes[0, 0].grid(True, alpha=0.3)\n        \n        # SNR\n        if 'train_snr' in history:\n            axes[0, 1].plot(epochs, history['train_snr'], 'b-o', label='Training SNR', markersize=4)\n            axes[0, 1].plot(epochs, history['val_snr'], 'r-s', label='Validation SNR', markersize=4)\n            axes[0, 1].set_xlabel('Epoch')\n            axes[0, 1].set_ylabel('SNR (dB)')\n            axes[0, 1].set_title('Training vs Validation SNR')\n            axes[0, 1].legend()\n            axes[0, 1].grid(True, alpha=0.3)\n        \n        # MAE\n        if 'train_mae' in history:\n            axes[1, 0].plot(epochs, history['train_mae'], 'b-o', label='Training MAE', markersize=4)\n            axes[1, 0].plot(epochs, history['val_mae'], 'r-s', label='Validation MAE', markersize=4)\n            axes[1, 0].set_xlabel('Epoch')\n            axes[1, 0].set_ylabel('Mean Absolute Error')\n            axes[1, 0].set_title('Training vs Validation MAE')\n            axes[1, 0].legend()\n            axes[1, 0].grid(True, alpha=0.3)\n        \n        # Generalization gap\n        axes[1, 1].plot(epochs, np.array(history['val_loss']) - np.array(history['train_loss']), \n                       'g-^', label='Generalization Gap', markersize=4)\n        axes[1, 1].axhline(y=0, color='k', linestyle='--', alpha=0.3)\n        axes[1, 1].set_xlabel('Epoch')\n        axes[1, 1].set_ylabel('Val Loss - Train Loss')\n        axes[1, 1].set_title('Generalization Gap (Overfitting Indicator)')\n        axes[1, 1].legend()\n        axes[1, 1].grid(True, alpha=0.3)\n        axes[1, 1].fill_between(epochs, 0, np.array(history['val_loss']) - np.array(history['train_loss']), \n                                alpha=0.3, color='green')\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Training metrics saved to: {filepath}\")\n    \n    def plot_prediction_quality(self, predictions: List[Dict], filename: str = \"prediction_quality.png\"):\n        \"\"\"Plot prediction quality metrics\"\"\"\n        fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n        fig.suptitle('Prediction Quality Analysis', fontsize=16, fontweight='bold')\n        axes = axes.flatten()\n        \n        # Extract metrics\n        snr_scores = [p['snr'] for p in predictions if 'snr' in p]\n        mae_scores = [p['mae'] for p in predictions if 'mae' in p]\n        rmse_scores = [p['rmse'] for p in predictions if 'rmse' in p]\n        correlation_scores = [p['correlation'] for p in predictions if 'correlation' in p]\n        \n        # SNR distribution\n        if snr_scores:\n            axes[0].hist(snr_scores, bins=30, color=self.color_palette[0], edgecolor='black', alpha=0.7)\n            axes[0].set_xlabel('SNR (dB)')\n            axes[0].set_ylabel('Frequency')\n            axes[0].set_title(f'SNR Distribution (Mean: {np.mean(snr_scores):.2f} dB)')\n            axes[0].axvline(np.mean(snr_scores), color='red', linestyle='--', linewidth=2, label='Mean')\n            axes[0].legend()\n        \n        # MAE distribution\n        if mae_scores:\n            axes[1].hist(mae_scores, bins=30, color=self.color_palette[1], edgecolor='black', alpha=0.7)\n            axes[1].set_xlabel('Mean Absolute Error')\n            axes[1].set_ylabel('Frequency')\n            axes[1].set_title(f'MAE Distribution (Mean: {np.mean(mae_scores):.4f})')\n            axes[1].axvline(np.mean(mae_scores), color='red', linestyle='--', linewidth=2, label='Mean')\n            axes[1].legend()\n        \n        # RMSE distribution\n        if rmse_scores:\n            axes[2].hist(rmse_scores, bins=30, color=self.color_palette[2], edgecolor='black', alpha=0.7)\n            axes[2].set_xlabel('Root Mean Squared Error')\n            axes[2].set_ylabel('Frequency')\n            axes[2].set_title(f'RMSE Distribution (Mean: {np.mean(rmse_scores):.4f})')\n            axes[2].axvline(np.mean(rmse_scores), color='red', linestyle='--', linewidth=2, label='Mean')\n            axes[2].legend()\n        \n        # Correlation distribution\n        if correlation_scores:\n            axes[3].hist(correlation_scores, bins=30, color=self.color_palette[3], edgecolor='black', alpha=0.7)\n            axes[3].set_xlabel('Correlation Coefficient')\n            axes[3].set_ylabel('Frequency')\n            axes[3].set_title(f'Correlation Distribution (Mean: {np.mean(correlation_scores):.4f})')\n            axes[3].axvline(np.mean(correlation_scores), color='red', linestyle='--', linewidth=2, label='Mean')\n            axes[3].legend()\n        \n        # SNR vs MAE scatter\n        if snr_scores and mae_scores and len(snr_scores) == len(mae_scores):\n            axes[4].scatter(snr_scores, mae_scores, alpha=0.5, color=self.color_palette[4])\n            axes[4].set_xlabel('SNR (dB)')\n            axes[4].set_ylabel('MAE')\n            axes[4].set_title('SNR vs MAE Relationship')\n            axes[4].grid(True, alpha=0.3)\n            \n            # Add trend line\n            z = np.polyfit(snr_scores, mae_scores, 1)\n            p = np.poly1d(z)\n            axes[4].plot(sorted(snr_scores), p(sorted(snr_scores)), \"r--\", alpha=0.8, linewidth=2)\n        \n        # Performance summary\n        axes[5].axis('off')\n        summary_data = [\n            ['Mean SNR', f'{np.mean(snr_scores):.2f} dB' if snr_scores else 'N/A'],\n            ['Median SNR', f'{np.median(snr_scores):.2f} dB' if snr_scores else 'N/A'],\n            ['Mean MAE', f'{np.mean(mae_scores):.4f}' if mae_scores else 'N/A'],\n            ['Mean RMSE', f'{np.mean(rmse_scores):.4f}' if rmse_scores else 'N/A'],\n            ['Mean Correlation', f'{np.mean(correlation_scores):.4f}' if correlation_scores else 'N/A'],\n            ['Num Predictions', f'{len(predictions)}']\n        ]\n        table = axes[5].table(cellText=summary_data, colLabels=['Metric', 'Value'],\n                             loc='center', cellLoc='left')\n        table.auto_set_font_size(False)\n        table.set_fontsize(11)\n        table.scale(1, 2.5)\n        axes[5].set_title('Performance Summary')\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Prediction quality analysis saved to: {filepath}\")\n    \n    def plot_cross_validation_results(self, cv_results: Dict, filename: str = \"cv_results.png\"):\n        \"\"\"Plot cross-validation results\"\"\"\n        fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n        fig.suptitle('Cross-Validation Performance Analysis', fontsize=16, fontweight='bold')\n        \n        folds = list(cv_results.keys())\n        train_scores = [cv_results[fold]['train_score'] for fold in folds]\n        val_scores = [cv_results[fold]['val_score'] for fold in folds]\n        \n        # Scores per fold\n        x = np.arange(len(folds))\n        width = 0.35\n        axes[0].bar(x - width/2, train_scores, width, label='Train Score', color=self.color_palette[0])\n        axes[0].bar(x + width/2, val_scores, width, label='Validation Score', color=self.color_palette[1])\n        axes[0].set_xlabel('Fold')\n        axes[0].set_ylabel('SNR (dB)')\n        axes[0].set_title('Train vs Validation Scores Across Folds')\n        axes[0].set_xticks(x)\n        axes[0].set_xticklabels(folds)\n        axes[0].legend()\n        axes[0].grid(True, alpha=0.3, axis='y')\n        \n        # Boxplot\n        axes[1].boxplot([train_scores, val_scores], labels=['Train', 'Validation'])\n        axes[1].set_ylabel('SNR (dB)')\n        axes[1].set_title('Score Distribution')\n        axes[1].grid(True, alpha=0.3, axis='y')\n        \n        # Add mean lines\n        axes[1].axhline(np.mean(train_scores), color=self.color_palette[0], linestyle='--', \n                       label=f'Train Mean: {np.mean(train_scores):.2f}')\n        axes[1].axhline(np.mean(val_scores), color=self.color_palette[1], linestyle='--',\n                       label=f'Val Mean: {np.mean(val_scores):.2f}')\n        axes[1].legend()\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Cross-validation results saved to: {filepath}\")\n    \n    def plot_signal_comparison(self, ground_truth: np.ndarray, predicted: np.ndarray,\n                              lead_name: str = \"Unknown\", filename: str = \"signal_comparison.png\"):\n        \"\"\"Plot ground truth vs predicted signal\"\"\"\n        fig, axes = plt.subplots(3, 1, figsize=(14, 10))\n        fig.suptitle(f'Signal Comparison: Lead {lead_name}', fontsize=16, fontweight='bold')\n        \n        time = np.arange(len(ground_truth)) / 500\n        \n        # Ground truth\n        axes[0].plot(time, ground_truth, color='blue', linewidth=1, label='Ground Truth')\n        axes[0].set_ylabel('Amplitude (mV)')\n        axes[0].set_title('Ground Truth Signal')\n        axes[0].legend()\n        axes[0].grid(True, alpha=0.3)\n        \n        # Predicted\n        axes[1].plot(time[:len(predicted)], predicted, color='red', linewidth=1, label='Predicted')\n        axes[1].set_ylabel('Amplitude (mV)')\n        axes[1].set_title('Predicted Signal')\n        axes[1].legend()\n        axes[1].grid(True, alpha=0.3)\n        \n        # Overlay\n        min_len = min(len(ground_truth), len(predicted))\n        axes[2].plot(time[:min_len], ground_truth[:min_len], color='blue', linewidth=1.5, \n                    alpha=0.7, label='Ground Truth')\n        axes[2].plot(time[:min_len], predicted[:min_len], color='red', linewidth=1, \n                    alpha=0.7, label='Predicted', linestyle='--')\n        axes[2].set_xlabel('Time (s)')\n        axes[2].set_ylabel('Amplitude (mV)')\n        axes[2].set_title('Overlay Comparison')\n        axes[2].legend()\n        axes[2].grid(True, alpha=0.3)\n        \n        # Calculate metrics\n        mae = np.mean(np.abs(ground_truth[:min_len] - predicted[:min_len]))\n        rmse = np.sqrt(np.mean((ground_truth[:min_len] - predicted[:min_len])**2))\n        corr = np.corrcoef(ground_truth[:min_len], predicted[:min_len])[0, 1]\n        \n        # Add text box with metrics\n        textstr = f'MAE: {mae:.4f}\\nRMSE: {rmse:.4f}\\nCorrelation: {corr:.4f}'\n        props = dict(boxstyle='round', facecolor='wheat', alpha=0.5)\n        axes[2].text(0.02, 0.98, textstr, transform=axes[2].transAxes, fontsize=10,\n                    verticalalignment='top', bbox=props)\n        \n        plt.tight_layout()\n        filepath = os.path.join(self.output_dir, filename)\n        plt.savefig(filepath, dpi=300, bbox_inches='tight')\n        plt.show()\n        print(f\"Signal comparison saved to: {filepath}\")\n\n# =============================================================================\n# ENHANCED IMAGE PREPROCESSING MODULE\n# =============================================================================\nclass EnhancedECGFrame:\n    \"\"\"Enhanced ECG image preprocessing with better accuracy\"\"\"\n    \n    def __init__(self, image_path: str):\n        self.image_path = image_path\n        self.image = cv2.imread(image_path)\n        if self.image is None:\n            raise ValueError(f\"Could not read image: {image_path}\")\n        \n        self.gray = cv2.cvtColor(self.image, cv2.COLOR_BGR2GRAY)\n        self.height, self.width = self.gray.shape\n        \n        # Enhanced detection\n        self.rotation_angle = self._detect_rotation()\n        self.grid_spacing = self._detect_grid_spacing()\n        self.grid_thickness = self._detect_grid_thickness()\n    \n    def _detect_rotation(self) -> float:\n        \"\"\"Enhanced rotation detection\"\"\"\n        edges = cv2.Canny(self.gray, 50, 150, apertureSize=3)\n        lines = cv2.HoughLines(edges, 1, np.pi/180, threshold=100)\n        \n        if lines is None:\n            return 0.0\n        \n        angles = []\n        for line in lines[:50]:\n            rho, theta = line[0]\n            angle = np.degrees(theta) - 90\n            if abs(angle) < 10:\n                angles.append(angle)\n        \n        return np.median(angles) if angles else 0.0\n    \n    def _detect_grid_spacing(self) -> Tuple[float, float]:\n        \"\"\"Enhanced grid spacing detection\"\"\"\n        # Use autocorrelation for robust grid detection\n        center_h = self.height // 4\n        center_w = self.width // 4\n        roi = self.gray[center_h:3*center_h, center_w:3*center_w]\n        \n        # Row-wise autocorrelation\n        row_profile = np.mean(roi, axis=1)\n        row_autocorr = np.correlate(row_profile, row_profile, mode='full')\n        row_autocorr = row_autocorr[len(row_autocorr)//2:]\n        row_peaks, _ = find_peaks(row_autocorr, distance=20, prominence=np.max(row_autocorr)*0.1)\n        \n        # Column-wise autocorrelation\n        col_profile = np.mean(roi, axis=0)\n        col_autocorr = np.correlate(col_profile, col_profile, mode='full')\n        col_autocorr = col_autocorr[len(col_autocorr)//2:]\n        col_peaks, _ = find_peaks(col_autocorr, distance=20, prominence=np.max(col_autocorr)*0.1)\n        \n        row_spacing = row_peaks[0] if len(row_peaks) > 0 else 50.0\n        col_spacing = col_peaks[0] if len(col_peaks) > 0 else 50.0\n        \n        return (row_spacing, col_spacing)\n    \n    def _detect_grid_thickness(self) -> Tuple[float, float]:\n        \"\"\"Detect grid line thickness\"\"\"\n        # Estimate thickness from horizontal and vertical profiles\n        h_profile = np.mean(self.gray, axis=0)\n        v_profile = np.mean(self.gray, axis=1)\n        \n        # Find transitions (grid lines)\n        h_diff = np.abs(np.diff(h_profile))\n        v_diff = np.abs(np.diff(v_profile))\n        \n        h_threshold = np.percentile(h_diff, 95)\n        v_threshold = np.percentile(v_diff, 95)\n        \n        h_transitions = h_diff > h_threshold\n        v_transitions = v_diff > v_threshold\n        \n        # Estimate thickness from transition clusters\n        h_thickness = np.median(np.diff(np.where(h_transitions)[0])) if np.any(h_transitions) else 1.0\n        v_thickness = np.median(np.diff(np.where(v_transitions)[0])) if np.any(v_transitions) else 1.0\n        \n        return (max(1.0, min(h_thickness, 3.0)), max(1.0, min(v_thickness, 3.0)))\n    \n    def preprocess(self) -> np.ndarray:\n        \"\"\"Complete preprocessing pipeline\"\"\"\n        # 1. Rotation correction\n        if abs(self.rotation_angle) > 0.1:\n            center = (self.width // 2, self.height // 2)\n            rotation_matrix = cv2.getRotationMatrix2D(center, self.rotation_angle, 1.0)\n            rotated = cv2.warpAffine(self.gray, rotation_matrix, (self.width, self.height),\n                                    flags=cv2.INTER_CUBIC, borderMode=cv2.BORDER_REPLICATE)\n        else:\n            rotated = self.gray.copy()\n        \n        # 2. Noise reduction\n        denoised = cv2.fastNlMeansDenoising(rotated, None, h=10, templateWindowSize=7, searchWindowSize=21)\n        \n        # 3. Contrast enhancement\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        enhanced = clahe.apply(denoised)\n        \n        # 4. Grid removal (morphological approach)\n        grid_removed = self._remove_grid_morphological(enhanced)\n        \n        # 5. Binarization\n        _, binary = cv2.threshold(grid_removed, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)\n        \n        # 6. Morphological cleaning\n        kernel = np.ones((2, 2), np.uint8)\n        cleaned = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel)\n        cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_OPEN, kernel)\n        \n        return cleaned\n    \n    def _remove_grid_morphological(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Remove grid using morphological operations\"\"\"\n        # Detect horizontal lines\n        horizontal_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (40, 1))\n        horizontal = cv2.morphologyEx(image, cv2.MORPH_OPEN, horizontal_kernel)\n        \n        # Detect vertical lines\n        vertical_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (1, 40))\n        vertical = cv2.morphologyEx(image, cv2.MORPH_OPEN, vertical_kernel)\n        \n        # Combine grid components\n        grid = cv2.add(horizontal, vertical)\n        \n        # Remove grid from original\n        grid_removed = cv2.subtract(image, grid)\n        \n        return grid_removed\n    \n    def extract_leads(self) -> Dict[str, np.ndarray]:\n        \"\"\"Extract individual lead signals from preprocessed image\"\"\"\n        processed = self.preprocess()\n        \n        # Detect lead regions using connected components\n        num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(processed, connectivity=8)\n        \n        leads = {}\n        lead_names = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n        \n        # Filter components by size and aspect ratio\n        valid_components = []\n        for i in range(1, num_labels):  # Skip background (0)\n            x, y, w, h, area = stats[i]\n            aspect_ratio = w / h if h > 0 else 0\n            \n            # Lead signals are typically wide and short\n            if area > 1000 and aspect_ratio > 5 and w > self.width * 0.3:\n                valid_components.append((y, i, x, w, h))  # Sort by y-position\n        \n        # Sort by vertical position (top to bottom)\n        valid_components.sort()\n        \n        # Extract signals for each lead\n        for idx, (y_pos, comp_id, x, w, h) in enumerate(valid_components[:len(lead_names)]):\n            lead_name = lead_names[idx]\n            \n            # Extract region\n            mask = (labels == comp_id).astype(np.uint8) * 255\n            lead_region = cv2.bitwise_and(processed, processed, mask=mask)\n            \n            # Extract signal trace\n            signal = self._extract_signal_from_region(lead_region[y_pos:y_pos+h, x:x+w])\n            \n            if signal is not None and len(signal) > config.min_signal_length:\n                leads[lead_name] = signal\n        \n        return leads\n    \n    def _extract_signal_from_region(self, region: np.ndarray) -> Optional[np.ndarray]:\n        \"\"\"Extract 1D signal from 2D region\"\"\"\n        if region.size == 0:\n            return None\n        \n        h, w = region.shape\n        signal = np.zeros(w)\n        \n        # For each column, find the center of mass of white pixels\n        for col in range(w):\n            column_data = region[:, col]\n            white_pixels = np.where(column_data > 0)[0]\n            \n            if len(white_pixels) > 0:\n                # Use center of mass for more accurate trace\n                signal[col] = np.mean(white_pixels)\n            elif col > 0:\n                # Interpolate if no pixels found\n                signal[col] = signal[col - 1]\n        \n        # Normalize signal (invert and scale)\n        signal = h - signal  # Invert (ECG is typically drawn top-down)\n        signal = signal - np.mean(signal)  # Center\n        \n        return signal\n\n# =============================================================================\n# ENHANCED SIGNAL PROCESSING MODULE\n# =============================================================================\nclass EnhancedSignalProcessor:\n    \"\"\"Advanced signal processing for ECG data\"\"\"\n    \n    def __init__(self, sampling_rate: int = 500):\n        self.sampling_rate = sampling_rate\n        self.nyquist = sampling_rate / 2\n    \n    def bandpass_filter(self, signal: np.ndarray, lowcut: float = 0.5, \n                       highcut: float = 40.0, order: int = 4) -> np.ndarray:\n        \"\"\"Apply bandpass filter to remove noise\"\"\"\n        low = lowcut / self.nyquist\n        high = highcut / self.nyquist\n        b, a = butter(order, [low, high], btype='band')\n        filtered = filtfilt(b, a, signal)\n        return filtered\n    \n    def remove_baseline_wander(self, signal: np.ndarray, window_size: int = None) -> np.ndarray:\n        \"\"\"Remove baseline wander using median filter\"\"\"\n        if window_size is None:\n            window_size = int(self.sampling_rate * config.baseline_window)\n        \n        baseline = median_filter(signal, size=window_size)\n        corrected = signal - baseline\n        return corrected\n    \n    def notch_filter(self, ecg_signal: np.ndarray, freq: float = 60.0, Q: float = 30.0) -> np.ndarray:\n        \"\"\"Remove power line interference\"\"\"\n        w0 = freq / self.nyquist\n        b, a = iirnotch(w0, Q)  # CORRECTED: Use imported iirnotch function\n        filtered = filtfilt(b, a, ecg_signal)\n        return filtered\n    \n    def detect_r_peaks(self, signal: np.ndarray) -> np.ndarray:\n        \"\"\"Detect R-peaks in ECG signal\"\"\"\n        # Use Pan-Tompkins algorithm\n        # 1. Bandpass filter\n        filtered = self.bandpass_filter(signal, 5, 15)\n        \n        # 2. Derivative\n        diff_signal = np.diff(filtered)\n        \n        # 3. Squaring\n        squared = diff_signal ** 2\n        \n        # 4. Moving average\n        window = int(0.12 * self.sampling_rate)\n        integrated = np.convolve(squared, np.ones(window)/window, mode='same')\n        \n        # 5. Find peaks\n        threshold = np.mean(integrated) + 0.5 * np.std(integrated)\n        peaks, _ = find_peaks(integrated, height=threshold, distance=int(0.2 * self.sampling_rate))\n        \n        return peaks\n    \n    def compute_heart_rate(self, signal: np.ndarray) -> float:\n        \"\"\"Compute heart rate from signal\"\"\"\n        peaks = self.detect_r_peaks(signal)\n        \n        if len(peaks) < 2:\n            return 0.0\n        \n        rr_intervals = np.diff(peaks) / self.sampling_rate  # in seconds\n        mean_rr = np.mean(rr_intervals)\n        heart_rate = 60.0 / mean_rr if mean_rr > 0 else 0.0\n        \n        return heart_rate\n    \n    def compute_snr(self, signal: np.ndarray, noise_signal: np.ndarray = None) -> float:\n        \"\"\"Compute Signal-to-Noise Ratio\"\"\"\n        if noise_signal is None:\n            # Estimate noise from high-frequency components\n            filtered = self.bandpass_filter(signal, 40, 100)\n            noise_signal = filtered\n        \n        signal_power = np.mean(signal ** 2)\n        noise_power = np.mean(noise_signal ** 2)\n        \n        if noise_power == 0:\n            return float('inf')\n        \n        snr = 10 * np.log10(signal_power / noise_power)\n        return snr\n    \n    def process_signal(self, signal: np.ndarray, apply_filters: bool = True) -> np.ndarray:\n        \"\"\"Complete signal processing pipeline\"\"\"\n        processed = signal.copy()\n        \n        if apply_filters:\n            # 1. Remove baseline wander\n            processed = self.remove_baseline_wander(processed)\n            \n            # 2. Bandpass filter\n            processed = self.bandpass_filter(processed)\n            \n            # 3. Notch filter (power line) - CORRECTED: parameter name changed\n            processed = self.notch_filter(processed)\n        \n        # 4. Normalize\n        processed = (processed - np.mean(processed)) / (np.std(processed) + 1e-8)\n        \n        return processed\n    \n    def resample_signal(self, signal: np.ndarray, target_length: int) -> np.ndarray:\n        \"\"\"Resample signal to target length\"\"\"\n        if len(signal) == target_length:\n            return signal\n        \n        x_old = np.linspace(0, 1, len(signal))\n        x_new = np.linspace(0, 1, target_length)\n        \n        f = interp1d(x_old, signal, kind=config.interpolation_method, \n                    bounds_error=False, fill_value='extrapolate')\n        resampled = f(x_new)\n        \n        return resampled\n\n# =============================================================================\n# DATA LOADING MODULE\n# =============================================================================\nclass DataLoader:\n    \"\"\"Load and manage training and test data\"\"\"\n    \n    def __init__(self, config: Config):\n        self.config = config\n        \n    def load_training_data(self) -> Dict:\n        \"\"\"Load all training data\"\"\"\n        train_data = {}\n        \n        # Load CSV metadata - CORRECTED: Using 'id' not 'patient_id'\n        train_csv = pd.read_csv(self.config.train_csv)\n        \n        # Load each patient's data\n        for idx, row in train_csv.iterrows():\n            patient_id = row['id']  # CORRECTED: Changed from 'patient_id' to 'id'\n            \n            # Check if patient folder exists\n            patient_folder = os.path.join(self.config.train_images, str(patient_id))\n            if not os.path.exists(patient_folder):\n                print(f\"Warning: Patient folder not found: {patient_folder}\")\n                continue\n            \n            # Look for PNG images in the patient folder\n            image_files = glob.glob(os.path.join(patient_folder, \"*.png\"))\n            if not image_files:\n                print(f\"Warning: No PNG images found for patient {patient_id}\")\n                continue\n            \n            # Use the first image (could be extended to process all images)\n            image_path = image_files[0]\n            \n            # Load corresponding CSV signal data\n            csv_path = os.path.join(patient_folder, f\"{patient_id}.csv\")\n            csv_data = None\n            if os.path.exists(csv_path):\n                csv_data = pd.read_csv(csv_path)\n            \n            train_data[patient_id] = {\n                'image_path': image_path,\n                'csv_data': csv_data,\n                'metadata': row\n            }\n        \n        print(f\"Loaded {len(train_data)} training samples\")\n        return train_data\n    \n    def load_test_data(self) -> Dict:\n        \"\"\"Load test data\"\"\"\n        test_data = {}\n        \n        # Load test CSV - CORRECTED: Using 'id' not 'patient_id'\n        test_csv = pd.read_csv(self.config.test_csv)\n        \n        # Group by patient ID since test.csv has multiple rows per patient\n        grouped = test_csv.groupby('id')\n        \n        for patient_id, group in grouped:\n            # Load image\n            image_path = os.path.join(self.config.test_images, f\"{patient_id}.png\")\n            if not os.path.exists(image_path):\n                print(f\"Warning: Test image not found: {image_path}\")\n                continue\n            \n            # Get leads for this patient\n            leads = group['lead'].tolist()\n            fs_values = group['fs'].unique()\n            number_of_rows = group['number_of_rows'].unique()\n            \n            test_data[patient_id] = {\n                'image_path': image_path,\n                'leads': leads,\n                'fs': fs_values[0] if len(fs_values) > 0 else config.default_sampling_rate,\n                'number_of_rows': number_of_rows[0] if len(number_of_rows) > 0 else 1000,\n                'metadata': group\n            }\n        \n        print(f\"Loaded {len(test_data)} test samples\")\n        return test_data\n\n# =============================================================================\n# ENHANCED EVALUATION METRIC\n# =============================================================================\nclass ECGSNRMetric:\n    \"\"\"Implementation of the modified SNR metric for ECG reconstruction\"\"\"\n    \n    @staticmethod\n    def align_signals(predicted: np.ndarray, ground_truth: np.ndarray, \n                     max_shift_samples: int = 100) -> Tuple[np.ndarray, int, float]:\n        \"\"\"\n        Align predicted signal with ground truth using cross-correlation\n        \n        Args:\n            predicted: Predicted ECG signal\n            ground_truth: Ground truth ECG signal\n            max_shift_samples: Maximum shift in samples (0.2 seconds at 500 Hz = 100 samples)\n        \n        Returns:\n            aligned_predicted: Aligned predicted signal\n            optimal_shift: Optimal time shift applied\n            vertical_offset: Vertical offset removed\n        \"\"\"\n        # Ensure signals have same length\n        min_len = min(len(predicted), len(ground_truth))\n        predicted = predicted[:min_len]\n        ground_truth = ground_truth[:min_len]\n        \n        # Find optimal time shift using cross-correlation\n        correlation = np.correlate(predicted, ground_truth, mode='full')\n        lags = correlation_lags(len(predicted), len(ground_truth))\n        \n        # Restrict to max shift\n        valid_indices = np.where(np.abs(lags) <= max_shift_samples)[0]\n        correlation = correlation[valid_indices]\n        lags = lags[valid_indices]\n        \n        # Find optimal shift\n        optimal_shift_idx = np.argmax(correlation)\n        optimal_shift = lags[optimal_shift_idx]\n        \n        # Apply time shift\n        if optimal_shift > 0:\n            aligned_predicted = np.concatenate([predicted[optimal_shift:], \n                                              np.zeros(min(optimal_shift, len(predicted)))])\n            ground_truth_aligned = ground_truth[:len(aligned_predicted)]\n        elif optimal_shift < 0:\n            aligned_predicted = np.concatenate([np.zeros(-optimal_shift), \n                                              predicted[:optimal_shift]])\n            ground_truth_aligned = ground_truth[:len(aligned_predicted)]\n        else:\n            aligned_predicted = predicted\n            ground_truth_aligned = ground_truth\n        \n        # Adjust lengths\n        min_len = min(len(aligned_predicted), len(ground_truth_aligned))\n        aligned_predicted = aligned_predicted[:min_len]\n        ground_truth_aligned = ground_truth_aligned[:min_len]\n        \n        # Remove vertical offset\n        vertical_offset = np.mean(aligned_predicted - ground_truth_aligned)\n        aligned_predicted = aligned_predicted - vertical_offset\n        \n        return aligned_predicted, optimal_shift, vertical_offset\n    \n    @staticmethod\n    def compute_snr(predicted: np.ndarray, ground_truth: np.ndarray) -> float:\n        \"\"\"\n        Compute modified SNR between predicted and ground truth signals\n        \n        Args:\n            predicted: Predicted ECG signal (after alignment)\n            ground_truth: Ground truth ECG signal\n        \n        Returns:\n            SNR value in decibels\n        \"\"\"\n        # Align signals\n        aligned_predicted, _, _ = ECGSNRMetric.align_signals(predicted, ground_truth)\n        \n        # Ensure same length\n        min_len = min(len(aligned_predicted), len(ground_truth))\n        aligned_predicted = aligned_predicted[:min_len]\n        ground_truth = ground_truth[:min_len]\n        \n        # Compute signal power\n        signal_power = np.sum(ground_truth ** 2)\n        \n        # Compute error power\n        error = aligned_predicted - ground_truth\n        error_power = np.sum(error ** 2)\n        \n        # Compute SNR\n        if error_power == 0:\n            return float('inf')\n        \n        snr = 10 * np.log10(signal_power / error_power)\n        return snr\n    \n    @staticmethod\n    def compute_snr_all_leads(predicted_signals: Dict[str, np.ndarray], \n                            ground_truth_signals: Dict[str, np.ndarray]) -> float:\n        \"\"\"\n        Compute SNR across all 12 leads for an entire ECG record\n        \n        Args:\n            predicted_signals: Dictionary of predicted signals for each lead\n            ground_truth_signals: Dictionary of ground truth signals for each lead\n        \n        Returns:\n            Combined SNR across all leads in decibels\n        \"\"\"\n        total_signal_power = 0\n        total_error_power = 0\n        \n        for lead in ground_truth_signals.keys():\n            if lead in predicted_signals:\n                pred = predicted_signals[lead]\n                gt = ground_truth_signals[lead]\n                \n                # Align signals\n                aligned_pred, _, _ = ECGSNRMetric.align_signals(pred, gt)\n                \n                # Ensure same length\n                min_len = min(len(aligned_pred), len(gt))\n                aligned_pred = aligned_pred[:min_len]\n                gt = gt[:min_len]\n                \n                # Accumulate powers\n                total_signal_power += np.sum(gt ** 2)\n                total_error_power += np.sum((aligned_pred - gt) ** 2)\n        \n        # Compute combined SNR\n        if total_error_power == 0:\n            return float('inf')\n        \n        snr = 10 * np.log10(total_signal_power / total_error_power)\n        return snr\n\n# =============================================================================\n# SUBMISSION GENERATOR\n# =============================================================================\nclass SubmissionGenerator:\n    \"\"\"Generate submission file in required format\"\"\"\n    \n    def __init__(self, config: Config):\n        self.config = config\n        \n    def create_submission_dataframe(self, predictions: List[Dict]) -> pd.DataFrame:\n        \"\"\"\n        Create submission dataframe in required format\n        \n        Format should be:\n        id,value\n        '62_0_I',0.0\n        '62_1_II',0.3\n        '62_2_I',0.4\n        etc.\n        \"\"\"\n        submission_records = []\n        \n        for pred in predictions:\n            patient_id = pred['patient_id']\n            lead_name = pred['lead_name']\n            signal = pred['signal']\n            \n            # Create ID in format: patient_id_timestep_lead\n            for time_idx, value in enumerate(signal):\n                record_id = f\"{patient_id}_{time_idx}_{lead_name}\"\n                submission_records.append({\n                    'id': record_id,\n                    'value': float(value)\n                })\n        \n        return pd.DataFrame(submission_records)\n    \n    def save_submission(self, predictions: List[Dict], output_path: str = None):\n        \"\"\"Save predictions to submission file\"\"\"\n        if output_path is None:\n            output_path = self.config.submission_file\n        \n        submission_df = self.create_submission_dataframe(predictions)\n        submission_df.to_parquet(output_path, index=False)\n        print(f\"Submission saved to: {output_path}\")\n        print(f\"Submission shape: {submission_df.shape}\")\n        print(f\"Sample records:\")\n        print(submission_df.head())\n        \n        return submission_df\n\n# =============================================================================\n# MAIN PIPELINE\n# =============================================================================\ndef main():\n    \"\"\"Main execution pipeline\"\"\"\n    \n    print(\"=\"*80)\n    print(\"ECG IMAGE DIGITIZATION PIPELINE - ENHANCED VERSION\")\n    print(\"=\"*80)\n    \n    # Initialize components\n    visualizer = ECGVisualizer(config.plots_dir)\n    data_loader = DataLoader(config)\n    signal_processor = EnhancedSignalProcessor(config.default_sampling_rate)\n    submission_generator = SubmissionGenerator(config)\n    \n    # Load data\n    print(\"\\n[1/6] Loading data...\")\n    train_data = data_loader.load_training_data()\n    test_data = data_loader.load_test_data()\n    \n    # EDA and visualization\n    print(\"\\n[2/6] Performing exploratory data analysis...\")\n    if train_data:\n        visualizer.plot_dataset_statistics(train_data)\n        visualizer.plot_signal_examples(train_data)\n    else:\n        print(\"No training data loaded, skipping EDA\")\n    \n    # Process training samples\n    print(\"\\n[3/6] Processing training images...\")\n    processed_train = []\n    \n    for patient_id, data in list(train_data.items())[:5]:  # Process first 5 for demo\n        try:\n            print(f\"  Processing {patient_id}...\")\n            ecg_frame = EnhancedECGFrame(data['image_path'])\n            extracted_leads = ecg_frame.extract_leads()\n            \n            # Visualize processing pipeline for first patient\n            if len(processed_train) == 0:\n                processed_stages = {\n                    'Rotation Corrected': ecg_frame.gray,\n                    'Denoised': cv2.fastNlMeansDenoising(ecg_frame.gray, None, h=10),\n                    'Grid Removed': ecg_frame.preprocess(),\n                    'Leads Extracted': cv2.cvtColor(ecg_frame.preprocess(), cv2.COLOR_GRAY2BGR)\n                }\n                visualizer.plot_processing_pipeline(ecg_frame.image, processed_stages)\n            \n            for lead_name, signal in extracted_leads.items():\n                processed_signal = signal_processor.process_signal(signal)\n                \n                # If ground truth available, compare\n                if data['csv_data'] is not None and lead_name in data['csv_data'].columns:\n                    gt_signal = data['csv_data'][lead_name].dropna().values\n                    \n                    if len(gt_signal) > 0:\n                        # Compute metrics using the ECG SNR metric\n                        snr_metric = ECGSNRMetric()\n                        snr = snr_metric.compute_snr(processed_signal, gt_signal)\n                        \n                        # Also compute traditional metrics\n                        min_len = min(len(processed_signal), len(gt_signal))\n                        mae = mean_absolute_error(gt_signal[:min_len], processed_signal[:min_len])\n                        rmse = np.sqrt(mean_squared_error(gt_signal[:min_len], processed_signal[:min_len]))\n                        corr = np.corrcoef(gt_signal[:min_len], processed_signal[:min_len])[0, 1]\n                        \n                        processed_train.append({\n                            'patient_id': patient_id,\n                            'lead': lead_name,\n                            'mae': mae,\n                            'rmse': rmse,\n                            'correlation': corr,\n                            'snr': snr\n                        })\n                        \n                        print(f\"    {lead_name}: SNR={snr:.2f}dB, MAE={mae:.4f}, RMSE={rmse:.4f}, Corr={corr:.4f}\")\n                        \n                        # Visualize signal comparison for first few leads\n                        if len(processed_train) <= 3:\n                            visualizer.plot_signal_comparison(\n                                gt_signal[:min_len], \n                                processed_signal[:min_len],\n                                lead_name,\n                                f\"signal_comparison_{patient_id}_{lead_name}.png\"\n                            )\n        \n        except Exception as e:\n            print(f\"  Error processing {patient_id}: {str(e)}\")\n            continue\n    \n    # Visualization of results\n    print(\"\\n[4/6] Visualizing prediction quality...\")\n    if processed_train:\n        visualizer.plot_prediction_quality(processed_train)\n        avg_snr = np.mean([p['snr'] for p in processed_train if not np.isinf(p['snr'])])\n        avg_mae = np.mean([p['mae'] for p in processed_train])\n        print(f\"Average SNR: {avg_snr:.2f} dB\")\n        print(f\"Average MAE: {avg_mae:.4f}\")\n    else:\n        print(\"No training predictions to visualize\")\n    \n    # Process test images\n    print(\"\\n[5/6] Processing test images...\")\n    test_results = []\n    \n    for patient_id, data in list(test_data.items())[:10]:  # Process first 10 for demo\n        try:\n            print(f\"  Processing {patient_id}...\")\n            ecg_frame = EnhancedECGFrame(data['image_path'])\n            extracted_leads = ecg_frame.extract_leads()\n            \n            for lead_name, signal in extracted_leads.items():\n                processed_signal = signal_processor.process_signal(signal)\n                \n                # Resample to required length if needed\n                if 'number_of_rows' in data:\n                    target_length = data['number_of_rows']\n                    if len(processed_signal) != target_length:\n                        processed_signal = signal_processor.resample_signal(processed_signal, target_length)\n                \n                test_results.append({\n                    'patient_id': str(patient_id),\n                    'lead_name': lead_name,\n                    'signal': processed_signal\n                })\n                print(f\"    Extracted {lead_name}: {len(processed_signal)} samples\")\n        \n        except Exception as e:\n            print(f\"  Error processing {patient_id}: {str(e)}\")\n            continue\n    \n    # Create submission\n    print(\"\\n[6/6] Creating submission file...\")\n    if test_results:\n        submission_df = submission_generator.save_submission(test_results)\n        \n        # Load sample submission for comparison\n        try:\n            sample_submission = pd.read_parquet(config.sample_submission)\n            print(f\"\\nSample submission shape: {sample_submission.shape}\")\n            print(\"Sample submission head:\")\n            print(sample_submission.head())\n        except Exception as e:\n            print(f\"Could not load sample submission: {e}\")\n    else:\n        print(\"Warning: No test data processed, creating empty submission\")\n        # Create empty submission with correct format\n        empty_df = pd.DataFrame(columns=['id', 'value'])\n        empty_df.to_parquet(config.submission_file, index=False)\n        print(f\"Empty submission saved to: {config.submission_file}\")\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"PIPELINE COMPLETE!\")\n    print(\"=\"*80)\n    print(f\"Total training samples processed: {len(processed_train)}\")\n    print(f\"Total test samples processed: {len(test_results)}\")\n    if processed_train:\n        valid_snr = [p['snr'] for p in processed_train if not np.isinf(p['snr'])]\n        if valid_snr:\n            print(f\"Average SNR: {np.mean(valid_snr):.2f} dB\")\n        print(f\"Average MAE: {np.mean([p['mae'] for p in processed_train]):.4f}\")\n        print(f\"Average Correlation: {np.mean([p['correlation'] for p in processed_train]):.4f}\")\n    print(\"=\"*80)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-22T22:26:32.007675Z","iopub.execute_input":"2026-01-22T22:26:32.008347Z","iopub.status.idle":"2026-01-22T22:30:10.047330Z","shell.execute_reply.started":"2026-01-22T22:26:32.008308Z","shell.execute_reply":"2026-01-22T22:30:10.046191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}