#!/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()