#!/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