ecg-digitization-experiments / code /scripts /gen_synthetic_0001_v3.py
Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
9 kB
#!/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()