Files
LORKT_AI_Service/service/predictors/shoulder.py
T
2025-12-03 16:50:06 +05:00

495 lines
17 KiB
Python

import warnings
from typing import Optional, NamedTuple
import cv2
import numpy as np
import torch
from torch import Tensor
import torchvision.ops as ops
from torchvision.transforms import v2 as T
from skimage.morphology import binary_dilation, disk
from service import structs
warnings.filterwarnings('ignore', category=UserWarning)
device = torch.device('cuda')
models_root = "service/models/shoulder"
model_frac = torch.load(f'{models_root}/frac_model.pth',
weights_only=False).to(device)
model_lr = torch.load(f'{models_root}/lr_model.pth',
weights_only=False).to(device)
model_parts = torch.load(f'{models_root}/parts_model.pth',
weights_only=False).to(device)
model_move = torch.load(f'{models_root}/move_model.pth',
weights_only=False).to(device)
model_frac.eval()
model_lr.eval()
model_parts.eval()
model_move.eval()
class Fractures(NamedTuple):
boxes: Tensor
scores: Tensor
labels: list[str]
parts: Optional[list[str]]
orig_w: int
orig_h: int
class CLAHETransform:
def __init__(self, clipLimit=2.0, tileGridSize=(8, 8)):
self.clipLimit = clipLimit
self.tileGridSize = tileGridSize
self.clahe = cv2.createCLAHE(clipLimit=self.clipLimit,
tileGridSize=self.tileGridSize)
def __call__(self, img, target=None):
img_np = img.cpu().numpy()
if img_np.ndim == 3 and img_np.shape[0] == 1:
img_np = img_np[0]
cl_img = self.clahe.apply(img_np)
cl_img_tensor = torch.from_numpy(cl_img).unsqueeze(0)
cl_img_tensor = cl_img_tensor.to(img.device)
return cl_img_tensor
transform_parts = T.Compose([
T.Resize((256, 256)),
T.Grayscale(num_output_channels=1),
CLAHETransform(clipLimit=2.0, tileGridSize=(8, 8)),
T.ToDtype(torch.float, scale=True),
T.ToTensor(),
T.Normalize(mean=[0.0773] * 3, std=[0.0516] * 3)
])
transform_frac = T.Compose([
T.Resize((512, 512)),
CLAHETransform(clipLimit=2.0, tileGridSize=(8, 8)),
T.ToDtype(torch.float, scale=True),
T.ToTensor(),
T.Normalize(mean=[0.40526121854782104], std=[0.23242981731891632])
])
transform_lr = T.Compose([
T.Resize((256, 256)),
CLAHETransform(clipLimit=2.0, tileGridSize=(8, 8)),
T.ToDtype(torch.float, scale=True),
T.ToPureTensor(),
T.Normalize(mean=[0.0773] * 3, std=[0.0516] * 3)
])
transform_move = transform_test = T.Compose([
T.Resize((224, 224)),
T.Grayscale(num_output_channels=1),
CLAHETransform(clipLimit=2.0, tileGridSize=(8, 8)),
T.ToDtype(torch.float, scale=True),
T.ToTensor(),
T.Normalize(mean=[0.4172] * 3, std=[0.2612] * 3)
])
parts_map = {
0: 'акромиального отростка',
1: 'клювовидного отростка',
2: 'плечевой кости',
3: 'суставной впадины',
4: 'тела',
5: 'шейки',
6: 'Фон',
7: 'головки',
8: 'анатомической шейки',
9: 'хирургической шейки',
10: 'диафиза'
}
parts2bones = {
0: 'Лопатка',
1: 'Лопатка',
2: 'Плечевая кость',
3: 'Лопатка',
4: 'Лопатка',
5: 'Лопатка',
6: 'Фон',
7: 'Плечевая кость',
8: 'Плечевая кость',
9: 'Плечевая кость',
10: 'Плечевая кость'
}
def _convert_bboxes(bboxes: Tensor, orig_width: int, orig_height: int,
transformed_width=512,
transformed_height=512) -> Tensor:
"""Масштабирует координаты боксов обратно к
размерам исходного изображения."""
conv_bboxes = bboxes.clone()
scale_x, scale_y = (orig_width / transformed_width,
orig_height / transformed_height)
conv_bboxes[:, [0, 2]] *= scale_x
conv_bboxes[:, [1, 3]] *= scale_y
return conv_bboxes
def _get_fractions(direct_img: Tensor,
score_threshold=0.07) -> Fractures:
"""Получение переломов"""
original = direct_img.clone()
image = transform_frac(direct_img).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model_frac(image)[0]
scores = outputs['scores'].cpu()
valid = scores >= score_threshold
boxes = outputs['boxes'][valid].cpu()
scores = scores[valid]
keep = ops.nms(boxes, scores, 0.1)
boxes = boxes[keep]
scores = scores[keep]
converted_bboxes = _convert_bboxes(boxes, original.shape[2],
original.shape[1])
converted_bboxes = converted_bboxes[converted_bboxes[:, 0].argsort()]
fractures = Fractures(converted_bboxes, scores, [],
None, original.shape[2], original.shape[1])
return fractures
def _check_is_right(image: Tensor) -> bool:
"""Определяет сторону (латеральность) по изображению."""
image = transform_lr(image).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model_lr(image)[0]
return bool(torch.argmax(outputs).item())
def _smooth_segmentation_mask(mask: np.ndarray, kernel_size=5,
min_area=10) -> np.ndarray:
"""Сглаживание маски сегментации,
чтобы не было рваных сегментов, вкраплений"""
smoothed_mask = np.zeros_like(mask)
n_classes = int(mask.max()) + 1
for cls in range(n_classes):
class_mask = (mask == cls).astype(np.uint8)
closed = cv2.morphologyEx(class_mask, cv2.MORPH_CLOSE,
np.ones((kernel_size, kernel_size),
np.uint8))
opened = cv2.morphologyEx(closed, cv2.MORPH_OPEN,
np.ones((kernel_size, kernel_size),
np.uint8))
num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(opened,
connectivity=4)
clean = np.zeros_like(opened)
for i in range(1, num_labels):
if stats[i, cv2.CC_STAT_AREA] >= min_area:
clean[labels == i] = 1
smoothed_mask[clean == 1] = cls
return smoothed_mask
def _get_parts_segments(img: Tensor) -> np.ndarray:
img = img.repeat(3, 1, 1)
image = transform_parts(img).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model_parts(image)[0]
output = outputs.detach().cpu().numpy()
predicted_classes = np.argmax(output, axis=0)
if np.sum(predicted_classes == 6) > 64000:
raise structs.ImagesError("Снимок иной анатомической области или низкого диагностического качества")
predicted_classes = _smooth_segmentation_mask(predicted_classes)
return predicted_classes
def _compute_pca_direction(coords: np.ndarray) -> np.ndarray:
mean = coords.mean(axis=0)
centered = coords - mean
cov = np.cov(centered, rowvar=False)
eigvals, eigvecs = np.linalg.eigh(cov)
principal = eigvecs[:, np.argmax(eigvals)]
return principal / np.linalg.norm(principal)
def _rotate_vector(vec: np.ndarray, angle_deg: float) -> np.ndarray:
theta = np.deg2rad(angle_deg)
rot = np.array([[np.cos(theta), -np.sin(theta)],
[np.sin(theta), np.cos(theta)]])
return rot.dot(vec)
def _get_segment2_coords(mask: np.ndarray, class_idx=2) -> np.ndarray:
ys, xs = np.where(mask == class_idx)
return np.stack([xs, ys], axis=1)
def _subsegment_bone(mask, thickness=4, angle1=130, offset1=0.02, angle2=90,
offset2=0.25):
"""Разбиение плечевой кости на микро-сегменты"""
mask = mask.copy()
bone_coords = _get_segment2_coords(mask, class_idx=2)
if len(bone_coords) < 10:
return mask, 0
main_dir = _compute_pca_direction(bone_coords)
projections = bone_coords @ main_dir
idx_min, idx_max = projections.argmin(), projections.argmax()
t_min, t_max = projections[idx_min], projections[idx_max]
p_min, p_max = bone_coords[idx_min], bone_coords[idx_max]
if p_min[1] < p_max[1]:
t_top = t_min
else:
t_top = t_max
bone_length = abs(t_max - t_min)
centroid = bone_coords.mean(axis=0)
centroid_proj = centroid.dot(main_dir)
H, W = mask.shape
y_grid, x_grid = np.mgrid[0:H, 0:W]
pix = np.stack([x_grid.ravel(), y_grid.ravel()], axis=1)
seg2_mask_flat = (mask.ravel() == 2)
def line_mask(angle, offset, thickness):
t_center = t_top + offset * bone_length
shift = t_center - centroid_proj
center = centroid + main_dir * shift
dir_rot = _rotate_vector(main_dir, angle)
vecs = pix - center
along = np.dot(vecs, dir_rot)
ortho_vecs = vecs - np.outer(along, dir_rot)
dists = np.linalg.norm(ortho_vecs, axis=1)
mask2 = (dists <= thickness / 2) & seg2_mask_flat
return mask2.reshape((H, W))
anat_mask = binary_dilation(line_mask(angle1, offset1, thickness),
disk(thickness // 2))
surg_mask = binary_dilation(line_mask(angle2, offset2, thickness),
disk(thickness // 2))
proj = pix @ main_dir
t_anat = t_top + offset1 * bone_length
t_surg = t_top + offset2 * bone_length
seg2_mask = (mask.ravel() == 2)
diaf_mask = (proj > t_surg) & seg2_mask
diaf_mask = diaf_mask.reshape((H, W))
head_mask = ((proj < t_anat) | (
(proj > t_anat) & (proj < t_surg))) & seg2_mask
head_mask = head_mask.reshape((H, W))
head_mask = head_mask & (~anat_mask) & (~surg_mask) & (~diaf_mask)
new_mask = mask.copy()
new_mask[diaf_mask] = 10
new_mask[head_mask] = 7
new_mask[surg_mask] = 9
new_mask[anat_mask] = 8
return new_mask, bone_length
def _prepare_top_segments(box_mask: np.ndarray, weights: np.ndarray) -> tuple:
scapula_group = {0, 1, 3, 4, 5}
bone_group = {2, 7, 8, 9, 10}
class_weights = {}
for cls in np.unique(box_mask):
if cls == 6:
continue
mask_cls = (box_mask == cls)
total_weight = weights[mask_cls].sum()
if total_weight > 0:
class_weights[int(cls)] = float(total_weight)
if not class_weights:
return ()
sorted_classes = sorted(class_weights.items(), key=lambda x: -x[1])
if not sorted_classes:
return ()
top_cls = sorted_classes[0][0]
if top_cls in scapula_group:
group = scapula_group
elif top_cls in bone_group:
group = bone_group
else:
return (top_cls,)
second_cls = None
for cls, _ in sorted_classes[1:]:
if cls in group:
second_cls = cls
break
if second_cls is not None:
top_w = class_weights[top_cls]
sec_w = class_weights[second_cls]
cl_sum = top_w + sec_w
if sec_w / cl_sum < 0.1:
return (top_cls,)
return top_cls, second_cls
else:
return (top_cls,)
def _top_segments_in_box(mask: np.ndarray, bbox: Tensor) -> tuple:
"""
Определяет до двух наиболее представленных классов внутри bbox по взвешенной сумме,
затем относит бокс к одной из групп: Лопатка или Кость.
"""
bbox = bbox.cpu().numpy().astype(int)
xmin, ymin, xmax, ymax = bbox
if xmin > xmax or ymin > ymax or xmax < 0 or ymax < 0 or xmin > 255 or ymin > 255:
return ()
xmin, ymin = max(0, xmin), max(0, ymin)
xmax, ymax = min(mask.shape[1] - 1, xmax), min(mask.shape[0] - 1, ymax)
box_mask = mask[ymin:ymax + 1, xmin:xmax + 1]
h, w = box_mask.shape
ys, xs = np.mgrid[0:h, 0:w]
center_y, center_x = (h - 1) / 2, (w - 1) / 2
dists = np.sqrt((ys - center_y) ** 2 + (xs - center_x) ** 2)
max_dist = dists.max() if dists.max() > 0 else 1.0
weights = np.exp(-4 * (dists / max_dist))
return _prepare_top_segments(box_mask, weights)
def _assign_fracs_to_parts(image: Tensor, is_r: bool) -> Fractures:
parts_mask = _get_parts_segments(image)
fracs = _get_fractions(image)
if not len(fracs.boxes):
return fracs, 0
conv_boxes = _convert_bboxes(fracs.boxes, 256, 256,
fracs.orig_w, fracs.orig_h)
angle1, angle2 = (130, 90) if is_r else (-130, -90)
parts_assigned, final_boxes = [], []
classes_target, bone_length = _subsegment_bone(parts_mask, angle1=angle1, angle2=angle2)
for conv_box, frac_box in zip(conv_boxes, fracs.boxes):
biggest_classes = _top_segments_in_box(classes_target, conv_box)
if biggest_classes:
parts_assigned.append(biggest_classes)
final_boxes.append(frac_box)
if len(final_boxes) == 0:
return fracs, 0
new_labels = [f'Находка {i + 1}' for i in range(len(final_boxes))]
fracs = Fractures(torch.stack(final_boxes), fracs.scores, new_labels,
parts_assigned, fracs.orig_w, fracs.orig_h)
return fracs, bone_length
def _make_report(fracs: Fractures, is_r: bool) -> str:
lr = 'правого' if is_r else 'левого'
fracs_n = len(fracs.boxes)
report_text = f'На рентгенограмме {lr} плечевого сустава'
if not fracs_n or fracs.parts is None:
report_text += ' переломов не выявлено.'
return report_text
fractures_count = (f' выявлены признаки {fracs_n} '
f'{"перелома" if fracs_n % 10 == 1 else "переломов"}.')
report_text += fractures_count
for label, parts in zip(fracs.labels, fracs.parts):
parts_text = f' {label}: перелом '
parts_text += ', '.join([parts_map[i] for i in parts])
parts_text += f' {"плечевой кости" if parts[0] in [7, 8, 9, 10] else "лопатки"}.'
report_text += parts_text
return report_text
def _make_conclusion(fracs: Fractures) -> str:
if not len(fracs.boxes) or fracs.parts is None:
return 'Признаков перелома не выявлено.'
finds = {}
for parts in fracs.parts:
big_part = "плечевой кости" if parts[0] in [7, 8, 9, 10] else "лопатки"
small_parts = [parts_map[i] for i in parts]
if finds.get(big_part):
finds[big_part].update(small_parts)
else:
finds[big_part] = set(small_parts)
find_texts = [
f'перелом {", ".join(sorted(parts))} {bone}' for bone, parts in
finds.items()
]
conclusion = '; '.join(find_texts) + '.'
return conclusion.replace(' ', ' ').capitalize()
def _get_move_confidence(image):
image = image.repeat(3, 1, 1)
image = transform_move(image).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model_move(image)[0]
if torch.argmax(outputs).item() == 1:
return 0
move_prob = torch.softmax(outputs, 0)[0]
return float(move_prob.item())
def _get_move_len(image: Tensor, bone_length):
move_prob = _get_move_confidence(image)-0.5
coef = 0.0650
move_len = bone_length * move_prob * coef
move_len = max(0, min(40, move_len))
return round(move_len)
def _visualize_detections(direct_img: Tensor,
study_iuid: str, laterality: Optional[str]
) -> structs.Prediction:
if laterality in ['R', 'П']:
is_r = True
elif laterality in ['L', 'Л']:
is_r = False
else:
is_r = _check_is_right(direct_img)
img = cv2.cvtColor(direct_img[0].numpy(), cv2.COLOR_GRAY2RGB)
font = cv2.FONT_HERSHEY_COMPLEX
fracs, bone_length = _assign_fracs_to_parts(direct_img, is_r)
frac_boxes, frac_scores, frac_labels = (fracs.boxes, fracs.scores,
fracs.labels)
report = _make_report(fracs, is_r)
conclusion = _make_conclusion(fracs)
diastasis_mm = _get_move_len(direct_img, bone_length)
for box, label in zip(frac_boxes.cpu().numpy(), frac_labels):
xmin, ymin, xmax, ymax = box.astype(int)
cv2.rectangle(img, (xmin, ymin), (xmax, ymax), color=(255, 0, 0),
thickness=2)
cv2.putText(img, label, (xmin, max(ymin - 5, 0)),
font, 0.5, (255, 0, 0), thickness=1, lineType=cv2.LINE_AA)
cv2.imwrite(f"client/static/{study_iuid}.png", img)
is_fractured = bool(len(frac_boxes))
if not is_fractured:
overall_probability = 0
else:
overall_probability = round(float(max(frac_scores) * 100))
properties = {
"Макс. величина диастаза отломков, мм": diastasis_mm
}
return structs.Prediction(overall_probability, is_fractured,
report, conclusion, img, properties)
def predict(input: structs.PredictorInput) -> structs.Prediction:
return _visualize_detections(input.image, input.study_uid, input.laterality)