ecg-digitization-experiments / code /scripts /compute_baseline_offsets_v19.py
Ubuntu
Add training scripts and notebooks
b69e447
Raw
History Blame Contribute Delete
17.4 kB
#!/usr/bin/env python3
"""
Compute per-segment baseline offsets for V19 model.
V19: ConvNeXt-Base + BiLSTM + Deformable Conv (per-row, no cross-row attention)
"""
import os
import sys
import argparse
import random
import numpy as np
import pandas as pd
from pathlib import Path
from tqdm import tqdm
from scipy.signal import savgol_filter
from scipy import stats
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
import cv2
import timm
try:
from torchvision.ops import DeformConv2d
HAS_DEFORM_CONV = True
except ImportError:
HAS_DEFORM_CONV = False
# =============================================================================
# Constants
# =============================================================================
TARGET_HEIGHT, TARGET_WIDTH = 1696, 4352
ZERO_MV = np.array([703.5, 987.5, 1271.5, 1531.5])
MV_TO_PIXEL = 78.5
T0, T1 = 235, 4161
X0, X1 = 0, 2176
Y0, Y1 = 0, 1696
OUTPUT_WIDTH = T1 - T0
CROP_HALF_HEIGHT = 250
ROW_HEIGHT = 500
ECG_MV_MIN, ECG_MV_MAX = -10.0, 10.0
VALID_VARIANTS = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']
LEAD_LAYOUT = [
['I', 'aVR', 'V1', 'V4'],
['II', 'aVL', 'V2', 'V5'],
['III', 'aVF', 'V3', 'V6'],
]
ALL_LEADS = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']
# =============================================================================
# V19 Model Architecture
# =============================================================================
class DeformableConvBlock(nn.Module):
def __init__(self, in_ch, out_ch, kernel_size=3, stride=1, padding=1):
super().__init__()
self.kernel_size = kernel_size
if HAS_DEFORM_CONV:
self.offset_conv = nn.Sequential(
nn.Conv2d(in_ch, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.Conv2d(64, 2 * kernel_size * kernel_size, 3, padding=1),
)
self.deform_conv = DeformConv2d(in_ch, out_ch, kernel_size, stride=stride, padding=padding)
else:
self.conv = nn.Conv2d(in_ch, out_ch, kernel_size, stride=stride, padding=padding)
self.norm = nn.BatchNorm2d(out_ch)
self.act = nn.GELU()
def forward(self, x):
if HAS_DEFORM_CONV:
offset = self.offset_conv(x)
out = self.deform_conv(x, offset)
else:
out = self.conv(x)
return self.act(self.norm(out))
class BiLSTMHead(nn.Module):
def __init__(self, input_dim, hidden_dim=128, num_layers=2, dropout=0.1):
super().__init__()
self.lstm = nn.LSTM(input_size=input_dim, hidden_size=hidden_dim, num_layers=num_layers,
batch_first=True, bidirectional=True, dropout=dropout if num_layers > 1 else 0)
self.output_proj = nn.Sequential(nn.Linear(hidden_dim * 2, hidden_dim), nn.LayerNorm(hidden_dim), nn.GELU())
self.output_dim = hidden_dim
def forward(self, x):
x = x.permute(0, 2, 1)
lstm_out, _ = self.lstm(x)
return self.output_proj(lstm_out)
class AuxiliaryHeads(nn.Module):
def __init__(self, feature_dim):
super().__init__()
self.grid_head = nn.Sequential(nn.Conv2d(32, 16, 3, padding=1), nn.BatchNorm2d(16), nn.ReLU(inplace=True),
nn.Conv2d(16, 1, 1), nn.Sigmoid())
self.gradient_head = nn.Sequential(nn.Linear(feature_dim, 64), nn.GELU(), nn.Linear(64, 1), nn.Tanh())
self.uncertainty_head = nn.Sequential(nn.Linear(feature_dim, 64), nn.GELU(), nn.Linear(64, 1))
def forward(self, features_2d, features_1d):
return self.grid_head(features_2d), self.gradient_head(features_1d).squeeze(-1), self.uncertainty_head(features_1d).squeeze(-1)
class CoordConv2d(nn.Module):
def __init__(self, in_channels, out_channels, kernel_size, **kwargs):
super().__init__()
self.conv = nn.Conv2d(in_channels + 2, out_channels, kernel_size, **kwargs)
def forward(self, x):
B, C, H, W = x.shape
yy = torch.linspace(-1, 1, H, device=x.device).view(1, 1, H, 1).expand(B, 1, H, W)
xx = torch.linspace(-1, 1, W, device=x.device).view(1, 1, 1, W).expand(B, 1, H, W)
return self.conv(torch.cat([x, yy, xx], dim=1))
class UNetDecoderBlockV19(nn.Module):
def __init__(self, in_ch, skip_ch, out_ch, use_deform=False):
super().__init__()
if use_deform and HAS_DEFORM_CONV:
self.conv1 = DeformableConvBlock(in_ch + skip_ch, out_ch)
else:
self.conv1 = nn.Sequential(nn.Conv2d(in_ch + skip_ch, out_ch, 3, padding=1, bias=False),
nn.BatchNorm2d(out_ch), nn.GELU())
self.conv2 = nn.Sequential(nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.GELU())
self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
def forward(self, x, skip=None):
x = self.upsample(x)
if skip is not None:
if x.shape[2:] != skip.shape[2:]:
x = F.interpolate(x, size=skip.shape[2:], mode='bilinear', align_corners=True)
x = torch.cat([x, skip], dim=1)
return self.conv2(self.conv1(x))
class PerLeadNetV19(nn.Module):
def __init__(self, encoder_name='convnext_base.fb_in22k_ft_in1k', pretrained=True):
super().__init__()
self.encoder = timm.create_model(encoder_name, pretrained=pretrained, features_only=True, out_indices=(0, 1, 2, 3))
enc_channels = self.encoder.feature_info.channels()
decoder_dims = [256, 128, 64, 32]
self.dec_blocks = nn.ModuleList()
in_ch = enc_channels[-1]
skip_channels = enc_channels[:-1][::-1] + [0]
for i, (skip_ch, out_ch) in enumerate(zip(skip_channels, decoder_dims)):
self.dec_blocks.append(UNetDecoderBlockV19(in_ch, skip_ch, out_ch, use_deform=(i >= 2)))
in_ch = out_ch
self.final_up = nn.Sequential(nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
nn.Conv2d(decoder_dims[-1], decoder_dims[-1], 3, padding=1, bias=False),
nn.BatchNorm2d(decoder_dims[-1]), nn.GELU())
self.height_attention = nn.Sequential(CoordConv2d(decoder_dims[-1], 64, 3, padding=1),
nn.BatchNorm2d(64), nn.GELU(), nn.Conv2d(64, 1, 1))
self.bilstm = BiLSTMHead(decoder_dims[-1], 128, 2, 0.1)
self.regression_head = nn.Sequential(nn.Linear(self.bilstm.output_dim, 64), nn.GELU(), nn.Linear(64, 1), nn.Sigmoid())
self.aux_heads = AuxiliaryHeads(self.bilstm.output_dim)
def forward(self, x, return_aux=False):
B, C, H, W = x.shape
features = self.encoder(x)
d = features[-1]
skips = features[:-1][::-1] + [None]
for block, skip in zip(self.dec_blocks, skips):
d = block(d, skip)
features_2d = d
d = self.final_up(d)
if d.shape[3] != W:
d = F.interpolate(d, size=(d.shape[2], W), mode='bilinear', align_corners=True)
attn = F.softmax(self.height_attention(d), dim=2)
pooled = (d * attn).sum(dim=2)
temporal_features = self.bilstm(pooled)
y_pred = self.regression_head(temporal_features).squeeze(-1)
if return_aux:
return y_pred, self.aux_heads(features_2d, temporal_features)
return y_pred
# =============================================================================
# Helper Functions
# =============================================================================
def apply_savgol_smoothing(signal_mv, window=7, polyorder=2):
if len(signal_mv) >= window:
return savgol_filter(signal_mv, window_length=window, polyorder=polyorder)
return signal_mv
def apply_einthoven_correction(pred_mv_rows, alpha=0.33):
segment_width = len(pred_mv_rows[0]) // 4
lead_I = pred_mv_rows[0][:segment_width].copy()
lead_II_short = pred_mv_rows[1][:segment_width].copy()
lead_III = pred_mv_rows[2][:segment_width].copy()
derived_II = lead_I + lead_III
error = lead_II_short - derived_II
pred_mv_rows[0][:segment_width] = lead_I + alpha * error
pred_mv_rows[2][:segment_width] = lead_III + alpha * error
return pred_mv_rows
def clamp_ecg_amplitude(signal_mv):
return np.clip(signal_mv, ECG_MV_MIN, ECG_MV_MAX)
def interpolate_nan(signal_1d):
valid_mask = np.isfinite(signal_1d)
if valid_mask.all():
return signal_1d
if not valid_mask.any():
return np.zeros_like(signal_1d)
x = np.arange(len(signal_1d))
signal_1d[~valid_mask] = np.interp(x[~valid_mask], x[valid_mask], signal_1d[valid_mask])
return signal_1d
def crop_row(image, row_idx):
baseline_y = int(ZERO_MV[row_idx])
y_start = max(0, baseline_y - CROP_HALF_HEIGHT)
y_end = min(TARGET_HEIGHT, baseline_y + CROP_HALF_HEIGHT)
row_crop = image[y_start:y_end, T0:T1, :].copy()
if row_crop.shape[0] < ROW_HEIGHT:
pad_top = max(0, CROP_HALF_HEIGHT - baseline_y)
pad_bottom = max(0, (baseline_y + CROP_HALF_HEIGHT) - TARGET_HEIGHT)
row_crop = np.pad(row_crop, ((pad_top, pad_bottom), (0, 0), (0, 0)), mode='edge')
return row_crop
def convert_crop_to_full(pred_y_crop, row_idx):
baseline_y = int(ZERO_MV[row_idx])
y_start = max(0, baseline_y - CROP_HALF_HEIGHT)
pad_top = max(0, CROP_HALF_HEIGHT - baseline_y)
return pred_y_crop - pad_top + y_start
def load_model(checkpoint_path, device):
model = PerLeadNetV19(encoder_name='convnext_base.fb_in22k_ft_in1k', pretrained=False)
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
state_dict = checkpoint['model']
if any(k.startswith('module.') for k in state_dict.keys()):
state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
model.load_state_dict(state_dict)
model = model.to(device)
model.eval()
print(f"Loaded V19 from epoch {checkpoint.get('epoch', 'N/A')}, SNR: {checkpoint.get('snr', checkpoint.get('best_snr', 'N/A')):.2f} dB")
return model
def compute_baseline_offset(pred_mv_segment, gt_mv):
x_pred = np.linspace(0, 1, len(pred_mv_segment))
x_gt = np.linspace(0, 1, len(gt_mv))
pred_resampled = np.interp(x_gt, x_pred, pred_mv_segment)
return np.median(pred_resampled - gt_mv)
def process_image(model, image, df, device):
segment_width = OUTPUT_WIDTH // 4
pred_mv_rows = {}
for row_idx in range(3):
row_crop = crop_row(image, row_idx)
image_tensor = torch.from_numpy(row_crop.astype(np.float32) / 255.0)
image_tensor = image_tensor.permute(2, 0, 1).unsqueeze(0).to(device)
with torch.no_grad():
with torch.amp.autocast('cuda'):
pred = model(image_tensor, return_aux=False)
pred_y_crop = pred[0].cpu().numpy() * ROW_HEIGHT
pred_y_full = convert_crop_to_full(pred_y_crop, row_idx)
baseline_y = ZERO_MV[row_idx]
pred_mv = (baseline_y - pred_y_full) / MV_TO_PIXEL
pred_mv = apply_savgol_smoothing(pred_mv, window=7, polyorder=2)
pred_mv = clamp_ecg_amplitude(pred_mv)
pred_mv = interpolate_nan(pred_mv.copy())
pred_mv_rows[row_idx] = pred_mv
pred_mv_rows = apply_einthoven_correction(pred_mv_rows, alpha=0.33)
offsets = {}
for row_idx in range(3):
lead_names = LEAD_LAYOUT[row_idx]
for seg_idx, lead_name in enumerate(lead_names):
if lead_name not in df.columns:
offsets[lead_name] = np.nan
continue
gt_mv = df[lead_name].dropna().values
if len(gt_mv) == 0:
offsets[lead_name] = np.nan
continue
if lead_name == 'II':
ref_lead = 'I' if 'I' in df.columns else 'III'
if ref_lead in df.columns:
ref_len = len(df[ref_lead].dropna().values)
if len(gt_mv) > ref_len:
gt_mv = gt_mv[:ref_len]
seg_start = seg_idx * segment_width
seg_end = (seg_idx + 1) * segment_width
pred_segment = pred_mv_rows[row_idx][seg_start:seg_end]
offsets[lead_name] = compute_baseline_offset(pred_segment, gt_mv)
return offsets
def describe_distribution(values, lead_name):
values = np.array([v for v in values if not np.isnan(v)])
if len(values) == 0:
return None
binned = np.round(values * 100) / 100
mode_result = stats.mode(binned, keepdims=True)
return {
'lead': lead_name, 'count': len(values), 'mean': float(np.mean(values)),
'median': float(np.median(values)), 'mode': float(mode_result.mode[0]),
'mode_count': int(mode_result.count[0]), 'std': float(np.std(values)),
'min': float(np.min(values)), 'max': float(np.max(values)),
'range': float(np.max(values) - np.min(values)),
'p5': float(np.percentile(values, 5)), 'p25': float(np.percentile(values, 25)),
'p75': float(np.percentile(values, 75)), 'p95': float(np.percentile(values, 95)),
'iqr': float(np.percentile(values, 75) - np.percentile(values, 25)),
}
def main():
parser = argparse.ArgumentParser(description='Compute per-lead baseline offsets for V19')
parser.add_argument('--checkpoint', type=str, required=True)
parser.add_argument('--kaggle_data', type=str, default='/data/ecg-digitization/stage1_data/train')
parser.add_argument('--num_samples', type=int, default=4000)
parser.add_argument('--output_dir', type=str, default='/home/azureuser/tmp/baseline_analysis_v19')
parser.add_argument('--seed', type=int, default=42)
args = parser.parse_args()
random.seed(args.seed)
np.random.seed(args.seed)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
print(f"{'='*80}")
print(f"Baseline Offset Analysis for V19")
print(f"{'='*80}")
print(f"Checkpoint: {args.checkpoint}")
print(f"Num samples: {args.num_samples}")
print(f"Device: {device}")
print(f"{'='*80}")
model = load_model(args.checkpoint, device)
kaggle_dir = Path(args.kaggle_data)
sample_paths = []
sample_dirs = list(kaggle_dir.iterdir())
random.shuffle(sample_dirs)
for sample_dir in sample_dirs:
if not sample_dir.is_dir():
continue
csv_files = list(sample_dir.glob('*.csv'))
if len(csv_files) != 1:
continue
csv_path = csv_files[0]
for img_path in sample_dir.glob('*.png'):
variant = img_path.stem.split('-')[-1] if '-' in img_path.stem else '0000'
if variant in VALID_VARIANTS:
sample_paths.append((img_path, csv_path))
break
if len(sample_paths) >= args.num_samples:
break
print(f"Found {len(sample_paths)} samples")
all_offsets = {lead: [] for lead in ALL_LEADS}
raw_data = []
for img_path, csv_path in tqdm(sample_paths, desc="Computing offsets"):
try:
image = cv2.imread(str(img_path), cv2.IMREAD_COLOR)
if image is None:
continue
image = image[Y0:Y1, X0:X1]
image = cv2.resize(image, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR)
df = pd.read_csv(csv_path)
offsets = process_image(model, image, df, device)
row = {'image': img_path.stem}
for lead in ALL_LEADS:
offset = offsets.get(lead, np.nan)
row[lead] = offset
if not np.isnan(offset):
all_offsets[lead].append(offset)
raw_data.append(row)
except Exception as e:
print(f"Error: {e}")
raw_df = pd.DataFrame(raw_data)
raw_df.to_csv(output_dir / 'baseline_offsets_raw.csv', index=False)
print(f"\n{'='*80}")
print(f"BASELINE OFFSET STATISTICS (mV) - V19")
print(f"{'='*80}")
print(f"{'Lead':<8} {'Count':>6} {'Mean':>8} {'Median':>8} {'Std':>8}")
print("-" * 50)
stats_data = []
for lead in ALL_LEADS:
s = describe_distribution(all_offsets[lead], lead)
if s:
stats_data.append(s)
print(f"{lead:<8} {s['count']:>6} {s['mean']:>8.4f} {s['median']:>8.4f} {s['std']:>8.4f}")
print(f"\n{'='*80}")
print(f"RECOMMENDED HARDCODED BASELINE OFFSETS (mV) - V19")
print(f"{'='*80}\n")
recommended = {}
print("BASELINE_OFFSETS_V19 = {")
for lead in ALL_LEADS:
if all_offsets[lead]:
median_offset = float(np.median(all_offsets[lead]))
recommended[lead] = median_offset
print(f" '{lead}': {median_offset:.4f},")
print("}")
with open(output_dir / 'baseline_stats.json', 'w') as f:
json.dump({'stats': stats_data, 'recommended_offsets': recommended,
'num_samples': len(sample_paths), 'checkpoint': args.checkpoint}, f, indent=2)
print(f"\nSaved to: {output_dir}")
if __name__ == '__main__':
main()