bone_2026/src/dxa/model.py

507 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
Модель классификации качества DXA исследований.
Бинарная классификация: 0 — изображение годно, 1 — есть нарушение.
Вспомогательная голова предсказывает анатомическую область (spine / hip_right /
hip_left / unknown). Её предсказание используется как дополнительный сигнал
(auxiliary loss) и складывается с основным логитом, что заставляет backbone
учитывать область при оценке качества. На входе модели область НЕ известна —
иначе инференс на закрытых данных зависел бы от соглашения об именах файлов.
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Any, Dict, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
import torchvision.models as models
from sklearn.metrics import (
average_precision_score,
confusion_matrix,
f1_score,
precision_score,
recall_score,
roc_auc_score,
)
from src.dxa.labels import REGIONS
from src.dxa.preprocess import PreprocessConfig
BACKBONES = ("resnet18", "resnet34")
QUALITY_CLASSES = 2
# 0 — неизвестная область, далее по порядку из labels.REGIONS
REGION_CLASSES = 4
_WEIGHTS = {
"resnet18": ("IMAGENET1K_V1", 512),
"resnet34": ("IMAGENET1K_V1", 512),
}
logger = logging.getLogger(__name__)
class DXAQualityClassifier(nn.Module):
"""
ResNet backbone + голова качества + вспомогательная голова области.
Голова выбирается параметром `head`:
- ``linear`` (по умолчанию) — один линейный слой на pooled-признаках.
Это линейный зонд: обучаются 2×(512+1) параметра, поэтому на выборке
из ~250 снимков он не переобучается. На валидации даёт AUC ≈ 0.80,
тогда как MLP-голова уходит в переобучение (train F1 → 0.9 при val AUC
≈ 0.5), что проверено экспериментально на этом датасете.
- ``mlp`` — двухслойная голова; требует существенно больше данных.
При ``feature_norm=True`` вход головы стандартизуется по статистикам
обучающей выборки, зафиксированным в буферах `feat_mean`/`feat_std`
(см. `fit_feature_norm`). Без этого логиты смещены, вероятности
скучены у нуля и подобранный порог лишается смысла.
"""
def __init__(
self,
backbone: str = "resnet18",
pretrained: bool = True,
dropout: float = 0.3,
head: str = "linear",
hidden_dim: int = 256,
feature_norm: bool = True,
):
super().__init__()
if backbone not in _WEIGHTS:
raise ValueError(f"Unknown backbone {backbone!r}; expected one of {BACKBONES}")
if head not in ("linear", "mlp"):
raise ValueError(f"head must be 'linear' or 'mlp', got {head!r}")
self.backbone_name = backbone
self.head_type = head
self.feature_norm = feature_norm
weights_name, feature_dim = _WEIGHTS[backbone]
net = getattr(models, backbone)(weights=weights_name if pretrained else None)
self.features = nn.Sequential(*list(net.children())[:-1]) # всё, кроме fc
self.feature_dim = feature_dim
self.register_buffer("feat_mean", torch.zeros(feature_dim))
self.register_buffer("feat_std", torch.ones(feature_dim))
if head == "linear":
# Dropout в линейном зонде только вредит: он добавляет шум в
# единственный линейный слой, который и так сильно регуляризован
# weight decay. Проверено на датасете: с dropout AUC ≈ 0.71,
# без него ≈ 0.80.
self.quality_head = nn.Sequential(nn.Flatten(), nn.Linear(feature_dim, QUALITY_CLASSES))
else:
self.quality_head = nn.Sequential(
nn.Flatten(),
nn.Dropout(dropout),
nn.Linear(feature_dim, hidden_dim),
nn.ReLU(inplace=True),
nn.Dropout(dropout),
nn.Linear(hidden_dim, QUALITY_CLASSES),
)
self.region_head = nn.Sequential(
nn.Flatten(),
nn.Dropout(dropout) if head == "mlp" else nn.Identity(),
nn.Linear(feature_dim, REGION_CLASSES),
)
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
return self.features(x)
def _head_input(self, feats: torch.Tensor) -> torch.Tensor:
"""Сгладить признаки до (B, feature_dim) и применить стандартизацию."""
feats = feats.flatten(1)
if self.feature_norm:
feats = (feats - self.feat_mean) / self.feat_std
return feats
def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]:
"""Вернуть логиты качества и логиты анатомической области.
Голова области обучается как вспомогательная задача (auxiliary loss):
она заставляет backbone различать анатомию, но НЕ сдвигает логиты
качества. Осторожный сдвиг к «нарушению» при неопределённой анатомии
применяется на этапе инференса (см. `src.dxa.inference`), чтобы
обучение оставалось устойчивым.
"""
feats = self._head_input(self.extract_features(x))
return {"quality_logits": self.quality_head(feats), "region_logits": self.region_head(feats)}
def predict_region(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Предсказать анатомическую область: (region_id, вероятность уверенности)."""
with torch.no_grad():
logits = self.forward(x)["region_logits"]
probs = torch.softmax(logits, dim=1)
confidence, region_id = probs.max(dim=1)
return region_id, confidence
def compute_metrics(
logits: torch.Tensor,
labels: torch.Tensor,
threshold: float = 0.0,
) -> Dict[str, Any]:
"""
Метрики бинарной классификации качества.
Решение принимается по логиту (log-odds) класса «нарушение», а не по
вероятности. После обучения линейный зонд разделяет обучающую выборку
почти идеально, поэтому вероятности насыщаются в 0/1: порог в единицах
вероятности вырождается (например 1e-7), а логиты остаются умеренными.
Порог по логитам численно устойчив; для отчётности он переводится в
вероятность через сигмоиду.
ROC-AUC и PR-AUC считаются по вероятностям и от порога не зависят.
"""
logits = logits.detach().cpu().flatten().float()
labels = labels.detach().cpu().flatten().long()
probs = torch.sigmoid(logits)
predicted = (logits >= threshold).long()
tp = int(((predicted == 1) & (labels == 1)).sum())
tn = int(((predicted == 0) & (labels == 0)).sum())
fp = int(((predicted == 1) & (labels == 0)).sum())
fn = int(((predicted == 0) & (labels == 1)).sum())
metrics: Dict[str, Any] = {
"accuracy": (tp + tn) / max(len(labels), 1),
"precision": precision_score(labels, predicted, zero_division=0),
"recall": recall_score(labels, predicted, zero_division=0),
"f1": f1_score(labels, predicted, zero_division=0),
"tp": tp, "tn": tn, "fp": fp, "fn": fn,
"threshold_logit": threshold,
"threshold_prob": float(torch.sigmoid(torch.tensor(threshold))),
"n": int(len(labels)),
"n_pos": int((labels == 1).sum()),
}
if len(set(labels.tolist())) > 1:
metrics["roc_auc"] = roc_auc_score(labels, probs)
metrics["pr_auc"] = average_precision_score(labels, probs)
metrics["confusion_matrix"] = confusion_matrix(labels, predicted).tolist()
else:
metrics["roc_auc"] = None
metrics["pr_auc"] = None
return metrics
def select_threshold(
logits: torch.Tensor,
labels: torch.Tensor,
min_recall: float = 0.0,
) -> Tuple[float, Dict[str, Any]]:
"""
Подобрать порог по логитам, максимизирующий F1 на валидации.
При доле брака ~15 % порог 0 (вероятность 0.5) почти всегда даёт нулевой
recall, поэтому порог калибруется по валидации и сохраняется в чекпоинт.
F1 по всем порогам считается векторно: вызов sklearn на каждый кандидат
делает подбор на порядки медленнее и заметно замедляет обучение.
Args:
min_recall: если задано, порог берётся максимальным среди дающих
recall не ниже указанного (снижает пропуск брака).
"""
logits = logits.detach().cpu().flatten().numpy().astype(np.float64)
labels = labels.detach().cpu().flatten().long().numpy()
if len(set(labels.tolist())) < 2:
return 0.0, compute_metrics(torch.from_numpy(logits), torch.from_numpy(labels), 0.0)
# Кандидаты: все наблюдаемые логиты плюс края диапазона.
candidates = np.unique(np.concatenate([logits, [logits.min() - 1.0, logits.max() + 1.0]]))
predicted = logits[None, :] >= candidates[:, None]
tp = (predicted & (labels == 1)[None, :]).sum(axis=1).astype(np.float64)
fp = (predicted & (labels == 0)[None, :]).sum(axis=1).astype(np.float64)
fn = ((~predicted) & (labels == 1)[None, :]).sum(axis=1).astype(np.float64)
with np.errstate(divide="ignore", invalid="ignore"):
precision = np.where(tp + fp > 0, tp / (tp + fp), 0.0)
recall = np.where(tp + fn > 0, tp / (tp + fn), 0.0)
f1 = np.where(precision + recall > 0, 2 * precision * recall / (precision + recall), 0.0)
eligible = np.ones_like(f1, dtype=bool)
if min_recall > 0:
eligible = recall >= min_recall
if not eligible.any():
eligible = np.ones_like(f1, dtype=bool)
# Среди порогов с равным F1 берём минимальный, чтобы не терять recall.
masked = np.where(eligible, f1, -1.0)
best_threshold = float(candidates[int(np.argmax(np.round(masked, 10)))])
return best_threshold, compute_metrics(
torch.from_numpy(logits), torch.from_numpy(labels), best_threshold
)
def per_region_metrics(
logits: torch.Tensor,
labels: torch.Tensor,
region_ids: torch.Tensor,
threshold: float,
) -> Dict[str, Dict[str, Any]]:
"""
Метрики отдельно по анатомическим областям.
Разбивка нужна для честной интерпретации: в этом датасете нарушения резко
неравномерны (в позвоночнике ~29 % снимков с нарушением против ~4–5 % у
бёдер), а область почти однозначно определяется по ширине кадра. Поэтому
высокий общий AUC может отражать не распознавание дефекта, а различение
области исследования.
Args:
region_ids: истинные идентификаторы областей (1 — позвоночник, 2 —
правый, 3 — левый). Группировка по истинной области обязательна:
по предсказанной метрики смещались бы в сторону тех областей,
которые модель путает, и перестали бы показывать реальную картину.
"""
logits = logits.detach().cpu().flatten()
labels = labels.detach().cpu().flatten().long()
region_ids = region_ids.detach().cpu().flatten().long()
result: Dict[str, Dict[str, Any]] = {}
for index, name in enumerate(REGIONS, start=1):
mask = region_ids == index
if int(mask.sum()) == 0:
continue
result[name] = compute_metrics(logits[mask], labels[mask], threshold)
return result
class DXAQualityModel:
"""Обёртка вокруг сети: обучение, валидация, предсказание, сохранение/загрузка."""
def __init__(
self,
model: DXAQualityClassifier,
device: str = "cpu",
learning_rate: float = 1e-4,
weight_decay: float = 1e-4,
region_loss_weight: float = 0.3,
pos_weight: Optional[float] = None,
):
self.model = model
self.device = torch.device(device)
self.model.to(self.device)
self._backbone_trainable = True
self.quality_criterion = nn.CrossEntropyLoss(
weight=None if pos_weight is None else torch.tensor([1.0, float(pos_weight)], device=self.device)
)
self.region_criterion = nn.CrossEntropyLoss(ignore_index=-1)
self.region_loss_weight = region_loss_weight
self.optimizer = torch.optim.AdamW(
model.parameters(), lr=learning_rate, weight_decay=weight_decay
)
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
self.optimizer, mode="min", factor=0.5, patience=3
)
self.history: Dict[str, list] = {"train_loss": [], "val_loss": [], "val_f1": [], "val_roc_auc": []}
def _step(self, batch, train: bool) -> Tuple[float, torch.Tensor, torch.Tensor]:
images, meta = batch
images = images.to(self.device, non_blocking=True)
labels = meta["label"].to(self.device)
region_ids = meta["region_id"].to(self.device)
with torch.set_grad_enabled(train):
out = self.model(images)
loss = self.quality_criterion(out["quality_logits"], labels)
if self.region_loss_weight > 0:
loss = loss + self.region_loss_weight * self.region_criterion(out["region_logits"], region_ids)
if train:
self.optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=5.0)
self.optimizer.step()
logits = out["quality_logits"][:, 1]
return float(loss.item()) * len(labels), logits.detach(), labels
def set_backbone_trainable(self, trainable: bool) -> None:
"""
Включить или отключить обучение backbone.
При ~250 обучающих снимках полный fine-tune ResNet18 быстро
переобучается (train F1 -> 1.0 при случайном val AUC). Поэтому backbone
заморожен и обучается только голова — это линейный зонд на признаках
ImageNet.
Кроме requires_grad отключается и режим train для слоёв BatchNorm:
иначе бегущие статистики продолжают обновляться на обучающих батчах и
признаки «уезжают» от тех, на которых оценивалась стандартизация в
`fit_feature_norm`. Слой остаётся в eval, поэтому признаки стабильны.
"""
self._backbone_trainable = trainable
for module in self.model.features.modules():
if isinstance(module, torch.nn.modules.batchnorm._BatchNorm):
module.train(trainable)
for param in self.model.features.parameters():
param.requires_grad = trainable
def _apply_train_mode(self) -> None:
"""
Перевести модель в режим обучения с учётом заморозки backbone.
`Module.train()` включает train и для слоёв BatchNorm, что при
замороженном backbone сдвигало бы бегущие статистики. Поэтому после
перевода модели в train слои backbone возвращаются в eval.
"""
self.model.train()
if not getattr(self, "_backbone_trainable", True):
for module in self.model.features.modules():
if isinstance(module, torch.nn.modules.batchnorm._BatchNorm):
module.eval()
def set_learning_rate(self, lr: float) -> None:
"""Задать learning rate всем группам параметров."""
for group in self.optimizer.param_groups:
group["lr"] = lr
@torch.no_grad()
def fit_feature_norm(self, loader) -> None:
"""
Оценить среднее и СКО признаков по обучающей выборке и зафиксировать их.
Стандартизация входа нужна линейной голове: без неё логиты смещены,
вероятности скучены у нуля, а подобранный порог теряет смысл.
Дисперсия считается в два прохода (сначала среднее, затем сумма
квадратов отклонений). Формула E[x²]−E[x]² при float32 на признаках
порядка 10 даёт погрешность, сопоставимую с самой дисперсией: СКО
выходило случайным, из-за чего логиты насыщались и порог вырождался.
Буферы не обучаемые, поэтому статистики не «подглядывают» в валидацию.
"""
if not self.model.feature_norm:
return
self.model.eval()
chunks = []
for images, _ in loader:
chunks.append(self.model.extract_features(images.to(self.device)).flatten(1).cpu().float())
if not chunks:
return
# float64 на CPU: размерность мала, а точность здесь критична.
feats = torch.cat(chunks).double()
mean = feats.mean(dim=0)
std = torch.sqrt(((feats - mean) ** 2).mean(dim=0))
# Нижняя граница СКО: у части размерностей разброс близок к нулю, а
# деление на него усиливает шум в десятки раз и насыщает логиты.
floor = max(float(std.median()) * 0.25, 1e-6)
std = std.clamp(min=floor)
self.model.feat_mean.copy_(mean.float().to(self.model.feat_mean.device))
self.model.feat_std.copy_(std.float().to(self.model.feat_std.device))
logger.debug(
"Feature norm fitted: %d samples, mean norm %.2f, std median %.4f, floor %.5f",
feats.shape[0], float(mean.norm()), float(std.median()), floor,
)
def train_epoch(self, loader) -> Tuple[float, Dict[str, Any]]:
self._apply_train_mode()
total_loss = 0.0
logits, labels = [], []
for batch in loader:
loss, l, y = self._step(batch, train=True)
total_loss += loss
logits.append(l)
labels.append(y)
n = max(len(loader.dataset), 1)
return total_loss / n, compute_metrics(torch.cat(logits), torch.cat(labels))
def validate(self, loader) -> Tuple[float, Dict[str, Any]]:
self.model.eval()
total_loss = 0.0
logits, labels = [], []
for batch in loader:
loss, l, y = self._step(batch, train=False)
total_loss += loss
logits.append(l)
labels.append(y)
n = max(len(loader.dataset), 1)
metrics = compute_metrics(torch.cat(logits), torch.cat(labels))
self.scheduler.step(total_loss / n)
return total_loss / n, metrics
def predict_logits(self, loader) -> Tuple[torch.Tensor, torch.Tensor]:
"""Логиты класса «нарушение» и истинные метки для всего набора."""
logits, labels, _ = self.predict_logits_with_regions(loader)
return logits, labels
def predict_logits_with_regions(self, loader) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Логиты, метки и ИСТИННЫЕ идентификаторы областей для всего набора.
Истинные области берутся из метаданных датасета (они известны при
обучении) и нужны для метрик по областям, чтобы оценка не смещалась
ошибками головы области.
"""
self.model.eval()
logits, labels, regions = [], [], []
with torch.no_grad():
for batch in loader:
images, meta = batch
out = self.model(images.to(self.device))
logits.append(out["quality_logits"][:, 1].cpu())
labels.append(meta["label"])
regions.append(meta["region_id"])
return torch.cat(logits), torch.cat(labels), torch.cat(regions)
def save(self, path: str | Path, preprocess: Optional[PreprocessConfig] = None, **extra) -> None:
"""Сохранить чекпоинт вместе с архитектурой и параметрами предобработки."""
payload: Dict[str, Any] = {
"model_state_dict": self.model.state_dict(),
"optimizer_state_dict": self.optimizer.state_dict(),
"history": self.history,
"backbone": self.model.backbone_name,
"head": self.model.head_type,
"format_version": 2,
"region_loss_weight": self.region_loss_weight,
}
if preprocess is not None:
payload["preprocess"] = preprocess.to_dict()
payload.update(extra)
Path(path).parent.mkdir(parents=True, exist_ok=True)
torch.save(payload, path)
def load(self, path: str | Path) -> Dict[str, Any]:
"""Загрузить чекпоинт. Возвращает метаданные (пустой dict для старых файлов)."""
# weights_only=False: чекпоинт содержит метрики и конфиг, а не только тензоры.
# Файлы создаются самим проектом, поэтому источник считается доверенным.
checkpoint = torch.load(path, map_location=self.device, weights_only=False)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
self.model.load_state_dict(checkpoint["model_state_dict"])
if "optimizer_state_dict" in checkpoint:
try:
self.optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
except (ValueError, KeyError):
pass # старый чекпоинт с другой архитектурой оптимизатора
self.history = checkpoint.get("history", self.history)
return {k: v for k, v in checkpoint.items() if k != "model_state_dict"}
raise ValueError(f"Checkpoint {path} has no 'model_state_dict' (raw state_dict is not supported)")
def create_model(
backbone: str = "resnet18",
pretrained: bool = True,
device: str = "cpu",
head: str = "linear",
**kwargs,
) -> DXAQualityModel:
"""Создать обёртку модели с заданным backbone и головой."""
return DXAQualityModel(
DXAQualityClassifier(backbone=backbone, pretrained=pretrained, head=head),
device=device,
**kwargs,
)