507 lines
25 KiB
Python
507 lines
25 KiB
Python
"""
|
||
Модель классификации качества 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,
|
||
)
|