495 lines
17 KiB
Python
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)
|