""" ECG Dataset V3 - Matches actual data structure Data locations: - Synthetic: /data/ecg-digitization/synthetic_ecgkit/ - stage1/ (images: syn_XXXXXXXX-YYYY.png) - gt/ (targets: syn_XXXXXXXX.npy) - Kaggle stage1: /data/ecg-digitization/stage1_data/train/ - {id}/ - {id}-0001.png through {id}-0012.png (image versions) - {id}.csv (ground truth) Expert guidance: - "Augmentation: heavily needed, add augmentation according to different type of images (0001-0012)" - "I just do augmentation on all images, but exclude some bad data" """ 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, ConcatDataset import albumentations as A from albumentations.pytorch import ToTensorV2 import random # Target specifications (matching competition Stage 2 input) TARGET_HEIGHT = 1696 TARGET_WIDTH = 4352 # Calibration constants (from winning solution) ZERO_MV = [703.5, 987.5, 1271.5, 1531.5] MV_TO_PIXEL = 78.5 T0, T1 = 235, 4161 # Time crop boundaries def get_heavy_train_transforms(height=TARGET_HEIGHT, width=TARGET_WIDTH): """ Heavy augmentation for ECG images. Expert: "Augmentation: heavily needed" """ return A.Compose([ # Resize to target size A.Resize(height, width), # Geometric - careful not to distort too much A.OneOf([ A.Affine(scale=(0.95, 1.05), translate_percent=(-0.02, 0.02), rotate=(-2, 2), shear=(-2, 2), p=0.5), ], p=0.4), # Color augmentations - HEAVY (simulate different paper/grid colors) A.OneOf([ A.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1, p=1.0), A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3, p=1.0), A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=30, p=1.0), ], p=0.8), # Additional color manipulation A.OneOf([ A.RGBShift(r_shift_limit=20, g_shift_limit=20, b_shift_limit=20, p=0.5), A.ChannelShuffle(p=0.1), A.ToGray(p=0.1), # Some images might be grayscale ], p=0.3), # CLAHE and equalization A.OneOf([ A.CLAHE(clip_limit=4.0, p=0.5), A.Equalize(p=0.3), ], p=0.2), # Noise augmentations - HEAVY (simulate scanner/photo artifacts) A.OneOf([ A.GaussNoise(std_range=(0.02, 0.1), p=0.5), A.ISONoise(color_shift=(0.01, 0.05), intensity=(0.1, 0.5), p=0.5), A.MultiplicativeNoise(multiplier=(0.9, 1.1), p=0.3), ], p=0.5), # Blur (simulate focus issues / motion) A.OneOf([ A.GaussianBlur(blur_limit=(3, 7), p=0.5), A.MotionBlur(blur_limit=(3, 7), p=0.4), A.MedianBlur(blur_limit=5, p=0.2), ], p=0.3), # Quality reduction (simulate jpeg artifacts, fax quality) A.OneOf([ A.ImageCompression(quality_range=(50, 95), p=0.5), A.Downscale(scale_range=(0.5, 0.9), p=0.3), ], p=0.3), # Shadows and lighting A.OneOf([ A.RandomShadow(shadow_roi=(0, 0.5, 1, 1), num_shadows_limit=(1, 2), shadow_dimension=5, p=0.3), A.RandomToneCurve(scale=0.1, p=0.3), ], p=0.2), # Grid-like augmentations (simulate grid distortions) A.OneOf([ A.GridDistortion(num_steps=5, distort_limit=0.1, p=0.2), A.ElasticTransform(alpha=20, sigma=5, p=0.1), ], p=0.15), # Cutout (simulate missing parts / occlusions) A.CoarseDropout( num_holes_range=(1, 8), hole_height_range=(10, 50), hole_width_range=(10, 50), fill=(255, 255, 255), 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(), ]) class SyntheticECGDatasetV3(Dataset): """ Dataset for synthetic ECG images from ecg-image-kit. Data structure: synthetic_dir/ stage1/ (images: syn_XXXXXXXX-YYYY.png) gt/ (targets: syn_XXXXXXXX.npy, shape [4, 3926]) """ def __init__(self, synthetic_dir, transform=None, max_samples=None): self.synthetic_dir = Path(synthetic_dir) self.transform = transform or get_heavy_train_transforms() # Find all GT files (source of truth for samples) gt_dir = self.synthetic_dir / 'gt' self.gt_files = sorted(gt_dir.glob('*.npy')) if max_samples is not None: self.gt_files = self.gt_files[:max_samples] # Build mapping from GT to images self.samples = [] stage1_dir = self.synthetic_dir / 'stage1' for gt_file in self.gt_files: # GT file: syn_00000000.npy -> images: syn_00000000-0001.png, etc. base_name = gt_file.stem # syn_00000000 # Find all image versions for this sample image_files = sorted(stage1_dir.glob(f'{base_name}-*.png')) for img_file in image_files: self.samples.append({ 'image_path': img_file, 'gt_path': gt_file, 'sample_id': base_name, 'image_version': img_file.stem.split('-')[-1], # 0001, 0002, etc. }) print(f"SyntheticECGDatasetV3: {len(self.gt_files)} samples, {len(self.samples)} images") def __len__(self): return len(self.samples) def __getitem__(self, idx): sample = self.samples[idx] # Load image image = cv2.imread(str(sample['image_path'])) if image is None: raise ValueError(f"Could not load image: {sample['image_path']}") image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # Load GT signal (shape: [4, 3926] in mV) gt_signal = np.load(sample['gt_path']) # [4, 3926] # Apply transforms if self.transform: transformed = self.transform(image=image) image = transformed['image'] # Convert GT from mV to normalized pixel Y-coordinates target = self._signal_to_target(gt_signal) return { 'image': image, 'target': torch.from_numpy(target.astype(np.float32)), 'id': sample['sample_id'], 'version': sample['image_version'], } def _signal_to_target(self, signal_mv): """ Convert mV signal to normalized pixel Y-coordinates. signal_mv: [4, W] in mV returns: [4, T1-T0] normalized to [0, 1] """ output_width = T1 - T0 # 3926 # Ensure correct width if signal_mv.shape[1] != output_width: # Resample new_signal = np.zeros((4, output_width), dtype=np.float32) for row in range(4): x_old = np.linspace(0, 1, signal_mv.shape[1]) x_new = np.linspace(0, 1, output_width) new_signal[row] = np.interp(x_new, x_old, signal_mv[row]) signal_mv = new_signal # Convert mV to pixel Y-coordinate target_pixel = np.zeros_like(signal_mv) for row_idx in range(4): target_pixel[row_idx] = ZERO_MV[row_idx] - signal_mv[row_idx] * MV_TO_PIXEL # Normalize to [0, 1] target_normalized = target_pixel / TARGET_HEIGHT target_normalized = np.clip(target_normalized, 0, 1) return target_normalized class KaggleStage1Dataset(Dataset): """ Dataset for Kaggle competition data after Stage 0/1 processing. Data structure: stage1_dir/ {id}/ {id}-0001.png through {id}-0012.png (image versions) {id}.csv (ground truth) Expert: "add augmentation according to different type of images (0001-0012)" """ def __init__(self, stage1_dir, transform=None, image_versions=None, max_samples=None, exclude_bad=True): """ Args: stage1_dir: Path to stage1_data/train/ transform: Augmentation transforms image_versions: List of versions to use (e.g., ['0001', '0002']) or None for all max_samples: Max number of sample IDs (None for all) exclude_bad: Whether to exclude bad samples """ self.stage1_dir = Path(stage1_dir) self.transform = transform or get_heavy_train_transforms() self.image_versions = image_versions # Bad samples to exclude (can be populated after analysis) self.bad_samples = set() # Find all sample directories sample_dirs = sorted([d for d in self.stage1_dir.iterdir() if d.is_dir()]) if max_samples is not None: sample_dirs = sample_dirs[:max_samples] # Build samples list self.samples = [] for sample_dir in sample_dirs: sample_id = sample_dir.name if exclude_bad and sample_id in self.bad_samples: continue # Check for CSV ground truth csv_path = sample_dir / f'{sample_id}.csv' if not csv_path.exists(): continue # Find image versions if image_versions: # Use specified versions for ver in image_versions: img_path = sample_dir / f'{sample_id}-{ver}.png' if img_path.exists(): self.samples.append({ 'image_path': img_path, 'csv_path': csv_path, 'sample_id': sample_id, 'image_version': ver, }) else: # Use all available versions for img_path in sorted(sample_dir.glob(f'{sample_id}-*.png')): ver = img_path.stem.split('-')[-1] self.samples.append({ 'image_path': img_path, 'csv_path': csv_path, 'sample_id': sample_id, 'image_version': ver, }) # Group by sample_id for reference self.sample_ids = list(set(s['sample_id'] for s in self.samples)) print(f"KaggleStage1Dataset: {len(self.sample_ids)} samples, {len(self.samples)} images") def __len__(self): return len(self.samples) def __getitem__(self, idx): sample = self.samples[idx] # Load image image = cv2.imread(str(sample['image_path'])) if image is None: raise ValueError(f"Could not load image: {sample['image_path']}") image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # Load ground truth CSV target = self._csv_to_target(sample['csv_path']) # 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': sample['sample_id'], 'version': sample['image_version'], } def _csv_to_target(self, csv_path): """ Convert ground truth CSV to target array. CSV contains 12 leads in mV. We arrange them into 4 rows matching the ECG layout: - Row 0: [I, aVR, V1, V4] (each 1/4 width) - Row 1: [II, aVL, V2, V5] - Row 2: [III, aVF, V3, V6] - Row 3: [II full rhythm strip] """ df = pd.read_csv(csv_path) lead_layout = [ ['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6'], ] output_width = T1 - T0 # 3926 quarter_width = output_width // 4 # 981 remainder = output_width - (quarter_width * 4) # Handle any remainder target = np.zeros((4, output_width), dtype=np.float32) # Rows 0-2: 4 leads each (2.5 seconds each) for row_idx in range(3): row_signals = [] for lead_idx, lead in enumerate(lead_layout[row_idx]): # Last segment gets the remainder seg_width = quarter_width + (remainder if lead_idx == 3 else 0) if lead in df.columns: signal = df[lead].dropna().values if len(signal) > 0: # Get ~2.5 seconds worth of data (first quarter for short leads) seg_len = len(signal) // 4 if len(signal) > 100 else len(signal) seg = signal[:seg_len] x_old = np.linspace(0, 1, len(seg)) x_new = np.linspace(0, 1, seg_width) signal_resampled = np.interp(x_new, x_old, seg) else: signal_resampled = np.zeros(seg_width) else: signal_resampled = np.zeros(seg_width) row_signals.append(signal_resampled) target[row_idx] = np.concatenate(row_signals) # Row 3: Full II rhythm strip (10 seconds) 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 Y-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 def create_dataloaders_v3( synthetic_dir=None, kaggle_stage1_dir=None, batch_size=4, num_workers=8, synthetic_ratio=0.5, distributed=False, val_split=0.1, image_versions=None, ): """ Create DataLoaders for training. Args: synthetic_dir: Path to synthetic data (e.g., /data/ecg-digitization/synthetic_ecgkit) kaggle_stage1_dir: Path to Kaggle stage1 data (e.g., /data/ecg-digitization/stage1_data/train) batch_size: Batch size per GPU num_workers: Number of data loading workers synthetic_ratio: Ratio of synthetic data (0=all Kaggle, 1=all synthetic) distributed: Whether using DDP val_split: Validation split ratio image_versions: Which image versions to use (None=all) Returns: train_loader, val_loader """ train_transform = get_heavy_train_transforms() val_transform = get_val_transforms() train_datasets = [] val_dataset = None # Load synthetic data if synthetic_dir and synthetic_ratio > 0: synthetic_dataset = SyntheticECGDatasetV3(synthetic_dir, train_transform) train_datasets.append(('synthetic', synthetic_dataset)) # Load Kaggle data if kaggle_stage1_dir: kaggle_dataset = KaggleStage1Dataset( kaggle_stage1_dir, train_transform, image_versions=image_versions ) # Split for validation n_samples = len(kaggle_dataset.sample_ids) n_val = int(n_samples * val_split) # Shuffle sample IDs all_ids = kaggle_dataset.sample_ids.copy() random.seed(42) random.shuffle(all_ids) val_ids = set(all_ids[:n_val]) train_ids = set(all_ids[n_val:]) # Filter samples train_samples = [s for s in kaggle_dataset.samples if s['sample_id'] in train_ids] val_samples = [s for s in kaggle_dataset.samples if s['sample_id'] in val_ids] # Create train/val datasets kaggle_train = KaggleStage1Dataset.__new__(KaggleStage1Dataset) kaggle_train.stage1_dir = kaggle_dataset.stage1_dir kaggle_train.transform = train_transform kaggle_train.samples = train_samples kaggle_train.sample_ids = list(train_ids) kaggle_train.bad_samples = kaggle_dataset.bad_samples kaggle_val = KaggleStage1Dataset.__new__(KaggleStage1Dataset) kaggle_val.stage1_dir = kaggle_dataset.stage1_dir kaggle_val.transform = val_transform kaggle_val.samples = val_samples kaggle_val.sample_ids = list(val_ids) kaggle_val.bad_samples = kaggle_dataset.bad_samples train_datasets.append(('kaggle', kaggle_train)) val_dataset = kaggle_val print(f"Kaggle split: {len(train_samples)} train, {len(val_samples)} val images") # Combine datasets based on ratio if len(train_datasets) == 2: # Mixed: synthetic + kaggle syn_dataset = train_datasets[0][1] kag_dataset = train_datasets[1][1] # Oversample based on ratio if synthetic_ratio >= 0.5: # More synthetic n_kaggle = len(kag_dataset) n_synthetic_target = int(n_kaggle * synthetic_ratio / (1 - synthetic_ratio)) # Repeat synthetic if needed syn_indices = list(range(len(syn_dataset))) * (n_synthetic_target // len(syn_dataset) + 1) syn_indices = syn_indices[:n_synthetic_target] combined_dataset = ConcatDataset([ torch.utils.data.Subset(syn_dataset, syn_indices), kag_dataset ]) else: # More kaggle combined_dataset = ConcatDataset([syn_dataset, kag_dataset]) train_dataset = combined_dataset print(f"Combined dataset: {len(train_dataset)} total samples") elif len(train_datasets) == 1: train_dataset = train_datasets[0][1] else: raise ValueError("No training data provided!") # Create samplers train_sampler = None val_sampler = None if distributed: train_sampler = torch.utils.data.distributed.DistributedSampler( train_dataset, shuffle=True ) 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, persistent_workers=True if num_workers > 0 else False, ) 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, persistent_workers=True if num_workers > 0 else False, ) return train_loader, val_loader if __name__ == '__main__': # Test the datasets print("Testing Dataset V3 classes...") # Test synthetic dataset synthetic_dir = '/data/ecg-digitization/synthetic_ecgkit' if Path(synthetic_dir).exists(): print("\n1. Testing SyntheticECGDatasetV3...") ds = SyntheticECGDatasetV3(synthetic_dir, max_samples=10) sample = ds[0] print(f" Image shape: {sample['image'].shape}") print(f" Target shape: {sample['target'].shape}") print(f" Sample ID: {sample['id']}") print(f" Version: {sample['version']}") # Test Kaggle dataset kaggle_dir = '/data/ecg-digitization/stage1_data/train' if Path(kaggle_dir).exists(): print("\n2. Testing KaggleStage1Dataset...") ds = KaggleStage1Dataset(kaggle_dir, max_samples=10) sample = ds[0] print(f" Image shape: {sample['image'].shape}") print(f" Target shape: {sample['target'].shape}") print(f" Sample ID: {sample['id']}") print(f" Version: {sample['version']}") print("\nDataset V3 tests completed!")