| |
| """ |
| 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 |
|
|
| |
| |
| |
| |
| |
|
|
| LEAD_ORDER = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6'] |
| TARGET_SAMPLES = 5000 |
|
|
| 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: |
| |
| record_path = source_path.replace('.hea', '') |
| record = wfdb.rdrecord(record_path) |
| |
| |
| signals = record.p_signal |
| sig_names = record.sig_name |
| |
| |
| 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) |
| |
| |
| lead_signals = {} |
| for i, name in enumerate(sig_names): |
| lead_signals[name.upper()] = signals[:, i] |
| |
| |
| |
| |
| |
| |
| output_width = 3926 |
| segment_len = output_width // 3 |
| |
| gt_signal = np.zeros((4, output_width), dtype=np.float32) |
| |
| row_leads = [ |
| ['I', 'AVR', 'V1', 'V4'], |
| ['II', 'AVL', 'V2', 'V5'], |
| ['III', 'AVF', 'V3', 'V6'], |
| ] |
| |
| |
| for row_idx in range(4): |
| row_signal = np.zeros(output_width, dtype=np.float32) |
| |
| |
| 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] |
| |
| start_sample = col_idx * 1250 |
| end_sample = min(start_sample + 1250, len(lead_data)) |
| segment = lead_data[start_sample:end_sample] |
| |
| |
| 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) |
| |
| |
| with open(synthetic_dir / 'manifest.json') as f: |
| manifest = json.load(f) |
| |
| print(f"Processing {len(manifest)} samples...") |
| |
| |
| 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() |
|
|