#!/usr/bin/env python3 """ Generate ONLY 0001-type synthetic ECG images (clean red grid, no augmentation). Uses the same ecg-image-kit approach as generate_synthetic_ecgkit.py but ONLY style 0. This is based on Kaggle discussion: train only on 0001-type images. Usage: python gen_synthetic_0001_v3.py --n_samples 21799 --n_workers 32 """ import os import sys import argparse import subprocess 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 # Paths 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') # 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_single(hea_file, dat_file, output_dir, sample_idx, temp_base): """Generate a single 0001-type image (clean red grid, no augmentation).""" sample_id = f'syn_{sample_idx: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(): return 'skip', sample_idx, "exists" temp_dir = temp_base / f't_{sample_idx}' try: temp_dir.mkdir(parents=True, exist_ok=True) # Generate 0001-style image: clean red grid, NO augmentation, NO wrinkles cmd = [ 'python', 'gen_ecg_image_from_data.py', '-i', dat_file, '-hea', hea_file, '-o', str(temp_dir), '-st', str(sample_idx), '-se', str(sample_idx), '-r', '200', '--num_columns', '4', '--full_mode', 'II', '--standard_grid_color', '5', # Red grid = 0001 style # NO --augment, NO --wrinkles ] result = subprocess.run( cmd, cwd=str(ECG_IMAGE_KIT), capture_output=True, text=True, timeout=120 ) if result.returncode != 0: shutil.rmtree(temp_dir, ignore_errors=True) err = result.stderr[:200] if result.stderr else "unknown" return 'fail', sample_idx, f"ecg-kit: {err}" gen_imgs = list(temp_dir.glob('*.png')) if not gen_imgs: shutil.rmtree(temp_dir, ignore_errors=True) return 'fail', sample_idx, "no image" img = cv2.imread(str(gen_imgs[0])) if img is None: shutil.rmtree(temp_dir, ignore_errors=True) return 'fail', sample_idx, "read failed" # Resize to Kaggle format img_resized = cv2.resize(img, (RAW_WIDTH, RAW_HEIGHT), interpolation=cv2.INTER_LINEAR) out_img.parent.mkdir(parents=True, exist_ok=True) cv2.imwrite(str(out_img), img_resized) gt_signal = extract_gt_signal(hea_file, OUTPUT_GT_WIDTH) if gt_signal is not None: out_gt.parent.mkdir(parents=True, exist_ok=True) np.save(out_gt, gt_signal) else: shutil.rmtree(temp_dir, ignore_errors=True) return 'fail', sample_idx, "gt extraction failed" shutil.rmtree(temp_dir, ignore_errors=True) return 'ok', sample_idx, "success" except subprocess.TimeoutExpired: shutil.rmtree(temp_dir, ignore_errors=True) return 'fail', sample_idx, "timeout" except Exception as e: shutil.rmtree(temp_dir, ignore_errors=True) return 'fail', sample_idx, str(e)[:100] 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_v3') parser.add_argument('--n_samples', type=int, default=None) parser.add_argument('--start_idx', type=int, default=0) 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) temp_base = output_dir / 'temp' temp_base.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 # Process sequentially - ecg-image-kit imports are slow, # but once loaded it's faster to stay in one process 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] status, idx, msg = generate_single(hea, dat, output_dir, i, temp_base) if status == 'ok': ok_count += 1 elif status == 'skip': skip_count += 1 else: fail_count += 1 if fail_count <= 5: print(f"\n Failed {i}: {msg}") 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()