develop - hack_2026

This commit is contained in:
denis 2026-09-22 15:49:54 +03:00
parent dccade80fd
commit a5dff35ce4
18 changed files with 2220 additions and 1626 deletions

View File

@ -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"]
# Default command
CMD ["python3", "-m", "uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000"]

193
QWEN.md Normal file
View File

@ -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)

View File

@ -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
```

View File

@ -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()

View File

@ -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

134
run.sh Normal file
View File

@ -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

View File

@ -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();

View File

@ -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)
}

View File

@ -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 на хакатоне
# DXA Quality Assessment Configuration
SEGMENTATION_MODEL = "dxa_model.pth"
DATA_PATH = "dataset_hack/"
INPUT_SIZE = (224, 224)
NUM_CLASSES = 2 # quality OK / violation

13
src/dxa/__init__.py Normal file
View File

@ -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'
]

258
src/dxa/dataset.py Normal file
View File

@ -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)}')

261
src/dxa/inference.py Normal file
View File

@ -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()

199
src/dxa/model.py Normal file
View File

@ -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)

199
src/dxa/train.py Normal file
View File

@ -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()

View File

@ -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": [
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",
"/redoc",
"/openapi.json",
"/api/v1/health",
"/api/v1/analyze (POST)",
"/api/v1/switch_mode",
"/api/v1/info",
"/static/"
"/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)

View File

@ -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
View File

@ -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()