ecg-digitization-experiments / code /scripts /generate_all_variants.py
Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
13.2 kB
#!/usr/bin/env python3
"""
Generate all 12 style variants for each synthetic ECG record.
This script takes existing PTB-XL records and generates 12 different style variants
(0001-0012) for each, similar to the Kaggle dataset format.
Usage:
python generate_all_variants.py --n_workers 8
"""
import os
import sys
import subprocess
import numpy as np
import cv2
import json
import shutil
from pathlib import Path
from multiprocessing import Pool, cpu_count
from tqdm import tqdm
import argparse
# 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_DIR = Path('/data/ecg-digitization/synthetic_ecgkit')
# Target specifications - match Kaggle format
TARGET_HEIGHT = 1700
TARGET_WIDTH = 2200
# 12 style configurations matching Kaggle variants 0001-0012
STYLE_CONFIGS = {
1: {'grid': True, 'grid_color': 5, 'augment': False, 'wrinkles': False, 'hw_text': False}, # Clean red grid
2: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': True, 'hw_text': False}, # Aged red
3: {'grid': True, 'grid_color': 2, 'augment': False, 'wrinkles': False, 'hw_text': True}, # Green grid + handwriting
4: {'grid': True, 'grid_color': 2, 'augment': True, 'wrinkles': True, 'hw_text': False}, # Aged green
5: {'grid': True, 'grid_color': 1, 'augment': False, 'wrinkles': False, 'hw_text': False}, # Blue grid
6: {'grid': True, 'grid_color': 3, 'augment': True, 'wrinkles': False, 'hw_text': True}, # Grey + handwriting
7: {'grid': True, 'grid_color': 4, 'augment': True, 'wrinkles': True, 'hw_text': False}, # Yellow aged
8: {'grid': True, 'grid_color': 2, 'augment': True, 'wrinkles': False, 'hw_text': False}, # Faded green
9: {'grid': True, 'grid_color': 6, 'augment': False, 'wrinkles': False, 'hw_text': True}, # Orange + handwriting
10: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': False, 'hw_text': False}, # High contrast red
11: {'grid': True, 'grid_color': 5, 'augment': True, 'wrinkles': True, 'hw_text': True}, # Low contrast + text
12: {'grid': False, 'grid_color': 0, 'augment': False, 'wrinkles': False, 'hw_text': False}, # No grid
}
def find_ptbxl_records(ptbxl_dir, max_records=None):
"""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 sorted(subdir.glob('*.hea')):
dat = hea.with_suffix('.dat')
if dat.exists():
records.append((str(hea), str(dat)))
if max_records and len(records) >= max_records:
return records
return records
def generate_variant(args):
"""Generate a single style variant for one record."""
hea_file, dat_file, sample_idx, style_idx, output_dir = args
try:
style = STYLE_CONFIGS[style_idx]
sample_id = f'syn_{sample_idx:08d}'
style_str = f'{style_idx:04d}'
# Create temp directory
temp_dir = output_dir / 'temp' / f'{sample_id}_{style_str}'
temp_dir.mkdir(parents=True, exist_ok=True)
# Build ecg-image-kit command
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')
if style['hw_text']:
cmd.append('--print_txt')
# Run with isolated matplotlib config
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:
shutil.rmtree(temp_dir, ignore_errors=True)
return None
# Find generated image
generated_images = list(temp_dir.glob('*.png'))
if not generated_images:
shutil.rmtree(temp_dir, ignore_errors=True)
return None
gen_img_path = generated_images[0]
# Read and resize
img = cv2.imread(str(gen_img_path))
if img is None:
shutil.rmtree(temp_dir, ignore_errors=True)
return None
img = cv2.resize(img, (TARGET_WIDTH, TARGET_HEIGHT))
# Save raw image
raw_dir = output_dir / 'raw'
raw_path = raw_dir / f'{sample_id}-{style_str}.png'
cv2.imwrite(str(raw_path), img)
# Clean up
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:
try:
shutil.rmtree(temp_dir, ignore_errors=True)
except:
pass
return None
def process_through_stage01(raw_dir, output_dir, batch_size=50):
"""Process raw images through Stage 0 and Stage 1."""
sys.path.insert(0, str(PROJECT_ROOT / 'data' / 'hengck23-submit-physionet' / 'hengck23-submit-physionet'))
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)
stage1_dir = output_dir / 'stage1'
stage1_dir.mkdir(parents=True, exist_ok=True)
# Load models
weight_dir = PROJECT_ROOT / 'data' / 'hengck23-submit-physionet' / 'hengck23-submit-physionet' / 'weight'
device = 'cuda:0'
stage0_net = Stage0Net(pretrained=False)
stage0_net = load_net(stage0_net, str(weight_dir / 'stage0-last.checkpoint.pth'))
stage0_net.to(device).eval()
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'))
# Filter to only process new images (not already in stage1)
existing = set(p.stem for p in stage1_dir.glob('*.png'))
to_process = [p for p in raw_images if p.stem not in existing]
print(f"Processing {len(to_process)} images through Stage 0/1 (skipping {len(existing)} existing)")
for raw_path in tqdm(to_process, desc="Stage 0/1"):
try:
image = cv2.imread(str(raw_path))
if image is None:
continue
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
# Stage 0
with torch.no_grad():
batch = image_to_batch(image)
output = stage0_net(batch)
rotated, keypoint = output_to_predict(image, batch, output)
normalized, _, _ = normalise_by_homography(rotated, keypoint)
# Stage 1
with torch.no_grad():
batch = {'image': torch.from_numpy(normalized.transpose(2, 0, 1)).unsqueeze(0)}
output = stage1_net(batch)
gridpoint_xy, _ = stage1_output_to_predict(normalized, batch, output)
rectified = rectify_image(normalized, gridpoint_xy)
# Save
rectified_bgr = cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR)
cv2.imwrite(str(stage1_dir / raw_path.name), rectified_bgr)
except Exception as e:
continue
print(f"Stage 0/1 processing complete. {len(list(stage1_dir.glob('*.png')))} images in stage1/")
def main():
parser = argparse.ArgumentParser()
parser.add_argument('--n_samples', type=int, default=None,
help='Number of PTB-XL records to use (default: all ~9000)')
parser.add_argument('--n_workers', type=int, default=8)
parser.add_argument('--output_dir', type=str, default=str(OUTPUT_DIR))
parser.add_argument('--skip_generation', action='store_true',
help='Skip generation, only run Stage 0/1')
parser.add_argument('--skip_stage01', action='store_true',
help='Skip Stage 0/1 processing')
parser.add_argument('--styles', type=str, default='1-12',
help='Style range to generate, e.g., "1-12" or "2,4,6,8"')
args = parser.parse_args()
output_dir = Path(args.output_dir)
raw_dir = output_dir / 'raw'
raw_dir.mkdir(parents=True, exist_ok=True)
# Parse styles
if '-' in args.styles:
start, end = map(int, args.styles.split('-'))
styles = list(range(start, end + 1))
else:
styles = [int(s) for s in args.styles.split(',')]
print(f"Generating styles: {styles}")
if not args.skip_generation:
# Find PTB-XL records
records = find_ptbxl_records(PTBXL_PATH, args.n_samples)
print(f"Found {len(records)} PTB-XL records")
# Check existing raw images to avoid regenerating
existing_raw = set()
for p in raw_dir.glob('*.png'):
# syn_00000000-0001.png -> (0, 1)
parts = p.stem.split('-')
if len(parts) == 2:
sample_num = int(parts[0].replace('syn_', ''))
style_num = int(parts[1])
existing_raw.add((sample_num, style_num))
print(f"Found {len(existing_raw)} existing raw images")
# Generate all variants
tasks = []
for sample_idx, (hea, dat) in enumerate(records):
for style_idx in styles:
if (sample_idx, style_idx) not in existing_raw:
tasks.append((hea, dat, sample_idx, style_idx, output_dir))
print(f"Generating {len(tasks)} new images...")
if tasks:
with Pool(args.n_workers) as pool:
results = list(tqdm(
pool.imap_unordered(generate_variant, tasks),
total=len(tasks),
desc="Generating"
))
successful = sum(1 for r in results if r is not None)
print(f"Generated {successful}/{len(tasks)} images")
# Process through Stage 0/1
if not args.skip_stage01:
import torch
process_through_stage01(raw_dir, output_dir)
# Generate CSV files for new stage1 images
print("\nConverting GT to CSV format...")
stage1_dir = output_dir / 'stage1'
gt_dir = output_dir / 'gt'
# Count new stage1 images without CSV
stage1_images = set(p.stem for p in stage1_dir.glob('*.png'))
existing_csv = set(p.stem for p in gt_dir.glob('*.csv'))
new_images = stage1_images - existing_csv
if new_images:
print(f"Need to create {len(new_images)} new CSV files")
# For each new image, find the base NPY and create CSV
for img_stem in tqdm(new_images, desc="Creating CSVs"):
# syn_00000000-0003 -> syn_00000000
base_id = img_stem.rsplit('-', 1)[0]
npy_path = gt_dir / f'{base_id}.npy'
if npy_path.exists():
# Load NPY and convert to CSV
mv_signal = np.load(npy_path)
# Convert mV to pixels
ZERO_MV = np.array([703.5, 987.5, 1271.5, 1531.5])
MV_TO_PIXEL = 78.5
pixel_signal = np.zeros_like(mv_signal)
for row in range(4):
pixel_signal[row] = ZERO_MV[row] - mv_signal[row] * MV_TO_PIXEL
pixel_signal = np.clip(pixel_signal, 0, 1695)
# Save CSV
import pandas as pd
df = pd.DataFrame(pixel_signal.T, columns=['row0', 'row1', 'row2', 'row3'])
df.to_csv(gt_dir / f'{img_stem}.csv', index=False)
# Summary
print("\n" + "="*60)
print("Summary:")
print(f" Raw images: {len(list(raw_dir.glob('*.png')))}")
print(f" Stage1 images: {len(list(stage1_dir.glob('*.png')))}")
print(f" GT CSV files: {len(list(gt_dir.glob('*.csv')))}")
print("="*60)
if __name__ == '__main__':
main()