| |
| """ |
| Synthetic ECG Image Generator using ecg-image-kit |
| |
| Generates synthetic ECG images from PTB-XL records, then processes through Stage 0/1. |
| Based on expert advice: use ecg-image-kit for high-res GT, then train on Stage 1 recovered images. |
| |
| Usage: |
| python generate_synthetic_ecgkit.py --n_samples 50000 --n_workers 8 |
| """ |
|
|
| import os |
| import sys |
| import argparse |
| import subprocess |
| import numpy as np |
| import cv2 |
| import torch |
| import random |
| import json |
| import traceback |
| from pathlib import Path |
| from multiprocessing import Pool, cpu_count |
| from tqdm import tqdm |
| from concurrent.futures import ThreadPoolExecutor, as_completed |
| import shutil |
|
|
| |
| PROJECT_ROOT = Path(__file__).parent.parent |
| ECG_IMAGE_KIT = PROJECT_ROOT / 'ecg-image-kit' / 'codes' / 'ecg-image-generator' |
| BASELINE_PATH = PROJECT_ROOT / 'data' / 'hengck23-submit-physionet' / 'hengck23-submit-physionet' |
| PTBXL_PATH = Path('/data/ecg-digitization/ptbxl/physionet.org/files/ptb-xl/1.0.3') |
|
|
| sys.path.insert(0, str(BASELINE_PATH)) |
|
|
| |
| TARGET_HEIGHT = 1700 |
| TARGET_WIDTH = 2200 |
|
|
|
|
| def find_ptbxl_records(ptbxl_dir): |
| """Find all PTB-XL record files.""" |
| records = [] |
| |
| |
| records500 = ptbxl_dir / 'records500' |
| if records500.exists(): |
| for subdir in sorted(records500.iterdir()): |
| if subdir.is_dir(): |
| for hea in subdir.glob('*.hea'): |
| dat = hea.with_suffix('.dat') |
| if dat.exists(): |
| records.append((str(hea), str(dat))) |
| |
| |
| if not records: |
| records100 = ptbxl_dir / 'records100' |
| if records100.exists(): |
| for subdir in sorted(records100.iterdir()): |
| if subdir.is_dir(): |
| for hea in subdir.glob('*.hea'): |
| dat = hea.with_suffix('.dat') |
| if dat.exists(): |
| records.append((str(hea), str(dat))) |
| |
| return records |
|
|
|
|
| def generate_single_image_ecgkit(args): |
| """Generate a single synthetic image using ecg-image-kit.""" |
| hea_file, dat_file, output_dir, sample_idx, style_idx = args |
| |
| try: |
| |
| temp_dir = Path(output_dir) / 'temp' / f'sample_{sample_idx}' |
| temp_dir.mkdir(parents=True, exist_ok=True) |
| |
| |
| style_configs = { |
| 0: {'grid': True, 'grid_color': 5, 'augment': False, 'wrinkles': False}, |
| 1: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': True}, |
| 2: {'grid': True, 'grid_color': 2, 'augment': False, 'wrinkles': False}, |
| 3: {'grid': True, 'grid_color': 2, 'augment': True, 'wrinkles': True}, |
| 4: {'grid': True, 'grid_color': 1, 'augment': False, 'wrinkles': False}, |
| 5: {'grid': True, 'grid_color': 3, 'augment': True, 'wrinkles': False}, |
| 6: {'grid': True, 'grid_color': 4, 'augment': True, 'wrinkles': True}, |
| 7: {'grid': True, 'grid_color': 2, 'augment': True, 'wrinkles': False}, |
| 8: {'grid': True, 'grid_color': 6, 'augment': False, 'wrinkles': False}, |
| 9: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': False}, |
| 10: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': True}, |
| 11: {'grid': False, 'grid_color': 0, 'augment': False, 'wrinkles': False}, |
| } |
| |
| style = style_configs.get(style_idx % 12, style_configs[0]) |
| |
| |
| 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', |
| '--store_config', '1', |
| ] |
| |
| if style['grid']: |
| cmd.extend(['--standard_grid_color', str(style['grid_color'])]) |
| else: |
| cmd.extend(['--random_grid_present', '0']) |
| |
| if style['augment']: |
| cmd.append('--augment') |
| |
| if style['wrinkles']: |
| cmd.append('--wrinkles') |
| |
| |
| env = os.environ.copy() |
| env['MPLBACKEND'] = 'Agg' |
| env['MPLCONFIGDIR'] = str(temp_dir) |
| |
| result = subprocess.run( |
| cmd, |
| cwd=str(ECG_IMAGE_KIT), |
| capture_output=True, |
| text=True, |
| timeout=120, |
| env=env |
| ) |
| |
| if result.returncode != 0: |
| return None |
| |
| |
| generated_images = list(temp_dir.glob('*.png')) |
| if not generated_images: |
| return None |
| |
| gen_img_path = generated_images[0] |
| |
| |
| img = cv2.imread(str(gen_img_path)) |
| if img is None: |
| return None |
| |
| img = cv2.resize(img, (TARGET_WIDTH, TARGET_HEIGHT)) |
| |
| |
| sample_id = f'syn_{sample_idx:08d}' |
| style_str = '0001' |
| |
| raw_dir = Path(output_dir) / 'raw' |
| raw_dir.mkdir(parents=True, exist_ok=True) |
| raw_path = raw_dir / f'{sample_id}-{style_str}.png' |
| cv2.imwrite(str(raw_path), img) |
| |
| |
| config_files = list(temp_dir.glob('*.json')) |
| if config_files: |
| with open(config_files[0]) as f: |
| config = json.load(f) |
| else: |
| config = {} |
| |
| |
| shutil.rmtree(temp_dir, ignore_errors=True) |
| |
| return { |
| 'sample_id': sample_id, |
| 'style': style_str, |
| 'raw_path': str(raw_path), |
| 'source_record': hea_file, |
| } |
| |
| except Exception as e: |
| return None |
|
|
|
|
| def process_through_stage01(raw_dir, output_dir, device='cuda:0'): |
| """Process all raw images through Stage 0 and Stage 1.""" |
| from stage0_model import Net as Stage0Net |
| from stage1_model import Net as Stage1Net |
| from stage0_common import image_to_batch, output_to_predict, normalise_by_homography, load_net |
| from stage1_common import output_to_predict as stage1_output_to_predict, rectify_image |
| |
| raw_dir = Path(raw_dir) |
| output_dir = Path(output_dir) |
| |
| |
| stage0_dir = output_dir / 'stage0' |
| stage1_dir = output_dir / 'stage1' |
| stage0_dir.mkdir(parents=True, exist_ok=True) |
| stage1_dir.mkdir(parents=True, exist_ok=True) |
| |
| |
| weight_dir = BASELINE_PATH / 'weight' |
| |
| print("Loading Stage 0 model...") |
| stage0_net = Stage0Net(pretrained=False) |
| stage0_net = load_net(stage0_net, str(weight_dir / 'stage0-last.checkpoint.pth')) |
| stage0_net.to(device).eval() |
| |
| print("Loading Stage 1 model...") |
| stage1_net = Stage1Net(pretrained=False) |
| stage1_net = load_net(stage1_net, str(weight_dir / 'stage1-last.checkpoint.pth')) |
| stage1_net.to(device).eval() |
| |
| |
| raw_images = sorted(raw_dir.glob('*.png')) |
| print(f"Processing {len(raw_images)} images through Stage 0/1...") |
| |
| for img_path in tqdm(raw_images, desc="Stage 0/1"): |
| sample_name = img_path.stem |
| stage1_path = stage1_dir / f'{sample_name}.png' |
| |
| if stage1_path.exists(): |
| continue |
| |
| try: |
| |
| image = cv2.imread(str(img_path)) |
| if image is None: |
| continue |
| image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) |
| |
| |
| batch = image_to_batch(image) |
| with torch.no_grad(): |
| output = stage0_net(batch) |
| |
| rotated, keypoint = output_to_predict(image, batch, output) |
| |
| if keypoint is not None and len(keypoint) >= 4: |
| try: |
| normalized, _, _ = normalise_by_homography(rotated, keypoint) |
| except: |
| normalized = rotated |
| else: |
| normalized = rotated |
| |
| |
| stage0_path = stage0_dir / f'{sample_name}.png' |
| cv2.imwrite(str(stage0_path), cv2.cvtColor(normalized, cv2.COLOR_RGB2BGR)) |
| |
| |
| batch = {'image': torch.from_numpy( |
| np.ascontiguousarray(normalized.transpose(2, 0, 1)) |
| ).unsqueeze(0)} |
| |
| with torch.no_grad(): |
| output = stage1_net(batch) |
| |
| try: |
| gridpoint_xy, _ = stage1_output_to_predict(normalized, batch, output) |
| rectified = rectify_image(normalized, gridpoint_xy) |
| except: |
| rectified = normalized |
| |
| |
| rectified = cv2.resize(rectified, (TARGET_WIDTH, TARGET_HEIGHT)) |
| |
| |
| cv2.imwrite(str(stage1_path), cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR)) |
| |
| except Exception as e: |
| |
| shutil.copy(img_path, stage1_path) |
| |
| print(f"Stage 0/1 processing complete. Output: {stage1_dir}") |
|
|
|
|
| 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_ecgkit') |
| parser.add_argument('--n_samples', type=int, default=50000) |
| parser.add_argument('--n_workers', type=int, default=8) |
| parser.add_argument('--skip_stage01', action='store_true') |
| parser.add_argument('--device', type=str, default='cuda:0') |
| args = parser.parse_args() |
| |
| output_dir = Path(args.output_dir) |
| output_dir.mkdir(parents=True, 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 |
| |
| |
| print(f"Preparing {args.n_samples} generation tasks...") |
| tasks = [] |
| for i in range(args.n_samples): |
| hea, dat = random.choice(records) |
| style_idx = 0 |
| tasks.append((hea, dat, str(output_dir), i, style_idx)) |
| |
| |
| print(f"Generating {args.n_samples} synthetic images...") |
| |
| results = [] |
| with Pool(args.n_workers) as pool: |
| for result in tqdm(pool.imap_unordered(generate_single_image_ecgkit, tasks), |
| total=len(tasks), desc="Generating"): |
| if result: |
| results.append(result) |
| |
| print(f"Successfully generated {len(results)} images") |
| |
| |
| manifest_path = output_dir / 'manifest.json' |
| with open(manifest_path, 'w') as f: |
| json.dump(results, f, indent=2) |
| |
| |
| if not args.skip_stage01: |
| print("\nProcessing through Stage 0/1...") |
| process_through_stage01(output_dir / 'raw', output_dir, args.device) |
| |
| |
| temp_dir = output_dir / 'temp' |
| if temp_dir.exists(): |
| shutil.rmtree(temp_dir, ignore_errors=True) |
| |
| print("\nDone!") |
| print(f"Raw images: {output_dir / 'raw'}") |
| print(f"Stage 1 images: {output_dir / 'stage1'}") |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|