Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
21 kB
"""
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!")