| |
| """ |
| V18+V19 Ensemble Inference |
| |
| Ensembles V18 (V16 + Cross-Row Refiner) and V19 (BiLSTM + Deformable Conv) models. |
| Allows configurable blending ratio between the two models. |
| |
| Usage: |
| python scripts/infer_v18_v19_ensemble.py --v18_weight 0.5 --v19_weight 0.5 |
| python scripts/infer_v18_v19_ensemble.py --v18_weight 0.3 --v19_weight 0.7 |
| """ |
|
|
| 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 |
|
|
| 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 |
|
|
|
|
| |
| |
| |
| 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 |
|
|
| |
| LOW_SNR_THRESHOLD = 5.0 |
|
|
| VALID_VARIANTS = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012'] |
|
|
| VAL_SAMPLE_IDS = [ |
| '1006427285', '1006867983', '1012423188', '10140238', '1015663939', |
| '102150619', '1026034238', '1041099777', '104573050', '1048962695', |
| '1052007218', '1053922973', '1059602762', '1063816858', |
| '1067371646', '1067975047', '1068062585', '1072767337', |
| '1084993373', '108599929' |
| ] |
|
|
| LEAD_LAYOUT = [ |
| ['I', 'aVR', 'V1', 'V4'], |
| ['II', 'aVL', 'V2', 'V5'], |
| ['III', 'aVF', 'V3', 'V6'], |
| ] |
|
|
| |
| V18_BASELINE_OFFSETS = { |
| 'I': 0.0003, 'II': 0.0005, 'III': 0.0002, |
| 'aVR': 0.0003, 'aVL': 0.0003, 'aVF': 0.0003, |
| 'V1': 0.0003, 'V2': 0.0003, 'V3': 0.0003, |
| 'V4': 0.0003, 'V5': 0.0003, 'V6': 0.0003, |
| } |
|
|
| |
| V19_BASELINE_OFFSETS = { |
| 'I': 0.0037, 'II': 0.0057, 'III': 0.0029, |
| 'aVR': 0.0043, 'aVL': 0.0037, 'aVF': 0.0038, |
| 'V1': 0.0045, 'V2': 0.0036, 'V3': 0.0038, |
| 'V4': 0.0045, 'V5': 0.0040, 'V6': 0.0038, |
| } |
|
|
| |
| LOCAL_CHECKPOINT_DIR = '/data/ecg-digitization/checkpoints' |
|
|
|
|
| |
| |
| |
| 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) |
| x = torch.cat([x, yy, xx], dim=1) |
| return self.conv(x) |
|
|
|
|
| class UNetDecoderBlock(nn.Module): |
| def __init__(self, in_ch, skip_ch, out_ch): |
| super().__init__() |
| self.conv = nn.Sequential( |
| nn.Conv2d(in_ch + skip_ch, out_ch, 3, padding=1, bias=False), |
| nn.BatchNorm2d(out_ch), |
| nn.GELU(), |
| 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.conv(x) |
|
|
|
|
| class PerLeadNetV16(nn.Module): |
| """V16 Per-Lead ECG Network (frozen, used as prior for V18).""" |
| |
| 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 skip_ch, out_ch in zip(skip_channels, decoder_dims): |
| self.dec_blocks.append(UNetDecoderBlock(in_ch, skip_ch, out_ch)) |
| 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.regression_head = nn.Sequential( |
| nn.Conv1d(decoder_dims[-1], 128, 7, padding=3), |
| nn.BatchNorm1d(128), |
| nn.GELU(), |
| nn.Conv1d(128, 64, 5, padding=2), |
| nn.BatchNorm1d(64), |
| nn.GELU(), |
| nn.Conv1d(64, 1, 1), |
| ) |
| |
| def forward(self, x): |
| 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) |
| |
| 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 = self.height_attention(d) |
| attn = F.softmax(attn, dim=2) |
| d = (d * attn).sum(dim=2) |
| out = self.regression_head(d) |
| out = torch.sigmoid(out) |
| |
| return out.squeeze(1) |
|
|
|
|
| |
| |
| |
| class CrossRowAttention(nn.Module): |
| def __init__(self, embed_dim, num_heads=8, dropout=0.1): |
| super().__init__() |
| self.num_heads = num_heads |
| self.head_dim = embed_dim // num_heads |
| self.scale = self.head_dim ** -0.5 |
| |
| self.qkv = nn.Linear(embed_dim, embed_dim * 3) |
| self.proj = nn.Linear(embed_dim, embed_dim) |
| self.dropout = nn.Dropout(dropout) |
| self.norm = nn.LayerNorm(embed_dim) |
| |
| def forward(self, x): |
| B, num_rows, C, W = x.shape |
| x = x.permute(0, 3, 1, 2).reshape(B * W, num_rows, C) |
| |
| residual = x |
| x = self.norm(x) |
| |
| qkv = self.qkv(x).reshape(B * W, num_rows, 3, self.num_heads, self.head_dim) |
| qkv = qkv.permute(2, 0, 3, 1, 4) |
| q, k, v = qkv[0], qkv[1], qkv[2] |
| |
| attn = (q @ k.transpose(-2, -1)) * self.scale |
| attn = attn.softmax(dim=-1) |
| attn = self.dropout(attn) |
| |
| out = (attn @ v).transpose(1, 2).reshape(B * W, num_rows, C) |
| out = self.proj(out) |
| out = self.dropout(out) |
| out = out + residual |
| |
| out = out.reshape(B, W, num_rows, C).permute(0, 2, 3, 1) |
| return out |
|
|
|
|
| class CrossRowTransformerBlock(nn.Module): |
| def __init__(self, embed_dim, num_heads=8, mlp_ratio=4.0, dropout=0.1): |
| super().__init__() |
| self.attn = CrossRowAttention(embed_dim, num_heads, dropout) |
| self.norm = nn.LayerNorm(embed_dim) |
| hidden_dim = int(embed_dim * mlp_ratio) |
| self.ffn = nn.Sequential( |
| nn.Linear(embed_dim, hidden_dim), |
| nn.GELU(), |
| nn.Dropout(dropout), |
| nn.Linear(hidden_dim, embed_dim), |
| nn.Dropout(dropout), |
| ) |
| |
| def forward(self, x): |
| x = self.attn(x) |
| B, num_rows, C, W = x.shape |
| residual = x |
| x = x.permute(0, 3, 1, 2).reshape(B * W, num_rows, C) |
| x = self.norm(x) |
| x = self.ffn(x) + residual.permute(0, 3, 1, 2).reshape(B * W, num_rows, C) |
| x = x.reshape(B, W, num_rows, C).permute(0, 2, 3, 1) |
| return x |
|
|
|
|
| class RefinerEncoder(nn.Module): |
| def __init__(self, encoder_name='efficientnet_b0', pretrained=True): |
| super().__init__() |
| self.encoder = timm.create_model( |
| encoder_name, |
| pretrained=pretrained, |
| features_only=True, |
| out_indices=(1, 2, 3), |
| in_chans=4, |
| ) |
| self.channels = self.encoder.feature_info.channels() |
| |
| def forward(self, x): |
| return self.encoder(x) |
|
|
|
|
| class V18RefinerNet(nn.Module): |
| """V18 Refiner: Takes V16 predictions and refines them.""" |
| |
| def __init__(self, encoder_name='efficientnet_b0', pretrained=True, |
| cross_row_layers=3, cross_row_dim=128, num_heads=4): |
| super().__init__() |
| |
| self.row_encoder = RefinerEncoder(encoder_name, pretrained) |
| enc_channels = self.row_encoder.channels |
| |
| self.feature_proj = nn.Sequential( |
| nn.AdaptiveAvgPool2d((1, None)), |
| nn.Flatten(1, 2), |
| ) |
| self.channel_proj = nn.Conv1d(enc_channels[-1], cross_row_dim, 1) |
| |
| self.cross_row_blocks = nn.ModuleList([ |
| CrossRowTransformerBlock( |
| embed_dim=cross_row_dim, |
| num_heads=num_heads, |
| mlp_ratio=2.0, |
| dropout=0.1 |
| ) |
| for _ in range(cross_row_layers) |
| ]) |
| |
| self.residual_head = nn.Sequential( |
| nn.Conv1d(cross_row_dim, 64, 5, padding=2), |
| nn.BatchNorm1d(64), |
| nn.GELU(), |
| nn.Conv1d(64, 32, 3, padding=1), |
| nn.BatchNorm1d(32), |
| nn.GELU(), |
| nn.Conv1d(32, 1, 1), |
| nn.Tanh(), |
| ) |
| |
| self.residual_scale = nn.Parameter(torch.tensor(0.1)) |
| |
| def create_guide_channel(self, v16_pred, height, sigma=15.0): |
| B, W = v16_pred.shape |
| device = v16_pred.device |
| y_pred = v16_pred * height |
| y_grid = torch.arange(height, device=device, dtype=torch.float32).view(1, height, 1) |
| y_pred = y_pred.unsqueeze(1) |
| guide = torch.exp(-0.5 * ((y_grid - y_pred) / sigma) ** 2) |
| return guide.unsqueeze(1) |
| |
| def forward(self, images, v16_preds): |
| B, num_rows, C, H, W = images.shape |
| all_features = [] |
| |
| for row_idx in range(num_rows): |
| row_image = images[:, row_idx] |
| row_v16 = v16_preds[:, row_idx] |
| guide = self.create_guide_channel(row_v16, H) |
| row_input = torch.cat([row_image, guide], dim=1) |
| enc_features = self.row_encoder(row_input) |
| row_feat = enc_features[-1] |
| row_feat = self.feature_proj(row_feat) |
| row_feat = F.interpolate(row_feat, size=W, mode='linear', align_corners=True) |
| row_feat = self.channel_proj(row_feat) |
| all_features.append(row_feat) |
| |
| features = torch.stack(all_features, dim=1) |
| |
| for block in self.cross_row_blocks: |
| features = block(features) |
| |
| residuals = [] |
| for row_idx in range(num_rows): |
| row_feat = features[:, row_idx] |
| residual = self.residual_head(row_feat) |
| residuals.append(residual.squeeze(1)) |
| |
| residuals = torch.stack(residuals, dim=1) |
| scaled_residuals = residuals * self.residual_scale * 0.1 |
| refined = torch.clamp(v16_preds + scaled_residuals, 0, 1) |
| |
| return refined, scaled_residuals |
|
|
|
|
| |
| |
| |
| 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 |
| self.padding = padding |
| self.stride = stride |
| |
| 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), |
| ) |
| nn.init.zeros_(self.offset_conv[-1].weight) |
| nn.init.zeros_(self.offset_conv[-1].bias) |
| 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) |
| out = self.norm(out) |
| out = self.act(out) |
| return 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): |
| grid_pred = self.grid_head(features_2d) |
| gradient_pred = self.gradient_head(features_1d).squeeze(-1) |
| log_var = self.uncertainty_head(features_1d).squeeze(-1) |
| return grid_pred, gradient_pred, log_var |
|
|
|
|
| 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) |
| x = self.conv1(x) |
| x = self.conv2(x) |
| return x |
|
|
|
|
| class PerLeadNetV19(nn.Module): |
| """V19 Per-Lead ECG Network with BiLSTM + Deformable Conv.""" |
| |
| 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)): |
| use_deform = (i >= 2) |
| self.dec_blocks.append(UNetDecoderBlockV19(in_ch, skip_ch, out_ch, use_deform)) |
| 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( |
| input_dim=decoder_dims[-1], |
| hidden_dim=128, |
| num_layers=2, |
| dropout=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 = self.height_attention(d) |
| attn = F.softmax(attn, 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: |
| grid_pred, gradient_pred, log_var = self.aux_heads(features_2d, temporal_features) |
| return y_pred, {'grid': grid_pred, 'gradient': gradient_pred, 'log_var': log_var} |
| |
| return y_pred |
|
|
|
|
| |
| |
| |
| 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 apply_baseline_correction(pred_mv_rows, offsets): |
| segment_width = len(pred_mv_rows[0]) // 4 |
| |
| for row_idx in range(3): |
| lead_names = LEAD_LAYOUT[row_idx] |
| for seg_idx, lead_name in enumerate(lead_names): |
| offset = offsets.get(lead_name, 0.0) |
| seg_start = seg_idx * segment_width |
| seg_end = (seg_idx + 1) * segment_width |
| pred_mv_rows[row_idx][seg_start:seg_end] -= offset |
| |
| pred_mv_rows[3] -= offsets.get('II', 0.0) |
| return pred_mv_rows |
|
|
|
|
| |
| |
| |
| def load_v18_models(v16_checkpoint, v18_checkpoint, device): |
| """Load V16 + V18 refiner models.""" |
| |
| v16_model = PerLeadNetV16(encoder_name='convnext_base.fb_in22k_ft_in1k', pretrained=False) |
| v16_ckpt = torch.load(v16_checkpoint, map_location=device, weights_only=False) |
| state_dict = v16_ckpt['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()} |
| v16_model.load_state_dict(state_dict) |
| v16_model = v16_model.to(device).eval() |
| |
| |
| v18_refiner = V18RefinerNet(encoder_name='efficientnet_b0', pretrained=False) |
| v18_ckpt = torch.load(v18_checkpoint, map_location=device, weights_only=False) |
| v18_refiner.load_state_dict(v18_ckpt['refiner']) |
| v18_refiner = v18_refiner.to(device).eval() |
| |
| v16_epoch = v16_ckpt.get('epoch', '?') |
| v16_snr = v16_ckpt.get('snr', 0) |
| v18_epoch = v18_ckpt.get('epoch', '?') |
| v18_snr = v18_ckpt.get('snr', 0) |
| |
| print(f"Loaded V16 epoch {v16_epoch}, SNR: {v16_snr:.2f} dB") |
| print(f"Loaded V18 epoch {v18_epoch}, SNR: {v18_snr:.2f} dB") |
| |
| return v16_model, v18_refiner |
|
|
|
|
| def load_v19_model(checkpoint, device): |
| """Load V19 model.""" |
| model = PerLeadNetV19(encoder_name='convnext_base.fb_in22k_ft_in1k', pretrained=False) |
| ckpt = torch.load(checkpoint, map_location=device, weights_only=False) |
| |
| state_dict = ckpt['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).eval() |
| |
| epoch = ckpt.get('epoch', '?') |
| snr = ckpt.get('snr', ckpt.get('best_snr', 0)) |
| print(f"Loaded V19 epoch {epoch}, SNR: {snr:.2f} dB") |
| |
| return model |
|
|
|
|
| |
| |
| |
| 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 |
|
|
|
|
| @torch.no_grad() |
| def predict_v18(v16_model, v18_refiner, row_crops, device): |
| """Run V16 + V18 refiner on all 4 row crops.""" |
| row_tensors = [] |
| for crop in row_crops: |
| t = torch.from_numpy(crop.astype(np.float32) / 255.0).permute(2, 0, 1) |
| row_tensors.append(t) |
| images = torch.stack(row_tensors, dim=0).unsqueeze(0).to(device) |
| |
| with torch.amp.autocast('cuda', dtype=torch.float16): |
| v16_preds = [] |
| for row_idx in range(4): |
| row_pred = v16_model(images[:, row_idx]) |
| v16_preds.append(row_pred) |
| v16_preds = torch.stack(v16_preds, dim=1) |
| |
| refined, _ = v18_refiner(images, v16_preds) |
| |
| refined_np = refined[0].cpu().numpy() |
| pred_mv_rows = {} |
| |
| for row_idx in range(4): |
| pred_y_crop = refined_np[row_idx] * ROW_HEIGHT |
| 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) |
| pred_y_full = pred_y_crop - pad_top + y_start |
| pred_mv = (ZERO_MV[row_idx] - pred_y_full) / MV_TO_PIXEL |
| pred_mv_rows[row_idx] = pred_mv |
| |
| return pred_mv_rows |
|
|
|
|
| @torch.no_grad() |
| def predict_v19(model, row_crops, device): |
| """Run V19 on all 4 row crops.""" |
| pred_mv_rows = {} |
| |
| for row_idx in range(4): |
| row_crop = row_crops[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.amp.autocast('cuda'): |
| output = model(image_tensor, return_aux=False) |
| |
| pred_y_crop = output[0].cpu().numpy() * ROW_HEIGHT |
| |
| 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) |
| pred_y_full = pred_y_crop - pad_top + y_start |
| pred_mv = (ZERO_MV[row_idx] - pred_y_full) / MV_TO_PIXEL |
| pred_mv_rows[row_idx] = pred_mv |
| |
| return pred_mv_rows |
|
|
|
|
| def ensemble_predictions(v18_pred, v19_pred, v18_weight, v19_weight): |
| """Ensemble V18 and V19 predictions with given weights.""" |
| ensemble = {} |
| for row_idx in range(4): |
| ensemble[row_idx] = v18_weight * v18_pred[row_idx] + v19_weight * v19_pred[row_idx] |
| return ensemble |
|
|
|
|
| def visualize_predictions(image, predictions, output_path): |
| """Draw predictions on the image.""" |
| vis_image = image[:, T0:T1, :].copy() |
| colors = [(255, 0, 0), (0, 255, 0), (255, 0, 255), (0, 165, 255)] |
| |
| for row_idx, pred_mv in predictions.items(): |
| color = colors[row_idx] |
| baseline_y = ZERO_MV[row_idx] |
| pred_y = baseline_y - pred_mv * MV_TO_PIXEL |
| |
| for x in range(len(pred_y)): |
| y = int(round(pred_y[x])) |
| if 0 <= y < vis_image.shape[0]: |
| cv2.circle(vis_image, (x, y), 1, color, -1) |
| |
| cv2.imwrite(str(output_path), vis_image) |
|
|
|
|
| def compute_snr_per_lead(pred_mv_rows, df, epsilon=1e-10): |
| """Compute SNR in dB for each lead.""" |
| segment_width = OUTPUT_WIDTH // 4 |
| all_lead_snrs = {} |
| |
| for row_idx in range(3): |
| lead_names = LEAD_LAYOUT[row_idx] |
| baseline_y = ZERO_MV[row_idx] |
| |
| for seg_idx, lead_name in enumerate(lead_names): |
| if lead_name not in df.columns: |
| all_lead_snrs[lead_name] = None |
| continue |
| |
| gt_mv = df[lead_name].dropna().values |
| if len(gt_mv) == 0: |
| all_lead_snrs[lead_name] = None |
| 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_mv = pred_mv_rows[row_idx][seg_start:seg_end] |
| |
| x_pred = np.linspace(0, 1, len(pred_mv)) |
| x_gt = np.linspace(0, 1, len(gt_mv)) |
| pred_mv_resampled = np.interp(x_gt, x_pred, pred_mv) |
| |
| signal_power = (gt_mv ** 2).mean() |
| noise_power = ((pred_mv_resampled - gt_mv) ** 2).mean() |
| |
| if noise_power < epsilon: |
| all_lead_snrs[lead_name] = 50.0 |
| else: |
| all_lead_snrs[lead_name] = 10 * np.log10(signal_power / noise_power) |
| |
| |
| if 'II' in df.columns: |
| gt_mv = df['II'].dropna().values |
| if len(gt_mv) > 0: |
| baseline_y = ZERO_MV[3] |
| pred_mv = pred_mv_rows[3] |
| |
| x_pred = np.linspace(0, 1, len(pred_mv)) |
| x_gt = np.linspace(0, 1, len(gt_mv)) |
| pred_mv_resampled = np.interp(x_gt, x_pred, pred_mv) |
| |
| signal_power = (gt_mv ** 2).mean() |
| noise_power = ((pred_mv_resampled - gt_mv) ** 2).mean() |
| |
| if noise_power < epsilon: |
| all_lead_snrs['II_rhythm'] = 50.0 |
| else: |
| all_lead_snrs['II_rhythm'] = 10 * np.log10(signal_power / noise_power) |
| |
| return all_lead_snrs |
|
|
|
|
| def process_image(v16_model, v18_refiner, v19_model, image_path, csv_path, output_dir, device, |
| v18_weight=0.5, v19_weight=0.5, |
| apply_smoothing=True, apply_einthoven=True, apply_baseline_fix=True, |
| negative_dir=None): |
| """Process a single image with ensemble.""" |
| image = cv2.imread(str(image_path), cv2.IMREAD_COLOR) |
| if image is None: |
| print(f"Failed to load: {image_path}") |
| return None, None, None |
| |
| image = image[Y0:Y1, X0:X1] |
| image = cv2.resize(image, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR) |
| |
| df = pd.read_csv(csv_path) |
| |
| |
| row_crops = [crop_row(image, row_idx) for row_idx in range(4)] |
| |
| |
| v18_pred = predict_v18(v16_model, v18_refiner, row_crops, device) |
| v19_pred = predict_v19(v19_model, row_crops, device) |
| |
| |
| pred_mv_rows = ensemble_predictions(v18_pred, v19_pred, v18_weight, v19_weight) |
| |
| |
| for row_idx in range(4): |
| if apply_smoothing: |
| pred_mv_rows[row_idx] = apply_savgol_smoothing(pred_mv_rows[row_idx], window=7, polyorder=2) |
| pred_mv_rows[row_idx] = clamp_ecg_amplitude(pred_mv_rows[row_idx]) |
| pred_mv_rows[row_idx] = interpolate_nan(pred_mv_rows[row_idx].copy()) |
| |
| if apply_einthoven: |
| pred_mv_rows = apply_einthoven_correction(pred_mv_rows, alpha=0.33) |
| |
| if apply_baseline_fix: |
| |
| blended_offsets = {} |
| for lead in V18_BASELINE_OFFSETS.keys(): |
| blended_offsets[lead] = (v18_weight * V18_BASELINE_OFFSETS[lead] + |
| v19_weight * V19_BASELINE_OFFSETS[lead]) |
| pred_mv_rows = apply_baseline_correction(pred_mv_rows, blended_offsets) |
| |
| |
| all_lead_snrs = compute_snr_per_lead(pred_mv_rows, df) |
| |
| valid_snrs = [v for v in all_lead_snrs.values() if v is not None] |
| min_snr = min(valid_snrs) if valid_snrs else 0.0 |
| |
| |
| sample_id = image_path.parent.name |
| variant = image_path.stem.split('-')[-1] if '-' in image_path.stem else '0000' |
| output_path = output_dir / f"{sample_id}_{variant}.png" |
| |
| visualize_predictions(image, pred_mv_rows, output_path) |
| |
| if negative_dir is not None and min_snr < LOW_SNR_THRESHOLD: |
| neg_output_path = negative_dir / f"{sample_id}_{variant}_snr{min_snr:.1f}.png" |
| visualize_predictions(image, pred_mv_rows, neg_output_path) |
| |
| return all_lead_snrs, min_snr, output_path |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description='V18+V19 Ensemble Inference') |
| parser.add_argument('--v16_checkpoint', type=str, |
| default=f'{LOCAL_CHECKPOINT_DIR}/v16_perlead_best_snr.pth', |
| help='V16 checkpoint path') |
| parser.add_argument('--v18_checkpoint', type=str, |
| default=f'{LOCAL_CHECKPOINT_DIR}/v18_refiner_best.pth', |
| help='V18 refiner checkpoint path') |
| parser.add_argument('--v19_checkpoint', type=str, |
| default=f'{LOCAL_CHECKPOINT_DIR}/v19_enhanced_epoch010.pth', |
| help='V19 checkpoint path') |
| parser.add_argument('--v18_weight', type=float, default=0.5, |
| help='Weight for V18 predictions (0 to 1)') |
| parser.add_argument('--v19_weight', type=float, default=0.5, |
| help='Weight for V19 predictions (0 to 1)') |
| parser.add_argument('--kaggle_data', type=str, |
| default='/data/ecg-digitization/stage1_data/train') |
| parser.add_argument('--output_dir', type=str, |
| default=os.path.expanduser('~/tmp/pred/v18v19_ensemble')) |
| parser.add_argument('--num_samples', type=int, default=None) |
| parser.add_argument('--use_holdout', action='store_true', default=True) |
| parser.add_argument('--no_smoothing', action='store_true') |
| parser.add_argument('--no_einthoven', action='store_true') |
| parser.add_argument('--no_baseline_fix', action='store_true') |
| parser.add_argument('--seed', type=int, default=42) |
| args = parser.parse_args() |
| |
| |
| total_weight = args.v18_weight + args.v19_weight |
| if total_weight <= 0: |
| print("Error: Weights must sum to > 0") |
| sys.exit(1) |
| v18_weight = args.v18_weight / total_weight |
| v19_weight = args.v19_weight / total_weight |
| |
| random.seed(args.seed) |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| output_dir = Path(args.output_dir) |
| negative_dir = output_dir / 'negpreds' |
| |
| print(f"{'='*70}") |
| print(f"V18 + V19 Ensemble Inference") |
| print(f"{'='*70}") |
| print(f"V18 weight: {v18_weight:.2%}") |
| print(f"V19 weight: {v19_weight:.2%}") |
| print(f"Device: {device}") |
| print(f"Deformable Conv: {HAS_DEFORM_CONV}") |
| print(f"{'='*70}") |
| |
| |
| print("\nLoading V18 (V16 + Refiner)...") |
| v16_model, v18_refiner = load_v18_models(args.v16_checkpoint, args.v18_checkpoint, device) |
| |
| print("\nLoading V19...") |
| v19_model = load_v19_model(args.v19_checkpoint, device) |
| |
| |
| if output_dir.exists(): |
| import shutil |
| shutil.rmtree(output_dir) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| negative_dir.mkdir(parents=True, exist_ok=True) |
| |
| print(f"\nOutput: {output_dir}") |
| print(f"Smoothing: {'OFF' if args.no_smoothing else 'ON'}") |
| print(f"Einthoven: {'OFF' if args.no_einthoven else 'ON'}") |
| print(f"Baseline fix: {'OFF' if args.no_baseline_fix else 'ON'}") |
| print(f"{'='*70}") |
| |
| |
| kaggle_dir = Path(args.kaggle_data) |
| val_sample_set = set(VAL_SAMPLE_IDS) |
| all_samples = [] |
| |
| for sample_dir in kaggle_dir.iterdir(): |
| if not sample_dir.is_dir(): |
| continue |
| if args.use_holdout and sample_dir.name not in val_sample_set: |
| 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: |
| all_samples.append((img_path, csv_path)) |
| |
| print(f"Found {len(all_samples)} valid images") |
| |
| if args.num_samples is not None and len(all_samples) > args.num_samples: |
| selected = random.sample(all_samples, args.num_samples) |
| else: |
| selected = all_samples |
| |
| print(f"Processing {len(selected)} images...") |
| |
| all_snrs = {lead: [] for lead in ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', |
| 'V1', 'V2', 'V3', 'V4', 'V5', 'V6', 'II_rhythm']} |
| |
| for img_path, csv_path in tqdm(selected): |
| try: |
| lead_snrs, min_snr, _ = process_image( |
| v16_model, v18_refiner, v19_model, |
| img_path, csv_path, output_dir, device, |
| v18_weight=v18_weight, v19_weight=v19_weight, |
| apply_smoothing=not args.no_smoothing, |
| apply_einthoven=not args.no_einthoven, |
| apply_baseline_fix=not args.no_baseline_fix, |
| negative_dir=negative_dir |
| ) |
| if lead_snrs: |
| for lead, snr in lead_snrs.items(): |
| if snr is not None: |
| all_snrs[lead].append(snr) |
| except Exception as e: |
| print(f"Error processing {img_path}: {e}") |
| import traceback |
| traceback.print_exc() |
| |
| |
| print(f"\n{'='*70}") |
| print(f"Per-Lead SNR Statistics (dB) - V18:{v18_weight:.0%} + V19:{v19_weight:.0%}") |
| print(f"{'='*70}") |
| print(f"{'Lead':<12} {'Mean':>8} {'Std':>8} {'Min':>8} {'Max':>8} {'Count':>6}") |
| print(f"{'-'*70}") |
| |
| total_snrs = [] |
| for lead in ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6', 'II_rhythm']: |
| snrs = all_snrs[lead] |
| if len(snrs) > 0: |
| print(f"{lead:<12} {np.mean(snrs):>8.2f} {np.std(snrs):>8.2f} " |
| f"{np.min(snrs):>8.2f} {np.max(snrs):>8.2f} {len(snrs):>6}") |
| total_snrs.extend(snrs) |
| |
| print(f"{'-'*70}") |
| if len(total_snrs) > 0: |
| print(f"{'OVERALL':<12} {np.mean(total_snrs):>8.2f} {np.std(total_snrs):>8.2f} " |
| f"{np.min(total_snrs):>8.2f} {np.max(total_snrs):>8.2f} {len(total_snrs):>6}") |
| print(f"{'='*70}") |
| |
| print(f"\nDone! Saved {len(selected)} visualizations to {output_dir}") |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|