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 initializeSGRNetwork16.UNet16V2:A variant of U-Net specifically designed for femoral diaphysis segmentation.
Repository Structure
BaseModels.pyβ Definitions of U-Net architecturesSGRModel.pyβ Implementation of the SGR coordinate regression modelWeights/β 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:
- ROI detection (hip, femoral diaphysis, knee, ankle)
- Cropping of detected regions
- Landmark detection using SGR models
- 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.