""" ECG Dataset and DataLoader for Training Handles: 1. Synthetic data loading (generated from PTB-XL) 2. Competition data loading (Kaggle + Stage 0/1 rectified) 3. Data augmentation pipeline 4. Mixed precision support """ import os import torch import numpy as np import cv2 import pandas as pd from pathlib import Path from torch.utils.data import Dataset, DataLoader import albumentations as A from albumentations.pytorch import ToTensorV2 # Target specifications (matching competition) TARGET_HEIGHT = 1696 TARGET_WIDTH = 4352 # Calibration constants ZERO_MV = [703.5, 987.5, 1271.5, 1531.5] MV_TO_PIXEL = 78.5 T0, T1 = 235, 4161 # Time crop boundaries def get_train_transforms(height=TARGET_HEIGHT, width=TARGET_WIDTH): """ Training augmentations for ECG images. Careful not to distort the signal too much - we need to preserve the vertical (mV) relationship. """ return A.Compose([ # Resize to target size A.Resize(height, width), # Color augmentations (simulate different paper/ink qualities) A.OneOf([ A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05, p=0.8), A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.8), A.HueSaturationValue(hue_shift_limit=10, sat_shift_limit=20, val_shift_limit=20, p=0.8), ], p=0.7), # Noise augmentations (simulate scanner artifacts) A.OneOf([ A.GaussNoise(var_limit=(5, 30), p=0.5), A.ISONoise(color_shift=(0.01, 0.05), intensity=(0.1, 0.3), p=0.5), ], p=0.3), # Blur (simulate focus issues) A.OneOf([ A.GaussianBlur(blur_limit=(3, 5), p=0.5), A.MotionBlur(blur_limit=3, p=0.3), ], p=0.2), # Quality reduction A.OneOf([ A.ImageCompression(quality_lower=70, quality_upper=100, p=0.3), A.Downscale(scale_min=0.8, scale_max=0.95, p=0.3), ], p=0.2), # Normalize and convert to tensor A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2(), ]) def get_val_transforms(height=TARGET_HEIGHT, width=TARGET_WIDTH): """ Validation transforms - just resize and normalize. """ return A.Compose([ A.Resize(height, width), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2(), ]) def get_simple_transforms(height=TARGET_HEIGHT, width=TARGET_WIDTH): """ Simple transforms without ImageNet normalization. Used when model expects [0, 1] normalized input. """ return A.Compose([ A.Resize(height, width), A.ToFloat(max_value=255.0), ToTensorV2(), ]) class SyntheticECGDataset(Dataset): """ Dataset for synthetic ECG images generated from PTB-XL. Expected directory structure: synthetic_dir/ images/ synth_000001_xxxxx.png ... targets/ synth_000001_xxxxx.npy ... Args: synthetic_dir: Path to synthetic data directory transform: Albumentations transform max_samples: Maximum number of samples (None for all) """ def __init__(self, synthetic_dir, transform=None, max_samples=None): self.synthetic_dir = Path(synthetic_dir) self.transform = transform or get_train_transforms() # Find all image files image_dir = self.synthetic_dir / 'images' self.image_files = sorted(image_dir.glob('*.png')) if max_samples is not None: self.image_files = self.image_files[:max_samples] print(f"Found {len(self.image_files)} synthetic images") def __len__(self): return len(self.image_files) def __getitem__(self, idx): # Load image image_path = self.image_files[idx] image = cv2.imread(str(image_path)) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # Load target target_path = self.synthetic_dir / 'targets' / f"{image_path.stem}.npy" target = np.load(target_path) # Shape: [4, W] # Apply transforms if self.transform: transformed = self.transform(image=image) image = transformed['image'] # Convert target to normalized coordinates (0-1) # Target is in mV, need to convert to pixel Y-coordinates then normalize target_pixel = np.zeros_like(target) for row_idx in range(4): # Convert mV to pixel Y-coordinate target_pixel[row_idx] = ZERO_MV[row_idx] - target[row_idx] * MV_TO_PIXEL # Normalize to [0, 1] target_normalized = target_pixel / TARGET_HEIGHT target_normalized = np.clip(target_normalized, 0, 1) # Resize target to match output width output_width = T1 - T0 # Cropped width if target_normalized.shape[1] != output_width: target_resized = np.zeros((4, output_width), dtype=np.float32) for row_idx in range(4): x_old = np.linspace(0, 1, target_normalized.shape[1]) x_new = np.linspace(0, 1, output_width) target_resized[row_idx] = np.interp(x_new, x_old, target_normalized[row_idx]) target_normalized = target_resized return { 'image': image, 'target': torch.from_numpy(target_normalized.astype(np.float32)), 'id': image_path.stem, } class KaggleECGDataset(Dataset): """ Dataset for Kaggle competition data (after Stage 0/1 rectification). Expected structure: kaggle_dir/ train/ {id}/ {id}-0001.png {id}.csv stage1_output/ {id}.png (rectified images) Args: kaggle_dir: Path to Kaggle data directory stage1_dir: Path to Stage 1 rectified images df: DataFrame with image IDs and metadata transform: Albumentations transform """ def __init__(self, kaggle_dir, stage1_dir, df, transform=None): self.kaggle_dir = Path(kaggle_dir) self.stage1_dir = Path(stage1_dir) self.df = df self.transform = transform or get_train_transforms() # Get unique image IDs self.image_ids = df['id'].astype(str).unique().tolist() print(f"Found {len(self.image_ids)} Kaggle images") def __len__(self): return len(self.image_ids) def __getitem__(self, idx): image_id = self.image_ids[idx] # Try to load rectified image first stage1_path = self.stage1_dir / f"{image_id}.png" if stage1_path.exists(): image = cv2.imread(str(stage1_path)) else: # Fall back to original image original_path = self.kaggle_dir / 'train' / image_id / f"{image_id}-0001.png" image = cv2.imread(str(original_path)) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # Load ground truth CSV csv_path = self.kaggle_dir / 'train' / image_id / f"{image_id}.csv" if csv_path.exists(): gt_df = pd.read_csv(csv_path) target = self._csv_to_target(gt_df) else: # No ground truth, return zeros target = np.zeros((4, T1 - T0), dtype=np.float32) # Apply transforms if self.transform: transformed = self.transform(image=image) image = transformed['image'] return { 'image': image, 'target': torch.from_numpy(target.astype(np.float32)), 'id': image_id, } def _csv_to_target(self, df): """Convert ground truth CSV to target array.""" lead_layout = [ ['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6'], ] output_width = T1 - T0 target = np.zeros((4, output_width), dtype=np.float32) # Rows 0-2: 4 leads each for row_idx in range(3): row_signals = [] for lead in lead_layout[row_idx]: if lead in df.columns: signal = df[lead].dropna().values # Resample to quarter width quarter_width = output_width // 4 if len(signal) > 0: x_old = np.linspace(0, 1, len(signal)) x_new = np.linspace(0, 1, quarter_width) signal_resampled = np.interp(x_new, x_old, signal) else: signal_resampled = np.zeros(quarter_width) row_signals.append(signal_resampled) else: row_signals.append(np.zeros(output_width // 4)) # Concatenate leads for this row target[row_idx] = np.concatenate(row_signals) # Row 3: Full II rhythm strip if 'II' in df.columns: signal_ii = df['II'].dropna().values if len(signal_ii) > 0: x_old = np.linspace(0, 1, len(signal_ii)) x_new = np.linspace(0, 1, output_width) target[3] = np.interp(x_new, x_old, signal_ii) # Convert mV to normalized pixel coordinates target_pixel = np.zeros_like(target) for row_idx in range(4): target_pixel[row_idx] = ZERO_MV[row_idx] - target[row_idx] * MV_TO_PIXEL target_normalized = target_pixel / TARGET_HEIGHT target_normalized = np.clip(target_normalized, 0, 1) return target_normalized class MixedECGDataset(Dataset): """ Combined dataset for training on both synthetic and real data. Useful for fine-tuning: mix real data with synthetic to prevent catastrophic forgetting. Args: synthetic_dataset: SyntheticECGDataset instance kaggle_dataset: KaggleECGDataset instance synthetic_ratio: Ratio of synthetic samples (0-1) """ def __init__(self, synthetic_dataset, kaggle_dataset, synthetic_ratio=0.2): self.synthetic_dataset = synthetic_dataset self.kaggle_dataset = kaggle_dataset self.synthetic_ratio = synthetic_ratio # Create indices for mixed sampling self.n_synthetic = len(synthetic_dataset) self.n_kaggle = len(kaggle_dataset) # Calculate number of samples to draw from each total_samples = self.n_kaggle + int(self.n_kaggle * synthetic_ratio / (1 - synthetic_ratio)) self.n_synthetic_samples = int(total_samples * synthetic_ratio) self.n_kaggle_samples = total_samples - self.n_synthetic_samples print(f"Mixed dataset: {self.n_kaggle_samples} Kaggle + {self.n_synthetic_samples} synthetic") def __len__(self): return self.n_kaggle_samples + self.n_synthetic_samples def __getitem__(self, idx): if idx < self.n_kaggle_samples: # Real data real_idx = idx % self.n_kaggle return self.kaggle_dataset[real_idx] else: # Synthetic data synthetic_idx = (idx - self.n_kaggle_samples) % self.n_synthetic return self.synthetic_dataset[synthetic_idx] def create_dataloaders( synthetic_dir=None, kaggle_dir=None, stage1_dir=None, train_df=None, val_df=None, batch_size=8, num_workers=4, synthetic_ratio=0.2, distributed=False, ): """ Create DataLoaders for training. Args: synthetic_dir: Path to synthetic data kaggle_dir: Path to Kaggle competition data stage1_dir: Path to Stage 1 rectified images train_df: Training DataFrame val_df: Validation DataFrame batch_size: Batch size per GPU num_workers: Number of data loading workers synthetic_ratio: Ratio of synthetic data in mixed training distributed: Whether using DDP (adds DistributedSampler) Returns: train_loader, val_loader """ train_transform = get_train_transforms() val_transform = get_val_transforms() # Create training dataset if synthetic_dir and kaggle_dir: # Mixed training synthetic_dataset = SyntheticECGDataset(synthetic_dir, train_transform) kaggle_dataset = KaggleECGDataset(kaggle_dir, stage1_dir, train_df, train_transform) train_dataset = MixedECGDataset(synthetic_dataset, kaggle_dataset, synthetic_ratio) elif synthetic_dir: # Synthetic only (pre-training) train_dataset = SyntheticECGDataset(synthetic_dir, train_transform) elif kaggle_dir: # Kaggle only (fine-tuning) train_dataset = KaggleECGDataset(kaggle_dir, stage1_dir, train_df, train_transform) else: raise ValueError("Must provide either synthetic_dir or kaggle_dir") # Create validation dataset (always Kaggle if available) val_dataset = None if val_df is not None and kaggle_dir: val_dataset = KaggleECGDataset(kaggle_dir, stage1_dir, val_df, val_transform) # Create samplers for distributed training train_sampler = None val_sampler = None if distributed: train_sampler = torch.utils.data.distributed.DistributedSampler(train_dataset) if val_dataset: val_sampler = torch.utils.data.distributed.DistributedSampler(val_dataset, shuffle=False) # Create loaders train_loader = DataLoader( train_dataset, batch_size=batch_size, shuffle=(train_sampler is None), sampler=train_sampler, num_workers=num_workers, pin_memory=True, drop_last=True, ) val_loader = None if val_dataset: val_loader = DataLoader( val_dataset, batch_size=batch_size, shuffle=False, sampler=val_sampler, num_workers=num_workers, pin_memory=True, ) return train_loader, val_loader if __name__ == '__main__': # Test datasets print("Testing Dataset classes...") # Test transforms print("\n1. Testing transforms...") train_tf = get_train_transforms(424, 1088) # Smaller for testing val_tf = get_val_transforms(424, 1088) # Create dummy image dummy_image = np.random.randint(0, 255, (424, 1088, 3), dtype=np.uint8) result = train_tf(image=dummy_image) print(f" Transform output shape: {result['image'].shape}") print("\nDataset tests completed!")