develop - hack_2026
This commit is contained in:
parent
1b82992a6d
commit
5a8dbb43a5
|
|
@ -1,5 +1,9 @@
|
|||
# .dockerignore
|
||||
# Цель: контекст сборки должен быть маленьким. Образ содержит только src,
|
||||
# models и run.sh (см. Dockerfile), поэтому данные, тесты и служебные файлы
|
||||
# исключаются — DICOM-датасет не должен попадать в образ с медицинскими данными.
|
||||
__pycache__
|
||||
**/__pycache__
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
|
|
@ -23,8 +27,11 @@ env
|
|||
*.t7
|
||||
data/
|
||||
datasets/
|
||||
.DS_Store
|
||||
dataset_hack/
|
||||
tests/
|
||||
.qwen/
|
||||
.idea/
|
||||
.vscode/
|
||||
.continue/
|
||||
*.swp
|
||||
*.swo
|
||||
71
Dockerfile
71
Dockerfile
|
|
@ -1,38 +1,59 @@
|
|||
# DXA Quality Assessment - Docker Container (CPU only, Python 3.11)
|
||||
# Build: docker build -t dxa-quality .
|
||||
# Run: docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
||||
# DXA Quality Assessment — контейнер для инференса (CPU)
|
||||
#
|
||||
# Сборка: 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
|
||||
|
||||
RUN apt-get update && apt-get install -y \
|
||||
build-essential \
|
||||
ninja-build \
|
||||
# libgl1/libglib2.0-0 нужны opencv (импортируется через зависимости проекта),
|
||||
# libgomp1 — для параллельных циклов torch.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
libgl1 \
|
||||
libglib2.0-0 \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender1 \
|
||||
libgomp1 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Свежий meson (0.64.0+) — ДО установки Python-пакетов
|
||||
RUN pip3 install --no-cache-dir "meson>=0.64.0"
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Устанавливаем пакеты без жёстких пинов — pip сам разрешит зависимости
|
||||
RUN pip3 install --no-cache-dir \
|
||||
fastapi uvicorn torch torchvision monai TotalSegmentator \
|
||||
nibabel pydicom pydicom-seg opencv-python-headless \
|
||||
scikit-learn scipy pandas numpy openpyxl matplotlib timm \
|
||||
python-multipart
|
||||
# --- Слой зависимостей: отдельно от кода, чтобы кэшировался ---
|
||||
# Все версии зафиксированы, torch ставится с CPU-индексом: образ для инференса
|
||||
# не требует GPU, а версия из PyPI на Linux подтянула бы CUDA-колёса.
|
||||
COPY requirements.txt ./
|
||||
RUN pip install --no-cache-dir --upgrade pip \
|
||||
&& 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
|
||||
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV PYTHONPATH=/app
|
||||
# HEALTHCHECK опирается на /api/v1/health, который отдаёт model_loaded.
|
||||
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
|
||||
|
||||
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)
|
||||
- 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
|
||||
### Ключевые требования (condition.txt)
|
||||
|
||||
### Tech Stack
|
||||
| Component | Technology |
|
||||
|-----------|------------|
|
||||
| Backend | Python 3.10, FastAPI, Uvicorn |
|
||||
| ML/Deep Learning | PyTorch, torchvision (ResNet18) |
|
||||
| Image Processing | PIL, OpenCV, pydicom, scipy |
|
||||
| Data Handling | pandas, openpyxl |
|
||||
| Containerization | Docker |
|
||||
- Области: поясничный отдел позвоночника и проксимальный отдел бедренной кости.
|
||||
- Обязательна контейнеризация и скрипт сборки/запуска в Linux.
|
||||
- Результат: XLSX/CSV, одна строка на изображение, столбцы
|
||||
`path_to_study, study_uid, image_uid, anatomical_region, quality_class,
|
||||
violation_type, processing_status, time_of_processing`.
|
||||
- Приоритетные метрики: F1 и ROC-AUC (с 95 % доверительными интервалами).
|
||||
- Работа офлайн, без передачи изображений во внешние сервисы.
|
||||
- Время обработки одного исследования — не более 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/
|
||||
│ ├── main.py # FastAPI app (DXA mode)
|
||||
│ ├── run.py # Server runner
|
||||
│ ├── dxa/ # DXA Quality module
|
||||
│ │ ├── dataset.py # DXADataset class
|
||||
│ │ ├── model.py # ResNet18 classifier
|
||||
│ │ ├── 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
|
||||
src/dxa/
|
||||
├── labels.py # имена -> метки, склейка дублей, разбиение по исследованиям
|
||||
├── preprocess.py # DICOM -> CHW-тензор (единый путь для обучения и API)
|
||||
├── dataset.py # DXADataset, DataLoader
|
||||
├── model.py # сеть, метрики, подбор порога, сохранение/загрузка
|
||||
├── train.py # обучение + отчёт (md/json)
|
||||
└── inference.py # пакетный инференс, определение области, визуализация
|
||||
```
|
||||
|
||||
---
|
||||
### Ключевые решения (проверены экспериментально)
|
||||
|
||||
## 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`)
|
||||
- **Motion detection**: Laplacian variance, FFT blur analysis
|
||||
- **Artifact detection**: Metal, implants, cement, calcifications
|
||||
- **Spine completeness**: Vertebrae count, alignment, spacing
|
||||
- **Hip completeness**: Full visibility, aspect ratio
|
||||
- **Hip rotation**: Major axis angle calculation
|
||||
- **ROI validation**: Boundary check, margin, size
|
||||
- **Violation types**: correct, position_error, artifact_motion, artifact_other, labeling_error, incomplete_view, roi_error, rotation
|
||||
- ROC-AUC по 5 разбиениям: **0.76 ± 0.08** (0.64 – 0.84)
|
||||
- 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
|
||||
|
||||
### ✅ Web Interface
|
||||
- Drag-and-drop DICOM upload
|
||||
- Table with results (filter, sort, search)
|
||||
- **NEW: Detail panel** - click on row to see:
|
||||
- Violation type
|
||||
- Reason (human-readable)
|
||||
- Motion metrics
|
||||
- Artifact detection
|
||||
- ROI validation
|
||||
Разбивка по областям обязательна: в позвоночнике ~29 % нарушений против ~4–5 %
|
||||
у бёдер, а область почти однозначно определяется по ширине кадра, поэтому общий
|
||||
AUC частично отражает различение области.
|
||||
|
||||
---
|
||||
|
||||
## 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)
|
||||
- Task: Binary classification (quality OK vs violation)
|
||||
- Input: 224x224 RGB images
|
||||
- Output: class probabilities
|
||||
```
|
||||
dataset_hack/
|
||||
├── Для теста/ # bad.dcm, l_hip.dcm, r_hip.dcm, spine.dcm
|
||||
└── НД_для_обучения/
|
||||
├── разметка.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
|
||||
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
|
||||
python src/dxa/inference.py \
|
||||
--input-path dataset_hack/Для\ теста \
|
||||
--output-path results.xlsx \
|
||||
--model-path models/dxa_model.pth
|
||||
./run.sh test
|
||||
python -m pytest tests/ -q # 65 тестов
|
||||
```
|
||||
|
||||
`tests/test_labels.py` — разбор имён, склейка дублей, отсутствие утечки при
|
||||
разбиении. `tests/test_preprocess_and_model.py` — предобработка, метрики,
|
||||
подбор порога, контракт модели, BatchNorm при заморозке, roundtrip чекпоинта.
|
||||
|
||||
---
|
||||
|
||||
## Running the Project
|
||||
## Известные ограничения
|
||||
|
||||
### Training
|
||||
```bash
|
||||
python src/dxa/train.py --epochs 10
|
||||
```
|
||||
1. Разметка на уровне исследования → шум в метках снимков.
|
||||
2. Мало данных: 252 снимка, 37 нарушений; доверительные интервалы широкие.
|
||||
3. Тип нарушения определяется эвристиками, а не обученной моделью.
|
||||
4. Порог `SPINE_MIN_WIDTH` привязан к текущему оборудованию.
|
||||
5. Grad-CAM (`src/models/visualization/gradcam.py`) есть, но не подключён.
|
||||
|
||||
### Inference
|
||||
```bash
|
||||
python src/dxa/inference.py --input-path file.dcm --output-path result.xlsx
|
||||
```
|
||||
## Устаревший код (не подключён к API)
|
||||
|
||||
### API Server
|
||||
```bash
|
||||
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
||||
```
|
||||
Эти модули не импортируются из `src/main.py` и `src/dxa/*`; их зависимости
|
||||
закомментированы в `requirements.txt`:
|
||||
|
||||
Web UI: http://localhost:8000
|
||||
|
||||
---
|
||||
|
||||
## Output Format (per Hackathon Requirements)
|
||||
|
||||
### 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` |
|
||||
- `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-заглушки.
|
||||
- `src/quality/artifact_detector.py`, `position_validator.py`, `medical_quality.py` — заглушки.
|
||||
- `src/dataloaders/pet_dataset.py` — остаток прототипа (Oxford-IIIT Pet).
|
||||
|
||||
---
|
||||
|
||||
|
|
@ -240,29 +176,9 @@ Pipeline:
|
|||
|
||||
```bash
|
||||
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
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Development Notes
|
||||
|
||||
### 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)
|
||||
Dockerfile ставит зафиксированные версии, копирует только `src/`, `models/` и
|
||||
`run.sh`, проверяет чекпоинт на этапе сборки и имеет HEALTHCHECK. Данные и тесты
|
||||
в образ не попадают (`.dockerignore`).
|
||||
|
|
|
|||
471
README.md
471
README.md
|
|
@ -1,328 +1,259 @@
|
|||
# 🦴 DXA Quality Assessment
|
||||
|
||||
[](https://www.python.org/)
|
||||
[](https://fastapi.tiangolo.com/)
|
||||
[](https://pytorch.org/)
|
||||
Сервис автоматизированного контроля качества денситометрических исследований (DXA):
|
||||
принимает DICOM, определяет анатомическую область, оценивает, пригодно ли изображение
|
||||
для клинической интерпретации, и формирует структурированный отчёт.
|
||||
|
||||
## Описание
|
||||
## Что делает решение
|
||||
|
||||
Сервис для автоматизированной оценки качества денситометрических исследований (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** — загрузка и обработка медицинских изображений
|
||||
- 🧠 **Классификация** — бинарная оценка качества (OK / Violation)
|
||||
- 🔍 **Детекция нарушений** — определение типа нарушения:
|
||||
- `correct` — качество соответствует норме
|
||||
- `artifact_motion` — артефакт движения
|
||||
- `artifact_other` — прочие артефакты
|
||||
- `position_error` — ошибка позиционирования
|
||||
- `rotation` — нарушение ротации
|
||||
- `incomplete_view` — неполный вид
|
||||
- `roi_error` — ошибка ROI
|
||||
- `labeling_error` — ошибка разметки
|
||||
- 🌐 **REST API** — интеграция с внешними системами
|
||||
- 📊 **Веб-интерфейс** — загрузка и визуализация результатов
|
||||
- 📈 **Экспорт** — выгрузка результатов в XLSX
|
||||
./run.sh train # обучить модель качества
|
||||
./run.sh infer "dataset_hack/Для теста" results.xlsx # пакетная обработка
|
||||
./run.sh serve # API и веб-интерфейс на :8000
|
||||
./run.sh test # тесты
|
||||
```
|
||||
|
||||
Для инференса только на CPU (образ меньше, без CUDA-колёс):
|
||||
|
||||
```bash
|
||||
pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cpu
|
||||
```
|
||||
|
||||
Обучение и инференс можно вызывать напрямую:
|
||||
|
||||
```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-х ответов в рантайме.
|
||||
Вес модели внутрь образа зашит, из сети ничего не скачивается.
|
||||
|
||||
---
|
||||
|
||||
## Архитектура
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────────┐
|
||||
│ FastAPI Server │
|
||||
│ (port 8000) │
|
||||
├─────────────────────────────────────────────────────────────┤
|
||||
│ /api/v1/analyze → Basic quality prediction │
|
||||
│ /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│ │
|
||||
│ └─────────────────────┘ │
|
||||
└─────────────────────────────────────────────────────────────┘
|
||||
DICOM ──▶ предобработка ──▶ ResNet18 (заморожен) ──▶ линейная голова ──▶ логит
|
||||
│ │
|
||||
│ └──▶ голова области (вспомогательная)
|
||||
│
|
||||
├──▶ геометрия яркой зоны: область, ROI, геометрия кадра
|
||||
└──▶ эвристики: резкость, «плотные» включения
|
||||
```
|
||||
|
||||
Ключевые решения и почему они такие:
|
||||
|
||||
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
|
||||
- (опционально) GPU CUDA/MPS для ускорения
|
||||
| Столбец | Описание |
|
||||
|---|---|
|
||||
| `path_to_study` | Путь к исследованию (для HTTP-загрузки — `upload://<имя>`) |
|
||||
| `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
|
||||
# Клонирование
|
||||
git clone https://github.com/yourusername/bone-quality-assessment.git
|
||||
cd bone-quality-assessment
|
||||
## API
|
||||
|
||||
# Создание виртуального окружения
|
||||
python -m venv venv
|
||||
source venv/bin/activate # Linux/Mac
|
||||
# venv\Scripts\activate # Windows
|
||||
|
||||
# Установка зависимостей
|
||||
pip install -r requirements.txt
|
||||
|
||||
# Загрузка модели (опционально)
|
||||
# Поместите файл модели в 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 отчёт |
|
||||
| Метод | Путь | Назначение |
|
||||
|---|---|---|
|
||||
| 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 |
|
||||
|
||||
### Пример использования
|
||||
| POST | `/api/v1/export` | Пакетный анализ и выгрузка в XLSX |
|
||||
|
||||
```bash
|
||||
# Анализ файла
|
||||
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"
|
||||
curl -X POST http://localhost:8000/api/v1/analyze -F "file=@study/spine.dcm"
|
||||
```
|
||||
|
||||
### Ответ детального анализа
|
||||
|
||||
```json
|
||||
{
|
||||
"study_uid": "1.2.643...",
|
||||
"image_uid": "1.2.643...",
|
||||
"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,
|
||||
"blur_laplacian": 0.0008,
|
||||
"severity": "HIGH"
|
||||
},
|
||||
"artifacts": {
|
||||
"any_detected": false,
|
||||
"metal_detected": false
|
||||
},
|
||||
"roi_check": {
|
||||
"valid": true
|
||||
}
|
||||
},
|
||||
"overall_quality": "POOR",
|
||||
"severity": "HIGH"
|
||||
"violation_type": "quality_violation_detected",
|
||||
"reason": "Выявлено нарушение качества изображения",
|
||||
"confidence": 0.72,
|
||||
"threshold_probability": 0.6154,
|
||||
"processing_status": "Success"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Метрики
|
||||
|
||||
Метрики зависят от выбранного разбиения по исследованиям, поэтому приводятся
|
||||
с разбросом. Оценка на валидационной части (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/
|
||||
├── src/
|
||||
│ ├── main.py # FastAPI приложение
|
||||
│ ├── run.py # Запуск сервера
|
||||
│ ├── dxa/ # DXA модуль
|
||||
│ │ ├── model.py # ResNet18 классификатор
|
||||
│ │ ├── dataset.py # Загрузчик данных
|
||||
│ │ ├── train.py # Обучение модели
|
||||
│ │ └── inference.py # Инференс и batch-обработка
|
||||
│ ├── quality/ # Оценка качества
|
||||
│ │ ├── quality_scorer.py # Базовый скорer
|
||||
│ │ └── detailed_assessment.py # Детальный анализ
|
||||
│ ├── api/ # REST API
|
||||
│ │ ├── endpoints.py # Дополнительные эндпоинты
|
||||
│ │ └── static/ # Веб-интерфейс
|
||||
│ └── utils/ # Утилиты
|
||||
├── models/ # Обученные модели
|
||||
│ └── dxa_model.pth # Модель классификатора
|
||||
├── dataset_hack/ # Датасет для обучения/тестирования
|
||||
├── docs/ # Документация
|
||||
│ ├── main.py # FastAPI: маршруты и загрузка модели
|
||||
│ ├── dxa/ # действующий модуль оценки качества
|
||||
│ │ ├── labels.py # разбор имён, метки, склейка дублей, сплит
|
||||
│ │ ├── preprocess.py # DICOM -> тензор (общий для обучения и API)
|
||||
│ │ ├── dataset.py # Dataset и DataLoader
|
||||
│ │ ├── model.py # сеть, метрики, подбор порога
|
||||
│ │ ├── train.py # обучение и отчёт
|
||||
│ │ └── inference.py # пакетный инференс, определение области
|
||||
│ ├── quality/ # эвристики (частично используются API)
|
||||
│ ├── api/static/ # веб-интерфейс
|
||||
│ └── model/, core/, pipeline/ # устаревшие модули, не подключены к API
|
||||
├── models/dxa_model.pth # чекпоинт (+ train_report.md)
|
||||
├── tests/ # pytest: метки, сплит, метрики, модель
|
||||
├── dataset_hack/ # данные (в git не хранятся)
|
||||
├── Dockerfile
|
||||
├── requirements.txt
|
||||
└── README.md
|
||||
└── run.sh
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Обучение модели
|
||||
|
||||
### Подготовка данных
|
||||
|
||||
1. Разместите DICOM-файлы в `dataset_hack/НД_для_обучения/Исследования/`
|
||||
2. Подготовьте Excel-файл разметки `dataset_hack/НД_для_обучения/разметка.xlsx`
|
||||
|
||||
Столбцы разметки:
|
||||
- `study_uid` — ID исследования
|
||||
- `позвоночник_укладка`, `позвоночник_ось`, `позвоночник_артефакты` — критерии для позвоночника
|
||||
- `бедро_позиция_лев`, `бедро_roi_лев` — критерии для левого бедра
|
||||
- `бедро_позиция_прав`, `бедро_roi_прав` — критерии для правого бедра
|
||||
- `итог_позвоночник`, `итог_бедро_лев`, `итог_бедро_прав` — итоговая оценка (0/1)
|
||||
|
||||
### Запуск обучения
|
||||
## Запуск обучения
|
||||
|
||||
```bash
|
||||
python src/dxa/train.py \
|
||||
--epochs 20 \
|
||||
--batch-size 8 \
|
||||
--backbone resnet18 \
|
||||
--output-dir models
|
||||
python -m src.dxa.train --epochs 100 --output-dir models
|
||||
```
|
||||
|
||||
### Аргументы
|
||||
|
||||
| Параметр | По умолчанию | Описание |
|
||||
|----------|-------------|----------|
|
||||
| `--data-root` | `dataset_hack` | Путь к директории с данными |
|
||||
| `--annotation-path` | `dataset_hack/НД_для_обучения/разметка.xlsx` | Путь к файлу разметки |
|
||||
| `--epochs` | 20 | Количество эпох |
|
||||
| `--batch-size` | 8 | Размер батча |
|
||||
| `--backbone` | `resnet18` | Архитектура (resnet18/resnet34/efficientnet_b0) |
|
||||
| `--input-size` | 224 | Размер входного изображения |
|
||||
| `--output-dir` | `models` | Директория для сохранения модели |
|
||||
|---|---|---|
|
||||
| `--data-root` | `dataset_hack` | Каталог датасета |
|
||||
| `--annotation-path` | `dataset_hack/НД_для_обучения/разметка.xlsx` | Excel с разметкой (только отчёт о расхождениях) |
|
||||
| `--backbone` | `resnet18` | `resnet18` / `resnet34` |
|
||||
| `--head` | `linear` | `linear` (линейный зонд) / `mlp` |
|
||||
| `--freeze-epochs` | `-1` | `-1` — backbone заморожен всегда; `0` — обучать всю сеть |
|
||||
| `--epochs`, `--batch-size`, `--learning-rate`, `--weight-decay` | 100 / 16 / 3e-4 / 5e-2 | Оптимизация |
|
||||
| `--balance` | `none` | `loss` / `sampler` для компенсации дисбаланса |
|
||||
| `--val-fraction`, `--seed` | 0.2 / 42 | Разбиение по исследованиям |
|
||||
| `--augment` | выключено | Включает яркостную аугментацию (ухудшает метрики, см. п. 5) |
|
||||
| `--output-dir` | `models` | Куда сохранять чекпоинт и отчёты |
|
||||
| `--dry-run` | — | Проверить разбор данных и разбиение без обучения |
|
||||
|
||||
---
|
||||
|
||||
## Инференс
|
||||
|
||||
### Одиночный файл
|
||||
|
||||
```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
|
||||
После обучения в `--output-dir` появляются `dxa_model.pth`, `train_report.md`
|
||||
и `train_report.json`; отчёт удобно приложить к презентации.
|
||||
|
||||
## Команда
|
||||
|
||||
- **Грачев Денис** — Разработка
|
||||
- **Грачев Татьяна** — Капитан
|
||||
|
||||
---
|
||||
- **Грачев Денис** — разработка
|
||||
- **Грачев Татьяна** — капитан
|
||||
|
||||
<div align="center">
|
||||
<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
|
||||
certifi==2026.7.22
|
||||
click==8.1.8
|
||||
contourpy==1.3.0
|
||||
cycler==0.12.1
|
||||
et_xmlfile==2.0.0
|
||||
exceptiongroup==1.3.1
|
||||
# Зависимости проекта с зафиксированными версиями.
|
||||
#
|
||||
# Установка:
|
||||
# pip install -r requirements.txt
|
||||
# Для инференса на CPU (меньше образ, без CUDA-колёс):
|
||||
# pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cpu
|
||||
#
|
||||
# Состав разбит на две части:
|
||||
# 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
|
||||
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
|
||||
sympy==1.14.0
|
||||
threadpoolctl==3.7.0
|
||||
timm==1.0.29
|
||||
uvicorn==0.39.0
|
||||
python-multipart==0.0.20
|
||||
pydantic==2.13.4
|
||||
|
||||
# --- Глубокое обучение ---
|
||||
torch==2.8.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
|
||||
typer==0.23.2
|
||||
typing-inspection==0.4.2
|
||||
typing_extensions==4.16.0
|
||||
tzdata==2026.3
|
||||
unicorn==2.1.4
|
||||
uvicorn==0.39.0
|
||||
zipp==3.23.1
|
||||
|
||||
# --- Тесты ---
|
||||
pytest==8.4.2
|
||||
|
||||
# --- Устаревшие модули (не используются действующим API) ---
|
||||
# TotalSegmentator при первом запуске скачивает собственные веса из сети,
|
||||
# поэтому в офлайн-контейнер инференса его включать не следует.
|
||||
# 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
|
||||
|
|
|
|||
190
run.sh
190
run.sh
|
|
@ -1,134 +1,90 @@
|
|||
#!/bin/bash
|
||||
# DXA Quality Assessment - Main entry point
|
||||
# Usage: bash run.sh [command] [args...]
|
||||
#!/usr/bin/env bash
|
||||
#
|
||||
# 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
|
||||
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}
|
||||
PYTHON=${PYTHON:-python3}
|
||||
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}
|
||||
EPOCHS=${EPOCHS:-10}
|
||||
BATCH_SIZE=${BATCH_SIZE:-16}
|
||||
PORT=${PORT:-8000}
|
||||
|
||||
case "$COMMAND" in
|
||||
command=${1:-help}
|
||||
shift || true
|
||||
|
||||
case "$command" in
|
||||
train)
|
||||
echo -e "${YELLOW}Training model...${NC}"
|
||||
|
||||
python3 -c "
|
||||
import sys
|
||||
sys.path.insert(0, '.')
|
||||
|
||||
import torch
|
||||
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}"
|
||||
echo -e "${YELLOW}Обучение классификатора качества DXA${NC}"
|
||||
"$PYTHON" -m src.dxa.train \
|
||||
--data-root "$DATA_ROOT" \
|
||||
--annotation-path "$ANNOTATION_PATH" \
|
||||
--output-dir "$(dirname "$MODEL_PATH")" \
|
||||
"$@"
|
||||
echo -e "${GREEN}Готово. Чекпоинт: ${MODEL_PATH}${NC}"
|
||||
;;
|
||||
|
||||
infer)
|
||||
INPUT_PATH=${2:-dataset_hack/Для теста}
|
||||
OUTPUT_PATH=${3:-results.xlsx}
|
||||
|
||||
echo -e "${YELLOW}Running inference...${NC}"
|
||||
echo "Input: ${INPUT_PATH}"
|
||||
echo "Output: ${OUTPUT_PATH}"
|
||||
|
||||
python3 -c "
|
||||
import sys
|
||||
sys.path.insert(0, '.')
|
||||
from src.dxa.inference import process_dicom_files
|
||||
import argparse
|
||||
|
||||
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}"
|
||||
input_path=${1:-dataset_hack/Для теста}
|
||||
output_path=${2:-results.xlsx}
|
||||
shift 2 2>/dev/null || true
|
||||
echo -e "${YELLOW}Пакетная обработка${NC}"
|
||||
echo " вход: $input_path"
|
||||
echo " выход: $output_path"
|
||||
"$PYTHON" -m src.dxa.inference \
|
||||
--input-path "$input_path" \
|
||||
--output-path "$output_path" \
|
||||
--model-path "$MODEL_PATH" \
|
||||
"$@"
|
||||
echo -e "${GREEN}Готово: ${output_path}${NC}"
|
||||
;;
|
||||
|
||||
serve)
|
||||
echo -e "${YELLOW}Starting API server...${NC}"
|
||||
python3 -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
||||
echo -e "${YELLOW}Запуск API на порту ${PORT}${NC}"
|
||||
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|*)
|
||||
echo "Usage: $0 [command] [options]"
|
||||
echo ""
|
||||
echo "Commands:"
|
||||
echo " train Train the model"
|
||||
echo " infer <input> <output> Run inference"
|
||||
echo " serve Start API server"
|
||||
echo ""
|
||||
echo "Environment variables:"
|
||||
echo " DATA_ROOT Data directory (default: dataset_hack)"
|
||||
echo " ANNOTATION_PATH Annotation Excel file"
|
||||
echo " MODEL_PATH Model output path"
|
||||
echo " EPOCHS Training epochs (default: 10)"
|
||||
echo " BATCH_SIZE Batch size (default: 16)"
|
||||
echo ""
|
||||
echo "Examples:"
|
||||
echo " $0 train"
|
||||
echo " EPOCHS=50 $0 train"
|
||||
echo " $0 infer dataset_hack/Для теста results.xlsx"
|
||||
cat <<'USAGE'
|
||||
DXA Quality Assessment — управление запуском
|
||||
|
||||
Команды:
|
||||
train [аргументы] Обучить классификатор качества.
|
||||
Пример: ./run.sh train --epochs 100 --head mlp
|
||||
infer <вход> <выход> [арг.] Пакетная обработка DICOM (файл или каталог).
|
||||
Пример: ./run.sh infer "dataset_hack/Для теста" results.xlsx
|
||||
Дополнительно: --zip-out masks.zip
|
||||
serve [порт] Запустить HTTP API и веб-интерфейс.
|
||||
test Запустить тесты.
|
||||
|
||||
Переменные окружения:
|
||||
DATA_ROOT Каталог датасета (по умолчанию dataset_hack)
|
||||
ANNOTATION_PATH Excel с разметкой (только для отчёта о расхождениях)
|
||||
MODEL_PATH Путь к чекпоинту (по умолчанию models/dxa_model.pth)
|
||||
PORT Порт API (по умолчанию 8000)
|
||||
PYTHON Интерпретатор (по умолчанию python3)
|
||||
|
||||
Типовой порядок работы:
|
||||
./run.sh train # обучить модель
|
||||
./run.sh infer dataset_hack results.xlsx
|
||||
./run.sh serve # веб-интерфейс на http://localhost:8000
|
||||
USAGE
|
||||
;;
|
||||
esac
|
||||
|
|
|
|||
|
|
@ -1,363 +1,205 @@
|
|||
"""
|
||||
Dataset for DXA (bone densitometry) quality assessment
|
||||
Датасет DXA для обучения классификатора качества.
|
||||
|
||||
Метки формируются из имён DICOM-файлов (см. `src.dxa.labels`), где суффикс
|
||||
`_good`/`_bad` кодирует экспертную оценку; отсутствие суффикса — «хорошее»
|
||||
изображение. Одинаковые по содержимому файлы склеиваются в один пример.
|
||||
|
||||
Разбиение на train/val выполняется по исследованиям, чтобы снимки одного
|
||||
исследования не попадали одновременно в обучение и валидацию.
|
||||
"""
|
||||
import os
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Tuple, Optional
|
||||
import pydicom
|
||||
from PIL import Image
|
||||
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
||||
|
||||
import numpy as np
|
||||
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):
|
||||
"""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:
|
||||
"""
|
||||
Determine anatomical region from DICOM image content.
|
||||
Uses the same algorithm as inference.py
|
||||
"""
|
||||
import pydicom
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
ds = pydicom.dcmread(dcm_path)
|
||||
img = ds.pixel_array.astype(np.float32)
|
||||
|
||||
h, w = img.shape
|
||||
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 __init__(
|
||||
self,
|
||||
records: Sequence[ImageRecord],
|
||||
preprocess: Optional[PreprocessConfig] = None,
|
||||
train: bool = False,
|
||||
seed: int = 0,
|
||||
augment: bool = False,
|
||||
):
|
||||
self.records = list(records)
|
||||
self.preprocess = preprocess or PreprocessConfig()
|
||||
self.train = train
|
||||
self.augment = augment
|
||||
self.seed = seed
|
||||
|
||||
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]
|
||||
@property
|
||||
def labels(self) -> List[int]:
|
||||
return [r.label for r in self.records]
|
||||
|
||||
# Load DICOM
|
||||
ds = pydicom.dcmread(sample['dcm_path'])
|
||||
img = ds.pixel_array.astype(np.float32)
|
||||
def region_index(self, region: Optional[str]) -> int:
|
||||
"""Индекс области для эмбеддинга (0 — неизвестная область)."""
|
||||
return REGIONS.index(region) + 1 if region in REGIONS else 0
|
||||
|
||||
# Normalize to 0-1
|
||||
img_min = img.min()
|
||||
img_max = img.max()
|
||||
if img_max > img_min:
|
||||
img = (img - img_min) / (img_max - img_min)
|
||||
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict[str, Any]]:
|
||||
rec = self.records[idx]
|
||||
img = preprocess_dicom(rec.path, self.preprocess)
|
||||
|
||||
# Convert to 3-channel for pretrained models
|
||||
img = np.stack([img] * 3, axis=0)
|
||||
if self.train and self.augment:
|
||||
rng = np.random.default_rng(self.seed + idx)
|
||||
img = _augment(img, rng)
|
||||
|
||||
# Convert to uint8 for PIL
|
||||
img = (img * 255).astype(np.uint8)
|
||||
tensor = torch.from_numpy(img).float()
|
||||
target = torch.tensor(rec.label, dtype=torch.long)
|
||||
|
||||
# Resize
|
||||
img_pil = Image.fromarray(img.transpose(1, 2, 0))
|
||||
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
|
||||
img = img.astype(np.float32) / 255.0
|
||||
|
||||
# Apply transforms
|
||||
if self.transform:
|
||||
img = self.transform(img)
|
||||
|
||||
# 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']
|
||||
return tensor, {
|
||||
"label": target,
|
||||
"region_id": torch.tensor(self.region_index(rec.region), dtype=torch.long),
|
||||
"study_uid": rec.study,
|
||||
"anatomical_region": rec.region or "unknown",
|
||||
"dcm_path": str(rec.path),
|
||||
}
|
||||
|
||||
|
||||
def create_dataloaders(data_root: str,
|
||||
annotation_path: str,
|
||||
batch_size: int = 8,
|
||||
input_size: Tuple[int, int] = (224, 224),
|
||||
num_workers: int = 4):
|
||||
"""Create train and validation dataloaders"""
|
||||
def build_records(
|
||||
data_root: str | Path,
|
||||
annotation_path: Optional[str | Path] = None,
|
||||
dedup: bool = True,
|
||||
) -> List[ImageRecord]:
|
||||
"""Найти и разметить все уникальные снимки датасета."""
|
||||
records = scan_dataset(data_root, with_pixel_dedup=dedup, annotation_path=annotation_path)
|
||||
logger.info("Dataset scan complete:\n%s", format_summary(records))
|
||||
return records
|
||||
|
||||
train_dataset = DXADataset(
|
||||
|
||||
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,
|
||||
annotation_path=annotation_path,
|
||||
input_size=input_size,
|
||||
mode='train'
|
||||
val_fraction=val_fraction,
|
||||
seed=seed,
|
||||
preprocess=preprocess,
|
||||
)
|
||||
|
||||
val_dataset = DXADataset(
|
||||
data_root=data_root,
|
||||
annotation_path=annotation_path,
|
||||
input_size=input_size,
|
||||
mode='val'
|
||||
pin = torch.cuda.is_available()
|
||||
train_loader = DataLoader(
|
||||
train_ds, batch_size=batch_size, shuffle=True,
|
||||
num_workers=num_workers, pin_memory=pin, drop_last=False,
|
||||
)
|
||||
|
||||
train_loader = torch.utils.data.DataLoader(
|
||||
train_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=True,
|
||||
num_workers=num_workers,
|
||||
pin_memory=True
|
||||
val_loader = DataLoader(
|
||||
val_ds, batch_size=batch_size, shuffle=False,
|
||||
num_workers=num_workers, pin_memory=pin, drop_last=False,
|
||||
)
|
||||
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
|
||||
|
||||
if __name__ == "__main__":
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||
train_loader, val_loader, cfg = create_dataloaders(
|
||||
data_root="dataset_hack",
|
||||
annotation_path="dataset_hack/НД_для_обучения/разметка.xlsx",
|
||||
batch_size=4,
|
||||
)
|
||||
|
||||
return train_loader, val_loader
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# Test
|
||||
train_loader, val_loader = create_dataloaders(
|
||||
data_root='dataset_hack',
|
||||
annotation_path='dataset_hack/НД_для_обучения/разметка.xlsx'
|
||||
)
|
||||
print(f'Train samples: {len(train_loader.dataset)}')
|
||||
print(f'Val samples: {len(val_loader.dataset)}')
|
||||
print(f"Preprocess: {cfg.to_dict()}")
|
||||
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
|
||||
"""
|
||||
DXA Quality Inference - Batch processing with Excel output
|
||||
Пакетный инференс классификатора качества DXA
|
||||
=============================================
|
||||
|
||||
Этот модуль выполняет инференс модели классификации качества DXA исследований.
|
||||
Основные функции:
|
||||
- Загрузка и предобработка DICOM изображений
|
||||
- Определение анатомической области (позвоночник/бедро)
|
||||
- Бинарная классификация качества (OK/Violation)
|
||||
- Пакетная обработка с экспортом в Excel
|
||||
Обрабатывает DICOM-исследования и формирует таблицу в формате требований:
|
||||
`path_to_study, study_uid, image_uid, anatomical_region, quality_class,
|
||||
violation_type, processing_status, time_of_processing`.
|
||||
|
||||
Анатомическая область определяется по содержимому изображения
|
||||
(`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
|
||||
import sys
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
import warnings
|
||||
warnings.filterwarnings('ignore')
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
import warnings
|
||||
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 pandas as pd
|
||||
import pydicom
|
||||
import torch
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
|
||||
# Add src to path
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||
|
||||
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
|
||||
width: int = 0
|
||||
height: int = 0
|
||||
|
||||
Returns:
|
||||
str: Устройство ('mps', 'cuda' или 'cpu')
|
||||
"""
|
||||
|
||||
@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():
|
||||
return 'mps'
|
||||
elif torch.cuda.is_available():
|
||||
return 'cuda'
|
||||
else:
|
||||
return 'cpu'
|
||||
return "mps"
|
||||
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 и
|
||||
добавляет классификационную голову для бинарной классификации
|
||||
(качество OK vs Violation).
|
||||
|
||||
Args:
|
||||
model_path: Путь к файлу модели (.pth)
|
||||
backbone: Архитектура backbone (resnet18/resnet34/efficientnet_b0)
|
||||
device: Устройство для загрузки модели
|
||||
Архитектура (backbone, тип головы) и параметры предобработки берутся из
|
||||
самого чекпоинта, поэтому вызывающей стороне не нужно их дублировать и
|
||||
невозможно рассинхронизировать обучение и инференс. Явно переданные
|
||||
`backbone`/`head` проверяются на совместимость с сохранёнными.
|
||||
|
||||
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)
|
||||
model.load(model_path)
|
||||
saved_backbone = (meta_only or {}).get("backbone", "resnet18")
|
||||
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()
|
||||
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 изображения для модели.
|
||||
Геометрические признаки изображения.
|
||||
|
||||
Этапы предобработки:
|
||||
1. Чтение DICOM и извлечение pixel_array
|
||||
2. Нормализация интенсивности в диапазон [0, 1]
|
||||
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)
|
||||
Порог яркой области берётся по 95-му перцентилю. Размеры кадра входят в
|
||||
набор сигналов, потому что у аппарата позвоночные и бедренные снимки имеют
|
||||
разную ширину кадра (300 против 280 пикселей), и это самый надёжный признак
|
||||
области на данном оборудовании.
|
||||
"""
|
||||
ds = pydicom.dcmread(dcm_path)
|
||||
img = ds.pixel_array.astype(np.float32)
|
||||
norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
||||
h, w = norm.shape
|
||||
binary = norm > np.percentile(norm, 95)
|
||||
if not binary.any():
|
||||
return HeuristicSignals(width=w, height=h)
|
||||
|
||||
# Normalize to 0-1
|
||||
img_min = img.min()
|
||||
img_max = img.max()
|
||||
if img_max > img_min:
|
||||
img = (img - img_min) / (img_max - img_min)
|
||||
rows, cols = np.any(binary, axis=1), np.any(binary, axis=0)
|
||||
rmin, rmax = np.where(rows)[0][[0, -1]]
|
||||
cmin, cmax = np.where(cols)[0][[0, -1]]
|
||||
bbox_aspect = (rmax - rmin) / ((cmax - cmin) + 1e-6)
|
||||
|
||||
# Convert to 3-channel
|
||||
img = np.stack([img] * 3, axis=0)
|
||||
left = binary[:, :w // 2].sum()
|
||||
right = binary[:, w // 2:].sum()
|
||||
ratio = left / (right + 1e-6)
|
||||
|
||||
# Convert to uint8 for PIL
|
||||
img = (img * 255).astype(np.uint8)
|
||||
left_half = norm[:, :w // 2]
|
||||
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
|
||||
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
|
||||
return HeuristicSignals(float(bbox_aspect), float(symmetry), float(ratio), int(w), int(h))
|
||||
|
||||
|
||||
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()
|
||||
which uses image analysis.
|
||||
Порядок решений:
|
||||
1. Ширина кадра: у позвоночника кадр шире (300 px против 280 px у бёдер).
|
||||
На этом оборудовании признак разделяет области безошибочно, поэтому
|
||||
используется первым.
|
||||
2. При нетипичной ширине — предсказание обученной головы области.
|
||||
3. Если голова неуверена — форма яркой области и перевес светимости.
|
||||
|
||||
Важно, что область НЕ берётся из имени файла: на закрытом наборе имена
|
||||
могут не содержать разметки региона.
|
||||
"""
|
||||
import os
|
||||
filename = os.path.basename(dcm_path).lower()
|
||||
if signals.width:
|
||||
if signals.width >= SPINE_MIN_WIDTH:
|
||||
return "spine", 0.8
|
||||
# Бедро: ширину кадра делят левый и правый снимки, поэтому сторону
|
||||
# определяем по перевесу светимости яркой области. Голова обучена на
|
||||
# обе стороны, но различает их хуже, чем асимметрия.
|
||||
if signals.left_right_ratio > 1.3:
|
||||
return "hip_right", 0.6
|
||||
if signals.left_right_ratio < 0.7:
|
||||
return "hip_left", 0.6
|
||||
if predicted_region in ("hip_left", "hip_right") and region_confidence >= min_confidence:
|
||||
return predicted_region, region_confidence
|
||||
return "hip", 0.4
|
||||
|
||||
if 'spine' in filename:
|
||||
return 'spine'
|
||||
elif 'l_hip' in filename or 'left_hip' in filename:
|
||||
return 'hip_left'
|
||||
elif 'r_hip' in filename or 'right_hip' in filename:
|
||||
return 'hip_right'
|
||||
if predicted_region in REGIONS and region_confidence >= min_confidence:
|
||||
return predicted_region, region_confidence
|
||||
|
||||
return 'unknown'
|
||||
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
|
||||
img_norm = (img - img.min()) / (img.max() - img.min() + 1e-6)
|
||||
if motion is not None and motion < metrics.get("motion_threshold", 0.0):
|
||||
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
|
||||
threshold = np.percentile(img_norm, 95)
|
||||
binary = img_norm > threshold
|
||||
return "quality_violation_detected", REASON_BY_TYPE["quality_violation_detected"]
|
||||
|
||||
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
|
||||
com = ndimage.center_of_mass(binary)
|
||||
bright_x = com[1] / w
|
||||
def image_samples(img: np.ndarray) -> Dict[str, float]:
|
||||
"""Числовые характеристики изображения для отчёта и выбора типа нарушения."""
|
||||
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
|
||||
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)
|
||||
def predict_from_array(arr: np.ndarray, model: DXAQualityModel,
|
||||
cfg: PreprocessConfig, device: str, threshold: float) -> Prediction:
|
||||
"""
|
||||
Предсказание качества по уже прочитанному массиву пикселей.
|
||||
|
||||
# Feature 3: Vertical/horizontal edges
|
||||
dx = np.diff(img_norm, axis=1)
|
||||
dy = np.diff(img_norm, axis=0)
|
||||
v_edges = np.abs(dx).mean()
|
||||
h_edges = np.abs(dy).mean()
|
||||
v_h_ratio = v_edges / (h_edges + 1e-6)
|
||||
Общая точка входа для файлового инференса и HTTP API: гарантирует, что
|
||||
предобработка и решающее правило совпадают во всех режимах.
|
||||
"""
|
||||
tensor = torch.from_numpy(preprocess_from_array(arr, cfg)).float().unsqueeze(0).to(device)
|
||||
|
||||
# Classification rules based on analysis:
|
||||
# spine: bbox_aspect ~1.2, symmetry > 0.4, v_h_ratio < 1.5
|
||||
# hip: bbox_aspect > 1.5, symmetry < 0.4, v_h_ratio > 1.5
|
||||
with torch.no_grad():
|
||||
out = model.model(tensor)
|
||||
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
|
||||
if bbox_aspect < 1.5:
|
||||
# More square bright region -> spine
|
||||
return 'spine'
|
||||
elif bbox_aspect < 1.8:
|
||||
# Check symmetry as secondary
|
||||
if symmetry > 0.35:
|
||||
return 'spine'
|
||||
else:
|
||||
return 'hip' # unknown side
|
||||
signals = heuristic_signals(arr)
|
||||
region, region_conf = resolve_region(REGION_BY_ID.get(region_id), region_conf, signals)
|
||||
quality_class = 1 if logit >= threshold else 0
|
||||
samples = image_samples(arr)
|
||||
|
||||
if quality_class == 1:
|
||||
violation_type, reason = classify_violation_type(region, {}, samples)
|
||||
else:
|
||||
# Highly elongated bright region -> hip
|
||||
# Determine left vs right based on left/right brightness ratio
|
||||
# left_right_ratio > 1.3 -> right hip (right side brighter)
|
||||
# left_right_ratio < 0.7 -> left hip (left side brighter)
|
||||
if left_right_ratio > 1.3:
|
||||
return 'hip_right'
|
||||
elif left_right_ratio < 0.7:
|
||||
return 'hip_left'
|
||||
else:
|
||||
return 'hip' # unclear side
|
||||
violation_type, reason = "", ""
|
||||
|
||||
return Prediction(
|
||||
quality_class=quality_class,
|
||||
prob=prob,
|
||||
logit=logit,
|
||||
threshold=threshold,
|
||||
anatomical_region=region,
|
||||
region_confidence=region_conf,
|
||||
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.
|
||||
Falls back to height-based heuristic only if image cannot be analyzed.
|
||||
API принимает файлы потоком, поэтому чтение идёт через BytesIO — временные
|
||||
файлы не создаются.
|
||||
"""
|
||||
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:
|
||||
# Load image
|
||||
ds = pydicom.dcmread(dcm_path)
|
||||
img = ds.pixel_array.astype(np.float32)
|
||||
arr = load_dicom_array(dcm_path)
|
||||
prediction = predict_from_array(arr, model, cfg, device, threshold)
|
||||
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 'hip'
|
||||
else:
|
||||
return 'spine'
|
||||
except:
|
||||
return 'unknown'
|
||||
return ImageResult(
|
||||
path_to_study=path_to_study,
|
||||
study_uid=study_uid_for(ds, dcm_path),
|
||||
image_uid=str(getattr(ds, "SOPInstanceUID", "") or ""),
|
||||
anatomical_region=prediction.anatomical_region,
|
||||
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 изображений.
|
||||
|
||||
Этапы обработки:
|
||||
1. Определение устройства (MPS/CUDA/CPU)
|
||||
2. Загрузка обученной модели
|
||||
3. Поиск DICOM файлов в указанной директории
|
||||
4. Инференс для каждого файла:
|
||||
- Предобработка изображения
|
||||
- Предикт модели (бинарная классификация)
|
||||
- Определение анатомической области
|
||||
- Сохранение метаданных DICOM
|
||||
5. Экспорт результатов в Excel/CSV
|
||||
|
||||
Args:
|
||||
args: Аргументы командной строки (input_path, output_path, model_path и т.д.)
|
||||
Дополнительный функционал: zip-архив с изображениями, где выделена
|
||||
зона интереса (порог по 90-му перцентилю яркости).
|
||||
"""
|
||||
# Setup
|
||||
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!")
|
||||
if not enabled:
|
||||
return
|
||||
|
||||
# Process each file
|
||||
results = []
|
||||
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as archive:
|
||||
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"):
|
||||
start_time = datetime.now()
|
||||
results.append(process_one(
|
||||
dcm_path, checkpoint.model, checkpoint.preprocess, device, checkpoint.threshold
|
||||
))
|
||||
|
||||
try:
|
||||
# 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
|
||||
df = results_to_dataframe(results)
|
||||
output_path = Path(args.output_path)
|
||||
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)
|
||||
else:
|
||||
df.to_excel(output_path, index=False)
|
||||
|
||||
print(f"\n{'='*50}")
|
||||
print(f"Results saved to {output_path}")
|
||||
print(f"{'='*50}")
|
||||
print(f"\nSummary:")
|
||||
print(f" Total files: {len(df)}")
|
||||
print(f" Successful: {(df['processing_status'] == 'Success').sum()}")
|
||||
print(f" Quality OK (class 0): {(df['quality_class'] == 0).sum()}")
|
||||
print(f" Quality Issues (class 1): {(df['quality_class'] == 1).sum()}")
|
||||
if args.zip_out:
|
||||
write_visualizations(results, checkpoint.preprocess, Path(args.zip_out), enabled=True)
|
||||
logger.info("Visualizations: %s", args.zip_out)
|
||||
|
||||
# Show sample output
|
||||
print(f"\nSample output:")
|
||||
print(df.head().to_string())
|
||||
successful = int((df["processing_status"] == "Success").sum())
|
||||
logger.info(
|
||||
"Done. %d/%d processed (%.1f%%), quality_class=1 in %d rows, median time %.3fs -> %s",
|
||||
successful, len(df), 100 * successful / max(len(df), 1),
|
||||
int((df["quality_class"] == 1).sum()),
|
||||
float(df["time_of_processing"].median()) if len(df) else 0.0,
|
||||
output_path,
|
||||
)
|
||||
return df
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='DXA Quality Inference')
|
||||
|
||||
# Input/Output
|
||||
parser.add_argument('--input-path', type=str, required=True,
|
||||
help='Path to DICOM file or directory')
|
||||
parser.add_argument('--output-path', type=str, required=True,
|
||||
help='Output CSV or Excel file')
|
||||
parser.add_argument('--model-path', type=str,
|
||||
default='models/dxa_model.pth',
|
||||
help='Path to trained model')
|
||||
|
||||
# Model arguments
|
||||
parser.add_argument('--backbone', type=str, default='resnet18',
|
||||
choices=['resnet18', 'resnet34', 'efficientnet_b0'],
|
||||
help='Backbone architecture')
|
||||
parser.add_argument('--input-size', type=int, default=224,
|
||||
help='Input image size')
|
||||
|
||||
args = parser.parse_args()
|
||||
process_dicom_files(args)
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Пакетный инференс классификатора качества DXA",
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
parser.add_argument("--input-path", required=True, help="DICOM файл или каталог")
|
||||
parser.add_argument("--output-path", required=True, help="Путь к .xlsx или .csv")
|
||||
parser.add_argument("--model-path", default="models/dxa_model.pth", help="Чекпоинт модели")
|
||||
parser.add_argument("--backbone", default=None, choices=["resnet18", "resnet34"],
|
||||
help="Проверить соответствие backbone в чекпоинте (по умолчанию — из чекпоинта)")
|
||||
parser.add_argument("--head", default=None, choices=["linear", "mlp"],
|
||||
help="Проверить соответствие головы в чекпоинте (по умолчанию — из чекпоинта)")
|
||||
parser.add_argument("--input-size", type=int, default=None,
|
||||
help="Переопределить размер входа (по умолчанию — из чекпоинта)")
|
||||
parser.add_argument("--device", default=None, help="cpu / cuda / mps")
|
||||
parser.add_argument("--zip-out", default=None,
|
||||
help="Zip-архив с визуализацией зоны интереса (дополнительный функционал)")
|
||||
return parser
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
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)
|
||||
try:
|
||||
run(args)
|
||||
except (FileNotFoundError, ValueError) as exc:
|
||||
logger.error("%s", exc)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
|
|||
635
src/dxa/model.py
635
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.nn as nn
|
||||
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):
|
||||
"""
|
||||
CNN classifier for DXA image quality assessment
|
||||
Uses pretrained backbone (ResNet/EfficientNet)
|
||||
ResNet backbone + голова качества + вспомогательная голова области.
|
||||
|
||||
Голова выбирается параметром `head`:
|
||||
|
||||
- ``linear`` (по умолчанию) — один линейный слой на pooled-признаках.
|
||||
Это линейный зонд: обучаются 2×(512+1) параметра, поэтому на выборке
|
||||
из ~250 снимков он не переобучается. На валидации даёт AUC ≈ 0.80,
|
||||
тогда как MLP-голова уходит в переобучение (train F1 → 0.9 при val AUC
|
||||
≈ 0.5), что проверено экспериментально на этом датасете.
|
||||
- ``mlp`` — двухслойная голова; требует существенно больше данных.
|
||||
|
||||
При ``feature_norm=True`` вход головы стандартизуется по статистикам
|
||||
обучающей выборки, зафиксированным в буферах `feat_mean`/`feat_std`
|
||||
(см. `fit_feature_norm`). Без этого логиты смещены, вероятности
|
||||
скучены у нуля и подобранный порог лишается смысла.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
backbone: str = 'resnet18',
|
||||
num_classes: int = 2,
|
||||
pretrained: bool = True,
|
||||
dropout: float = 0.3):
|
||||
def __init__(
|
||||
self,
|
||||
backbone: str = "resnet18",
|
||||
pretrained: bool = True,
|
||||
dropout: float = 0.3,
|
||||
head: str = "linear",
|
||||
hidden_dim: int = 256,
|
||||
feature_norm: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
if backbone not in _WEIGHTS:
|
||||
raise ValueError(f"Unknown backbone {backbone!r}; expected one of {BACKBONES}")
|
||||
if head not in ("linear", "mlp"):
|
||||
raise ValueError(f"head must be 'linear' or 'mlp', got {head!r}")
|
||||
|
||||
self.backbone_name = backbone
|
||||
|
||||
# Load pretrained backbone
|
||||
if backbone == 'resnet18':
|
||||
self.backbone = models.resnet18(weights='IMAGENET1K_V1' if pretrained else None)
|
||||
feature_dim = 512
|
||||
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.head_type = head
|
||||
self.feature_norm = feature_norm
|
||||
weights_name, feature_dim = _WEIGHTS[backbone]
|
||||
net = getattr(models, backbone)(weights=weights_name if pretrained else None)
|
||||
self.features = nn.Sequential(*list(net.children())[:-1]) # всё, кроме fc
|
||||
self.feature_dim = feature_dim
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Forward pass"""
|
||||
features = self.backbone(x)
|
||||
return self.classifier(features)
|
||||
self.register_buffer("feat_mean", torch.zeros(feature_dim))
|
||||
self.register_buffer("feat_std", torch.ones(feature_dim))
|
||||
|
||||
if head == "linear":
|
||||
# Dropout в линейном зонде только вредит: он добавляет шум в
|
||||
# единственный линейный слой, который и так сильно регуляризован
|
||||
# weight decay. Проверено на датасете: с dropout AUC ≈ 0.71,
|
||||
# без него ≈ 0.80.
|
||||
self.quality_head = nn.Sequential(nn.Flatten(), nn.Linear(feature_dim, QUALITY_CLASSES))
|
||||
else:
|
||||
self.quality_head = nn.Sequential(
|
||||
nn.Flatten(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(feature_dim, hidden_dim),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(hidden_dim, QUALITY_CLASSES),
|
||||
)
|
||||
|
||||
self.region_head = nn.Sequential(
|
||||
nn.Flatten(),
|
||||
nn.Dropout(dropout) if head == "mlp" else nn.Identity(),
|
||||
nn.Linear(feature_dim, REGION_CLASSES),
|
||||
)
|
||||
|
||||
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Extract features without classification"""
|
||||
return self.backbone(x)
|
||||
return self.features(x)
|
||||
|
||||
def get_info(self) -> Dict:
|
||||
"""Get model info"""
|
||||
return {
|
||||
'backbone': self.backbone_name,
|
||||
'num_classes': 2,
|
||||
'feature_dim': self.feature_dim,
|
||||
'task': 'binary_quality_classification'
|
||||
}
|
||||
def _head_input(self, feats: torch.Tensor) -> torch.Tensor:
|
||||
"""Сгладить признаки до (B, feature_dim) и применить стандартизацию."""
|
||||
feats = feats.flatten(1)
|
||||
if self.feature_norm:
|
||||
feats = (feats - self.feat_mean) / self.feat_std
|
||||
return feats
|
||||
|
||||
def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]:
|
||||
"""Вернуть логиты качества и логиты анатомической области.
|
||||
|
||||
Голова области обучается как вспомогательная задача (auxiliary loss):
|
||||
она заставляет backbone различать анатомию, но НЕ сдвигает логиты
|
||||
качества. Осторожный сдвиг к «нарушению» при неопределённой анатомии
|
||||
применяется на этапе инференса (см. `src.dxa.inference`), чтобы
|
||||
обучение оставалось устойчивым.
|
||||
"""
|
||||
feats = self._head_input(self.extract_features(x))
|
||||
return {"quality_logits": self.quality_head(feats), "region_logits": self.region_head(feats)}
|
||||
|
||||
def predict_region(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Предсказать анатомическую область: (region_id, вероятность уверенности)."""
|
||||
with torch.no_grad():
|
||||
logits = self.forward(x)["region_logits"]
|
||||
probs = torch.softmax(logits, dim=1)
|
||||
confidence, region_id = probs.max(dim=1)
|
||||
return region_id, confidence
|
||||
|
||||
|
||||
def compute_metrics(
|
||||
logits: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
threshold: float = 0.0,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Метрики бинарной классификации качества.
|
||||
|
||||
Решение принимается по логиту (log-odds) класса «нарушение», а не по
|
||||
вероятности. После обучения линейный зонд разделяет обучающую выборку
|
||||
почти идеально, поэтому вероятности насыщаются в 0/1: порог в единицах
|
||||
вероятности вырождается (например 1e-7), а логиты остаются умеренными.
|
||||
Порог по логитам численно устойчив; для отчётности он переводится в
|
||||
вероятность через сигмоиду.
|
||||
|
||||
ROC-AUC и PR-AUC считаются по вероятностям и от порога не зависят.
|
||||
"""
|
||||
logits = logits.detach().cpu().flatten().float()
|
||||
labels = labels.detach().cpu().flatten().long()
|
||||
probs = torch.sigmoid(logits)
|
||||
predicted = (logits >= threshold).long()
|
||||
|
||||
tp = int(((predicted == 1) & (labels == 1)).sum())
|
||||
tn = int(((predicted == 0) & (labels == 0)).sum())
|
||||
fp = int(((predicted == 1) & (labels == 0)).sum())
|
||||
fn = int(((predicted == 0) & (labels == 1)).sum())
|
||||
|
||||
metrics: Dict[str, Any] = {
|
||||
"accuracy": (tp + tn) / max(len(labels), 1),
|
||||
"precision": precision_score(labels, predicted, zero_division=0),
|
||||
"recall": recall_score(labels, predicted, zero_division=0),
|
||||
"f1": f1_score(labels, predicted, zero_division=0),
|
||||
"tp": tp, "tn": tn, "fp": fp, "fn": fn,
|
||||
"threshold_logit": threshold,
|
||||
"threshold_prob": float(torch.sigmoid(torch.tensor(threshold))),
|
||||
"n": int(len(labels)),
|
||||
"n_pos": int((labels == 1).sum()),
|
||||
}
|
||||
if len(set(labels.tolist())) > 1:
|
||||
metrics["roc_auc"] = roc_auc_score(labels, probs)
|
||||
metrics["pr_auc"] = average_precision_score(labels, probs)
|
||||
metrics["confusion_matrix"] = confusion_matrix(labels, predicted).tolist()
|
||||
else:
|
||||
metrics["roc_auc"] = None
|
||||
metrics["pr_auc"] = None
|
||||
return metrics
|
||||
|
||||
|
||||
def select_threshold(
|
||||
logits: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
min_recall: float = 0.0,
|
||||
) -> Tuple[float, Dict[str, Any]]:
|
||||
"""
|
||||
Подобрать порог по логитам, максимизирующий F1 на валидации.
|
||||
|
||||
При доле брака ~15 % порог 0 (вероятность 0.5) почти всегда даёт нулевой
|
||||
recall, поэтому порог калибруется по валидации и сохраняется в чекпоинт.
|
||||
F1 по всем порогам считается векторно: вызов sklearn на каждый кандидат
|
||||
делает подбор на порядки медленнее и заметно замедляет обучение.
|
||||
|
||||
Args:
|
||||
min_recall: если задано, порог берётся максимальным среди дающих
|
||||
recall не ниже указанного (снижает пропуск брака).
|
||||
"""
|
||||
logits = logits.detach().cpu().flatten().numpy().astype(np.float64)
|
||||
labels = labels.detach().cpu().flatten().long().numpy()
|
||||
|
||||
if len(set(labels.tolist())) < 2:
|
||||
return 0.0, compute_metrics(torch.from_numpy(logits), torch.from_numpy(labels), 0.0)
|
||||
|
||||
# Кандидаты: все наблюдаемые логиты плюс края диапазона.
|
||||
candidates = np.unique(np.concatenate([logits, [logits.min() - 1.0, logits.max() + 1.0]]))
|
||||
predicted = logits[None, :] >= candidates[:, None]
|
||||
|
||||
tp = (predicted & (labels == 1)[None, :]).sum(axis=1).astype(np.float64)
|
||||
fp = (predicted & (labels == 0)[None, :]).sum(axis=1).astype(np.float64)
|
||||
fn = ((~predicted) & (labels == 1)[None, :]).sum(axis=1).astype(np.float64)
|
||||
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
precision = np.where(tp + fp > 0, tp / (tp + fp), 0.0)
|
||||
recall = np.where(tp + fn > 0, tp / (tp + fn), 0.0)
|
||||
f1 = np.where(precision + recall > 0, 2 * precision * recall / (precision + recall), 0.0)
|
||||
|
||||
eligible = np.ones_like(f1, dtype=bool)
|
||||
if min_recall > 0:
|
||||
eligible = recall >= min_recall
|
||||
if not eligible.any():
|
||||
eligible = np.ones_like(f1, dtype=bool)
|
||||
|
||||
# Среди порогов с равным F1 берём минимальный, чтобы не терять recall.
|
||||
masked = np.where(eligible, f1, -1.0)
|
||||
best_threshold = float(candidates[int(np.argmax(np.round(masked, 10)))])
|
||||
|
||||
return best_threshold, compute_metrics(
|
||||
torch.from_numpy(logits), torch.from_numpy(labels), best_threshold
|
||||
)
|
||||
|
||||
|
||||
def per_region_metrics(
|
||||
logits: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
region_ids: torch.Tensor,
|
||||
threshold: float,
|
||||
) -> Dict[str, Dict[str, Any]]:
|
||||
"""
|
||||
Метрики отдельно по анатомическим областям.
|
||||
|
||||
Разбивка нужна для честной интерпретации: в этом датасете нарушения резко
|
||||
неравномерны (в позвоночнике ~29 % снимков с нарушением против ~4–5 % у
|
||||
бёдер), а область почти однозначно определяется по ширине кадра. Поэтому
|
||||
высокий общий AUC может отражать не распознавание дефекта, а различение
|
||||
области исследования.
|
||||
|
||||
Args:
|
||||
region_ids: истинные идентификаторы областей (1 — позвоночник, 2 —
|
||||
правый, 3 — левый). Группировка по истинной области обязательна:
|
||||
по предсказанной метрики смещались бы в сторону тех областей,
|
||||
которые модель путает, и перестали бы показывать реальную картину.
|
||||
"""
|
||||
logits = logits.detach().cpu().flatten()
|
||||
labels = labels.detach().cpu().flatten().long()
|
||||
region_ids = region_ids.detach().cpu().flatten().long()
|
||||
|
||||
result: Dict[str, Dict[str, Any]] = {}
|
||||
for index, name in enumerate(REGIONS, start=1):
|
||||
mask = region_ids == index
|
||||
if int(mask.sum()) == 0:
|
||||
continue
|
||||
result[name] = compute_metrics(logits[mask], labels[mask], threshold)
|
||||
return result
|
||||
|
||||
|
||||
class DXAQualityModel:
|
||||
"""Wrapper for training and inference"""
|
||||
"""Обёртка вокруг сети: обучение, валидация, предсказание, сохранение/загрузка."""
|
||||
|
||||
def __init__(self,
|
||||
model: DXAQualityClassifier,
|
||||
device: str = 'cpu',
|
||||
learning_rate: float = 1e-4):
|
||||
def __init__(
|
||||
self,
|
||||
model: DXAQualityClassifier,
|
||||
device: str = "cpu",
|
||||
learning_rate: float = 1e-4,
|
||||
weight_decay: float = 1e-4,
|
||||
region_loss_weight: float = 0.3,
|
||||
pos_weight: Optional[float] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.device = device
|
||||
self.model.to(device)
|
||||
self.device = torch.device(device)
|
||||
self.model.to(self.device)
|
||||
self._backbone_trainable = True
|
||||
|
||||
self.quality_criterion = nn.CrossEntropyLoss(
|
||||
weight=None if pos_weight is None else torch.tensor([1.0, float(pos_weight)], device=self.device)
|
||||
)
|
||||
self.region_criterion = nn.CrossEntropyLoss(ignore_index=-1)
|
||||
self.region_loss_weight = region_loss_weight
|
||||
|
||||
# Loss and optimizer
|
||||
self.criterion = nn.CrossEntropyLoss()
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
model.parameters(),
|
||||
lr=learning_rate,
|
||||
weight_decay=1e-5
|
||||
model.parameters(), lr=learning_rate, weight_decay=weight_decay
|
||||
)
|
||||
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||
self.optimizer, mode='min', factor=0.5, patience=3
|
||||
self.optimizer, mode="min", factor=0.5, patience=3
|
||||
)
|
||||
self.history: Dict[str, list] = {"train_loss": [], "val_loss": [], "val_f1": [], "val_roc_auc": []}
|
||||
|
||||
def _step(self, batch, train: bool) -> Tuple[float, torch.Tensor, torch.Tensor]:
|
||||
images, meta = batch
|
||||
images = images.to(self.device, non_blocking=True)
|
||||
labels = meta["label"].to(self.device)
|
||||
region_ids = meta["region_id"].to(self.device)
|
||||
|
||||
with torch.set_grad_enabled(train):
|
||||
out = self.model(images)
|
||||
loss = self.quality_criterion(out["quality_logits"], labels)
|
||||
if self.region_loss_weight > 0:
|
||||
loss = loss + self.region_loss_weight * self.region_criterion(out["region_logits"], region_ids)
|
||||
if train:
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=5.0)
|
||||
self.optimizer.step()
|
||||
|
||||
logits = out["quality_logits"][:, 1]
|
||||
return float(loss.item()) * len(labels), logits.detach(), labels
|
||||
|
||||
def set_backbone_trainable(self, trainable: bool) -> None:
|
||||
"""
|
||||
Включить или отключить обучение backbone.
|
||||
|
||||
При ~250 обучающих снимках полный fine-tune ResNet18 быстро
|
||||
переобучается (train F1 -> 1.0 при случайном val AUC). Поэтому backbone
|
||||
заморожен и обучается только голова — это линейный зонд на признаках
|
||||
ImageNet.
|
||||
|
||||
Кроме requires_grad отключается и режим train для слоёв BatchNorm:
|
||||
иначе бегущие статистики продолжают обновляться на обучающих батчах и
|
||||
признаки «уезжают» от тех, на которых оценивалась стандартизация в
|
||||
`fit_feature_norm`. Слой остаётся в eval, поэтому признаки стабильны.
|
||||
"""
|
||||
self._backbone_trainable = trainable
|
||||
for module in self.model.features.modules():
|
||||
if isinstance(module, torch.nn.modules.batchnorm._BatchNorm):
|
||||
module.train(trainable)
|
||||
for param in self.model.features.parameters():
|
||||
param.requires_grad = trainable
|
||||
|
||||
def _apply_train_mode(self) -> None:
|
||||
"""
|
||||
Перевести модель в режим обучения с учётом заморозки backbone.
|
||||
|
||||
`Module.train()` включает train и для слоёв BatchNorm, что при
|
||||
замороженном backbone сдвигало бы бегущие статистики. Поэтому после
|
||||
перевода модели в train слои backbone возвращаются в eval.
|
||||
"""
|
||||
self.model.train()
|
||||
if not getattr(self, "_backbone_trainable", True):
|
||||
for module in self.model.features.modules():
|
||||
if isinstance(module, torch.nn.modules.batchnorm._BatchNorm):
|
||||
module.eval()
|
||||
|
||||
def set_learning_rate(self, lr: float) -> None:
|
||||
"""Задать learning rate всем группам параметров."""
|
||||
for group in self.optimizer.param_groups:
|
||||
group["lr"] = lr
|
||||
|
||||
@torch.no_grad()
|
||||
def fit_feature_norm(self, loader) -> None:
|
||||
"""
|
||||
Оценить среднее и СКО признаков по обучающей выборке и зафиксировать их.
|
||||
|
||||
Стандартизация входа нужна линейной голове: без неё логиты смещены,
|
||||
вероятности скучены у нуля, а подобранный порог теряет смысл.
|
||||
|
||||
Дисперсия считается в два прохода (сначала среднее, затем сумма
|
||||
квадратов отклонений). Формула E[x²]−E[x]² при float32 на признаках
|
||||
порядка 10 даёт погрешность, сопоставимую с самой дисперсией: СКО
|
||||
выходило случайным, из-за чего логиты насыщались и порог вырождался.
|
||||
Буферы не обучаемые, поэтому статистики не «подглядывают» в валидацию.
|
||||
"""
|
||||
if not self.model.feature_norm:
|
||||
return
|
||||
self.model.eval()
|
||||
|
||||
chunks = []
|
||||
for images, _ in loader:
|
||||
chunks.append(self.model.extract_features(images.to(self.device)).flatten(1).cpu().float())
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
# float64 на CPU: размерность мала, а точность здесь критична.
|
||||
feats = torch.cat(chunks).double()
|
||||
mean = feats.mean(dim=0)
|
||||
std = torch.sqrt(((feats - mean) ** 2).mean(dim=0))
|
||||
|
||||
# Нижняя граница СКО: у части размерностей разброс близок к нулю, а
|
||||
# деление на него усиливает шум в десятки раз и насыщает логиты.
|
||||
floor = max(float(std.median()) * 0.25, 1e-6)
|
||||
std = std.clamp(min=floor)
|
||||
|
||||
self.model.feat_mean.copy_(mean.float().to(self.model.feat_mean.device))
|
||||
self.model.feat_std.copy_(std.float().to(self.model.feat_std.device))
|
||||
logger.debug(
|
||||
"Feature norm fitted: %d samples, mean norm %.2f, std median %.4f, floor %.5f",
|
||||
feats.shape[0], float(mean.norm()), float(std.median()), floor,
|
||||
)
|
||||
|
||||
# Training history
|
||||
self.history = {
|
||||
'train_loss': [],
|
||||
'val_loss': [],
|
||||
'train_acc': [],
|
||||
'val_acc': []
|
||||
def train_epoch(self, loader) -> Tuple[float, Dict[str, Any]]:
|
||||
self._apply_train_mode()
|
||||
total_loss = 0.0
|
||||
logits, labels = [], []
|
||||
for batch in loader:
|
||||
loss, l, y = self._step(batch, train=True)
|
||||
total_loss += loss
|
||||
logits.append(l)
|
||||
labels.append(y)
|
||||
n = max(len(loader.dataset), 1)
|
||||
return total_loss / n, compute_metrics(torch.cat(logits), torch.cat(labels))
|
||||
|
||||
def validate(self, loader) -> Tuple[float, Dict[str, Any]]:
|
||||
self.model.eval()
|
||||
total_loss = 0.0
|
||||
logits, labels = [], []
|
||||
for batch in loader:
|
||||
loss, l, y = self._step(batch, train=False)
|
||||
total_loss += loss
|
||||
logits.append(l)
|
||||
labels.append(y)
|
||||
n = max(len(loader.dataset), 1)
|
||||
metrics = compute_metrics(torch.cat(logits), torch.cat(labels))
|
||||
self.scheduler.step(total_loss / n)
|
||||
return total_loss / n, metrics
|
||||
|
||||
def predict_logits(self, loader) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Логиты класса «нарушение» и истинные метки для всего набора."""
|
||||
logits, labels, _ = self.predict_logits_with_regions(loader)
|
||||
return logits, labels
|
||||
|
||||
def predict_logits_with_regions(self, loader) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Логиты, метки и ИСТИННЫЕ идентификаторы областей для всего набора.
|
||||
|
||||
Истинные области берутся из метаданных датасета (они известны при
|
||||
обучении) и нужны для метрик по областям, чтобы оценка не смещалась
|
||||
ошибками головы области.
|
||||
"""
|
||||
self.model.eval()
|
||||
logits, labels, regions = [], [], []
|
||||
with torch.no_grad():
|
||||
for batch in loader:
|
||||
images, meta = batch
|
||||
out = self.model(images.to(self.device))
|
||||
logits.append(out["quality_logits"][:, 1].cpu())
|
||||
labels.append(meta["label"])
|
||||
regions.append(meta["region_id"])
|
||||
return torch.cat(logits), torch.cat(labels), torch.cat(regions)
|
||||
|
||||
def save(self, path: str | Path, preprocess: Optional[PreprocessConfig] = None, **extra) -> None:
|
||||
"""Сохранить чекпоинт вместе с архитектурой и параметрами предобработки."""
|
||||
payload: Dict[str, Any] = {
|
||||
"model_state_dict": self.model.state_dict(),
|
||||
"optimizer_state_dict": self.optimizer.state_dict(),
|
||||
"history": self.history,
|
||||
"backbone": self.model.backbone_name,
|
||||
"head": self.model.head_type,
|
||||
"format_version": 2,
|
||||
"region_loss_weight": self.region_loss_weight,
|
||||
}
|
||||
if preprocess is not None:
|
||||
payload["preprocess"] = preprocess.to_dict()
|
||||
payload.update(extra)
|
||||
Path(path).parent.mkdir(parents=True, exist_ok=True)
|
||||
torch.save(payload, path)
|
||||
|
||||
def train_epoch(self, train_loader) -> Tuple[float, float]:
|
||||
"""Train one epoch"""
|
||||
self.model.train()
|
||||
total_loss = 0
|
||||
correct = 0
|
||||
total = 0
|
||||
|
||||
for images, labels in train_loader:
|
||||
images = images.to(self.device)
|
||||
|
||||
# Handle dict format from dataset
|
||||
if isinstance(labels, dict):
|
||||
labels_tensor = labels['label'].to(self.device)
|
||||
else:
|
||||
labels_tensor = labels.to(self.device)
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
outputs = self.model(images)
|
||||
loss = self.criterion(outputs, labels_tensor)
|
||||
|
||||
loss.backward()
|
||||
self.optimizer.step()
|
||||
|
||||
total_loss += loss.item() * images.size(0)
|
||||
_, predicted = outputs.max(1)
|
||||
correct += predicted.eq(labels_tensor).sum().item()
|
||||
total += labels_tensor.size(0)
|
||||
|
||||
return total_loss / total, correct / total
|
||||
|
||||
def validate(self, val_loader) -> Tuple[float, float]:
|
||||
"""Validate"""
|
||||
self.model.eval()
|
||||
total_loss = 0
|
||||
correct = 0
|
||||
total = 0
|
||||
|
||||
with torch.no_grad():
|
||||
for images, labels in val_loader:
|
||||
images = images.to(self.device)
|
||||
|
||||
# Handle dict format from dataset
|
||||
if isinstance(labels, dict):
|
||||
labels_tensor = labels['label'].to(self.device)
|
||||
else:
|
||||
labels_tensor = labels.to(self.device)
|
||||
|
||||
outputs = self.model(images)
|
||||
loss = self.criterion(outputs, labels_tensor)
|
||||
|
||||
total_loss += loss.item() * images.size(0)
|
||||
_, predicted = outputs.max(1)
|
||||
correct += predicted.eq(labels_tensor).sum().item()
|
||||
total += labels_tensor.size(0)
|
||||
|
||||
return total_loss / total, correct / total
|
||||
|
||||
def predict(self, images: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Predict on batch of images"""
|
||||
self.model.eval()
|
||||
with torch.no_grad():
|
||||
images = images.to(self.device)
|
||||
outputs = self.model(images)
|
||||
probs = torch.softmax(outputs, dim=1)
|
||||
preds = outputs.argmax(dim=1)
|
||||
return preds, probs
|
||||
|
||||
def save(self, path: str):
|
||||
"""Save model"""
|
||||
torch.save({
|
||||
'model_state_dict': self.model.state_dict(),
|
||||
'optimizer_state_dict': self.optimizer.state_dict(),
|
||||
'history': self.history
|
||||
}, path)
|
||||
|
||||
def load(self, path: str):
|
||||
"""Load model"""
|
||||
checkpoint = torch.load(path, map_location=self.device)
|
||||
self.model.load_state_dict(checkpoint['model_state_dict'])
|
||||
self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
|
||||
self.history = checkpoint.get('history', self.history)
|
||||
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',
|
||||
num_classes: int = 2,
|
||||
pretrained: bool = True,
|
||||
device: str = 'cpu') -> DXAQualityModel:
|
||||
"""Create model instance"""
|
||||
model = DXAQualityClassifier(
|
||||
backbone=backbone,
|
||||
num_classes=num_classes,
|
||||
pretrained=pretrained
|
||||
def create_model(
|
||||
backbone: str = "resnet18",
|
||||
pretrained: bool = True,
|
||||
device: str = "cpu",
|
||||
head: str = "linear",
|
||||
**kwargs,
|
||||
) -> DXAQualityModel:
|
||||
"""Создать обёртку модели с заданным backbone и головой."""
|
||||
return DXAQualityModel(
|
||||
DXAQualityClassifier(backbone=backbone, pretrained=pretrained, head=head),
|
||||
device=device,
|
||||
**kwargs,
|
||||
)
|
||||
return DXAQualityModel(model, device=device)
|
||||
|
|
|
|||
661
src/dxa/train.py
661
src/dxa/train.py
|
|
@ -1,248 +1,529 @@
|
|||
#!/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
|
||||
Особенности, отличающие этот скрипт от прежней версии:
|
||||
|
||||
Использование:
|
||||
python src/dxa/train.py --epochs 20 --batch-size 8
|
||||
1. **Метки из имён файлов.** Раньше метка бралась из Excel по исследованию и
|
||||
раздавалась всем снимкам этого исследования; файлы с явной меткой (`_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
|
||||
import sys
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
from torch.utils.data import WeightedRandomSampler
|
||||
|
||||
# Add src to path
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||
|
||||
from src.dxa.dataset import create_dataloaders
|
||||
from src.dxa.model import create_model
|
||||
from src.dxa.dataset import DXADataset, make_datasets
|
||||
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():
|
||||
"""Get best available device"""
|
||||
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():
|
||||
return 'mps'
|
||||
elif torch.cuda.is_available():
|
||||
return 'cuda'
|
||||
else:
|
||||
return 'cpu'
|
||||
return "mps"
|
||||
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.
|
||||
|
||||
Рассчитывает:
|
||||
- Accuracy: доля правильных предсказаний
|
||||
- Precision: точность (доля TP среди предсказанных positive)
|
||||
- Recall: полнота (доля TP среди реальных positive)
|
||||
- F1: гармоническое среднее precision и recall
|
||||
|
||||
Args:
|
||||
preds: Предсказания модели (numpy array)
|
||||
labels: Истинные метки (numpy array)
|
||||
|
||||
Returns:
|
||||
Dict с метриками
|
||||
balance='sampler' выравнивает классы взвешенной выборкой: важно при доле
|
||||
брака ~15 %, иначе модель сходится к «всё хорошее» и даёт высокую accuracy
|
||||
при нулевом recall.
|
||||
"""
|
||||
preds = np.array(preds)
|
||||
labels = np.array(labels)
|
||||
sampler = None
|
||||
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
|
||||
accuracy = (preds == labels).mean()
|
||||
|
||||
# True/False positives/negatives
|
||||
tp = ((preds == 1) & (labels == 1)).sum()
|
||||
tn = ((preds == 0) & (labels == 0)).sum()
|
||||
fp = ((preds == 1) & (labels == 0)).sum()
|
||||
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
|
||||
}
|
||||
return torch.utils.data.DataLoader(
|
||||
dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=shuffle,
|
||||
sampler=sampler,
|
||||
num_workers=num_workers,
|
||||
pin_memory=torch.cuda.is_available(),
|
||||
)
|
||||
|
||||
|
||||
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.
|
||||
|
||||
Этапы:
|
||||
1. Определение устройства (MPS/CUDA/CPU)
|
||||
2. Создание директории для сохранения модели
|
||||
3. Загрузка данных (DataLoader)
|
||||
4. Создание модели
|
||||
5. Цикл обучения по эпохам:
|
||||
- Обучение на train set
|
||||
- Валидация на val set
|
||||
- Расчет метрик (accuracy, precision, recall, F1)
|
||||
- Сохранение лучшей модели по F1
|
||||
6. Сохранение финальной модели
|
||||
|
||||
Args:
|
||||
args: Аргументы командной строки
|
||||
ROC-AUC выбран первичным критерием отбора, потому что на валидации всего
|
||||
~8 изображений с нарушением, и F1 принимает лишь несколько значений —
|
||||
выбор эпохи по F1 шумит и переобучает порог. AUC использует ранжирование
|
||||
всех изображений и заметно стабильнее. Сам порог всё равно подбирается
|
||||
по F1 (см. `select_threshold`).
|
||||
"""
|
||||
auc = metrics.get("roc_auc") or 0.0
|
||||
if auc > best_auc + min_delta:
|
||||
return True
|
||||
if abs(auc - best_auc) <= min_delta and metrics["f1"] > best_f1 + min_delta:
|
||||
return True
|
||||
return False
|
||||
|
||||
# Setup
|
||||
device = get_device()
|
||||
print(f"Using device: {device}")
|
||||
|
||||
# Create output directory
|
||||
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.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Create dataloaders
|
||||
print("Loading data...")
|
||||
train_loader, val_loader = create_dataloaders(
|
||||
preprocess = PreprocessConfig(norm=args.norm, imagenet_norm=not args.no_imagenet_norm)
|
||||
logger.info("Loading dataset from %s", args.data_root)
|
||||
train_ds, val_ds, preprocess = make_datasets(
|
||||
data_root=args.data_root,
|
||||
annotation_path=args.annotation_path,
|
||||
batch_size=args.batch_size,
|
||||
input_size=(args.input_size, args.input_size),
|
||||
num_workers=args.num_workers
|
||||
input_size=args.input_size,
|
||||
val_fraction=args.val_fraction,
|
||||
seed=args.seed,
|
||||
preprocess=preprocess,
|
||||
)
|
||||
|
||||
print(f"Train samples: {len(train_loader.dataset)}")
|
||||
print(f"Val samples: {len(val_loader.dataset)}")
|
||||
if len(train_ds) == 0 or len(val_ds) == 0:
|
||||
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))
|
||||
|
||||
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)
|
||||
|
||||
# Create model
|
||||
print(f"Creating model: {args.backbone}")
|
||||
model = create_model(
|
||||
backbone=args.backbone,
|
||||
num_classes=2,
|
||||
pretrained=True,
|
||||
device=device
|
||||
pretrained=not args.no_pretrained,
|
||||
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
|
||||
threshold = 0.0
|
||||
best_score, best_f1, best_epoch = -1.0, -1.0, 0
|
||||
epochs_without_improvement = 0
|
||||
history: List[Dict] = []
|
||||
recent_auc: deque = deque(maxlen=args.select_window)
|
||||
|
||||
for epoch in range(args.epochs):
|
||||
print(f"\n{'='*50}")
|
||||
print(f"Epoch {epoch+1}/{args.epochs}")
|
||||
print(f"{'='*50}")
|
||||
# Фаза 1: backbone заморожен, обучается только голова. При ~250 уникальных
|
||||
# снимках полный fine-tune даёт переобучение (train F1 -> 1.0, val AUC ~0.5),
|
||||
# а линейный зонд на признаках ImageNet держит val AUC ~0.80.
|
||||
# freeze_epochs = -1 означает «заморозить навсегда».
|
||||
freeze_epochs = args.epochs if args.freeze_epochs < 0 else args.freeze_epochs
|
||||
if freeze_epochs > 0:
|
||||
model.set_backbone_trainable(False)
|
||||
# Стандартизация признаков по обучающей выборке: без неё логиты смещены,
|
||||
# вероятности скучены у нуля и подобранный порог теряет смысл.
|
||||
model.fit_feature_norm(train_loader)
|
||||
logger.info("Phase 1 (epochs 1..%d): backbone frozen, head only, lr=%.1e",
|
||||
freeze_epochs, args.learning_rate)
|
||||
|
||||
# Train
|
||||
train_loss, train_acc = model.train_epoch(train_loader)
|
||||
for epoch in range(1, args.epochs + 1):
|
||||
if epoch == freeze_epochs + 1 and 0 < freeze_epochs < args.epochs:
|
||||
model.set_backbone_trainable(True)
|
||||
# При размораживании backbone нужен меньший шаг, иначе предобученные
|
||||
# признаки разрушаются за несколько эпох.
|
||||
model.set_learning_rate(args.learning_rate / 10)
|
||||
logger.info("Phase 2: backbone unfrozen, lr=%.1e", args.learning_rate / 10)
|
||||
|
||||
# Validate
|
||||
val_loss, val_acc = model.validate(val_loader)
|
||||
started = time.time()
|
||||
train_loss, train_metrics = model.train_epoch(train_loader)
|
||||
val_loss, _ = model.validate(val_loader)
|
||||
|
||||
# Compute detailed metrics
|
||||
model.model.eval()
|
||||
all_preds = []
|
||||
all_labels = []
|
||||
# Порог по логитам подбирается на валидации на каждой эпохе: при доле
|
||||
# брака ~15 % порог 0 (вероятность 0.5) почти всегда даёт нулевой recall.
|
||||
logits, labels = model.predict_logits(val_loader)
|
||||
threshold, val_metrics = select_threshold(logits, labels, min_recall=args.min_recall)
|
||||
elapsed = time.time() - started
|
||||
|
||||
with torch.no_grad():
|
||||
for images, labels in val_loader:
|
||||
images = images.to(device)
|
||||
# labels is a dict with 'label' key from our dataset
|
||||
if isinstance(labels, dict):
|
||||
labels_arr = labels['label'].to(device)
|
||||
else:
|
||||
labels_arr = labels.to(device)
|
||||
model.history["train_loss"].append(train_loss)
|
||||
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"))
|
||||
|
||||
preds, _ = model.predict(images)
|
||||
all_preds.extend(preds.cpu().numpy())
|
||||
all_labels.extend(labels_arr.cpu().numpy())
|
||||
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,
|
||||
)
|
||||
|
||||
metrics = compute_metrics(all_preds, all_labels)
|
||||
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))}})
|
||||
|
||||
print(f"\nResults:")
|
||||
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:")
|
||||
print(f" Precision: {metrics['precision']:.4f}")
|
||||
print(f" Recall: {metrics['recall']:.4f}")
|
||||
print(f" F1: {metrics['f1']:.4f}")
|
||||
recent_auc.append(val_metrics.get("roc_auc"))
|
||||
score = smoothed_score(recent_auc, args.select_window)
|
||||
# Прогрев: до накопления окна чекпоинт не сохраняется, иначе им станет
|
||||
# случайно удачная ранняя эпоха с ещё не обученной моделью.
|
||||
warmed_up = epoch >= args.select_window
|
||||
|
||||
# Save best model
|
||||
if metrics['f1'] > best_val_f1:
|
||||
best_val_f1 = metrics['f1']
|
||||
best_epoch = epoch + 1
|
||||
model.save(str(output_dir / 'best_model.pth'))
|
||||
print(f" ✅ Saved best model (F1: {best_val_f1:.4f})")
|
||||
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
|
||||
|
||||
# Save checkpoint
|
||||
if (epoch + 1) % args.save_every == 0:
|
||||
model.save(str(output_dir / f'checkpoint_epoch_{epoch+1}.pth'))
|
||||
# Финальная валидация лучшим чекпоинтом, а не последним.
|
||||
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)
|
||||
|
||||
print(f"\n{'='*50}")
|
||||
print(f"Training complete!")
|
||||
print(f"Best F1: {best_val_f1:.4f} at epoch {best_epoch}")
|
||||
print(f"{'='*50}")
|
||||
model.save(
|
||||
output_dir / "dxa_model_final.pth",
|
||||
preprocess=preprocess,
|
||||
threshold=threshold,
|
||||
epoch=epoch,
|
||||
)
|
||||
|
||||
# Save final model
|
||||
model.save(str(output_dir / 'final_model.pth'))
|
||||
print(f"Final model saved to {output_dir / 'final_model.pth'}")
|
||||
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():
|
||||
parser = argparse.ArgumentParser(description='Train DXA Quality Classifier')
|
||||
|
||||
# Data arguments
|
||||
parser.add_argument('--data-root', type=str,
|
||||
default='dataset_hack',
|
||||
help='Path to data directory')
|
||||
parser.add_argument('--annotation-path', type=str,
|
||||
default='dataset_hack/НД_для_обучения/разметка.xlsx',
|
||||
help='Path to annotation Excel file')
|
||||
parser.add_argument('--input-size', type=int, default=224,
|
||||
help='Input image size')
|
||||
parser.add_argument('--batch-size', type=int, default=8,
|
||||
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)
|
||||
def _to_builtin(value):
|
||||
"""Привести значения numpy к встроенным типам, чтобы отчёт сериализовался в JSON."""
|
||||
if isinstance(value, dict):
|
||||
return {k: _to_builtin(v) for k, v in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_to_builtin(v) for v in value]
|
||||
if isinstance(value, (np.integer,)):
|
||||
return int(value)
|
||||
if isinstance(value, (np.floating,)):
|
||||
return float(value)
|
||||
if isinstance(value, np.bool_):
|
||||
return bool(value)
|
||||
return value
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
def _write_markdown_report(path: Path, report: Dict) -> None:
|
||||
"""Краткий отчёт об обучении для документации и презентации."""
|
||||
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())
|
||||
|
|
|
|||
498
src/main.py
498
src/main.py
|
|
@ -20,11 +20,13 @@ FastAPI сервер для оценки качества DXA исследова
|
|||
- POST /api/v1/batch - Пакетный анализ
|
||||
- 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 pathlib import Path
|
||||
from fastapi.responses import StreamingResponse
|
||||
from typing import Dict
|
||||
from typing import Dict, Optional
|
||||
import numpy as np
|
||||
import pydicom
|
||||
import pandas as pd
|
||||
|
|
@ -35,11 +37,13 @@ import base64
|
|||
|
||||
from starlette.responses import JSONResponse, FileResponse
|
||||
|
||||
from src.dxa.model import create_model
|
||||
from src.dxa.inference import determine_anatomical_region
|
||||
from src.dxa.inference import (
|
||||
REASON_BY_TYPE,
|
||||
load_model as load_dxa_checkpoint,
|
||||
predict_from_bytes,
|
||||
)
|
||||
from src.utils.utils import get_device
|
||||
from src.quality.detailed_assessment import generate_quality_report
|
||||
from src.quality.quality_scorer import QualityScorer, convert_to_serializable
|
||||
from src.quality.quality_scorer import convert_to_serializable
|
||||
|
||||
# Create app
|
||||
app = FastAPI(
|
||||
|
|
@ -56,8 +60,11 @@ static_dir = Path(__file__).parent / "api/static"
|
|||
if static_dir.exists():
|
||||
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_preprocess = None
|
||||
dxa_threshold = 0.0
|
||||
device = None
|
||||
|
||||
|
||||
|
|
@ -66,26 +73,27 @@ def load_model():
|
|||
Загрузка модели классификатора качества DXA.
|
||||
|
||||
Модель загружается глобально при первом запросе и сохраняется в памяти.
|
||||
Это позволяет избежать повторной загрузки при каждом запросе.
|
||||
Путь к чекпоинту берётся из переменной окружения DXA_MODEL_PATH, чтобы
|
||||
контейнер не зависел от текущего рабочего каталога.
|
||||
|
||||
Returns:
|
||||
DXAQualityModel: Обученная модель или None при ошибке
|
||||
"""
|
||||
global dxa_model, device
|
||||
global dxa_model, dxa_preprocess, dxa_threshold, device
|
||||
|
||||
if dxa_model is None:
|
||||
device = get_device()
|
||||
print(f"Loading DXA model on {device}...")
|
||||
|
||||
try:
|
||||
dxa_model = create_model(
|
||||
backbone='resnet18',
|
||||
pretrained=False,
|
||||
device=device
|
||||
checkpoint = load_dxa_checkpoint(MODEL_PATH, device=device)
|
||||
dxa_model = checkpoint.model
|
||||
dxa_preprocess = checkpoint.preprocess
|
||||
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:
|
||||
print(f"Error loading model: {e}")
|
||||
dxa_model = None
|
||||
|
|
@ -97,14 +105,9 @@ def preprocess_dicom(dcm_bytes: bytes, input_size: int = 224):
|
|||
"""
|
||||
Предобработка DICOM изображения для модели.
|
||||
|
||||
Этапы:
|
||||
1. Сохранение байтов во временный файл
|
||||
2. Чтение DICOM (pydicom)
|
||||
3. Нормализация значений пикселей
|
||||
4. Создание 3-канального изображения
|
||||
5. Изменение размера до input_size x input_size
|
||||
6. Нормализация для PyTorch (деление на 255)
|
||||
7. Преобразование в тензор
|
||||
Сохранена для обратной совместимости: эндпоинты используют
|
||||
`predict_from_bytes`, который применяет ту же предобработку, что и при
|
||||
обучении (параметры берутся из чекпоинта, а не задаются заново).
|
||||
|
||||
Args:
|
||||
dcm_bytes: Байты DICOM файла
|
||||
|
|
@ -113,32 +116,52 @@ def preprocess_dicom(dcm_bytes: bytes, input_size: int = 224):
|
|||
Returns:
|
||||
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:
|
||||
f.write(dcm_bytes)
|
||||
dcm_path = f.name
|
||||
|
||||
ds = pydicom.dcmread(dcm_path)
|
||||
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]
|
||||
|
||||
# Normalize
|
||||
img = (img - img.min()) / (img.max() - img.min() + 1e-8)
|
||||
tensor = torch.from_numpy(preprocess_from_array(img, PreprocessConfig(input_size=input_size)))
|
||||
return tensor.float().unsqueeze(0), ds
|
||||
|
||||
# 3-channel
|
||||
img = np.stack([img] * 3, axis=0)
|
||||
|
||||
# Resize
|
||||
img = (img * 255).astype(np.uint8)
|
||||
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)
|
||||
img = img.astype(np.float32) / 255.0
|
||||
def predict_upload(dcm_bytes: bytes):
|
||||
"""
|
||||
Единая точка предсказания для HTTP-эндпоинтов.
|
||||
|
||||
# Tensor
|
||||
img = torch.from_numpy(img).float().unsqueeze(0)
|
||||
Возвращает (prediction, dataset_metadata) или (None, None), если модель не
|
||||
загружена. Использование общего пути гарантирует, что предобработка и
|
||||
решающий порог совпадают с 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
|
||||
|
|
@ -191,37 +214,14 @@ async def analyze_dicom(file: UploadFile = File(...)):
|
|||
content={"error": "Model not loaded"}
|
||||
)
|
||||
|
||||
# Read file
|
||||
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"})
|
||||
|
||||
# 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()
|
||||
|
||||
# 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
|
||||
return prediction_to_result(prediction, ds, file.filename)
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
return JSONResponse(
|
||||
status_code=500,
|
||||
content={
|
||||
|
|
@ -259,65 +259,27 @@ async def analyze_dicom_detailed(
|
|||
content={"error": "Model not loaded"}
|
||||
)
|
||||
|
||||
# Read file
|
||||
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
|
||||
img_array = ds.pixel_array.astype(np.float32)
|
||||
img_array = (img_array - img_array.min()) / (img_array.max() - img_array.min() + 1e-8)
|
||||
|
||||
# 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 = prediction_to_result(prediction, ds, file.filename)
|
||||
result["quality_label"] = (
|
||||
"Качественное изображение" if prediction.quality_class == 0 else "Есть нарушение качества"
|
||||
)
|
||||
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', '')
|
||||
quality_report["image_uid"] = getattr(ds, 'SOPInstanceUID', '')
|
||||
|
||||
# Add visualization if requested
|
||||
# Дополнительный функционал: схематичная маска костной ткани
|
||||
# (90-й перцентиль яркости) как основа для визуализации нарушения.
|
||||
if include_visualization:
|
||||
# Create mask visualization
|
||||
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')
|
||||
result["mask"] = _mask_base64(dcm_bytes)
|
||||
else:
|
||||
quality_report["mask"] = None
|
||||
result["mask"] = None
|
||||
|
||||
# Convert numpy types to Python types for JSON serialization
|
||||
quality_report = convert_to_serializable(quality_report)
|
||||
|
||||
return quality_report
|
||||
return convert_to_serializable(result)
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
|
@ -331,11 +293,48 @@ 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")
|
||||
async def batch_analyze(files: list[UploadFile] = File(...)):
|
||||
"""Batch analyze multiple DICOM files"""
|
||||
results = []
|
||||
|
||||
"""Пакетный анализ нескольких DICOM файлов."""
|
||||
model = load_model()
|
||||
if model is None:
|
||||
return JSONResponse(
|
||||
|
|
@ -343,35 +342,19 @@ async def batch_analyze(files: list[UploadFile] = File(...)):
|
|||
content={"error": "Model not loaded"}
|
||||
)
|
||||
|
||||
results = []
|
||||
for file in files:
|
||||
try:
|
||||
dcm_bytes = await file.read()
|
||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
||||
|
||||
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()
|
||||
|
||||
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"
|
||||
})
|
||||
prediction, ds = predict_upload(dcm_bytes)
|
||||
if prediction is None:
|
||||
raise RuntimeError("Model not loaded")
|
||||
results.append(prediction_to_result(prediction, ds, file.filename))
|
||||
except Exception as e:
|
||||
results.append({
|
||||
"filename": file.filename,
|
||||
"error": str(e),
|
||||
"processing_status": "Failure"
|
||||
"processing_status": "Failure",
|
||||
})
|
||||
|
||||
return {"results": results}
|
||||
|
|
@ -379,9 +362,13 @@ async def batch_analyze(files: list[UploadFile] = File(...)):
|
|||
|
||||
@app.post("/api/v1/export")
|
||||
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()
|
||||
if model is None:
|
||||
return JSONResponse(
|
||||
|
|
@ -389,66 +376,55 @@ async def export_results(files: list[UploadFile] = File(...)):
|
|||
content={"error": "Model not loaded"}
|
||||
)
|
||||
|
||||
results = []
|
||||
for file in files:
|
||||
started = time.time()
|
||||
try:
|
||||
dcm_bytes = await file.read()
|
||||
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
||||
|
||||
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()
|
||||
|
||||
h, w = ds.pixel_array.shape
|
||||
region = determine_anatomical_region(file.filename) if file.filename else ('hip' if h < 270 else 'spine')
|
||||
|
||||
prediction, ds = predict_upload(dcm_bytes)
|
||||
if prediction is None:
|
||||
raise RuntimeError("Model not loaded")
|
||||
results.append({
|
||||
'filename': file.filename,
|
||||
'path_to_study': str(Path(file.filename).parent if file.filename else ''),
|
||||
'study_uid': getattr(ds, 'StudyInstanceUID', ''),
|
||||
'image_uid': getattr(ds, 'SOPInstanceUID', ''),
|
||||
'anatomical_region': region,
|
||||
'quality_class': int(pred),
|
||||
'violation_type': 'quality_violation_detected' if pred == 1 else '',
|
||||
'confidence': round(confidence, 4),
|
||||
'processing_status': 'Success'
|
||||
"path_to_study": f"upload://{file.filename or 'unknown'}",
|
||||
"study_uid": str(getattr(ds, "StudyInstanceUID", "") or ""),
|
||||
"image_uid": str(getattr(ds, "SOPInstanceUID", "") or ""),
|
||||
"anatomical_region": prediction.anatomical_region,
|
||||
"quality_class": prediction.quality_class,
|
||||
"violation_type": prediction.violation_type,
|
||||
"processing_status": "Success",
|
||||
"time_of_processing": round(time.time() - started, 4),
|
||||
"confidence": round(prediction.prob, 4),
|
||||
"violation_reason": prediction.violation_reason,
|
||||
})
|
||||
except Exception as e:
|
||||
results.append({
|
||||
'filename': file.filename,
|
||||
'path_to_study': '',
|
||||
'study_uid': '',
|
||||
'image_uid': '',
|
||||
'anatomical_region': 'unknown',
|
||||
'quality_class': -1,
|
||||
'violation_type': '',
|
||||
'confidence': 0.0,
|
||||
'processing_status': f'Failure: {str(e)[:80]}'
|
||||
"path_to_study": f"upload://{file.filename or 'unknown'}",
|
||||
"study_uid": "",
|
||||
"image_uid": "",
|
||||
"anatomical_region": "unknown",
|
||||
"quality_class": -1,
|
||||
"violation_type": "",
|
||||
"processing_status": f"Failure: {str(e)[:80]}",
|
||||
"time_of_processing": round(time.time() - started, 4),
|
||||
"confidence": 0.0,
|
||||
"violation_reason": "",
|
||||
})
|
||||
|
||||
# Create DataFrame and export to Excel
|
||||
df = pd.DataFrame(results)
|
||||
columns = [
|
||||
"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()
|
||||
df.to_excel(buffer, index=False, engine='openpyxl')
|
||||
df.to_excel(buffer, index=False, engine="openpyxl")
|
||||
buffer.seek(0)
|
||||
|
||||
return StreamingResponse(
|
||||
buffer,
|
||||
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'}
|
||||
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'},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -460,8 +436,6 @@ async def analyze_dicom_sr(file: UploadFile = File(...)):
|
|||
Returns a text representation of DICOM SR with standardized codes.
|
||||
"""
|
||||
try:
|
||||
# Use the detailed analysis
|
||||
# Reuse the detailed analysis logic
|
||||
model = load_model()
|
||||
if model is None:
|
||||
return JSONResponse(
|
||||
|
|
@ -469,60 +443,40 @@ async def analyze_dicom_sr(file: UploadFile = File(...)):
|
|||
content={"error": "Model not loaded"}
|
||||
)
|
||||
|
||||
# Read file
|
||||
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
|
||||
img_array = ds.pixel_array.astype(np.float32)
|
||||
img_array = (img_array - img_array.min()) / (img_array.max() - img_array.min() + 1e-8)
|
||||
|
||||
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()
|
||||
|
||||
# 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 = prediction_to_result(prediction, ds)
|
||||
snomed_map = {
|
||||
"artifact_motion": ("Motion artifact", "WARNING"),
|
||||
"artifact_other": ("Foreign object artifact", "WARNING"),
|
||||
"rotation": ("Rotational misalignment", "WARNING"),
|
||||
"roi_error": ("Region of interest mismatch", "WARNING"),
|
||||
"incomplete_view": ("Incomplete anatomy", "WARNING"),
|
||||
"position_error": ("Positioning deviation", "WARNING"),
|
||||
}
|
||||
label, completion = snomed_map.get(
|
||||
prediction.violation_type, ("DXA image quality acceptable", "FINAL")
|
||||
)
|
||||
quality_report["violation_label"] = label
|
||||
quality_report["completion"] = completion
|
||||
|
||||
# Generate DICOM SR text representation
|
||||
sr_content = generate_dicom_sr_text(
|
||||
study_uid=getattr(ds, 'StudyInstanceUID', ''),
|
||||
image_uid=getattr(ds, 'SOPInstanceUID', ''),
|
||||
quality_report=quality_report
|
||||
study_uid=quality_report["study_uid"],
|
||||
image_uid=quality_report["image_uid"],
|
||||
quality_report=quality_report,
|
||||
)
|
||||
|
||||
return {
|
||||
"format": "DICOM SR (Text)",
|
||||
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
||||
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
||||
"sr_content": sr_content
|
||||
"study_uid": quality_report["study_uid"],
|
||||
"image_uid": quality_report["image_uid"],
|
||||
"sr_content": sr_content,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
return JSONResponse(
|
||||
status_code=500,
|
||||
content={
|
||||
|
|
@ -554,20 +508,25 @@ def generate_dicom_sr_text(study_uid: str, image_uid: str, quality_report: Dict)
|
|||
str: Текстовое представление SR отчета
|
||||
"""
|
||||
|
||||
# Map violation types to DICOM codes (simplified)
|
||||
# Соответствие типа нарушения коду отчёта. Коды условные (локальные), так
|
||||
# как полноценного справочника SNOMED/DICOM для контролёра качества DXA в
|
||||
# наборе нет; текстовая формулировка приводится рядом.
|
||||
violation_code_map = {
|
||||
"correct": ("113001", "DXA Quality Assessment", "FINAL"),
|
||||
"position_error": ("123456", "Positioning Error", "WARNING"),
|
||||
"artifact_motion": ("234567", "Motion Artifact", "WARNING"),
|
||||
"artifact_other": ("234568", "Other Artifact", "WARNING"),
|
||||
"labeling_error": ("345678", "Labeling Error", "WARNING"),
|
||||
"incomplete_view": ("456789", "Incomplete View", "WARNING"),
|
||||
"roi_error": ("567890", "ROI Error", "WARNING"),
|
||||
"rotation": ("678901", "Rotation Error", "WARNING")
|
||||
"": ("113001", "DXA image quality acceptable", "FINAL"),
|
||||
"position_error": ("123456", "Positioning deviation", "WARNING"),
|
||||
"artifact_motion": ("234567", "Motion artifact", "WARNING"),
|
||||
"artifact_other": ("234568", "Other artifact", "WARNING"),
|
||||
"labeling_error": ("345678", "Labeling error", "WARNING"),
|
||||
"incomplete_view": ("456789", "Incomplete anatomy", "WARNING"),
|
||||
"roi_error": ("567890", "Region of interest mismatch", "WARNING"),
|
||||
"rotation": ("678901", "Rotational misalignment", "WARNING"),
|
||||
"quality_violation_detected": ("999001", "Image quality violation", "WARNING"),
|
||||
}
|
||||
|
||||
violation_type = quality_report.get("violation_type", "correct")
|
||||
code, label, completion = violation_code_map.get(violation_type, ("999999", "Unknown", "UNKNOWN"))
|
||||
violation_type = (quality_report.get("violation_type") or "").strip()
|
||||
code, label, completion = violation_code_map.get(
|
||||
violation_type, ("999999", "Unknown finding", "UNKNOWN")
|
||||
)
|
||||
|
||||
sr_lines = [
|
||||
"DICOM Structured Report - DXA Quality Assessment",
|
||||
|
|
@ -578,43 +537,36 @@ def generate_dicom_sr_text(study_uid: str, image_uid: str, quality_report: Dict)
|
|||
"Procedure Report:",
|
||||
f" - Anatomical Region: {quality_report.get('anatomical_region', '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:",
|
||||
f" - Violation Type Code: {code}",
|
||||
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')}",
|
||||
"",
|
||||
"Detailed Metrics:",
|
||||
"Image Characteristics:",
|
||||
]
|
||||
|
||||
# Add metrics
|
||||
metrics = quality_report.get("metrics", {})
|
||||
# Числовые характеристики изображения, использованные при решении.
|
||||
metrics = quality_report.get("metrics") or {}
|
||||
for key, label_ru in (
|
||||
("laplacian_variance", "Local sharpness metric"),
|
||||
("bright_fraction", "Dense-pixel fraction"),
|
||||
("bbox_aspect", "Bright region aspect ratio"),
|
||||
):
|
||||
if key in metrics:
|
||||
sr_lines.append(f" - {label_ru}: {float(metrics[key]):.5f}")
|
||||
|
||||
motion = metrics.get("motion", {})
|
||||
if motion:
|
||||
sr_lines.append(f" Motion Detection:")
|
||||
sr_lines.append(f" - Motion Detected: {motion.get('motion_detected', False)}")
|
||||
sr_lines.append(f" - Severity: {motion.get('severity', 'NONE')}")
|
||||
if not metrics:
|
||||
sr_lines.append(" - No image metrics available")
|
||||
|
||||
artifacts = metrics.get("artifacts", {})
|
||||
if artifacts:
|
||||
sr_lines.append(f" Artifact Detection:")
|
||||
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([
|
||||
"",
|
||||
"Completion Flag: " + completion,
|
||||
"Verification Flag: UNVERIFIED"
|
||||
"Verification Flag: UNVERIFIED",
|
||||
])
|
||||
|
||||
return "\n".join(sr_lines)
|
||||
|
|
|
|||
Loading…
Reference in New Issue