135 lines
3.9 KiB
Bash
135 lines
3.9 KiB
Bash
#!/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
|