| """ |
| 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_HEIGHT = 1696 |
| TARGET_WIDTH = 4352 |
|
|
| |
| ZERO_MV = [703.5, 987.5, 1271.5, 1531.5] |
| MV_TO_PIXEL = 78.5 |
| T0, T1 = 235, 4161 |
|
|
|
|
| def get_heavy_train_transforms(height=TARGET_HEIGHT, width=TARGET_WIDTH): |
| """ |
| Heavy augmentation for ECG images. |
| Expert: "Augmentation: heavily needed" |
| """ |
| return A.Compose([ |
| |
| A.Resize(height, width), |
| |
| |
| 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), |
| |
| |
| 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), |
| |
| |
| 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), |
| ], p=0.3), |
| |
| |
| A.OneOf([ |
| A.CLAHE(clip_limit=4.0, p=0.5), |
| A.Equalize(p=0.3), |
| ], p=0.2), |
| |
| |
| 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), |
| |
| |
| 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), |
| |
| |
| 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), |
| |
| |
| 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), |
| |
| |
| 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), |
| |
| |
| A.CoarseDropout( |
| num_holes_range=(1, 8), hole_height_range=(10, 50), hole_width_range=(10, 50), |
| fill=(255, 255, 255), p=0.2 |
| ), |
| |
| |
| 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() |
| |
| |
| 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] |
| |
| |
| self.samples = [] |
| stage1_dir = self.synthetic_dir / 'stage1' |
| |
| for gt_file in self.gt_files: |
| |
| base_name = gt_file.stem |
| |
| |
| 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], |
| }) |
| |
| 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] |
| |
| |
| 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) |
| |
| |
| gt_signal = np.load(sample['gt_path']) |
| |
| |
| if self.transform: |
| transformed = self.transform(image=image) |
| image = transformed['image'] |
| |
| |
| 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 |
| |
| |
| if signal_mv.shape[1] != output_width: |
| |
| 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 |
| |
| |
| 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 |
| |
| |
| 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 |
| |
| |
| self.bad_samples = set() |
| |
| |
| 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] |
| |
| |
| self.samples = [] |
| for sample_dir in sample_dirs: |
| sample_id = sample_dir.name |
| |
| if exclude_bad and sample_id in self.bad_samples: |
| continue |
| |
| |
| csv_path = sample_dir / f'{sample_id}.csv' |
| if not csv_path.exists(): |
| continue |
| |
| |
| if image_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: |
| |
| 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, |
| }) |
| |
| |
| 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] |
| |
| |
| 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) |
| |
| |
| target = self._csv_to_target(sample['csv_path']) |
| |
| |
| 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 |
| quarter_width = output_width // 4 |
| remainder = output_width - (quarter_width * 4) |
| |
| target = np.zeros((4, output_width), dtype=np.float32) |
| |
| |
| for row_idx in range(3): |
| row_signals = [] |
| for lead_idx, lead in enumerate(lead_layout[row_idx]): |
| |
| 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: |
| |
| 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) |
| |
| |
| 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) |
| |
| |
| 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 |
| |
| |
| if synthetic_dir and synthetic_ratio > 0: |
| synthetic_dataset = SyntheticECGDatasetV3(synthetic_dir, train_transform) |
| train_datasets.append(('synthetic', synthetic_dataset)) |
| |
| |
| if kaggle_stage1_dir: |
| kaggle_dataset = KaggleStage1Dataset( |
| kaggle_stage1_dir, |
| train_transform, |
| image_versions=image_versions |
| ) |
| |
| |
| n_samples = len(kaggle_dataset.sample_ids) |
| n_val = int(n_samples * val_split) |
| |
| |
| 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:]) |
| |
| |
| 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] |
| |
| |
| 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") |
| |
| |
| if len(train_datasets) == 2: |
| |
| syn_dataset = train_datasets[0][1] |
| kag_dataset = train_datasets[1][1] |
| |
| |
| if synthetic_ratio >= 0.5: |
| |
| n_kaggle = len(kag_dataset) |
| n_synthetic_target = int(n_kaggle * synthetic_ratio / (1 - synthetic_ratio)) |
| |
| 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: |
| |
| 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!") |
| |
| |
| 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 |
| ) |
| |
| |
| 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__': |
| |
| print("Testing Dataset V3 classes...") |
| |
| |
| 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']}") |
| |
| |
| 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!") |
|
|