"""
PhysioNet ECG Digitization Challenge - Kaggle Submission Notebook
This notebook extracts time-series ECG data from images using classical computer vision
"""

import os
import numpy as np
import pandas as pd
import cv2
from pathlib import Path
from scipy import signal
from scipy.interpolate import interp1d

class ECGDigitizer:
    """ECG digitization using classical computer vision"""
    
    def __init__(self):
        self.leads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']
    
    def enhance_image(self, image):
        """Enhance ECG image"""
        if len(image.shape) == 3:
            gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
        else:
            gray = image.copy()
        
        # Adaptive thresholding
        binary = cv2.adaptiveThreshold(
            gray, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
            cv2.THRESH_BINARY_INV, 11, 2
        )
        
        # Noise removal
        kernel = np.ones((2,2), np.uint8)
        cleaned = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel)
        
        return cleaned
    
    def segment_lead_regions(self, image):
        """Segment image into lead regions"""
        height, width = image.shape[:2]
        
        regions = []
        lead_section_height = int(height * 0.75)
        row_height = lead_section_height // 4
        col_width = width // 3
        
        for row in range(4):
            for col in range(3):
                y1 = row * row_height
                y2 = (row + 1) * row_height
                x1 = col * col_width
                x2 = (col + 1) * col_width
                
                region = image[y1:y2, x1:x2]
                if region.size > 0:
                    regions.append(region)
        
        return regions[:12]
    
    def extract_signal(self, region, target_length):
        """Extract 1D signal from 2D region"""
        if region.size == 0:
            return np.zeros(target_length)
        
        height, width = region.shape[:2]
        
        # Invert if needed
        if np.mean(region) > 127:
            region = 255 - region
        
        # Column-wise signal extraction
        signal_raw = []
        for col in range(width):
            column = region[:, col]
            dark_pixels = np.where(column > 128)[0]
            if len(dark_pixels) > 0:
                center = np.mean(dark_pixels)
                value = (height - center) / height
                signal_raw.append(value)
            else:
                signal_raw.append(0.5)
        
        # Normalize to ECG range
        signal_raw = np.array(signal_raw)
        signal_raw = (signal_raw - 0.5) * 4.0
        
        # Resample
        if len(signal_raw) != target_length:
            x_old = np.linspace(0, 1, len(signal_raw))
            x_new = np.linspace(0, 1, target_length)
            f = interp1d(x_old, signal_raw, kind='cubic', fill_value='extrapolate')
            signal_resampled = f(x_new)
        else:
            signal_resampled = signal_raw
        
        # Smooth
        window_length = min(11, max(3, len(signal_resampled) // 50))
        if window_length % 2 == 0:
            window_length += 1
        
        try:
            signal_smooth = signal.savgol_filter(signal_resampled, window_length, 3)
        except:
            signal_smooth = signal_resampled
        
        return signal_smooth
    
    def process_image(self, image_path, fs, sig_len):
        """Process a single ECG image"""
        image = cv2.imread(str(image_path))
        if image is None:
            return {lead: np.zeros(sig_len // 4 if lead != 'II' else sig_len) 
                    for lead in self.leads}
        
        enhanced = self.enhance_image(image)
        regions = self.segment_lead_regions(enhanced)
        
        signals_dict = {}
        for i, lead_name in enumerate(self.leads):
            if i < len(regions):
                if lead_name == 'II':
                    target_len = sig_len
                else:
                    target_len = sig_len // 4
                
                ecg_signal = self.extract_signal(regions[i], target_len)
                signals_dict[lead_name] = ecg_signal
            else:
                if lead_name == 'II':
                    signals_dict[lead_name] = np.zeros(sig_len)
                else:
                    signals_dict[lead_name] = np.zeros(sig_len // 4)
        
        return signals_dict


# Main execution
print("PhysioNet ECG Digitization - Starting...")

# Setup paths
data_dir = Path('/kaggle/input/physionet-ecg-image-digitization')
test_csv = data_dir / 'test.csv'
test_dir = data_dir / 'test'

# Initialize
digitizer = ECGDigitizer()

# Load test data
print("Loading test.csv...")
test_df = pd.read_csv(test_csv)
print(f"Test entries: {len(test_df)}")

# Process images
print("Processing ECG images...")
results = []
processed_images = {}

for idx, row in test_df.iterrows():
    record_id = row['id']
    lead = row['lead']
    fs = row['fs']
    num_rows = row['number_of_rows']
    
    # Process each image once
    if record_id not in processed_images:
        image_path = test_dir / f"{record_id}.png"
        sig_len = fs * 10  # 10 seconds for full ECG
        
        if image_path.exists():
            try:
                signals = digitizer.process_image(image_path, fs, sig_len)
                processed_images[record_id] = signals
            except Exception as e:
                print(f"Error on {record_id}: {e}")
                processed_images[record_id] = {
                    l: np.zeros(sig_len // 4 if l != 'II' else sig_len)
                    for l in digitizer.leads
                }
        else:
            processed_images[record_id] = {
                l: np.zeros(num_rows) for l in digitizer.leads
            }
    
    # Get signal for this lead
    ecg_signal = processed_images[record_id].get(lead, np.zeros(num_rows))
    
    # Resample to exact length if needed
    if len(ecg_signal) != num_rows:
        x_old = np.linspace(0, 1, len(ecg_signal))
        x_new = np.linspace(0, 1, num_rows)
        f = interp1d(x_old, ecg_signal, kind='linear', fill_value='extrapolate')
        ecg_signal = f(x_new)
    
    # Create submission rows
    for row_id in range(num_rows):
        results.append({
            'id': f"{record_id}_{row_id}_{lead}",
            'value': float(ecg_signal[row_id])
        })
    
    if (idx + 1) % 100 == 0:
        print(f"Processed {idx + 1}/{len(test_df)} entries")

# Create submission
print("Creating submission file...")
submission_df = pd.DataFrame(results)

# Save submission
submission_df.to_csv('submission.csv', index=False)

print(f"✓ Submission complete!")
print(f"  Shape: {submission_df.shape}")
print(f"  Images processed: {len(processed_images)}")
print(f"  File: submission.csv")

