From a5dff35ce412a467f3bc5d72b59c8a8c3b226b93 Mon Sep 17 00:00:00 2001 From: denis Date: Tue, 22 Sep 2026 15:49:54 +0300 Subject: [PATCH] develop - hack_2026 --- Dockerfile | 41 +- QWEN.md | 193 ++++ README.md | 38 +- inference.py | 140 --- requirements.txt | 84 +- run.sh | 134 +++ src/api/static/index.html | 1076 +++++------------------ src/api/static/js/dxa-app.js | 356 ++++++++ src/classifiers/pet_breed_classifier.py | 174 ---- src/config.py | 13 +- src/dxa/__init__.py | 13 + src/dxa/dataset.py | 258 ++++++ src/dxa/inference.py | 261 ++++++ src/dxa/model.py | 199 +++++ src/dxa/train.py | 199 +++++ src/main.py | 321 ++++++- src/segmentators/unet_segmentator.py | 61 -- train.py | 285 ------ 18 files changed, 2220 insertions(+), 1626 deletions(-) create mode 100644 QWEN.md delete mode 100644 inference.py create mode 100644 run.sh create mode 100644 src/api/static/js/dxa-app.js delete mode 100644 src/classifiers/pet_breed_classifier.py create mode 100644 src/dxa/__init__.py create mode 100644 src/dxa/dataset.py create mode 100644 src/dxa/inference.py create mode 100644 src/dxa/model.py create mode 100644 src/dxa/train.py delete mode 100644 src/segmentators/unet_segmentator.py delete mode 100644 train.py diff --git a/Dockerfile b/Dockerfile index 5d87819..baf3c32 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,22 +1,16 @@ -# --- Этап сборки --- -FROM python:3.10-slim AS builder +# DXA Quality Assessment - Docker Container +# 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 \ - && rm -rf /var/lib/apt/lists/* - -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 \ + libgl1-mesa-glx \ libglib2.0-0 \ libsm6 \ libxext6 \ @@ -26,15 +20,20 @@ RUN apt-get update && apt-get install -y --no-install-recommends \ WORKDIR /app -# Копируем установленные пакеты из builder-этапа -COPY --from=builder /root/.local /root/.local -ENV PATH=/root/.local/bin:$PATH +# Copy requirements and install Python dependencies +COPY requirements.txt . +RUN pip3 install --no-cache-dir -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu118 +# Copy application code COPY . . RUN mkdir -p models +# Expose API port EXPOSE 8000 + +# Environment variables ENV PYTHONUNBUFFERED=1 ENV PYTHONPATH=/app -CMD ["python", "src/run.py"] \ No newline at end of file +# Default command +CMD ["python3", "-m", "uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/QWEN.md b/QWEN.md new file mode 100644 index 0000000..d204199 --- /dev/null +++ b/QWEN.md @@ -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) diff --git a/README.md b/README.md index 5a9ef6a..6d5a345 100644 --- a/README.md +++ b/README.md @@ -26,6 +26,18 @@ ![!img](public/static/arch.png) +## 🎯 Доступные режимы + +### 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 # 4. Загрузка обученной модели (опционально) -# Поместите модель в папку models/unet_cats_dogs.pth +# Поместите модель в папку models/dxa_model.pth # 5. Запуск сервера python run.py @@ -145,22 +157,22 @@ curl -X POST "http://localhost:8000/api/v1/analyze" \ 📁 Структура проекта ```text -bone-quality-assessment/ +bone_2026/ ├── src/ -│ ├── api/ -│ │ ├── endpoints.py # FastAPI эндпоинты -│ │ └── static/ -│ │ └── index.html # Web-интерфейс -│ ├── models/ -│ │ └── unet.py # U-Net архитектура -│ └── quality/ -│ └── quality_scorer.py # Оценка качества +│ ├── dxa/ # DXA Quality модуль +│ │ ├── model.py # ResNet18 классификатор +│ │ ├── dataset.py # Загрузчик данных +│ │ ├── train.py # Обучение +│ │ └── inference.py # Инференс +│ ├── api/ # REST API +│ ├── quality/ # Оценка качества +│ └── main.py # FastAPI приложение ├── models/ -│ └── unet_cats_dogs.pth # Обученная модель +│ └── dxa_model.pth # Обученная модель +├── dataset_hack/ # DICOM датасет ├── Dockerfile -├── docker-compose.yml ├── requirements.txt -├── run.py +├── run.sh └── README.md ``` diff --git a/inference.py b/inference.py deleted file mode 100644 index 2e0b3f0..0000000 --- a/inference.py +++ /dev/null @@ -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() diff --git a/requirements.txt b/requirements.txt index bbafbbf..31f914e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,47 +1,37 @@ -annotated-doc==0.0.5 -annotated-types==0.7.0 -anyio==4.12.1 -click==8.1.8 -contourpy==1.3.0 -cycler==0.12.1 -exceptiongroup==1.3.1 -fastapi==0.128.8 -filelock==3.19.1 -fonttools==4.60.2 -fsspec==2025.10.0 -h11==0.16.0 -idna==3.18 -importlib_resources==6.5.2 -Jinja2==3.1.6 -kiwisolver==1.4.7 -MarkupSafe==3.0.3 -matplotlib==3.9.4 -monai==1.5.2 -mpmath==1.3.0 -networkx==3.2.1 -nibabel==5.3.3 -numpy==2.0.2 -opencv-python-headless==5.0.0.93 -packaging==26.3 -pandas==2.3.3 -pillow==11.3.0 -pydantic==2.13.4 -pydantic_core==2.46.4 -pydicom==2.4.4 -pyparsing==3.3.2 -python-dateutil==2.9.0.post0 -python-multipart==0.0.20 -pytz==2026.3.post1 -scipy==1.13.1 -six==1.17.0 -starlette==0.49.3 -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 +# Core dependencies +torch>=2.0.0 +torchvision>=0.15.0 +numpy>=1.24.0 + +# Image processing +Pillow>=10.0.0 +opencv-python-headless>=4.8.0 + +# Medical imaging +pydicom>=2.4.0 +pydicom-seg>=0.4.0 +nibabel>=5.0.0 +SimpleITK>=2.2.0 + +# Deep learning / models +torchvision>=0.15.0 +timm>=0.9.0 + +# Data handling +pandas>=2.0.0 +openpyxl>=3.1.0 + +# Visualization +matplotlib>=3.7.0 + +# API / serving +fastapi>=0.100.0 +uvicorn>=0.23.0 +python-multipart>=0.0.6 + +# Progress bars +tqdm>=4.65.0 + +# Utils +scipy>=1.10.0 +scikit-learn>=1.3.0 diff --git a/run.sh b/run.sh new file mode 100644 index 0000000..4aac654 --- /dev/null +++ b/run.sh @@ -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 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 diff --git a/src/api/static/index.html b/src/api/static/index.html index dd2a24b..14adafa 100644 --- a/src/api/static/index.html +++ b/src/api/static/index.html @@ -3,873 +3,261 @@ - Bone Quality Assessment - Анализ качества костной ткани + DXA Quality Assessment + + + - -
-

🦴 Bone Quality Assessment

-

Загрузите медицинское изображение для анализа качества костной ткани

- - -
- 📤 -

Перетащите изображение сюда

-

или кликните для выбора файла

-

Поддерживаются: PNG, JPG, JPEG

- -
- - -
- Preview -
-
- - • - -
-
- - - - -
- - -
-

🧩 Визуализация сегментации

-
- - - -
-
- Mask Visualization - -
-
- 🧊 Объектов: - - 📏 Размер: - - 🎯 Уверенность: - -
-
- - -
-
-
Общее качество
-
-
- -
-
Серьезность проблемы
-
-
- -
-
Выявленные проблемы
-
-
- -
-
-
📍 Позиционирование
-
-
+ + +
+
+
+
+
-
-
🔬 Артефакты
-
-
+
+

DXA Quality Assessment

+

Оценка качества денситометрических исследований

+
+
+
+ + + + API Docs + +
+
+
+ + +
+ +
+
+
+ Проверка статуса сервиса... +
+
+ + +
+

+ + Загрузка DICOM файлов +

+ +
+ +
+ +

+ Перетащите DICOM файлы сюда или нажмите для выбора +

+

+ Поддерживаются файлы .dcm +

-
-
Уверенность модели
-
-
-
+ + +
+ + +
- + - \ No newline at end of file + diff --git a/src/api/static/js/dxa-app.js b/src/api/static/js/dxa-app.js new file mode 100644 index 0000000..058c282 --- /dev/null +++ b/src/api/static/js/dxa-app.js @@ -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 ` + + +
+ + ${r.filename || '—'} +
+ ${r.study_uid ? `
UID: ${r.study_uid.substring(0, 20)}...
` : ''} + + + + ${regionLabel} + + + + + ${qualityLabel} + + + +
+
+
+
+ ${confidence} +
+ + + `; + }).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 = ' Экспорт...'; + 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(); diff --git a/src/classifiers/pet_breed_classifier.py b/src/classifiers/pet_breed_classifier.py deleted file mode 100644 index 807473c..0000000 --- a/src/classifiers/pet_breed_classifier.py +++ /dev/null @@ -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) - } \ No newline at end of file diff --git a/src/config.py b/src/config.py index e7c4808..7b6d5d4 100644 --- a/src/config.py +++ b/src/config.py @@ -1,9 +1,6 @@ class Config: - # На хакатоне просто меняешь эти пути! - SEGMENTATION_MODEL = "unet_cats_dogs.pth" # → "totalsegmentator" - DATA_PATH = "data/cats_dogs/" # → "data/dicom/" - INPUT_SIZE = (512, 512) - NUM_CLASSES = 2 # фон + объект - - # Абстрактный интерфейс - USE_TOTAL_SEGMENTATOR = False # → True на хакатоне \ No newline at end of file + # DXA Quality Assessment Configuration + SEGMENTATION_MODEL = "dxa_model.pth" + DATA_PATH = "dataset_hack/" + INPUT_SIZE = (224, 224) + NUM_CLASSES = 2 # quality OK / violation diff --git a/src/dxa/__init__.py b/src/dxa/__init__.py new file mode 100644 index 0000000..7f2205c --- /dev/null +++ b/src/dxa/__init__.py @@ -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' +] diff --git a/src/dxa/dataset.py b/src/dxa/dataset.py new file mode 100644 index 0000000..b76aacf --- /dev/null +++ b/src/dxa/dataset.py @@ -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)}') diff --git a/src/dxa/inference.py b/src/dxa/inference.py new file mode 100644 index 0000000..426ba08 --- /dev/null +++ b/src/dxa/inference.py @@ -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() diff --git a/src/dxa/model.py b/src/dxa/model.py new file mode 100644 index 0000000..93baf12 --- /dev/null +++ b/src/dxa/model.py @@ -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) diff --git a/src/dxa/train.py b/src/dxa/train.py new file mode 100644 index 0000000..d779c0d --- /dev/null +++ b/src/dxa/train.py @@ -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() diff --git a/src/main.py b/src/main.py index eca67b0..be77aa1 100644 --- a/src/main.py +++ b/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 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( - title="Bone Quality Assessment API", - description="API для оценки качества медицинских изображений", + title="DXA Quality Assessment API", + description="API for automated quality assessment of bone densitometry (DXA) studies", version="1.0.0", docs_url="/docs", redoc_url="/redoc", openapi_url="/openapi.json" ) -# Подключаем статические файлы +# Mount static files if exists static_dir = Path(__file__).parent / "api/static" if static_dir.exists(): app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") - print(f"✅ Статика подключена: {static_dir}") -else: - print(f"⚠️ Папка static не найдена: {static_dir}") -app.include_router(endpoints) -app.include_router(root) -app.include_router(annotation) +# Global model +dxa_model = None +device = None -# Глобальный обработчик 404 -@app.exception_handler(404) -async def not_found_handler(request, exc): - return JSONResponse( - status_code=404, - content={ - "error": "Not Found", - "message": f"Endpoint {request.url.path} not found", - "available_endpoints": [ - "/", - "/docs", - "/redoc", - "/openapi.json", - "/api/v1/health", - "/api/v1/analyze (POST)", - "/api/v1/switch_mode", - "/api/v1/info", - "/static/" - ] + +def load_model(): + """Load DXA model""" + global dxa_model, device + + if dxa_model is None: + device = get_device() + print(f"Loading DXA model on {device}...") + + try: + dxa_model = create_model( + backbone='resnet18', + pretrained=False, + device=device + ) + dxa_model.load('../models/dxa_model.pth') + dxa_model.model.eval() + print("DXA model loaded successfully") + except Exception as e: + print(f"Error loading model: {e}") + dxa_model = None + + 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" } - ) \ No newline at end of file + + 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) diff --git a/src/segmentators/unet_segmentator.py b/src/segmentators/unet_segmentator.py deleted file mode 100644 index 3030077..0000000 --- a/src/segmentators/unet_segmentator.py +++ /dev/null @@ -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 - } \ No newline at end of file diff --git a/train.py b/train.py deleted file mode 100644 index a3e70ee..0000000 --- a/train.py +++ /dev/null @@ -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() \ No newline at end of file