ecg-digitization-experiments / code /scripts /generate_synthetic_0001_only.py
Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
10.3 kB
#!/usr/bin/env python3
"""
Generate ONLY 0001-type synthetic ECG images using ecg-image-kit.
Based on Kaggle discussion: train only on 0001 type (clean red grid) and use
heavy augmentation to simulate other variants.
Usage:
python generate_synthetic_0001_only.py --output_dir /data/ecg-digitization/synthetic_0001_all
"""
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 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'
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 specifications - match Kaggle format exactly
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)
# Use records500 (500Hz) for better quality
records500 = ptbxl_dir / 'records500'
if records500.exists():
for subdir in sorted(records500.iterdir()):
if subdir.is_dir():
for hea in subdir.glob('*_hr.hea'): # High-res files only
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 # [samples, leads]
sig_names = record.sig_name
fs = record.fs
# Resample to 500Hz if needed
if fs != 500:
num_samples = int(signals.shape[0] * 500 / fs)
signals = scipy_signal.resample(signals, num_samples, axis=0)
# Map lead names
lead_signals = {}
for i, name in enumerate(sig_names):
lead_signals[name.upper()] = signals[:, i]
# 4-row output: each row shows 3 leads, 2.5s each (standard 12-lead layout)
# Row 0: I, aVR, V1, V4
# Row 1: II, aVL, V2, V5
# Row 2: III, aVF, V3, V6
# Row 3: Long II strip
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)
# First 3 rows: 3 leads each, 2.5s segments
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]
# 2.5s at 500Hz = 1250 samples
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)
# Row 4: Long Lead II strip (10s)
if 'II' in lead_signals:
lead_ii = lead_signals['II']
# Full 10 seconds
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:
print(f"Error extracting GT from {hea_file}: {e}")
return None
def generate_single_0001(hea_file, dat_file, output_dir, sample_idx, temp_base):
"""Generate a single 0001-type image (clean red grid, no distortion)."""
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'
# Skip if exists
if out_img.exists() and out_gt.exists():
return True, sample_idx, "exists"
temp_dir = temp_base / f'temp_{sample_idx}'
try:
temp_dir.mkdir(parents=True, exist_ok=True)
# Generate image with ecg-image-kit
# 0001 = 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', # 200 DPI
'--num_columns', '4',
'--full_mode', 'II',
'--standard_grid_color', '5', # Red grid (0001 style)
# NO --augment, NO --wrinkles for clean baseline
]
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)
return False, sample_idx, f"ecg-kit error: {result.stderr[:200]}"
# Find generated image
gen_imgs = list(temp_dir.glob('*.png'))
if not gen_imgs:
shutil.rmtree(temp_dir, ignore_errors=True)
return False, sample_idx, "no image generated"
# Load image
img = cv2.imread(str(gen_imgs[0]))
if img is None:
shutil.rmtree(temp_dir, ignore_errors=True)
return False, sample_idx, "failed to read image"
# Resize to match Kaggle format (1700x2200)
img_resized = cv2.resize(img, (RAW_WIDTH, RAW_HEIGHT), interpolation=cv2.INTER_LINEAR)
# Save image
out_img.parent.mkdir(parents=True, exist_ok=True)
cv2.imwrite(str(out_img), img_resized)
# Extract and save GT signal
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 False, sample_idx, "failed to extract GT"
# Cleanup
shutil.rmtree(temp_dir, ignore_errors=True)
return True, sample_idx, "success"
except subprocess.TimeoutExpired:
shutil.rmtree(temp_dir, ignore_errors=True)
return False, sample_idx, "timeout"
except Exception as e:
shutil.rmtree(temp_dir, ignore_errors=True)
return False, sample_idx, str(e)
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_all')
parser.add_argument('--start_idx', type=int, default=0)
parser.add_argument('--max_samples', type=int, default=None, help='Max samples to generate (default: all)')
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)
# Find PTB-XL records
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
# Determine how many to generate
n_samples = len(records)
if args.max_samples:
n_samples = min(n_samples, args.max_samples)
print(f"\nGenerating {n_samples} 0001-type synthetic images...")
print(f" Output: {output_dir / 'stage1'}")
print(f" GT: {output_dir / 'gt'}")
print(f" Format: {RAW_WIDTH}x{RAW_HEIGHT} (matches Kaggle)")
print()
# Generate one by one (simple, reliable)
success_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]
success, idx, msg = generate_single_0001(hea, dat, output_dir, i, temp_base)
if success:
if msg == "exists":
skip_count += 1
else:
success_count += 1
else:
fail_count += 1
if fail_count <= 10:
print(f"\n Failed {i}: {msg}")
pbar.set_postfix({'ok': success_count, 'skip': skip_count, 'fail': fail_count})
print(f"\n\nGeneration complete!")
print(f" Success: {success_count}")
print(f" Skipped (existing): {skip_count}")
print(f" Failed: {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"\nManifest saved with {len(manifest)} samples")
print(f"Output directory: {output_dir}")
if __name__ == '__main__':
main()