Segmentation-Guided Coordinate Regression (SGR)

This repository provides PyTorch model weights for segmentation-guided coordinate regression (SGR) applied to anatomical landmark detection in full lower-limb X-rays. Refer to the associated paper for full methodological details and achieved performance. The provided weights are intended to be used in conjunction with the repository for automatic lower-limb malalignment assessment.

If you use any of the provided models, please cite:

@article{sanchez2024segmentation,
  title={Segmentation-guided coordinate regression for robust landmark detection on X-rays: application to automated assessment of lower limb alignment},
  author={Sanchez, Sebastian Amador and Van Overschelde, Philippe and Vandemeulebroucke, Jef},
  journal={IEEE Access},
  volume={12},
  pages={61484--61497},
  year={2024},
  publisher={IEEE}
}

Architectures

The repository includes the following model architectures:

  • SGRNetwork16: Segmentation-guided coordinate regression architecture. The segmentation branch is based on a U-Net architecture, whose encoder topology is equivalent to the coordinate regression branch, where both follow a VGG-16 backbone.
  • UNet16: U-Net architecture with a VGG-16 backbone on its encoder used to initialize SGRNetwork16.
  • UNet16V2: A variant of U-Net specifically designed for femoral diaphysis segmentation.

Repository Structure

  • BaseModels.py β€” Definitions of U-Net architectures
  • SGRModel.py β€” Implementation of the SGR coordinate regression model
  • Weights/ β€” Pretrained PyTorch model weights

Available Models

Region # Landmarks # Coordinates # Regions Weights Path
Hip 1 2 β€” Weights/Femur/...
Femoral diaphysis β€” β€” β€” Weights/Diaphysis/...
Knee 5 10 β€” Weights/Knee/...
Ankle 2 4 β€” Weights/Ankle/...
ROI detection β€” β€” 5 (including background) Weights/ROI_Detection/...

Requirements

Developed and tested with Python 3.9.12. The following dependencies are required:

torch==2.2.2
torchvision==0.17.2
huggingface_hub==1.8.0

Usage Examples

A. Pretrained SGR Models

I. Expected input and outputs:

  • Input: Cropped and resized (512 x 512) regions of full lower-limb X-ray images.
  • Output: 2D landmark coordinates (x, y) in image pixel space.
  • The number of output coordinates depends on the anatomical region:
    • Hip: 1 landmark (2 values)
    • Knee: 5 landmarks (10 values)
    • Ankle: 2 landmarks (4 values)

II. Typical inference pipeline involves:

  1. ROI detection (hip, femoral diaphysis, knee, ankle)
  2. Cropping of detected regions
  3. Landmark detection using SGR models
  4. Downstream alignment metric computation

III. Load of the weights and inference:

import torch
from BaseModels import UNet16
from SGRModel import SGRNetwork16

device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

# Paths to pretrained model weights
path_to_hip_weights = 'Weights/Femur/Best_Run_ep100_bs8_lr1e-05_a0.2.pth'
path_to_knee_weights = 'Weights/Knee/Best_Run_ep100_bs8_lr1e-05_a0.2.pth'
path_to_ankle_weights = 'Weights/Ankle/Best_Run_ep100_bs8_lr1e-05_a0.2.pth'

# Hip model
hip_unet = UNet16()
hip_sgr = SGRNetwork16(unet16=hip_unet, num_coordinates=2)
hip_sgr.load_state_dict(torch.load(path_to_hip_weights, map_location=device))

torch_hip_images = torch.rand(1, 3, 512, 512)  # Format: batch, channel, height, width
hip_sgr.eval()  # Estimate the landmarks
with torch.no_grad():
    _, hip_landmarks = hip_sgr(torch_hip_images)  # Output landmarks: batch, num_landmarks, 2 (x-, y-coordinates)

# Knee model
knee_unet = UNet16()
knee_unet.final = torch.nn.Conv2d(32, 5, kernel_size=1)  # 5 landmark segmentations
knee_sgr = SGRNetwork16(unet16=knee_unet, num_coordinates=10)
knee_sgr.load_state_dict(torch.load(path_to_knee_weights, map_location=device))

torch_knee_images = torch.rand(1, 3, 512, 512)  # Format: batch, channel, height, width
knee_sgr.eval()  # Estimate the landmarks
with torch.no_grad():
    _, knee_landmarks = knee_sgr(torch_knee_images)  # Output landmarks: batch, num_landmarks, 2 (x-, y-coordinates)

# Ankle model
ankle_unet = UNet16()
ankle_unet.final = torch.nn.Conv2d(32, 2, kernel_size=1)  # 2 landmark segmentations
ankle_sgr = SGRNetwork16(unet16=ankle_unet, num_coordinates=4)
ankle_sgr.load_state_dict(torch.load(path_to_ankle_weights, map_location=device))

torch_ankle_images = torch.rand(1, 3, 512, 512)  # Format: batch, channel, height, width
ankle_sgr.eval()  # Estimate the landmarks
with torch.no_grad():
    _, ankle_landmarks = ankle_sgr(torch_ankle_images)  # Output landmarks: batch, num_landmarks, 2 (x-, y-coordinates)

Note: Input images should be preprocessed consistently with the training setup (e.g., normalization, resizing). Refer to the main repository for details.


B. Initialize a new SGR Model

import torch
from BaseModels import UNet16
from SGRModel import SGRNetwork16

device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

# Example setup:
# num_classes: number of segmentation classes (landmarks)
# num_coordinates: 2 Γ— number of landmarks (x, y per landmark)
# net_pretrained: whether to use ImageNet initialization

unet = UNet16(num_classes=5)
sgr_model = SGRNetwork16(
    unet16=unet,
    net_pretrained=True,
    num_coordinates=10
).to(device)

C. ROI (Joint) Detection

import torch
from Detection import get_model_instance_segmentation

device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

path_to_weights = 'Weights/ROI_Detection/fll_detection_joints_cv_5_31.pth'

# Classes: background, hip, femoral diaphysis, knee, ankle
num_classes = 5

detection_model = get_model_instance_segmentation(num_classes)
detection_model.load_state_dict(torch.load(path_to_weights, map_location=device))
detection_model.to(device)

D. Femoral Diaphysis Segmentation

import torch
from BaseModels import UNet16V2

device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')

path_to_weights = 'Weights/Diaphysis/Best_DiceCELoss_lr1e5.pth'

unet = UNet16V2()
unet.load_state_dict(torch.load(path_to_weights, map_location=device))

Limitations

The provided models are intended for research purposes only and have not been validated for clinical use; hence, they must not be used for diagnostic or medical decision-making.


License

The model weights in this repository are licensed under the Creative Commons Attribution Non Commercial Share Alike 4.0 License. See the LICENSE file for the full legal text.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support