knee-landmarks / landmark_predictor.py
Pakawat Nakwijit
calculate landmarks
8fcc1ae
Raw
History Blame Contribute Delete
8.79 kB
import os
import gc
import cv2
import matplotlib.cm as cm
import numpy as np
import torch
import segmentation_models_pytorch as smp
from huggingface_hub import hf_hub_download
from landmark_geometry import plot_landmarks
# Hugging Face space model details
HF_TOKEN = os.environ.get("HF_TOKEN")
REPO_ID = "imtk/knee-landmarks"
# Each fine-grained landmark has its own heatmap-regression model
LANDMARK_MODELS = {
"AL": "knee-heatmap-regression/MITB05-HEATMAP-AL-1779305935",
"LTL": "knee-heatmap-regression/MITB05-HEATMAP-LTL-1779297640",
"FHL": "knee-heatmap-regression/MITB05-HEATMAP-FHL-1779259956",
"MFL": "knee-heatmap-regression/MITB05-HEATMAP-MFL-1779293571",
"FL": "knee-heatmap-regression/MITB05-HEATMAP-FL-1779263939",
"MTL": "knee-heatmap-regression/MITB05-HEATMAP-MTL-1779301793",
"LFL": "knee-heatmap-regression/MITB05-HEATMAP-LFL-1779289548",
"FHR": "knee-heatmap-regression/MITB05-HEATMAP-FHR-1779258660",
"MFR": "knee-heatmap-regression/MITB05-HEATMAP-MFR-1779289493",
"AR": "knee-heatmap-regression/MITB05-HEATMAP-AR-1779301913",
"LTR": "knee-heatmap-regression/MITB05-HEATMAP-LTR-1779293588",
"FR": "knee-heatmap-regression/MITB05-HEATMAP-FR-1779262758",
"MTR": "knee-heatmap-regression/MITB05-HEATMAP-MTR-1779297757",
"LFR": "knee-heatmap-regression/MITB05-HEATMAP-LFR-1779266866",
}
# Each fine-grained landmark is cropped from its parent ROI region
RELATED_REGION = {
"FHR": "FHR",
"FHL": "FHL",
"FR": "DFR",
"LFR": "DFR",
"MFR": "DFR",
"FL": "DFL",
"LFL": "DFL",
"MFL": "DFL",
"LTR": "PTR",
"MTR": "PTR",
"LTL": "PTL",
"MTL": "PTL",
"AR": "TR",
"AL": "TL",
}
CROP_SIZE = 256
# Std-dev (in crop pixels) used to spread the raw softargmax probability map
# into a visible blob for heatmap visualization.
HEATMAP_SPREAD_SIGMA = 8.0
def _get_device():
if torch.cuda.is_available():
return torch.device("cuda")
return torch.device("cpu")
def _spatial_softargmax2d(logits):
"""logits: [B, 1, H, W] -> (xy: [B, 2], prob: [B, 1, H, W])"""
b, _, h, w = logits.shape
prob = torch.softmax(logits.view(b, -1), dim=-1).view(b, 1, h, w)
xs = torch.linspace(0, w - 1, w, device=logits.device)
ys = torch.linspace(0, h - 1, h, device=logits.device)
prob_2d = prob[:, 0, :, :]
expected_x = (prob_2d.sum(dim=1) * xs).sum(dim=1)
expected_y = (prob_2d.sum(dim=2) * ys).sum(dim=1)
return torch.stack([expected_x, expected_y], dim=1), prob
def _spread_heatmap(heatmap, sigma=HEATMAP_SPREAD_SIGMA):
"""Blur a sharply-peaked probability heatmap into a visible blob (for
display only) and renormalize so it still sums to 1."""
blurred = cv2.GaussianBlur(heatmap, ksize=(0, 0), sigmaX=sigma)
total = blurred.sum()
if total > 0:
blurred = blurred / total
return blurred
def _crop_fixed_roi(image, minx, miny, maxx, maxy, size=CROP_SIZE):
"""Crop a fixed size*size window centered on the box, padding with
black if it falls outside the image, matching how the training
dataset was generated."""
h, w = image.shape[:2]
cx = (minx + maxx) / 2.0
cy = (miny + maxy) / 2.0
crop_x1 = int(round(cx - size / 2))
crop_y1 = int(round(cy - size / 2))
crop_x2 = crop_x1 + size
crop_y2 = crop_y1 + size
pad_left = max(0, -crop_x1)
pad_top = max(0, -crop_y1)
pad_right = max(0, crop_x2 - w)
pad_bottom = max(0, crop_y2 - h)
src_x1 = max(0, crop_x1)
src_y1 = max(0, crop_y1)
src_x2 = min(w, crop_x2)
src_y2 = min(h, crop_y2)
crop = image[src_y1:src_y2, src_x1:src_x2]
if pad_left or pad_top or pad_right or pad_bottom:
crop = cv2.copyMakeBorder(
crop, pad_top, pad_bottom, pad_left, pad_right,
borderType=cv2.BORDER_CONSTANT, value=(0, 0, 0),
)
src = (src_x1, src_y1, src_x2, src_y2)
pad = (pad_left, pad_top, pad_right, pad_bottom)
return crop, src, pad
def _load_landmark_model(model_name):
config_path = hf_hub_download(
repo_id=REPO_ID,
filename=f"{model_name}/config.json",
token=HF_TOKEN,
)
hf_hub_download(
repo_id=REPO_ID,
filename=f"{model_name}/model.safetensors",
token=HF_TOKEN,
)
model = smp.from_pretrained(os.path.dirname(config_path))
model.eval()
return model
def _clear_model(model, device):
del model
gc.collect()
if device.type == "cuda":
torch.cuda.empty_cache()
elif device.type == "mps":
torch.mps.empty_cache()
def predict_landmark(image, rois):
if image is None or not rois:
return []
device = _get_device()
landmarks = []
for lm_name, model_name in LANDMARK_MODELS.items():
region = RELATED_REGION[lm_name]
if region not in rois:
print(f"Landmark Not Found: {lm_name} (region {region} not detected)")
continue
roi = rois[region]
crop, src, pad = _crop_fixed_roi(
image, roi["minx"], roi["miny"], roi["maxx"], roi["maxy"]
)
src_x1, src_y1, src_x2, src_y2 = src
pad_left, pad_top, pad_right, pad_bottom = pad
model = _load_landmark_model(model_name)
model.to(device)
input_tensor = crop.astype(np.float32) / 255.0
input_tensor = np.transpose(input_tensor, (2, 0, 1))
# image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
input_tensor = torch.from_numpy(input_tensor).unsqueeze(0).to(device)
with torch.inference_mode():
logits = model(input_tensor)
pred_xy, prob = _spatial_softargmax2d(logits)
x, y = pred_xy[0].cpu().numpy()
confidence = float(prob[0, 0].max().cpu())
heatmap = _spread_heatmap(prob[0, 0].cpu().numpy())
_clear_model(model, device)
landmarks.append({
"name": lm_name,
"region": region,
"orig_x": float(x) - pad_left + src_x1,
"orig_y": float(y) - pad_top + src_y1,
"crop_x": float(x),
"crop_y": float(y),
"crop_minx": src_x1 - pad_left,
"crop_miny": src_y1 - pad_top,
"heatmap": heatmap,
"confidence": confidence,
})
return landmarks
def _normalize_heatmap(heatmap):
return heatmap / (heatmap.max() + 1e-8)
def _overlay_heatmap(image_rgb, mask_hw):
"""Blend a plasma-colored heatmap mask (H, W) in [0, 1] over an RGB image."""
image_hwc = image_rgb.astype(np.float32) / 255.0
heatmap_rgba = cm.plasma(mask_hw)
heatmap_rgb = heatmap_rgba[..., :3].astype(np.float32)
alpha = 0.8 * mask_hw[..., None]
overlay_image = (1 - alpha) * image_hwc + alpha * heatmap_rgb
return (np.clip(overlay_image, 0.0, 1.0) * 255).astype(np.uint8)
def _paste_max(canvas_hw, mask_hw, minx, miny):
"""Composite mask_hw onto canvas_hw at (minx, miny), clipped to bounds,
keeping the max value where masks from different landmarks overlap."""
h, w = canvas_hw.shape[:2]
size = mask_hw.shape[0]
dst_x1, dst_y1 = max(0, minx), max(0, miny)
dst_x2, dst_y2 = min(w, minx + size), min(h, miny + size)
if dst_x2 <= dst_x1 or dst_y2 <= dst_y1:
return
src_x1, src_y1 = dst_x1 - minx, dst_y1 - miny
src_x2, src_y2 = src_x1 + (dst_x2 - dst_x1), src_y1 + (dst_y2 - dst_y1)
region = canvas_hw[dst_y1:dst_y2, dst_x1:dst_x2]
canvas_hw[dst_y1:dst_y2, dst_x1:dst_x2] = np.maximum(
region, mask_hw[src_y1:src_y2, src_x1:src_x2]
)
def _build_heatmap_image(image, landmarks):
"""Overlay every landmark's heatmap, in its original-image position, on
the full input image."""
h, w = image.shape[:2]
canvas = np.zeros((h, w), dtype=np.float32)
for lm in landmarks:
mask = _normalize_heatmap(lm["heatmap"])
_paste_max(canvas, mask, int(round(lm["crop_minx"])), int(round(lm["crop_miny"])))
return _overlay_heatmap(image, canvas)
def _build_heatmap_gallery(image, rois, landmarks):
"""Overlay each landmark's heatmap on its own ROI crop."""
gallery = []
for lm in landmarks:
roi = rois.get(lm["region"])
if roi is None:
continue
crop, _, _ = _crop_fixed_roi(image, roi["minx"], roi["miny"], roi["maxx"], roi["maxy"])
mask = _normalize_heatmap(lm["heatmap"])
gallery.append((_overlay_heatmap(crop, mask), lm["name"]))
return gallery
def show_landmarks(image, rois, landmarks):
heatmap_image = _build_heatmap_image(image, landmarks)
heatmap_gallery = _build_heatmap_gallery(image, rois, landmarks)
landmark_image = plot_landmarks(np.ascontiguousarray(image.copy()), landmarks)
return heatmap_image, heatmap_gallery, landmark_image