#!/usr/bin/env python3 """ Find Best V18 + V19 + Net3 Ensemble Ratio Runs inference ONCE for all three models, saves predictions, then sweeps through different blend ratios to find the optimal ensemble weights. """ 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 pickle 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'], ] LOCAL_CHECKPOINT_DIR = '/data/ecg-digitization/checkpoints' BASELINE_PATH = '/data/ecg-digitization/hengck23-submit-physionet' # ============================================================================= # V16 Model # ============================================================================= 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): 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) return torch.sigmoid(out).squeeze(1) # ============================================================================= # V18 Refiner # ============================================================================= 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).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) + residual return out.reshape(B, W, num_rows, C).permute(0, 2, 3, 1) 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.ffn(self.norm(x)) + residual.permute(0, 3, 1, 2).reshape(B * W, num_rows, C) return x.reshape(B, W, num_rows, C).permute(0, 2, 3, 1) 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): 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(cross_row_dim, num_heads, 2.0, 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 y_pred = v16_pred * height y_grid = torch.arange(height, device=v16_pred.device, dtype=torch.float32).view(1, height, 1) return torch.exp(-0.5 * ((y_grid - y_pred.unsqueeze(1)) / sigma) ** 2).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): guide = self.create_guide_channel(v16_preds[:, row_idx], H) row_input = torch.cat([images[:, row_idx], guide], dim=1) row_feat = self.row_encoder(row_input)[-1] row_feat = self.feature_proj(row_feat) row_feat = F.interpolate(row_feat, size=W, mode='linear', align_corners=True) all_features.append(self.channel_proj(row_feat)) features = torch.stack(all_features, dim=1) for block in self.cross_row_blocks: features = block(features) residuals = torch.stack([self.residual_head(features[:, i]).squeeze(1) for i in range(num_rows)], dim=1) scaled_residuals = residuals * self.residual_scale * 0.1 return torch.clamp(v16_preds + scaled_residuals, 0, 1), scaled_residuals # ============================================================================= # V19 Model # ============================================================================= class DeformableConvBlock(nn.Module): def __init__(self, in_ch, out_ch, kernel_size=3, stride=1, padding=1): super().__init__() 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: out = self.deform_conv(x, self.offset_conv(x)) 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_dim, hidden_dim, 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): return self.output_proj(self.lstm(x.permute(0, 2, 1))[0]) 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 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, 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, {'grid': self.aux_heads.grid_head(features_2d), 'gradient': self.aux_heads.gradient_head(temporal_features).squeeze(-1), 'log_var': self.aux_heads.uncertainty_head(temporal_features).squeeze(-1)} return y_pred # ============================================================================= # Net3 Model (from public solution - matches checkpoint structure) # ============================================================================= class MyCoordUnetBlock(nn.Module): """Matches checkpoint structure: conv1 and conv2 are Sequential with [Conv, BN]""" def __init__(self, in_channel, skip_channel, out_channel, scale): super().__init__() self.scale = scale in_ch = in_channel + skip_channel + 2 # +2 for coord channels # conv1: Sequential[0=Conv2d, 1=BatchNorm2d] self.conv1 = nn.Sequential( nn.Conv2d(in_ch, out_channel, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channel) ) # conv2: Sequential[0=Conv2d, 1=BatchNorm2d] self.conv2 = nn.Sequential( nn.Conv2d(out_channel, out_channel, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channel) ) def add_coord(self, x): B, C, H, W = x.shape h = torch.linspace(-1, 1, H, device=x.device).view(1, 1, H, 1).expand(B, 1, H, W) w = torch.linspace(-1, 1, W, device=x.device).view(1, 1, 1, W).expand(B, 1, H, W) return torch.cat([x, h, w], 1) def forward(self, x, skip=None): x = F.interpolate(x, scale_factor=self.scale, mode='bilinear', align_corners=False) if skip is not None: x = torch.cat([x, skip], dim=1) x = self.add_coord(x) x = F.relu(self.conv1(x), inplace=True) x = F.relu(self.conv2(x), inplace=True) return x class MyCoordUnetDecoder(nn.Module): def __init__(self, in_channel, skip_channel, out_channel, scale): super().__init__() self.block = nn.ModuleList() for i, (s, o, sc) in enumerate(zip(scale, out_channel, skip_channel)): self.block.append(MyCoordUnetBlock(in_channel if i == 0 else out_channel[i-1], sc, o, s)) def forward(self, feature, skip): d = feature for i, b in enumerate(self.block): d = b(d, skip[i]) return d, None def encode_with_resnet(encoder, image): x = encoder.conv1(image) x = encoder.bn1(x) x = encoder.act1(x) x = encoder.maxpool(x) e1 = encoder.layer1(x) e2 = encoder.layer2(e1) e3 = encoder.layer3(e2) e4 = encoder.layer4(e3) return [e1, e2, e3, e4] class Net3(nn.Module): def __init__(self, pretrained=True): super().__init__() encoder_dim = [64, 128, 256, 512] decoder_dim = [128, 64, 32, 16] self.encoder = timm.create_model('resnet34.a3_in1k', pretrained=pretrained, in_chans=3, num_classes=0, global_pool='') self.decoder = MyCoordUnetDecoder(in_channel=encoder_dim[-1], skip_channel=encoder_dim[:-1][::-1] + [0], out_channel=decoder_dim, scale=[2, 2, 2, 2]) self.pixel = nn.Conv2d(decoder_dim[-1], 4, 1) def forward(self, image): encode = encode_with_resnet(self.encoder, image) last, _ = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + [None]) return self.pixel(last) # ============================================================================= # Helper Functions # ============================================================================= 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 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 pixel_to_series(pixel, zero_mv, length): """Convert Net3 pixel predictions to series.""" num_row, W = pixel.shape[0], pixel.shape[-1] series = np.zeros((num_row, length), dtype=np.float32) for r in range(num_row): p = pixel[r] if p.ndim == 2: p = p.argmax(axis=0) if p.shape[0] > 1 else p[0] x_old = np.linspace(0, 1, len(p)) x_new = np.linspace(0, 1, length) series[r] = np.interp(x_new, x_old, p.astype(np.float32)) return series def compute_snr(pred_mv_rows, df, epsilon=1e-10): """Compute overall SNR for a prediction.""" segment_width = OUTPUT_WIDTH // 4 all_snrs = [] 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: continue gt_mv = df[lead_name].dropna().values if len(gt_mv) == 0: 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_snrs.append(10 * np.log10(signal_power / noise_power)) # Lead II rhythm strip if 'II' in df.columns: gt_mv = df['II'].dropna().values if len(gt_mv) > 0: 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_snrs.append(10 * np.log10(signal_power / noise_power)) return np.mean(all_snrs) if all_snrs else 0.0 def load_models(v16_ckpt, v18_ckpt, v19_ckpt, net3_ckpt, device): # V16 v16_model = PerLeadNetV16(pretrained=False) ckpt = torch.load(v16_ckpt, map_location=device, weights_only=False) state = ckpt['model'] if any(k.startswith('module.') for k in state.keys()): state = {k.replace('module.', ''): v for k, v in state.items()} v16_model.load_state_dict(state) v16_model = v16_model.to(device).eval() print(f"V16: epoch {ckpt.get('epoch', '?')}, SNR {ckpt.get('snr', 0):.2f} dB") # V18 refiner v18_refiner = V18RefinerNet(pretrained=False) ckpt = torch.load(v18_ckpt, map_location=device, weights_only=False) v18_refiner.load_state_dict(ckpt['refiner']) v18_refiner = v18_refiner.to(device).eval() print(f"V18: epoch {ckpt.get('epoch', '?')}, SNR {ckpt.get('snr', 0):.2f} dB") # V19 v19_model = PerLeadNetV19(pretrained=False) ckpt = torch.load(v19_ckpt, map_location=device, weights_only=False) state = ckpt['model'] if any(k.startswith('module.') for k in state.keys()): state = {k.replace('module.', ''): v for k, v in state.items()} v19_model.load_state_dict(state) v19_model = v19_model.to(device).eval() print(f"V19: epoch {ckpt.get('epoch', '?')}, SNR {ckpt.get('snr', ckpt.get('best_snr', 0)):.2f} dB") # Net3 net3_model = Net3(pretrained=False) ckpt = torch.load(net3_ckpt, map_location=device, weights_only=False) if isinstance(ckpt, dict) and 'state_dict' in ckpt: ckpt = ckpt['state_dict'] net3_model.load_state_dict(ckpt) net3_model = net3_model.to(device).eval() print(f"Net3: loaded from {net3_ckpt}") return v16_model, v18_refiner, v19_model, net3_model @torch.no_grad() def run_inference(v16_model, v18_refiner, v19_model, net3_model, samples, device): """Run inference once for all samples, return raw mV predictions for all models.""" all_predictions = [] resize = torch.nn.functional.interpolate for img_path, csv_path in tqdm(samples, desc="Running inference"): 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) row_crops = [crop_row(image, i) for i in range(4)] # V18 prediction (V16 + refiner) row_tensors = [torch.from_numpy(c.astype(np.float32) / 255.0).permute(2, 0, 1) for c in row_crops] images = torch.stack(row_tensors, dim=0).unsqueeze(0).to(device) with torch.amp.autocast('cuda', dtype=torch.float16): v16_preds = torch.stack([v16_model(images[:, i]) for i in range(4)], dim=1) refined, _ = v18_refiner(images, v16_preds) v18_mv = {} refined_np = refined[0].cpu().numpy() 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 v18_mv[row_idx] = (ZERO_MV[row_idx] - pred_y_full) / MV_TO_PIXEL # V19 prediction v19_mv = {} for row_idx in range(4): row_tensor = torch.from_numpy(row_crops[row_idx].astype(np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0).to(device) with torch.amp.autocast('cuda'): output = v19_model(row_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 v19_mv[row_idx] = (ZERO_MV[row_idx] - pred_y_full) / MV_TO_PIXEL # Net3 prediction net3_input = torch.from_numpy(image.astype(np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0).to(device) with torch.amp.autocast('cuda', dtype=torch.float16): net3_output = net3_model(net3_input) pixel = torch.sigmoid(net3_output).float().cpu().numpy()[0] # [4, H, W] # Convert Net3 pixel to mV net3_mv = {} for row_idx in range(4): row_pixel = pixel[row_idx, :, T0:T1] # [H, W] # Find peak per column pred_y = row_pixel.argmax(axis=0).astype(np.float32) # Resample to OUTPUT_WIDTH x_old = np.linspace(0, 1, len(pred_y)) x_new = np.linspace(0, 1, OUTPUT_WIDTH) pred_y_resampled = np.interp(x_new, x_old, pred_y) # Scale to full image coordinates scale_h = TARGET_HEIGHT / pixel.shape[1] pred_y_full = pred_y_resampled * scale_h net3_mv[row_idx] = (ZERO_MV[row_idx] - pred_y_full) / MV_TO_PIXEL all_predictions.append({ 'img_path': str(img_path), 'csv_path': str(csv_path), 'v18_mv': v18_mv, 'v19_mv': v19_mv, 'net3_mv': net3_mv, 'df': df, }) return all_predictions def evaluate_ratio(predictions, w_v18, w_v19, w_net3, apply_smoothing=True, apply_einthoven=True): """Evaluate a specific ensemble ratio.""" snrs = [] for pred in predictions: # Blend predictions ensemble = {} for row_idx in range(4): ensemble[row_idx] = ( w_v18 * pred['v18_mv'][row_idx] + w_v19 * pred['v19_mv'][row_idx] + w_net3 * pred['net3_mv'][row_idx] ) # Post-processing for row_idx in range(4): if apply_smoothing: ensemble[row_idx] = apply_savgol_smoothing(ensemble[row_idx], window=7, polyorder=2) ensemble[row_idx] = np.clip(ensemble[row_idx], ECG_MV_MIN, ECG_MV_MAX) if apply_einthoven: ensemble = apply_einthoven_correction(ensemble, alpha=0.33) # Compute SNR snr = compute_snr(ensemble, pred['df']) if snr > 0: snrs.append(snr) return np.mean(snrs) if snrs else 0.0 def main(): parser = argparse.ArgumentParser(description='Find best V18+V19+Net3 ensemble ratio') parser.add_argument('--v16_checkpoint', type=str, default=f'{LOCAL_CHECKPOINT_DIR}/v16_perlead_epoch020.pth') parser.add_argument('--v18_checkpoint', type=str, default=f'{LOCAL_CHECKPOINT_DIR}/v18_refiner_best.pth') parser.add_argument('--v19_checkpoint', type=str, default=f'{LOCAL_CHECKPOINT_DIR}/v19_enhanced_epoch010.pth') parser.add_argument('--net3_checkpoint', type=str, default=f'{LOCAL_CHECKPOINT_DIR}/net3_kaggle/iter_0004200.pt') parser.add_argument('--kaggle_data', type=str, default='/data/ecg-digitization/stage1_data/train') parser.add_argument('--num_samples', type=int, default=1500) parser.add_argument('--cache_file', type=str, default='/tmp/triple_ensemble_predictions.pkl') 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') print(f"{'='*70}") print(f"Finding Best V18 + V19 + Net3 Ensemble Ratio") print(f"{'='*70}") # Check if we have cached predictions cache_path = Path(args.cache_file) if cache_path.exists(): print(f"Loading cached predictions from {cache_path}...") with open(cache_path, 'rb') as f: predictions = pickle.load(f) print(f"Loaded {len(predictions)} cached predictions") else: # Load models print("\nLoading models...") v16_model, v18_refiner, v19_model, net3_model = load_models( args.v16_checkpoint, args.v18_checkpoint, args.v19_checkpoint, args.net3_checkpoint, device ) # Find samples kaggle_dir = Path(args.kaggle_data) all_samples = [] for sample_dir in kaggle_dir.iterdir(): 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: all_samples.append((img_path, csv_path)) print(f"Found {len(all_samples)} valid images") if len(all_samples) > args.num_samples: samples = random.sample(all_samples, args.num_samples) else: samples = all_samples print(f"Running inference on {len(samples)} samples...") predictions = run_inference(v16_model, v18_refiner, v19_model, net3_model, samples, device) # Cache predictions print(f"Caching predictions to {cache_path}...") with open(cache_path, 'wb') as f: pickle.dump(predictions, f) # ========================================================================= # Test all combinations # ========================================================================= print(f"\n{'='*70}") print(f"Testing All Model Combinations...") print(f"{'='*70}") # Individual models print("\n--- Individual Models ---") results = [] for name, (w18, w19, wn3) in [ ("V18 only", (1.0, 0.0, 0.0)), ("V19 only", (0.0, 1.0, 0.0)), ("Net3 only", (0.0, 0.0, 1.0)), ]: snr = evaluate_ratio(predictions, w18, w19, wn3) results.append((name, w18, w19, wn3, snr)) print(f" {name:15s}: {snr:>6.2f} dB") # Two-model ensembles print("\n--- Two-Model Ensembles ---") two_model_ratios = [0.0, 0.2, 0.4, 0.5, 0.6, 0.8, 1.0] print("\nV18 + V19:") for r in two_model_ratios: snr = evaluate_ratio(predictions, r, 1-r, 0.0) results.append((f"V18={r:.0%}+V19={1-r:.0%}", r, 1-r, 0.0, snr)) print(f" V18={r*100:>3.0f}%, V19={(1-r)*100:>3.0f}%: {snr:>6.2f} dB") print("\nV18 + Net3:") for r in two_model_ratios: snr = evaluate_ratio(predictions, r, 0.0, 1-r) results.append((f"V18={r:.0%}+Net3={1-r:.0%}", r, 0.0, 1-r, snr)) print(f" V18={r*100:>3.0f}%, Net3={(1-r)*100:>3.0f}%: {snr:>6.2f} dB") print("\nV19 + Net3:") for r in two_model_ratios: snr = evaluate_ratio(predictions, 0.0, r, 1-r) results.append((f"V19={r:.0%}+Net3={1-r:.0%}", 0.0, r, 1-r, snr)) print(f" V19={r*100:>3.0f}%, Net3={(1-r)*100:>3.0f}%: {snr:>6.2f} dB") # Three-model ensembles (grid search) print("\n--- Three-Model Ensembles ---") print(f"{'V18':>5} {'V19':>5} {'Net3':>5} {'SNR':>8}") print("-" * 30) three_model_grid = [0.0, 0.2, 0.4, 0.6, 0.8, 1.0] for w18 in three_model_grid: for w19 in three_model_grid: for wn3 in three_model_grid: if abs(w18 + w19 + wn3 - 1.0) < 0.01: # Must sum to 1 snr = evaluate_ratio(predictions, w18, w19, wn3) results.append((f"V18={w18:.0%}+V19={w19:.0%}+Net3={wn3:.0%}", w18, w19, wn3, snr)) print(f"{w18*100:>4.0f}% {w19*100:>4.0f}% {wn3*100:>4.0f}% {snr:>7.2f} dB") # Find best overall best = max(results, key=lambda x: x[4]) print(f"\n{'='*70}") print(f"BEST OVERALL: {best[0]}") print(f" V18={best[1]*100:.0f}%, V19={best[2]*100:.0f}%, Net3={best[3]*100:.0f}%") print(f" SNR = {best[4]:.2f} dB") print(f"{'='*70}") # Fine-tune around best print(f"\nFine-tuning around best...") fine_results = [] for d18 in [-0.1, -0.05, 0.0, 0.05, 0.1]: for d19 in [-0.1, -0.05, 0.0, 0.05, 0.1]: w18 = max(0, min(1, best[1] + d18)) w19 = max(0, min(1, best[2] + d19)) wn3 = max(0, 1.0 - w18 - w19) if wn3 >= 0 and abs(w18 + w19 + wn3 - 1.0) < 0.01: snr = evaluate_ratio(predictions, w18, w19, wn3) fine_results.append((w18, w19, wn3, snr)) if fine_results: final_best = max(fine_results, key=lambda x: x[3]) print(f"\nFINAL BEST: V18={final_best[0]*100:.0f}%, V19={final_best[1]*100:.0f}%, Net3={final_best[2]*100:.0f}%") print(f" SNR = {final_best[3]:.2f} dB") # Summary print(f"\n{'='*70}") print("TOP 10 CONFIGURATIONS:") print(f"{'='*70}") sorted_results = sorted(results, key=lambda x: x[4], reverse=True)[:10] for i, (name, w18, w19, wn3, snr) in enumerate(sorted_results, 1): print(f" {i:2d}. {name:35s} → {snr:.2f} dB") if __name__ == '__main__': main()