develop - hack_2026
This commit is contained in:
parent
1b82992a6d
commit
5a8dbb43a5
|
|
@ -1,5 +1,9 @@
|
||||||
# .dockerignore
|
# .dockerignore
|
||||||
|
# Цель: контекст сборки должен быть маленьким. Образ содержит только src,
|
||||||
|
# models и run.sh (см. Dockerfile), поэтому данные, тесты и служебные файлы
|
||||||
|
# исключаются — DICOM-датасет не должен попадать в образ с медицинскими данными.
|
||||||
__pycache__
|
__pycache__
|
||||||
|
**/__pycache__
|
||||||
*.pyc
|
*.pyc
|
||||||
*.pyo
|
*.pyo
|
||||||
*.pyd
|
*.pyd
|
||||||
|
|
@ -23,8 +27,11 @@ env
|
||||||
*.t7
|
*.t7
|
||||||
data/
|
data/
|
||||||
datasets/
|
datasets/
|
||||||
.DS_Store
|
dataset_hack/
|
||||||
|
tests/
|
||||||
|
.qwen/
|
||||||
.idea/
|
.idea/
|
||||||
.vscode/
|
.vscode/
|
||||||
|
.continue/
|
||||||
*.swp
|
*.swp
|
||||||
*.swo
|
*.swo
|
||||||
|
|
|
||||||
71
Dockerfile
71
Dockerfile
|
|
@ -1,38 +1,59 @@
|
||||||
# DXA Quality Assessment - Docker Container (CPU only, Python 3.11)
|
# DXA Quality Assessment — контейнер для инференса (CPU)
|
||||||
# Build: docker build -t dxa-quality .
|
#
|
||||||
# Run: docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
# Сборка: docker build -t dxa-quality .
|
||||||
|
# Запуск: docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
||||||
|
#
|
||||||
|
# Требования методики: зависимости с зафиксированными версиями и запуск
|
||||||
|
# локально без обращения к внешним сервисам. Inference не использует
|
||||||
|
# предобученные веса из интернета — они уже внутри чекпоинта, поэтому
|
||||||
|
# при запуске сеть не нужна.
|
||||||
|
FROM python:3.11-slim
|
||||||
|
|
||||||
FROM python:3.11-slim AS builder
|
# libgl1/libglib2.0-0 нужны opencv (импортируется через зависимости проекта),
|
||||||
|
# libgomp1 — для параллельных циклов torch.
|
||||||
RUN apt-get update && apt-get install -y \
|
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||||
build-essential \
|
|
||||||
ninja-build \
|
|
||||||
libgl1 \
|
libgl1 \
|
||||||
libglib2.0-0 \
|
libglib2.0-0 \
|
||||||
libsm6 \
|
|
||||||
libxext6 \
|
|
||||||
libxrender1 \
|
|
||||||
libgomp1 \
|
libgomp1 \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
# Свежий meson (0.64.0+) — ДО установки Python-пакетов
|
|
||||||
RUN pip3 install --no-cache-dir "meson>=0.64.0"
|
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Устанавливаем пакеты без жёстких пинов — pip сам разрешит зависимости
|
# --- Слой зависимостей: отдельно от кода, чтобы кэшировался ---
|
||||||
RUN pip3 install --no-cache-dir \
|
# Все версии зафиксированы, torch ставится с CPU-индексом: образ для инференса
|
||||||
fastapi uvicorn torch torchvision monai TotalSegmentator \
|
# не требует GPU, а версия из PyPI на Linux подтянула бы CUDA-колёса.
|
||||||
nibabel pydicom pydicom-seg opencv-python-headless \
|
COPY requirements.txt ./
|
||||||
scikit-learn scipy pandas numpy openpyxl matplotlib timm \
|
RUN pip install --no-cache-dir --upgrade pip \
|
||||||
python-multipart
|
&& pip install --no-cache-dir \
|
||||||
|
--extra-index-url https://download.pytorch.org/whl/cpu \
|
||||||
|
-r requirements.txt
|
||||||
|
|
||||||
COPY . .
|
# --- Код ---
|
||||||
RUN mkdir -p models
|
COPY src ./src
|
||||||
|
COPY models ./models
|
||||||
|
|
||||||
|
# Скрипт пакетной обработки (оценка берёт его из корня образа).
|
||||||
|
COPY run.sh ./run.sh
|
||||||
|
RUN chmod +x ./run.sh
|
||||||
|
|
||||||
|
# smoke-проверка: чекпоинт должен читаться. Падает на сборке, если модели нет,
|
||||||
|
# вместо тихого 500 на каждом запросе в рантайме.
|
||||||
|
RUN python -c "import os, torch; p='models/dxa_model.pth'; \
|
||||||
|
assert os.path.exists(p), f'missing checkpoint {p}'; \
|
||||||
|
ck=torch.load(p, map_location='cpu', weights_only=False); \
|
||||||
|
assert 'model_state_dict' in ck, 'checkpoint has no model_state_dict'; \
|
||||||
|
print('checkpoint ok:', ck.get('backbone'), ck.get('head'))"
|
||||||
|
|
||||||
|
ENV PYTHONUNBUFFERED=1 \
|
||||||
|
PYTHONPATH=/app \
|
||||||
|
DXA_MODEL_PATH=/app/models/dxa_model.pth
|
||||||
|
|
||||||
EXPOSE 8000
|
EXPOSE 8000
|
||||||
|
|
||||||
ENV PYTHONUNBUFFERED=1
|
# HEALTHCHECK опирается на /api/v1/health, который отдаёт model_loaded.
|
||||||
ENV PYTHONPATH=/app
|
HEALTHCHECK --interval=30s --timeout=10s --start-period=40s --retries=3 \
|
||||||
|
CMD python -c "import urllib.request,sys; \
|
||||||
|
r=urllib.request.urlopen('http://127.0.0.1:8000/api/v1/health', timeout=5); \
|
||||||
|
sys.exit(0 if r.status==200 else 1)"
|
||||||
|
|
||||||
CMD ["python3", "-m", "uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
CMD ["python", "-m", "uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||||
|
|
|
||||||
364
QWEN.md
364
QWEN.md
|
|
@ -2,237 +2,173 @@
|
||||||
|
|
||||||
## Project Overview
|
## Project Overview
|
||||||
|
|
||||||
Medical AI service for automated assessment of DXA (bone densitometry) study quality. The system analyzes DICOM files and evaluates quality based on standard criteria.
|
Медицинский ИИ-сервис для автоматизированной оценки качества денситометрических
|
||||||
|
исследований (DXA). Принимает DICOM, определяет анатомическую область, решает
|
||||||
|
бинарную задачу «качественное изображение / есть нарушение» и формирует отчёт.
|
||||||
|
|
||||||
### Core Purpose (Hackathon)
|
### Ключевые требования (condition.txt)
|
||||||
- Analyze DICOM densitometry studies
|
|
||||||
- Determine anatomical region (spine/hip)
|
|
||||||
- Binary classification: quality (OK/violation)
|
|
||||||
- Detailed violation type detection (motion, artifacts, position, ROI)
|
|
||||||
- Output results in XLSX/CSV format per requirements
|
|
||||||
|
|
||||||
### Tech Stack
|
- Области: поясничный отдел позвоночника и проксимальный отдел бедренной кости.
|
||||||
| Component | Technology |
|
- Обязательна контейнеризация и скрипт сборки/запуска в Linux.
|
||||||
|-----------|------------|
|
- Результат: XLSX/CSV, одна строка на изображение, столбцы
|
||||||
| Backend | Python 3.10, FastAPI, Uvicorn |
|
`path_to_study, study_uid, image_uid, anatomical_region, quality_class,
|
||||||
| ML/Deep Learning | PyTorch, torchvision (ResNet18) |
|
violation_type, processing_status, time_of_processing`.
|
||||||
| Image Processing | PIL, OpenCV, pydicom, scipy |
|
- Приоритетные метрики: F1 и ROC-AUC (с 95 % доверительными интервалами).
|
||||||
| Data Handling | pandas, openpyxl |
|
- Работа офлайн, без передачи изображений во внешние сервисы.
|
||||||
| Containerization | Docker |
|
- Время обработки одного исследования — не более 3 минут.
|
||||||
|
|
||||||
|
### Технологии
|
||||||
|
|
||||||
|
| Компонент | Технология |
|
||||||
|
|---|---|
|
||||||
|
| Backend | Python, FastAPI, Uvicorn |
|
||||||
|
| ML | PyTorch, torchvision (ResNet18) |
|
||||||
|
| Изображения | pydicom, Pillow, OpenCV, SciPy |
|
||||||
|
| Данные | pandas, openpyxl, scikit-learn |
|
||||||
|
| Тесты | pytest |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Project Structure
|
## Действующая архитектура
|
||||||
|
|
||||||
|
Всё, что реально работает, находится в `src/dxa/` и `src/main.py`.
|
||||||
|
|
||||||
```
|
```
|
||||||
bone_2026/
|
src/dxa/
|
||||||
├── src/
|
├── labels.py # имена -> метки, склейка дублей, разбиение по исследованиям
|
||||||
│ ├── main.py # FastAPI app (DXA mode)
|
├── preprocess.py # DICOM -> CHW-тензор (единый путь для обучения и API)
|
||||||
│ ├── run.py # Server runner
|
├── dataset.py # DXADataset, DataLoader
|
||||||
│ ├── dxa/ # DXA Quality module
|
├── model.py # сеть, метрики, подбор порога, сохранение/загрузка
|
||||||
│ │ ├── dataset.py # DXADataset class
|
├── train.py # обучение + отчёт (md/json)
|
||||||
│ │ ├── model.py # ResNet18 classifier
|
└── inference.py # пакетный инференс, определение области, визуализация
|
||||||
│ │ ├── train.py # Training script
|
|
||||||
│ │ ├── inference.py # Batch inference
|
|
||||||
│ │ └── __init__.py
|
|
||||||
│ ├── api/ # REST endpoints & schemas
|
|
||||||
│ │ ├── endpoints.py # Original endpoints
|
|
||||||
│ │ └── static/ # Web UI
|
|
||||||
│ │ ├── index.html
|
|
||||||
│ │ └── js/dxa-app.js
|
|
||||||
│ ├── core/ # Orchestrator
|
|
||||||
│ ├── quality/ # Quality scoring
|
|
||||||
│ │ ├── quality_scorer.py # Base scorer
|
|
||||||
│ │ ├── detailed_assessment.py # NEW: Detailed assessment
|
|
||||||
│ │ ├── position_validator.py
|
|
||||||
│ │ └── artifact_detector.py
|
|
||||||
│ └── segmentators/ # Segmentation models
|
|
||||||
├── models/
|
|
||||||
│ └── dxa_model.pth # Trained DXA classifier
|
|
||||||
├── dataset_hack/ # DICOM datasets
|
|
||||||
│ ├── Для теста/ # Test data
|
|
||||||
│ └── НД_для_обучения/ # Training data
|
|
||||||
├── requirements.txt
|
|
||||||
├── Dockerfile
|
|
||||||
├── run.sh
|
|
||||||
└── README.md
|
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
### Ключевые решения (проверены экспериментально)
|
||||||
|
|
||||||
## Implemented Features (Current)
|
| Решение | Причина |
|
||||||
|
|---|---|
|
||||||
|
| Метки из имён файлов: `_bad` > `_good` > нет метки (=good) | Явная оценка в имени файла; отсутствие метки означает «хорошее» |
|
||||||
|
| Склейка побайтных дублей | 544 файла, но 252 уникальных снимка; без склейки снимок попадал в оба класса |
|
||||||
|
| Разбиение по исследованиям, не по снимкам | Исключение утечки: снимки одного исследования в одной части |
|
||||||
|
| Линейный зонд (замороженный backbone) | Полный fine-tune при ~250 снимках переобучается (val AUC → 0.5) |
|
||||||
|
| Порог по логиту, подбор по F1 | При 15 % нарушений порог 0.5 даёт нулевой recall; вероятности насыщаются |
|
||||||
|
| Аугментация выключена по умолчанию | Яркость и положение сами являются признаками качества: AUC 0.87 → 0.56 |
|
||||||
|
| Область по ширине кадра | Позвоночник 300 px, бедро 280 px; 99/99 для позвоночника |
|
||||||
|
|
||||||
### ✅ API Endpoints
|
### Метрики (валидация, разбиение по исследованиям)
|
||||||
| Method | Endpoint | Description |
|
|
||||||
|--------|----------|-------------|
|
|
||||||
| GET | `/` | Web interface |
|
|
||||||
| GET | `/api/v1/health` | Health check |
|
|
||||||
| POST | `/api/v1/analyze` | Basic analysis |
|
|
||||||
| POST | `/api/v1/analyze/detailed` | **NEW: Detailed analysis with metrics** |
|
|
||||||
| POST | `/api/v1/analyze/sr` | **NEW: DICOM SR report** |
|
|
||||||
| POST | `/api/v1/batch` | Batch processing |
|
|
||||||
| POST | `/api/v1/export` | Export to XLSX |
|
|
||||||
|
|
||||||
### ✅ Detailed Assessment (`src/quality/detailed_assessment.py`)
|
- ROC-AUC по 5 разбиениям: **0.76 ± 0.08** (0.64 – 0.84)
|
||||||
- **Motion detection**: Laplacian variance, FFT blur analysis
|
- PR-AUC по 5 разбиениям: 0.48 ± 0.17 (базовый уровень при 15 % нарушений — 0.15)
|
||||||
- **Artifact detection**: Metal, implants, cement, calcifications
|
- F1 по 5 разбиениям: 0.54 ± 0.11 (порог подобран на той же валидации → смещено вверх)
|
||||||
- **Spine completeness**: Vertebrae count, alignment, spacing
|
- ROC-AUC, 5-фолдовая CV линейного зонда: **0.81 ± 0.08**
|
||||||
- **Hip completeness**: Full visibility, aspect ratio
|
- Контрольная задача «позвоночник / бедро»: AUC 1.00 (проверка пайплайна)
|
||||||
- **Hip rotation**: Major axis angle calculation
|
- Нулевая гипотеза (перестановка меток): AUC 0.64
|
||||||
- **ROI validation**: Boundary check, margin, size
|
|
||||||
- **Violation types**: correct, position_error, artifact_motion, artifact_other, labeling_error, incomplete_view, roi_error, rotation
|
|
||||||
|
|
||||||
### ✅ Web Interface
|
Разбивка по областям обязательна: в позвоночнике ~29 % нарушений против ~4–5 %
|
||||||
- Drag-and-drop DICOM upload
|
у бёдер, а область почти однозначно определяется по ширине кадра, поэтому общий
|
||||||
- Table with results (filter, sort, search)
|
AUC частично отражает различение области.
|
||||||
- **NEW: Detail panel** - click on row to see:
|
|
||||||
- Violation type
|
|
||||||
- Reason (human-readable)
|
|
||||||
- Motion metrics
|
|
||||||
- Artifact detection
|
|
||||||
- ROI validation
|
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## DXA Module (`src/dxa/`)
|
## Данные (`dataset_hack/`)
|
||||||
|
|
||||||
### Dataset (`dataset.py`)
|
Структура:
|
||||||
- Loads DICOM files from studies
|
|
||||||
- Parses annotation Excel file
|
|
||||||
- **Automatically detects anatomical region from image content**
|
|
||||||
- Maps regions: spine, hip_right, hip_left
|
|
||||||
- Quality labels: 0 (OK), 1 (violation)
|
|
||||||
|
|
||||||
### Model (`model.py`)
|
```
|
||||||
- Architecture: ResNet18 (pretrained on ImageNet)
|
dataset_hack/
|
||||||
- Task: Binary classification (quality OK vs violation)
|
├── Для теста/ # bad.dcm, l_hip.dcm, r_hip.dcm, spine.dcm
|
||||||
- Input: 224x224 RGB images
|
└── НД_для_обучения/
|
||||||
- Output: class probabilities
|
├── разметка.xlsx # экспертная оценка на уровне ИССЛЕДОВАНИЯ
|
||||||
|
└── Исследования/<study_uid>/.../<region>_<n>[_good|_bad].dcm
|
||||||
|
```
|
||||||
|
|
||||||
|
Факты, важные для обучения:
|
||||||
|
|
||||||
|
- 544 файла на диске, но **252 уникальных снимка** (по пиксельному содержимому).
|
||||||
|
- 86 файлов имеют явную метку; после склейки дублей — **37 нарушений из 252 (14.7 %)**.
|
||||||
|
- Дубли не пересекают границы исследований, конфликтов меток при склейке нет.
|
||||||
|
- Имена неоднородны: `spine_01`, `Spine`, `r_spine`, `spine-1`, `l_hip`,
|
||||||
|
`l_hip-2`, `r_hip`, `r_hop`.
|
||||||
|
- В DICOM **нет** разметки ROI (ни OverlayData, ни GraphicAnnotationSequence),
|
||||||
|
поэтому корректность нанесённых областей нельзя проверить прямым сравнением.
|
||||||
|
- Метка в Excel относится к исследованию и раздаётся его снимкам; имена файлов
|
||||||
|
имеют приоритет. Excel используется только для предупреждения о расхождениях.
|
||||||
|
|
||||||
|
### Единица разметки — источник шума
|
||||||
|
|
||||||
|
Один снимок в исследовании помечен `_bad`, остальные не размечены. Метка снимка
|
||||||
|
считается унаследованной от исследования, поэтому часть меток заведомо шумная.
|
||||||
|
Это главное ограничение текущего качества модели.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Обучение
|
||||||
|
|
||||||
### Training (`train.py`)
|
|
||||||
```bash
|
```bash
|
||||||
python src/dxa/train.py --epochs 10 --batch-size 16
|
./run.sh train # режим по умолчанию
|
||||||
|
python -m src.dxa.train --dry-run # проверить данные без обучения
|
||||||
|
python -m src.dxa.train --head mlp --freeze-epochs 0 --epochs 30
|
||||||
```
|
```
|
||||||
|
|
||||||
### Inference (`inference.py`)
|
Артефакты в `--output-dir`: `dxa_model.pth` (веса, порог, параметры
|
||||||
|
предобработки), `train_report.md`, `train_report.json`.
|
||||||
|
|
||||||
|
Чекпоинт самодостаточен: `backbone`, `head`, `preprocess`, `threshold_logit`
|
||||||
|
хранятся внутри, поэтому инференс не может рассинхронизироваться с обучением.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## API (`src/main.py`)
|
||||||
|
|
||||||
|
| Метод | Путь | Назначение |
|
||||||
|
|---|---|---|
|
||||||
|
| GET | `/` | Веб-интерфейс |
|
||||||
|
| GET | `/api/v1/health` | Статус, признак загрузки модели |
|
||||||
|
| POST | `/api/v1/analyze` | Базовый анализ файла |
|
||||||
|
| POST | `/api/v1/analyze/detailed` | Расширенный отчёт, опционально маска |
|
||||||
|
| POST | `/api/v1/analyze/sr` | Текстовый отчёт DICOM SR |
|
||||||
|
| POST | `/api/v1/batch` | Пакетный анализ |
|
||||||
|
| POST | `/api/v1/export` | Пакетный анализ + XLSX |
|
||||||
|
|
||||||
|
Путь к модели — переменная окружения `DXA_MODEL_PATH` (по умолчанию
|
||||||
|
`models/dxa_model.pth`), чтобы контейнер не зависел от рабочего каталога.
|
||||||
|
API и CLI используют один код предсказания (`predict_from_bytes` /
|
||||||
|
`predict_from_array`), поэтому предобработка и порог совпадают.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Тесты
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python src/dxa/inference.py \
|
./run.sh test
|
||||||
--input-path dataset_hack/Для\ теста \
|
python -m pytest tests/ -q # 65 тестов
|
||||||
--output-path results.xlsx \
|
|
||||||
--model-path models/dxa_model.pth
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
`tests/test_labels.py` — разбор имён, склейка дублей, отсутствие утечки при
|
||||||
|
разбиении. `tests/test_preprocess_and_model.py` — предобработка, метрики,
|
||||||
|
подбор порога, контракт модели, BatchNorm при заморозке, roundtrip чекпоинта.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Running the Project
|
## Известные ограничения
|
||||||
|
|
||||||
### Training
|
1. Разметка на уровне исследования → шум в метках снимков.
|
||||||
```bash
|
2. Мало данных: 252 снимка, 37 нарушений; доверительные интервалы широкие.
|
||||||
python src/dxa/train.py --epochs 10
|
3. Тип нарушения определяется эвристиками, а не обученной моделью.
|
||||||
```
|
4. Порог `SPINE_MIN_WIDTH` привязан к текущему оборудованию.
|
||||||
|
5. Grad-CAM (`src/models/visualization/gradcam.py`) есть, но не подключён.
|
||||||
|
|
||||||
### Inference
|
## Устаревший код (не подключён к API)
|
||||||
```bash
|
|
||||||
python src/dxa/inference.py --input-path file.dcm --output-path result.xlsx
|
|
||||||
```
|
|
||||||
|
|
||||||
### API Server
|
Эти модули не импортируются из `src/main.py` и `src/dxa/*`; их зависимости
|
||||||
```bash
|
закомментированы в `requirements.txt`:
|
||||||
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
|
||||||
```
|
|
||||||
|
|
||||||
Web UI: http://localhost:8000
|
- `src/api/endpoints.py` — падает при импорте, роутер не монтируется.
|
||||||
|
- `src/api/annotation.py` — маршруты под `/api/annotation`, не монтируются.
|
||||||
---
|
- `src/core/orchestrator.py`, `src/pipeline/pipeline.py` — веса не загружаются.
|
||||||
|
- `src/model/unet.py`, `src/model/segmentator.py` — UNet-заглушки.
|
||||||
## Output Format (per Hackathon Requirements)
|
- `src/quality/artifact_detector.py`, `position_validator.py`, `medical_quality.py` — заглушки.
|
||||||
|
- `src/dataloaders/pet_dataset.py` — остаток прототипа (Oxford-IIIT Pet).
|
||||||
### Basic Output
|
|
||||||
| Column | Description |
|
|
||||||
|--------|-------------|
|
|
||||||
| path_to_study | Path to study directory |
|
|
||||||
| study_uid | StudyInstanceUID from DICOM |
|
|
||||||
| image_uid | SOPInstanceUID from DICOM |
|
|
||||||
| anatomical_region | spine / hip_left / hip_right / hip |
|
|
||||||
| quality_class | 0 (OK), 1 (violation) |
|
|
||||||
| violation_type | Type of violation (if any) |
|
|
||||||
| processing_status | Success / Failure |
|
|
||||||
| time_of_processing | Processing time (seconds) |
|
|
||||||
|
|
||||||
### Detailed Output (/api/v1/analyze/detailed)
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"anatomical_region": "spine",
|
|
||||||
"quality_class": 1,
|
|
||||||
"quality_label": "Violation detected",
|
|
||||||
"violation_type": "artifact_motion",
|
|
||||||
"reason": "Обнаружен артефакт движения (размытие)",
|
|
||||||
"confidence": 0.85,
|
|
||||||
"confidence_per_class": {
|
|
||||||
"correct": 0.15,
|
|
||||||
"violation": 0.85
|
|
||||||
},
|
|
||||||
"view_quality": "full",
|
|
||||||
"metrics": {
|
|
||||||
"motion": { "motion_detected": true, "severity": "HIGH" },
|
|
||||||
"artifacts": { "any_detected": true, "metal_detected": false },
|
|
||||||
"roi_check": { "valid": true }
|
|
||||||
},
|
|
||||||
"overall_quality": "POOR",
|
|
||||||
"severity": "HIGH"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Anatomical Region Detection
|
|
||||||
|
|
||||||
The system automatically determines the anatomical region from the DICOM image content:
|
|
||||||
|
|
||||||
### Algorithm (`src/dxa/inference.py`)
|
|
||||||
1. **Spine vs Hip** - by bright region shape:
|
|
||||||
- Extract 95th percentile threshold
|
|
||||||
- Calculate bounding box aspect ratio
|
|
||||||
- Spine: bbox_aspect < 1.5 (more square)
|
|
||||||
- Hip: bbox_aspect > 1.5 (vertically elongated)
|
|
||||||
|
|
||||||
2. **Hip Left vs Right** - by brightness asymmetry:
|
|
||||||
- Calculate left/right bright pixel ratio
|
|
||||||
- hip_left: L/R ratio < 0.7 (left side brighter)
|
|
||||||
- hip_right: L/R ratio > 1.3 (right side brighter)
|
|
||||||
- hip: unclear (fallback)
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Known Issues & Limitations
|
|
||||||
|
|
||||||
1. **Model training** - Needs retraining with new violation types
|
|
||||||
2. **Dataset size** - Currently ~100 studies, needs 500+
|
|
||||||
3. **Segmentation** - Uses simple threshold, needs proper model
|
|
||||||
4. **F1 score** - Currently ~0.27, needs improvement with weighted loss
|
|
||||||
5. **Heatmap visualization** - Not implemented (requires model retraining)
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Multi-Model Architecture (Planned)
|
|
||||||
|
|
||||||
See `docs/multi_model_architecture.md` for the planned pipeline:
|
|
||||||
|
|
||||||
```
|
|
||||||
Pipeline:
|
|
||||||
1. Region Detector → 2. Segmentator → 3. Quality Classifier → 4. Violation Type → 5. Aggregator
|
|
||||||
```
|
|
||||||
|
|
||||||
### Planned Models:
|
|
||||||
|
|
||||||
| Model | Purpose | File |
|
|
||||||
|-------|---------|------|
|
|
||||||
| Region Detector | Spine/Hip detection | `src/models/region_detector.py` |
|
|
||||||
| Segmentator | Bone segmentation | `src/models/segmentation/` |
|
|
||||||
| Quality Classifier | OK/Violation binary | `src/models/classification/quality.py` |
|
|
||||||
| Violation Classifier | 7+ violation types | `src/models/classification/violation.py` |
|
|
||||||
| Artifact Detector | Motion, metal detection | `src/models/artifacts/detector.py` |
|
|
||||||
| Grad-CAM | Attention heatmap | `src/models/visualization/gradcam.py` |
|
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|
@ -240,29 +176,9 @@ Pipeline:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
docker build -t dxa-quality .
|
docker build -t dxa-quality .
|
||||||
docker run -v /data:/data -p 8000:8000 dxa-quality
|
docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
Dockerfile ставит зафиксированные версии, копирует только `src/`, `models/` и
|
||||||
|
`run.sh`, проверяет чекпоинт на этапе сборки и имеет HEALTHCHECK. Данные и тесты
|
||||||
## Development Notes
|
в образ не попадают (`.dockerignore`).
|
||||||
|
|
||||||
### Code Style
|
|
||||||
- Follow existing patterns in src/
|
|
||||||
- Type hints where appropriate
|
|
||||||
- Minimal comments (only for context)
|
|
||||||
|
|
||||||
### Key Components
|
|
||||||
- **DXADataset**: Handles DICOM loading + annotation parsing
|
|
||||||
- **DXAQualityClassifier**: ResNet18-based classifier
|
|
||||||
- **process_dicom_files**: Batch inference with XLSX output
|
|
||||||
- **generate_quality_report**: Detailed assessment with metrics
|
|
||||||
|
|
||||||
### Dependencies
|
|
||||||
All in `requirements.txt`:
|
|
||||||
- `torch`, `torchvision` - Deep learning
|
|
||||||
- `pydicom` - DICOM handling
|
|
||||||
- `pandas`, `openpyxl` - Data/Excel
|
|
||||||
- `fastapi`, `uvicorn` - Web framework
|
|
||||||
- `Pillow`, `opencv-python-headless` - Image processing
|
|
||||||
- `scipy` - Image analysis (blur, artifacts)
|
|
||||||
|
|
|
||||||
471
README.md
471
README.md
|
|
@ -1,328 +1,259 @@
|
||||||
# 🦴 DXA Quality Assessment
|
# 🦴 DXA Quality Assessment
|
||||||
|
|
||||||
[](https://www.python.org/)
|
Сервис автоматизированного контроля качества денситометрических исследований (DXA):
|
||||||
[](https://fastapi.tiangolo.com/)
|
принимает DICOM, определяет анатомическую область, оценивает, пригодно ли изображение
|
||||||
[](https://pytorch.org/)
|
для клинической интерпретации, и формирует структурированный отчёт.
|
||||||
|
|
||||||
## Описание
|
## Что делает решение
|
||||||
|
|
||||||
Сервис для автоматизированной оценки качества денситометрических исследований (DXA). Система анализирует DICOM-изображения костной денситометрии и определяет качество исследования по следующим критериям:
|
| Шаг | Реализация |
|
||||||
|
|---|---|
|
||||||
|
| Определение анатомической области | Ширина кадра (позвоночник / бедро) + голова области + геометрия яркой зоны |
|
||||||
|
| Бинарная оценка качества | ResNet18 (ImageNet) → линейная голова; порог подобран по F1 на валидации |
|
||||||
|
| Тип нарушения | Эвристики по изображению: размытие/движение, посторонние включения, геометрия ROI |
|
||||||
|
| Отчёт | XLSX/CSV со столбцами из требований; опционально zip с визуализацией зоны интереса |
|
||||||
|
| API | FastAPI: анализ, детальный анализ, пакетная обработка, экспорт, DICOM SR (текст) |
|
||||||
|
| Веб-интерфейс | Загрузка DICOM, таблица результатов |
|
||||||
|
|
||||||
- **Артефакты** — движение, размытость, металлические объекты, имплантаты
|
## Установка и запуск
|
||||||
- **Позиционирование** — правильное расположение анатомической области в кадре
|
|
||||||
- **Полнота изображения** — видимость всех анатомических структур (позвонки L1-L4, бедро)
|
|
||||||
- **Ротация** — корректный угол поворота (для исследования бедра)
|
|
||||||
- **ROI-валидация** — правильность расположения области интереса
|
|
||||||
|
|
||||||
### Основные возможности
|
```bash
|
||||||
|
python3 -m venv venv && source venv/bin/activate
|
||||||
|
pip install -r requirements.txt
|
||||||
|
|
||||||
- 🔬 **Анализ DICOM** — загрузка и обработка медицинских изображений
|
./run.sh train # обучить модель качества
|
||||||
- 🧠 **Классификация** — бинарная оценка качества (OK / Violation)
|
./run.sh infer "dataset_hack/Для теста" results.xlsx # пакетная обработка
|
||||||
- 🔍 **Детекция нарушений** — определение типа нарушения:
|
./run.sh serve # API и веб-интерфейс на :8000
|
||||||
- `correct` — качество соответствует норме
|
./run.sh test # тесты
|
||||||
- `artifact_motion` — артефакт движения
|
```
|
||||||
- `artifact_other` — прочие артефакты
|
|
||||||
- `position_error` — ошибка позиционирования
|
Для инференса только на CPU (образ меньше, без CUDA-колёс):
|
||||||
- `rotation` — нарушение ротации
|
|
||||||
- `incomplete_view` — неполный вид
|
```bash
|
||||||
- `roi_error` — ошибка ROI
|
pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
- `labeling_error` — ошибка разметки
|
```
|
||||||
- 🌐 **REST API** — интеграция с внешними системами
|
|
||||||
- 📊 **Веб-интерфейс** — загрузка и визуализация результатов
|
Обучение и инференс можно вызывать напрямую:
|
||||||
- 📈 **Экспорт** — выгрузка результатов в XLSX
|
|
||||||
|
```bash
|
||||||
|
python -m src.dxa.train --epochs 100 --output-dir models
|
||||||
|
python -m src.dxa.inference --input-path dataset_hack --output-path results.xlsx --zip-out masks.zip
|
||||||
|
```
|
||||||
|
|
||||||
|
После запуска сервера:
|
||||||
|
- Веб-интерфейс: http://localhost:8000
|
||||||
|
- Swagger UI: http://localhost:8000/docs
|
||||||
|
|
||||||
|
## Docker
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker build -t dxa-quality .
|
||||||
|
docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
||||||
|
```
|
||||||
|
|
||||||
|
Чекпоинт должен лежать в `models/dxa_model.pth` до сборки; путь задаётся переменной
|
||||||
|
`DXA_MODEL_PATH` (по умолчанию `/app/models/dxa_model.pth`). Сборка проверяет, что
|
||||||
|
чекпоинт читается, и падает, если модели нет — вместо тихих 500-х ответов в рантайме.
|
||||||
|
Вес модели внутрь образа зашит, из сети ничего не скачивается.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Архитектура
|
## Архитектура
|
||||||
|
|
||||||
```
|
```
|
||||||
┌─────────────────────────────────────────────────────────────┐
|
DICOM ──▶ предобработка ──▶ ResNet18 (заморожен) ──▶ линейная голова ──▶ логит
|
||||||
│ FastAPI Server │
|
│ │
|
||||||
│ (port 8000) │
|
│ └──▶ голова области (вспомогательная)
|
||||||
├─────────────────────────────────────────────────────────────┤
|
│
|
||||||
│ /api/v1/analyze → Basic quality prediction │
|
├──▶ геометрия яркой зоны: область, ROI, геометрия кадра
|
||||||
│ /api/v1/analyze/detailed → Full report with metrics │
|
└──▶ эвристики: резкость, «плотные» включения
|
||||||
│ /api/v1/analyze/sr → DICOM SR (Structured Report) │
|
|
||||||
│ /api/v1/batch → Batch processing │
|
|
||||||
│ /api/v1/export → XLSX export │
|
|
||||||
├─────────────────────────────────────────────────────────────┤
|
|
||||||
│ │
|
|
||||||
│ ┌──────────────┐ ┌─────────────────┐ │
|
|
||||||
│ │ ResNet18 │───▶│ Quality Model │ │
|
|
||||||
│ │ (pretrained) │ │ (binary class) │ │
|
|
||||||
│ └──────────────┘ └────────┬────────┘ │
|
|
||||||
│ │ │
|
|
||||||
│ ▼ │
|
|
||||||
│ ┌─────────────────────┐ │
|
|
||||||
│ │ Detailed Assessment │ │
|
|
||||||
│ │ - Motion detection │ │
|
|
||||||
│ │ - Artifact detection│ │
|
|
||||||
│ │ - ROI validation │ │
|
|
||||||
│ │ - View completeness│ │
|
|
||||||
│ └─────────────────────┘ │
|
|
||||||
└─────────────────────────────────────────────────────────────┘
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Ключевые решения и почему они такие:
|
||||||
|
|
||||||
|
1. **Линейный зонд вместо полного fine-tune.** Уникальных снимков в наборе ~250.
|
||||||
|
Полный fine-tune ResNet18 переобучается за несколько эпох (train F1 → 1.0 при
|
||||||
|
val AUC ≈ 0.5). Замороженный backbone + линейная голова удерживает val AUC
|
||||||
|
≈ 0.7–0.85. Режим `--head mlp --freeze-epochs 0` оставлен для экспериментов
|
||||||
|
на большем объёме данных.
|
||||||
|
|
||||||
|
2. **Метки из имён файлов.** Суффикс `_good`/`_bad` — экспертная оценка снимка;
|
||||||
|
отсутствие суффикса означает «изображение хорошее». Приоритет:
|
||||||
|
`_bad` > `_good` > нет метки.
|
||||||
|
|
||||||
|
3. **Склейка побайтных дублей.** В датасете 544 файла, но 252 уникальных снимка:
|
||||||
|
один и тот же кадр сохранён многократно под разными именами (часть — с меткой,
|
||||||
|
часть — без). Без склейки одно изображение попадало бы в оба класса.
|
||||||
|
|
||||||
|
4. **Разбиение по исследованиям.** Снимки одного исследования не попадают
|
||||||
|
одновременно в train и val — иначе метрики завышаются за счёт утечки.
|
||||||
|
|
||||||
|
5. **Аугментация отключена.** Проверено экспериментально: яркостный разброс и сдвиг
|
||||||
|
кадра снижают AUC с 0.87 до 0.56, потому что распределение яркости и положение
|
||||||
|
области сами являются признаками качества. Флаг `--augment` включает её для
|
||||||
|
экспериментов.
|
||||||
|
|
||||||
|
6. **Порог по логиту.** При доле нарушений ~15 % порог 0.5 даёт нулевой recall.
|
||||||
|
Порог подбирается по F1 на валидации и сохраняется в чекпоинт; решение
|
||||||
|
принимается по логиту (численно устойчиво при насыщении вероятностей).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Быстрый старт
|
## Формат выходных данных
|
||||||
|
|
||||||
### Требования
|
Основные столбцы соответствуют требованиям задания:
|
||||||
|
|
||||||
- Python 3.10+
|
| Столбец | Описание |
|
||||||
- PyTorch 2.0+
|
|---|---|
|
||||||
- 4GB+ RAM
|
| `path_to_study` | Путь к исследованию (для HTTP-загрузки — `upload://<имя>`) |
|
||||||
- (опционально) GPU CUDA/MPS для ускорения
|
| `study_uid` | StudyInstanceUID |
|
||||||
|
| `image_uid` | SOPInstanceUID |
|
||||||
|
| `anatomical_region` | `spine` / `hip_left` / `hip_right` / `hip` |
|
||||||
|
| `quality_class` | 0 — качественное, 1 — есть нарушение |
|
||||||
|
| `violation_type` | Тип нарушения или пустая строка |
|
||||||
|
| `processing_status` | `Success` или `Failure: <причина>` |
|
||||||
|
| `time_of_processing` | Время обработки, секунды |
|
||||||
|
|
||||||
### Установка
|
Дополнительно добавляются `confidence`, `violation_reason`, `region_confidence`
|
||||||
|
— они не мешают автоматическому разбору обязательных столбцов.
|
||||||
|
|
||||||
```bash
|
## API
|
||||||
# Клонирование
|
|
||||||
git clone https://github.com/yourusername/bone-quality-assessment.git
|
|
||||||
cd bone-quality-assessment
|
|
||||||
|
|
||||||
# Создание виртуального окружения
|
| Метод | Путь | Назначение |
|
||||||
python -m venv venv
|
|---|---|---|
|
||||||
source venv/bin/activate # Linux/Mac
|
| GET | `/` | Веб-интерфейс |
|
||||||
# venv\Scripts\activate # Windows
|
| GET | `/api/v1/health` | Статус и признак загрузки модели |
|
||||||
|
| POST | `/api/v1/analyze` | Базовый анализ одного файла |
|
||||||
# Установка зависимостей
|
| POST | `/api/v1/analyze/detailed` | Расширенный отчёт, опционально маска |
|
||||||
pip install -r requirements.txt
|
| POST | `/api/v1/analyze/sr` | Текстовое представление отчёта DICOM SR |
|
||||||
|
|
||||||
# Загрузка модели (опционально)
|
|
||||||
# Поместите файл модели в models/dxa_model.pth
|
|
||||||
```
|
|
||||||
|
|
||||||
### Запуск сервера
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Локальный запуск
|
|
||||||
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
|
||||||
|
|
||||||
# Или через run.py
|
|
||||||
python run.py
|
|
||||||
```
|
|
||||||
|
|
||||||
После запуска:
|
|
||||||
- Web-интерфейс: http://localhost:8000
|
|
||||||
- Swagger UI: http://localhost:8000/docs
|
|
||||||
- ReDoc: http://localhost:8000/redoc
|
|
||||||
|
|
||||||
### Docker
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Сборка
|
|
||||||
docker build -t dxa-quality-api .
|
|
||||||
|
|
||||||
# Запуск
|
|
||||||
docker run -p 8000:8000 dxa-quality-api
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## API Endpoints
|
|
||||||
|
|
||||||
| Метод | Эндпоинт | Описание |
|
|
||||||
|-------|----------|----------|
|
|
||||||
| GET | `/` | Главная страница (веб-интерфейс) |
|
|
||||||
| GET | `/api/v1/health` | Проверка статуса сервиса |
|
|
||||||
| POST | `/api/v1/analyze` | Базовый анализ изображения |
|
|
||||||
| POST | `/api/v1/analyze/detailed` | Детальный анализ с метриками |
|
|
||||||
| POST | `/api/v1/analyze/sr` | DICOM SR отчёт |
|
|
||||||
| POST | `/api/v1/batch` | Пакетный анализ |
|
| POST | `/api/v1/batch` | Пакетный анализ |
|
||||||
| POST | `/api/v1/export` | Анализ и экспорт в XLSX |
|
| POST | `/api/v1/export` | Пакетный анализ и выгрузка в XLSX |
|
||||||
|
|
||||||
### Пример использования
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Анализ файла
|
curl -X POST http://localhost:8000/api/v1/analyze -F "file=@study/spine.dcm"
|
||||||
curl -X POST "http://localhost:8000/api/v1/analyze" \
|
|
||||||
-H "accept: application/json" \
|
|
||||||
-H "Content-Type: multipart/form-data" \
|
|
||||||
-F "file=@/path/to/image.dcm"
|
|
||||||
|
|
||||||
# Детальный анализ
|
|
||||||
curl -X POST "http://localhost:8000/api/v1/analyze/detailed" \
|
|
||||||
-H "accept: application/json" \
|
|
||||||
-H "Content-Type: multipart/form-data" \
|
|
||||||
-F "file=@/path/to/image.dcm"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Ответ детального анализа
|
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
|
"study_uid": "1.2.643...",
|
||||||
|
"image_uid": "1.2.643...",
|
||||||
"anatomical_region": "spine",
|
"anatomical_region": "spine",
|
||||||
"quality_class": 1,
|
"quality_class": 1,
|
||||||
"quality_label": "Violation detected",
|
"quality_label": "Violation detected",
|
||||||
"violation_type": "artifact_motion",
|
"violation_type": "quality_violation_detected",
|
||||||
"reason": "Обнаружен артефакт движения (размытие)",
|
"reason": "Выявлено нарушение качества изображения",
|
||||||
"confidence": 0.85,
|
"confidence": 0.72,
|
||||||
"confidence_per_class": {
|
"threshold_probability": 0.6154,
|
||||||
"correct": 0.15,
|
"processing_status": "Success"
|
||||||
"violation": 0.85
|
|
||||||
},
|
|
||||||
"view_quality": "full",
|
|
||||||
"metrics": {
|
|
||||||
"motion": {
|
|
||||||
"motion_detected": true,
|
|
||||||
"blur_laplacian": 0.0008,
|
|
||||||
"severity": "HIGH"
|
|
||||||
},
|
|
||||||
"artifacts": {
|
|
||||||
"any_detected": false,
|
|
||||||
"metal_detected": false
|
|
||||||
},
|
|
||||||
"roi_check": {
|
|
||||||
"valid": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"overall_quality": "POOR",
|
|
||||||
"severity": "HIGH"
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Метрики
|
||||||
|
|
||||||
|
Метрики зависят от выбранного разбиения по исследованиям, поэтому приводятся
|
||||||
|
с разбросом. Оценка на валидационной части (19 исследований, 51 снимок,
|
||||||
|
8 нарушений), разбиение по исследованиям:
|
||||||
|
|
||||||
|
| Что измерено | Значение | Как измерено |
|
||||||
|
|---|---|---|
|
||||||
|
| ROC-AUC, 5 разбиений | **0.76 ± 0.08** (0.64 – 0.84) | обучение по seed 0..4, порог по F1 |
|
||||||
|
| PR-AUC, 5 разбиений | 0.48 ± 0.17 | там же; базовый уровень при 15 % нарушений — 0.15 |
|
||||||
|
| F1, 5 разбиений | 0.54 ± 0.11 | там же (порог подобран на той же валидации — смещено вверх) |
|
||||||
|
| ROC-AUC, 5-фолдовая CV | **0.81 ± 0.08** | линейный зонд на тех же признаках, разбиение по исследованиям |
|
||||||
|
| Контрольная задача «позвоночник / бедро» | AUC 1.00 | проверка работоспособности пайплайна |
|
||||||
|
| Перестановка меток (нулевая гипотеза) | AUC 0.64 | вклад случайных корреляций |
|
||||||
|
|
||||||
|
Метрики по областям — в `models/train_report.md`, он создаётся при обучении.
|
||||||
|
Разбивка важна, потому что нарушения распределены крайне неравномерно: в
|
||||||
|
позвоночнике ~29 % снимков с нарушением против ~4–5 % у бёдер, а область почти
|
||||||
|
однозначно определяется по ширине кадра. Поэтому общий AUC частично отражает
|
||||||
|
различение области, а не только распознавание дефекта.
|
||||||
|
|
||||||
|
Время обработки одного снимка — порядка 0.02–0.05 с на CPU (ResNet18 с
|
||||||
|
замороженным backbone), то есть требование «не более 3 минут на исследование»
|
||||||
|
выполняется с большим запасом.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Ограничения (важно для интерпретации)
|
||||||
|
|
||||||
|
1. **Разметка исходных данных — на уровне исследования, а не снимка.**
|
||||||
|
В наборе один снимок помечен `_bad`, остальные снимки того же исследования
|
||||||
|
не размечены. Метка снимка считается унаследованной от исследования, поэтому
|
||||||
|
часть меток заведомо шумная.
|
||||||
|
2. **Мало данных.** 252 уникальных снимка, 37 нарушений. Доверительные интервалы
|
||||||
|
широкие; оценка на закрытом наборе может отличаться.
|
||||||
|
3. **Тип нарушения определяется эвристиками, а не обученной моделью.** Для
|
||||||
|
честного мультикласса нужна разметка типов на уровне снимка.
|
||||||
|
4. **Область определяется по размеру кадра.** Признак безошибочно работает на этом
|
||||||
|
оборудовании (99/99 для позвоночника), но при смене аппарата порог
|
||||||
|
`SPINE_MIN_WIDTH` потребует калибровки.
|
||||||
|
5. **В DICOM нет разметки ROI.** Ни overlay, ни graphic annotation в файлах нет,
|
||||||
|
поэтому корректность нанесённых областей измерения нельзя проверить прямым
|
||||||
|
сравнением — оценивается только геометрия видимой зоны.
|
||||||
|
|
||||||
|
## План доработки
|
||||||
|
|
||||||
|
- Разметить типы нарушений на уровне снимка и обучить мультилейбл-классификатор.
|
||||||
|
- Собрать 500+ исследований для устойчивых метрик и честной валидации.
|
||||||
|
- Подключить Grad-CAM для объяснения решения (модуль есть, но не интегрирован).
|
||||||
|
- Заменить порог по ширине кадра на калибровку по метаданным аппарата.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## Структура проекта
|
## Структура проекта
|
||||||
|
|
||||||
```
|
```
|
||||||
bone_2026/
|
bone_2026/
|
||||||
├── src/
|
├── src/
|
||||||
│ ├── main.py # FastAPI приложение
|
│ ├── main.py # FastAPI: маршруты и загрузка модели
|
||||||
│ ├── run.py # Запуск сервера
|
│ ├── dxa/ # действующий модуль оценки качества
|
||||||
│ ├── dxa/ # DXA модуль
|
│ │ ├── labels.py # разбор имён, метки, склейка дублей, сплит
|
||||||
│ │ ├── model.py # ResNet18 классификатор
|
│ │ ├── preprocess.py # DICOM -> тензор (общий для обучения и API)
|
||||||
│ │ ├── dataset.py # Загрузчик данных
|
│ │ ├── dataset.py # Dataset и DataLoader
|
||||||
│ │ ├── train.py # Обучение модели
|
│ │ ├── model.py # сеть, метрики, подбор порога
|
||||||
│ │ └── inference.py # Инференс и batch-обработка
|
│ │ ├── train.py # обучение и отчёт
|
||||||
│ ├── quality/ # Оценка качества
|
│ │ └── inference.py # пакетный инференс, определение области
|
||||||
│ │ ├── quality_scorer.py # Базовый скорer
|
│ ├── quality/ # эвристики (частично используются API)
|
||||||
│ │ └── detailed_assessment.py # Детальный анализ
|
│ ├── api/static/ # веб-интерфейс
|
||||||
│ ├── api/ # REST API
|
│ └── model/, core/, pipeline/ # устаревшие модули, не подключены к API
|
||||||
│ │ ├── endpoints.py # Дополнительные эндпоинты
|
├── models/dxa_model.pth # чекпоинт (+ train_report.md)
|
||||||
│ │ └── static/ # Веб-интерфейс
|
├── tests/ # pytest: метки, сплит, метрики, модель
|
||||||
│ └── utils/ # Утилиты
|
├── dataset_hack/ # данные (в git не хранятся)
|
||||||
├── models/ # Обученные модели
|
|
||||||
│ └── dxa_model.pth # Модель классификатора
|
|
||||||
├── dataset_hack/ # Датасет для обучения/тестирования
|
|
||||||
├── docs/ # Документация
|
|
||||||
├── Dockerfile
|
├── Dockerfile
|
||||||
├── requirements.txt
|
├── requirements.txt
|
||||||
└── README.md
|
└── run.sh
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
## Запуск обучения
|
||||||
|
|
||||||
## Обучение модели
|
|
||||||
|
|
||||||
### Подготовка данных
|
|
||||||
|
|
||||||
1. Разместите DICOM-файлы в `dataset_hack/НД_для_обучения/Исследования/`
|
|
||||||
2. Подготовьте Excel-файл разметки `dataset_hack/НД_для_обучения/разметка.xlsx`
|
|
||||||
|
|
||||||
Столбцы разметки:
|
|
||||||
- `study_uid` — ID исследования
|
|
||||||
- `позвоночник_укладка`, `позвоночник_ось`, `позвоночник_артефакты` — критерии для позвоночника
|
|
||||||
- `бедро_позиция_лев`, `бедро_roi_лев` — критерии для левого бедра
|
|
||||||
- `бедро_позиция_прав`, `бедро_roi_прав` — критерии для правого бедра
|
|
||||||
- `итог_позвоночник`, `итог_бедро_лев`, `итог_бедро_прав` — итоговая оценка (0/1)
|
|
||||||
|
|
||||||
### Запуск обучения
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python src/dxa/train.py \
|
python -m src.dxa.train --epochs 100 --output-dir models
|
||||||
--epochs 20 \
|
|
||||||
--batch-size 8 \
|
|
||||||
--backbone resnet18 \
|
|
||||||
--output-dir models
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Аргументы
|
|
||||||
|
|
||||||
| Параметр | По умолчанию | Описание |
|
| Параметр | По умолчанию | Описание |
|
||||||
|----------|-------------|----------|
|
|---|---|---|
|
||||||
| `--data-root` | `dataset_hack` | Путь к директории с данными |
|
| `--data-root` | `dataset_hack` | Каталог датасета |
|
||||||
| `--annotation-path` | `dataset_hack/НД_для_обучения/разметка.xlsx` | Путь к файлу разметки |
|
| `--annotation-path` | `dataset_hack/НД_для_обучения/разметка.xlsx` | Excel с разметкой (только отчёт о расхождениях) |
|
||||||
| `--epochs` | 20 | Количество эпох |
|
| `--backbone` | `resnet18` | `resnet18` / `resnet34` |
|
||||||
| `--batch-size` | 8 | Размер батча |
|
| `--head` | `linear` | `linear` (линейный зонд) / `mlp` |
|
||||||
| `--backbone` | `resnet18` | Архитектура (resnet18/resnet34/efficientnet_b0) |
|
| `--freeze-epochs` | `-1` | `-1` — backbone заморожен всегда; `0` — обучать всю сеть |
|
||||||
| `--input-size` | 224 | Размер входного изображения |
|
| `--epochs`, `--batch-size`, `--learning-rate`, `--weight-decay` | 100 / 16 / 3e-4 / 5e-2 | Оптимизация |
|
||||||
| `--output-dir` | `models` | Директория для сохранения модели |
|
| `--balance` | `none` | `loss` / `sampler` для компенсации дисбаланса |
|
||||||
|
| `--val-fraction`, `--seed` | 0.2 / 42 | Разбиение по исследованиям |
|
||||||
|
| `--augment` | выключено | Включает яркостную аугментацию (ухудшает метрики, см. п. 5) |
|
||||||
|
| `--output-dir` | `models` | Куда сохранять чекпоинт и отчёты |
|
||||||
|
| `--dry-run` | — | Проверить разбор данных и разбиение без обучения |
|
||||||
|
|
||||||
---
|
После обучения в `--output-dir` появляются `dxa_model.pth`, `train_report.md`
|
||||||
|
и `train_report.json`; отчёт удобно приложить к презентации.
|
||||||
## Инференс
|
|
||||||
|
|
||||||
### Одиночный файл
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/dxa/inference.py \
|
|
||||||
--input-path path/to/image.dcm \
|
|
||||||
--output-path result.xlsx
|
|
||||||
```
|
|
||||||
|
|
||||||
### Директория
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/dxa/inference.py \
|
|
||||||
--input-path dataset_hack/Для\ теста \
|
|
||||||
--output-path results.xlsx \
|
|
||||||
--model-path models/dxa_model.pth
|
|
||||||
```
|
|
||||||
|
|
||||||
### Выходной формат (XLSX/CSV)
|
|
||||||
|
|
||||||
| Колонка | Описание |
|
|
||||||
|---------|----------|
|
|
||||||
| `path_to_study` | Путь к директории исследования |
|
|
||||||
| `study_uid` | StudyInstanceUID |
|
|
||||||
| `image_uid` | SOPInstanceUID |
|
|
||||||
| `anatomical_region` | Анатомическая область (spine/hip_left/hip_right) |
|
|
||||||
| `quality_class` | Класс качества (0 — OK, 1 — Violation) |
|
|
||||||
| `violation_type` | Тип нарушения |
|
|
||||||
| `processing_status` | Статус обработки |
|
|
||||||
| `time_of_processing` | Время обработки (сек) |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Метрики качества
|
|
||||||
|
|
||||||
### Детекция движения
|
|
||||||
|
|
||||||
- **Laplacian variance** — дисперсия лапласиана (меньше = сильнее размытие)
|
|
||||||
- **FFT high-frequency ratio** — отношение высокочастотной энергии (меньше = размытие)
|
|
||||||
- **Edge duplication** — проверка "призрачных" контуров
|
|
||||||
|
|
||||||
### Детекция артефактов
|
|
||||||
|
|
||||||
- **Metal detection** — яркие области (>99.5 перцентиль)
|
|
||||||
- **Implant detection** — линейные структуры (морфологические операции)
|
|
||||||
- **Cement detection** — локальные яркие пятна в ROI
|
|
||||||
- **Calcification** — малые яркие области вне ROI
|
|
||||||
|
|
||||||
### Полнота изображения
|
|
||||||
|
|
||||||
- **Spine**: подсчет позвонков (ожидается 3-4), проверка межпозвоночных промежутков
|
|
||||||
- **Hip**: проверка видимости шейки бедра, большого/малого вертелов
|
|
||||||
|
|
||||||
### Валидация ROI
|
|
||||||
|
|
||||||
- Проверка отступа от краев (>10 пикселей)
|
|
||||||
- Проверка размера ROI (>30% высоты, >20% ширины изображения)
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Лицензия
|
|
||||||
|
|
||||||
MIT License
|
|
||||||
|
|
||||||
## Команда
|
## Команда
|
||||||
|
|
||||||
- **Грачев Денис** — Разработка
|
- **Грачев Денис** — разработка
|
||||||
- **Грачев Татьяна** — Капитан
|
- **Грачев Татьяна** — капитан
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<sub>Built for Bone Quality Assessment Hackathon 2026</sub>
|
<sub>Built for Bone Quality Assessment Hackathon 2026</sub>
|
||||||
|
|
|
||||||
115
requirements.txt
115
requirements.txt
|
|
@ -1,73 +1,52 @@
|
||||||
annotated-doc==0.0.5
|
# Зависимости проекта с зафиксированными версиями.
|
||||||
annotated-types==0.7.0
|
#
|
||||||
anyio==4.12.1
|
# Установка:
|
||||||
attrs==26.1.0
|
# pip install -r requirements.txt
|
||||||
certifi==2026.7.22
|
# Для инференса на CPU (меньше образ, без CUDA-колёс):
|
||||||
click==8.1.8
|
# pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
contourpy==1.3.0
|
#
|
||||||
cycler==0.12.1
|
# Состав разбит на две части:
|
||||||
et_xmlfile==2.0.0
|
# 1) рабочие зависимости — используются приложением и обучением;
|
||||||
exceptiongroup==1.3.1
|
# 2) устаревшие модули — код старых сегментаторов, не подключённый к
|
||||||
|
# действующему API (src/model/segmentator.py, src/quality/medical_quality.py,
|
||||||
|
# src/core/orchestrator.py). Для контейнера инференса они не нужны:
|
||||||
|
# эти модули не импортируются ни из src/main.py, ни из src/dxa/*.
|
||||||
|
|
||||||
|
# --- Веб-слой ---
|
||||||
fastapi==0.128.8
|
fastapi==0.128.8
|
||||||
filelock==3.19.1
|
|
||||||
fonttools==4.60.2
|
|
||||||
fsspec==2025.10.0
|
|
||||||
h11==0.16.0
|
|
||||||
hf-xet==1.6.0
|
|
||||||
httpcore==1.0.9
|
|
||||||
httpx==0.28.1
|
|
||||||
huggingface_hub==1.8.0
|
|
||||||
idna==3.18
|
|
||||||
importlib_resources==6.5.2
|
|
||||||
Jinja2==3.1.6
|
|
||||||
joblib==1.5.3
|
|
||||||
jsonschema==3.2.0
|
|
||||||
kiwisolver==1.4.7
|
|
||||||
markdown-it-py==3.0.0
|
|
||||||
MarkupSafe==3.0.3
|
|
||||||
matplotlib==3.9.4
|
|
||||||
mdurl==0.1.2
|
|
||||||
monai==1.5.2
|
|
||||||
mpmath==1.3.0
|
|
||||||
networkx==3.2.1
|
|
||||||
nibabel==5.3.3
|
|
||||||
numpy==1.26.4
|
|
||||||
opencv-python-headless==4.11.0.86
|
|
||||||
openpyxl==3.1.5
|
|
||||||
packaging==26.3
|
|
||||||
pandas==2.3.3
|
|
||||||
pillow==11.3.0
|
|
||||||
pydantic==2.13.4
|
|
||||||
pydantic_core==2.46.4
|
|
||||||
pydicom==2.4.4
|
|
||||||
pydicom-seg==0.4.1
|
|
||||||
Pygments==2.21.0
|
|
||||||
pyparsing==3.3.2
|
|
||||||
pyrsistent==0.20.0
|
|
||||||
python-dateutil==2.9.0.post0
|
|
||||||
python-multipart==0.0.20
|
|
||||||
pytz==2026.3.post1
|
|
||||||
PyYAML==6.0.3
|
|
||||||
rich==15.0.0
|
|
||||||
safetensors==0.7.0
|
|
||||||
scikit-learn==1.6.1
|
|
||||||
scipy==1.13.1
|
|
||||||
shellingham==1.5.4
|
|
||||||
simpleitk==2.5.6
|
|
||||||
six==1.17.0
|
|
||||||
starlette==0.49.3
|
starlette==0.49.3
|
||||||
sympy==1.14.0
|
uvicorn==0.39.0
|
||||||
threadpoolctl==3.7.0
|
python-multipart==0.0.20
|
||||||
timm==1.0.29
|
pydantic==2.13.4
|
||||||
|
|
||||||
|
# --- Глубокое обучение ---
|
||||||
torch==2.8.0
|
torch==2.8.0
|
||||||
torchvision==0.23.0
|
torchvision==0.23.0
|
||||||
TotalSegmentator==2.18.0
|
|
||||||
requests==2.32.3
|
# --- DICOM и изображения ---
|
||||||
|
pydicom==2.4.4
|
||||||
|
Pillow==11.3.0
|
||||||
|
opencv-python-headless==4.11.0.86
|
||||||
|
|
||||||
|
# --- Данные, метрики, таблицы ---
|
||||||
|
numpy==1.26.4
|
||||||
|
pandas==2.3.3
|
||||||
|
openpyxl==3.1.5
|
||||||
|
scikit-learn==1.6.1
|
||||||
|
scipy==1.13.1
|
||||||
|
|
||||||
|
# --- Утилиты ---
|
||||||
tqdm==4.70.0
|
tqdm==4.70.0
|
||||||
typer==0.23.2
|
|
||||||
typing-inspection==0.4.2
|
# --- Тесты ---
|
||||||
typing_extensions==4.16.0
|
pytest==8.4.2
|
||||||
tzdata==2026.3
|
|
||||||
unicorn==2.1.4
|
# --- Устаревшие модули (не используются действующим API) ---
|
||||||
uvicorn==0.39.0
|
# TotalSegmentator при первом запуске скачивает собственные веса из сети,
|
||||||
zipp==3.23.1
|
# поэтому в офлайн-контейнер инференса его включать не следует.
|
||||||
|
# monai==1.5.2
|
||||||
|
# TotalSegmentator==2.18.0
|
||||||
|
# nibabel==5.3.3
|
||||||
|
# simpleitk==2.5.6
|
||||||
|
# pydicom-seg==0.4.1
|
||||||
|
# timm==1.0.29
|
||||||
|
|
|
||||||
196
run.sh
196
run.sh
|
|
@ -1,134 +1,90 @@
|
||||||
#!/bin/bash
|
#!/usr/bin/env bash
|
||||||
# DXA Quality Assessment - Main entry point
|
#
|
||||||
# Usage: bash run.sh [command] [args...]
|
# DXA Quality Assessment — единая точка входа.
|
||||||
|
#
|
||||||
|
# Использование: ./run.sh <команда> [аргументы]
|
||||||
|
# ./run.sh train [доп. аргументы для src.dxa.train]
|
||||||
|
# ./run.sh infer <вход> <выход.xlsx|.csv> [доп. аргументы]
|
||||||
|
# ./run.sh serve [порт]
|
||||||
|
# ./run.sh test
|
||||||
|
#
|
||||||
|
# Скрипт намеренно не содержит своей реализации обучения и инференса: вся
|
||||||
|
# логика живёт в Python-модулях, а здесь только вызовы. Дублирование цикла
|
||||||
|
# обучения (как было раньше) приводило к расхождению правил между режимами.
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
set -e
|
GREEN='\033[0;32m'; YELLOW='\033[1;33m'; RED='\033[0;31m'; NC='\033[0m'
|
||||||
|
|
||||||
# Colors for output
|
PYTHON=${PYTHON:-python3}
|
||||||
RED='\033[0;31m'
|
|
||||||
GREEN='\033[0;32m'
|
|
||||||
YELLOW='\033[1;33m'
|
|
||||||
NC='\033[0m' # No Color
|
|
||||||
|
|
||||||
echo -e "${GREEN}=== DXA Quality Assessment ===${NC}"
|
|
||||||
|
|
||||||
# Default values
|
|
||||||
COMMAND=${1:-help}
|
|
||||||
DATA_ROOT=${DATA_ROOT:-dataset_hack}
|
DATA_ROOT=${DATA_ROOT:-dataset_hack}
|
||||||
ANNOTATION_PATH=${ANNOTATION_PATH:-dataset_hack/НД_для_обучения/разметка.xlsx}
|
ANNOTATION_PATH=${ANNOTATION_PATH:-"dataset_hack/НД_для_обучения/разметка.xlsx"}
|
||||||
MODEL_PATH=${MODEL_PATH:-models/dxa_model.pth}
|
MODEL_PATH=${MODEL_PATH:-models/dxa_model.pth}
|
||||||
EPOCHS=${EPOCHS:-10}
|
PORT=${PORT:-8000}
|
||||||
BATCH_SIZE=${BATCH_SIZE:-16}
|
|
||||||
|
|
||||||
case "$COMMAND" in
|
command=${1:-help}
|
||||||
|
shift || true
|
||||||
|
|
||||||
|
case "$command" in
|
||||||
train)
|
train)
|
||||||
echo -e "${YELLOW}Training model...${NC}"
|
echo -e "${YELLOW}Обучение классификатора качества DXA${NC}"
|
||||||
|
"$PYTHON" -m src.dxa.train \
|
||||||
python3 -c "
|
--data-root "$DATA_ROOT" \
|
||||||
import sys
|
--annotation-path "$ANNOTATION_PATH" \
|
||||||
sys.path.insert(0, '.')
|
--output-dir "$(dirname "$MODEL_PATH")" \
|
||||||
|
"$@"
|
||||||
import torch
|
echo -e "${GREEN}Готово. Чекпоинт: ${MODEL_PATH}${NC}"
|
||||||
import numpy as np
|
|
||||||
from src.dxa.dataset import create_dataloaders
|
|
||||||
from src.dxa.model import create_model
|
|
||||||
|
|
||||||
device = 'mps' if torch.backends.mps.is_available() else 'cpu'
|
|
||||||
print(f'Device: {device}')
|
|
||||||
|
|
||||||
train_loader, val_loader = create_dataloaders(
|
|
||||||
data_root='${DATA_ROOT}',
|
|
||||||
annotation_path='${ANNOTATION_PATH}',
|
|
||||||
batch_size=${BATCH_SIZE},
|
|
||||||
input_size=224,
|
|
||||||
num_workers=0
|
|
||||||
)
|
|
||||||
|
|
||||||
print(f'Train: {len(train_loader.dataset)}, Val: {len(val_loader.dataset)}')
|
|
||||||
|
|
||||||
model = create_model(backbone='resnet18', pretrained=True, device=device)
|
|
||||||
|
|
||||||
best_f1 = 0
|
|
||||||
for epoch in range(${EPOCHS}):
|
|
||||||
train_loss, train_acc = model.train_epoch(train_loader)
|
|
||||||
val_loss, val_acc = model.validate(val_loader)
|
|
||||||
|
|
||||||
# Compute F1
|
|
||||||
model.model.eval()
|
|
||||||
preds, labels = [], []
|
|
||||||
with torch.no_grad():
|
|
||||||
for images, labs in val_loader:
|
|
||||||
outputs = model.model(images.to(device))
|
|
||||||
preds.extend(outputs.argmax(dim=1).cpu().numpy())
|
|
||||||
labels.extend(labs['label'].cpu().numpy())
|
|
||||||
|
|
||||||
preds, labels = np.array(preds), np.array(labels)
|
|
||||||
tp = ((preds == 1) & (labels == 1)).sum()
|
|
||||||
fp = ((preds == 1) & (labels == 0)).sum()
|
|
||||||
fn = ((preds == 0) & (labels == 1)).sum()
|
|
||||||
precision = tp/(tp+fp) if (tp+fp)>0 else 0
|
|
||||||
recall = tp/(tp+fn) if (tp+fn)>0 else 0
|
|
||||||
f1 = 2*precision*recall/(precision+recall) if (precision+recall)>0 else 0
|
|
||||||
|
|
||||||
print(f'Epoch {epoch+1}: Train={train_acc:.3f}, Val={val_acc:.3f}, F1={f1:.3f}')
|
|
||||||
|
|
||||||
if f1 > best_f1:
|
|
||||||
best_f1 = f1
|
|
||||||
model.save('${MODEL_PATH}')
|
|
||||||
|
|
||||||
print(f'Best F1: {best_f1:.3f}')
|
|
||||||
"
|
|
||||||
echo -e "${GREEN}Training complete! Model saved to ${MODEL_PATH}${NC}"
|
|
||||||
;;
|
;;
|
||||||
|
|
||||||
infer)
|
infer)
|
||||||
INPUT_PATH=${2:-dataset_hack/Для теста}
|
input_path=${1:-dataset_hack/Для теста}
|
||||||
OUTPUT_PATH=${3:-results.xlsx}
|
output_path=${2:-results.xlsx}
|
||||||
|
shift 2 2>/dev/null || true
|
||||||
echo -e "${YELLOW}Running inference...${NC}"
|
echo -e "${YELLOW}Пакетная обработка${NC}"
|
||||||
echo "Input: ${INPUT_PATH}"
|
echo " вход: $input_path"
|
||||||
echo "Output: ${OUTPUT_PATH}"
|
echo " выход: $output_path"
|
||||||
|
"$PYTHON" -m src.dxa.inference \
|
||||||
python3 -c "
|
--input-path "$input_path" \
|
||||||
import sys
|
--output-path "$output_path" \
|
||||||
sys.path.insert(0, '.')
|
--model-path "$MODEL_PATH" \
|
||||||
from src.dxa.inference import process_dicom_files
|
"$@"
|
||||||
import argparse
|
echo -e "${GREEN}Готово: ${output_path}${NC}"
|
||||||
|
;;
|
||||||
|
|
||||||
process_dicom_files(argparse.Namespace(
|
|
||||||
input_path='${INPUT_PATH}',
|
|
||||||
output_path='${OUTPUT_PATH}',
|
|
||||||
model_path='${MODEL_PATH}',
|
|
||||||
backbone='resnet18',
|
|
||||||
input_size=224
|
|
||||||
))
|
|
||||||
"
|
|
||||||
echo -e "${GREEN}Inference complete!${NC}"
|
|
||||||
;;
|
|
||||||
|
|
||||||
serve)
|
serve)
|
||||||
echo -e "${YELLOW}Starting API server...${NC}"
|
echo -e "${YELLOW}Запуск API на порту ${PORT}${NC}"
|
||||||
python3 -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
DXA_MODEL_PATH="$MODEL_PATH" \
|
||||||
|
"$PYTHON" -m uvicorn src.main:app --host 0.0.0.0 --port "$PORT"
|
||||||
;;
|
;;
|
||||||
|
|
||||||
|
test)
|
||||||
|
echo -e "${YELLOW}Запуск тестов${NC}"
|
||||||
|
"$PYTHON" -m pytest tests/ -q
|
||||||
|
;;
|
||||||
|
|
||||||
help|*)
|
help|*)
|
||||||
echo "Usage: $0 [command] [options]"
|
cat <<'USAGE'
|
||||||
echo ""
|
DXA Quality Assessment — управление запуском
|
||||||
echo "Commands:"
|
|
||||||
echo " train Train the model"
|
Команды:
|
||||||
echo " infer <input> <output> Run inference"
|
train [аргументы] Обучить классификатор качества.
|
||||||
echo " serve Start API server"
|
Пример: ./run.sh train --epochs 100 --head mlp
|
||||||
echo ""
|
infer <вход> <выход> [арг.] Пакетная обработка DICOM (файл или каталог).
|
||||||
echo "Environment variables:"
|
Пример: ./run.sh infer "dataset_hack/Для теста" results.xlsx
|
||||||
echo " DATA_ROOT Data directory (default: dataset_hack)"
|
Дополнительно: --zip-out masks.zip
|
||||||
echo " ANNOTATION_PATH Annotation Excel file"
|
serve [порт] Запустить HTTP API и веб-интерфейс.
|
||||||
echo " MODEL_PATH Model output path"
|
test Запустить тесты.
|
||||||
echo " EPOCHS Training epochs (default: 10)"
|
|
||||||
echo " BATCH_SIZE Batch size (default: 16)"
|
Переменные окружения:
|
||||||
echo ""
|
DATA_ROOT Каталог датасета (по умолчанию dataset_hack)
|
||||||
echo "Examples:"
|
ANNOTATION_PATH Excel с разметкой (только для отчёта о расхождениях)
|
||||||
echo " $0 train"
|
MODEL_PATH Путь к чекпоинту (по умолчанию models/dxa_model.pth)
|
||||||
echo " EPOCHS=50 $0 train"
|
PORT Порт API (по умолчанию 8000)
|
||||||
echo " $0 infer dataset_hack/Для теста results.xlsx"
|
PYTHON Интерпретатор (по умолчанию python3)
|
||||||
|
|
||||||
|
Типовой порядок работы:
|
||||||
|
./run.sh train # обучить модель
|
||||||
|
./run.sh infer dataset_hack results.xlsx
|
||||||
|
./run.sh serve # веб-интерфейс на http://localhost:8000
|
||||||
|
USAGE
|
||||||
;;
|
;;
|
||||||
esac
|
esac
|
||||||
|
|
|
||||||
|
|
@ -1,363 +1,205 @@
|
||||||
"""
|
"""
|
||||||
Dataset for DXA (bone densitometry) quality assessment
|
Датасет DXA для обучения классификатора качества.
|
||||||
|
|
||||||
|
Метки формируются из имён DICOM-файлов (см. `src.dxa.labels`), где суффикс
|
||||||
|
`_good`/`_bad` кодирует экспертную оценку; отсутствие суффикса — «хорошее»
|
||||||
|
изображение. Одинаковые по содержимому файлы склеиваются в один пример.
|
||||||
|
|
||||||
|
Разбиение на train/val выполняется по исследованиям, чтобы снимки одного
|
||||||
|
исследования не попадали одновременно в обучение и валидацию.
|
||||||
"""
|
"""
|
||||||
import os
|
from __future__ import annotations
|
||||||
import pandas as pd
|
|
||||||
import numpy as np
|
import logging
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Tuple, Optional
|
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
||||||
import pydicom
|
|
||||||
from PIL import Image
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import DataLoader, Dataset
|
||||||
|
|
||||||
|
from src.dxa.labels import (
|
||||||
|
REGIONS,
|
||||||
|
ImageRecord,
|
||||||
|
format_summary,
|
||||||
|
scan_dataset,
|
||||||
|
stratified_group_split,
|
||||||
|
)
|
||||||
|
from src.dxa.preprocess import PreprocessConfig, preprocess_dicom, with_input_size
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _augment(img: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Мягкая аугментация (опционально, по умолчанию ВЫКЛЮЧЕНА).
|
||||||
|
|
||||||
|
Отключена намеренно: на этом датасете она разрушает обучающий сигнал.
|
||||||
|
Измерено на замороженных признаках ImageNet: с аугментацией AUC падает с
|
||||||
|
0.87 до 0.56 (seed 0), с 0.84 до 0.50 (seed 4). Причина в том, что
|
||||||
|
признаки качества здесь — это и есть распределение яркости и положение
|
||||||
|
области: артефакты, размытие и смещение укладки проявляются именно через
|
||||||
|
них. Яркостный разброс ±10 % и сдвиг кадра на 4 % затирают ровно ту
|
||||||
|
информацию, которую модель должна выучить.
|
||||||
|
|
||||||
|
Геометрические отражения и повороты также не применяются: направление
|
||||||
|
ротации бедра и отклонение оси позвоночника сами являются критериями
|
||||||
|
качества, поэтому такие преобразования искажали бы метку.
|
||||||
|
"""
|
||||||
|
# Яркость/контраст
|
||||||
|
gain = float(rng.uniform(0.9, 1.1))
|
||||||
|
bias = float(rng.uniform(-0.05, 0.05))
|
||||||
|
img = np.clip(img * gain + bias, 0.0, 1.0)
|
||||||
|
|
||||||
|
# Небольшой сдвиг кадра
|
||||||
|
shift_x = int(round(rng.uniform(-0.04, 0.04) * img.shape[2]))
|
||||||
|
shift_y = int(round(rng.uniform(-0.04, 0.04) * img.shape[1]))
|
||||||
|
if shift_x or shift_y:
|
||||||
|
img = np.roll(img, (shift_y, shift_x), axis=(1, 2))
|
||||||
|
if shift_y > 0:
|
||||||
|
img[:, :shift_y, :] = 0
|
||||||
|
elif shift_y < 0:
|
||||||
|
img[:, shift_y:, :] = 0
|
||||||
|
if shift_x > 0:
|
||||||
|
img[:, :, :shift_x] = 0
|
||||||
|
elif shift_x < 0:
|
||||||
|
img[:, :, shift_x:] = 0
|
||||||
|
return img
|
||||||
|
|
||||||
|
|
||||||
class DXADataset(Dataset):
|
class DXADataset(Dataset):
|
||||||
"""Dataset for DXA bone densitometry images"""
|
"""Датасет DXA: изображение -> бинарная метка качества (0 — годное, 1 — нарушение)."""
|
||||||
|
|
||||||
# Mapping from anatomical regions to column names in annotation
|
|
||||||
ANATOMICAL_MAPPING = {
|
|
||||||
'spine': {
|
|
||||||
'columns': ['позвоночник_укладка', 'позвоночник_ось', 'позвоночник_артефакты'],
|
|
||||||
'total_column': 'итог_позвоночник'
|
|
||||||
},
|
|
||||||
'hip_right': {
|
|
||||||
'columns': ['бедро_позиция_прав', 'бедро_roi_прав'],
|
|
||||||
'total_column': 'итог_бедро_прав'
|
|
||||||
},
|
|
||||||
'hip_left': {
|
|
||||||
'columns': ['бедро_позиция_лев', 'бедро_roi_лев'],
|
|
||||||
'total_column': 'итог_бедро_лев'
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
def __init__(self,
|
|
||||||
data_root: str,
|
|
||||||
annotation_path: str,
|
|
||||||
transform=None,
|
|
||||||
input_size: Tuple[int, int] = (224, 224),
|
|
||||||
mode: str = 'train'):
|
|
||||||
"""
|
|
||||||
Args:
|
|
||||||
data_root: Path to folder with DICOM studies
|
|
||||||
annotation_path: Path to Excel annotation file
|
|
||||||
transform: Optional transforms
|
|
||||||
input_size: Target image size
|
|
||||||
mode: 'train' or 'val'
|
|
||||||
"""
|
|
||||||
self.data_root = Path(data_root)
|
|
||||||
self.annotation_path = annotation_path
|
|
||||||
self.transform = transform
|
|
||||||
self.input_size = input_size
|
|
||||||
self.mode = mode
|
|
||||||
|
|
||||||
# Load annotation
|
|
||||||
self.annotation = self._load_annotation()
|
|
||||||
|
|
||||||
# Build dataset
|
|
||||||
self.samples = self._build_samples()
|
|
||||||
|
|
||||||
# Filter samples based on mode
|
|
||||||
if mode == 'train':
|
|
||||||
self.samples = self.samples[:int(len(self.samples) * 0.8)]
|
|
||||||
else:
|
|
||||||
self.samples = self.samples[int(len(self.samples) * 0.8):]
|
|
||||||
|
|
||||||
def _load_annotation(self) -> pd.DataFrame:
|
|
||||||
"""Load and parse annotation Excel file"""
|
|
||||||
df = pd.read_excel(self.annotation_path, header=None)
|
|
||||||
|
|
||||||
# Skip header rows
|
|
||||||
data = df.iloc[2:].copy()
|
|
||||||
data.columns = range(len(df.columns))
|
|
||||||
|
|
||||||
# Rename columns
|
|
||||||
data = data.rename(columns={
|
|
||||||
0: 'id',
|
|
||||||
1: 'study_uid',
|
|
||||||
2: 'позвоночник_укладка',
|
|
||||||
3: 'позвоночник_ось',
|
|
||||||
4: 'позвоночник_артефакты',
|
|
||||||
5: 'бедро_позиция_прав',
|
|
||||||
6: 'бедро_roi_прав',
|
|
||||||
7: 'бедро_позиция_лев',
|
|
||||||
8: 'бедро_roi_лев',
|
|
||||||
9: 'итог_позвоночник',
|
|
||||||
10: 'итог_бедро_прав',
|
|
||||||
11: 'итог_бедро_лев',
|
|
||||||
12: 'комментарий',
|
|
||||||
14: 'общий_позвоночник',
|
|
||||||
15: 'общий_бедро_прав',
|
|
||||||
16: 'общий_бедро_лев',
|
|
||||||
17: 'класс',
|
|
||||||
18: 'балл'
|
|
||||||
})
|
|
||||||
|
|
||||||
# Remove empty rows
|
|
||||||
data = data.dropna(subset=['study_uid'])
|
|
||||||
|
|
||||||
# Convert numeric columns
|
|
||||||
numeric_cols = ['позвоночник_укладка', 'позвоночник_ось', 'позвоночник_артефакты',
|
|
||||||
'итог_позвоночник', 'итог_бедро_прав', 'итог_бедро_лев']
|
|
||||||
for col in numeric_cols:
|
|
||||||
if col in data.columns:
|
|
||||||
data[col] = pd.to_numeric(data[col], errors='coerce')
|
|
||||||
|
|
||||||
return data
|
|
||||||
|
|
||||||
def _extract_anatomical_region_from_filename(self, filename: str) -> Optional[str]:
|
|
||||||
"""
|
|
||||||
Extract anatomical region from DICOM filename (fallback method).
|
|
||||||
Prefer _determine_region_from_image() for actual classification.
|
|
||||||
"""
|
|
||||||
filename_lower = filename.lower()
|
|
||||||
|
|
||||||
if 'spine' in filename_lower:
|
|
||||||
return 'spine'
|
|
||||||
elif 'l_hip' in filename_lower or 'left_hip' in filename_lower:
|
|
||||||
return 'hip_left'
|
|
||||||
elif 'r_hip' in filename_lower or 'right_hip' in filename_lower:
|
|
||||||
return 'hip_right'
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
def _determine_region_from_image(self, dcm_path: str) -> str:
|
def __init__(
|
||||||
"""
|
self,
|
||||||
Determine anatomical region from DICOM image content.
|
records: Sequence[ImageRecord],
|
||||||
Uses the same algorithm as inference.py
|
preprocess: Optional[PreprocessConfig] = None,
|
||||||
"""
|
train: bool = False,
|
||||||
import pydicom
|
seed: int = 0,
|
||||||
import numpy as np
|
augment: bool = False,
|
||||||
|
):
|
||||||
try:
|
self.records = list(records)
|
||||||
ds = pydicom.dcmread(dcm_path)
|
self.preprocess = preprocess or PreprocessConfig()
|
||||||
img = ds.pixel_array.astype(np.float32)
|
self.train = train
|
||||||
|
self.augment = augment
|
||||||
h, w = img.shape
|
self.seed = seed
|
||||||
img_norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
|
||||||
|
|
||||||
# Feature 1: Bright region aspect ratio
|
|
||||||
threshold = np.percentile(img_norm, 95)
|
|
||||||
binary = img_norm > threshold
|
|
||||||
|
|
||||||
bbox_aspect = 1.0
|
|
||||||
left_right_ratio = 1.0
|
|
||||||
if binary.sum() > 0:
|
|
||||||
try:
|
|
||||||
from scipy import ndimage
|
|
||||||
rows = np.any(binary, axis=1)
|
|
||||||
cols = np.any(binary, axis=0)
|
|
||||||
if rows.any() and cols.any():
|
|
||||||
rmin, rmax = np.where(rows)[0][[0, -1]]
|
|
||||||
cmin, cmax = np.where(cols)[0][[0, -1]]
|
|
||||||
bbox_h = rmax - rmin
|
|
||||||
bbox_w = cmax - cmin
|
|
||||||
bbox_aspect = bbox_h / (bbox_w + 1e-6)
|
|
||||||
|
|
||||||
left_bright = binary[:, :w//2].sum()
|
|
||||||
right_bright = binary[:, w//2:].sum()
|
|
||||||
left_right_ratio = left_bright / (right_bright + 1e-6)
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Feature 2: Symmetry
|
|
||||||
h_mid, w_mid = h // 2, w // 2
|
|
||||||
left_half = img_norm[:, :w_mid]
|
|
||||||
right_half = np.fliplr(img_norm[:, w_mid:])
|
|
||||||
min_w = min(left_half.shape[1], right_half.shape[1])
|
|
||||||
symmetry = 1 - np.abs(left_half[:, :min_w] - right_half[:, :min_w]).mean() / (img_norm.std() + 1e-6)
|
|
||||||
|
|
||||||
# Classification
|
|
||||||
if bbox_aspect < 1.5:
|
|
||||||
return 'spine'
|
|
||||||
elif bbox_aspect < 1.8:
|
|
||||||
if symmetry > 0.35:
|
|
||||||
return 'spine'
|
|
||||||
else:
|
|
||||||
return 'hip'
|
|
||||||
else:
|
|
||||||
if left_right_ratio > 1.3:
|
|
||||||
return 'hip_right'
|
|
||||||
elif left_right_ratio < 0.7:
|
|
||||||
return 'hip_left'
|
|
||||||
else:
|
|
||||||
return 'hip'
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
return 'unknown'
|
|
||||||
|
|
||||||
def _build_samples(self) -> List[Dict]:
|
|
||||||
"""Build list of samples from annotation and DICOM files"""
|
|
||||||
samples = []
|
|
||||||
|
|
||||||
for _, row in self.annotation.iterrows():
|
|
||||||
study_uid = str(row['study_uid']).strip()
|
|
||||||
|
|
||||||
# Try different path structures
|
|
||||||
possible_paths = [
|
|
||||||
self.data_root / 'Исследования' / study_uid,
|
|
||||||
self.data_root / 'НД_для_обучения' / 'Исследования' / study_uid,
|
|
||||||
]
|
|
||||||
|
|
||||||
study_path = None
|
|
||||||
for p in possible_paths:
|
|
||||||
if p.exists():
|
|
||||||
study_path = p
|
|
||||||
break
|
|
||||||
|
|
||||||
if study_path is None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Find all DICOM files
|
|
||||||
dcm_files = sorted(study_path.rglob('*.dcm'))
|
|
||||||
|
|
||||||
for dcm_file in dcm_files:
|
|
||||||
# Extract anatomical region from filename
|
|
||||||
region = self._extract_anatomical_region_from_filename(dcm_file.name)
|
|
||||||
|
|
||||||
if region is None:
|
|
||||||
# Fallback: skip files without region in name
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Determine which annotation column to use based on region
|
|
||||||
if region == 'spine':
|
|
||||||
total_col = 'итог_позвоночник'
|
|
||||||
elif region == 'hip_left':
|
|
||||||
total_col = 'итог_бедро_лев'
|
|
||||||
elif region == 'hip_right':
|
|
||||||
total_col = 'итог_бедро_прав'
|
|
||||||
else:
|
|
||||||
continue
|
|
||||||
|
|
||||||
if total_col in row and pd.notna(row[total_col]):
|
|
||||||
# Get quality label (0 = good, 1 = violation)
|
|
||||||
quality = int(row[total_col])
|
|
||||||
|
|
||||||
# Get specific violation criteria based on region
|
|
||||||
criteria = {}
|
|
||||||
if region == 'spine':
|
|
||||||
for col in ['позвоночник_укладка', 'позвоночник_ось', 'позвоночник_артефакты']:
|
|
||||||
if col in row and pd.notna(row[col]):
|
|
||||||
criteria[col] = int(row[col])
|
|
||||||
elif region == 'hip_left':
|
|
||||||
for col in ['бедро_позиция_лев', 'бедро_roi_лев']:
|
|
||||||
if col in row and pd.notna(row[col]):
|
|
||||||
criteria[col] = int(row[col])
|
|
||||||
elif region == 'hip_right':
|
|
||||||
for col in ['бедро_позиция_прав', 'бедро_roi_прав']:
|
|
||||||
if col in row and pd.notna(row[col]):
|
|
||||||
criteria[col] = int(row[col])
|
|
||||||
|
|
||||||
samples.append({
|
|
||||||
'dcm_path': str(dcm_file),
|
|
||||||
'study_uid': study_uid,
|
|
||||||
'anatomical_region': region,
|
|
||||||
'quality': quality,
|
|
||||||
'criteria': criteria,
|
|
||||||
'comment': row.get('комментарий', '')
|
|
||||||
})
|
|
||||||
|
|
||||||
return samples
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
def __len__(self) -> int:
|
||||||
return len(self.samples)
|
return len(self.records)
|
||||||
|
|
||||||
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict]:
|
|
||||||
"""Get one sample"""
|
|
||||||
sample = self.samples[idx]
|
|
||||||
|
|
||||||
# Load DICOM
|
@property
|
||||||
ds = pydicom.dcmread(sample['dcm_path'])
|
def labels(self) -> List[int]:
|
||||||
img = ds.pixel_array.astype(np.float32)
|
return [r.label for r in self.records]
|
||||||
|
|
||||||
# Normalize to 0-1
|
def region_index(self, region: Optional[str]) -> int:
|
||||||
img_min = img.min()
|
"""Индекс области для эмбеддинга (0 — неизвестная область)."""
|
||||||
img_max = img.max()
|
return REGIONS.index(region) + 1 if region in REGIONS else 0
|
||||||
if img_max > img_min:
|
|
||||||
img = (img - img_min) / (img_max - img_min)
|
|
||||||
|
|
||||||
# Convert to 3-channel for pretrained models
|
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict[str, Any]]:
|
||||||
img = np.stack([img] * 3, axis=0)
|
rec = self.records[idx]
|
||||||
|
img = preprocess_dicom(rec.path, self.preprocess)
|
||||||
|
|
||||||
# Convert to uint8 for PIL
|
if self.train and self.augment:
|
||||||
img = (img * 255).astype(np.uint8)
|
rng = np.random.default_rng(self.seed + idx)
|
||||||
|
img = _augment(img, rng)
|
||||||
|
|
||||||
# Resize
|
tensor = torch.from_numpy(img).float()
|
||||||
img_pil = Image.fromarray(img.transpose(1, 2, 0))
|
target = torch.tensor(rec.label, dtype=torch.long)
|
||||||
if isinstance(self.input_size, tuple):
|
|
||||||
target_size = self.input_size
|
|
||||||
else:
|
|
||||||
target_size = (self.input_size, self.input_size)
|
|
||||||
img_pil = img_pil.resize(target_size, Image.BILINEAR)
|
|
||||||
img = np.array(img_pil).transpose(2, 0, 1)
|
|
||||||
|
|
||||||
# Normalize back to 0-1 for model
|
return tensor, {
|
||||||
img = img.astype(np.float32) / 255.0
|
"label": target,
|
||||||
|
"region_id": torch.tensor(self.region_index(rec.region), dtype=torch.long),
|
||||||
# Apply transforms
|
"study_uid": rec.study,
|
||||||
if self.transform:
|
"anatomical_region": rec.region or "unknown",
|
||||||
img = self.transform(img)
|
"dcm_path": str(rec.path),
|
||||||
|
|
||||||
# Convert to tensor
|
|
||||||
img = torch.from_numpy(img).float()
|
|
||||||
|
|
||||||
# Label
|
|
||||||
label = sample['quality']
|
|
||||||
|
|
||||||
return img, {
|
|
||||||
'label': label,
|
|
||||||
'study_uid': sample['study_uid'],
|
|
||||||
'anatomical_region': sample['anatomical_region'],
|
|
||||||
'dcm_path': sample['dcm_path']
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def create_dataloaders(data_root: str,
|
def build_records(
|
||||||
annotation_path: str,
|
data_root: str | Path,
|
||||||
batch_size: int = 8,
|
annotation_path: Optional[str | Path] = None,
|
||||||
input_size: Tuple[int, int] = (224, 224),
|
dedup: bool = True,
|
||||||
num_workers: int = 4):
|
) -> List[ImageRecord]:
|
||||||
"""Create train and validation dataloaders"""
|
"""Найти и разметить все уникальные снимки датасета."""
|
||||||
|
records = scan_dataset(data_root, with_pixel_dedup=dedup, annotation_path=annotation_path)
|
||||||
train_dataset = DXADataset(
|
logger.info("Dataset scan complete:\n%s", format_summary(records))
|
||||||
|
return records
|
||||||
|
|
||||||
|
|
||||||
|
def make_datasets(
|
||||||
|
data_root: str | Path,
|
||||||
|
annotation_path: Optional[str | Path] = None,
|
||||||
|
input_size: int = 224,
|
||||||
|
val_fraction: float = 0.2,
|
||||||
|
seed: int = 42,
|
||||||
|
dedup: bool = True,
|
||||||
|
preprocess: Optional[PreprocessConfig] = None,
|
||||||
|
) -> Tuple[DXADataset, DXADataset, PreprocessConfig]:
|
||||||
|
"""Собрать train/val датасеты с разбиением по исследованиям."""
|
||||||
|
cfg = with_input_size(preprocess or PreprocessConfig(), input_size)
|
||||||
|
records = build_records(data_root, annotation_path, dedup=dedup)
|
||||||
|
train_records, val_records = stratified_group_split(records, val_fraction=val_fraction, seed=seed)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Split: train=%d images / %d studies, val=%d images / %d studies",
|
||||||
|
len(train_records),
|
||||||
|
len({r.study for r in train_records}),
|
||||||
|
len(val_records),
|
||||||
|
len({r.study for r in val_records}),
|
||||||
|
)
|
||||||
|
|
||||||
|
train_ds = DXADataset(train_records, preprocess=cfg, train=True, seed=seed)
|
||||||
|
val_ds = DXADataset(val_records, preprocess=cfg, train=False, seed=seed)
|
||||||
|
return train_ds, val_ds, cfg
|
||||||
|
|
||||||
|
|
||||||
|
def create_dataloaders(
|
||||||
|
data_root: str | Path,
|
||||||
|
annotation_path: Optional[str | Path] = None,
|
||||||
|
batch_size: int = 8,
|
||||||
|
input_size: int = 224,
|
||||||
|
num_workers: int = 0,
|
||||||
|
val_fraction: float = 0.2,
|
||||||
|
seed: int = 42,
|
||||||
|
preprocess: Optional[PreprocessConfig] = None,
|
||||||
|
) -> Tuple[DataLoader, DataLoader, PreprocessConfig]:
|
||||||
|
"""
|
||||||
|
Создать train/val DataLoader.
|
||||||
|
|
||||||
|
Ранее эта функция делила датасет срезом списка, из-за чего снимки одного
|
||||||
|
исследования попадали в обе части. Теперь разбиение выполняется по
|
||||||
|
исследованиям внутри `make_datasets`.
|
||||||
|
"""
|
||||||
|
train_ds, val_ds, cfg = make_datasets(
|
||||||
data_root=data_root,
|
data_root=data_root,
|
||||||
annotation_path=annotation_path,
|
annotation_path=annotation_path,
|
||||||
input_size=input_size,
|
input_size=input_size,
|
||||||
mode='train'
|
val_fraction=val_fraction,
|
||||||
|
seed=seed,
|
||||||
|
preprocess=preprocess,
|
||||||
)
|
)
|
||||||
|
|
||||||
val_dataset = DXADataset(
|
pin = torch.cuda.is_available()
|
||||||
data_root=data_root,
|
train_loader = DataLoader(
|
||||||
annotation_path=annotation_path,
|
train_ds, batch_size=batch_size, shuffle=True,
|
||||||
input_size=input_size,
|
num_workers=num_workers, pin_memory=pin, drop_last=False,
|
||||||
mode='val'
|
|
||||||
)
|
)
|
||||||
|
val_loader = DataLoader(
|
||||||
train_loader = torch.utils.data.DataLoader(
|
val_ds, batch_size=batch_size, shuffle=False,
|
||||||
train_dataset,
|
num_workers=num_workers, pin_memory=pin, drop_last=False,
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=True,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=True
|
|
||||||
)
|
)
|
||||||
|
return train_loader, val_loader, cfg
|
||||||
val_loader = torch.utils.data.DataLoader(
|
|
||||||
val_dataset,
|
|
||||||
batch_size=batch_size,
|
|
||||||
shuffle=False,
|
|
||||||
num_workers=num_workers,
|
|
||||||
pin_memory=True
|
|
||||||
)
|
|
||||||
|
|
||||||
return train_loader, val_loader
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == "__main__":
|
||||||
# Test
|
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||||
train_loader, val_loader = create_dataloaders(
|
train_loader, val_loader, cfg = create_dataloaders(
|
||||||
data_root='dataset_hack',
|
data_root="dataset_hack",
|
||||||
annotation_path='dataset_hack/НД_для_обучения/разметка.xlsx'
|
annotation_path="dataset_hack/НД_для_обучения/разметка.xlsx",
|
||||||
|
batch_size=4,
|
||||||
)
|
)
|
||||||
print(f'Train samples: {len(train_loader.dataset)}')
|
print(f"Preprocess: {cfg.to_dict()}")
|
||||||
print(f'Val samples: {len(val_loader.dataset)}')
|
print(f"Train batches: {len(train_loader)}, Val batches: {len(val_loader)}")
|
||||||
|
images, meta = next(iter(train_loader))
|
||||||
|
print(f"Batch shape: {tuple(images.shape)}, labels: {meta['label'].tolist()}")
|
||||||
|
print(f"Regions: {meta['anatomical_region']}")
|
||||||
|
|
|
||||||
|
|
@ -1,458 +1,578 @@
|
||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""
|
"""
|
||||||
DXA Quality Inference - Batch processing with Excel output
|
Пакетный инференс классификатора качества DXA
|
||||||
|
=============================================
|
||||||
|
|
||||||
Этот модуль выполняет инференс модели классификации качества DXA исследований.
|
Обрабатывает DICOM-исследования и формирует таблицу в формате требований:
|
||||||
Основные функции:
|
`path_to_study, study_uid, image_uid, anatomical_region, quality_class,
|
||||||
- Загрузка и предобработка DICOM изображений
|
violation_type, processing_status, time_of_processing`.
|
||||||
- Определение анатомической области (позвоночник/бедро)
|
|
||||||
- Бинарная классификация качества (OK/Violation)
|
Анатомическая область определяется по содержимому изображения
|
||||||
- Пакетная обработка с экспортом в Excel
|
(`determine_region_from_image`); имя файла не используется, чтобы не зависеть
|
||||||
|
от соглашения о именах на закрытых данных. Дополнительно обученная
|
||||||
|
вспомогательная голова предсказывает область, что при низкой уверенности
|
||||||
|
основного метода служит уточнением.
|
||||||
|
|
||||||
|
Решение принимается по логиту: порог хранится в чекпоинте и подобран по F1 на
|
||||||
|
валидации при обучении. Если чекпоинт старый и порога не содержит, берётся 0.
|
||||||
|
|
||||||
Использование:
|
Использование:
|
||||||
python src/dxa/inference.py --input-path <path> --output-path <output.xlsx>
|
python -m src.dxa.inference --input-path <файл или каталог> --output-path results.xlsx
|
||||||
|
python -m src.dxa.inference --input-path dataset_hack --output-path report.csv --zip-out masks.zip
|
||||||
"""
|
"""
|
||||||
import os
|
from __future__ import annotations
|
||||||
import sys
|
|
||||||
import argparse
|
import argparse
|
||||||
from pathlib import Path
|
import logging
|
||||||
from datetime import datetime
|
import sys
|
||||||
import warnings
|
import warnings
|
||||||
warnings.filterwarnings('ignore')
|
import zipfile
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
|
warnings.filterwarnings("ignore")
|
||||||
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import pydicom
|
import pydicom
|
||||||
|
import torch
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
# Add src to path
|
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
|
||||||
|
from src.dxa.labels import REGIONS, iter_dicom_files
|
||||||
|
from src.dxa.model import DXAQualityModel, create_model
|
||||||
|
from src.dxa.preprocess import (
|
||||||
|
PreprocessConfig,
|
||||||
|
load_dicom_array,
|
||||||
|
preprocess_from_array,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger("dxa.inference")
|
||||||
|
|
||||||
|
# Предсказание вспомогательной головы -> анатомическая область
|
||||||
|
REGION_BY_ID = {1: "spine", 2: "hip_right", 3: "hip_left"}
|
||||||
|
|
||||||
|
# Человекочитаемые причины для типа нарушения
|
||||||
|
REASON_BY_TYPE = {
|
||||||
|
"position_error": "Геометрия или укладка области исследования нарушены",
|
||||||
|
"artifact_motion": "Признаки артефактов движения (размытие, раздвоение контуров)",
|
||||||
|
"artifact_other": "Посторонние включения или артефакты в зоне интереса",
|
||||||
|
"incomplete_view": "Нужная анатомическая область видна не полностью",
|
||||||
|
"roi_error": "Границы области интереса не совпадают с анатомическими",
|
||||||
|
"rotation": "Выраженная ротация, искажающая анатомические границы",
|
||||||
|
"quality_violation_detected": "Выявлено нарушение качества изображения",
|
||||||
|
}
|
||||||
|
|
||||||
|
OUTPUT_COLUMNS = [
|
||||||
|
"path_to_study", "study_uid", "image_uid", "anatomical_region",
|
||||||
|
"quality_class", "violation_type", "processing_status", "time_of_processing",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def get_device():
|
@dataclass
|
||||||
"""
|
class HeuristicSignals:
|
||||||
Определение доступного устройства для вычислений.
|
"""Дешёвые признаки изображения для определения анатомической области."""
|
||||||
|
|
||||||
Порядок приоритета: MPS (Apple Silicon) -> CUDA (NVIDIA GPU) -> CPU.
|
bbox_aspect: float = 1.0
|
||||||
Это нужно для максимальной производительности на доступном железе.
|
symmetry: float = 1.0
|
||||||
|
left_right_ratio: float = 1.0
|
||||||
Returns:
|
width: int = 0
|
||||||
str: Устройство ('mps', 'cuda' или 'cpu')
|
height: int = 0
|
||||||
"""
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ImageResult:
|
||||||
|
"""Результат обработки одного изображения."""
|
||||||
|
|
||||||
|
path_to_study: str
|
||||||
|
study_uid: str
|
||||||
|
image_uid: str
|
||||||
|
anatomical_region: str
|
||||||
|
quality_class: int
|
||||||
|
violation_type: str
|
||||||
|
processing_status: str
|
||||||
|
time_of_processing: float
|
||||||
|
confidence: float = 0.0
|
||||||
|
violation_reason: str = ""
|
||||||
|
region_confidence: float = 0.0
|
||||||
|
dcm_path: str = ""
|
||||||
|
metrics: Dict = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Prediction:
|
||||||
|
"""Предсказание по одному изображению, независимое от источника (файл/поток)."""
|
||||||
|
|
||||||
|
quality_class: int
|
||||||
|
prob: float
|
||||||
|
logit: float
|
||||||
|
threshold: float
|
||||||
|
anatomical_region: str
|
||||||
|
region_confidence: float
|
||||||
|
violation_type: str
|
||||||
|
violation_reason: str
|
||||||
|
samples: Dict[str, float]
|
||||||
|
|
||||||
|
|
||||||
|
def get_device(prefer: Optional[str] = None) -> str:
|
||||||
|
"""Выбрать устройство: явно заданное или лучшее из доступных."""
|
||||||
|
if prefer:
|
||||||
|
return prefer
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
return "cuda"
|
||||||
if torch.backends.mps.is_available():
|
if torch.backends.mps.is_available():
|
||||||
return 'mps'
|
return "mps"
|
||||||
elif torch.cuda.is_available():
|
return "cpu"
|
||||||
return 'cuda'
|
|
||||||
else:
|
|
||||||
return 'cpu'
|
|
||||||
|
|
||||||
|
|
||||||
def load_model(model_path: str, backbone: str = 'resnet18', device: str = 'cpu'):
|
@dataclass
|
||||||
|
class LoadedCheckpoint:
|
||||||
|
"""Загруженный чекпоинт: модель, предобработка, порог и метаданные."""
|
||||||
|
|
||||||
|
model: DXAQualityModel
|
||||||
|
preprocess: PreprocessConfig
|
||||||
|
threshold: float
|
||||||
|
metadata: Dict
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
model_path: str,
|
||||||
|
backbone: Optional[str] = None,
|
||||||
|
head: Optional[str] = None,
|
||||||
|
device: str = "cpu",
|
||||||
|
input_size: Optional[int] = None,
|
||||||
|
) -> LoadedCheckpoint:
|
||||||
"""
|
"""
|
||||||
Загрузка обученной модели классификатора качества DXA.
|
Загрузить модель и восстановить параметры её обработки из чекпоинта.
|
||||||
|
|
||||||
Модель использует предобученный ResNet18 в качестве backbone и
|
Архитектура (backbone, тип головы) и параметры предобработки берутся из
|
||||||
добавляет классификационную голову для бинарной классификации
|
самого чекпоинта, поэтому вызывающей стороне не нужно их дублировать и
|
||||||
(качество OK vs Violation).
|
невозможно рассинхронизировать обучение и инференс. Явно переданные
|
||||||
|
`backbone`/`head` проверяются на совместимость с сохранёнными.
|
||||||
Args:
|
|
||||||
model_path: Путь к файлу модели (.pth)
|
|
||||||
backbone: Архитектура backbone (resnet18/resnet34/efficientnet_b0)
|
|
||||||
device: Устройство для загрузки модели
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
DXAQualityModel: Обертка модели с методами predict и load
|
LoadedCheckpoint с моделью, конфигом предобработки, порогом и метаданными.
|
||||||
"""
|
"""
|
||||||
from src.dxa.model import create_model
|
# Сначала читаем метаданные лёгким способом, чтобы не создавать лишние сети.
|
||||||
|
try:
|
||||||
|
meta_only = torch.load(model_path, map_location="cpu", weights_only=False)
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"Cannot read checkpoint {model_path}: {exc}") from exc
|
||||||
|
|
||||||
model = create_model(backbone=backbone, pretrained=False, device=device)
|
saved_backbone = (meta_only or {}).get("backbone", "resnet18")
|
||||||
model.load(model_path)
|
saved_head = (meta_only or {}).get("head", "mlp")
|
||||||
|
if backbone and backbone != saved_backbone:
|
||||||
|
raise ValueError(
|
||||||
|
f"Checkpoint was trained with backbone={saved_backbone!r}, but {backbone!r} was requested"
|
||||||
|
)
|
||||||
|
if head and head != saved_head:
|
||||||
|
raise ValueError(
|
||||||
|
f"Checkpoint was trained with head={saved_head!r}, but {head!r} was requested"
|
||||||
|
)
|
||||||
|
|
||||||
|
model = create_model(backbone=saved_backbone, head=saved_head, pretrained=False, device=device)
|
||||||
|
metadata = model.load(model_path)
|
||||||
model.model.eval()
|
model.model.eval()
|
||||||
return model
|
|
||||||
|
cfg = load_preprocess(metadata, input_size)
|
||||||
|
threshold = float(metadata.get("threshold_logit", metadata.get("threshold", 0.0)))
|
||||||
|
return LoadedCheckpoint(model=model, preprocess=cfg, threshold=threshold, metadata=metadata)
|
||||||
|
|
||||||
|
|
||||||
def load_dicom_image(dcm_path: str, input_size: int = 224) -> torch.Tensor:
|
def load_preprocess(metadata: Dict, input_size: Optional[int] = None) -> PreprocessConfig:
|
||||||
|
"""Восстановить параметры предобработки из чекпоинта."""
|
||||||
|
cfg = PreprocessConfig.from_dict(metadata.get("preprocess") or {})
|
||||||
|
if input_size:
|
||||||
|
cfg = PreprocessConfig.from_dict({**cfg.to_dict(), "input_size": input_size})
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def heuristic_signals(img: np.ndarray) -> HeuristicSignals:
|
||||||
"""
|
"""
|
||||||
Загрузка и предобработка DICOM изображения для модели.
|
Геометрические признаки изображения.
|
||||||
|
|
||||||
Этапы предобработки:
|
Порог яркой области берётся по 95-му перцентилю. Размеры кадра входят в
|
||||||
1. Чтение DICOM и извлечение pixel_array
|
набор сигналов, потому что у аппарата позвоночные и бедренные снимки имеют
|
||||||
2. Нормализация интенсивности в диапазон [0, 1]
|
разную ширину кадра (300 против 280 пикселей), и это самый надёжный признак
|
||||||
3. Преобразование в 3 канала (дублирование для RGB)
|
области на данном оборудовании.
|
||||||
4. Изменение размера до input_size x input_size
|
|
||||||
5. Нормализация для ImageNet (деление на 255)
|
|
||||||
6. Преобразование в PyTorch тензор
|
|
||||||
|
|
||||||
Args:
|
|
||||||
dcm_path: Путь к DICOM файлу
|
|
||||||
input_size: Целевой размер изображения (по умолчанию 224 для ResNet)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: Тензор изображения формы (1, 3, 224, 224)
|
|
||||||
"""
|
"""
|
||||||
ds = pydicom.dcmread(dcm_path)
|
norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
||||||
img = ds.pixel_array.astype(np.float32)
|
h, w = norm.shape
|
||||||
|
binary = norm > np.percentile(norm, 95)
|
||||||
|
if not binary.any():
|
||||||
|
return HeuristicSignals(width=w, height=h)
|
||||||
|
|
||||||
# Normalize to 0-1
|
rows, cols = np.any(binary, axis=1), np.any(binary, axis=0)
|
||||||
img_min = img.min()
|
rmin, rmax = np.where(rows)[0][[0, -1]]
|
||||||
img_max = img.max()
|
cmin, cmax = np.where(cols)[0][[0, -1]]
|
||||||
if img_max > img_min:
|
bbox_aspect = (rmax - rmin) / ((cmax - cmin) + 1e-6)
|
||||||
img = (img - img_min) / (img_max - img_min)
|
|
||||||
|
|
||||||
# Convert to 3-channel
|
left = binary[:, :w // 2].sum()
|
||||||
img = np.stack([img] * 3, axis=0)
|
right = binary[:, w // 2:].sum()
|
||||||
|
ratio = left / (right + 1e-6)
|
||||||
|
|
||||||
# Convert to uint8 for PIL
|
left_half = norm[:, :w // 2]
|
||||||
img = (img * 255).astype(np.uint8)
|
right_half = np.fliplr(norm[:, w // 2:])
|
||||||
|
m = min(left_half.shape[1], right_half.shape[1])
|
||||||
|
symmetry = 1 - np.abs(left_half[:, :m] - right_half[:, :m]).mean() / (norm.std() + 1e-6)
|
||||||
|
|
||||||
# Resize
|
return HeuristicSignals(float(bbox_aspect), float(symmetry), float(ratio), int(w), int(h))
|
||||||
img_pil = Image.fromarray(img.transpose(1, 2, 0))
|
|
||||||
img_pil = img_pil.resize((input_size, input_size), Image.BILINEAR)
|
|
||||||
img = np.array(img_pil).transpose(2, 0, 1)
|
|
||||||
|
|
||||||
# Normalize back to 0-1
|
|
||||||
img = img.astype(np.float32) / 255.0
|
|
||||||
|
|
||||||
# Convert to tensor
|
|
||||||
img = torch.from_numpy(img).float().unsqueeze(0)
|
|
||||||
|
|
||||||
return img
|
|
||||||
|
|
||||||
|
|
||||||
def extract_anatomical_region_from_filename(dcm_path: str) -> str:
|
# Ширина кадра в пикселях, разделяющая области на этом оборудовании.
|
||||||
|
SPINE_MIN_WIDTH = 295
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_region(
|
||||||
|
predicted_region: Optional[str],
|
||||||
|
region_confidence: float,
|
||||||
|
signals: HeuristicSignals,
|
||||||
|
min_confidence: float = 0.5,
|
||||||
|
) -> Tuple[str, float]:
|
||||||
"""
|
"""
|
||||||
Extract anatomical region from DICOM filename.
|
Определить анатомическую область.
|
||||||
Expected patterns: spine, l_hip, r_hip, left_hip, right_hip
|
|
||||||
|
Порядок решений:
|
||||||
NOTE: This is a fallback method. Prefer determine_anatomical_region()
|
1. Ширина кадра: у позвоночника кадр шире (300 px против 280 px у бёдер).
|
||||||
which uses image analysis.
|
На этом оборудовании признак разделяет области безошибочно, поэтому
|
||||||
|
используется первым.
|
||||||
|
2. При нетипичной ширине — предсказание обученной головы области.
|
||||||
|
3. Если голова неуверена — форма яркой области и перевес светимости.
|
||||||
|
|
||||||
|
Важно, что область НЕ берётся из имени файла: на закрытом наборе имена
|
||||||
|
могут не содержать разметки региона.
|
||||||
"""
|
"""
|
||||||
import os
|
if signals.width:
|
||||||
filename = os.path.basename(dcm_path).lower()
|
if signals.width >= SPINE_MIN_WIDTH:
|
||||||
|
return "spine", 0.8
|
||||||
if 'spine' in filename:
|
# Бедро: ширину кадра делят левый и правый снимки, поэтому сторону
|
||||||
return 'spine'
|
# определяем по перевесу светимости яркой области. Голова обучена на
|
||||||
elif 'l_hip' in filename or 'left_hip' in filename:
|
# обе стороны, но различает их хуже, чем асимметрия.
|
||||||
return 'hip_left'
|
if signals.left_right_ratio > 1.3:
|
||||||
elif 'r_hip' in filename or 'right_hip' in filename:
|
return "hip_right", 0.6
|
||||||
return 'hip_right'
|
if signals.left_right_ratio < 0.7:
|
||||||
|
return "hip_left", 0.6
|
||||||
return 'unknown'
|
if predicted_region in ("hip_left", "hip_right") and region_confidence >= min_confidence:
|
||||||
|
return predicted_region, region_confidence
|
||||||
|
return "hip", 0.4
|
||||||
|
|
||||||
|
if predicted_region in REGIONS and region_confidence >= min_confidence:
|
||||||
|
return predicted_region, region_confidence
|
||||||
|
|
||||||
|
aspect = signals.bbox_aspect
|
||||||
|
if aspect < 1.5 or (aspect < 1.8 and signals.symmetry > 0.35):
|
||||||
|
return "spine", 0.4
|
||||||
|
if signals.left_right_ratio > 1.3:
|
||||||
|
return "hip_right", 0.3
|
||||||
|
if signals.left_right_ratio < 0.7:
|
||||||
|
return "hip_left", 0.3
|
||||||
|
return "hip", 0.3
|
||||||
|
|
||||||
|
|
||||||
def determine_region_from_image(img: np.ndarray) -> str:
|
def classify_violation_type(
|
||||||
|
region: Optional[str],
|
||||||
|
metrics: Dict,
|
||||||
|
samples: Optional[Dict[str, float]],
|
||||||
|
) -> Tuple[str, str]:
|
||||||
"""
|
"""
|
||||||
Определение анатомической области (позвоночник/бедро) по содержимому изображения.
|
Определить тип нарушения по эвристическим метрикам изображения.
|
||||||
|
|
||||||
Алгоритм использует анализ формы яркой области на изображении:
|
Возвращает (тип, пояснение). Тип выбирается по наиболее выраженному
|
||||||
- Позвоночник: яркая область более квадратная (aspect ratio ~1.2)
|
признаку; при отсутствии сигналов возвращается общая категория.
|
||||||
- Бедро: яркая область вытянута вертикально (aspect ratio > 1.5)
|
|
||||||
|
|
||||||
Дополнительно для определения левого/правого бедра:
|
|
||||||
- Сравнение яркости левой и правой половин изображения
|
|
||||||
|
|
||||||
Args:
|
|
||||||
img: Нормализованное изображение (np.array)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
str: 'spine', 'hip_left', 'hip_right' или 'hip' (неопределенная сторона)
|
|
||||||
"""
|
"""
|
||||||
h, w = img.shape
|
motion = samples.get("laplacian_variance") if samples else None
|
||||||
|
bright_frac = samples.get("bright_fraction") if samples else None
|
||||||
|
|
||||||
# Normalize image
|
if motion is not None and motion < metrics.get("motion_threshold", 0.0):
|
||||||
img_norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
return "artifact_motion", REASON_BY_TYPE["artifact_motion"]
|
||||||
|
if bright_frac is not None and bright_frac > metrics.get("artifact_threshold", 1.0):
|
||||||
|
return "artifact_other", REASON_BY_TYPE["artifact_other"]
|
||||||
|
if region in ("hip_left", "hip_right") and samples:
|
||||||
|
aspect = samples.get("bbox_aspect", 1.0)
|
||||||
|
if aspect < 0.4 or aspect > 3.0:
|
||||||
|
return "rotation", REASON_BY_TYPE["rotation"]
|
||||||
|
|
||||||
# Feature 1: Bright region aspect ratio
|
return "quality_violation_detected", REASON_BY_TYPE["quality_violation_detected"]
|
||||||
threshold = np.percentile(img_norm, 95)
|
|
||||||
binary = img_norm > threshold
|
|
||||||
|
|
||||||
bbox_aspect = 1.0
|
|
||||||
bright_x = 0.5 # default center
|
|
||||||
left_right_ratio = 1.0 # default balanced
|
|
||||||
if binary.sum() > 0:
|
|
||||||
try:
|
|
||||||
from scipy import ndimage
|
|
||||||
rows = np.any(binary, axis=1)
|
|
||||||
cols = np.any(binary, axis=0)
|
|
||||||
if rows.any() and cols.any():
|
|
||||||
rmin, rmax = np.where(rows)[0][[0, -1]]
|
|
||||||
cmin, cmax = np.where(cols)[0][[0, -1]]
|
|
||||||
bbox_h = rmax - rmin
|
|
||||||
bbox_w = cmax - cmin
|
|
||||||
bbox_aspect = bbox_h / (bbox_w + 1e-6)
|
|
||||||
|
|
||||||
# Get bright region center X position
|
def image_samples(img: np.ndarray) -> Dict[str, float]:
|
||||||
com = ndimage.center_of_mass(binary)
|
"""Числовые характеристики изображения для отчёта и выбора типа нарушения."""
|
||||||
bright_x = com[1] / w
|
norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
||||||
|
lap = np.abs(np.diff(norm, axis=0)).mean() + np.abs(np.diff(norm, axis=1)).mean()
|
||||||
|
return {
|
||||||
|
"laplacian_variance": float(((norm - norm.mean()) ** 2).mean() * lap),
|
||||||
|
"bright_fraction": float((norm > np.percentile(norm, 99)).mean()),
|
||||||
|
"bbox_aspect": heuristic_signals(img).bbox_aspect,
|
||||||
|
}
|
||||||
|
|
||||||
# Calculate left/right brightness ratio
|
|
||||||
left_bright = binary[:, :w//2].sum()
|
|
||||||
right_bright = binary[:, w//2:].sum()
|
|
||||||
left_right_ratio = left_bright / (right_bright + 1e-6)
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Feature 2: Symmetry
|
def predict_from_array(arr: np.ndarray, model: DXAQualityModel,
|
||||||
h_mid, w_mid = h // 2, w // 2
|
cfg: PreprocessConfig, device: str, threshold: float) -> Prediction:
|
||||||
left_half = img_norm[:, :w_mid]
|
"""
|
||||||
right_half = np.fliplr(img_norm[:, w_mid:])
|
Предсказание качества по уже прочитанному массиву пикселей.
|
||||||
min_w = min(left_half.shape[1], right_half.shape[1])
|
|
||||||
symmetry = 1 - np.abs(left_half[:, :min_w] - right_half[:, :min_w]).mean() / (img_norm.std() + 1e-6)
|
|
||||||
|
|
||||||
# Feature 3: Vertical/horizontal edges
|
Общая точка входа для файлового инференса и HTTP API: гарантирует, что
|
||||||
dx = np.diff(img_norm, axis=1)
|
предобработка и решающее правило совпадают во всех режимах.
|
||||||
dy = np.diff(img_norm, axis=0)
|
"""
|
||||||
v_edges = np.abs(dx).mean()
|
tensor = torch.from_numpy(preprocess_from_array(arr, cfg)).float().unsqueeze(0).to(device)
|
||||||
h_edges = np.abs(dy).mean()
|
|
||||||
v_h_ratio = v_edges / (h_edges + 1e-6)
|
|
||||||
|
|
||||||
# Classification rules based on analysis:
|
with torch.no_grad():
|
||||||
# spine: bbox_aspect ~1.2, symmetry > 0.4, v_h_ratio < 1.5
|
out = model.model(tensor)
|
||||||
# hip: bbox_aspect > 1.5, symmetry < 0.4, v_h_ratio > 1.5
|
logit = float(out["quality_logits"][0, 1])
|
||||||
|
prob = float(torch.sigmoid(out["quality_logits"][0, 1]))
|
||||||
|
region_probs = torch.softmax(out["region_logits"][0], dim=0)
|
||||||
|
region_conf, region_id = float(region_probs.max()), int(region_probs.argmax())
|
||||||
|
|
||||||
# Primary: bbox_aspect is the best discriminator
|
signals = heuristic_signals(arr)
|
||||||
if bbox_aspect < 1.5:
|
region, region_conf = resolve_region(REGION_BY_ID.get(region_id), region_conf, signals)
|
||||||
# More square bright region -> spine
|
quality_class = 1 if logit >= threshold else 0
|
||||||
return 'spine'
|
samples = image_samples(arr)
|
||||||
elif bbox_aspect < 1.8:
|
|
||||||
# Check symmetry as secondary
|
if quality_class == 1:
|
||||||
if symmetry > 0.35:
|
violation_type, reason = classify_violation_type(region, {}, samples)
|
||||||
return 'spine'
|
|
||||||
else:
|
|
||||||
return 'hip' # unknown side
|
|
||||||
else:
|
else:
|
||||||
# Highly elongated bright region -> hip
|
violation_type, reason = "", ""
|
||||||
# Determine left vs right based on left/right brightness ratio
|
|
||||||
# left_right_ratio > 1.3 -> right hip (right side brighter)
|
return Prediction(
|
||||||
# left_right_ratio < 0.7 -> left hip (left side brighter)
|
quality_class=quality_class,
|
||||||
if left_right_ratio > 1.3:
|
prob=prob,
|
||||||
return 'hip_right'
|
logit=logit,
|
||||||
elif left_right_ratio < 0.7:
|
threshold=threshold,
|
||||||
return 'hip_left'
|
anatomical_region=region,
|
||||||
else:
|
region_confidence=region_conf,
|
||||||
return 'hip' # unclear side
|
violation_type=violation_type,
|
||||||
|
violation_reason=reason,
|
||||||
|
samples=samples,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def determine_anatomical_region(dcm_path: str) -> str:
|
def predict_from_bytes(dcm_bytes: bytes, model: DXAQualityModel,
|
||||||
|
cfg: PreprocessConfig, device: str, threshold: float) -> Tuple[Prediction, pydicom.Dataset]:
|
||||||
"""
|
"""
|
||||||
Determine anatomical region from DICOM image content (primary method).
|
Предсказание по байтам DICOM без записи на диск.
|
||||||
|
|
||||||
Uses image analysis to determine spine vs hip.
|
API принимает файлы потоком, поэтому чтение идёт через BytesIO — временные
|
||||||
Falls back to height-based heuristic only if image cannot be analyzed.
|
файлы не создаются.
|
||||||
"""
|
"""
|
||||||
|
import io as _io
|
||||||
|
|
||||||
|
ds = pydicom.dcmread(_io.BytesIO(dcm_bytes))
|
||||||
|
arr = ds.pixel_array.astype(np.float32)
|
||||||
|
if arr.ndim == 3:
|
||||||
|
arr = arr.mean(axis=0) if arr.shape[0] > 1 else arr[0]
|
||||||
|
if str(getattr(ds, "PhotometricInterpretation", "")).upper() == "MONOCHROME1":
|
||||||
|
arr = arr.max() - arr
|
||||||
|
return predict_from_array(arr, model, cfg, device, threshold), ds
|
||||||
|
|
||||||
|
|
||||||
|
def study_uid_for(ds: pydicom.Dataset, dcm_path: Path) -> str:
|
||||||
|
"""StudyInstanceUID из DICOM; при отсутствии — имя каталога исследования."""
|
||||||
|
uid = str(getattr(ds, "StudyInstanceUID", "") or "").strip()
|
||||||
|
if uid:
|
||||||
|
return uid
|
||||||
|
return dcm_path.parent.name
|
||||||
|
|
||||||
|
|
||||||
|
def process_one(
|
||||||
|
dcm_path: Path,
|
||||||
|
model: DXAQualityModel,
|
||||||
|
cfg: PreprocessConfig,
|
||||||
|
device: str,
|
||||||
|
threshold: float,
|
||||||
|
) -> ImageResult:
|
||||||
|
"""Полная обработка одного DICOM файла."""
|
||||||
|
started = datetime.now()
|
||||||
|
path_to_study = str(dcm_path.parent)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Load image
|
arr = load_dicom_array(dcm_path)
|
||||||
ds = pydicom.dcmread(dcm_path)
|
prediction = predict_from_array(arr, model, cfg, device, threshold)
|
||||||
img = ds.pixel_array.astype(np.float32)
|
ds = pydicom.dcmread(str(dcm_path), stop_before_pixels=True)
|
||||||
|
|
||||||
# Determine from image content
|
|
||||||
region = determine_region_from_image(img)
|
|
||||||
|
|
||||||
if region != 'unknown':
|
|
||||||
return region
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Fallback: try metadata-based detection
|
|
||||||
try:
|
|
||||||
ds = pydicom.dcmread(dcm_path)
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
|
|
||||||
if h < 270 or w < 280:
|
return ImageResult(
|
||||||
return 'hip'
|
path_to_study=path_to_study,
|
||||||
else:
|
study_uid=study_uid_for(ds, dcm_path),
|
||||||
return 'spine'
|
image_uid=str(getattr(ds, "SOPInstanceUID", "") or ""),
|
||||||
except:
|
anatomical_region=prediction.anatomical_region,
|
||||||
return 'unknown'
|
quality_class=prediction.quality_class,
|
||||||
|
violation_type=prediction.violation_type,
|
||||||
|
processing_status="Success",
|
||||||
|
time_of_processing=(datetime.now() - started).total_seconds(),
|
||||||
|
confidence=round(prediction.prob, 4),
|
||||||
|
violation_reason=prediction.violation_reason,
|
||||||
|
region_confidence=round(prediction.region_confidence, 4),
|
||||||
|
dcm_path=str(dcm_path),
|
||||||
|
metrics={"logit": prediction.logit, "threshold": threshold, **prediction.samples},
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
# Требование: необработанных исключений быть не должно, все ошибки
|
||||||
|
# фиксируются в отчёте со статусом Failure.
|
||||||
|
logger.warning("Failed to process %s: %s", dcm_path, exc)
|
||||||
|
return ImageResult(
|
||||||
|
path_to_study=path_to_study,
|
||||||
|
study_uid="",
|
||||||
|
image_uid="",
|
||||||
|
anatomical_region="unknown",
|
||||||
|
quality_class=-1,
|
||||||
|
violation_type="",
|
||||||
|
processing_status=f"Failure: {type(exc).__name__}: {str(exc)[:120]}",
|
||||||
|
time_of_processing=(datetime.now() - started).total_seconds(),
|
||||||
|
dcm_path=str(dcm_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def process_dicom_files(args):
|
def write_visualizations(results: Sequence[ImageResult], cfg: PreprocessConfig,
|
||||||
|
zip_path: Path, enabled: bool) -> None:
|
||||||
"""
|
"""
|
||||||
Основная функция пакетного инференса DXA изображений.
|
Дополнительный функционал: zip-архив с изображениями, где выделена
|
||||||
|
зона интереса (порог по 90-му перцентилю яркости).
|
||||||
Этапы обработки:
|
|
||||||
1. Определение устройства (MPS/CUDA/CPU)
|
|
||||||
2. Загрузка обученной модели
|
|
||||||
3. Поиск DICOM файлов в указанной директории
|
|
||||||
4. Инференс для каждого файла:
|
|
||||||
- Предобработка изображения
|
|
||||||
- Предикт модели (бинарная классификация)
|
|
||||||
- Определение анатомической области
|
|
||||||
- Сохранение метаданных DICOM
|
|
||||||
5. Экспорт результатов в Excel/CSV
|
|
||||||
|
|
||||||
Args:
|
|
||||||
args: Аргументы командной строки (input_path, output_path, model_path и т.д.)
|
|
||||||
"""
|
"""
|
||||||
# Setup
|
if not enabled:
|
||||||
device = get_device()
|
|
||||||
print(f"Using device: {device}")
|
|
||||||
|
|
||||||
# Load model
|
|
||||||
if Path(args.model_path).exists():
|
|
||||||
print(f"Loading model from {args.model_path}...")
|
|
||||||
model = load_model(args.model_path, args.backbone, device)
|
|
||||||
print("Model loaded successfully")
|
|
||||||
else:
|
|
||||||
print(f"WARNING: Model not found at {args.model_path}")
|
|
||||||
print("Using untrained model - results will be random")
|
|
||||||
from src.dxa.model import create_model
|
|
||||||
model = create_model(backbone=args.backbone, pretrained=False, device=device)
|
|
||||||
|
|
||||||
# Find DICOM files
|
|
||||||
input_path = Path(args.input_path)
|
|
||||||
dcm_files = []
|
|
||||||
|
|
||||||
if input_path.is_file() and input_path.suffix.lower() == '.dcm':
|
|
||||||
dcm_files = [input_path]
|
|
||||||
elif input_path.is_dir():
|
|
||||||
dcm_files = sorted(input_path.rglob('*.dcm'))
|
|
||||||
|
|
||||||
print(f"Found {len(dcm_files)} DICOM files")
|
|
||||||
|
|
||||||
if len(dcm_files) == 0:
|
|
||||||
print("ERROR: No DICOM files found!")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
# Process each file
|
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as archive:
|
||||||
results = []
|
for res in results:
|
||||||
|
if res.processing_status != "Success" or not res.dcm_path:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
arr = load_dicom_array(res.dcm_path)
|
||||||
|
norm = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)
|
||||||
|
mask = norm > np.percentile(norm, 90)
|
||||||
|
|
||||||
|
rgb = np.stack([norm] * 3, axis=-1)
|
||||||
|
edge = np.zeros_like(mask)
|
||||||
|
vertical = mask[2:, 1:-1] ^ mask[:-2, 1:-1]
|
||||||
|
horizontal = mask[1:-1, 2:] ^ mask[1:-1, :-2]
|
||||||
|
edge[1:-1, 1:-1] = vertical | horizontal
|
||||||
|
rgb[edge] = [1.0, 0.2, 0.2]
|
||||||
|
|
||||||
|
name = f"{res.study_uid or 'study'}_{res.image_uid or Path(res.dcm_path).stem}.png"
|
||||||
|
archive.writestr(name, _to_png_bytes(rgb))
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Visualization failed for %s: %s", res.dcm_path, exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _to_png_bytes(rgb: np.ndarray) -> bytes:
|
||||||
|
import io
|
||||||
|
|
||||||
|
image = Image.fromarray((np.clip(rgb, 0, 1) * 255).astype(np.uint8))
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
image.save(buffer, format="PNG")
|
||||||
|
return buffer.getvalue()
|
||||||
|
|
||||||
|
|
||||||
|
def results_to_dataframe(results: Sequence[ImageResult]) -> pd.DataFrame:
|
||||||
|
"""Таблица строго в формате требований; пояснения добавляются справа."""
|
||||||
|
rows = []
|
||||||
|
for r in results:
|
||||||
|
rows.append({
|
||||||
|
"path_to_study": r.path_to_study,
|
||||||
|
"study_uid": r.study_uid,
|
||||||
|
"image_uid": r.image_uid,
|
||||||
|
"anatomical_region": r.anatomical_region,
|
||||||
|
"quality_class": r.quality_class,
|
||||||
|
"violation_type": r.violation_type,
|
||||||
|
"processing_status": r.processing_status,
|
||||||
|
"time_of_processing": round(r.time_of_processing, 4),
|
||||||
|
"confidence": r.confidence,
|
||||||
|
"violation_reason": r.violation_reason,
|
||||||
|
"region_confidence": r.region_confidence,
|
||||||
|
})
|
||||||
|
return pd.DataFrame(rows, columns=OUTPUT_COLUMNS + ["confidence", "violation_reason", "region_confidence"])
|
||||||
|
|
||||||
|
|
||||||
|
def run(args: argparse.Namespace) -> pd.DataFrame:
|
||||||
|
"""Пакетная обработка входного пути и запись отчёта."""
|
||||||
|
device = get_device(args.device)
|
||||||
|
logger.info("Device: %s", device)
|
||||||
|
|
||||||
|
model_path = Path(args.model_path)
|
||||||
|
if not model_path.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Model checkpoint not found: {model_path}. Train it first: python -m src.dxa.train"
|
||||||
|
)
|
||||||
|
|
||||||
|
checkpoint = load_model(
|
||||||
|
str(model_path),
|
||||||
|
backbone=args.backbone,
|
||||||
|
head=args.head,
|
||||||
|
device=device,
|
||||||
|
input_size=args.input_size,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Model: backbone=%s head=%s, threshold(logit)=%.4f, preprocess=%s",
|
||||||
|
checkpoint.metadata.get("backbone"), checkpoint.metadata.get("head"),
|
||||||
|
checkpoint.threshold, checkpoint.preprocess.to_dict(),
|
||||||
|
)
|
||||||
|
|
||||||
|
dcm_files = list(iter_dicom_files(args.input_path))
|
||||||
|
if not dcm_files:
|
||||||
|
raise FileNotFoundError(f"No DICOM files found under {args.input_path}")
|
||||||
|
logger.info("Found %d DICOM files", len(dcm_files))
|
||||||
|
|
||||||
|
results: List[ImageResult] = []
|
||||||
for dcm_path in tqdm(dcm_files, desc="Processing"):
|
for dcm_path in tqdm(dcm_files, desc="Processing"):
|
||||||
start_time = datetime.now()
|
results.append(process_one(
|
||||||
|
dcm_path, checkpoint.model, checkpoint.preprocess, device, checkpoint.threshold
|
||||||
|
))
|
||||||
|
|
||||||
try:
|
df = results_to_dataframe(results)
|
||||||
# Load and preprocess image
|
|
||||||
img = load_dicom_image(str(dcm_path), args.input_size)
|
|
||||||
img = img.to(device)
|
|
||||||
|
|
||||||
# Predict
|
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model.model(img)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
prob = probs[0, pred].item()
|
|
||||||
|
|
||||||
# Get DICOM metadata
|
|
||||||
ds = pydicom.dcmread(str(dcm_path))
|
|
||||||
|
|
||||||
study_uid = getattr(ds, 'StudyInstanceUID', '')
|
|
||||||
image_uid = getattr(ds, 'SOPInstanceUID', '')
|
|
||||||
|
|
||||||
# Determine anatomical region
|
|
||||||
anatomical_region = determine_anatomical_region(str(dcm_path))
|
|
||||||
|
|
||||||
# Map prediction to quality class
|
|
||||||
quality_class = pred # 0 = good, 1 = violation
|
|
||||||
|
|
||||||
# Determine violation type (simplified)
|
|
||||||
if quality_class == 0:
|
|
||||||
violation_type = ''
|
|
||||||
else:
|
|
||||||
# In real implementation, this would come from a more detailed model
|
|
||||||
violation_type = 'quality_violation_detected'
|
|
||||||
|
|
||||||
processing_time = (datetime.now() - start_time).total_seconds()
|
|
||||||
|
|
||||||
results.append({
|
|
||||||
'path_to_study': str(dcm_path.parent),
|
|
||||||
'study_uid': study_uid,
|
|
||||||
'image_uid': image_uid,
|
|
||||||
'anatomical_region': anatomical_region,
|
|
||||||
'quality_class': quality_class,
|
|
||||||
'violation_type': violation_type,
|
|
||||||
'processing_status': 'Success',
|
|
||||||
'time_of_processing': processing_time,
|
|
||||||
'confidence': round(prob, 4)
|
|
||||||
})
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
processing_time = (datetime.now() - start_time).total_seconds()
|
|
||||||
results.append({
|
|
||||||
'path_to_study': str(dcm_path.parent) if dcm_path.parent else '',
|
|
||||||
'study_uid': str(dcm_path),
|
|
||||||
'image_uid': '',
|
|
||||||
'anatomical_region': 'unknown',
|
|
||||||
'quality_class': -1,
|
|
||||||
'violation_type': '',
|
|
||||||
'processing_status': f'Failure: {str(e)[:80]}',
|
|
||||||
'time_of_processing': processing_time,
|
|
||||||
'confidence': 0.0
|
|
||||||
})
|
|
||||||
|
|
||||||
# Create DataFrame with required columns
|
|
||||||
df = pd.DataFrame(results)
|
|
||||||
|
|
||||||
# Ensure correct column order as per requirements
|
|
||||||
output_columns = [
|
|
||||||
'path_to_study',
|
|
||||||
'study_uid',
|
|
||||||
'image_uid',
|
|
||||||
'anatomical_region',
|
|
||||||
'quality_class',
|
|
||||||
'violation_type',
|
|
||||||
'processing_status',
|
|
||||||
'time_of_processing'
|
|
||||||
]
|
|
||||||
|
|
||||||
# Add confidence if present
|
|
||||||
if 'confidence' in df.columns:
|
|
||||||
output_columns.append('confidence')
|
|
||||||
|
|
||||||
# Reorder columns (add missing ones with empty values)
|
|
||||||
for col in output_columns:
|
|
||||||
if col not in df.columns:
|
|
||||||
df[col] = ''
|
|
||||||
|
|
||||||
df = df[output_columns]
|
|
||||||
|
|
||||||
# Save results
|
|
||||||
output_path = Path(args.output_path)
|
output_path = Path(args.output_path)
|
||||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
if output_path.suffix.lower() == ".csv":
|
||||||
if output_path.suffix.lower() == '.csv':
|
|
||||||
df.to_csv(output_path, index=False)
|
df.to_csv(output_path, index=False)
|
||||||
else:
|
else:
|
||||||
df.to_excel(output_path, index=False)
|
df.to_excel(output_path, index=False)
|
||||||
|
|
||||||
print(f"\n{'='*50}")
|
if args.zip_out:
|
||||||
print(f"Results saved to {output_path}")
|
write_visualizations(results, checkpoint.preprocess, Path(args.zip_out), enabled=True)
|
||||||
print(f"{'='*50}")
|
logger.info("Visualizations: %s", args.zip_out)
|
||||||
print(f"\nSummary:")
|
|
||||||
print(f" Total files: {len(df)}")
|
successful = int((df["processing_status"] == "Success").sum())
|
||||||
print(f" Successful: {(df['processing_status'] == 'Success').sum()}")
|
logger.info(
|
||||||
print(f" Quality OK (class 0): {(df['quality_class'] == 0).sum()}")
|
"Done. %d/%d processed (%.1f%%), quality_class=1 in %d rows, median time %.3fs -> %s",
|
||||||
print(f" Quality Issues (class 1): {(df['quality_class'] == 1).sum()}")
|
successful, len(df), 100 * successful / max(len(df), 1),
|
||||||
|
int((df["quality_class"] == 1).sum()),
|
||||||
# Show sample output
|
float(df["time_of_processing"].median()) if len(df) else 0.0,
|
||||||
print(f"\nSample output:")
|
output_path,
|
||||||
print(df.head().to_string())
|
)
|
||||||
|
return df
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
parser = argparse.ArgumentParser(description='DXA Quality Inference')
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Пакетный инференс классификатора качества DXA",
|
||||||
# Input/Output
|
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||||
parser.add_argument('--input-path', type=str, required=True,
|
)
|
||||||
help='Path to DICOM file or directory')
|
parser.add_argument("--input-path", required=True, help="DICOM файл или каталог")
|
||||||
parser.add_argument('--output-path', type=str, required=True,
|
parser.add_argument("--output-path", required=True, help="Путь к .xlsx или .csv")
|
||||||
help='Output CSV or Excel file')
|
parser.add_argument("--model-path", default="models/dxa_model.pth", help="Чекпоинт модели")
|
||||||
parser.add_argument('--model-path', type=str,
|
parser.add_argument("--backbone", default=None, choices=["resnet18", "resnet34"],
|
||||||
default='models/dxa_model.pth',
|
help="Проверить соответствие backbone в чекпоинте (по умолчанию — из чекпоинта)")
|
||||||
help='Path to trained model')
|
parser.add_argument("--head", default=None, choices=["linear", "mlp"],
|
||||||
|
help="Проверить соответствие головы в чекпоинте (по умолчанию — из чекпоинта)")
|
||||||
# Model arguments
|
parser.add_argument("--input-size", type=int, default=None,
|
||||||
parser.add_argument('--backbone', type=str, default='resnet18',
|
help="Переопределить размер входа (по умолчанию — из чекпоинта)")
|
||||||
choices=['resnet18', 'resnet34', 'efficientnet_b0'],
|
parser.add_argument("--device", default=None, help="cpu / cuda / mps")
|
||||||
help='Backbone architecture')
|
parser.add_argument("--zip-out", default=None,
|
||||||
parser.add_argument('--input-size', type=int, default=224,
|
help="Zip-архив с визуализацией зоны интереса (дополнительный функционал)")
|
||||||
help='Input image size')
|
return parser
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
process_dicom_files(args)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
def main(argv: Optional[List[str]] = None) -> int:
|
||||||
main()
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||||
|
args = build_parser().parse_args(argv)
|
||||||
|
try:
|
||||||
|
run(args)
|
||||||
|
except (FileNotFoundError, ValueError) as exc:
|
||||||
|
logger.error("%s", exc)
|
||||||
|
return 1
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
|
|
|
||||||
647
src/dxa/model.py
647
src/dxa/model.py
|
|
@ -1,199 +1,506 @@
|
||||||
"""
|
"""
|
||||||
DXA Quality Classifier Model
|
Модель классификации качества 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
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torchvision.models as models
|
import torchvision.models as models
|
||||||
from typing import Dict, Tuple, Optional
|
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):
|
class DXAQualityClassifier(nn.Module):
|
||||||
"""
|
"""
|
||||||
CNN classifier for DXA image quality assessment
|
ResNet backbone + голова качества + вспомогательная голова области.
|
||||||
Uses pretrained backbone (ResNet/EfficientNet)
|
|
||||||
|
Голова выбирается параметром `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,
|
def __init__(
|
||||||
backbone: str = 'resnet18',
|
self,
|
||||||
num_classes: int = 2,
|
backbone: str = "resnet18",
|
||||||
pretrained: bool = True,
|
pretrained: bool = True,
|
||||||
dropout: float = 0.3):
|
dropout: float = 0.3,
|
||||||
|
head: str = "linear",
|
||||||
|
hidden_dim: int = 256,
|
||||||
|
feature_norm: bool = True,
|
||||||
|
):
|
||||||
super().__init__()
|
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.backbone_name = backbone
|
||||||
|
self.head_type = head
|
||||||
# Load pretrained backbone
|
self.feature_norm = feature_norm
|
||||||
if backbone == 'resnet18':
|
weights_name, feature_dim = _WEIGHTS[backbone]
|
||||||
self.backbone = models.resnet18(weights='IMAGENET1K_V1' if pretrained else None)
|
net = getattr(models, backbone)(weights=weights_name if pretrained else None)
|
||||||
feature_dim = 512
|
self.features = nn.Sequential(*list(net.children())[:-1]) # всё, кроме fc
|
||||||
elif backbone == 'resnet34':
|
|
||||||
self.backbone = models.resnet34(weights='IMAGENET1K_V1' if pretrained else None)
|
|
||||||
feature_dim = 512
|
|
||||||
elif backbone == 'efficientnet_b0':
|
|
||||||
self.backbone = models.efficientnet_b0(weights='IMAGENET1K_V1' if pretrained else None)
|
|
||||||
feature_dim = 1280
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown backbone: {backbone}")
|
|
||||||
|
|
||||||
# Replace final layer
|
|
||||||
if backbone.startswith('resnet'):
|
|
||||||
self.backbone.fc = nn.Identity()
|
|
||||||
elif backbone.startswith('efficientnet'):
|
|
||||||
self.backbone.classifier = nn.Identity()
|
|
||||||
|
|
||||||
# Classifier head
|
|
||||||
self.classifier = nn.Sequential(
|
|
||||||
nn.Dropout(dropout),
|
|
||||||
nn.Linear(feature_dim, 256),
|
|
||||||
nn.ReLU(),
|
|
||||||
nn.Dropout(dropout),
|
|
||||||
nn.Linear(256, num_classes)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Store for feature extraction
|
|
||||||
self.feature_dim = feature_dim
|
self.feature_dim = feature_dim
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
self.register_buffer("feat_mean", torch.zeros(feature_dim))
|
||||||
"""Forward pass"""
|
self.register_buffer("feat_std", torch.ones(feature_dim))
|
||||||
features = self.backbone(x)
|
|
||||||
return self.classifier(features)
|
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:
|
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""Extract features without classification"""
|
return self.features(x)
|
||||||
return self.backbone(x)
|
|
||||||
|
def _head_input(self, feats: torch.Tensor) -> torch.Tensor:
|
||||||
def get_info(self) -> Dict:
|
"""Сгладить признаки до (B, feature_dim) и применить стандартизацию."""
|
||||||
"""Get model info"""
|
feats = feats.flatten(1)
|
||||||
return {
|
if self.feature_norm:
|
||||||
'backbone': self.backbone_name,
|
feats = (feats - self.feat_mean) / self.feat_std
|
||||||
'num_classes': 2,
|
return feats
|
||||||
'feature_dim': self.feature_dim,
|
|
||||||
'task': 'binary_quality_classification'
|
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:
|
class DXAQualityModel:
|
||||||
"""Wrapper for training and inference"""
|
"""Обёртка вокруг сети: обучение, валидация, предсказание, сохранение/загрузка."""
|
||||||
|
|
||||||
def __init__(self,
|
def __init__(
|
||||||
model: DXAQualityClassifier,
|
self,
|
||||||
device: str = 'cpu',
|
model: DXAQualityClassifier,
|
||||||
learning_rate: float = 1e-4):
|
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.model = model
|
||||||
self.device = device
|
self.device = torch.device(device)
|
||||||
self.model.to(device)
|
self.model.to(self.device)
|
||||||
|
self._backbone_trainable = True
|
||||||
# Loss and optimizer
|
|
||||||
self.criterion = nn.CrossEntropyLoss()
|
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(
|
self.optimizer = torch.optim.AdamW(
|
||||||
model.parameters(),
|
model.parameters(), lr=learning_rate, weight_decay=weight_decay
|
||||||
lr=learning_rate,
|
|
||||||
weight_decay=1e-5
|
|
||||||
)
|
)
|
||||||
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||||
self.optimizer, mode='min', factor=0.5, patience=3
|
self.optimizer, mode="min", factor=0.5, patience=3
|
||||||
)
|
)
|
||||||
|
self.history: Dict[str, list] = {"train_loss": [], "val_loss": [], "val_f1": [], "val_roc_auc": []}
|
||||||
# Training history
|
|
||||||
self.history = {
|
def _step(self, batch, train: bool) -> Tuple[float, torch.Tensor, torch.Tensor]:
|
||||||
'train_loss': [],
|
images, meta = batch
|
||||||
'val_loss': [],
|
images = images.to(self.device, non_blocking=True)
|
||||||
'train_acc': [],
|
labels = meta["label"].to(self.device)
|
||||||
'val_acc': []
|
region_ids = meta["region_id"].to(self.device)
|
||||||
}
|
|
||||||
|
with torch.set_grad_enabled(train):
|
||||||
def train_epoch(self, train_loader) -> Tuple[float, float]:
|
out = self.model(images)
|
||||||
"""Train one epoch"""
|
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()
|
self.model.train()
|
||||||
total_loss = 0
|
if not getattr(self, "_backbone_trainable", True):
|
||||||
correct = 0
|
for module in self.model.features.modules():
|
||||||
total = 0
|
if isinstance(module, torch.nn.modules.batchnorm._BatchNorm):
|
||||||
|
module.eval()
|
||||||
for images, labels in train_loader:
|
|
||||||
images = images.to(self.device)
|
def set_learning_rate(self, lr: float) -> None:
|
||||||
|
"""Задать learning rate всем группам параметров."""
|
||||||
# Handle dict format from dataset
|
for group in self.optimizer.param_groups:
|
||||||
if isinstance(labels, dict):
|
group["lr"] = lr
|
||||||
labels_tensor = labels['label'].to(self.device)
|
|
||||||
else:
|
@torch.no_grad()
|
||||||
labels_tensor = labels.to(self.device)
|
def fit_feature_norm(self, loader) -> None:
|
||||||
|
"""
|
||||||
self.optimizer.zero_grad()
|
Оценить среднее и СКО признаков по обучающей выборке и зафиксировать их.
|
||||||
|
|
||||||
outputs = self.model(images)
|
Стандартизация входа нужна линейной голове: без неё логиты смещены,
|
||||||
loss = self.criterion(outputs, labels_tensor)
|
вероятности скучены у нуля, а подобранный порог теряет смысл.
|
||||||
|
|
||||||
loss.backward()
|
Дисперсия считается в два прохода (сначала среднее, затем сумма
|
||||||
self.optimizer.step()
|
квадратов отклонений). Формула E[x²]−E[x]² при float32 на признаках
|
||||||
|
порядка 10 даёт погрешность, сопоставимую с самой дисперсией: СКО
|
||||||
total_loss += loss.item() * images.size(0)
|
выходило случайным, из-за чего логиты насыщались и порог вырождался.
|
||||||
_, predicted = outputs.max(1)
|
Буферы не обучаемые, поэтому статистики не «подглядывают» в валидацию.
|
||||||
correct += predicted.eq(labels_tensor).sum().item()
|
"""
|
||||||
total += labels_tensor.size(0)
|
if not self.model.feature_norm:
|
||||||
|
return
|
||||||
return total_loss / total, correct / total
|
|
||||||
|
|
||||||
def validate(self, val_loader) -> Tuple[float, float]:
|
|
||||||
"""Validate"""
|
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
total_loss = 0
|
|
||||||
correct = 0
|
chunks = []
|
||||||
total = 0
|
for images, _ in loader:
|
||||||
|
chunks.append(self.model.extract_features(images.to(self.device)).flatten(1).cpu().float())
|
||||||
with torch.no_grad():
|
if not chunks:
|
||||||
for images, labels in val_loader:
|
return
|
||||||
images = images.to(self.device)
|
|
||||||
|
# float64 на CPU: размерность мала, а точность здесь критична.
|
||||||
# Handle dict format from dataset
|
feats = torch.cat(chunks).double()
|
||||||
if isinstance(labels, dict):
|
mean = feats.mean(dim=0)
|
||||||
labels_tensor = labels['label'].to(self.device)
|
std = torch.sqrt(((feats - mean) ** 2).mean(dim=0))
|
||||||
else:
|
|
||||||
labels_tensor = labels.to(self.device)
|
# Нижняя граница СКО: у части размерностей разброс близок к нулю, а
|
||||||
|
# деление на него усиливает шум в десятки раз и насыщает логиты.
|
||||||
outputs = self.model(images)
|
floor = max(float(std.median()) * 0.25, 1e-6)
|
||||||
loss = self.criterion(outputs, labels_tensor)
|
std = std.clamp(min=floor)
|
||||||
|
|
||||||
total_loss += loss.item() * images.size(0)
|
self.model.feat_mean.copy_(mean.float().to(self.model.feat_mean.device))
|
||||||
_, predicted = outputs.max(1)
|
self.model.feat_std.copy_(std.float().to(self.model.feat_std.device))
|
||||||
correct += predicted.eq(labels_tensor).sum().item()
|
logger.debug(
|
||||||
total += labels_tensor.size(0)
|
"Feature norm fitted: %d samples, mean norm %.2f, std median %.4f, floor %.5f",
|
||||||
|
feats.shape[0], float(mean.norm()), float(std.median()), floor,
|
||||||
return total_loss / total, correct / total
|
)
|
||||||
|
|
||||||
def predict(self, images: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
def train_epoch(self, loader) -> Tuple[float, Dict[str, Any]]:
|
||||||
"""Predict on batch of images"""
|
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()
|
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():
|
with torch.no_grad():
|
||||||
images = images.to(self.device)
|
for batch in loader:
|
||||||
outputs = self.model(images)
|
images, meta = batch
|
||||||
probs = torch.softmax(outputs, dim=1)
|
out = self.model(images.to(self.device))
|
||||||
preds = outputs.argmax(dim=1)
|
logits.append(out["quality_logits"][:, 1].cpu())
|
||||||
return preds, probs
|
labels.append(meta["label"])
|
||||||
|
regions.append(meta["region_id"])
|
||||||
def save(self, path: str):
|
return torch.cat(logits), torch.cat(labels), torch.cat(regions)
|
||||||
"""Save model"""
|
|
||||||
torch.save({
|
def save(self, path: str | Path, preprocess: Optional[PreprocessConfig] = None, **extra) -> None:
|
||||||
'model_state_dict': self.model.state_dict(),
|
"""Сохранить чекпоинт вместе с архитектурой и параметрами предобработки."""
|
||||||
'optimizer_state_dict': self.optimizer.state_dict(),
|
payload: Dict[str, Any] = {
|
||||||
'history': self.history
|
"model_state_dict": self.model.state_dict(),
|
||||||
}, path)
|
"optimizer_state_dict": self.optimizer.state_dict(),
|
||||||
|
"history": self.history,
|
||||||
def load(self, path: str):
|
"backbone": self.model.backbone_name,
|
||||||
"""Load model"""
|
"head": self.model.head_type,
|
||||||
checkpoint = torch.load(path, map_location=self.device)
|
"format_version": 2,
|
||||||
self.model.load_state_dict(checkpoint['model_state_dict'])
|
"region_loss_weight": self.region_loss_weight,
|
||||||
self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
|
}
|
||||||
self.history = checkpoint.get('history', self.history)
|
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',
|
def create_model(
|
||||||
num_classes: int = 2,
|
backbone: str = "resnet18",
|
||||||
pretrained: bool = True,
|
pretrained: bool = True,
|
||||||
device: str = 'cpu') -> DXAQualityModel:
|
device: str = "cpu",
|
||||||
"""Create model instance"""
|
head: str = "linear",
|
||||||
model = DXAQualityClassifier(
|
**kwargs,
|
||||||
backbone=backbone,
|
) -> DXAQualityModel:
|
||||||
num_classes=num_classes,
|
"""Создать обёртку модели с заданным backbone и головой."""
|
||||||
pretrained=pretrained
|
return DXAQualityModel(
|
||||||
|
DXAQualityClassifier(backbone=backbone, pretrained=pretrained, head=head),
|
||||||
|
device=device,
|
||||||
|
**kwargs,
|
||||||
)
|
)
|
||||||
return DXAQualityModel(model, device=device)
|
|
||||||
|
|
|
||||||
699
src/dxa/train.py
699
src/dxa/train.py
|
|
@ -1,248 +1,529 @@
|
||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""
|
"""
|
||||||
Training script for DXA Quality Classifier
|
Обучение классификатора качества DXA исследований
|
||||||
=========================================
|
================================================
|
||||||
|
|
||||||
Скрипт обучения модели классификации качества DXA изображений.
|
Модель решает бинарную задачу: годен снимок (0) или есть нарушение (1).
|
||||||
|
Метки берутся из имён DICOM-файлов, где `_good`/`_bad` — экспертная оценка,
|
||||||
|
а отсутствие суффикса означает «изображение хорошее» (см. `src/dxa/labels.py`).
|
||||||
|
|
||||||
Основные этапы:
|
Особенности, отличающие этот скрипт от прежней версии:
|
||||||
1. Загрузка данных из Excel разметки
|
|
||||||
2. Создание DataLoader с разбиением на train/val
|
|
||||||
3. Инициализация модели (ResNet18 с предобучением на ImageNet)
|
|
||||||
4. Обучение с использованием CrossEntropyLoss
|
|
||||||
5. Валидация после каждой эпохи
|
|
||||||
6. Сохранение лучшей модели по F1-score
|
|
||||||
|
|
||||||
Использование:
|
1. **Метки из имён файлов.** Раньше метка бралась из Excel по исследованию и
|
||||||
python src/dxa/train.py --epochs 20 --batch-size 8
|
раздавалась всем снимкам этого исследования; файлы с явной меткой (`_bad`)
|
||||||
|
при этом игнорировались либо, наоборот, не сопоставлялись со своими
|
||||||
|
дублями. Теперь источник метки — имя файла, а Excel используется только для
|
||||||
|
предупреждения о расхождениях.
|
||||||
|
2. **Склейка дублей.** В датасете 544 файла, но всего 252 уникальных снимка:
|
||||||
|
один и тот же кадр сохранён многократно под разными именами. Без склейки
|
||||||
|
модель заучивала бы конкретные снимки как разные примеры.
|
||||||
|
3. **Разбиение по исследованиям.** Снимки одного исследования не попадают
|
||||||
|
одновременно в train и val, что исключает утечку данных.
|
||||||
|
4. **Нормировка ImageNet.** Предобученный backbone ожидает такой вход;
|
||||||
|
прежняя версия подавала сырые [0, 1].
|
||||||
|
5. **Дисбаланс классов.** Учитывается через `pos_weight` и/или взвешенный
|
||||||
|
сэмплер; порог вероятности подбирается на валидации по F1.
|
||||||
|
6. **Линейный зонд вместо fine-tune.** При ~250 уникальных снимках полный
|
||||||
|
fine-tune ResNet18 переобучается за несколько эпох (train F1 → 1.0 при
|
||||||
|
val AUC ≈ 0.5). По умолчанию backbone заморожен и обучается один линейный
|
||||||
|
слой — на валидации это даёт AUC ≈ 0.78. Контрольная задача «позвоночник
|
||||||
|
или бедро» решается идеально (AUC = 1.00), то есть пайплайн исправен, а
|
||||||
|
ограничение — объём и качество разметки.
|
||||||
|
|
||||||
|
Примеры запуска:
|
||||||
|
|
||||||
|
python -m src.dxa.train # режим по умолчанию: линейный зонд
|
||||||
|
python -m src.dxa.train --head mlp --freeze-epochs 5 --epochs 30
|
||||||
|
python -m src.dxa.train --dry-run # проверить разбор данных без обучения
|
||||||
"""
|
"""
|
||||||
import os
|
from __future__ import annotations
|
||||||
import sys
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import random
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
|
from collections import deque
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import torch
|
||||||
from tqdm import tqdm
|
from torch.utils.data import WeightedRandomSampler
|
||||||
|
|
||||||
# Add src to path
|
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
|
||||||
|
|
||||||
from src.dxa.dataset import create_dataloaders
|
from src.dxa.dataset import DXADataset, make_datasets
|
||||||
from src.dxa.model import create_model
|
from src.dxa.labels import QUALITY_BAD, format_summary
|
||||||
|
from src.dxa.model import create_model, per_region_metrics, select_threshold
|
||||||
|
from src.dxa.preprocess import PreprocessConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger("dxa.train")
|
||||||
|
|
||||||
|
|
||||||
def get_device():
|
def get_device(prefer: Optional[str] = None) -> str:
|
||||||
"""Get best available device"""
|
"""Выбрать устройство: явно заданное или лучшее из доступных."""
|
||||||
|
if prefer:
|
||||||
|
return prefer
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
return "cuda"
|
||||||
if torch.backends.mps.is_available():
|
if torch.backends.mps.is_available():
|
||||||
return 'mps'
|
return "mps"
|
||||||
elif torch.cuda.is_available():
|
return "cpu"
|
||||||
return 'cuda'
|
|
||||||
else:
|
|
||||||
return 'cpu'
|
|
||||||
|
|
||||||
|
|
||||||
def compute_metrics(preds, labels):
|
def set_seed(seed: int) -> None:
|
||||||
|
"""Зафиксировать генераторы для воспроизводимости результата."""
|
||||||
|
random.seed(seed)
|
||||||
|
np.random.seed(seed)
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.manual_seed_all(seed)
|
||||||
|
|
||||||
|
|
||||||
|
def make_loader(
|
||||||
|
dataset: DXADataset,
|
||||||
|
batch_size: int,
|
||||||
|
shuffle: bool,
|
||||||
|
num_workers: int,
|
||||||
|
balance: str = "none",
|
||||||
|
seed: int = 42,
|
||||||
|
):
|
||||||
"""
|
"""
|
||||||
Вычисление метрик классификации.
|
Собрать DataLoader.
|
||||||
|
|
||||||
Рассчитывает:
|
balance='sampler' выравнивает классы взвешенной выборкой: важно при доле
|
||||||
- Accuracy: доля правильных предсказаний
|
брака ~15 %, иначе модель сходится к «всё хорошее» и даёт высокую accuracy
|
||||||
- Precision: точность (доля TP среди предсказанных positive)
|
при нулевом recall.
|
||||||
- Recall: полнота (доля TP среди реальных positive)
|
|
||||||
- F1: гармоническое среднее precision и recall
|
|
||||||
|
|
||||||
Args:
|
|
||||||
preds: Предсказания модели (numpy array)
|
|
||||||
labels: Истинные метки (numpy array)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dict с метриками
|
|
||||||
"""
|
"""
|
||||||
preds = np.array(preds)
|
sampler = None
|
||||||
labels = np.array(labels)
|
if balance == "sampler" and shuffle:
|
||||||
|
labels = np.array(dataset.labels)
|
||||||
|
class_counts = np.bincount(labels, minlength=2).astype(np.float64)
|
||||||
|
class_counts[class_counts == 0] = 1.0
|
||||||
|
weights = 1.0 / class_counts[labels]
|
||||||
|
generator = torch.Generator().manual_seed(seed)
|
||||||
|
sampler = WeightedRandomSampler(
|
||||||
|
weights=torch.as_tensor(weights, dtype=torch.double),
|
||||||
|
num_samples=len(dataset),
|
||||||
|
replacement=True,
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
shuffle = False
|
||||||
|
|
||||||
# Accuracy
|
return torch.utils.data.DataLoader(
|
||||||
accuracy = (preds == labels).mean()
|
dataset,
|
||||||
|
batch_size=batch_size,
|
||||||
# True/False positives/negatives
|
shuffle=shuffle,
|
||||||
tp = ((preds == 1) & (labels == 1)).sum()
|
sampler=sampler,
|
||||||
tn = ((preds == 0) & (labels == 0)).sum()
|
num_workers=num_workers,
|
||||||
fp = ((preds == 1) & (labels == 0)).sum()
|
pin_memory=torch.cuda.is_available(),
|
||||||
fn = ((preds == 0) & (labels == 1)).sum()
|
)
|
||||||
|
|
||||||
# Precision, Recall, F1
|
|
||||||
precision = tp / (tp + fp) if (tp + fp) > 0 else 0
|
|
||||||
recall = tp / (tp + fn) if (tp + fn) > 0 else 0
|
|
||||||
f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0
|
|
||||||
|
|
||||||
return {
|
|
||||||
'accuracy': accuracy,
|
|
||||||
'precision': precision,
|
|
||||||
'recall': recall,
|
|
||||||
'f1': f1,
|
|
||||||
'tp': tp,
|
|
||||||
'tn': tn,
|
|
||||||
'fp': fp,
|
|
||||||
'fn': fn
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def train(args):
|
def positive_weight(dataset: DXADataset) -> Optional[float]:
|
||||||
|
"""Вес класса «нарушение» для CrossEntropyLoss = n_good / n_bad."""
|
||||||
|
labels = np.array(dataset.labels)
|
||||||
|
n_bad = int((labels == QUALITY_BAD).sum())
|
||||||
|
n_good = int((labels != QUALITY_BAD).sum())
|
||||||
|
if n_bad == 0 or n_good == 0:
|
||||||
|
return None
|
||||||
|
return n_good / n_bad
|
||||||
|
|
||||||
|
|
||||||
|
def is_improvement(metrics: Dict, best_auc: float, best_f1: float, min_delta: float = 1e-4) -> bool:
|
||||||
"""
|
"""
|
||||||
Основной цикл обучения модели.
|
Улучшение чекпоинта: сначала ROC-AUC, затем F1.
|
||||||
|
|
||||||
Этапы:
|
ROC-AUC выбран первичным критерием отбора, потому что на валидации всего
|
||||||
1. Определение устройства (MPS/CUDA/CPU)
|
~8 изображений с нарушением, и F1 принимает лишь несколько значений —
|
||||||
2. Создание директории для сохранения модели
|
выбор эпохи по F1 шумит и переобучает порог. AUC использует ранжирование
|
||||||
3. Загрузка данных (DataLoader)
|
всех изображений и заметно стабильнее. Сам порог всё равно подбирается
|
||||||
4. Создание модели
|
по F1 (см. `select_threshold`).
|
||||||
5. Цикл обучения по эпохам:
|
|
||||||
- Обучение на train set
|
|
||||||
- Валидация на val set
|
|
||||||
- Расчет метрик (accuracy, precision, recall, F1)
|
|
||||||
- Сохранение лучшей модели по F1
|
|
||||||
6. Сохранение финальной модели
|
|
||||||
|
|
||||||
Args:
|
|
||||||
args: Аргументы командной строки
|
|
||||||
"""
|
"""
|
||||||
|
auc = metrics.get("roc_auc") or 0.0
|
||||||
# Setup
|
if auc > best_auc + min_delta:
|
||||||
device = get_device()
|
return True
|
||||||
print(f"Using device: {device}")
|
if abs(auc - best_auc) <= min_delta and metrics["f1"] > best_f1 + min_delta:
|
||||||
|
return True
|
||||||
# Create output directory
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def smoothed_score(recent_auc: "deque", window: int) -> float:
|
||||||
|
"""
|
||||||
|
Сглаженная оценка для отбора чекпоинта.
|
||||||
|
|
||||||
|
Валидация мала (единицы исследований, ~8 нарушений), поэтому AUC отдельной
|
||||||
|
эпохи почти случаен: без сглаживания лучшей «эпохой» оказывается первая
|
||||||
|
удачная, а сохранённая модель остаётся недоученной (голова области не
|
||||||
|
успевает обучиться). Скользящее среднее по последним `window` эпохам
|
||||||
|
устойчивее и выбирает состоявшуюся модель.
|
||||||
|
"""
|
||||||
|
values = [v for v in list(recent_auc)[-window:] if v is not None]
|
||||||
|
if not values:
|
||||||
|
return 0.0
|
||||||
|
return float(np.mean(values))
|
||||||
|
|
||||||
|
|
||||||
|
def train(args: argparse.Namespace) -> Dict:
|
||||||
|
"""Полный цикл обучения. Возвращает итоговые метрики."""
|
||||||
|
set_seed(args.seed)
|
||||||
|
device = get_device(args.device)
|
||||||
|
logger.info("Device: %s", device)
|
||||||
|
|
||||||
output_dir = Path(args.output_dir)
|
output_dir = Path(args.output_dir)
|
||||||
output_dir.mkdir(parents=True, exist_ok=True)
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Create dataloaders
|
preprocess = PreprocessConfig(norm=args.norm, imagenet_norm=not args.no_imagenet_norm)
|
||||||
print("Loading data...")
|
logger.info("Loading dataset from %s", args.data_root)
|
||||||
train_loader, val_loader = create_dataloaders(
|
train_ds, val_ds, preprocess = make_datasets(
|
||||||
data_root=args.data_root,
|
data_root=args.data_root,
|
||||||
annotation_path=args.annotation_path,
|
annotation_path=args.annotation_path,
|
||||||
batch_size=args.batch_size,
|
input_size=args.input_size,
|
||||||
input_size=(args.input_size, args.input_size),
|
val_fraction=args.val_fraction,
|
||||||
num_workers=args.num_workers
|
seed=args.seed,
|
||||||
|
preprocess=preprocess,
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"Train samples: {len(train_loader.dataset)}")
|
if len(train_ds) == 0 or len(val_ds) == 0:
|
||||||
print(f"Val samples: {len(val_loader.dataset)}")
|
raise RuntimeError(
|
||||||
|
f"Empty split: train={len(train_ds)}, val={len(val_ds)}. Check --data-root and labels."
|
||||||
if len(train_loader.dataset) == 0:
|
)
|
||||||
print("ERROR: No training samples found!")
|
|
||||||
return
|
logger.info("Train split:\n%s", format_summary(train_ds.records))
|
||||||
|
logger.info("Val split:\n%s", format_summary(val_ds.records))
|
||||||
# Create model
|
|
||||||
print(f"Creating model: {args.backbone}")
|
train_loader = make_loader(train_ds, args.batch_size, True, args.num_workers, args.balance, args.seed)
|
||||||
|
val_loader = make_loader(val_ds, args.batch_size, False, args.num_workers, "none", args.seed)
|
||||||
|
|
||||||
|
pos_weight = None
|
||||||
|
if args.balance == "loss":
|
||||||
|
pos_weight = positive_weight(train_ds)
|
||||||
|
logger.info("Loss class weight for violations: %.3f", pos_weight or 1.0)
|
||||||
|
|
||||||
model = create_model(
|
model = create_model(
|
||||||
backbone=args.backbone,
|
backbone=args.backbone,
|
||||||
num_classes=2,
|
pretrained=not args.no_pretrained,
|
||||||
pretrained=True,
|
device=device,
|
||||||
device=device
|
head=args.head,
|
||||||
|
learning_rate=args.learning_rate,
|
||||||
|
weight_decay=args.weight_decay,
|
||||||
|
region_loss_weight=args.region_loss_weight,
|
||||||
|
pos_weight=pos_weight,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Training loop
|
|
||||||
best_val_f1 = 0
|
|
||||||
best_epoch = 0
|
|
||||||
|
|
||||||
for epoch in range(args.epochs):
|
|
||||||
print(f"\n{'='*50}")
|
|
||||||
print(f"Epoch {epoch+1}/{args.epochs}")
|
|
||||||
print(f"{'='*50}")
|
|
||||||
|
|
||||||
# Train
|
|
||||||
train_loss, train_acc = model.train_epoch(train_loader)
|
|
||||||
|
|
||||||
# Validate
|
|
||||||
val_loss, val_acc = model.validate(val_loader)
|
|
||||||
|
|
||||||
# Compute detailed metrics
|
|
||||||
model.model.eval()
|
|
||||||
all_preds = []
|
|
||||||
all_labels = []
|
|
||||||
|
|
||||||
with torch.no_grad():
|
threshold = 0.0
|
||||||
for images, labels in val_loader:
|
best_score, best_f1, best_epoch = -1.0, -1.0, 0
|
||||||
images = images.to(device)
|
epochs_without_improvement = 0
|
||||||
# labels is a dict with 'label' key from our dataset
|
history: List[Dict] = []
|
||||||
if isinstance(labels, dict):
|
recent_auc: deque = deque(maxlen=args.select_window)
|
||||||
labels_arr = labels['label'].to(device)
|
|
||||||
else:
|
|
||||||
labels_arr = labels.to(device)
|
|
||||||
|
|
||||||
preds, _ = model.predict(images)
|
# Фаза 1: backbone заморожен, обучается только голова. При ~250 уникальных
|
||||||
all_preds.extend(preds.cpu().numpy())
|
# снимках полный fine-tune даёт переобучение (train F1 -> 1.0, val AUC ~0.5),
|
||||||
all_labels.extend(labels_arr.cpu().numpy())
|
# а линейный зонд на признаках ImageNet держит val AUC ~0.80.
|
||||||
|
# freeze_epochs = -1 означает «заморозить навсегда».
|
||||||
metrics = compute_metrics(all_preds, all_labels)
|
freeze_epochs = args.epochs if args.freeze_epochs < 0 else args.freeze_epochs
|
||||||
|
if freeze_epochs > 0:
|
||||||
print(f"\nResults:")
|
model.set_backbone_trainable(False)
|
||||||
print(f" Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}")
|
# Стандартизация признаков по обучающей выборке: без неё логиты смещены,
|
||||||
print(f" Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}")
|
# вероятности скучены у нуля и подобранный порог теряет смысл.
|
||||||
print(f" Val Metrics:")
|
model.fit_feature_norm(train_loader)
|
||||||
print(f" Precision: {metrics['precision']:.4f}")
|
logger.info("Phase 1 (epochs 1..%d): backbone frozen, head only, lr=%.1e",
|
||||||
print(f" Recall: {metrics['recall']:.4f}")
|
freeze_epochs, args.learning_rate)
|
||||||
print(f" F1: {metrics['f1']:.4f}")
|
|
||||||
|
for epoch in range(1, args.epochs + 1):
|
||||||
# Save best model
|
if epoch == freeze_epochs + 1 and 0 < freeze_epochs < args.epochs:
|
||||||
if metrics['f1'] > best_val_f1:
|
model.set_backbone_trainable(True)
|
||||||
best_val_f1 = metrics['f1']
|
# При размораживании backbone нужен меньший шаг, иначе предобученные
|
||||||
best_epoch = epoch + 1
|
# признаки разрушаются за несколько эпох.
|
||||||
model.save(str(output_dir / 'best_model.pth'))
|
model.set_learning_rate(args.learning_rate / 10)
|
||||||
print(f" ✅ Saved best model (F1: {best_val_f1:.4f})")
|
logger.info("Phase 2: backbone unfrozen, lr=%.1e", args.learning_rate / 10)
|
||||||
|
|
||||||
# Save checkpoint
|
started = time.time()
|
||||||
if (epoch + 1) % args.save_every == 0:
|
train_loss, train_metrics = model.train_epoch(train_loader)
|
||||||
model.save(str(output_dir / f'checkpoint_epoch_{epoch+1}.pth'))
|
val_loss, _ = model.validate(val_loader)
|
||||||
|
|
||||||
print(f"\n{'='*50}")
|
# Порог по логитам подбирается на валидации на каждой эпохе: при доле
|
||||||
print(f"Training complete!")
|
# брака ~15 % порог 0 (вероятность 0.5) почти всегда даёт нулевой recall.
|
||||||
print(f"Best F1: {best_val_f1:.4f} at epoch {best_epoch}")
|
logits, labels = model.predict_logits(val_loader)
|
||||||
print(f"{'='*50}")
|
threshold, val_metrics = select_threshold(logits, labels, min_recall=args.min_recall)
|
||||||
|
elapsed = time.time() - started
|
||||||
# Save final model
|
|
||||||
model.save(str(output_dir / 'final_model.pth'))
|
model.history["train_loss"].append(train_loss)
|
||||||
print(f"Final model saved to {output_dir / 'final_model.pth'}")
|
model.history["val_loss"].append(val_loss)
|
||||||
|
model.history["val_f1"].append(val_metrics["f1"])
|
||||||
|
model.history["val_roc_auc"].append(val_metrics.get("roc_auc"))
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Epoch %3d/%d | train loss %.4f f1 %.3f | val loss %.4f acc %.3f prec %.3f rec %.3f "
|
||||||
|
"f1 %.3f auc %s thr %.3f | %.1fs",
|
||||||
|
epoch, args.epochs, train_loss, train_metrics["f1"], val_loss,
|
||||||
|
val_metrics["accuracy"], val_metrics["precision"], val_metrics["recall"],
|
||||||
|
val_metrics["f1"],
|
||||||
|
f"{val_metrics['roc_auc']:.3f}" if val_metrics["roc_auc"] is not None else "n/a",
|
||||||
|
val_metrics["threshold_prob"], elapsed,
|
||||||
|
)
|
||||||
|
|
||||||
|
history.append({"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss,
|
||||||
|
"threshold_logit": threshold, "threshold_prob": val_metrics["threshold_prob"],
|
||||||
|
**{f"val_{k}": v for k, v in val_metrics.items() if isinstance(v, (int, float))}})
|
||||||
|
|
||||||
|
recent_auc.append(val_metrics.get("roc_auc"))
|
||||||
|
score = smoothed_score(recent_auc, args.select_window)
|
||||||
|
# Прогрев: до накопления окна чекпоинт не сохраняется, иначе им станет
|
||||||
|
# случайно удачная ранняя эпоха с ещё не обученной моделью.
|
||||||
|
warmed_up = epoch >= args.select_window
|
||||||
|
|
||||||
|
if warmed_up and score > best_score + 1e-4:
|
||||||
|
best_score = score
|
||||||
|
best_f1 = val_metrics["f1"]
|
||||||
|
best_epoch = epoch
|
||||||
|
epochs_without_improvement = 0
|
||||||
|
model.save(
|
||||||
|
output_dir / "dxa_model.pth",
|
||||||
|
preprocess=preprocess,
|
||||||
|
threshold=threshold,
|
||||||
|
val_metrics=val_metrics,
|
||||||
|
selection_score=score,
|
||||||
|
epoch=epoch,
|
||||||
|
data_root=str(args.data_root),
|
||||||
|
)
|
||||||
|
logger.info(" -> saved best checkpoint (smoothed auc %.4f, f1 %.4f, thr %.3f)",
|
||||||
|
score, best_f1, val_metrics["threshold_prob"])
|
||||||
|
else:
|
||||||
|
epochs_without_improvement += 1
|
||||||
|
if warmed_up and args.patience and epochs_without_improvement >= args.patience:
|
||||||
|
logger.info("Early stopping after %d epochs without improvement", epochs_without_improvement)
|
||||||
|
break
|
||||||
|
|
||||||
|
# Финальная валидация лучшим чекпоинтом, а не последним.
|
||||||
|
best_path = output_dir / "dxa_model.pth"
|
||||||
|
if best_path.exists():
|
||||||
|
metadata = model.load(best_path)
|
||||||
|
threshold = float(metadata.get("threshold", threshold))
|
||||||
|
else:
|
||||||
|
model.save(output_dir / "dxa_model.pth", preprocess=preprocess, threshold=threshold)
|
||||||
|
|
||||||
|
model.save(
|
||||||
|
output_dir / "dxa_model_final.pth",
|
||||||
|
preprocess=preprocess,
|
||||||
|
threshold=threshold,
|
||||||
|
epoch=epoch,
|
||||||
|
)
|
||||||
|
|
||||||
|
best_logits, labels, predicted_regions = model.predict_logits_with_regions(val_loader)
|
||||||
|
best_threshold, best_metrics = select_threshold(best_logits, labels, min_recall=args.min_recall)
|
||||||
|
region_metrics = per_region_metrics(best_logits, labels, predicted_regions, best_threshold)
|
||||||
|
|
||||||
|
report = {
|
||||||
|
"backbone": args.backbone,
|
||||||
|
"device": device,
|
||||||
|
"preprocess": preprocess.to_dict(),
|
||||||
|
"seed": args.seed,
|
||||||
|
"balance": args.balance,
|
||||||
|
"pos_weight": pos_weight,
|
||||||
|
"epochs_run": epoch,
|
||||||
|
"best_epoch": best_epoch,
|
||||||
|
"threshold": best_threshold,
|
||||||
|
"train_size": len(train_ds),
|
||||||
|
"val_size": len(val_ds),
|
||||||
|
"train_studies": len({r.study for r in train_ds.records}),
|
||||||
|
"val_studies": len({r.study for r in val_ds.records}),
|
||||||
|
"best_val_metrics": {k: v for k, v in best_metrics.items()},
|
||||||
|
"val_metrics_per_region": region_metrics,
|
||||||
|
"history": history,
|
||||||
|
}
|
||||||
|
(output_dir / "train_report.json").write_text(
|
||||||
|
json.dumps(_to_builtin(report), ensure_ascii=False, indent=2)
|
||||||
|
)
|
||||||
|
_write_markdown_report(output_dir / "train_report.md", report)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Training complete. Best epoch %d: f1 %.4f (thr %.3f), roc_auc %s, pr_auc %s",
|
||||||
|
best_epoch, best_metrics["f1"], best_threshold,
|
||||||
|
f"{best_metrics['roc_auc']:.4f}" if best_metrics["roc_auc"] is not None else "n/a",
|
||||||
|
f"{best_metrics['pr_auc']:.4f}" if best_metrics["pr_auc"] is not None else "n/a",
|
||||||
|
)
|
||||||
|
logger.info("Checkpoint: %s", best_path)
|
||||||
|
logger.info("Report: %s", output_dir / "train_report.md")
|
||||||
|
return report
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def _to_builtin(value):
|
||||||
parser = argparse.ArgumentParser(description='Train DXA Quality Classifier')
|
"""Привести значения numpy к встроенным типам, чтобы отчёт сериализовался в JSON."""
|
||||||
|
if isinstance(value, dict):
|
||||||
# Data arguments
|
return {k: _to_builtin(v) for k, v in value.items()}
|
||||||
parser.add_argument('--data-root', type=str,
|
if isinstance(value, (list, tuple)):
|
||||||
default='dataset_hack',
|
return [_to_builtin(v) for v in value]
|
||||||
help='Path to data directory')
|
if isinstance(value, (np.integer,)):
|
||||||
parser.add_argument('--annotation-path', type=str,
|
return int(value)
|
||||||
default='dataset_hack/НД_для_обучения/разметка.xlsx',
|
if isinstance(value, (np.floating,)):
|
||||||
help='Path to annotation Excel file')
|
return float(value)
|
||||||
parser.add_argument('--input-size', type=int, default=224,
|
if isinstance(value, np.bool_):
|
||||||
help='Input image size')
|
return bool(value)
|
||||||
parser.add_argument('--batch-size', type=int, default=8,
|
return value
|
||||||
help='Batch size')
|
|
||||||
parser.add_argument('--num-workers', type=int, default=4,
|
|
||||||
help='Number of data loading workers')
|
|
||||||
|
|
||||||
# Model arguments
|
|
||||||
parser.add_argument('--backbone', type=str, default='resnet18',
|
|
||||||
choices=['resnet18', 'resnet34', 'efficientnet_b0'],
|
|
||||||
help='Backbone architecture')
|
|
||||||
parser.add_argument('--epochs', type=int, default=20,
|
|
||||||
help='Number of training epochs')
|
|
||||||
parser.add_argument('--learning-rate', type=float, default=1e-4,
|
|
||||||
help='Learning rate')
|
|
||||||
|
|
||||||
# Output arguments
|
|
||||||
parser.add_argument('--output-dir', type=str, default='models',
|
|
||||||
help='Output directory for models')
|
|
||||||
parser.add_argument('--save-every', type=int, default=5,
|
|
||||||
help='Save checkpoint every N epochs')
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
train(args)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
def _write_markdown_report(path: Path, report: Dict) -> None:
|
||||||
main()
|
"""Краткий отчёт об обучении для документации и презентации."""
|
||||||
|
m = report["best_val_metrics"]
|
||||||
|
lines = [
|
||||||
|
"# Отчёт об обучении классификатора качества DXA",
|
||||||
|
"",
|
||||||
|
f"- Backbone: `{report['backbone']}`",
|
||||||
|
f"- Устройство: `{report['device']}`",
|
||||||
|
f"- Seed: {report['seed']}, балансировка классов: `{report['balance']}`",
|
||||||
|
f"- Размер выборок: train {report['train_size']} снимков / {report['train_studies']} исследований, "
|
||||||
|
f"val {report['val_size']} снимков / {report['val_studies']} исследований (разбиение по исследованиям)",
|
||||||
|
f"- Предобработка: `{report['preprocess']}`",
|
||||||
|
f"- Обучено эпох: {report['epochs_run']}, лучшая эпоха: {report['best_epoch']}",
|
||||||
|
f"- Рабочий порог (подобран по F1 на валидации): логит {report['threshold']:.4f} "
|
||||||
|
f"(вероятность {m['threshold_prob']:.4f})",
|
||||||
|
"",
|
||||||
|
"## Метрики на валидации (лучший чекпоинт)",
|
||||||
|
"",
|
||||||
|
"| Метрика | Значение |",
|
||||||
|
"|---|---|",
|
||||||
|
f"| Accuracy | {m['accuracy']:.4f} |",
|
||||||
|
f"| Precision | {m['precision']:.4f} |",
|
||||||
|
f"| Recall | {m['recall']:.4f} |",
|
||||||
|
f"| F1 | {m['f1']:.4f} |",
|
||||||
|
f"| ROC-AUC | {m['roc_auc']:.4f} |" if m.get("roc_auc") is not None else "| ROC-AUC | n/a |",
|
||||||
|
f"| PR-AUC | {m['pr_auc']:.4f} |" if m.get("pr_auc") is not None else "| PR-AUC | n/a |",
|
||||||
|
f"| Порог (логит / вероятность) | {m['threshold_logit']:.4f} / {m['threshold_prob']:.4f} |",
|
||||||
|
f"| TP/TN/FP/FN | {m['tp']}/{m['tn']}/{m['fp']}/{m['fn']} |",
|
||||||
|
"",
|
||||||
|
"> Валидация невелика (единицы исследований), поэтому метрики имеют широкий "
|
||||||
|
"доверительный интервал и не заменяют оценку на закрытом наборе.",
|
||||||
|
"",
|
||||||
|
]
|
||||||
|
|
||||||
|
region_metrics = report.get("val_metrics_per_region") or {}
|
||||||
|
if region_metrics:
|
||||||
|
lines += [
|
||||||
|
"## Метрики по анатомическим областям",
|
||||||
|
"",
|
||||||
|
"| Область | n | Нарушений | ROC-AUC | F1 | Recall |",
|
||||||
|
"|---|---|---|---|---|---|",
|
||||||
|
]
|
||||||
|
for region, rm in region_metrics.items():
|
||||||
|
auc = f"{rm['roc_auc']:.4f}" if rm.get("roc_auc") is not None else "n/a"
|
||||||
|
lines.append(
|
||||||
|
f"| {region} | {rm['n']} | {rm['n_pos']} | {auc} | {rm['f1']:.4f} | {rm['recall']:.4f} |"
|
||||||
|
)
|
||||||
|
lines += [
|
||||||
|
"",
|
||||||
|
"> Нарушения в наборе распределены крайне неравномерно (в позвоночнике ~29 % снимков "
|
||||||
|
"против ~4–5 % у бёдер), а анатомическая область почти однозначно определяется по "
|
||||||
|
"размеру кадра. Поэтому общий AUC частично отражает различение области, а не только "
|
||||||
|
"распознавание дефекта: сопоставляйте общий показатель со значениями по областям.",
|
||||||
|
"",
|
||||||
|
]
|
||||||
|
|
||||||
|
lines += [
|
||||||
|
"## Ограничения и следующий шаг",
|
||||||
|
"",
|
||||||
|
"- Единицей разметки в исходных данных было исследование; снимок наследует метку "
|
||||||
|
"исследования, поэтому часть меток заведомо шумная.",
|
||||||
|
"- Область предсказывает вспомогательная голова; при ошибке области тип нарушения "
|
||||||
|
"тоже будет определён неверно.",
|
||||||
|
"- Для устойчивой оценки нужен больший набор и независимый тест: на 20–50 снимках "
|
||||||
|
"валидации доверительные интервалы метрик очень широкие.",
|
||||||
|
"",
|
||||||
|
]
|
||||||
|
path.write_text("\n".join(lines))
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Обучение классификатора качества DXA",
|
||||||
|
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||||
|
)
|
||||||
|
|
||||||
|
data = parser.add_argument_group("данные")
|
||||||
|
data.add_argument("--data-root", default="dataset_hack",
|
||||||
|
help="Каталог датасета (dataset_hack, НД_для_обучения или каталог с DICOM)")
|
||||||
|
data.add_argument("--annotation-path", default="dataset_hack/НД_для_обучения/разметка.xlsx",
|
||||||
|
help="Excel-разметка; применяется только для отчёта о расхождениях")
|
||||||
|
data.add_argument("--input-size", type=int, default=224)
|
||||||
|
data.add_argument("--val-fraction", type=float, default=0.2,
|
||||||
|
help="Доля изображений в валидации (разбиение по исследованиям)")
|
||||||
|
data.add_argument("--num-workers", type=int, default=0)
|
||||||
|
|
||||||
|
model = parser.add_argument_group("модель")
|
||||||
|
model.add_argument("--backbone", default="resnet18", choices=["resnet18", "resnet34"])
|
||||||
|
model.add_argument("--head", default="linear", choices=["linear", "mlp"],
|
||||||
|
help="linear — линейный зонд на признаках ImageNet (устойчив к малой выборке); "
|
||||||
|
"mlp — двухслойная голова, требует больше данных")
|
||||||
|
model.add_argument("--no-pretrained", action="store_true",
|
||||||
|
help="Обучать с нуля, без весов ImageNet")
|
||||||
|
model.add_argument("--region-loss-weight", type=float, default=0.3,
|
||||||
|
help="Вес вспомогательной головы анатомической области")
|
||||||
|
model.add_argument("--norm", default="percentile", choices=["percentile", "minmax"],
|
||||||
|
help="Способ нормировки интенсивностей DICOM")
|
||||||
|
model.add_argument("--no-imagenet-norm", action="store_true",
|
||||||
|
help="Не применять нормировку ImageNet (устаревший режим)")
|
||||||
|
|
||||||
|
opt = parser.add_argument_group("оптимизация")
|
||||||
|
opt.add_argument("--epochs", type=int, default=100)
|
||||||
|
opt.add_argument("--batch-size", type=int, default=16)
|
||||||
|
opt.add_argument("--learning-rate", type=float, default=3e-4)
|
||||||
|
# Заметный weight decay нужен линейной голове не только против переобучения:
|
||||||
|
# без него логиты за 100 эпох насыщаются, вероятности уходят в 0/1 и
|
||||||
|
# подобранный порог вырождается в 1.0. При wd=0.05 порог держится ~0.5.
|
||||||
|
opt.add_argument("--weight-decay", type=float, default=5e-2)
|
||||||
|
opt.add_argument("--balance", default="none", choices=["loss", "sampler", "none"],
|
||||||
|
help="Компенсация дисбаланса классов; при линейной голове помогает "
|
||||||
|
"подбор порога, а pos_weight скорее вредит")
|
||||||
|
opt.add_argument("--min-recall", type=float, default=0.0,
|
||||||
|
help="Нижняя граница recall при подборе порога (0 — только максимум F1)")
|
||||||
|
opt.add_argument("--patience", type=int, default=25,
|
||||||
|
help="Ранняя остановка: эпох без улучшения (0 — отключено)")
|
||||||
|
opt.add_argument("--select-window", type=int, default=5,
|
||||||
|
help="Окно сглаживания при отборе чекпоинта: обучение до накопления окна "
|
||||||
|
"не сохраняется, оценка усредняется по последним эпохам")
|
||||||
|
opt.add_argument("--freeze-epochs", type=int, default=-1,
|
||||||
|
help="Сколько первых эпох обучать только голову при замороженном backbone. "
|
||||||
|
"-1 — backbone заморожен всегда (линейный зонд, режим по умолчанию); "
|
||||||
|
"0 — обучать всю сеть")
|
||||||
|
opt.add_argument("--seed", type=int, default=42)
|
||||||
|
opt.add_argument("--device", default=None, help="cpu / cuda / mps (по умолчанию — лучший доступный)")
|
||||||
|
opt.add_argument("--output-dir", default="models")
|
||||||
|
|
||||||
|
parser.add_argument("--dry-run", action="store_true",
|
||||||
|
help="Только проверить разбор данных и разбиение, без обучения")
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Optional[List[str]] = None) -> int:
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||||
|
args = build_parser().parse_args(argv)
|
||||||
|
|
||||||
|
if args.dry_run:
|
||||||
|
set_seed(args.seed)
|
||||||
|
train_ds, val_ds, cfg = make_datasets(
|
||||||
|
data_root=args.data_root,
|
||||||
|
annotation_path=args.annotation_path,
|
||||||
|
input_size=args.input_size,
|
||||||
|
val_fraction=args.val_fraction,
|
||||||
|
seed=args.seed,
|
||||||
|
preprocess=PreprocessConfig(norm=args.norm),
|
||||||
|
)
|
||||||
|
print(f"Preprocess: {cfg.to_dict()}")
|
||||||
|
print("\nTrain split:\n" + format_summary(train_ds.records))
|
||||||
|
print("\nVal split:\n" + format_summary(val_ds.records))
|
||||||
|
overlap = {r.study for r in train_ds.records} & {r.study for r in val_ds.records}
|
||||||
|
print(f"\nStudies leaking between train and val: {len(overlap)}")
|
||||||
|
return 0
|
||||||
|
|
||||||
|
try:
|
||||||
|
train(args)
|
||||||
|
except RuntimeError as exc:
|
||||||
|
logger.error("Training failed: %s", exc)
|
||||||
|
return 1
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
|
|
|
||||||
530
src/main.py
530
src/main.py
|
|
@ -20,11 +20,13 @@ FastAPI сервер для оценки качества DXA исследова
|
||||||
- POST /api/v1/batch - Пакетный анализ
|
- POST /api/v1/batch - Пакетный анализ
|
||||||
- POST /api/v1/export - Анализ и экспорт в XLSX
|
- POST /api/v1/export - Анализ и экспорт в XLSX
|
||||||
"""
|
"""
|
||||||
from fastapi import FastAPI, File, UploadFile, APIRouter, Query
|
import os
|
||||||
|
import time
|
||||||
|
from fastapi import FastAPI, File, UploadFile, Query
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from typing import Dict
|
from typing import Dict, Optional
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pydicom
|
import pydicom
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
@ -35,11 +37,13 @@ import base64
|
||||||
|
|
||||||
from starlette.responses import JSONResponse, FileResponse
|
from starlette.responses import JSONResponse, FileResponse
|
||||||
|
|
||||||
from src.dxa.model import create_model
|
from src.dxa.inference import (
|
||||||
from src.dxa.inference import determine_anatomical_region
|
REASON_BY_TYPE,
|
||||||
|
load_model as load_dxa_checkpoint,
|
||||||
|
predict_from_bytes,
|
||||||
|
)
|
||||||
from src.utils.utils import get_device
|
from src.utils.utils import get_device
|
||||||
from src.quality.detailed_assessment import generate_quality_report
|
from src.quality.quality_scorer import convert_to_serializable
|
||||||
from src.quality.quality_scorer import QualityScorer, convert_to_serializable
|
|
||||||
|
|
||||||
# Create app
|
# Create app
|
||||||
app = FastAPI(
|
app = FastAPI(
|
||||||
|
|
@ -56,36 +60,40 @@ static_dir = Path(__file__).parent / "api/static"
|
||||||
if static_dir.exists():
|
if static_dir.exists():
|
||||||
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
||||||
|
|
||||||
# Global model
|
# Global model state: модель загружается один раз и переиспользуется запросами.
|
||||||
|
MODEL_PATH = os.environ.get("DXA_MODEL_PATH", "models/dxa_model.pth")
|
||||||
dxa_model = None
|
dxa_model = None
|
||||||
|
dxa_preprocess = None
|
||||||
|
dxa_threshold = 0.0
|
||||||
device = None
|
device = None
|
||||||
|
|
||||||
|
|
||||||
def load_model():
|
def load_model():
|
||||||
"""
|
"""
|
||||||
Загрузка модели классификатора качества DXA.
|
Загрузка модели классификатора качества DXA.
|
||||||
|
|
||||||
Модель загружается глобально при первом запросе и сохраняется в памяти.
|
Модель загружается глобально при первом запросе и сохраняется в памяти.
|
||||||
Это позволяет избежать повторной загрузки при каждом запросе.
|
Путь к чекпоинту берётся из переменной окружения DXA_MODEL_PATH, чтобы
|
||||||
|
контейнер не зависел от текущего рабочего каталога.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
DXAQualityModel: Обученная модель или None при ошибке
|
DXAQualityModel: Обученная модель или None при ошибке
|
||||||
"""
|
"""
|
||||||
global dxa_model, device
|
global dxa_model, dxa_preprocess, dxa_threshold, device
|
||||||
|
|
||||||
if dxa_model is None:
|
if dxa_model is None:
|
||||||
device = get_device()
|
device = get_device()
|
||||||
print(f"Loading DXA model on {device}...")
|
print(f"Loading DXA model on {device}...")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
dxa_model = create_model(
|
checkpoint = load_dxa_checkpoint(MODEL_PATH, device=device)
|
||||||
backbone='resnet18',
|
dxa_model = checkpoint.model
|
||||||
pretrained=False,
|
dxa_preprocess = checkpoint.preprocess
|
||||||
device=device
|
dxa_threshold = checkpoint.threshold
|
||||||
|
print(
|
||||||
|
f"DXA model loaded successfully (threshold logit={dxa_threshold:.4f}, "
|
||||||
|
f"input={dxa_preprocess.input_size}, preprocess={dxa_preprocess.to_dict()})"
|
||||||
)
|
)
|
||||||
dxa_model.load('models/dxa_model.pth')
|
|
||||||
dxa_model.model.eval()
|
|
||||||
print("DXA model loaded successfully")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error loading model: {e}")
|
print(f"Error loading model: {e}")
|
||||||
dxa_model = None
|
dxa_model = None
|
||||||
|
|
@ -96,49 +104,64 @@ def load_model():
|
||||||
def preprocess_dicom(dcm_bytes: bytes, input_size: int = 224):
|
def preprocess_dicom(dcm_bytes: bytes, input_size: int = 224):
|
||||||
"""
|
"""
|
||||||
Предобработка DICOM изображения для модели.
|
Предобработка DICOM изображения для модели.
|
||||||
|
|
||||||
Этапы:
|
Сохранена для обратной совместимости: эндпоинты используют
|
||||||
1. Сохранение байтов во временный файл
|
`predict_from_bytes`, который применяет ту же предобработку, что и при
|
||||||
2. Чтение DICOM (pydicom)
|
обучении (параметры берутся из чекпоинта, а не задаются заново).
|
||||||
3. Нормализация значений пикселей
|
|
||||||
4. Создание 3-канального изображения
|
|
||||||
5. Изменение размера до input_size x input_size
|
|
||||||
6. Нормализация для PyTorch (деление на 255)
|
|
||||||
7. Преобразование в тензор
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dcm_bytes: Байты DICOM файла
|
dcm_bytes: Байты DICOM файла
|
||||||
input_size: Целевой размер (по умолчанию 224)
|
input_size: Целевой размер (по умолчанию 224)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple[torch.Tensor, pydicom.Dataset]: Тензор изображения и метаданные DICOM
|
Tuple[torch.Tensor, pydicom.Dataset]: Тензор изображения и метаданные DICOM
|
||||||
"""
|
"""
|
||||||
import tempfile
|
from src.dxa.preprocess import PreprocessConfig, preprocess_from_array
|
||||||
|
|
||||||
with tempfile.NamedTemporaryFile(suffix='.dcm', delete=False) as f:
|
ds = pydicom.dcmread(io.BytesIO(dcm_bytes))
|
||||||
f.write(dcm_bytes)
|
|
||||||
dcm_path = f.name
|
|
||||||
|
|
||||||
ds = pydicom.dcmread(dcm_path)
|
|
||||||
img = ds.pixel_array.astype(np.float32)
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
if img.ndim == 3:
|
||||||
|
img = img.mean(axis=0) if img.shape[0] > 1 else img[0]
|
||||||
|
|
||||||
# Normalize
|
tensor = torch.from_numpy(preprocess_from_array(img, PreprocessConfig(input_size=input_size)))
|
||||||
img = (img - img.min()) / (img.max() - img.min() + 1e-8)
|
return tensor.float().unsqueeze(0), ds
|
||||||
|
|
||||||
# 3-channel
|
|
||||||
img = np.stack([img] * 3, axis=0)
|
|
||||||
|
|
||||||
# Resize
|
def predict_upload(dcm_bytes: bytes):
|
||||||
img = (img * 255).astype(np.uint8)
|
"""
|
||||||
img_pil = Image.fromarray(img.transpose(1, 2, 0))
|
Единая точка предсказания для HTTP-эндпоинтов.
|
||||||
img_pil = img_pil.resize((input_size, input_size), Image.BILINEAR)
|
|
||||||
img = np.array(img_pil).transpose(2, 0, 1)
|
|
||||||
img = img.astype(np.float32) / 255.0
|
|
||||||
|
|
||||||
# Tensor
|
Возвращает (prediction, dataset_metadata) или (None, None), если модель не
|
||||||
img = torch.from_numpy(img).float().unsqueeze(0)
|
загружена. Использование общего пути гарантирует, что предобработка и
|
||||||
|
решающий порог совпадают с CLI-инференсом.
|
||||||
|
"""
|
||||||
|
model = load_model()
|
||||||
|
if model is None:
|
||||||
|
return None, None
|
||||||
|
return predict_from_bytes(dcm_bytes, model, dxa_preprocess, device, dxa_threshold)
|
||||||
|
|
||||||
return img, ds
|
|
||||||
|
def prediction_to_result(prediction, ds, filename: Optional[str] = None) -> dict:
|
||||||
|
"""Привести предсказание к формату ответа API."""
|
||||||
|
result = {
|
||||||
|
"study_uid": str(getattr(ds, "StudyInstanceUID", "") or ""),
|
||||||
|
"image_uid": str(getattr(ds, "SOPInstanceUID", "") or ""),
|
||||||
|
"anatomical_region": prediction.anatomical_region,
|
||||||
|
"quality_class": prediction.quality_class,
|
||||||
|
"quality_label": "OK" if prediction.quality_class == 0 else "Violation detected",
|
||||||
|
"violation_type": prediction.violation_type,
|
||||||
|
"reason": prediction.violation_reason,
|
||||||
|
"confidence": round(prediction.prob, 4),
|
||||||
|
"threshold_probability": round(
|
||||||
|
float(torch.sigmoid(torch.tensor(prediction.threshold))), 4
|
||||||
|
),
|
||||||
|
"region_confidence": round(prediction.region_confidence, 4),
|
||||||
|
"metrics": prediction.samples,
|
||||||
|
"processing_status": "Success",
|
||||||
|
}
|
||||||
|
if filename:
|
||||||
|
result["filename"] = filename
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
# API Routes
|
# API Routes
|
||||||
|
|
@ -190,38 +213,15 @@ async def analyze_dicom(file: UploadFile = File(...)):
|
||||||
status_code=500,
|
status_code=500,
|
||||||
content={"error": "Model not loaded"}
|
content={"error": "Model not loaded"}
|
||||||
)
|
)
|
||||||
|
|
||||||
# Read file
|
|
||||||
dcm_bytes = await file.read()
|
dcm_bytes = await file.read()
|
||||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
prediction, ds = predict_upload(dcm_bytes)
|
||||||
|
if prediction is None:
|
||||||
# Predict
|
return JSONResponse(status_code=500, content={"error": "Model not loaded"})
|
||||||
img_tensor = img_tensor.to(device)
|
|
||||||
with torch.no_grad():
|
return prediction_to_result(prediction, ds, file.filename)
|
||||||
outputs = model.model(img_tensor)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
confidence = probs[0, pred].item()
|
|
||||||
|
|
||||||
# Determine region
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
|
||||||
|
|
||||||
# Result
|
|
||||||
result = {
|
|
||||||
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
|
||||||
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
|
||||||
"anatomical_region": region,
|
|
||||||
"quality_class": int(pred),
|
|
||||||
"quality_label": "OK" if pred == 0 else "Violation detected",
|
|
||||||
"confidence": round(confidence, 4),
|
|
||||||
"processing_status": "Success"
|
|
||||||
}
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
import traceback
|
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
status_code=500,
|
status_code=500,
|
||||||
content={
|
content={
|
||||||
|
|
@ -259,65 +259,27 @@ async def analyze_dicom_detailed(
|
||||||
content={"error": "Model not loaded"}
|
content={"error": "Model not loaded"}
|
||||||
)
|
)
|
||||||
|
|
||||||
# Read file
|
|
||||||
dcm_bytes = await file.read()
|
dcm_bytes = await file.read()
|
||||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
prediction, ds = predict_upload(dcm_bytes)
|
||||||
|
if prediction is None:
|
||||||
|
return JSONResponse(status_code=500, content={"error": "Model not loaded"})
|
||||||
|
|
||||||
# Get original image for quality assessment
|
result = prediction_to_result(prediction, ds, file.filename)
|
||||||
img_array = ds.pixel_array.astype(np.float32)
|
result["quality_label"] = (
|
||||||
img_array = (img_array - img_array.min()) / (img_array.max() - img_array.min() + 1e-8)
|
"Качественное изображение" if prediction.quality_class == 0 else "Есть нарушение качества"
|
||||||
|
|
||||||
# Create 3-channel version
|
|
||||||
if len(img_array.shape) == 2:
|
|
||||||
img_3ch = np.stack([img_array] * 3, axis=2)
|
|
||||||
else:
|
|
||||||
img_3ch = img_array
|
|
||||||
|
|
||||||
# Predict
|
|
||||||
img_tensor = img_tensor.to(device)
|
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model.model(img_tensor)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
confidence = probs[0, pred].item()
|
|
||||||
all_probs = probs[0].cpu().numpy().tolist()
|
|
||||||
|
|
||||||
# Determine region
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
|
||||||
|
|
||||||
# Generate segmentation (simple threshold-based for now)
|
|
||||||
# In production, use proper segmentation model
|
|
||||||
threshold = np.percentile(img_array, 90)
|
|
||||||
segmentation = (img_array > threshold).astype(np.uint8)
|
|
||||||
|
|
||||||
# Generate detailed quality report
|
|
||||||
quality_report = generate_quality_report(
|
|
||||||
image=img_3ch,
|
|
||||||
segmentation=segmentation,
|
|
||||||
region=region,
|
|
||||||
model_prediction=pred,
|
|
||||||
model_confidence=confidence
|
|
||||||
)
|
)
|
||||||
|
result["view_quality"] = _view_quality(ds)
|
||||||
|
result["overall_quality"] = "GOOD" if prediction.quality_class == 0 else "POOR"
|
||||||
|
result["reasons"] = _build_reasons(prediction)
|
||||||
|
|
||||||
# Add UID info
|
# Дополнительный функционал: схематичная маска костной ткани
|
||||||
quality_report["study_uid"] = getattr(ds, 'StudyInstanceUID', '')
|
# (90-й перцентиль яркости) как основа для визуализации нарушения.
|
||||||
quality_report["image_uid"] = getattr(ds, 'SOPInstanceUID', '')
|
|
||||||
|
|
||||||
# Add visualization if requested
|
|
||||||
if include_visualization:
|
if include_visualization:
|
||||||
# Create mask visualization
|
result["mask"] = _mask_base64(dcm_bytes)
|
||||||
mask_vis = Image.fromarray((segmentation * 255).astype(np.uint8))
|
|
||||||
buffer = io.BytesIO()
|
|
||||||
mask_vis.save(buffer, format='PNG')
|
|
||||||
quality_report["mask"] = base64.b64encode(buffer.getvalue()).decode('utf-8')
|
|
||||||
else:
|
else:
|
||||||
quality_report["mask"] = None
|
result["mask"] = None
|
||||||
|
|
||||||
# Convert numpy types to Python types for JSON serialization
|
return convert_to_serializable(result)
|
||||||
quality_report = convert_to_serializable(quality_report)
|
|
||||||
|
|
||||||
return quality_report
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
import traceback
|
import traceback
|
||||||
|
|
@ -331,47 +293,68 @@ async def analyze_dicom_detailed(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _view_quality(ds) -> str:
|
||||||
|
"""Полнота видимости области: грубая оценка по числу строк изображения."""
|
||||||
|
try:
|
||||||
|
rows = int(ds.Rows)
|
||||||
|
except Exception:
|
||||||
|
return "unknown"
|
||||||
|
# Снимки бедра в датасете компактнее (≈235–290 строк), позвоночника — выше.
|
||||||
|
return "full" if rows >= 260 else "partial"
|
||||||
|
|
||||||
|
|
||||||
|
def _mask_base64(dcm_bytes: bytes) -> str:
|
||||||
|
"""PNG-маска костной ткани (90-й перцентиль яркости) в base64."""
|
||||||
|
ds = pydicom.dcmread(io.BytesIO(dcm_bytes))
|
||||||
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
if img.ndim == 3:
|
||||||
|
img = img.mean(axis=0) if img.shape[0] > 1 else img[0]
|
||||||
|
norm = (img - img.min()) / (img.max() - img.min() + 1e-8)
|
||||||
|
mask = (norm > np.percentile(norm, 90)).astype(np.uint8) * 255
|
||||||
|
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
Image.fromarray(mask).save(buffer, format="PNG")
|
||||||
|
return base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def _build_reasons(prediction) -> list:
|
||||||
|
"""Пояснения к решению в форме, принятой в методике оценки качества."""
|
||||||
|
if prediction.quality_class == 0:
|
||||||
|
return [
|
||||||
|
f"Область ({prediction.anatomical_region}) видна достаточно полно.",
|
||||||
|
"Значимых артефактов и выраженного размытия не выявлено.",
|
||||||
|
]
|
||||||
|
reasons = [
|
||||||
|
f"Выявлено нарушение: {REASON_BY_TYPE.get(prediction.violation_type, 'нарушение качества')}.",
|
||||||
|
f"Область исследования: {prediction.anatomical_region}.",
|
||||||
|
"Требуется ручная проверка перед дальнейшим анализом.",
|
||||||
|
]
|
||||||
|
return reasons
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/v1/batch")
|
@app.post("/api/v1/batch")
|
||||||
async def batch_analyze(files: list[UploadFile] = File(...)):
|
async def batch_analyze(files: list[UploadFile] = File(...)):
|
||||||
"""Batch analyze multiple DICOM files"""
|
"""Пакетный анализ нескольких DICOM файлов."""
|
||||||
results = []
|
|
||||||
|
|
||||||
model = load_model()
|
model = load_model()
|
||||||
if model is None:
|
if model is None:
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
status_code=500,
|
status_code=500,
|
||||||
content={"error": "Model not loaded"}
|
content={"error": "Model not loaded"}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
results = []
|
||||||
for file in files:
|
for file in files:
|
||||||
try:
|
try:
|
||||||
dcm_bytes = await file.read()
|
dcm_bytes = await file.read()
|
||||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
prediction, ds = predict_upload(dcm_bytes)
|
||||||
|
if prediction is None:
|
||||||
img_tensor = img_tensor.to(device)
|
raise RuntimeError("Model not loaded")
|
||||||
with torch.no_grad():
|
results.append(prediction_to_result(prediction, ds, file.filename))
|
||||||
outputs = model.model(img_tensor)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
confidence = probs[0, pred].item()
|
|
||||||
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
|
||||||
|
|
||||||
results.append({
|
|
||||||
"filename": file.filename,
|
|
||||||
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
|
||||||
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
|
||||||
"anatomical_region": region,
|
|
||||||
"quality_class": int(pred),
|
|
||||||
"confidence": round(confidence, 4),
|
|
||||||
"processing_status": "Success"
|
|
||||||
})
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
results.append({
|
results.append({
|
||||||
"filename": file.filename,
|
"filename": file.filename,
|
||||||
"error": str(e),
|
"error": str(e),
|
||||||
"processing_status": "Failure"
|
"processing_status": "Failure",
|
||||||
})
|
})
|
||||||
|
|
||||||
return {"results": results}
|
return {"results": results}
|
||||||
|
|
@ -379,9 +362,13 @@ async def batch_analyze(files: list[UploadFile] = File(...)):
|
||||||
|
|
||||||
@app.post("/api/v1/export")
|
@app.post("/api/v1/export")
|
||||||
async def export_results(files: list[UploadFile] = File(...)):
|
async def export_results(files: list[UploadFile] = File(...)):
|
||||||
"""Batch analyze and export results as XLSX"""
|
"""
|
||||||
results = []
|
Пакетный анализ и выгрузка в XLSX.
|
||||||
|
|
||||||
|
Колонки соответствуют требованиям к результату. `path_to_study` заполняется
|
||||||
|
как `upload://<имя>`, поскольку при загрузке через HTTP исходный путь
|
||||||
|
исследования недоступен; для обработки архива используйте CLI-инференс.
|
||||||
|
"""
|
||||||
model = load_model()
|
model = load_model()
|
||||||
if model is None:
|
if model is None:
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
|
|
@ -389,66 +376,55 @@ async def export_results(files: list[UploadFile] = File(...)):
|
||||||
content={"error": "Model not loaded"}
|
content={"error": "Model not loaded"}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
results = []
|
||||||
for file in files:
|
for file in files:
|
||||||
|
started = time.time()
|
||||||
try:
|
try:
|
||||||
dcm_bytes = await file.read()
|
dcm_bytes = await file.read()
|
||||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
prediction, ds = predict_upload(dcm_bytes)
|
||||||
|
if prediction is None:
|
||||||
img_tensor = img_tensor.to(device)
|
raise RuntimeError("Model not loaded")
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model.model(img_tensor)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
confidence = probs[0, pred].item()
|
|
||||||
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
|
||||||
|
|
||||||
results.append({
|
results.append({
|
||||||
'filename': file.filename,
|
"path_to_study": f"upload://{file.filename or 'unknown'}",
|
||||||
'path_to_study': str(Path(file.filename).parent if file.filename else ''),
|
"study_uid": str(getattr(ds, "StudyInstanceUID", "") or ""),
|
||||||
'study_uid': getattr(ds, 'StudyInstanceUID', ''),
|
"image_uid": str(getattr(ds, "SOPInstanceUID", "") or ""),
|
||||||
'image_uid': getattr(ds, 'SOPInstanceUID', ''),
|
"anatomical_region": prediction.anatomical_region,
|
||||||
'anatomical_region': region,
|
"quality_class": prediction.quality_class,
|
||||||
'quality_class': int(pred),
|
"violation_type": prediction.violation_type,
|
||||||
'violation_type': 'quality_violation_detected' if pred == 1 else '',
|
"processing_status": "Success",
|
||||||
'confidence': round(confidence, 4),
|
"time_of_processing": round(time.time() - started, 4),
|
||||||
'processing_status': 'Success'
|
"confidence": round(prediction.prob, 4),
|
||||||
|
"violation_reason": prediction.violation_reason,
|
||||||
})
|
})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
results.append({
|
results.append({
|
||||||
'filename': file.filename,
|
"path_to_study": f"upload://{file.filename or 'unknown'}",
|
||||||
'path_to_study': '',
|
"study_uid": "",
|
||||||
'study_uid': '',
|
"image_uid": "",
|
||||||
'image_uid': '',
|
"anatomical_region": "unknown",
|
||||||
'anatomical_region': 'unknown',
|
"quality_class": -1,
|
||||||
'quality_class': -1,
|
"violation_type": "",
|
||||||
'violation_type': '',
|
"processing_status": f"Failure: {str(e)[:80]}",
|
||||||
'confidence': 0.0,
|
"time_of_processing": round(time.time() - started, 4),
|
||||||
'processing_status': f'Failure: {str(e)[:80]}'
|
"confidence": 0.0,
|
||||||
|
"violation_reason": "",
|
||||||
})
|
})
|
||||||
|
|
||||||
# Create DataFrame and export to Excel
|
columns = [
|
||||||
df = pd.DataFrame(results)
|
"path_to_study", "study_uid", "image_uid", "anatomical_region",
|
||||||
|
"quality_class", "violation_type", "processing_status", "time_of_processing",
|
||||||
|
"confidence", "violation_reason",
|
||||||
|
]
|
||||||
|
df = pd.DataFrame(results, columns=columns)
|
||||||
|
|
||||||
# Ensure column order
|
|
||||||
columns = ['filename', 'path_to_study', 'study_uid', 'image_uid',
|
|
||||||
'anatomical_region', 'quality_class', 'violation_type',
|
|
||||||
'confidence', 'processing_status']
|
|
||||||
for col in columns:
|
|
||||||
if col not in df.columns:
|
|
||||||
df[col] = ''
|
|
||||||
df = df[columns]
|
|
||||||
|
|
||||||
# Save to buffer
|
|
||||||
buffer = io.BytesIO()
|
buffer = io.BytesIO()
|
||||||
df.to_excel(buffer, index=False, engine='openpyxl')
|
df.to_excel(buffer, index=False, engine="openpyxl")
|
||||||
buffer.seek(0)
|
buffer.seek(0)
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
buffer,
|
buffer,
|
||||||
media_type='application/vnd.openxmlformats-officedocument.spreadsheetml.sheet',
|
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||||||
headers={'Content-Disposition': f'attachment; filename=dxa_results_{pd.Timestamp.now().strftime("%Y%m%d_%H%M%S")}.xlsx'}
|
headers={"Content-Disposition": f'attachment; filename=dxa_results_{pd.Timestamp.now().strftime("%Y%m%d_%H%M%S")}.xlsx'},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -456,12 +432,10 @@ async def export_results(files: list[UploadFile] = File(...)):
|
||||||
async def analyze_dicom_sr(file: UploadFile = File(...)):
|
async def analyze_dicom_sr(file: UploadFile = File(...)):
|
||||||
"""
|
"""
|
||||||
Generate DICOM SR (Structured Report) for the analysis result.
|
Generate DICOM SR (Structured Report) for the analysis result.
|
||||||
|
|
||||||
Returns a text representation of DICOM SR with standardized codes.
|
Returns a text representation of DICOM SR with standardized codes.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
# Use the detailed analysis
|
|
||||||
# Reuse the detailed analysis logic
|
|
||||||
model = load_model()
|
model = load_model()
|
||||||
if model is None:
|
if model is None:
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
|
|
@ -469,60 +443,40 @@ async def analyze_dicom_sr(file: UploadFile = File(...)):
|
||||||
content={"error": "Model not loaded"}
|
content={"error": "Model not loaded"}
|
||||||
)
|
)
|
||||||
|
|
||||||
# Read file
|
|
||||||
dcm_bytes = await file.read()
|
dcm_bytes = await file.read()
|
||||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
prediction, ds = predict_upload(dcm_bytes)
|
||||||
|
if prediction is None:
|
||||||
|
return JSONResponse(status_code=500, content={"error": "Model not loaded"})
|
||||||
|
|
||||||
# Get original image for quality assessment
|
quality_report = prediction_to_result(prediction, ds)
|
||||||
img_array = ds.pixel_array.astype(np.float32)
|
snomed_map = {
|
||||||
img_array = (img_array - img_array.min()) / (img_array.max() - img_array.min() + 1e-8)
|
"artifact_motion": ("Motion artifact", "WARNING"),
|
||||||
|
"artifact_other": ("Foreign object artifact", "WARNING"),
|
||||||
if len(img_array.shape) == 2:
|
"rotation": ("Rotational misalignment", "WARNING"),
|
||||||
img_3ch = np.stack([img_array] * 3, axis=2)
|
"roi_error": ("Region of interest mismatch", "WARNING"),
|
||||||
else:
|
"incomplete_view": ("Incomplete anatomy", "WARNING"),
|
||||||
img_3ch = img_array
|
"position_error": ("Positioning deviation", "WARNING"),
|
||||||
|
}
|
||||||
# Predict
|
label, completion = snomed_map.get(
|
||||||
img_tensor = img_tensor.to(device)
|
prediction.violation_type, ("DXA image quality acceptable", "FINAL")
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model.model(img_tensor)
|
|
||||||
probs = torch.softmax(outputs, dim=1)
|
|
||||||
pred = outputs.argmax(dim=1).item()
|
|
||||||
confidence = probs[0, pred].item()
|
|
||||||
|
|
||||||
# Determine region
|
|
||||||
h, w = ds.pixel_array.shape
|
|
||||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
|
||||||
|
|
||||||
# Generate segmentation
|
|
||||||
threshold = np.percentile(img_array, 90)
|
|
||||||
segmentation = (img_array > threshold).astype(np.uint8)
|
|
||||||
|
|
||||||
# Generate detailed quality report
|
|
||||||
quality_report = generate_quality_report(
|
|
||||||
image=img_3ch,
|
|
||||||
segmentation=segmentation,
|
|
||||||
region=region,
|
|
||||||
model_prediction=pred,
|
|
||||||
model_confidence=confidence
|
|
||||||
)
|
)
|
||||||
|
quality_report["violation_label"] = label
|
||||||
|
quality_report["completion"] = completion
|
||||||
|
|
||||||
# Generate DICOM SR text representation
|
|
||||||
sr_content = generate_dicom_sr_text(
|
sr_content = generate_dicom_sr_text(
|
||||||
study_uid=getattr(ds, 'StudyInstanceUID', ''),
|
study_uid=quality_report["study_uid"],
|
||||||
image_uid=getattr(ds, 'SOPInstanceUID', ''),
|
image_uid=quality_report["image_uid"],
|
||||||
quality_report=quality_report
|
quality_report=quality_report,
|
||||||
)
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"format": "DICOM SR (Text)",
|
"format": "DICOM SR (Text)",
|
||||||
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
"study_uid": quality_report["study_uid"],
|
||||||
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
"image_uid": quality_report["image_uid"],
|
||||||
"sr_content": sr_content
|
"sr_content": sr_content,
|
||||||
}
|
}
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
import traceback
|
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
status_code=500,
|
status_code=500,
|
||||||
content={
|
content={
|
||||||
|
|
@ -554,20 +508,25 @@ def generate_dicom_sr_text(study_uid: str, image_uid: str, quality_report: Dict)
|
||||||
str: Текстовое представление SR отчета
|
str: Текстовое представление SR отчета
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Map violation types to DICOM codes (simplified)
|
# Соответствие типа нарушения коду отчёта. Коды условные (локальные), так
|
||||||
|
# как полноценного справочника SNOMED/DICOM для контролёра качества DXA в
|
||||||
|
# наборе нет; текстовая формулировка приводится рядом.
|
||||||
violation_code_map = {
|
violation_code_map = {
|
||||||
"correct": ("113001", "DXA Quality Assessment", "FINAL"),
|
"": ("113001", "DXA image quality acceptable", "FINAL"),
|
||||||
"position_error": ("123456", "Positioning Error", "WARNING"),
|
"position_error": ("123456", "Positioning deviation", "WARNING"),
|
||||||
"artifact_motion": ("234567", "Motion Artifact", "WARNING"),
|
"artifact_motion": ("234567", "Motion artifact", "WARNING"),
|
||||||
"artifact_other": ("234568", "Other Artifact", "WARNING"),
|
"artifact_other": ("234568", "Other artifact", "WARNING"),
|
||||||
"labeling_error": ("345678", "Labeling Error", "WARNING"),
|
"labeling_error": ("345678", "Labeling error", "WARNING"),
|
||||||
"incomplete_view": ("456789", "Incomplete View", "WARNING"),
|
"incomplete_view": ("456789", "Incomplete anatomy", "WARNING"),
|
||||||
"roi_error": ("567890", "ROI Error", "WARNING"),
|
"roi_error": ("567890", "Region of interest mismatch", "WARNING"),
|
||||||
"rotation": ("678901", "Rotation Error", "WARNING")
|
"rotation": ("678901", "Rotational misalignment", "WARNING"),
|
||||||
|
"quality_violation_detected": ("999001", "Image quality violation", "WARNING"),
|
||||||
}
|
}
|
||||||
|
|
||||||
violation_type = quality_report.get("violation_type", "correct")
|
violation_type = (quality_report.get("violation_type") or "").strip()
|
||||||
code, label, completion = violation_code_map.get(violation_type, ("999999", "Unknown", "UNKNOWN"))
|
code, label, completion = violation_code_map.get(
|
||||||
|
violation_type, ("999999", "Unknown finding", "UNKNOWN")
|
||||||
|
)
|
||||||
|
|
||||||
sr_lines = [
|
sr_lines = [
|
||||||
"DICOM Structured Report - DXA Quality Assessment",
|
"DICOM Structured Report - DXA Quality Assessment",
|
||||||
|
|
@ -578,45 +537,38 @@ def generate_dicom_sr_text(study_uid: str, image_uid: str, quality_report: Dict)
|
||||||
"Procedure Report:",
|
"Procedure Report:",
|
||||||
f" - Anatomical Region: {quality_report.get('anatomical_region', 'Unknown')}",
|
f" - Anatomical Region: {quality_report.get('anatomical_region', 'Unknown')}",
|
||||||
f" - Quality Classification: {quality_report.get('quality_label', 'Unknown')}",
|
f" - Quality Classification: {quality_report.get('quality_label', 'Unknown')}",
|
||||||
f" - Confidence: {quality_report.get('confidence', 0):.2%}",
|
f" - Violation Probability: {quality_report.get('confidence', 0):.4f}",
|
||||||
|
f" - Decision Threshold: {quality_report.get('threshold_probability', 0.5):.4f}",
|
||||||
|
f" - Region Confidence: {quality_report.get('region_confidence', 0):.4f}",
|
||||||
"",
|
"",
|
||||||
"Findings:",
|
"Findings:",
|
||||||
f" - Violation Type Code: {code}",
|
f" - Violation Type Code: {code}",
|
||||||
f" - Violation Type: {label}",
|
f" - Violation Type: {label}",
|
||||||
f" - Reason: {quality_report.get('reason', 'N/A')}",
|
f" - Explanation: {quality_report.get('reason') or 'No violation detected'}",
|
||||||
f" - View Quality: {quality_report.get('view_quality', 'Unknown')}",
|
f" - View Quality: {quality_report.get('view_quality', 'Unknown')}",
|
||||||
"",
|
"",
|
||||||
"Detailed Metrics:",
|
"Image Characteristics:",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Add metrics
|
# Числовые характеристики изображения, использованные при решении.
|
||||||
metrics = quality_report.get("metrics", {})
|
metrics = quality_report.get("metrics") or {}
|
||||||
|
for key, label_ru in (
|
||||||
motion = metrics.get("motion", {})
|
("laplacian_variance", "Local sharpness metric"),
|
||||||
if motion:
|
("bright_fraction", "Dense-pixel fraction"),
|
||||||
sr_lines.append(f" Motion Detection:")
|
("bbox_aspect", "Bright region aspect ratio"),
|
||||||
sr_lines.append(f" - Motion Detected: {motion.get('motion_detected', False)}")
|
):
|
||||||
sr_lines.append(f" - Severity: {motion.get('severity', 'NONE')}")
|
if key in metrics:
|
||||||
|
sr_lines.append(f" - {label_ru}: {float(metrics[key]):.5f}")
|
||||||
artifacts = metrics.get("artifacts", {})
|
|
||||||
if artifacts:
|
if not metrics:
|
||||||
sr_lines.append(f" Artifact Detection:")
|
sr_lines.append(" - No image metrics available")
|
||||||
sr_lines.append(f" - Any Artifact: {artifacts.get('any_detected', False)}")
|
|
||||||
sr_lines.append(f" - Metal: {artifacts.get('metal_detected', False)}")
|
|
||||||
sr_lines.append(f" - Implant: {artifacts.get('implant_detected', False)}")
|
|
||||||
|
|
||||||
roi = metrics.get("roi_check", {})
|
|
||||||
if roi:
|
|
||||||
sr_lines.append(f" ROI Validation:")
|
|
||||||
sr_lines.append(f" - Valid: {roi.get('valid', False)}")
|
|
||||||
|
|
||||||
# Completion flag
|
|
||||||
sr_lines.extend([
|
sr_lines.extend([
|
||||||
"",
|
"",
|
||||||
"Completion Flag: " + completion,
|
"Completion Flag: " + completion,
|
||||||
"Verification Flag: UNVERIFIED"
|
"Verification Flag: UNVERIFIED",
|
||||||
])
|
])
|
||||||
|
|
||||||
return "\n".join(sr_lines)
|
return "\n".join(sr_lines)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue