| |
| """ |
| V16 Inference with Visualization |
| |
| Runs V16 per-lead model on Kaggle images and plots predictions as dots on the strips. |
| V16 uses signal-region-only input (T0:T1) and 500px crop height. |
| """ |
|
|
| import os |
| import sys |
| import argparse |
| import random |
| import numpy as np |
| import pandas as pd |
| from pathlib import Path |
| from tqdm import tqdm |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
| import cv2 |
| import timm |
|
|
|
|
| |
| |
| |
| 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 |
|
|
| VALID_VARIANTS = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012'] |
|
|
|
|
| |
| |
| |
| 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 PerLeadNet(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) |
| out = torch.sigmoid(out) |
| |
| return out.squeeze(1) |
|
|
|
|
| |
| |
| |
| def load_model(checkpoint_path, device): |
| """Load trained V15 model.""" |
| model = PerLeadNet(encoder_name='convnext_base.fb_in22k_ft_in1k', pretrained=False) |
| |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) |
| model.load_state_dict(checkpoint['model']) |
| model = model.to(device) |
| model.eval() |
| |
| print(f"Loaded checkpoint from epoch {checkpoint['epoch']}") |
| print(f" SNR: {checkpoint.get('snr', 'N/A'):.2f} dB") |
| print(f" MAE: {checkpoint.get('mae', 'N/A'):.2f} px") |
| |
| return model |
|
|
|
|
| def crop_row(image, row_idx): |
| """Crop a single row centered on its baseline, signal region only (T0:T1).""" |
| 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 predict_row(model, row_crop, device): |
| """Run inference on a single row crop.""" |
| |
| 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.cuda.amp.autocast(): |
| output = model(image_tensor) |
| |
| |
| pred_y_crop = output[0].cpu().numpy() * ROW_HEIGHT |
| |
| return pred_y_crop |
|
|
|
|
| def convert_crop_to_full(pred_y_crop, row_idx): |
| """ |
| Convert crop-relative y-coordinates to full image coordinates. |
| V16: Output is OUTPUT_WIDTH (3926) values, maps to x positions T0:T1. |
| """ |
| 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 |
| |
| return pred_y_full |
|
|
|
|
| def visualize_predictions(image, predictions, output_path): |
| """ |
| Draw predictions as dots on the image. |
| |
| predictions: list of 4 arrays, one per row, each [OUTPUT_WIDTH] in full image y-coords |
| V16: predictions are already OUTPUT_WIDTH (3926) values for x positions 0 to OUTPUT_WIDTH-1. |
| Output is cropped to signal region. |
| Plots EVERY predicted point. |
| """ |
| |
| vis_image = image[:, T0:T1, :].copy() |
| |
| |
| colors = [ |
| (255, 0, 0), |
| (0, 255, 0), |
| (255, 0, 255), |
| (0, 165, 255), |
| ] |
| |
| for row_idx, pred_y in enumerate(predictions): |
| color = colors[row_idx] |
| |
| |
| for x in range(len(pred_y)): |
| y = int(np.clip(pred_y[x], 0, TARGET_HEIGHT - 1)) |
| |
| cv2.circle(vis_image, (x, y), 1, color, -1) |
| |
| cv2.imwrite(str(output_path), vis_image) |
|
|
|
|
| def process_image(model, image_path, output_dir, device): |
| """Process a single image and save visualization.""" |
| |
| image = cv2.imread(str(image_path), cv2.IMREAD_COLOR) |
| if image is None: |
| print(f"Failed to load: {image_path}") |
| return |
| |
| |
| image = image[Y0:Y1, X0:X1] |
| image = cv2.resize(image, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR) |
| |
| |
| predictions = [] |
| for row_idx in range(4): |
| row_crop = crop_row(image, row_idx) |
| pred_y_crop = predict_row(model, row_crop, device) |
| pred_y_full = convert_crop_to_full(pred_y_crop, row_idx) |
| predictions.append(pred_y_full) |
| |
| |
| 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, predictions, output_path) |
| |
| return predictions |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument('--checkpoint', type=str, |
| default='/data/ecg-digitization/checkpoints/v16_perlead_latest.pth') |
| 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/v16')) |
| parser.add_argument('--num_samples', type=int, default=50) |
| parser.add_argument('--seed', type=int, default=42) |
| args = parser.parse_args() |
| |
| |
| random.seed(args.seed) |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| output_dir = Path(args.output_dir) |
| |
| |
| if output_dir.exists(): |
| import shutil |
| shutil.rmtree(output_dir) |
| print(f"Deleted old predictions in {output_dir}") |
| output_dir.mkdir(parents=True, exist_ok=True) |
| |
| print(f"{'='*70}") |
| print(f"V16 Inference Visualization") |
| print(f"{'='*70}") |
| print(f"Checkpoint: {args.checkpoint}") |
| print(f"Output: {output_dir}") |
| print(f"Device: {device}") |
| print(f"{'='*70}") |
| |
| |
| model = load_model(args.checkpoint, device) |
| |
| |
| kaggle_dir = Path(args.kaggle_data) |
| all_images = [] |
| |
| for sample_dir in kaggle_dir.iterdir(): |
| if not sample_dir.is_dir(): |
| continue |
| 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_images.append(img_path) |
| |
| print(f"Found {len(all_images)} valid images") |
| |
| |
| if len(all_images) > args.num_samples: |
| selected = random.sample(all_images, args.num_samples) |
| else: |
| selected = all_images |
| |
| print(f"Processing {len(selected)} images...") |
| |
| |
| for img_path in tqdm(selected): |
| try: |
| process_image(model, img_path, output_dir, device) |
| except Exception as e: |
| print(f"Error processing {img_path}: {e}") |
| |
| print(f"\n{'='*70}") |
| print(f"Done! Saved {len(selected)} visualizations to {output_dir}") |
| print(f"{'='*70}") |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|