#!/usr/bin/env python3 """ Fast 0001-type synthetic ECG image generator. Imports ecg-image-kit dependencies ONCE, then generates images in a loop. This avoids the ~15s TensorFlow import overhead per image. Usage: python gen_0001_fast.py --n_samples 21799 --n_workers 8 """ import os os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' # Suppress TF warnings import sys import argparse import numpy as np import cv2 import random import json from pathlib import Path from tqdm import tqdm import shutil import wfdb from scipy import signal as scipy_signal import warnings warnings.filterwarnings('ignore') # Set up paths BEFORE importing ecg-image-kit PROJECT_ROOT = Path(__file__).parent.parent ECG_IMAGE_KIT = PROJECT_ROOT / 'ecg-image-kit' / 'codes' / 'ecg-image-generator' PTBXL_PATH = Path('/data/ecg-digitization/ptbxl/physionet.org/files/ptb-xl/1.0.3') # Add ecg-image-kit to path sys.path.insert(0, str(ECG_IMAGE_KIT)) # Now import ecg-image-kit components print("Loading ecg-image-kit (this takes ~30s due to TensorFlow)...") try: from ecg_plot import ECGPlot print(" ECGPlot loaded") except ImportError as e: print(f" Warning: Could not import ECGPlot: {e}") ECGPlot = None # Output specs - match Kaggle format RAW_HEIGHT = 1700 RAW_WIDTH = 2200 # GT extraction parameters T0 = 235 T1 = 4161 OUTPUT_GT_WIDTH = T1 - T0 # 3926 def find_ptbxl_records(ptbxl_dir): """Find all PTB-XL record files (500Hz high-res).""" records = [] ptbxl_dir = Path(ptbxl_dir) records500 = ptbxl_dir / 'records500' if records500.exists(): for subdir in sorted(records500.iterdir()): if subdir.is_dir(): for hea in subdir.glob('*_hr.hea'): dat = hea.with_suffix('.dat') if dat.exists(): records.append((str(hea), str(dat))) return records def extract_gt_signal(hea_file, output_width=3926): """Extract ground truth signal from PTB-XL record.""" try: record_path = hea_file.replace('.hea', '') record = wfdb.rdrecord(record_path) signals = record.p_signal sig_names = record.sig_name fs = record.fs if fs != 500: num_samples = int(signals.shape[0] * 500 / fs) signals = scipy_signal.resample(signals, num_samples, axis=0) lead_signals = {} for i, name in enumerate(sig_names): lead_signals[name.upper()] = signals[:, i] row_leads = [ ['I', 'AVR', 'V1', 'V4'], ['II', 'AVL', 'V2', 'V5'], ['III', 'AVF', 'V3', 'V6'], ] segment_len = output_width // 3 gt_signal = np.zeros((4, output_width), dtype=np.float32) for row_idx in range(3): for col_idx in range(3): lead_name = row_leads[row_idx][col_idx] if lead_name in lead_signals: lead_data = lead_signals[lead_name] start = col_idx * 1250 end = min(start + 1250, len(lead_data)) segment = lead_data[start:end] col_start = col_idx * segment_len col_end = (col_idx + 1) * segment_len if col_idx < 2 else output_width if len(segment) > 0: x_old = np.linspace(0, 1, len(segment)) x_new = np.linspace(0, 1, col_end - col_start) gt_signal[row_idx, col_start:col_end] = np.interp(x_new, x_old, segment) if 'II' in lead_signals: lead_ii = lead_signals['II'] x_old = np.linspace(0, 1, len(lead_ii)) x_new = np.linspace(0, 1, output_width) gt_signal[3, :] = np.interp(x_new, x_old, lead_ii) return gt_signal except Exception as e: return None def generate_ecg_image_native(hea_file, dat_file, output_path, resolution=200): """Generate ECG image using native ecg-image-kit plotting. This is a simplified version that generates clean red-grid 0001-style images. """ try: # Read the ECG data record_path = hea_file.replace('.hea', '') record = wfdb.rdrecord(record_path) signals = record.p_signal sig_names = [n.upper() for n in record.sig_name] fs = record.fs # Create the plot using matplotlib directly (faster than subprocess) import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt # Standard 12-lead layout fig, axes = plt.subplots(4, 1, figsize=(11, 8.5), dpi=resolution) # Grid color: red (0001 style) grid_color = '#ffcccc' line_color = 'black' row_leads = [ ['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6'], ] # Map lead names lead_map = {} for i, name in enumerate(sig_names): lead_map[name] = signals[:, i] # Plot first 3 rows (3 leads each, 2.5s segments) for row in range(3): ax = axes[row] ax.set_facecolor('white') ax.grid(True, color=grid_color, linewidth=0.5) ax.set_xlim(0, 10) # 10 seconds total (4 leads * 2.5s) ax.set_ylim(-2, 2) # mV range for col, lead_name in enumerate(row_leads[row]): lead_key = lead_name.upper() if lead_key in lead_map: data = lead_map[lead_key] # 2.5s segment start_sample = int(col * 2.5 * fs) end_sample = int((col + 1) * 2.5 * fs) segment = data[start_sample:min(end_sample, len(data))] t = np.linspace(col * 2.5, (col + 1) * 2.5, len(segment)) ax.plot(t, segment, color=line_color, linewidth=0.5) ax.set_ylabel('') ax.tick_params(left=False, labelleft=False, bottom=False, labelbottom=False) # Row 4: Long Lead II (10 seconds) ax = axes[3] ax.set_facecolor('white') ax.grid(True, color=grid_color, linewidth=0.5) ax.set_xlim(0, 10) ax.set_ylim(-2, 2) if 'II' in lead_map: data = lead_map['II'] t = np.linspace(0, 10, len(data)) ax.plot(t, data, color=line_color, linewidth=0.5) ax.tick_params(left=False, labelleft=False, bottom=False, labelbottom=False) plt.tight_layout() plt.savefig(output_path, dpi=resolution, bbox_inches='tight', pad_inches=0) plt.close(fig) return True except Exception as e: print(f"Error generating image: {e}") return False def main(): parser = argparse.ArgumentParser() parser.add_argument('--ptbxl_dir', type=str, default=str(PTBXL_PATH)) parser.add_argument('--output_dir', type=str, default='/data/ecg-digitization/synthetic_0001_fast') parser.add_argument('--n_samples', type=int, default=None) parser.add_argument('--start_idx', type=int, default=0) parser.add_argument('--resolution', type=int, default=200) args = parser.parse_args() output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) (output_dir / 'stage1').mkdir(exist_ok=True) (output_dir / 'gt').mkdir(exist_ok=True) print(f"Finding PTB-XL records in {args.ptbxl_dir}...") records = find_ptbxl_records(Path(args.ptbxl_dir)) print(f"Found {len(records)} PTB-XL records") if not records: print("No records found!") return n_samples = args.n_samples or len(records) n_samples = min(n_samples, len(records)) print(f"\nGenerating {n_samples} 0001-type synthetic images...") print(f" Output: {output_dir}") print(f" Format: {RAW_WIDTH}x{RAW_HEIGHT} (matches Kaggle)") print() ok_count = 0 skip_count = 0 fail_count = 0 pbar = tqdm(range(args.start_idx, args.start_idx + n_samples), desc="Generating") for i in pbar: record_idx = i % len(records) hea, dat = records[record_idx] sample_id = f'syn_{i:08d}' out_img = output_dir / 'stage1' / f'{sample_id}-0001.png' out_gt = output_dir / 'gt' / f'{sample_id}.npy' if out_img.exists() and out_gt.exists(): skip_count += 1 pbar.set_postfix({'ok': ok_count, 'skip': skip_count, 'fail': fail_count}) continue # Generate image temp_img = output_dir / 'temp' / f'{sample_id}.png' temp_img.parent.mkdir(exist_ok=True) success = generate_ecg_image_native(hea, dat, str(temp_img), args.resolution) if success and temp_img.exists(): # Read and resize to Kaggle format img = cv2.imread(str(temp_img)) if img is not None: img_resized = cv2.resize(img, (RAW_WIDTH, RAW_HEIGHT), interpolation=cv2.INTER_LINEAR) cv2.imwrite(str(out_img), img_resized) # Extract GT gt_signal = extract_gt_signal(hea, OUTPUT_GT_WIDTH) if gt_signal is not None: np.save(out_gt, gt_signal) ok_count += 1 else: fail_count += 1 else: fail_count += 1 temp_img.unlink(missing_ok=True) else: fail_count += 1 pbar.set_postfix({'ok': ok_count, 'skip': skip_count, 'fail': fail_count}) print(f"\n\nDone! ok={ok_count}, skip={skip_count}, fail={fail_count}") # Create manifest manifest = [] for img in sorted((output_dir / 'stage1').glob('*.png')): sample_id = img.stem.replace('-0001', '') gt_path = output_dir / 'gt' / f'{sample_id}.npy' if gt_path.exists(): manifest.append({ 'sample_id': sample_id, 'image': str(img), 'gt': str(gt_path), }) with open(output_dir / 'manifest.json', 'w') as f: json.dump(manifest, f, indent=2) print(f"Manifest: {len(manifest)} samples") if __name__ == '__main__': main()