develop - hack_2026
This commit is contained in:
parent
dccade80fd
commit
a5dff35ce4
41
Dockerfile
41
Dockerfile
|
|
@ -1,22 +1,16 @@
|
||||||
# --- Этап сборки ---
|
# DXA Quality Assessment - Docker Container
|
||||||
FROM python:3.10-slim AS builder
|
# Build: docker build -t dxa-quality .
|
||||||
|
# Run: docker run -v /path/to/data:/data -p 8000:8000 dxa-quality
|
||||||
|
|
||||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
# Base image with GPU support
|
||||||
|
FROM nvidia/cuda:11.8.0-cudnn8-runtime-ubuntu22.04 AS builder
|
||||||
|
|
||||||
|
# Install Python and build tools
|
||||||
|
RUN apt-get update && apt-get install -y \
|
||||||
|
python3.10 \
|
||||||
|
python3-pip \
|
||||||
build-essential \
|
build-essential \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
libgl1-mesa-glx \
|
||||||
|
|
||||||
WORKDIR /app
|
|
||||||
COPY requirements.txt .
|
|
||||||
RUN pip install --no-cache-dir --user -r requirements.txt \
|
|
||||||
--extra-index-url https://download.pytorch.org/whl/cpu
|
|
||||||
|
|
||||||
# --- Финальный этап ---
|
|
||||||
FROM python:3.10-slim
|
|
||||||
|
|
||||||
# Только runtime-библиотеки, без gcc
|
|
||||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
|
||||||
libgl1 \
|
|
||||||
libglx-mesa0 \
|
|
||||||
libglib2.0-0 \
|
libglib2.0-0 \
|
||||||
libsm6 \
|
libsm6 \
|
||||||
libxext6 \
|
libxext6 \
|
||||||
|
|
@ -26,15 +20,20 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Копируем установленные пакеты из builder-этапа
|
# Copy requirements and install Python dependencies
|
||||||
COPY --from=builder /root/.local /root/.local
|
COPY requirements.txt .
|
||||||
ENV PATH=/root/.local/bin:$PATH
|
RUN pip3 install --no-cache-dir -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu118
|
||||||
|
|
||||||
|
# Copy application code
|
||||||
COPY . .
|
COPY . .
|
||||||
RUN mkdir -p models
|
RUN mkdir -p models
|
||||||
|
|
||||||
|
# Expose API port
|
||||||
EXPOSE 8000
|
EXPOSE 8000
|
||||||
|
|
||||||
|
# Environment variables
|
||||||
ENV PYTHONUNBUFFERED=1
|
ENV PYTHONUNBUFFERED=1
|
||||||
ENV PYTHONPATH=/app
|
ENV PYTHONPATH=/app
|
||||||
|
|
||||||
CMD ["python", "src/run.py"]
|
# Default command
|
||||||
|
CMD ["python3", "-m", "uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"]
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,193 @@
|
||||||
|
# Bone Quality Assessment Project
|
||||||
|
|
||||||
|
## 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.
|
||||||
|
|
||||||
|
### Core Purpose (Hackathon)
|
||||||
|
- Analyze DICOM densitometry studies
|
||||||
|
- Determine anatomical region (spine/hip)
|
||||||
|
- Binary classification: quality (OK/violation)
|
||||||
|
- Output results in XLSX/CSV format per requirements
|
||||||
|
|
||||||
|
### Tech Stack
|
||||||
|
| Component | Technology |
|
||||||
|
|-----------|------------|
|
||||||
|
| Backend | Python 3.10, FastAPI, Uvicorn |
|
||||||
|
| ML/Deep Learning | PyTorch, torchvision (ResNet18) |
|
||||||
|
| Image Processing | PIL, OpenCV, pydicom |
|
||||||
|
| Data Handling | pandas, openpyxl |
|
||||||
|
| Containerization | Docker |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Project Structure
|
||||||
|
|
||||||
|
```
|
||||||
|
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
|
||||||
|
│ ├── core/ # Orchestrator
|
||||||
|
│ ├── quality/ # Quality scoring
|
||||||
|
│ ├── segmentators/ # Segmentation models
|
||||||
|
│ └── classifiers/ # Classification models
|
||||||
|
├── models/
|
||||||
|
│ └── dxa_model.pth # Trained DXA classifier
|
||||||
|
├── dataset_hack/ # DICOM datasets
|
||||||
|
│ ├── Для теста/ # Test data (3 files)
|
||||||
|
│ └── НД_для_обучения/ # Training data (100 studies, 499 DICOMs)
|
||||||
|
│ └── разметка.xlsx # Annotation file
|
||||||
|
├── requirements.txt
|
||||||
|
├── Dockerfile
|
||||||
|
├── run.sh # Main entry script
|
||||||
|
└── README.md
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## DXA Module (`src/dxa/`)
|
||||||
|
|
||||||
|
### Dataset (`dataset.py`)
|
||||||
|
- Loads DICOM files from studies
|
||||||
|
- Parses annotation Excel file
|
||||||
|
- Maps anatomical regions: spine, hip_right, hip_left
|
||||||
|
- Quality labels: 0 (OK), 1 (violation)
|
||||||
|
- Total: ~1433 samples (train: 1146, val: 287)
|
||||||
|
|
||||||
|
### Model (`model.py`)
|
||||||
|
- Architecture: ResNet18 (pretrained on ImageNet)
|
||||||
|
- Task: Binary classification (quality OK vs violation)
|
||||||
|
- Input: 224x224 RGB images
|
||||||
|
- Output: class probabilities
|
||||||
|
|
||||||
|
### Training (`train.py`)
|
||||||
|
```bash
|
||||||
|
python src/dxa/train.py --epochs 10 --batch-size 16
|
||||||
|
```
|
||||||
|
|
||||||
|
### Inference (`inference.py`)
|
||||||
|
```bash
|
||||||
|
python src/dxa/inference.py \
|
||||||
|
--input-path dataset_hack/Для\ теста \
|
||||||
|
--output-path results.xlsx \
|
||||||
|
--model-path models/dxa_model.pth
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Running the Project
|
||||||
|
|
||||||
|
### Training
|
||||||
|
```bash
|
||||||
|
# Option 1: Direct Python
|
||||||
|
python src/dxa/train.py --epochs 10
|
||||||
|
|
||||||
|
# Option 2: Via run.sh
|
||||||
|
bash run.sh train
|
||||||
|
```
|
||||||
|
|
||||||
|
### Inference
|
||||||
|
```bash
|
||||||
|
# Single file
|
||||||
|
python src/dxa/inference.py --input-path file.dcm --output-path result.xlsx
|
||||||
|
|
||||||
|
# Directory (batch)
|
||||||
|
python src/dxa/inference.py --input-path dataset_hack/Для\ теста --output-path results.xlsx
|
||||||
|
```
|
||||||
|
|
||||||
|
### API Server
|
||||||
|
```bash
|
||||||
|
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Output Format (per Hackathon Requirements)
|
||||||
|
|
||||||
|
| Column | Description |
|
||||||
|
|--------|-------------|
|
||||||
|
| path_to_study | Path to study directory |
|
||||||
|
| study_uid | StudyInstanceUID from DICOM |
|
||||||
|
| image_uid | SOPInstanceUID from DICOM |
|
||||||
|
| anatomical_region | spine / hip |
|
||||||
|
| quality_class | 0 (OK), 1 (violation) |
|
||||||
|
| violation_type | Type of violation (if any) |
|
||||||
|
| processing_status | Success / Failure |
|
||||||
|
| time_of_processing | Processing time (seconds) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Model Performance
|
||||||
|
|
||||||
|
```
|
||||||
|
Training data: 1146 samples
|
||||||
|
Validation data: 287 samples
|
||||||
|
|
||||||
|
Training (10 epochs):
|
||||||
|
- Best F1: ~0.27 (imbalanced classes: ~65% OK, ~35% violation)
|
||||||
|
- Accuracy: ~84%
|
||||||
|
|
||||||
|
Note: Need more epochs (50+) and class balancing for production
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Annotation Format
|
||||||
|
|
||||||
|
The annotation Excel (`разметка.xlsx`) contains:
|
||||||
|
- Study UID
|
||||||
|
- Spine columns: укладка, ось, артефакты
|
||||||
|
- Hip columns: позиция, ROI (left/right)
|
||||||
|
- Total columns: итого
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Development Conventions
|
||||||
|
|
||||||
|
### 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
|
||||||
|
|
||||||
|
### 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
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Docker
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Build
|
||||||
|
docker build -t dxa-quality .
|
||||||
|
|
||||||
|
# Run
|
||||||
|
docker run -v /data:/data -p 8000:8000 dxa-quality
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Notes
|
||||||
|
|
||||||
|
- This is a **hackathon project** for DXA quality assessment
|
||||||
|
- Model trained on limited data (100 studies)
|
||||||
|
- Binary classification (quality OK / violation)
|
||||||
|
- Anatomical region detection via image size heuristic
|
||||||
|
- Output format matches hackathon requirements (XLSX/CSV)
|
||||||
38
README.md
38
README.md
|
|
@ -26,6 +26,18 @@
|
||||||
|
|
||||||

|

|
||||||
|
|
||||||
|
## 🎯 Доступные режимы
|
||||||
|
|
||||||
|
### DXA Режим (хакатон)
|
||||||
|
Анализ денситометрических исследований:
|
||||||
|
```bash
|
||||||
|
# Обучение
|
||||||
|
python src/dxa/train.py --epochs 10
|
||||||
|
|
||||||
|
# Инференс
|
||||||
|
python src/dxa/inference.py --input-path dataset_hack/Для\ теста --output-path results.xlsx
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
## 🚀 Быстрый старт
|
## 🚀 Быстрый старт
|
||||||
|
|
||||||
|
|
@ -52,7 +64,7 @@ venv\Scripts\activate # Windows
|
||||||
pip install -r requirements.txt
|
pip install -r requirements.txt
|
||||||
|
|
||||||
# 4. Загрузка обученной модели (опционально)
|
# 4. Загрузка обученной модели (опционально)
|
||||||
# Поместите модель в папку models/unet_cats_dogs.pth
|
# Поместите модель в папку models/dxa_model.pth
|
||||||
|
|
||||||
# 5. Запуск сервера
|
# 5. Запуск сервера
|
||||||
python run.py
|
python run.py
|
||||||
|
|
@ -145,22 +157,22 @@ curl -X POST "http://localhost:8000/api/v1/analyze" \
|
||||||
|
|
||||||
📁 Структура проекта
|
📁 Структура проекта
|
||||||
```text
|
```text
|
||||||
bone-quality-assessment/
|
bone_2026/
|
||||||
├── src/
|
├── src/
|
||||||
│ ├── api/
|
│ ├── dxa/ # DXA Quality модуль
|
||||||
│ │ ├── endpoints.py # FastAPI эндпоинты
|
│ │ ├── model.py # ResNet18 классификатор
|
||||||
│ │ └── static/
|
│ │ ├── dataset.py # Загрузчик данных
|
||||||
│ │ └── index.html # Web-интерфейс
|
│ │ ├── train.py # Обучение
|
||||||
│ ├── models/
|
│ │ └── inference.py # Инференс
|
||||||
│ │ └── unet.py # U-Net архитектура
|
│ ├── api/ # REST API
|
||||||
│ └── quality/
|
│ ├── quality/ # Оценка качества
|
||||||
│ └── quality_scorer.py # Оценка качества
|
│ └── main.py # FastAPI приложение
|
||||||
├── models/
|
├── models/
|
||||||
│ └── unet_cats_dogs.pth # Обученная модель
|
│ └── dxa_model.pth # Обученная модель
|
||||||
|
├── dataset_hack/ # DICOM датасет
|
||||||
├── Dockerfile
|
├── Dockerfile
|
||||||
├── docker-compose.yml
|
|
||||||
├── requirements.txt
|
├── requirements.txt
|
||||||
├── run.py
|
├── run.sh
|
||||||
└── README.md
|
└── README.md
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
|
||||||
140
inference.py
140
inference.py
|
|
@ -1,140 +0,0 @@
|
||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from PIL import Image
|
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
import argparse
|
|
||||||
|
|
||||||
from src import UNet
|
|
||||||
from src.quality.quality_scorer import QualityScorer
|
|
||||||
|
|
||||||
|
|
||||||
# from src.quality.quality_scorer_old import QualityScorer
|
|
||||||
|
|
||||||
|
|
||||||
def load_model(model_path, in_channels=3, out_classes=2, device='cpu'):
|
|
||||||
"""Загрузка модели"""
|
|
||||||
model = UNet(in_channels=in_channels, out_classes=out_classes)
|
|
||||||
model.load_state_dict(torch.load(model_path, map_location=device))
|
|
||||||
model.to(device)
|
|
||||||
model.eval()
|
|
||||||
return model
|
|
||||||
|
|
||||||
|
|
||||||
def predict_single(model, image_path, device='cpu', input_size=(256, 256)):
|
|
||||||
"""Предсказание для одного изображения"""
|
|
||||||
# Загрузка изображения
|
|
||||||
image = Image.open(image_path).convert('RGB')
|
|
||||||
original_image = np.array(image)
|
|
||||||
original_size = image.size # (width, height)
|
|
||||||
|
|
||||||
# Подготовка для модели (ресайз)
|
|
||||||
image_resized = image.resize(input_size)
|
|
||||||
image_array = np.array(image_resized, dtype=np.float32) / 255.0
|
|
||||||
image_array = image_array.transpose(2, 0, 1)
|
|
||||||
image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(device)
|
|
||||||
|
|
||||||
# Инференс
|
|
||||||
with torch.no_grad():
|
|
||||||
output = model(image_tensor)
|
|
||||||
pred = torch.softmax(output, dim=1)
|
|
||||||
mask_resized = pred.argmax(dim=1).squeeze().cpu().numpy()
|
|
||||||
|
|
||||||
# Ресайз маски обратно к оригинальному размеру
|
|
||||||
# Используем NEAREST интерполяцию для сохранения бинарности
|
|
||||||
mask_pil = Image.fromarray(mask_resized.astype(np.uint8))
|
|
||||||
mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST)
|
|
||||||
mask = np.array(mask_pil)
|
|
||||||
|
|
||||||
return original_image, mask
|
|
||||||
|
|
||||||
|
|
||||||
def visualize_prediction(image, mask, quality_scores, save_path=None):
|
|
||||||
"""Визуализация результатов"""
|
|
||||||
fig, axes = plt.subplots(1, 3, figsize=(15, 5))
|
|
||||||
|
|
||||||
# 1. Оригинальное изображение
|
|
||||||
axes[0].imshow(image)
|
|
||||||
axes[0].set_title('Оригинал')
|
|
||||||
axes[0].axis('off')
|
|
||||||
|
|
||||||
# 2. Сегментация (маска наложена на оригинал)
|
|
||||||
axes[1].imshow(image)
|
|
||||||
axes[1].imshow(mask, cmap='jet', alpha=0.5)
|
|
||||||
axes[1].set_title('Сегментация')
|
|
||||||
axes[1].axis('off')
|
|
||||||
|
|
||||||
# 3. Оценка качества
|
|
||||||
quality_text = f"Качество: {quality_scores['overall_quality']}\n"
|
|
||||||
quality_text += f"Уверенность: {quality_scores['confidence']:.2f}\n"
|
|
||||||
quality_text += f"Нарушений: {len(quality_scores['issues'])}\n\n"
|
|
||||||
|
|
||||||
# Добавляем детали
|
|
||||||
if quality_scores['issues']:
|
|
||||||
quality_text += "Проблемы:\n"
|
|
||||||
for issue in quality_scores['issues']:
|
|
||||||
quality_text += f" • {issue['type']}: {issue['details'][:30]}...\n"
|
|
||||||
else:
|
|
||||||
quality_text += "✅ Нарушений не обнаружено"
|
|
||||||
|
|
||||||
axes[2].text(0.05, 0.95, quality_text, fontsize=10,
|
|
||||||
verticalalignment='top', transform=axes[2].transAxes,
|
|
||||||
bbox=dict(boxstyle="round", facecolor="white", alpha=0.8))
|
|
||||||
axes[2].set_title('Оценка качества')
|
|
||||||
axes[2].axis('off')
|
|
||||||
plt.rcParams['font.family'] = 'sans-serif'
|
|
||||||
plt.rcParams['font.sans-serif'] = ['Apple Color Emoji', 'Segoe UI Emoji', 'Noto Color Emoji', 'DejaVu Sans']
|
|
||||||
plt.tight_layout()
|
|
||||||
|
|
||||||
if save_path:
|
|
||||||
plt.savefig(save_path, dpi=150, bbox_inches='tight')
|
|
||||||
print(f"✅ Визуализация сохранена: {save_path}")
|
|
||||||
|
|
||||||
plt.show()
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
parser = argparse.ArgumentParser(description='Инференс модели сегментации')
|
|
||||||
parser.add_argument('--image', type=str, default="example.jpg", required=False, help='Путь к изображению')
|
|
||||||
parser.add_argument('--model', type=str, default='models/unet_cats_dogs.pth',
|
|
||||||
help='Путь к модели')
|
|
||||||
parser.add_argument('--device', type=str, default='cpu',
|
|
||||||
choices=['cpu', 'mps', 'cuda'], help='Устройство')
|
|
||||||
parser.add_argument('--save', type=str, default=None, help='Сохранить результат')
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
# Проверка устройства
|
|
||||||
if args.device == 'mps' and not torch.backends.mps.is_available():
|
|
||||||
print("⚠️ MPS не доступен, используем CPU")
|
|
||||||
args.device = 'cpu'
|
|
||||||
|
|
||||||
print(f"🚀 Загрузка модели с {args.device}...")
|
|
||||||
model = load_model(args.model, device=args.device)
|
|
||||||
|
|
||||||
print(f"📸 Анализ изображения: {args.image}")
|
|
||||||
image, mask = predict_single(model, args.image, device=args.device)
|
|
||||||
|
|
||||||
print(f" Изображение: {image.shape}")
|
|
||||||
print(f" Маска: {mask.shape}")
|
|
||||||
|
|
||||||
print("🔍 Оценка качества...")
|
|
||||||
scorer = QualityScorer()
|
|
||||||
quality_scores = scorer.evaluate(image=image, segmentation=mask)
|
|
||||||
|
|
||||||
print(f"\n📊 Результаты:")
|
|
||||||
print(f" Качество: {quality_scores['overall_quality']}")
|
|
||||||
print(f" Уверенность: {quality_scores['confidence']:.2f}")
|
|
||||||
print(f" Количество проблем: {len(quality_scores['issues'])}")
|
|
||||||
|
|
||||||
if quality_scores['issues']:
|
|
||||||
print("\n Проблемы:")
|
|
||||||
for issue in quality_scores['issues']:
|
|
||||||
print(f" • {issue['type']}: {issue['details'][:50]}...")
|
|
||||||
|
|
||||||
# Визуализация
|
|
||||||
save_path = args.save or 'prediction_result.png'
|
|
||||||
visualize_prediction(image, mask, quality_scores, save_path)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
|
|
@ -1,47 +1,37 @@
|
||||||
annotated-doc==0.0.5
|
# Core dependencies
|
||||||
annotated-types==0.7.0
|
torch>=2.0.0
|
||||||
anyio==4.12.1
|
torchvision>=0.15.0
|
||||||
click==8.1.8
|
numpy>=1.24.0
|
||||||
contourpy==1.3.0
|
|
||||||
cycler==0.12.1
|
# Image processing
|
||||||
exceptiongroup==1.3.1
|
Pillow>=10.0.0
|
||||||
fastapi==0.128.8
|
opencv-python-headless>=4.8.0
|
||||||
filelock==3.19.1
|
|
||||||
fonttools==4.60.2
|
# Medical imaging
|
||||||
fsspec==2025.10.0
|
pydicom>=2.4.0
|
||||||
h11==0.16.0
|
pydicom-seg>=0.4.0
|
||||||
idna==3.18
|
nibabel>=5.0.0
|
||||||
importlib_resources==6.5.2
|
SimpleITK>=2.2.0
|
||||||
Jinja2==3.1.6
|
|
||||||
kiwisolver==1.4.7
|
# Deep learning / models
|
||||||
MarkupSafe==3.0.3
|
torchvision>=0.15.0
|
||||||
matplotlib==3.9.4
|
timm>=0.9.0
|
||||||
monai==1.5.2
|
|
||||||
mpmath==1.3.0
|
# Data handling
|
||||||
networkx==3.2.1
|
pandas>=2.0.0
|
||||||
nibabel==5.3.3
|
openpyxl>=3.1.0
|
||||||
numpy==2.0.2
|
|
||||||
opencv-python-headless==5.0.0.93
|
# Visualization
|
||||||
packaging==26.3
|
matplotlib>=3.7.0
|
||||||
pandas==2.3.3
|
|
||||||
pillow==11.3.0
|
# API / serving
|
||||||
pydantic==2.13.4
|
fastapi>=0.100.0
|
||||||
pydantic_core==2.46.4
|
uvicorn>=0.23.0
|
||||||
pydicom==2.4.4
|
python-multipart>=0.0.6
|
||||||
pyparsing==3.3.2
|
|
||||||
python-dateutil==2.9.0.post0
|
# Progress bars
|
||||||
python-multipart==0.0.20
|
tqdm>=4.65.0
|
||||||
pytz==2026.3.post1
|
|
||||||
scipy==1.13.1
|
# Utils
|
||||||
six==1.17.0
|
scipy>=1.10.0
|
||||||
starlette==0.49.3
|
scikit-learn>=1.3.0
|
||||||
sympy==1.14.0
|
|
||||||
torch==2.8.0
|
|
||||||
torchvision==0.23.0
|
|
||||||
TotalSegmentator==2.18.0
|
|
||||||
tqdm==4.70.0
|
|
||||||
typing-inspection==0.4.2
|
|
||||||
typing_extensions==4.16.0
|
|
||||||
tzdata==2026.3
|
|
||||||
uvicorn==0.39.0
|
|
||||||
zipp==3.23.1
|
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,134 @@
|
||||||
|
#!/bin/bash
|
||||||
|
# DXA Quality Assessment - Main entry point
|
||||||
|
# Usage: bash run.sh [command] [args...]
|
||||||
|
|
||||||
|
set -e
|
||||||
|
|
||||||
|
# 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}
|
||||||
|
DATA_ROOT=${DATA_ROOT:-dataset_hack}
|
||||||
|
ANNOTATION_PATH=${ANNOTATION_PATH:-dataset_hack/НД_для_обучения/разметка.xlsx}
|
||||||
|
MODEL_PATH=${MODEL_PATH:-models/dxa_model.pth}
|
||||||
|
EPOCHS=${EPOCHS:-10}
|
||||||
|
BATCH_SIZE=${BATCH_SIZE:-16}
|
||||||
|
|
||||||
|
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}"
|
||||||
|
;;
|
||||||
|
|
||||||
|
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}"
|
||||||
|
;;
|
||||||
|
|
||||||
|
serve)
|
||||||
|
echo -e "${YELLOW}Starting API server...${NC}"
|
||||||
|
python3 -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
||||||
|
;;
|
||||||
|
|
||||||
|
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"
|
||||||
|
;;
|
||||||
|
esac
|
||||||
File diff suppressed because it is too large
Load Diff
|
|
@ -0,0 +1,356 @@
|
||||||
|
// DXA Quality Assessment Web Interface
|
||||||
|
|
||||||
|
const API_BASE = '';
|
||||||
|
let results = [];
|
||||||
|
let sortColumn = null;
|
||||||
|
let sortDirection = 'asc';
|
||||||
|
|
||||||
|
// DOM Elements
|
||||||
|
const dropZone = document.getElementById('dropZone');
|
||||||
|
const fileInput = document.getElementById('fileInput');
|
||||||
|
const progressSection = document.getElementById('progressSection');
|
||||||
|
const progressBar = document.getElementById('progressBar');
|
||||||
|
const progressText = document.getElementById('progressText');
|
||||||
|
const progressPercent = document.getElementById('progressPercent');
|
||||||
|
const resultsSection = document.getElementById('resultsSection');
|
||||||
|
const resultsTable = document.getElementById('resultsTable');
|
||||||
|
const emptyResults = document.getElementById('emptyResults');
|
||||||
|
const errorSection = document.getElementById('errorSection');
|
||||||
|
const errorMessage = document.getElementById('errorMessage');
|
||||||
|
const statusBanner = document.getElementById('statusBanner');
|
||||||
|
const statusIcon = document.getElementById('statusIcon');
|
||||||
|
const statusText = document.getElementById('statusText');
|
||||||
|
|
||||||
|
// Theme toggle
|
||||||
|
const themeToggle = document.getElementById('themeToggle');
|
||||||
|
if (themeToggle) {
|
||||||
|
themeToggle.addEventListener('click', () => {
|
||||||
|
document.documentElement.classList.toggle('dark');
|
||||||
|
localStorage.setItem('theme', document.documentElement.classList.contains('dark') ? 'dark' : 'light');
|
||||||
|
});
|
||||||
|
|
||||||
|
// Load saved theme
|
||||||
|
if (localStorage.getItem('theme') === 'dark' || (!localStorage.getItem('theme') && window.matchMedia('(prefers-color-scheme: dark)').matches)) {
|
||||||
|
document.documentElement.classList.add('dark');
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check API health
|
||||||
|
async function checkHealth() {
|
||||||
|
try {
|
||||||
|
const response = await fetch(`${API_BASE}/api/v1/health`);
|
||||||
|
const data = await response.json();
|
||||||
|
|
||||||
|
if (data.status === 'ok') {
|
||||||
|
statusIcon.className = 'w-3 h-3 rounded-full bg-green-500';
|
||||||
|
statusText.textContent = `Сервис готов • Модель: ${data.model_loaded ? 'загружена' : 'не загружена'} • Устройство: ${data.device}`;
|
||||||
|
} else {
|
||||||
|
statusIcon.className = 'w-3 h-3 rounded-full bg-yellow-500';
|
||||||
|
statusText.textContent = 'Проблемы с сервисом';
|
||||||
|
}
|
||||||
|
} catch (e) {
|
||||||
|
statusIcon.className = 'w-3 h-3 rounded-full bg-red-500';
|
||||||
|
statusText.textContent = 'Сервис недоступен';
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Drag and drop
|
||||||
|
dropZone.addEventListener('click', () => fileInput.click());
|
||||||
|
|
||||||
|
dropZone.addEventListener('dragover', (e) => {
|
||||||
|
e.preventDefault();
|
||||||
|
dropZone.classList.add('drag-over');
|
||||||
|
});
|
||||||
|
|
||||||
|
dropZone.addEventListener('dragleave', () => {
|
||||||
|
dropZone.classList.remove('drag-over');
|
||||||
|
});
|
||||||
|
|
||||||
|
dropZone.addEventListener('drop', (e) => {
|
||||||
|
e.preventDefault();
|
||||||
|
dropZone.classList.remove('drag-over');
|
||||||
|
handleFiles(e.dataTransfer.files);
|
||||||
|
});
|
||||||
|
|
||||||
|
fileInput.addEventListener('change', (e) => {
|
||||||
|
handleFiles(e.target.files);
|
||||||
|
});
|
||||||
|
|
||||||
|
// Filter inputs
|
||||||
|
const searchInput = document.getElementById('searchInput');
|
||||||
|
const filterRegion = document.getElementById('filterRegion');
|
||||||
|
const filterQuality = document.getElementById('filterQuality');
|
||||||
|
|
||||||
|
searchInput.addEventListener('input', renderTable);
|
||||||
|
filterRegion.addEventListener('change', renderTable);
|
||||||
|
filterQuality.addEventListener('change', renderTable);
|
||||||
|
|
||||||
|
// Sort headers
|
||||||
|
document.querySelectorAll('th[data-sort]').forEach(th => {
|
||||||
|
th.addEventListener('click', () => {
|
||||||
|
const column = th.dataset.sort;
|
||||||
|
if (sortColumn === column) {
|
||||||
|
sortDirection = sortDirection === 'asc' ? 'desc' : 'asc';
|
||||||
|
} else {
|
||||||
|
sortColumn = column;
|
||||||
|
sortDirection = 'asc';
|
||||||
|
}
|
||||||
|
renderTable();
|
||||||
|
updateSortIcons();
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
function updateSortIcons() {
|
||||||
|
document.querySelectorAll('th[data-sort]').forEach(th => {
|
||||||
|
const icon = th.querySelector('i');
|
||||||
|
if (th.dataset.sort === sortColumn) {
|
||||||
|
icon.className = `fas fa-sort-${sortDirection === 'asc' ? 'up' : 'down'} ml-1`;
|
||||||
|
} else {
|
||||||
|
icon.className = 'fas fa-sort ml-1';
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle files
|
||||||
|
async function handleFiles(files) {
|
||||||
|
if (files.length === 0) return;
|
||||||
|
|
||||||
|
results = [];
|
||||||
|
errorSection.classList.add('hidden');
|
||||||
|
resultsSection.classList.remove('hidden');
|
||||||
|
progressSection.classList.remove('hidden');
|
||||||
|
|
||||||
|
const dcmFiles = Array.from(files).filter(f =>
|
||||||
|
f.name.toLowerCase().endsWith('.dcm') ||
|
||||||
|
f.type === 'application/dicom' ||
|
||||||
|
f.name.toLowerCase().includes('dcm')
|
||||||
|
);
|
||||||
|
|
||||||
|
if (dcmFiles.length === 0) {
|
||||||
|
showError('DICOM файлы не найдены. Пожалуйста, выберите файлы с расширением .dcm');
|
||||||
|
progressSection.classList.add('hidden');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
progressText.textContent = `Анализ ${dcmFiles.length} файлов...`;
|
||||||
|
progressBar.style.width = '0%';
|
||||||
|
progressPercent.textContent = '0%';
|
||||||
|
|
||||||
|
let processed = 0;
|
||||||
|
|
||||||
|
for (const file of dcmFiles) {
|
||||||
|
try {
|
||||||
|
const formData = new FormData();
|
||||||
|
formData.append('file', file);
|
||||||
|
|
||||||
|
const response = await fetch(`${API_BASE}/api/v1/analyze`, {
|
||||||
|
method: 'POST',
|
||||||
|
body: formData
|
||||||
|
});
|
||||||
|
|
||||||
|
if (!response.ok) {
|
||||||
|
throw new Error(`HTTP ${response.status}`);
|
||||||
|
}
|
||||||
|
|
||||||
|
const data = await response.json();
|
||||||
|
results.push({
|
||||||
|
filename: file.name,
|
||||||
|
...data
|
||||||
|
});
|
||||||
|
} catch (e) {
|
||||||
|
results.push({
|
||||||
|
filename: file.name,
|
||||||
|
quality_class: -1,
|
||||||
|
quality_label: 'Error',
|
||||||
|
processing_status: `Failure: ${e.message}`,
|
||||||
|
confidence: 0
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
processed++;
|
||||||
|
const percent = Math.round((processed / dcmFiles.length) * 100);
|
||||||
|
progressBar.style.width = `${percent}%`;
|
||||||
|
progressPercent.textContent = `${percent}%`;
|
||||||
|
}
|
||||||
|
|
||||||
|
progressSection.classList.add('hidden');
|
||||||
|
updateStats();
|
||||||
|
renderTable();
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update stats
|
||||||
|
function updateStats() {
|
||||||
|
const total = results.length;
|
||||||
|
const ok = results.filter(r => r.quality_class === 0).length;
|
||||||
|
const violation = results.filter(r => r.quality_class === 1).length;
|
||||||
|
const avgConfidence = results.length > 0
|
||||||
|
? (results.reduce((sum, r) => sum + (r.confidence || 0), 0) / results.length * 100).toFixed(1)
|
||||||
|
: 0;
|
||||||
|
|
||||||
|
document.getElementById('totalCount').textContent = total;
|
||||||
|
document.getElementById('okCount').textContent = ok;
|
||||||
|
document.getElementById('violationCount').textContent = violation;
|
||||||
|
document.getElementById('accuracy').textContent = `${avgConfidence}%`;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Filter and sort results
|
||||||
|
function getFilteredResults() {
|
||||||
|
let filtered = [...results];
|
||||||
|
|
||||||
|
// Search filter
|
||||||
|
const search = searchInput.value.toLowerCase();
|
||||||
|
if (search) {
|
||||||
|
filtered = filtered.filter(r =>
|
||||||
|
(r.filename || '').toLowerCase().includes(search) ||
|
||||||
|
(r.study_uid || '').toLowerCase().includes(search) ||
|
||||||
|
(r.image_uid || '').toLowerCase().includes(search)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Region filter
|
||||||
|
const region = filterRegion.value;
|
||||||
|
if (region) {
|
||||||
|
filtered = filtered.filter(r => r.anatomical_region === region);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Quality filter
|
||||||
|
const quality = filterQuality.value;
|
||||||
|
if (quality !== '') {
|
||||||
|
filtered = filtered.filter(r => String(r.quality_class) === quality);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sort
|
||||||
|
if (sortColumn) {
|
||||||
|
filtered.sort((a, b) => {
|
||||||
|
let aVal = a[sortColumn] || '';
|
||||||
|
let bVal = b[sortColumn] || '';
|
||||||
|
|
||||||
|
if (typeof aVal === 'number' && typeof bVal === 'number') {
|
||||||
|
return sortDirection === 'asc' ? aVal - bVal : bVal - aVal;
|
||||||
|
}
|
||||||
|
|
||||||
|
aVal = String(aVal).toLowerCase();
|
||||||
|
bVal = String(bVal).toLowerCase();
|
||||||
|
|
||||||
|
if (sortDirection === 'asc') {
|
||||||
|
return aVal.localeCompare(bVal);
|
||||||
|
}
|
||||||
|
return bVal.localeCompare(aVal);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
return filtered;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Render table
|
||||||
|
function renderTable() {
|
||||||
|
const filtered = getFilteredResults();
|
||||||
|
|
||||||
|
if (filtered.length === 0) {
|
||||||
|
resultsTable.innerHTML = '';
|
||||||
|
emptyResults.classList.remove('hidden');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
emptyResults.classList.add('hidden');
|
||||||
|
|
||||||
|
resultsTable.innerHTML = filtered.map(r => {
|
||||||
|
const qualityClass = r.quality_class === 0 ? 'bg-green-100 text-green-800 dark:bg-green-900/30 dark:text-green-400' :
|
||||||
|
r.quality_class === 1 ? 'bg-red-100 text-red-800 dark:bg-red-900/30 dark:text-red-400' :
|
||||||
|
'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-300';
|
||||||
|
|
||||||
|
const qualityLabel = r.quality_class === 0 ? 'OK' :
|
||||||
|
r.quality_class === 1 ? 'Нарушение' :
|
||||||
|
'Ошибка';
|
||||||
|
|
||||||
|
const regionLabel = r.anatomical_region === 'spine' ? 'Позвоночник' :
|
||||||
|
r.anatomical_region === 'hip' ? 'Бедро' :
|
||||||
|
r.anatomical_region || '—';
|
||||||
|
|
||||||
|
const confidence = r.confidence ? (r.confidence * 100).toFixed(1) + '%' : '—';
|
||||||
|
|
||||||
|
return `
|
||||||
|
<tr class="hover:bg-gray-50 dark:hover:bg-gray-700/50 fade-in">
|
||||||
|
<td class="px-4 py-3 text-sm text-gray-900 dark:text-white">
|
||||||
|
<div class="flex items-center gap-2">
|
||||||
|
<i class="fas fa-file-medical text-gray-400"></i>
|
||||||
|
<span class="font-medium">${r.filename || '—'}</span>
|
||||||
|
</div>
|
||||||
|
${r.study_uid ? `<div class="text-xs text-gray-500 dark:text-gray-400 mt-1">UID: ${r.study_uid.substring(0, 20)}...</div>` : ''}
|
||||||
|
</td>
|
||||||
|
<td class="px-4 py-3 text-sm text-gray-600 dark:text-gray-300">
|
||||||
|
<span class="px-2 py-1 bg-gray-100 dark:bg-gray-700 rounded text-xs font-medium">
|
||||||
|
${regionLabel}
|
||||||
|
</span>
|
||||||
|
</td>
|
||||||
|
<td class="px-4 py-3">
|
||||||
|
<span class="px-2 py-1 rounded text-xs font-medium ${qualityClass}">
|
||||||
|
${qualityLabel}
|
||||||
|
</span>
|
||||||
|
</td>
|
||||||
|
<td class="px-4 py-3 text-sm text-gray-600 dark:text-gray-300">
|
||||||
|
<div class="flex items-center gap-2">
|
||||||
|
<div class="w-16 bg-gray-200 dark:bg-gray-700 rounded-full h-1.5">
|
||||||
|
<div class="bg-primary-600 h-1.5 rounded-full" style="width: ${r.confidence ? r.confidence * 100 : 0}%"></div>
|
||||||
|
</div>
|
||||||
|
<span class="text-xs">${confidence}</span>
|
||||||
|
</div>
|
||||||
|
</td>
|
||||||
|
</tr>
|
||||||
|
`;
|
||||||
|
}).join('');
|
||||||
|
}
|
||||||
|
|
||||||
|
// Export XLSX
|
||||||
|
document.getElementById('exportXlsx')?.addEventListener('click', async () => {
|
||||||
|
if (results.length === 0) return;
|
||||||
|
|
||||||
|
// Show loading state
|
||||||
|
const btn = document.getElementById('exportXlsx');
|
||||||
|
const originalText = btn.innerHTML;
|
||||||
|
btn.innerHTML = '<i class="fas fa-spinner fa-spin"></i> Экспорт...';
|
||||||
|
btn.disabled = true;
|
||||||
|
|
||||||
|
try {
|
||||||
|
// For proper XLSX export, we'd need to re-upload files
|
||||||
|
// For now, export as CSV which Excel can open
|
||||||
|
const headers = ['filename', 'study_uid', 'image_uid', 'anatomical_region', 'quality_class', 'quality_label', 'confidence', 'processing_status'];
|
||||||
|
const csvContent = [
|
||||||
|
headers.join(','),
|
||||||
|
...results.map(r => headers.map(h => {
|
||||||
|
let val = r[h] || '';
|
||||||
|
if (typeof val === 'string' && val.includes(',')) {
|
||||||
|
val = `"${val}"`;
|
||||||
|
}
|
||||||
|
return val;
|
||||||
|
}).join(','))
|
||||||
|
].join('\n');
|
||||||
|
|
||||||
|
const blob = new Blob([csvContent], { type: 'text/csv;charset=utf-8;' });
|
||||||
|
const link = document.createElement('a');
|
||||||
|
link.href = URL.createObjectURL(blob);
|
||||||
|
link.download = `dxa_results_${new Date().toISOString().slice(0, 10)}.csv`;
|
||||||
|
link.click();
|
||||||
|
} catch (e) {
|
||||||
|
showError('Ошибка экспорта: ' + e.message);
|
||||||
|
} finally {
|
||||||
|
btn.innerHTML = originalText;
|
||||||
|
btn.disabled = false;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// Clear results
|
||||||
|
document.getElementById('clearResults')?.addEventListener('click', () => {
|
||||||
|
results = [];
|
||||||
|
resultsSection.classList.add('hidden');
|
||||||
|
fileInput.value = '';
|
||||||
|
});
|
||||||
|
|
||||||
|
// Show error
|
||||||
|
function showError(message) {
|
||||||
|
errorMessage.textContent = message;
|
||||||
|
errorSection.classList.remove('hidden');
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize
|
||||||
|
checkHealth();
|
||||||
|
|
@ -1,174 +0,0 @@
|
||||||
# src/models/pet_breed_classifier.py
|
|
||||||
from typing import Dict, Any, Union, Optional
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torchvision.models as models
|
|
||||||
from PIL import Image
|
|
||||||
import numpy as np
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from src.classifiers.base import BaseClassifier
|
|
||||||
|
|
||||||
# Маппинг ID породы в название
|
|
||||||
BREED_MAP = {
|
|
||||||
1: "Abyssinian", 2: "Bengal", 3: "Birman", 4: "Bombay",
|
|
||||||
5: "British Shorthair", 6: "Egyptian Mau", 7: "Maine Coon",
|
|
||||||
8: "Persian", 9: "Ragdoll", 10: "Russian Blue", 11: "Siamese",
|
|
||||||
12: "Sphynx", 13: "American Bulldog", 14: "American Pit Bull Terrier",
|
|
||||||
15: "American Staffordshire Terrier", 16: "Australian Shepherd",
|
|
||||||
17: "Beagle", 18: "Border Collie", 19: "Boxer", 20: "Chihuahua",
|
|
||||||
21: "Cocker Spaniel", 22: "Dachshund", 23: "Doberman Pinscher",
|
|
||||||
24: "English Cocker Spaniel", 25: "German Shepherd", 26: "Golden Retriever",
|
|
||||||
27: "Great Dane", 28: "Jack Russell Terrier", 29: "Labrador Retriever",
|
|
||||||
30: "Poodle", 31: "Rottweiler", 32: "Siberian Husky",
|
|
||||||
33: "Staffordshire Bull Terrier", 34: "Yorkshire Terrier",
|
|
||||||
35: "Dalmatian", 36: "Pug", 37: "Shih Tzu"
|
|
||||||
}
|
|
||||||
|
|
||||||
SPECIES_MAP = {
|
|
||||||
1: "Cat",
|
|
||||||
2: "Dog"
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class PetBreedClassifier(BaseClassifier):
|
|
||||||
"""
|
|
||||||
Классификатор породы для кошек и собак
|
|
||||||
Использует ResNet50 с дообучением на 37 пород
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, device: str = 'cpu'):
|
|
||||||
"""
|
|
||||||
Args:
|
|
||||||
model_path: путь к обученной модели (опционально)
|
|
||||||
device: устройство для инференса ('cpu', 'mps', 'cuda')
|
|
||||||
"""
|
|
||||||
self.device = device
|
|
||||||
|
|
||||||
# Создаем модель ResNet50 с предобученными весами на ImageNet
|
|
||||||
self.model = models.resnet50(pretrained=True)
|
|
||||||
|
|
||||||
# Заменяем последний слой на 37 классов (породы)
|
|
||||||
num_features = self.model.fc.in_features
|
|
||||||
# self.model.fc = nn.Linear(num_features, 37)
|
|
||||||
model_path = "../models/unet_cats_dogs.pth"
|
|
||||||
# Загружаем обученные веса если есть
|
|
||||||
self.loaded = False
|
|
||||||
if model_path and Path(model_path).exists():
|
|
||||||
try:
|
|
||||||
self.model.load_state_dict(torch.load(model_path, map_location=device))
|
|
||||||
self.loaded = True
|
|
||||||
print(f"✅ Загружен классификатор из {model_path}")
|
|
||||||
except Exception as e:
|
|
||||||
print(f"⚠️ Ошибка загрузки модели из {model_path}: {e}")
|
|
||||||
print(" Используем неподготовленную модель (будет работать плохо)")
|
|
||||||
else:
|
|
||||||
print(f"⚠️ Классификатор не найден ({model_path}), используем неподготовленную модель ")
|
|
||||||
print(" Для точной работы обучите модель: python train_classifier.py")
|
|
||||||
|
|
||||||
# Перемещаем на устройство
|
|
||||||
self.model.to(device)
|
|
||||||
self.model.eval()
|
|
||||||
|
|
||||||
print(f"📊 Модель на устройстве: {device}")
|
|
||||||
|
|
||||||
def predict(self, image: Union[np.ndarray, Image.Image]) -> Dict[str, Any]:
|
|
||||||
"""
|
|
||||||
Предсказание породы по изображению
|
|
||||||
|
|
||||||
Args:
|
|
||||||
image: изображение в формате numpy array или PIL Image
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Словарь с результатами:
|
|
||||||
- breed_id: ID породы (1-37)
|
|
||||||
- breed_name: название породы
|
|
||||||
- species_id: ID вида (1-кошка, 2-собака)
|
|
||||||
- species_name: название вида
|
|
||||||
- confidence: уверенность (0-1)
|
|
||||||
- loaded: загружена ли модель
|
|
||||||
"""
|
|
||||||
# Проверка загрузки модели
|
|
||||||
if not self.loaded:
|
|
||||||
return {
|
|
||||||
"breed_id": 0,
|
|
||||||
"breed_name": "Unknown",
|
|
||||||
"species_id": 0,
|
|
||||||
"species_name": "Unknown",
|
|
||||||
"confidence": 0.0,
|
|
||||||
"loaded": False,
|
|
||||||
"error": "Model not loaded"
|
|
||||||
}
|
|
||||||
|
|
||||||
try:
|
|
||||||
# Подготовка изображения
|
|
||||||
if isinstance(image, np.ndarray):
|
|
||||||
image = Image.fromarray(image)
|
|
||||||
|
|
||||||
# Ресайз до 224x224 (стандартный вход ResNet)
|
|
||||||
image = image.resize((224, 224))
|
|
||||||
image_array = np.array(image, dtype=np.float32) / 255.0
|
|
||||||
image_array = image_array.transpose(2, 0, 1) # HWC -> CHW
|
|
||||||
|
|
||||||
# Нормализация для ImageNet (как у ResNet)
|
|
||||||
mean = np.array([0.485, 0.456, 0.406])
|
|
||||||
std = np.array([0.229, 0.224, 0.225])
|
|
||||||
for i in range(3):
|
|
||||||
image_array[i] = (image_array[i] - mean[i]) / std[i]
|
|
||||||
|
|
||||||
# Создаем тензор
|
|
||||||
image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(self.device)
|
|
||||||
|
|
||||||
# Инференс
|
|
||||||
with torch.no_grad():
|
|
||||||
outputs = self.model(image_tensor)
|
|
||||||
probabilities = torch.softmax(outputs, dim=1)
|
|
||||||
confidence, predicted = torch.max(probabilities, 1)
|
|
||||||
|
|
||||||
# Формируем результат
|
|
||||||
breed_id = predicted.item() + 1 # ID начинаются с 1
|
|
||||||
confidence_score = confidence.item()
|
|
||||||
|
|
||||||
# Определяем вид (кошка или собака)
|
|
||||||
# 1-25 кошки, 26-37 собаки
|
|
||||||
species_id = 1 if breed_id <= 25 else 2
|
|
||||||
|
|
||||||
return {
|
|
||||||
"breed_id": breed_id,
|
|
||||||
# "breed_name": BREED_MAP.get(breed_id, "Unknown"),
|
|
||||||
"species_id": species_id,
|
|
||||||
# "species_name": SPECIES_MAP.get(species_id, "Unknown"),
|
|
||||||
"confidence": confidence_score,
|
|
||||||
"loaded": True,
|
|
||||||
"error": None
|
|
||||||
}
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
return {
|
|
||||||
"breed_id": 0,
|
|
||||||
"breed_name": "Error",
|
|
||||||
"species_id": 0,
|
|
||||||
"species_name": "Error",
|
|
||||||
"confidence": 0.0,
|
|
||||||
"loaded": self.loaded,
|
|
||||||
"error": str(e)
|
|
||||||
}
|
|
||||||
|
|
||||||
def classify(self, image: Union[np.ndarray, Image.Image]) -> Dict[str, Any]:
|
|
||||||
"""Алиас для predict (для совместимости с BaseClassifier)"""
|
|
||||||
return self.predict(image)
|
|
||||||
|
|
||||||
def get_info(self) -> Dict[str, Any]:
|
|
||||||
"""Информация о классификаторе"""
|
|
||||||
return {
|
|
||||||
"name": self.__class__.__name__,
|
|
||||||
"type": "classification",
|
|
||||||
"target": "pet_breeds",
|
|
||||||
"num_classes": 37,
|
|
||||||
"loaded": self.loaded,
|
|
||||||
"device": self.device,
|
|
||||||
"model_architecture": "ResNet50",
|
|
||||||
"pretrained": True,
|
|
||||||
"breeds": list(BREED_MAP.values())[:5] + ["..."], # Показываем 5 пород
|
|
||||||
"total_breeds": len(BREED_MAP)
|
|
||||||
}
|
|
||||||
|
|
@ -1,9 +1,6 @@
|
||||||
class Config:
|
class Config:
|
||||||
# На хакатоне просто меняешь эти пути!
|
# DXA Quality Assessment Configuration
|
||||||
SEGMENTATION_MODEL = "unet_cats_dogs.pth" # → "totalsegmentator"
|
SEGMENTATION_MODEL = "dxa_model.pth"
|
||||||
DATA_PATH = "data/cats_dogs/" # → "data/dicom/"
|
DATA_PATH = "dataset_hack/"
|
||||||
INPUT_SIZE = (512, 512)
|
INPUT_SIZE = (224, 224)
|
||||||
NUM_CLASSES = 2 # фон + объект
|
NUM_CLASSES = 2 # quality OK / violation
|
||||||
|
|
||||||
# Абстрактный интерфейс
|
|
||||||
USE_TOTAL_SEGMENTATOR = False # → True на хакатоне
|
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,13 @@
|
||||||
|
"""
|
||||||
|
DXA Quality Assessment Package
|
||||||
|
"""
|
||||||
|
from src.dxa.dataset import DXADataset, create_dataloaders
|
||||||
|
from src.dxa.model import DXAQualityClassifier, DXAQualityModel, create_model
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'DXADataset',
|
||||||
|
'create_dataloaders',
|
||||||
|
'DXAQualityClassifier',
|
||||||
|
'DXAQualityModel',
|
||||||
|
'create_model'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,258 @@
|
||||||
|
"""
|
||||||
|
Dataset for DXA (bone densitometry) quality assessment
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import pandas as pd
|
||||||
|
import numpy as np
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Tuple, Optional
|
||||||
|
import pydicom
|
||||||
|
from PIL import Image
|
||||||
|
import torch
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
|
||||||
|
class DXADataset(Dataset):
|
||||||
|
"""Dataset for DXA bone densitometry images"""
|
||||||
|
|
||||||
|
# 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 _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:
|
||||||
|
# Determine anatomical region from file index
|
||||||
|
# For now, we'll create samples for all regions based on annotation
|
||||||
|
|
||||||
|
# Check each anatomical region
|
||||||
|
for region, config in self.ANATOMICAL_MAPPING.items():
|
||||||
|
total_col = config['total_column']
|
||||||
|
|
||||||
|
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
|
||||||
|
criteria = {}
|
||||||
|
for col in config['columns']:
|
||||||
|
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, # 0 = good, 1 = violation
|
||||||
|
'criteria': criteria,
|
||||||
|
'comment': row.get('комментарий', '')
|
||||||
|
})
|
||||||
|
|
||||||
|
return samples
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.samples)
|
||||||
|
|
||||||
|
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict]:
|
||||||
|
"""Get one sample"""
|
||||||
|
sample = self.samples[idx]
|
||||||
|
|
||||||
|
# Load DICOM
|
||||||
|
ds = pydicom.dcmread(sample['dcm_path'])
|
||||||
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# Convert to 3-channel for pretrained models
|
||||||
|
img = np.stack([img] * 3, axis=0)
|
||||||
|
|
||||||
|
# Convert to uint8 for PIL
|
||||||
|
img = (img * 255).astype(np.uint8)
|
||||||
|
|
||||||
|
# Resize
|
||||||
|
img_pil = Image.fromarray(img.transpose(1, 2, 0))
|
||||||
|
img_pil = img_pil.resize((self.input_size, self.input_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']
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
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"""
|
||||||
|
|
||||||
|
train_dataset = DXADataset(
|
||||||
|
data_root=data_root,
|
||||||
|
annotation_path=annotation_path,
|
||||||
|
input_size=input_size,
|
||||||
|
mode='train'
|
||||||
|
)
|
||||||
|
|
||||||
|
val_dataset = DXADataset(
|
||||||
|
data_root=data_root,
|
||||||
|
annotation_path=annotation_path,
|
||||||
|
input_size=input_size,
|
||||||
|
mode='val'
|
||||||
|
)
|
||||||
|
|
||||||
|
train_loader = torch.utils.data.DataLoader(
|
||||||
|
train_dataset,
|
||||||
|
batch_size=batch_size,
|
||||||
|
shuffle=True,
|
||||||
|
num_workers=num_workers,
|
||||||
|
pin_memory=True
|
||||||
|
)
|
||||||
|
|
||||||
|
val_loader = torch.utils.data.DataLoader(
|
||||||
|
val_dataset,
|
||||||
|
batch_size=batch_size,
|
||||||
|
shuffle=False,
|
||||||
|
num_workers=num_workers,
|
||||||
|
pin_memory=True
|
||||||
|
)
|
||||||
|
|
||||||
|
return train_loader, val_loader
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
# 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)}')
|
||||||
|
|
@ -0,0 +1,261 @@
|
||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
DXA Quality Inference - Batch processing with Excel output
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import argparse
|
||||||
|
from pathlib import Path
|
||||||
|
from datetime import datetime
|
||||||
|
import warnings
|
||||||
|
warnings.filterwarnings('ignore')
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import pydicom
|
||||||
|
from PIL import Image
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
# Add src to path
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||||
|
|
||||||
|
|
||||||
|
def get_device():
|
||||||
|
"""Get best available device"""
|
||||||
|
if torch.backends.mps.is_available():
|
||||||
|
return 'mps'
|
||||||
|
elif torch.cuda.is_available():
|
||||||
|
return 'cuda'
|
||||||
|
else:
|
||||||
|
return 'cpu'
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(model_path: str, backbone: str = 'resnet18', device: str = 'cpu'):
|
||||||
|
"""Load trained model"""
|
||||||
|
from src.dxa.model import create_model
|
||||||
|
|
||||||
|
model = create_model(backbone=backbone, pretrained=False, device=device)
|
||||||
|
model.load(model_path)
|
||||||
|
model.model.eval()
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def load_dicom_image(dcm_path: str, input_size: int = 224) -> torch.Tensor:
|
||||||
|
"""Load and preprocess DICOM image"""
|
||||||
|
ds = pydicom.dcmread(dcm_path)
|
||||||
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# Convert to 3-channel
|
||||||
|
img = np.stack([img] * 3, axis=0)
|
||||||
|
|
||||||
|
# Convert to uint8 for PIL
|
||||||
|
img = (img * 255).astype(np.uint8)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|
||||||
|
|
||||||
|
def determine_anatomical_region(dcm_path: str) -> str:
|
||||||
|
"""Determine anatomical region from DICOM metadata"""
|
||||||
|
try:
|
||||||
|
ds = pydicom.dcmread(dcm_path)
|
||||||
|
h, w = ds.pixel_array.shape
|
||||||
|
|
||||||
|
# Heuristic: hip images are typically smaller
|
||||||
|
if h < 270 or w < 280:
|
||||||
|
return 'hip'
|
||||||
|
else:
|
||||||
|
return 'spine'
|
||||||
|
except:
|
||||||
|
return 'unknown'
|
||||||
|
|
||||||
|
|
||||||
|
def process_dicom_files(args):
|
||||||
|
"""Run inference on DICOM files"""
|
||||||
|
|
||||||
|
# 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!")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Process each file
|
||||||
|
results = []
|
||||||
|
|
||||||
|
for dcm_path in tqdm(dcm_files, desc="Processing"):
|
||||||
|
start_time = datetime.now()
|
||||||
|
|
||||||
|
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
|
||||||
|
output_path = Path(args.output_path)
|
||||||
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
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()}")
|
||||||
|
|
||||||
|
# Show sample output
|
||||||
|
print(f"\nSample output:")
|
||||||
|
print(df.head().to_string())
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
main()
|
||||||
|
|
@ -0,0 +1,199 @@
|
||||||
|
"""
|
||||||
|
DXA Quality Classifier Model
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torchvision.models as models
|
||||||
|
from typing import Dict, Tuple, Optional
|
||||||
|
|
||||||
|
|
||||||
|
class DXAQualityClassifier(nn.Module):
|
||||||
|
"""
|
||||||
|
CNN classifier for DXA image quality assessment
|
||||||
|
Uses pretrained backbone (ResNet/EfficientNet)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
num_classes: int = 2,
|
||||||
|
pretrained: bool = True,
|
||||||
|
dropout: float = 0.3):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
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.feature_dim = feature_dim
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Forward pass"""
|
||||||
|
features = self.backbone(x)
|
||||||
|
return self.classifier(features)
|
||||||
|
|
||||||
|
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Extract features without classification"""
|
||||||
|
return self.backbone(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'
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class DXAQualityModel:
|
||||||
|
"""Wrapper for training and inference"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
model: DXAQualityClassifier,
|
||||||
|
device: str = 'cpu',
|
||||||
|
learning_rate: float = 1e-4):
|
||||||
|
self.model = model
|
||||||
|
self.device = device
|
||||||
|
self.model.to(device)
|
||||||
|
|
||||||
|
# Loss and optimizer
|
||||||
|
self.criterion = nn.CrossEntropyLoss()
|
||||||
|
self.optimizer = torch.optim.AdamW(
|
||||||
|
model.parameters(),
|
||||||
|
lr=learning_rate,
|
||||||
|
weight_decay=1e-5
|
||||||
|
)
|
||||||
|
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||||
|
self.optimizer, mode='min', factor=0.5, patience=3
|
||||||
|
)
|
||||||
|
|
||||||
|
# Training history
|
||||||
|
self.history = {
|
||||||
|
'train_loss': [],
|
||||||
|
'val_loss': [],
|
||||||
|
'train_acc': [],
|
||||||
|
'val_acc': []
|
||||||
|
}
|
||||||
|
|
||||||
|
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 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
|
||||||
|
)
|
||||||
|
return DXAQualityModel(model, device=device)
|
||||||
|
|
@ -0,0 +1,199 @@
|
||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
Training script for DXA Quality Classifier
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import argparse
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
# Add src to path
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||||
|
|
||||||
|
from src.dxa.dataset import create_dataloaders
|
||||||
|
from src.dxa.model import create_model
|
||||||
|
|
||||||
|
|
||||||
|
def get_device():
|
||||||
|
"""Get best available device"""
|
||||||
|
if torch.backends.mps.is_available():
|
||||||
|
return 'mps'
|
||||||
|
elif torch.cuda.is_available():
|
||||||
|
return 'cuda'
|
||||||
|
else:
|
||||||
|
return 'cpu'
|
||||||
|
|
||||||
|
|
||||||
|
def compute_metrics(preds, labels):
|
||||||
|
"""Compute classification metrics"""
|
||||||
|
preds = np.array(preds)
|
||||||
|
labels = np.array(labels)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def train(args):
|
||||||
|
"""Main training loop"""
|
||||||
|
|
||||||
|
# Setup
|
||||||
|
device = get_device()
|
||||||
|
print(f"Using device: {device}")
|
||||||
|
|
||||||
|
# Create output directory
|
||||||
|
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(
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
print(f"Train samples: {len(train_loader.dataset)}")
|
||||||
|
print(f"Val samples: {len(val_loader.dataset)}")
|
||||||
|
|
||||||
|
if len(train_loader.dataset) == 0:
|
||||||
|
print("ERROR: No training samples found!")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Create model
|
||||||
|
print(f"Creating model: {args.backbone}")
|
||||||
|
model = create_model(
|
||||||
|
backbone=args.backbone,
|
||||||
|
num_classes=2,
|
||||||
|
pretrained=True,
|
||||||
|
device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Training loop
|
||||||
|
best_val_f1 = 0
|
||||||
|
best_epoch = 0
|
||||||
|
|
||||||
|
for epoch in range(args.epochs):
|
||||||
|
print(f"\n{'='*50}")
|
||||||
|
print(f"Epoch {epoch+1}/{args.epochs}")
|
||||||
|
print(f"{'='*50}")
|
||||||
|
|
||||||
|
# Train
|
||||||
|
train_loss, train_acc = model.train_epoch(train_loader)
|
||||||
|
|
||||||
|
# Validate
|
||||||
|
val_loss, val_acc = model.validate(val_loader)
|
||||||
|
|
||||||
|
# Compute detailed metrics
|
||||||
|
model.model.eval()
|
||||||
|
all_preds = []
|
||||||
|
all_labels = []
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for images, labels in val_loader:
|
||||||
|
images = images.to(device)
|
||||||
|
labels_dict = labels if isinstance(labels, dict) else {'label': labels}
|
||||||
|
labels_arr = torch.tensor([l['label'] for l in labels_dict]).to(device)
|
||||||
|
|
||||||
|
preds, _ = model.predict(images)
|
||||||
|
all_preds.extend(preds.cpu().numpy())
|
||||||
|
all_labels.extend(labels_arr.cpu().numpy())
|
||||||
|
|
||||||
|
metrics = compute_metrics(all_preds, all_labels)
|
||||||
|
|
||||||
|
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}")
|
||||||
|
|
||||||
|
# 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})")
|
||||||
|
|
||||||
|
# Save checkpoint
|
||||||
|
if (epoch + 1) % args.save_every == 0:
|
||||||
|
model.save(str(output_dir / f'checkpoint_epoch_{epoch+1}.pth'))
|
||||||
|
|
||||||
|
print(f"\n{'='*50}")
|
||||||
|
print(f"Training complete!")
|
||||||
|
print(f"Best F1: {best_val_f1:.4f} at epoch {best_epoch}")
|
||||||
|
print(f"{'='*50}")
|
||||||
|
|
||||||
|
# Save final model
|
||||||
|
model.save(str(output_dir / 'final_model.pth'))
|
||||||
|
print(f"Final model saved to {output_dir / 'final_model.pth'}")
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
main()
|
||||||
321
src/main.py
321
src/main.py
|
|
@ -1,51 +1,306 @@
|
||||||
from fastapi import FastAPI
|
"""
|
||||||
|
Main FastAPI application for DXA Quality Assessment
|
||||||
|
"""
|
||||||
|
from fastapi import FastAPI, File, UploadFile, APIRouter
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
import numpy as np
|
||||||
|
import pydicom
|
||||||
|
import pandas as pd
|
||||||
|
from PIL import Image
|
||||||
|
import io
|
||||||
|
import torch
|
||||||
|
import base64
|
||||||
|
|
||||||
from starlette.responses import JSONResponse
|
from starlette.responses import JSONResponse, FileResponse
|
||||||
|
|
||||||
from src.api import endpoints, root, annotation
|
from src.dxa.model import create_model
|
||||||
|
from src.dxa.dataset import DXADataset
|
||||||
|
from src.utils.utils import get_device
|
||||||
|
|
||||||
# Создаем приложение
|
# Create app
|
||||||
app = FastAPI(
|
app = FastAPI(
|
||||||
title="Bone Quality Assessment API",
|
title="DXA Quality Assessment API",
|
||||||
description="API для оценки качества медицинских изображений",
|
description="API for automated quality assessment of bone densitometry (DXA) studies",
|
||||||
version="1.0.0",
|
version="1.0.0",
|
||||||
docs_url="/docs",
|
docs_url="/docs",
|
||||||
redoc_url="/redoc",
|
redoc_url="/redoc",
|
||||||
openapi_url="/openapi.json"
|
openapi_url="/openapi.json"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Подключаем статические файлы
|
# Mount static files if exists
|
||||||
static_dir = Path(__file__).parent / "api/static"
|
static_dir = Path(__file__).parent / "api/static"
|
||||||
if static_dir.exists():
|
if static_dir.exists():
|
||||||
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
||||||
print(f"✅ Статика подключена: {static_dir}")
|
|
||||||
else:
|
|
||||||
print(f"⚠️ Папка static не найдена: {static_dir}")
|
|
||||||
|
|
||||||
app.include_router(endpoints)
|
# Global model
|
||||||
app.include_router(root)
|
dxa_model = None
|
||||||
app.include_router(annotation)
|
device = None
|
||||||
|
|
||||||
# Глобальный обработчик 404
|
|
||||||
@app.exception_handler(404)
|
def load_model():
|
||||||
async def not_found_handler(request, exc):
|
"""Load DXA model"""
|
||||||
return JSONResponse(
|
global dxa_model, device
|
||||||
status_code=404,
|
|
||||||
content={
|
if dxa_model is None:
|
||||||
"error": "Not Found",
|
device = get_device()
|
||||||
"message": f"Endpoint {request.url.path} not found",
|
print(f"Loading DXA model on {device}...")
|
||||||
"available_endpoints": [
|
|
||||||
"/",
|
try:
|
||||||
"/docs",
|
dxa_model = create_model(
|
||||||
"/redoc",
|
backbone='resnet18',
|
||||||
"/openapi.json",
|
pretrained=False,
|
||||||
"/api/v1/health",
|
device=device
|
||||||
"/api/v1/analyze (POST)",
|
)
|
||||||
"/api/v1/switch_mode",
|
dxa_model.load('../models/dxa_model.pth')
|
||||||
"/api/v1/info",
|
dxa_model.model.eval()
|
||||||
"/static/"
|
print("DXA model loaded successfully")
|
||||||
]
|
except Exception as e:
|
||||||
|
print(f"Error loading model: {e}")
|
||||||
|
dxa_model = None
|
||||||
|
|
||||||
|
return dxa_model
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_dicom(dcm_bytes: bytes, input_size: int = 224):
|
||||||
|
"""Preprocess DICOM for model input"""
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
with tempfile.NamedTemporaryFile(suffix='.dcm', delete=False) as f:
|
||||||
|
f.write(dcm_bytes)
|
||||||
|
dcm_path = f.name
|
||||||
|
|
||||||
|
ds = pydicom.dcmread(dcm_path)
|
||||||
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
|
||||||
|
# Normalize
|
||||||
|
img = (img - img.min()) / (img.max() - img.min() + 1e-8)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|
||||||
|
# Tensor
|
||||||
|
img = torch.from_numpy(img).float().unsqueeze(0)
|
||||||
|
|
||||||
|
return img, ds
|
||||||
|
|
||||||
|
|
||||||
|
# API Routes
|
||||||
|
@app.get("/")
|
||||||
|
async def root():
|
||||||
|
"""Root endpoint - serve web interface"""
|
||||||
|
static_dir = Path(__file__).parent / "api/static"
|
||||||
|
index_path = static_dir / "index.html"
|
||||||
|
|
||||||
|
if index_path.exists():
|
||||||
|
return FileResponse(str(index_path))
|
||||||
|
|
||||||
|
return {
|
||||||
|
"service": "DXA Quality Assessment API",
|
||||||
|
"version": "1.0.0",
|
||||||
|
"status": "operational",
|
||||||
|
"endpoints": [
|
||||||
|
"/",
|
||||||
|
"/docs",
|
||||||
|
"/api/v1/health",
|
||||||
|
"/api/v1/analyze",
|
||||||
|
"/api/v1/batch"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/api/v1/health")
|
||||||
|
async def health_check():
|
||||||
|
"""Health check"""
|
||||||
|
model = load_model()
|
||||||
|
return {
|
||||||
|
"status": "ok",
|
||||||
|
"model_loaded": model is not None,
|
||||||
|
"device": str(device) if device else "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/api/v1/analyze")
|
||||||
|
async def analyze_dicom(file: UploadFile = File(...)):
|
||||||
|
"""Analyze single DICOM file"""
|
||||||
|
try:
|
||||||
|
# Load model
|
||||||
|
model = load_model()
|
||||||
|
if model is None:
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={"error": "Model not loaded"}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Read file
|
||||||
|
dcm_bytes = await file.read()
|
||||||
|
img_tensor, ds = preprocess_dicom(dcm_bytes)
|
||||||
|
|
||||||
|
# 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 = 'hip' if h < 270 else 'spine'
|
||||||
|
|
||||||
|
# Result
|
||||||
|
result = {
|
||||||
|
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
||||||
|
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
||||||
|
"anatomical_region": region,
|
||||||
|
"quality_class": int(pred),
|
||||||
|
"quality_label": "OK" if pred == 0 else "Violation detected",
|
||||||
|
"confidence": round(confidence, 4),
|
||||||
|
"processing_status": "Success"
|
||||||
}
|
}
|
||||||
)
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
import traceback
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={
|
||||||
|
"error": str(e),
|
||||||
|
"processing_status": f"Failure: {str(e)[:50]}"
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/api/v1/batch")
|
||||||
|
async def batch_analyze(files: list[UploadFile] = File(...)):
|
||||||
|
"""Batch analyze multiple DICOM files"""
|
||||||
|
results = []
|
||||||
|
|
||||||
|
model = load_model()
|
||||||
|
if model is None:
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={"error": "Model not loaded"}
|
||||||
|
)
|
||||||
|
|
||||||
|
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 = 'hip' if h < 270 else 'spine'
|
||||||
|
|
||||||
|
results.append({
|
||||||
|
"filename": file.filename,
|
||||||
|
"study_uid": getattr(ds, 'StudyInstanceUID', ''),
|
||||||
|
"image_uid": getattr(ds, 'SOPInstanceUID', ''),
|
||||||
|
"anatomical_region": region,
|
||||||
|
"quality_class": int(pred),
|
||||||
|
"confidence": round(confidence, 4),
|
||||||
|
"processing_status": "Success"
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
results.append({
|
||||||
|
"filename": file.filename,
|
||||||
|
"error": str(e),
|
||||||
|
"processing_status": "Failure"
|
||||||
|
})
|
||||||
|
|
||||||
|
return {"results": results}
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/api/v1/export")
|
||||||
|
async def export_results(files: list[UploadFile] = File(...)):
|
||||||
|
"""Batch analyze and export results as XLSX"""
|
||||||
|
results = []
|
||||||
|
|
||||||
|
model = load_model()
|
||||||
|
if model is None:
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={"error": "Model not loaded"}
|
||||||
|
)
|
||||||
|
|
||||||
|
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 = 'hip' if h < 270 else 'spine'
|
||||||
|
|
||||||
|
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'
|
||||||
|
})
|
||||||
|
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]}'
|
||||||
|
})
|
||||||
|
|
||||||
|
# Create DataFrame and export to Excel
|
||||||
|
df = pd.DataFrame(results)
|
||||||
|
|
||||||
|
# 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')
|
||||||
|
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'}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import uvicorn
|
||||||
|
uvicorn.run(app, host="0.0.0.0", port=8000)
|
||||||
|
|
|
||||||
|
|
@ -1,61 +0,0 @@
|
||||||
# src/segmentators/unet_segmentator.py
|
|
||||||
from pathlib import Path
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from PIL import Image
|
|
||||||
|
|
||||||
from src import UNet
|
|
||||||
from src.segmentators.base import BaseSegmentator
|
|
||||||
|
|
||||||
|
|
||||||
class UNetSegmentator(BaseSegmentator):
|
|
||||||
"""Сегментатор на основе U-Net (для кошек/собак)"""
|
|
||||||
|
|
||||||
def __init__(self, model_path="../models/unet_cats_dogs.pth", device='cpu'):
|
|
||||||
self.device = device
|
|
||||||
self.model = UNet(in_channels=3, out_classes=2)
|
|
||||||
|
|
||||||
if Path(model_path).exists():
|
|
||||||
self.model.load_state_dict(torch.load(model_path, map_location=device))
|
|
||||||
self.model.to(device)
|
|
||||||
self.model.eval()
|
|
||||||
self.loaded = True
|
|
||||||
print(f"✅ U-Net загружен из {model_path}")
|
|
||||||
else:
|
|
||||||
print(f"⚠️ U-Net не найден: {model_path}")
|
|
||||||
self.loaded = False
|
|
||||||
|
|
||||||
def segment(self, image):
|
|
||||||
if not self.loaded:
|
|
||||||
return np.zeros((256, 256), dtype=np.int64)
|
|
||||||
|
|
||||||
# Подготовка изображения
|
|
||||||
if isinstance(image, np.ndarray):
|
|
||||||
image = Image.fromarray(image)
|
|
||||||
|
|
||||||
original_size = image.size
|
|
||||||
image_resized = image.resize((256, 256))
|
|
||||||
image_array = np.array(image_resized, dtype=np.float32) / 255.0
|
|
||||||
image_array = image_array.transpose(2, 0, 1)
|
|
||||||
image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(self.device)
|
|
||||||
|
|
||||||
# Инференс
|
|
||||||
with torch.no_grad():
|
|
||||||
output = self.model(image_tensor)
|
|
||||||
probs = torch.softmax(output, dim=1)
|
|
||||||
mask = probs.argmax(dim=1).squeeze().cpu().numpy()
|
|
||||||
|
|
||||||
# Ресайз к оригиналу
|
|
||||||
mask_pil = Image.fromarray(mask.astype(np.uint8))
|
|
||||||
mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST)
|
|
||||||
|
|
||||||
return np.array(mask_pil)
|
|
||||||
|
|
||||||
def get_info(self):
|
|
||||||
return {
|
|
||||||
"name": "U-Net",
|
|
||||||
"type": "segmentation",
|
|
||||||
"target": "cats_and_dogs",
|
|
||||||
"classes": 2,
|
|
||||||
"loaded": self.loaded
|
|
||||||
}
|
|
||||||
285
train.py
285
train.py
|
|
@ -1,285 +0,0 @@
|
||||||
# train.py
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.optim as optim
|
|
||||||
import numpy as np
|
|
||||||
from pathlib import Path
|
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
from tqdm import tqdm
|
|
||||||
import time
|
|
||||||
|
|
||||||
from src import UNet
|
|
||||||
from src.dataloaders.pet_dataset import PetDataset, create_dataloaders
|
|
||||||
|
|
||||||
|
|
||||||
def get_device():
|
|
||||||
"""Автоматический выбор устройства"""
|
|
||||||
if torch.backends.mps.is_available():
|
|
||||||
device = torch.device("mps")
|
|
||||||
print(f"✅ Используем MPS (Apple Silicon) - {torch.backends.mps.is_built()}")
|
|
||||||
elif torch.cuda.is_available():
|
|
||||||
device = torch.device("cuda")
|
|
||||||
print(f"✅ Используем CUDA - {torch.cuda.get_device_name(0)}")
|
|
||||||
else:
|
|
||||||
device = torch.device("cpu")
|
|
||||||
print("⚠️ Используем CPU (медленно)")
|
|
||||||
|
|
||||||
# Проверка производительности
|
|
||||||
if device.type == "mps":
|
|
||||||
# MPS иногда тормозит на некоторых операциях, проверяем
|
|
||||||
test_tensor = torch.randn(1000, 1000).to(device)
|
|
||||||
start = time.time()
|
|
||||||
_ = test_tensor @ test_tensor.T
|
|
||||||
elapsed = time.time() - start
|
|
||||||
print(f" MPS скорость: {elapsed:.4f} сек для matmul 1000x1000")
|
|
||||||
|
|
||||||
return device
|
|
||||||
|
|
||||||
|
|
||||||
class DiceLoss(nn.Module):
|
|
||||||
"""Dice Loss для сегментации"""
|
|
||||||
|
|
||||||
def __init__(self, smooth=1e-6):
|
|
||||||
super().__init__()
|
|
||||||
self.smooth = smooth
|
|
||||||
|
|
||||||
def forward(self, pred, target):
|
|
||||||
pred = torch.softmax(pred, dim=1)
|
|
||||||
target_one_hot = torch.nn.functional.one_hot(target, num_classes=pred.shape[1])
|
|
||||||
target_one_hot = target_one_hot.permute(0, 3, 1, 2).float()
|
|
||||||
|
|
||||||
intersection = (pred * target_one_hot).sum(dim=(2, 3))
|
|
||||||
union = pred.sum(dim=(2, 3)) + target_one_hot.sum(dim=(2, 3))
|
|
||||||
dice = (2. * intersection + self.smooth) / (union + self.smooth)
|
|
||||||
|
|
||||||
return 1 - dice.mean()
|
|
||||||
|
|
||||||
|
|
||||||
class Trainer:
|
|
||||||
def __init__(self, model, train_loader, val_loader, device,
|
|
||||||
learning_rate=1e-4, use_amp=False):
|
|
||||||
self.model = model.to(device)
|
|
||||||
self.train_loader = train_loader
|
|
||||||
self.val_loader = val_loader
|
|
||||||
self.device = device
|
|
||||||
self.use_amp = use_amp and device.type == "cuda"
|
|
||||||
|
|
||||||
# Потери
|
|
||||||
self.criterion_ce = nn.CrossEntropyLoss()
|
|
||||||
self.criterion_dice = DiceLoss()
|
|
||||||
|
|
||||||
# Оптимизатор
|
|
||||||
self.optimizer = optim.AdamW(model.parameters(), lr=learning_rate,
|
|
||||||
weight_decay=1e-5)
|
|
||||||
self.scheduler = optim.lr_scheduler.ReduceLROnPlateau(
|
|
||||||
self.optimizer, mode='min', factor=0.5, patience=5
|
|
||||||
)
|
|
||||||
|
|
||||||
# AMP
|
|
||||||
if self.use_amp:
|
|
||||||
self.scaler = torch.cuda.amp.GradScaler()
|
|
||||||
|
|
||||||
# История
|
|
||||||
self.train_losses = []
|
|
||||||
self.val_losses = []
|
|
||||||
self.val_dice_scores = []
|
|
||||||
|
|
||||||
print(f"✅ Trainer инициализирован на {device}")
|
|
||||||
|
|
||||||
def train_epoch(self):
|
|
||||||
self.model.train()
|
|
||||||
total_loss = 0
|
|
||||||
|
|
||||||
for images, masks in tqdm(self.train_loader, desc='Training'):
|
|
||||||
images = images.to(self.device)
|
|
||||||
masks = masks.to(self.device)
|
|
||||||
|
|
||||||
self.optimizer.zero_grad()
|
|
||||||
|
|
||||||
if self.use_amp:
|
|
||||||
with torch.cuda.amp.autocast():
|
|
||||||
outputs = self.model(images)
|
|
||||||
loss_ce = self.criterion_ce(outputs, masks)
|
|
||||||
loss_dice = self.criterion_dice(outputs, masks)
|
|
||||||
loss = loss_ce + loss_dice
|
|
||||||
|
|
||||||
self.scaler.scale(loss).backward()
|
|
||||||
self.scaler.step(self.optimizer)
|
|
||||||
self.scaler.update()
|
|
||||||
else:
|
|
||||||
outputs = self.model(images)
|
|
||||||
loss_ce = self.criterion_ce(outputs, masks)
|
|
||||||
loss_dice = self.criterion_dice(outputs, masks)
|
|
||||||
loss = loss_ce + loss_dice
|
|
||||||
|
|
||||||
loss.backward()
|
|
||||||
self.optimizer.step()
|
|
||||||
|
|
||||||
total_loss += loss.item()
|
|
||||||
|
|
||||||
return total_loss / len(self.train_loader)
|
|
||||||
|
|
||||||
def validate(self):
|
|
||||||
self.model.eval()
|
|
||||||
total_loss = 0
|
|
||||||
dice_scores = []
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
for images, masks in tqdm(self.val_loader, desc='Validation'):
|
|
||||||
images = images.to(self.device)
|
|
||||||
masks = masks.to(self.device)
|
|
||||||
|
|
||||||
outputs = self.model(images)
|
|
||||||
|
|
||||||
# Потери
|
|
||||||
loss_ce = self.criterion_ce(outputs, masks)
|
|
||||||
loss_dice = self.criterion_dice(outputs, masks)
|
|
||||||
loss = loss_ce + loss_dice
|
|
||||||
total_loss += loss.item()
|
|
||||||
|
|
||||||
# Исправленный Dice Score
|
|
||||||
pred = torch.softmax(outputs, dim=1)
|
|
||||||
pred_mask = pred.argmax(dim=1)
|
|
||||||
|
|
||||||
# Вычисляем Dice правильно
|
|
||||||
dice = self.compute_dice_score(pred_mask, masks)
|
|
||||||
dice_scores.append(dice)
|
|
||||||
|
|
||||||
avg_loss = total_loss / len(self.val_loader)
|
|
||||||
avg_dice = np.mean(dice_scores)
|
|
||||||
|
|
||||||
return avg_loss, avg_dice
|
|
||||||
|
|
||||||
def compute_dice_score(self, pred, target):
|
|
||||||
"""
|
|
||||||
Правильный Dice Score
|
|
||||||
pred: [B, H, W] - предсказанные маски (0 или 1)
|
|
||||||
target: [B, H, W] - истинные маски (0 или 1)
|
|
||||||
"""
|
|
||||||
smooth = 1e-6
|
|
||||||
|
|
||||||
# Преобразуем в float
|
|
||||||
pred = pred.float()
|
|
||||||
target = target.float()
|
|
||||||
|
|
||||||
# Пересечение
|
|
||||||
intersection = (pred * target).sum(dim=(1, 2))
|
|
||||||
|
|
||||||
# Суммы
|
|
||||||
pred_sum = pred.sum(dim=(1, 2))
|
|
||||||
target_sum = target.sum(dim=(1, 2))
|
|
||||||
|
|
||||||
# Dice = 2 * |A∩B| / (|A| + |B|)
|
|
||||||
dice = (2. * intersection + smooth) / (pred_sum + target_sum + smooth)
|
|
||||||
|
|
||||||
return dice.mean().item()
|
|
||||||
|
|
||||||
def train(self, epochs, save_best=True):
|
|
||||||
best_dice = 0
|
|
||||||
|
|
||||||
for epoch in range(epochs):
|
|
||||||
print(f"\n{'=' * 50}")
|
|
||||||
print(f"Epoch {epoch + 1}/{epochs}")
|
|
||||||
print(f"{'=' * 50}")
|
|
||||||
|
|
||||||
# Train
|
|
||||||
train_loss = self.train_epoch()
|
|
||||||
self.train_losses.append(train_loss)
|
|
||||||
|
|
||||||
# Validate
|
|
||||||
val_loss, val_dice = self.validate()
|
|
||||||
self.val_losses.append(val_loss)
|
|
||||||
self.val_dice_scores.append(val_dice)
|
|
||||||
|
|
||||||
# Scheduler
|
|
||||||
self.scheduler.step(val_loss)
|
|
||||||
|
|
||||||
current_lr = self.optimizer.param_groups[0]['lr']
|
|
||||||
|
|
||||||
print(f"\n📊 Результаты:")
|
|
||||||
print(f" Train Loss: {train_loss:.4f}")
|
|
||||||
print(f" Val Loss: {val_loss:.4f}")
|
|
||||||
print(f" Val Dice: {val_dice:.4f}") # Теперь должно быть 0-1
|
|
||||||
print(f" LR: {current_lr:.6f}")
|
|
||||||
|
|
||||||
# Сохраняем лучшую модель
|
|
||||||
if save_best and val_dice > best_dice:
|
|
||||||
best_dice = val_dice
|
|
||||||
torch.save(self.model.state_dict(), 'best_model_old.pth')
|
|
||||||
print(f" ✅ Saved best model (Dice: {val_dice:.4f})")
|
|
||||||
|
|
||||||
def plot_history(self):
|
|
||||||
fig, axes = plt.subplots(1, 2, figsize=(12, 4))
|
|
||||||
|
|
||||||
axes[0].plot(self.train_losses, label='Train Loss')
|
|
||||||
axes[0].plot(self.val_losses, label='Val Loss')
|
|
||||||
axes[0].set_xlabel('Epoch')
|
|
||||||
axes[0].set_ylabel('Loss')
|
|
||||||
axes[0].set_title('Training History')
|
|
||||||
axes[0].legend()
|
|
||||||
axes[0].grid(True)
|
|
||||||
|
|
||||||
axes[1].plot(self.val_dice_scores, label='Val Dice', color='green')
|
|
||||||
axes[1].set_xlabel('Epoch')
|
|
||||||
axes[1].set_ylabel('Dice Score')
|
|
||||||
axes[1].set_title('Validation Dice Score')
|
|
||||||
axes[1].legend()
|
|
||||||
axes[1].grid(True)
|
|
||||||
|
|
||||||
plt.tight_layout()
|
|
||||||
plt.savefig('training_history.png', dpi=150)
|
|
||||||
plt.show()
|
|
||||||
|
|
||||||
def main():
|
|
||||||
# Пути
|
|
||||||
data_root = "/Users/denis/workspace/bone_2026/datasets/oxford-iiit-pet"
|
|
||||||
model_save_path = "models/unet_cats_dogs.pth"
|
|
||||||
|
|
||||||
# Создаем папку для моделей
|
|
||||||
Path("models").mkdir(exist_ok=True)
|
|
||||||
|
|
||||||
# Определяем устройство
|
|
||||||
device = get_device()
|
|
||||||
|
|
||||||
# Настройки для быстрого обучения на MPS
|
|
||||||
batch_size = 16 if device.type == "mps" else 8 # MPS может больше
|
|
||||||
input_size = (256, 256)
|
|
||||||
epochs = 20
|
|
||||||
|
|
||||||
print(f"\n📋 Параметры обучения:")
|
|
||||||
print(f" Batch size: {batch_size}")
|
|
||||||
print(f" Input size: {input_size}")
|
|
||||||
print(f" Epochs: {epochs}")
|
|
||||||
print(f" Device: {device}")
|
|
||||||
|
|
||||||
# Загрузка данных
|
|
||||||
print("\n📂 Загрузка данных...")
|
|
||||||
train_loader, val_loader = create_dataloaders(
|
|
||||||
data_root=data_root,
|
|
||||||
batch_size=batch_size,
|
|
||||||
input_size=input_size
|
|
||||||
)
|
|
||||||
|
|
||||||
# Модель
|
|
||||||
model = UNet(in_channels=3, out_classes=2)
|
|
||||||
|
|
||||||
# Тренировка
|
|
||||||
trainer = Trainer(
|
|
||||||
model,
|
|
||||||
train_loader,
|
|
||||||
val_loader,
|
|
||||||
device=device,
|
|
||||||
learning_rate=1e-4
|
|
||||||
)
|
|
||||||
trainer.train(epochs=epochs)
|
|
||||||
|
|
||||||
# Сохраняем модель
|
|
||||||
torch.save(model.state_dict(), model_save_path)
|
|
||||||
print(f"\n✅ Model saved to {model_save_path}")
|
|
||||||
|
|
||||||
# Визуализация
|
|
||||||
trainer.plot_history()
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
Loading…
Reference in New Issue