bone_2026/run.sh

135 lines
3.9 KiB
Bash
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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