#!/usr/bin/env python3 """ Synthetic ECG Image Generator for Training This script generates synthetic ECG images from PTB-XL signals using ecg-image-kit. The generated images match the "Stage 1 Rectified" format of the competition data. Target Format: - Image Size: 1696 (H) x 4352 (W) - Layout: 4 rows (3 leads per row + 1 rhythm strip) - Standard 12-lead ECG grid Usage: python generate_synthetic.py --n_samples 200000 --n_workers 16 --output_dir ../data/synthetic """ import os import sys import argparse import numpy as np import pandas as pd import cv2 import wfdb from pathlib import Path from multiprocessing import Pool, cpu_count from tqdm import tqdm import json import random # Add ecg-image-kit to path ECG_IMAGE_KIT_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'ecg-image-kit') sys.path.insert(0, ECG_IMAGE_KIT_PATH) # ecg-image-kit imports (will be imported in worker processes) try: from ecg_image_generator.ecg_plot import ecg_plot except ImportError: print(f"Warning: ecg-image-kit not found at {ECG_IMAGE_KIT_PATH}") print("Please ensure the repository is cloned correctly") # Target image specifications (matching competition Stage 1 output) TARGET_HEIGHT = 1696 TARGET_WIDTH = 4352 LEAD_NAMES = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6'] # Standard ECG row layout ROW_LAYOUT = [ ['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6'], ] # Calibration constants (from competition baseline) ZERO_MV = [703.5, 987.5, 1271.5, 1531.5] # Baseline Y positions for 4 rows MV_TO_PIXEL = 78.5 # mV to pixel conversion class AugmentationConfig: """Configuration for image augmentations to make synthetic data realistic""" def __init__(self, difficulty='medium'): if difficulty == 'easy': self.config = { 'grid_color': 'red', 'grid_opacity': 0.8, 'add_noise': False, 'noise_level': 0, 'add_wrinkles': False, 'add_shadows': False, 'rotate': False, 'add_text': True, 'text_opacity': 0.7, } elif difficulty == 'medium': self.config = { 'grid_color': random.choice(['red', 'green', 'lightred', 'lightgreen']), 'grid_opacity': random.uniform(0.5, 0.9), 'add_noise': True, 'noise_level': random.uniform(0.01, 0.05), 'add_wrinkles': random.random() > 0.5, 'add_shadows': random.random() > 0.5, 'rotate': False, # We want rectified images 'add_text': True, 'text_opacity': random.uniform(0.5, 0.9), } elif difficulty == 'hard': self.config = { 'grid_color': random.choice(['red', 'green', 'lightred', 'lightgreen', 'faded']), 'grid_opacity': random.uniform(0.3, 1.0), 'add_noise': True, 'noise_level': random.uniform(0.02, 0.10), 'add_wrinkles': random.random() > 0.3, 'add_shadows': random.random() > 0.3, 'rotate': False, 'add_text': True, 'text_opacity': random.uniform(0.3, 1.0), } else: self.config = {} def get(self, key, default=None): return self.config.get(key, default) def load_ptbxl_record(record_path): """ Load a PTB-XL record and return the 12-lead signals Args: record_path: Path to the record (without extension) Returns: dict: Dictionary with lead names as keys and signal arrays as values int: Sampling frequency """ try: record = wfdb.rdrecord(record_path) signals = record.p_signal fs = record.fs # PTB-XL uses standard lead names lead_names = record.sig_name signal_dict = {} for i, lead in enumerate(lead_names): # Normalize lead names lead_normalized = lead.upper().replace(' ', '') if lead_normalized in LEAD_NAMES: signal_dict[lead_normalized] = signals[:, i] return signal_dict, int(fs) except Exception as e: print(f"Error loading record {record_path}: {e}") return None, None def generate_ecg_image_matplotlib(signal_dict, fs, output_path, config=None): """ Generate an ECG image using matplotlib (fallback method) Args: signal_dict: Dictionary with lead signals fs: Sampling frequency output_path: Path to save the image config: AugmentationConfig instance """ import matplotlib.pyplot as plt import matplotlib matplotlib.use('Agg') # Create figure with ECG paper background fig, axes = plt.subplots(4, 1, figsize=(43.52, 16.96), dpi=100) # ECG paper colors if config: grid_color = config.get('grid_color', 'red') if grid_color in ['red', 'lightred']: major_color = '#ffcccc' minor_color = '#ffe6e6' else: major_color = '#ccffcc' minor_color = '#e6ffe6' else: major_color = '#ffcccc' minor_color = '#ffe6e6' for row_idx, ax in enumerate(axes[:3]): ax.set_facecolor('white') # Draw grid ax.grid(True, which='major', color=major_color, linewidth=0.5) ax.grid(True, which='minor', color=minor_color, linewidth=0.25) ax.minorticks_on() # Plot 4 leads per row leads_in_row = ROW_LAYOUT[row_idx] samples_per_lead = int(2.5 * fs) # 2.5 seconds per lead for lead_idx, lead in enumerate(leads_in_row): if lead in signal_dict: signal = signal_dict[lead][:samples_per_lead] x_offset = lead_idx * samples_per_lead t = np.arange(len(signal)) + x_offset ax.plot(t, signal, 'k-', linewidth=0.8) ax.set_xlim(0, 4 * samples_per_lead) ax.set_ylim(-2, 2) ax.set_ylabel(f'Row {row_idx + 1}') ax.tick_params(labelbottom=False) # Row 4: Full II rhythm strip ax = axes[3] ax.set_facecolor('white') ax.grid(True, which='major', color=major_color, linewidth=0.5) ax.grid(True, which='minor', color=minor_color, linewidth=0.25) ax.minorticks_on() if 'II' in signal_dict: signal = signal_dict['II'][:int(10 * fs)] # 10 seconds ax.plot(signal, 'k-', linewidth=0.8) ax.set_xlim(0, int(10 * fs)) ax.set_ylim(-2, 2) ax.set_ylabel('II (Rhythm)') ax.set_xlabel('Samples') plt.tight_layout() plt.savefig(output_path, dpi=100, bbox_inches='tight', facecolor='white') plt.close(fig) # Resize to target size img = cv2.imread(output_path) if img is not None: img_resized = cv2.resize(img, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR) cv2.imwrite(output_path, img_resized) def generate_ecg_image_ecgkit(signal_dict, fs, output_path, config=None): """ Generate an ECG image using ecg-image-kit Args: signal_dict: Dictionary with lead signals fs: Sampling frequency output_path: Path to save the image config: AugmentationConfig instance """ try: from ecg_image_generator.ecg_plot import ecg_plot # Prepare signals in correct order signals = [] for lead in LEAD_NAMES: if lead in signal_dict: signals.append(signal_dict[lead]) else: # Fill missing leads with zeros if signals: signals.append(np.zeros_like(signals[0])) else: signals.append(np.zeros(int(10 * fs))) signals = np.array(signals).T # Shape: (samples, 12) # Generate plot with ecg-image-kit ecg_plot( signals, sample_rate=fs, lead_index=LEAD_NAMES, columns=4, row_height=4, show_lead_name=True, show_grid=True, show_separate_line=True, output_file=str(output_path), dpi=100, ) # Resize to target size img = cv2.imread(str(output_path)) if img is not None: img_resized = cv2.resize(img, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR) cv2.imwrite(str(output_path), img_resized) except Exception as e: print(f"ecg-image-kit failed, using matplotlib fallback: {e}") generate_ecg_image_matplotlib(signal_dict, fs, output_path, config) def apply_augmentations(image_path, config): """ Apply augmentations to make synthetic images more realistic Args: image_path: Path to the image config: AugmentationConfig instance """ img = cv2.imread(str(image_path)) if img is None: return # Add Gaussian noise if config.get('add_noise', False): noise_level = config.get('noise_level', 0.02) noise = np.random.normal(0, noise_level * 255, img.shape).astype(np.float32) img = np.clip(img.astype(np.float32) + noise, 0, 255).astype(np.uint8) # Add subtle shadows (wrinkles simulation) if config.get('add_shadows', False): shadow_mask = np.ones(img.shape[:2], dtype=np.float32) for _ in range(random.randint(1, 3)): x1 = random.randint(0, img.shape[1]) x2 = random.randint(0, img.shape[1]) y1 = random.randint(0, img.shape[0]) y2 = random.randint(0, img.shape[0]) # Create gradient shadow shadow_intensity = random.uniform(0.85, 0.95) cv2.line(shadow_mask, (x1, y1), (x2, y2), shadow_intensity, random.randint(20, 50)) shadow_mask = cv2.GaussianBlur(shadow_mask, (21, 21), 0) img = (img * shadow_mask[:, :, np.newaxis]).astype(np.uint8) # Add wrinkle-like texture if config.get('add_wrinkles', False): # Create subtle texture pattern texture = np.random.normal(1.0, 0.02, img.shape[:2]).astype(np.float32) texture = cv2.GaussianBlur(texture, (5, 5), 0) img = np.clip(img * texture[:, :, np.newaxis], 0, 255).astype(np.uint8) cv2.imwrite(str(image_path), img) def signal_to_target(signal_dict, fs, target_width): """ Convert signal dictionary to normalized target values for training This creates the ground truth Y-coordinates that the model should predict. Args: signal_dict: Dictionary with lead signals fs: Sampling frequency target_width: Width of the target output Returns: np.ndarray: Shape (4, target_width) with normalized signal values """ # Time crop boundaries (from competition baseline) T0, T1 = 235, 4161 output_width = T1 - T0 # 3926 pixels target = np.zeros((4, output_width), dtype=np.float32) # Process rows 0-2 (4 leads each, 2.5s per lead) for row_idx in range(3): leads_in_row = ROW_LAYOUT[row_idx] quarter_width = output_width // 4 for lead_idx, lead in enumerate(leads_in_row): if lead in signal_dict: signal = signal_dict[lead] # Resample to fit the quarter width samples_needed = int(2.5 * fs) if len(signal) >= samples_needed: signal_segment = signal[:samples_needed] else: signal_segment = np.pad(signal, (0, samples_needed - len(signal))) # Interpolate to target width x_old = np.linspace(0, 1, len(signal_segment)) x_new = np.linspace(0, 1, quarter_width) signal_resampled = np.interp(x_new, x_old, signal_segment) start_idx = lead_idx * quarter_width end_idx = start_idx + quarter_width target[row_idx, start_idx:end_idx] = signal_resampled # Row 3: Full II rhythm strip (10 seconds) if 'II' in signal_dict: signal = signal_dict['II'] samples_needed = int(10 * fs) if len(signal) >= samples_needed: signal_segment = signal[:samples_needed] else: signal_segment = np.pad(signal, (0, samples_needed - len(signal))) x_old = np.linspace(0, 1, len(signal_segment)) x_new = np.linspace(0, 1, output_width) target[3, :] = np.interp(x_new, x_old, signal_segment) return target def process_single_record(args): """ Process a single PTB-XL record and generate synthetic image + target Args: args: Tuple of (record_path, output_dir, idx, difficulty) Returns: bool: Success status """ record_path, output_dir, idx, difficulty = args try: # Load record signal_dict, fs = load_ptbxl_record(record_path) if signal_dict is None or fs is None: return False # Create unique output name record_name = Path(record_path).stem output_name = f"synth_{idx:06d}_{record_name}" image_path = Path(output_dir) / 'images' / f"{output_name}.png" target_path = Path(output_dir) / 'targets' / f"{output_name}.npy" # Skip if already exists if image_path.exists() and target_path.exists(): return True # Create augmentation config config = AugmentationConfig(difficulty) # Generate image try: generate_ecg_image_ecgkit(signal_dict, fs, image_path, config) except: generate_ecg_image_matplotlib(signal_dict, fs, image_path, config) # Apply augmentations apply_augmentations(image_path, config) # Generate target (ground truth signal in pixel coordinates) target = signal_to_target(signal_dict, fs, TARGET_WIDTH) np.save(target_path, target) return True except Exception as e: print(f"Error processing {record_path}: {e}") return False def find_ptbxl_records(ptbxl_dir): """ Find all PTB-XL records Args: ptbxl_dir: Path to PTB-XL directory Returns: list: List of record paths (without extensions) """ records = [] ptbxl_path = Path(ptbxl_dir) # PTB-XL structure: records/XXX00000/ for folder in sorted(ptbxl_path.rglob('*')): if folder.is_dir() and folder.name.startswith('records'): continue # Look for .hea files for hea_file in folder.glob('*.hea'): record_path = str(hea_file)[:-4] # Remove .hea extension records.append(record_path) # Also check physionet structure physionet_path = ptbxl_path / 'physionet.org' / 'files' / 'ptb-xl' / '1.0.3' if physionet_path.exists(): for folder in sorted(physionet_path.rglob('records*')): for subfolder in sorted(folder.iterdir()): if subfolder.is_dir(): for hea_file in subfolder.glob('*.hea'): record_path = str(hea_file)[:-4] if record_path not in records: records.append(record_path) return records def main(): parser = argparse.ArgumentParser(description='Generate synthetic ECG images from PTB-XL') parser.add_argument('--ptbxl_dir', type=str, default='../data/ptbxl', help='Path to PTB-XL dataset directory') parser.add_argument('--output_dir', type=str, default='../data/synthetic', help='Output directory for synthetic data') parser.add_argument('--n_samples', type=int, default=200000, help='Number of synthetic samples to generate') parser.add_argument('--n_workers', type=int, default=None, help='Number of parallel workers (default: CPU count)') parser.add_argument('--difficulty', type=str, default='medium', choices=['easy', 'medium', 'hard', 'mixed'], help='Augmentation difficulty level') args = parser.parse_args() # Setup paths output_dir = Path(args.output_dir) (output_dir / 'images').mkdir(parents=True, exist_ok=True) (output_dir / 'targets').mkdir(parents=True, exist_ok=True) # Find PTB-XL records print(f"Searching for PTB-XL records in {args.ptbxl_dir}...") records = find_ptbxl_records(args.ptbxl_dir) print(f"Found {len(records)} PTB-XL records") if len(records) == 0: print("No records found! Please check the PTB-XL directory path.") print("Expected structure: ptbxl/physionet.org/files/ptb-xl/1.0.3/records*/") return # Create task list (repeat records to reach n_samples) n_workers = args.n_workers or cpu_count() tasks = [] for i in range(args.n_samples): record_idx = i % len(records) if args.difficulty == 'mixed': difficulty = random.choice(['easy', 'medium', 'hard']) else: difficulty = args.difficulty tasks.append((records[record_idx], str(output_dir), i, difficulty)) print(f"Generating {args.n_samples} synthetic images using {n_workers} workers...") # Process in parallel with Pool(n_workers) as pool: results = list(tqdm( pool.imap(process_single_record, tasks), total=len(tasks), desc="Generating" )) success_count = sum(results) print(f"\nGeneration complete!") print(f"Successfully generated: {success_count}/{len(tasks)} images") # Save metadata metadata = { 'n_samples': success_count, 'target_height': TARGET_HEIGHT, 'target_width': TARGET_WIDTH, 'lead_layout': ROW_LAYOUT, 'zero_mv': ZERO_MV, 'mv_to_pixel': MV_TO_PIXEL, } with open(output_dir / 'metadata.json', 'w') as f: json.dump(metadata, f, indent=2) print(f"Metadata saved to {output_dir / 'metadata.json'}") if __name__ == '__main__': main()