ecg-digitization-experiments / code /scripts /extract_synthetic_gt.py
Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
4.84 kB
#!/usr/bin/env python3
"""
Extract ground truth signals from PTB-XL for synthetic images.
Creates GT numpy files from the manifest.
"""
import os
import json
import numpy as np
import wfdb
from pathlib import Path
from tqdm import tqdm
from concurrent.futures import ProcessPoolExecutor, as_completed
# Competition standard: 12-lead ECG with 4 rows
# Row 0: I, II, III
# Row 1: aVR, aVL, aVF
# Row 2: V1, V2, V3
# Row 3: V4, V5, V6
LEAD_ORDER = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']
TARGET_SAMPLES = 5000 # Standard 10s at 500Hz
def extract_gt(sample_info, output_dir):
"""Extract GT signal from PTB-XL for one sample."""
sample_id = sample_info['sample_id']
source_path = sample_info['source']
output_path = output_dir / f'{sample_id}.npy'
if output_path.exists():
return sample_id, True, "exists"
try:
# Load WFDB record (remove .hea extension)
record_path = source_path.replace('.hea', '')
record = wfdb.rdrecord(record_path)
# Get signals and lead names
signals = record.p_signal # Shape: [samples, leads]
sig_names = record.sig_name
# Resample to 500Hz if needed
fs = record.fs
if fs != 500:
from scipy import signal as scipy_signal
num_samples = int(signals.shape[0] * 500 / fs)
signals = scipy_signal.resample(signals, num_samples, axis=0)
# Map leads to standard order
lead_signals = {}
for i, name in enumerate(sig_names):
lead_signals[name.upper()] = signals[:, i]
# Create 4-row output (3 leads per row, concatenated)
# Each row shows 10 seconds = 5000 samples at 500Hz
# But images show 2.5s per lead segment (1250 samples)
# 4 rows x 3 leads = 12 leads, 2.5s each
output_width = 3926 # T1 - T0 = 4161 - 235
segment_len = output_width // 3 # ~1308 samples per lead
gt_signal = np.zeros((4, output_width), dtype=np.float32)
row_leads = [
['I', 'AVR', 'V1', 'V4'], # First column of each rhythm strip
['II', 'AVL', 'V2', 'V5'], # Second column
['III', 'AVF', 'V3', 'V6'], # Third column
]
# Standard 12-lead format: each row shows 3 leads, 2.5s each
for row_idx in range(4):
row_signal = np.zeros(output_width, dtype=np.float32)
# Get 3 leads for this row
leads_for_row = [row_leads[col][row_idx] for col in range(3)]
for col_idx, lead_name in enumerate(leads_for_row):
if lead_name in lead_signals:
lead_data = lead_signals[lead_name]
# Take 2.5s segment
start_sample = col_idx * 1250 # 2.5s at 500Hz
end_sample = min(start_sample + 1250, len(lead_data))
segment = lead_data[start_sample:end_sample]
# Resample to fit column width
col_start = col_idx * segment_len
col_end = (col_idx + 1) * segment_len
if col_idx == 2:
col_end = output_width
x_old = np.linspace(0, 1, len(segment))
x_new = np.linspace(0, 1, col_end - col_start)
row_signal[col_start:col_end] = np.interp(x_new, x_old, segment)
gt_signal[row_idx] = row_signal
np.save(output_path, gt_signal)
return sample_id, True, "created"
except Exception as e:
return sample_id, False, str(e)
def main():
synthetic_dir = Path('/home/azureuser/data/ecg-digitization/synthetic_ecgkit')
output_dir = synthetic_dir / 'gt'
output_dir.mkdir(exist_ok=True)
# Load manifest
with open(synthetic_dir / 'manifest.json') as f:
manifest = json.load(f)
print(f"Processing {len(manifest)} samples...")
# Get unique sample_ids (each sample_id may have multiple styles)
sample_map = {}
for entry in manifest:
sample_id = entry['sample_id']
if sample_id not in sample_map:
sample_map[sample_id] = entry
print(f"Unique samples: {len(sample_map)}")
success = 0
failed = 0
for sample_id, info in tqdm(sample_map.items()):
sid, ok, msg = extract_gt(info, output_dir)
if ok:
success += 1
else:
failed += 1
if failed <= 5:
print(f"Failed {sid}: {msg}")
print(f"\nDone! Success: {success}, Failed: {failed}")
print(f"GT files: {output_dir}")
if __name__ == '__main__':
main()