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