Harden WebUI training lifecycle, security, and Docker builds
This commit is contained in:
parent
c86c23cd0d
commit
7cd7b01f76
19 changed files with 1799 additions and 226 deletions
600
.agents/PROJECT_CONTEXT.md
Normal file
600
.agents/PROJECT_CONTEXT.md
Normal file
|
|
@ -0,0 +1,600 @@
|
||||||
|
# Контекст проекта YOLO Train WebUI
|
||||||
|
|
||||||
|
Дата актуализации: 2026-07-18
|
||||||
|
|
||||||
|
Этот файл — основной технический контекст проекта для разработчиков и агентов.
|
||||||
|
История найденных и исправленных дефектов находится в `PROJECT_ISSUES.md`.
|
||||||
|
|
||||||
|
## 1. Назначение и границы проекта
|
||||||
|
|
||||||
|
YOLO Train WebUI — локальное веб-приложение для настройки и запуска обучения
|
||||||
|
Ultralytics YOLO. Оно предоставляет форму конфигурации, профили запусков, live-логи,
|
||||||
|
прогресс по эпохам, графики метрик, мягкую остановку и интеграцию с MLflow.
|
||||||
|
|
||||||
|
Поддерживаемые задачи:
|
||||||
|
|
||||||
|
- `detect` — детекция объектов;
|
||||||
|
- `segment` — сегментация;
|
||||||
|
- `classify` — классификация;
|
||||||
|
- `pose` — оценка поз;
|
||||||
|
- `obb` — ориентированные bounding boxes.
|
||||||
|
|
||||||
|
Приложение рассчитано на локального доверенного пользователя и один активный запуск
|
||||||
|
обучения. Это не многопользовательская платформа, не планировщик задач и не сервис
|
||||||
|
хранения датасетов. В нём нет встроенных учётных записей, ролей или аутентификации.
|
||||||
|
|
||||||
|
## 2. Технологии
|
||||||
|
|
||||||
|
| Область | Технология |
|
||||||
|
|---|---|
|
||||||
|
| Backend/API | Python 3.11+, FastAPI, Uvicorn |
|
||||||
|
| Обучение | Ultralytics YOLO, PyTorch |
|
||||||
|
| Эксперименты | MLflow |
|
||||||
|
| Frontend | HTML, CSS, vanilla JavaScript |
|
||||||
|
| Графики | Chart.js из CDN |
|
||||||
|
| Real-time | WebSocket |
|
||||||
|
| Зависимости | `uv`, frozen-набор в `uv.lock` |
|
||||||
|
| Упаковка | Hatchling |
|
||||||
|
| Тесты | pytest, FastAPI TestClient/httpx, Node.js smoke-test |
|
||||||
|
| Контейнер | Docker, Docker Compose |
|
||||||
|
|
||||||
|
Основные зависимости объявлены в `pyproject.toml`: `fastapi`, `uvicorn`,
|
||||||
|
`websockets`, `ultralytics`, `mlflow`. Dev-группа содержит `pytest` и `httpx`.
|
||||||
|
Python package называется `yolo-train-webui`, текущая версия — `0.1.0`; wheel
|
||||||
|
собирается Hatchling только из `src/yolo_webui`.
|
||||||
|
|
||||||
|
`yolo_webui.__init__` публично экспортирует `TrainingConfig`, `TrainingEvent` и
|
||||||
|
`TrainingRunner`.
|
||||||
|
|
||||||
|
## 3. Структура репозитория
|
||||||
|
|
||||||
|
```text
|
||||||
|
.
|
||||||
|
├── .agents/
|
||||||
|
│ ├── PROJECT_CONTEXT.md # этот технический контекст
|
||||||
|
│ └── PROJECT_ISSUES.md # аудит и история исправлений
|
||||||
|
├── src/yolo_webui/
|
||||||
|
│ ├── __init__.py # публичные Python-экспорты
|
||||||
|
│ ├── __main__.py # запуск `python -m yolo_webui`
|
||||||
|
│ ├── app.py # FastAPI, TrainingManager, REST и WebSocket
|
||||||
|
│ ├── config.py # dataclass-конфигурации и валидация
|
||||||
|
│ ├── dataset_splitter.py # detection-style train/val splitter
|
||||||
|
│ ├── subprocess_runner.py # дочерний процесс обучения и stdout-протокол
|
||||||
|
│ ├── trainer.py # Ultralytics callbacks, MLflow, cancellation
|
||||||
|
│ └── static/
|
||||||
|
│ ├── index.html # форма и панель мониторинга
|
||||||
|
│ ├── app.js # browser state, API, WebSocket, Chart.js
|
||||||
|
│ └── style.css # всё визуальное оформление
|
||||||
|
├── tests/
|
||||||
|
│ ├── test_app.py # API и TrainingManager
|
||||||
|
│ ├── test_config.py # конфигурация, безопасность, MLflow env
|
||||||
|
│ ├── test_splitter.py # классы и разбиение датасета
|
||||||
|
│ ├── test_subprocess_runner.py
|
||||||
|
│ ├── test_trainer.py # callbacks и остановка
|
||||||
|
│ ├── test_frontend.py # запуск Node-проверок из pytest
|
||||||
|
│ └── frontend_smoke.js # browser stubs, форма и графики
|
||||||
|
├── Dockerfile
|
||||||
|
├── docker-compose.yml
|
||||||
|
├── pyproject.toml
|
||||||
|
├── uv.lock
|
||||||
|
└── README.md
|
||||||
|
```
|
||||||
|
|
||||||
|
Рабочие каталоги не входят в Git:
|
||||||
|
|
||||||
|
- `datasets/` — локальные датасеты;
|
||||||
|
- `models/` — локальные веса и YAML моделей;
|
||||||
|
- `runs/` — результаты Ultralytics и JSON-профили;
|
||||||
|
- `.yolo-webui/` — сгенерированные split-файлы внутри датасетов;
|
||||||
|
- `mlflow.db`, `mlruns/`, `mlflow/` — локальные данные MLflow;
|
||||||
|
- `.venv/`, кэши Python и pytest.
|
||||||
|
|
||||||
|
## 4. Архитектура во время выполнения
|
||||||
|
|
||||||
|
```text
|
||||||
|
Browser
|
||||||
|
├── HTTP JSON ───────────────┐
|
||||||
|
└── WebSocket /api/ws ───────┤
|
||||||
|
v
|
||||||
|
FastAPI / TrainingManager (основной процесс Uvicorn)
|
||||||
|
├── хранит LiveState и WebSocket-клиентов
|
||||||
|
├── сохраняет профили в runs/sessions
|
||||||
|
└── запускает background thread
|
||||||
|
|
|
||||||
|
v
|
||||||
|
Python subprocess: yolo_webui.subprocess_runner
|
||||||
|
├── читает временный JSON config
|
||||||
|
├── ставит SIGTERM/SIGINT handlers
|
||||||
|
├── создаёт TrainingRunner
|
||||||
|
├── запускает Ultralytics YOLO.train()
|
||||||
|
└── пишет события, логи и результат в stdout
|
||||||
|
|
|
||||||
|
├── dataset / generated split
|
||||||
|
├── models / official model download
|
||||||
|
├── runs / training artifacts
|
||||||
|
└── MLflow storage
|
||||||
|
```
|
||||||
|
|
||||||
|
Изоляция обучения в subprocess нужна, чтобы тяжёлый Ultralytics/PyTorch не блокировал
|
||||||
|
ASGI event loop, stdout можно было транслировать в браузер, а зависший запуск —
|
||||||
|
принудительно завершить.
|
||||||
|
|
||||||
|
Важная деталь: `TrainingRunner` используется в двух процессах.
|
||||||
|
|
||||||
|
- В родительском `TrainingManager` он хранит ссылку на subprocess и управляет
|
||||||
|
сигналами остановки.
|
||||||
|
- В дочернем процессе отдельный экземпляр владеет моделью Ultralytics и выставляет
|
||||||
|
`trainer.stop = True`.
|
||||||
|
|
||||||
|
## 5. Точки входа и запуск
|
||||||
|
|
||||||
|
CLI entry point из `pyproject.toml`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run yolo-train-webui
|
||||||
|
```
|
||||||
|
|
||||||
|
Альтернативный модульный запуск:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run -m yolo_webui
|
||||||
|
```
|
||||||
|
|
||||||
|
Оба варианта вызывают `yolo_webui.app:main`. CLI принимает:
|
||||||
|
|
||||||
|
- `--host`, default `127.0.0.1`;
|
||||||
|
- `--port`, default `8000`.
|
||||||
|
|
||||||
|
Перед запуском Uvicorn устанавливается `MPLBACKEND=Agg`. В Docker приложение слушает
|
||||||
|
`0.0.0.0:8000` внутри контейнера, но Compose публикует его только как
|
||||||
|
`127.0.0.1:8000` на host.
|
||||||
|
|
||||||
|
## 6. Backend и состояние обучения
|
||||||
|
|
||||||
|
### `LiveState`
|
||||||
|
|
||||||
|
Глобальный `TrainingManager` хранит единственное состояние:
|
||||||
|
|
||||||
|
- `status`;
|
||||||
|
- текущую и общую эпохи;
|
||||||
|
- список логов;
|
||||||
|
- историю метрик;
|
||||||
|
- `output_dir`;
|
||||||
|
- `stop_requested`;
|
||||||
|
- `last_event_kind` для классификации результата.
|
||||||
|
|
||||||
|
Состояния:
|
||||||
|
|
||||||
|
```text
|
||||||
|
idle
|
||||||
|
└── start -> preparing
|
||||||
|
├── event started -> training
|
||||||
|
│ ├── normal exit 0 -> succeeded
|
||||||
|
│ ├── stop -> stopping -> cancelled
|
||||||
|
│ └── error -> failed
|
||||||
|
├── stop -> stopping -> cancelled/failed
|
||||||
|
└── setup error -> failed
|
||||||
|
```
|
||||||
|
|
||||||
|
`finished` больше не создаётся backend-ом; frontend понимает его только для
|
||||||
|
совместимости со старым состоянием. Новый запуск разрешён лишь когда нет активного
|
||||||
|
статуса и предыдущий background thread уже завершён.
|
||||||
|
|
||||||
|
### Потоки и lock
|
||||||
|
|
||||||
|
`TrainingManager._lock` защищает `LiveState`, ссылку на thread и набор WebSocket.
|
||||||
|
Нельзя выполнять `broadcast()` внутри `with self._lock`: broadcast сам читает
|
||||||
|
защищённые данные, и повторный захват обычного `threading.Lock` вызовет deadlock.
|
||||||
|
|
||||||
|
WebSocket привязывается к event loop Uvicorn при подключении. Вызовы broadcast из
|
||||||
|
фонового потока передаются через `asyncio.run_coroutine_threadsafe()`. Отправки
|
||||||
|
сериализуются `asyncio.Lock`; failed-клиенты логируются и удаляются.
|
||||||
|
|
||||||
|
### Запуск subprocess
|
||||||
|
|
||||||
|
`TrainingManager._run_subprocess()`:
|
||||||
|
|
||||||
|
1. сериализует `TrainingConfig.to_dict()` во временный JSON;
|
||||||
|
2. запускает текущий interpreter с `-u -m yolo_webui.subprocess_runner`;
|
||||||
|
3. объединяет stderr со stdout;
|
||||||
|
4. читает поток посимвольно, различая `\r` и `\n` для progress-строк;
|
||||||
|
5. обновляет состояние и транслирует события;
|
||||||
|
6. ждёт return code, удаляет временный JSON и очищает ссылку на процесс.
|
||||||
|
|
||||||
|
### Внутренний stdout-протокол
|
||||||
|
|
||||||
|
Дочерний процесс печатает специальные маркеры:
|
||||||
|
|
||||||
|
```text
|
||||||
|
__YOLO_WEBUI_READY__
|
||||||
|
__YOLO_WEBUI_EVENT__:{"kind":"epoch","message":"...","epoch":1,"total_epochs":100}
|
||||||
|
__YOLO_WEBUI_RESULT__:/absolute/path/to/run
|
||||||
|
```
|
||||||
|
|
||||||
|
- `READY` означает, что signal handlers уже установлены и отложенный stop можно
|
||||||
|
безопасно доставить.
|
||||||
|
- `EVENT` несёт `kind`, `message`, `epoch`, `total_epochs`.
|
||||||
|
- `RESULT` передаёт каталог результатов.
|
||||||
|
- Любая другая строка считается обычным логом.
|
||||||
|
|
||||||
|
События от `TrainingRunner`: `info`, `started`, `epoch`, `success`, `cancelled`,
|
||||||
|
`warning`. Исключение выводится traceback-ом в stderr/stdout и даёт return code `1`.
|
||||||
|
|
||||||
|
### Классификация завершения
|
||||||
|
|
||||||
|
Backend различает:
|
||||||
|
|
||||||
|
- `succeeded` — return code `0` без подтверждённой остановки;
|
||||||
|
- `cancelled` — был stop и subprocess завершился с `0`, прислал `cancelled` или был
|
||||||
|
убит force-stop таймером;
|
||||||
|
- `failed` — ненулевой код без подтверждённой отмены, в том числе реальная ошибка,
|
||||||
|
случившаяся после нажатия Stop.
|
||||||
|
|
||||||
|
## 7. Остановка обучения
|
||||||
|
|
||||||
|
Остановка кооперативная и двухуровневая:
|
||||||
|
|
||||||
|
1. `POST /api/train/stop` ставит `stop_requested` и статус `stopping`.
|
||||||
|
2. Если subprocess ещё не прислал `READY`, запрос сохраняется.
|
||||||
|
3. После `READY` родитель отправляет `SIGTERM`.
|
||||||
|
4. Signal handler дочернего процесса вызывает `TrainingRunner.request_stop()`.
|
||||||
|
5. Runner выставляет `trainer.stop = True` сразу либо в ближайшем callback.
|
||||||
|
6. Ultralytics штатно завершает callbacks и сохранение результатов.
|
||||||
|
7. Если subprocess не завершился за 30 секунд, parent вызывает `kill()`.
|
||||||
|
|
||||||
|
`prepare_run()` перед каждым новым запуском очищает stop-флаги и старый таймер.
|
||||||
|
|
||||||
|
## 8. REST и WebSocket API
|
||||||
|
|
||||||
|
| Метод | Путь | Назначение |
|
||||||
|
|---|---|---|
|
||||||
|
| GET | `/` | Возвращает `static/index.html` |
|
||||||
|
| GET | `/static/*` | CSS и JavaScript |
|
||||||
|
| GET | `/api/config/defaults` | Полный default `TrainingConfig` |
|
||||||
|
| GET | `/api/datasets` | Верхнеуровневые каталоги и YAML из `datasets/` |
|
||||||
|
| GET | `/api/models` | Верхнеуровневые `.pt/.pth/.yaml/.yml` из `models/` |
|
||||||
|
| GET | `/api/sessions` | Список профилей без `last_run` |
|
||||||
|
| GET | `/api/sessions/{name}` | Загрузить профиль; `last_run` читать можно |
|
||||||
|
| POST | `/api/sessions/{name}` | Сохранить произвольный JSON профиля |
|
||||||
|
| DELETE | `/api/sessions/{name}` | Удалить профиль |
|
||||||
|
| GET | `/api/train/status` | Текущее состояние, эпохи, результат, метрики, логи |
|
||||||
|
| POST | `/api/train/start` | Провалидировать config, сохранить `last_run`, запустить |
|
||||||
|
| POST | `/api/train/stop` | Запросить остановку |
|
||||||
|
| WS | `/api/ws` | Init-снимок и live-события |
|
||||||
|
|
||||||
|
FastAPI также оставляет включёнными стандартные OpenAPI endpoints: `/docs`,
|
||||||
|
`/redoc`, `/openapi.json`.
|
||||||
|
|
||||||
|
WebSocket server → browser сообщения:
|
||||||
|
|
||||||
|
- `init`: полный snapshot состояния, логов и метрик при подключении;
|
||||||
|
- `status`: новое состояние и опциональный `output_dir`;
|
||||||
|
- `log`: `message` и `level`;
|
||||||
|
- `progress`: эпоха, total, извлечённые метрики и сообщение.
|
||||||
|
|
||||||
|
Browser → server сообщения не используются; endpoint только читает и отбрасывает их,
|
||||||
|
поддерживая соединение открытым.
|
||||||
|
|
||||||
|
Профили хранятся в `runs/sessions/{name}.json`. Имя: 1–64 символа из латинских
|
||||||
|
букв, цифр, `_`, `-`. `last_run` зарезервирован для автосохранения при старте: его
|
||||||
|
можно прочитать, но нельзя создать или удалить через profile endpoints.
|
||||||
|
|
||||||
|
## 9. Конфигурация обучения
|
||||||
|
|
||||||
|
### `TrainingConfig`
|
||||||
|
|
||||||
|
| Поле | Default | Передача в Ultralytics |
|
||||||
|
|---|---:|---|
|
||||||
|
| `dataset` | обязательно; API default `coco8.yaml` | `data` |
|
||||||
|
| `model` | обязательно; API default `yolo11n.pt` | аргумент конструктора `YOLO()` |
|
||||||
|
| `task` | `detect` | аргумент конструктора `YOLO()` |
|
||||||
|
| `epochs` | `100` | `epochs` |
|
||||||
|
| `image_size` | `640` | `imgsz` |
|
||||||
|
| `batch_size` | `16` | `batch` |
|
||||||
|
| `device` | пусто | `device`, только если задано |
|
||||||
|
| `workers` | `8` | `workers`; `0` допустим |
|
||||||
|
| `patience` | `100` | `patience`; `0` допустим |
|
||||||
|
| `project` | `runs/train` | `project` |
|
||||||
|
| `run_name` | пусто | `name`, только если задано |
|
||||||
|
| `augmentation` | включена | набор augmentation kwargs |
|
||||||
|
| `mlflow` | включён | Ultralytics setting и env |
|
||||||
|
| `split` | выключен | preprocessing до `YOLO.train()` |
|
||||||
|
|
||||||
|
Основная валидация:
|
||||||
|
|
||||||
|
- `epochs >= 1`, `image_size >= 32`;
|
||||||
|
- batch положительный или `-1`; `0` и значения `< -1` запрещены;
|
||||||
|
- `workers >= 0`, `patience >= 0`;
|
||||||
|
- задача входит в фиксированный список;
|
||||||
|
- detection-style auto split запрещён для `classify`;
|
||||||
|
- dataset/model/project проходят security path validation.
|
||||||
|
|
||||||
|
Модель без `/` или `\` считается именем и резолвится как `models/{name}`.
|
||||||
|
|
||||||
|
### `DatasetSplitConfig`
|
||||||
|
|
||||||
|
- `enabled=False`;
|
||||||
|
- `train_ratio=0.8`, допустимо `0.1…0.95`;
|
||||||
|
- `classes_path=""`, пустое значение включает автопоиск.
|
||||||
|
|
||||||
|
### `MlflowConfig`
|
||||||
|
|
||||||
|
- `enabled=True`;
|
||||||
|
- `tracking_uri="sqlite:///mlflow.db"`;
|
||||||
|
- `experiment_name="yolo-webui"`;
|
||||||
|
- `run_name=""`.
|
||||||
|
|
||||||
|
Если MLflow включён, tracking URI и experiment name не могут быть пустыми.
|
||||||
|
`mlflow_environment()` временно выставляет:
|
||||||
|
|
||||||
|
- `MLFLOW_TRACKING_URI`;
|
||||||
|
- `MLFLOW_EXPERIMENT_NAME`;
|
||||||
|
- `MLFLOW_RUN`;
|
||||||
|
- `MLFLOW_KEEP_RUN_ACTIVE=False`.
|
||||||
|
|
||||||
|
После обучения предыдущие значения окружения восстанавливаются. В Ultralytics
|
||||||
|
глобальная настройка `mlflow` включается/выключается через `settings.update()`.
|
||||||
|
|
||||||
|
### `AugmentationConfig`
|
||||||
|
|
||||||
|
Default-параметры:
|
||||||
|
|
||||||
|
```text
|
||||||
|
hsv_h=0.015 hsv_s=0.7 hsv_v=0.4
|
||||||
|
degrees=0.0 translate=0.1 scale=0.5
|
||||||
|
shear=0.0 perspective=0.0
|
||||||
|
flipud=0.0 fliplr=0.5 bgr=0.0
|
||||||
|
mosaic=1.0 mixup=0.0 cutmix=0.0
|
||||||
|
copy_paste=0.0 erasing=0.4 close_mosaic=10
|
||||||
|
copy_paste_mode=flip
|
||||||
|
auto_augment=randaugment
|
||||||
|
```
|
||||||
|
|
||||||
|
Вероятности и доли валидируются в диапазоне `0…1`; `degrees`, `shear` и
|
||||||
|
`close_mosaic` не могут быть отрицательными. Режимы copy-paste: `flip`, `mixup`.
|
||||||
|
Политики AutoAugment: `randaugment`, `autoaugment`, `augmix`. Если augmentation
|
||||||
|
выключена, эти kwargs вообще не передаются в Ultralytics.
|
||||||
|
|
||||||
|
## 10. Работа с датасетами
|
||||||
|
|
||||||
|
Auto split предназначен только для detection-style структуры:
|
||||||
|
|
||||||
|
```text
|
||||||
|
dataset/
|
||||||
|
├── images/
|
||||||
|
│ └── **/*.{jpg,jpeg,png,bmp,webp,tif,tiff}
|
||||||
|
└── labels/
|
||||||
|
└── **/*.txt
|
||||||
|
```
|
||||||
|
|
||||||
|
Изображения и labels могут быть вложенными. Для split нужно минимум два изображения.
|
||||||
|
Shuffle детерминирован seed-ом `42`; train и val всегда получают минимум по одному
|
||||||
|
изображению.
|
||||||
|
|
||||||
|
Порядок определения классов:
|
||||||
|
|
||||||
|
1. явно заданный `classes_path` — авторитетный, без fallback при ошибке;
|
||||||
|
2. корневой `classes.txt`;
|
||||||
|
3. `labels/classes.txt`;
|
||||||
|
4. первый по имени корневой `.yaml/.yml` с полем `names`;
|
||||||
|
5. вывод диапазона `0…max_id` из всех label-файлов с именами `class_N`.
|
||||||
|
|
||||||
|
Поддерживаются text, YAML list и YAML dict. ID должны быть целыми,
|
||||||
|
неповторяющимися и последовательными от `0`; пустые имена запрещены.
|
||||||
|
|
||||||
|
Каждый split создаётся эксклюзивно:
|
||||||
|
|
||||||
|
```text
|
||||||
|
dataset/.yolo-webui/splits/{uuid}/
|
||||||
|
├── train.txt # абсолютные пути изображений
|
||||||
|
├── val.txt # абсолютные пути изображений
|
||||||
|
└── dataset.yaml
|
||||||
|
```
|
||||||
|
|
||||||
|
Итоговый YAML сохраняет дополнительные ключи исходного YAML, например `kpt_shape`
|
||||||
|
и `flip_idx`, но перезаписывает `path`, `train`, `val` и `names`. Проверенный
|
||||||
|
результат `read_classes()` всегда авторитетен.
|
||||||
|
|
||||||
|
Для `classify` auto split отключён: пользователь должен предоставить готовую
|
||||||
|
структуру `train/`, `val/` или `test/` с подкаталогами классов.
|
||||||
|
|
||||||
|
## 11. TrainingRunner и метрики
|
||||||
|
|
||||||
|
Перед обучением Runner:
|
||||||
|
|
||||||
|
1. повторно валидирует config;
|
||||||
|
2. при необходимости создаёт split и заменяет `data` на generated YAML;
|
||||||
|
3. импортирует Ultralytics;
|
||||||
|
4. включает/выключает MLflow integration;
|
||||||
|
5. создаёт `YOLO(config.resolved_model, task=config.task)`;
|
||||||
|
6. подключает callbacks `on_train_start`, `on_train_epoch_end`, `on_train_end`;
|
||||||
|
7. вызывает `model.train(**config.train_kwargs())`.
|
||||||
|
|
||||||
|
Epoch callback берёт numeric metrics из `trainer.metrics`, форматирует максимум три
|
||||||
|
первых значения и отправляет их в текстовом сообщении. Parent разбирает пары
|
||||||
|
`key=value`, поэтому live chart сейчас показывает не более трёх метрик на эпоху.
|
||||||
|
|
||||||
|
Если `trainer.save_dir` существует, его путь передаётся parent-у как результат.
|
||||||
|
|
||||||
|
Restricted checkpoint loading принудительно включён:
|
||||||
|
|
||||||
|
```text
|
||||||
|
ULTRALYTICS_SAFE_LOAD=1
|
||||||
|
```
|
||||||
|
|
||||||
|
## 12. Frontend
|
||||||
|
|
||||||
|
Frontend не имеет сборщика и framework: `index.html`, `style.css` и `app.js`
|
||||||
|
отдаются FastAPI как статические файлы. Chart.js загружается с jsDelivr CDN.
|
||||||
|
|
||||||
|
Левая панель содержит профили и вкладки:
|
||||||
|
|
||||||
|
- «Основное» — task, model, dataset, auto split;
|
||||||
|
- «Обучение» — epochs, image size, batch, device, workers, patience, output;
|
||||||
|
- «Аугментация» — все поля `AugmentationConfig`;
|
||||||
|
- «MLflow» — enabled, tracking URI, experiment и run name.
|
||||||
|
|
||||||
|
Правая панель содержит статус, timer, progress bar, Start/Stop, live chart и журнал.
|
||||||
|
|
||||||
|
Browser state:
|
||||||
|
|
||||||
|
- активная вкладка хранится в `localStorage.active_tab`;
|
||||||
|
- выбранный профиль — `localStorage.selected_profile`;
|
||||||
|
- черновик формы — `localStorage.draft_config`;
|
||||||
|
- при старте загрузки приоритет: draft → `last_run` → API defaults;
|
||||||
|
- список датасетов и моделей запрашивается у backend;
|
||||||
|
- стандартные модели YOLO11 выбираются динамически по task;
|
||||||
|
- при disconnect WebSocket переподключается через 5 секунд;
|
||||||
|
- `init` восстанавливает status, логи, progress и историю графика.
|
||||||
|
|
||||||
|
Числа читаются через `Number.parseInt/parseFloat` и проверку `Number.isNaN`. Нельзя
|
||||||
|
заменять это на `value || default`: допустимые `0` для workers, patience,
|
||||||
|
close_mosaic и augmentation-параметров должны сохраняться.
|
||||||
|
|
||||||
|
Chart datasets создаются по фактически пришедшим ключам. Новая метрика может
|
||||||
|
появиться на поздней эпохе; пропущенные точки заполняются `null`, чтобы серии не
|
||||||
|
сдвигались относительно labels.
|
||||||
|
|
||||||
|
## 13. Безопасность и доверенная модель
|
||||||
|
|
||||||
|
Приложение не имеет аутентификации. Безопасность по умолчанию строится на локальной
|
||||||
|
публикации и ограничении файловых путей.
|
||||||
|
|
||||||
|
Default доверенные корни:
|
||||||
|
|
||||||
|
| Назначение | Корни |
|
||||||
|
|---|---|
|
||||||
|
| Dataset и classes | `./datasets` |
|
||||||
|
| Model/checkpoint | `./models`, `./runs` |
|
||||||
|
| Training output | `./runs` |
|
||||||
|
|
||||||
|
Дополнительные корни перечисляются через системный `os.pathsep`:
|
||||||
|
|
||||||
|
- `YOLO_WEBUI_DATA_ROOTS`;
|
||||||
|
- `YOLO_WEBUI_MODEL_ROOTS`;
|
||||||
|
- `YOLO_WEBUI_RUN_ROOTS`.
|
||||||
|
|
||||||
|
Проверка запрещает URL, нормализует путь через `resolve(strict=False)` и проверяет
|
||||||
|
принадлежность корню, включая существующие symlink-компоненты. Безопасные bare
|
||||||
|
identifiers разрешены для официальных имён, но существующий одноимённый файл вне
|
||||||
|
доверенного root отвергается. Model-файлы ограничены расширениями `.pt`, `.pth`,
|
||||||
|
`.yaml`, `.yml`.
|
||||||
|
|
||||||
|
Compose публикует только `127.0.0.1:8000:8000`. Для доступа из сети обязателен
|
||||||
|
аутентифицирующий reverse proxy и явная оценка риска: API может запускать тяжёлое
|
||||||
|
обучение, останавливать его и управлять профилями.
|
||||||
|
|
||||||
|
## 14. Docker
|
||||||
|
|
||||||
|
Dockerfile:
|
||||||
|
|
||||||
|
- основан на `python:3.11-slim`;
|
||||||
|
- устанавливает системные библиотеки для OpenCV/PyTorch/Ultralytics;
|
||||||
|
- фиксирует `uv==0.10.6`;
|
||||||
|
- копирует `pyproject.toml`, `uv.lock`, README;
|
||||||
|
- выполняет `uv sync --locked --no-dev` в `/opt/venv`;
|
||||||
|
- включает `ULTRALYTICS_SAFE_LOAD=1`;
|
||||||
|
- запускает `yolo-train-webui --host 0.0.0.0 --port 8000`.
|
||||||
|
|
||||||
|
Compose монтирует:
|
||||||
|
|
||||||
|
```text
|
||||||
|
./datasets -> /workspace/datasets
|
||||||
|
./runs -> /workspace/runs
|
||||||
|
./models -> /workspace/models
|
||||||
|
./models/.config -> /root/.config/Ultralytics
|
||||||
|
```
|
||||||
|
|
||||||
|
Порт 8000 опубликован только на loopback. Порт 5000 объявлен в image, но Compose не
|
||||||
|
запускает и не публикует MLflow UI. GPU reservation оставлена как закомментированный
|
||||||
|
пример для NVIDIA/Linux.
|
||||||
|
|
||||||
|
## 15. Тесты и проверки
|
||||||
|
|
||||||
|
Текущий regression suite содержит 49 pytest-тестов.
|
||||||
|
|
||||||
|
- `test_app.py`: defaults/status, profiles, deadlock, background WebSocket loop,
|
||||||
|
финальные состояния.
|
||||||
|
- `test_config.py`: kwargs, validation, zero-compatible параметры, MLflow env,
|
||||||
|
security roots и URL.
|
||||||
|
- `test_splitter.py`: форматы классов, приоритеты, nested data, уникальные outputs,
|
||||||
|
сохранение YAML metadata.
|
||||||
|
- `test_subprocess_runner.py`: return codes, READY/RESULT и traceback.
|
||||||
|
- `test_trainer.py`: Ultralytics callbacks, metrics, ранний stop, сигналы.
|
||||||
|
- `test_frontend.py` + `frontend_smoke.js`: syntax, сохранение нулей, динамические
|
||||||
|
chart series в fake browser environment.
|
||||||
|
|
||||||
|
Основные команды:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv sync --locked
|
||||||
|
uv run pytest -q
|
||||||
|
uv run python -m compileall -q src tests
|
||||||
|
node --check src/yolo_webui/static/app.js
|
||||||
|
node tests/frontend_smoke.js
|
||||||
|
uv lock --check
|
||||||
|
docker compose config
|
||||||
|
git diff --check
|
||||||
|
```
|
||||||
|
|
||||||
|
Python-команды проекта следует выполнять через `uv run`, чтобы использовать
|
||||||
|
зафиксированное окружение.
|
||||||
|
|
||||||
|
## 16. Согласованное изменение проекта
|
||||||
|
|
||||||
|
При добавлении или переименовании config-поля обычно нужно изменить вместе:
|
||||||
|
|
||||||
|
1. dataclass, default, validation и `train_kwargs()` в `config.py`;
|
||||||
|
2. `TrainingConfig.from_dict()`;
|
||||||
|
3. поле в `static/index.html`;
|
||||||
|
4. чтение в `getFormConfig()` и восстановление в `applyConfig()` в `app.js`;
|
||||||
|
5. backend/frontend regression tests;
|
||||||
|
6. README и этот контекст, если меняется пользовательский контракт.
|
||||||
|
|
||||||
|
При добавлении нового состояния обучения нужно обновить:
|
||||||
|
|
||||||
|
1. backend state machine и финальную классификацию;
|
||||||
|
2. WebSocket status payload;
|
||||||
|
3. `updateUIStatus()`;
|
||||||
|
4. CSS-селекторы `status-*`;
|
||||||
|
5. тесты переходов и reconnect snapshot.
|
||||||
|
|
||||||
|
При изменении subprocess-протокола синхронно меняются `subprocess_runner.py` и parser
|
||||||
|
в `TrainingManager._handle_subprocess_line()`. Префиксы протокола нельзя печатать в
|
||||||
|
обычных логах.
|
||||||
|
|
||||||
|
Критические инварианты:
|
||||||
|
|
||||||
|
- не вызывать WebSocket send из нового или чужого event loop;
|
||||||
|
- не вызывать `broadcast()` под `TrainingManager._lock`;
|
||||||
|
- не объединять `succeeded`, `cancelled`, `failed` в общий `finished`;
|
||||||
|
- не использовать JS truthiness для числовых полей;
|
||||||
|
- явно указанный classes-файл всегда авторитетен;
|
||||||
|
- не создавать split поверх пользовательских файлов;
|
||||||
|
- не снимать `--locked` с Docker/CI установки;
|
||||||
|
- не расширять сетевую публикацию без аутентификации;
|
||||||
|
- сохранять traceback и ошибки доставки в наблюдаемых логах.
|
||||||
|
|
||||||
|
## 17. Текущие ограничения
|
||||||
|
|
||||||
|
- Только один активный training job и один глобальный in-memory `LiveState`.
|
||||||
|
- После перезапуска server live state теряется; сохраняются лишь JSON-профили,
|
||||||
|
`last_run`, training artifacts и MLflow data.
|
||||||
|
- Нет очереди, scheduler, истории runs API, upload API и файлового браузера.
|
||||||
|
- Нет встроенной аутентификации и multi-user isolation.
|
||||||
|
- Discovery просматривает только верхний уровень `datasets/` и `models/`.
|
||||||
|
- Live chart зависит от внешнего Chart.js CDN.
|
||||||
|
- В график попадают максимум три numeric metrics, выбранные callback-ом.
|
||||||
|
- Реальное длительное YOLO-обучение и Docker image build не входят в быстрый test
|
||||||
|
suite; unit-тесты подменяют Ultralytics и subprocess там, где это возможно.
|
||||||
|
- В репозитории нет `.dockerignore` и CI-конфигурации; Docker build context зависит
|
||||||
|
от содержимого рабочей копии.
|
||||||
|
- FastAPI TestClient выдаёт deprecation warning для текущей связки Starlette/httpx;
|
||||||
|
тесты при этом проходят.
|
||||||
|
|
||||||
|
Отдельного файла лицензии проекта в репозитории нет. README напоминает, что
|
||||||
|
Ultralytics распространяется по AGPL-3.0 и предлагает отдельно проверить условия
|
||||||
|
Enterprise-лицензии для закрытого коммерческого использования.
|
||||||
|
|
||||||
|
Перед работой с известными дефектами сверяйтесь с `PROJECT_ISSUES.md`: на дату этого
|
||||||
|
контекста перечисленные там 10 проблем исправлены.
|
||||||
|
|
@ -1,73 +1,105 @@
|
||||||
# Исправленные проблемы проекта YOLO Train TUI
|
# Исправленные проблемы проекта YOLO Train WebUI
|
||||||
|
|
||||||
Дата исправления и повторной проверки: 2026-07-17
|
Дата исправления и повторной проверки: 2026-07-18
|
||||||
|
|
||||||
## Итог
|
## Итог
|
||||||
|
|
||||||
Все 11 ранее зафиксированных дефектов исправлены и покрыты регрессионными
|
Все 10 дефектов аудита от 2026-07-17 исправлены. WebUI снова проходит
|
||||||
проверками.
|
синтаксическую проверку, серверный обработчик события начала обучения не зависает,
|
||||||
|
WebSocket-сообщения отправляются в event loop ASGI-сервера, а успешное завершение,
|
||||||
|
отмена и ошибка представлены отдельными состояниями.
|
||||||
|
|
||||||
|
Регрессионный набор расширен с 37 до 49 тестов.
|
||||||
|
|
||||||
| ID | Приоритет | Статус | Исправление |
|
| ID | Приоритет | Статус | Исправление |
|
||||||
|---|---|---|---|
|
|---|---|---|---|
|
||||||
| BUG-001 | Критический | Исправлено | `subprocess_runner.main()` возвращает код, а `SystemExit` создаётся только снаружи обрабатывающего блока |
|
| BUG-001 | Критический | Исправлено | Закрыт `try/catch`, удалено повторное объявление `configForm`, добавлен `node --check` в тесты |
|
||||||
| BUG-002 | Высокий | Исправлено | Перед каждым запуском `prepare_run()` сбрасывает состояние остановки |
|
| BUG-002 | Критический | Исправлено | Status broadcast вынесен за пределы `threading.Lock`; добавлен тест на отсутствие deadlock |
|
||||||
| BUG-003 | Высокий | Исправлено | Родитель отправляет кооперативный сигнал; принудительный `kill()` используется только после таймаута |
|
| SEC-001 | Критический при сетевой публикации | Исправлено | Compose публикует loopback, URL запрещены, пути ограничены доверенными корнями, restricted checkpoint loading включён |
|
||||||
| BUG-004 | Высокий | Исправлено | Запрос, сделанный до готовности subprocess, сохраняется и доставляется после маркера `READY` |
|
| BUG-003 | Высокий | Исправлено | Все WebSocket send выполняются в ASGI loop через `run_coroutine_threadsafe`; ошибки логируются, сломанные сокеты удаляются |
|
||||||
| BUG-005 | Высокий | Исправлено | Detection-style авторазбиение запрещено для `classify` в UI и конфигурации |
|
| BUG-004 | Высокий | Исправлено | Введены состояния `succeeded`, `cancelled`, `failed`; ошибка после stop больше не маскируется как отмена |
|
||||||
| BUG-006 | Средний | Исправлено | `.yaml`/`.yml` разбираются через `yaml.safe_load()`, поле `names` валидируется |
|
| DOC-001 | Высокий | Исправлено | README полностью обновлён для WebUI, актуальных CLI-команд, Docker и модели безопасности |
|
||||||
| BUG-007 | Средний | Исправлено | Датасет с одним изображением отклоняется с понятной ошибкой |
|
| BUG-005 | Средний | Исправлено | Провалидированные классы всегда записываются в итоговый YAML и имеют приоритет над случайным корневым YAML |
|
||||||
| BUG-008 | Средний | Исправлено | Изображения и метки ищутся рекурсивно с сохранением вложенных путей |
|
| BUG-006 | Средний | Исправлено | Числа разбираются с проверкой `Number.isNaN`; нули сохраняются при чтении и восстановлении формы |
|
||||||
| BUG-009 | Средний | Исправлено | Каждый результат создаётся в уникальном `.yolo-tui/splits/<id>` без перезаписи пользовательского `split/` |
|
| BUG-007 | Средний | Исправлено | Серии графика добавляются динамически и выравниваются по эпохам, включая новые ключи метрик |
|
||||||
| BUG-010 | Низкий | Исправлено | Traceback выводится в журнал TUI; абсолютный путь другого пользователя удалён |
|
| BUILD-001 | Средний | Исправлено | Docker устанавливает frozen-набор из `uv.lock`; версия `uv` также зафиксирована |
|
||||||
| BUG-011 | Низкий | Исправлено | Явно указанный отсутствующий или некорректный файл классов вызывает точную ошибку без fallback |
|
|
||||||
|
|
||||||
## Жизненный цикл обучения
|
## Жизненный цикл и WebSocket
|
||||||
|
|
||||||
- Дочерний процесс устанавливает обработчики остановки и только затем печатает
|
- Событие `started` меняет состояние под lock, но отправляет статус только после
|
||||||
`__YOLO_TUI_READY__`.
|
освобождения lock.
|
||||||
- Если пользователь нажал «Остановить» раньше, родитель запоминает запрос и
|
- Event loop запоминается при подключении WebSocket. Вызовы из фонового потока
|
||||||
отправляет его после получения маркера готовности.
|
передаются в него через `asyncio.run_coroutine_threadsafe()`.
|
||||||
- Дочерний `TrainingRunner` устанавливает `trainer.stop = True`; Ultralytics
|
- Отправки сериализуются `asyncio.Lock`, поэтому сообщения одного запуска сохраняют
|
||||||
останавливается между пакетами данных, затем выполняет штатную финализацию и
|
порядок. Ошибка доставки попадает в журнал, а нерабочий клиент удаляется.
|
||||||
завершающие callbacks.
|
- Финальная классификация учитывает return code, stop-флаг, последнее
|
||||||
- Если процесс не завершился за 30 секунд, используется принудительный fallback.
|
структурированное событие и факт принудительной остановки.
|
||||||
- После завершения ссылка на subprocess и таймер очищаются; перед следующим
|
- Штатная кооперативная остановка даёт `cancelled`; ненулевой код после stop без
|
||||||
запуском флаг остановки сбрасывается.
|
подтверждённой отмены даёт `failed`.
|
||||||
|
|
||||||
## Работа с датасетами
|
## Безопасность
|
||||||
|
|
||||||
- Текущий splitter предназначен для задач `detect`, `segment`, `pose` и `obb`
|
- `docker-compose.yml` публикует `127.0.0.1:8000:8000`.
|
||||||
со структурой `images/` + `labels/`.
|
- Dataset, model и project не принимают URL.
|
||||||
- Для `classify` требуется готовый каталог с `train`/`val` и подкаталогами
|
- Локальные пути ограничены `datasets`, `models` и `runs`; дополнительные доверенные
|
||||||
классов; несовместимый переключатель в UI отключён.
|
корни задаются переменными `YOLO_WEBUI_DATA_ROOTS`,
|
||||||
- Поддерживаются `classes.txt`, `.yaml` и `.yml`; YAML может хранить `names` как
|
`YOLO_WEBUI_MODEL_ROOTS`, `YOLO_WEBUI_RUN_ROOTS`.
|
||||||
список или словарь с последовательными ID от 0.
|
- Проверка использует разрешённые абсолютные пути после `resolve()`, поэтому
|
||||||
- Явный путь к классам считается обязательным и не заменяется автопоиском при
|
symlink/`..` не позволяют выйти из доверенного корня.
|
||||||
опечатке или ошибке формата.
|
- Имена профилей валидируются на сервере, а `last_run` нельзя перезаписать через
|
||||||
- Split требует минимум два изображения и рекурсивно обрабатывает вложенные
|
публичный endpoint профилей.
|
||||||
каталоги.
|
- `ULTRALYTICS_SAFE_LOAD=1` включён и в Python-процессе, и в Docker-образе.
|
||||||
- Файлы каждого запуска создаются эксклюзивно в отдельном управляемом каталоге.
|
- Для намеренной удалённой публикации по-прежнему нужен аутентифицирующий reverse
|
||||||
|
proxy; это явно указано в README.
|
||||||
|
|
||||||
## Проверка
|
## Frontend
|
||||||
|
|
||||||
Выполнены команды:
|
- `app.js` снова является валидным JavaScript.
|
||||||
|
- `workers=0`, `patience=0` и `close_mosaic=0` проходят полный цикл
|
||||||
|
form → JSON → localStorage → form без замены default-значениями.
|
||||||
|
- График создаёт dataset при первом ключе метрики и добавляет новые серии в следующих
|
||||||
|
эпохах. Пропущенные значения дополняются `null`, поэтому точки не сдвигаются.
|
||||||
|
- UI и CSS отдельно отображают `succeeded`, `cancelled` и `failed`; старый
|
||||||
|
`finished` оставлен только как frontend-совместимость.
|
||||||
|
|
||||||
|
## Датасеты, Docker и документация
|
||||||
|
|
||||||
|
- Результат `read_classes()` безусловно становится `dataset_data["names"]`, сохраняя
|
||||||
|
при этом остальные ключи выбранного YAML (`kpt_shape`, `flip_idx` и другие).
|
||||||
|
- Docker копирует `pyproject.toml` вместе с `uv.lock` и выполняет
|
||||||
|
`uv sync --locked --no-dev`; обход lock-файла удалён.
|
||||||
|
- README описывает `uv run yolo-train-webui`, `uv run -m yolo_webui`, Compose,
|
||||||
|
структуру датасетов, MLflow и ограничения доверенных путей.
|
||||||
|
|
||||||
|
## Добавленные регрессионные проверки
|
||||||
|
|
||||||
|
Тесты теперь покрывают:
|
||||||
|
|
||||||
|
1. синтаксис browser JavaScript;
|
||||||
|
2. сохранение допустимых нулей и динамические серии Chart.js в Node smoke-test;
|
||||||
|
3. отсутствие deadlock на событии `started`;
|
||||||
|
4. доставку сообщения из background thread в loop WebSocket-сервера;
|
||||||
|
5. различие `succeeded` / `cancelled` / `failed`;
|
||||||
|
6. запрет URL и выходов за разрешённые корни;
|
||||||
|
7. защиту зарезервированного профиля `last_run`;
|
||||||
|
8. приоритет явно указанного `classes.txt` над корневым YAML.
|
||||||
|
|
||||||
|
## Выполненные проверки
|
||||||
|
|
||||||
```text
|
```text
|
||||||
uv run pytest -q
|
uv run pytest -q -> 49 passed, 1 warning
|
||||||
uv run python -m compileall -q src tests
|
uv run python -m compileall -q src tests -> успешно
|
||||||
git diff --check
|
node --check src/yolo_webui/static/app.js -> успешно
|
||||||
|
node tests/frontend_smoke.js -> успешно
|
||||||
|
uv lock --check -> успешно
|
||||||
|
docker compose config -> успешно, host_ip=127.0.0.1
|
||||||
|
git diff --check -> успешно
|
||||||
```
|
```
|
||||||
|
|
||||||
Результат: `38 passed`; ошибок компиляции и форматирования diff нет.
|
Полная сборка Docker-образа локально не запускалась: Docker daemon недоступен.
|
||||||
|
Конфигурация Compose проверена отдельно, а соответствие lock-файла — через
|
||||||
|
`uv lock --check`.
|
||||||
|
|
||||||
Регрессионные тесты проверяют:
|
Оставшееся предупреждение pytest относится к deprecated-связке
|
||||||
|
`fastapi.testclient`/`starlette.testclient` с `httpx`; оно не связано с исправленными
|
||||||
1. успешный и ошибочный коды `subprocess_runner.main()`;
|
дефектами и не ломает тесты.
|
||||||
2. сброс остановки между запусками;
|
|
||||||
3. кооперативный сигнал вместо немедленного `terminate()`;
|
|
||||||
4. доставку раннего запроса после готовности subprocess;
|
|
||||||
5. запрет авторазбиения для `classify`;
|
|
||||||
6. пользовательские YAML-файлы классов и ошибочный явный путь;
|
|
||||||
7. датасеты из одного и двух изображений;
|
|
||||||
8. вложенные изображения и метки;
|
|
||||||
9. сохранность пользовательского каталога `split/` и уникальность результатов.
|
|
||||||
|
|
|
||||||
35
Dockerfile
35
Dockerfile
|
|
@ -1,8 +1,5 @@
|
||||||
FROM python:3.11-slim
|
FROM python:3.11-slim
|
||||||
|
|
||||||
# Build argument: 'cpu' for Mac/CPU-only environments, 'gpu' for CUDA/NVIDIA GPU support
|
|
||||||
ARG DEVICE=gpu
|
|
||||||
|
|
||||||
# Install system dependencies needed for OpenCV, PyTorch, and Ultralytics
|
# Install system dependencies needed for OpenCV, PyTorch, and Ultralytics
|
||||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||||
build-essential \
|
build-essential \
|
||||||
|
|
@ -12,33 +9,29 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||||
git \
|
git \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
# Install uv for fast dependency resolution using pip (avoids ghcr.io network issues)
|
# Pin the installer as well as application dependencies.
|
||||||
RUN pip install --no-cache-dir uv
|
RUN pip install --no-cache-dir uv==0.10.6
|
||||||
|
|
||||||
# Set working directory
|
# Set working directory
|
||||||
WORKDIR /workspace
|
WORKDIR /workspace
|
||||||
|
|
||||||
# Copy dependency definition
|
ENV UV_COMPILE_BYTECODE=1 \
|
||||||
COPY pyproject.toml ./
|
UV_LINK_MODE=copy \
|
||||||
|
UV_PROJECT_ENVIRONMENT=/opt/venv \
|
||||||
|
ULTRALYTICS_SAFE_LOAD=1
|
||||||
|
|
||||||
# Install dependencies using uv pip in system python to bypass uv.lock file hashes
|
# Install the exact dependency set recorded in uv.lock. Keeping the project out of
|
||||||
# and fetch the correct PyTorch package based on the target DEVICE (CPU or GPU)
|
# this layer allows dependency caching while source files change.
|
||||||
|
COPY pyproject.toml uv.lock README.md ./
|
||||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||||
if [ "$DEVICE" = "cpu" ]; then \
|
uv sync --locked --no-dev --no-install-project
|
||||||
echo "Installing CPU-only PyTorch..." && \
|
|
||||||
uv pip install --system --extra-index-url https://download.pytorch.org/whl/cpu -r pyproject.toml; \
|
|
||||||
else \
|
|
||||||
echo "Installing GPU (CUDA) PyTorch..." && \
|
|
||||||
uv pip install --system -r pyproject.toml; \
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Copy source code and files
|
# Copy source code and install the project without re-resolving dependencies.
|
||||||
COPY src ./src
|
COPY src ./src
|
||||||
COPY README.md ./
|
|
||||||
|
|
||||||
# Install the project itself without re-installing dependencies
|
|
||||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||||
uv pip install --system --no-deps -e .
|
uv sync --locked --no-dev
|
||||||
|
|
||||||
|
ENV PATH="/opt/venv/bin:$PATH"
|
||||||
|
|
||||||
# Expose Web UI port and MLflow port
|
# Expose Web UI port and MLflow port
|
||||||
EXPOSE 8000
|
EXPOSE 8000
|
||||||
|
|
|
||||||
110
README.md
110
README.md
|
|
@ -1,59 +1,117 @@
|
||||||
# YOLO Train TUI
|
# YOLO Train WebUI
|
||||||
|
|
||||||
Терминальный интерфейс для обучения моделей Ultralytics YOLO с автоматической
|
Локальный веб-интерфейс для обучения моделей Ultralytics YOLO с журналом,
|
||||||
регистрацией параметров, метрик и артефактов в MLflow.
|
графиками метрик, мягкой остановкой и интеграцией MLflow.
|
||||||
|
|
||||||
## Возможности
|
## Возможности
|
||||||
|
|
||||||
- задачи `detect`, `segment`, `classify`, `pose` и `obb`;
|
- задачи `detect`, `segment`, `classify`, `pose` и `obb`;
|
||||||
- локальные пути, YAML-конфигурации и официальные имена моделей/датасетов;
|
- локальные датасеты и официальные имена моделей Ultralytics;
|
||||||
- настройка эпох, размера изображения, batch, устройства, workers и patience;
|
- настройка эпох, размера изображения, batch, устройства, workers и patience;
|
||||||
- настройка цветовых и геометрических аугментаций, flip, Mosaic, MixUp,
|
- цветовые и геометрические аугментации, Mosaic, MixUp, CutMix, copy-paste,
|
||||||
CutMix, copy-paste, erasing и AutoAugment;
|
erasing и AutoAugment;
|
||||||
- обучение в фоновом потоке, прогресс по эпохам, журнал и мягкая остановка;
|
- live-прогресс, журнал, графики метрик и восстановление состояния после
|
||||||
- встроенная интеграция Ultralytics ↔ MLflow;
|
переподключения браузера;
|
||||||
- локальное MLflow-хранилище по умолчанию или внешний tracking server.
|
- сохранение профилей запуска и кооперативная остановка обучения;
|
||||||
|
- локальное MLflow-хранилище или внешний tracking server.
|
||||||
|
|
||||||
## Установка и запуск
|
## Локальная установка и запуск
|
||||||
|
|
||||||
|
Нужны Python 3.11+ и [uv](https://docs.astral.sh/uv/).
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv sync
|
uv sync --locked
|
||||||
uv run yolo-train-tui
|
uv run yolo-train-webui
|
||||||
```
|
```
|
||||||
|
|
||||||
Также приложение можно запустить как модуль:
|
Альтернативный запуск как Python-модуля:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv run -m yolo_tui
|
uv run -m yolo_webui
|
||||||
```
|
```
|
||||||
|
|
||||||
При первом использовании официального имени модели (например, `yolo11n.pt`)
|
Откройте `http://127.0.0.1:8000`. Сервер по умолчанию слушает только loopback.
|
||||||
Ultralytics автоматически скачает веса. Для полностью локальной работы укажите
|
|
||||||
путь к уже загруженному `.pt` или `.yaml` файлу.
|
При первом использовании официального имени модели, например `yolo11n.pt`,
|
||||||
|
Ultralytics скачает веса. Пользовательские модели размещайте в `./models` или в
|
||||||
|
`./runs`, а датасеты — в `./datasets`. Результаты записываются в `./runs`.
|
||||||
|
|
||||||
|
## Docker Compose
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker compose up --build
|
||||||
|
```
|
||||||
|
|
||||||
|
WebUI будет доступен по `http://127.0.0.1:8000`. Compose намеренно публикует порт
|
||||||
|
только на loopback. Не заменяйте адрес на `0.0.0.0` без аутентифицирующего reverse
|
||||||
|
proxy: API позволяет запускать и останавливать ресурсоёмкие задачи.
|
||||||
|
|
||||||
|
Для NVIDIA GPU раскомментируйте секцию `deploy.resources.reservations.devices` в
|
||||||
|
`docker-compose.yml`. Образ устанавливает зафиксированные в `uv.lock` зависимости;
|
||||||
|
для другого варианта PyTorch используйте отдельно сгенерированный и проверенный
|
||||||
|
lock-файл.
|
||||||
|
|
||||||
|
## Разрешённые пути
|
||||||
|
|
||||||
|
API отклоняет URL и не разрешает обучению читать или записывать произвольные пути:
|
||||||
|
|
||||||
|
- датасеты и файлы классов — `./datasets`;
|
||||||
|
- модели — `./models` и `./runs`;
|
||||||
|
- результаты — `./runs`.
|
||||||
|
|
||||||
|
Дополнительные доверенные корни можно перечислить через системный разделитель путей
|
||||||
|
в `YOLO_WEBUI_DATA_ROOTS`, `YOLO_WEBUI_MODEL_ROOTS` и
|
||||||
|
`YOLO_WEBUI_RUN_ROOTS`. Например, в Linux/macOS:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
YOLO_WEBUI_DATA_ROOTS=/mnt/datasets:/data/shared uv run yolo-train-webui
|
||||||
|
```
|
||||||
|
|
||||||
|
PyTorch checkpoints загружаются с включённым restricted-режимом Ultralytics
|
||||||
|
(`ULTRALYTICS_SAFE_LOAD=1`). Используйте только модели из доверенных источников.
|
||||||
|
|
||||||
## Датасеты
|
## Датасеты
|
||||||
|
|
||||||
Для `detect`, `segment`, `pose` и `obb` укажите путь к YAML-файлу датасета.
|
Для `detect`, `segment`, `pose` и `obb` укажите YAML-файл либо каталог со структурой
|
||||||
Для `classify` укажите каталог с подкаталогами `train`, `test`/`val`, внутри
|
`images/` + `labels/`. WebUI может детерминированно разделить такой каталог на
|
||||||
которых изображения разложены по классам.
|
train/val. Для `classify` нужен готовый каталог с `train` и `val`/`test`, внутри
|
||||||
|
которых изображения разложены по классам; автоматическое detection-style разбиение
|
||||||
|
для этой задачи отключено.
|
||||||
|
|
||||||
## MLflow
|
## MLflow
|
||||||
|
|
||||||
По умолчанию метаданные записываются в локальную SQLite-базу `./mlflow.db`,
|
По умолчанию метаданные записываются в `./mlflow.db`. Открыть интерфейс просмотра:
|
||||||
сервер для обучения не требуется. Артефакты сохраняются локально средствами MLflow.
|
|
||||||
Открыть интерфейс просмотра:
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv run mlflow ui --backend-store-uri sqlite:///mlflow.db
|
uv run mlflow ui --backend-store-uri sqlite:///mlflow.db
|
||||||
```
|
```
|
||||||
|
|
||||||
Затем откройте `http://127.0.0.1:5000`. Для удаленного MLflow-сервера включите
|
Затем откройте `http://127.0.0.1:5000`. Для внешнего tracking server укажите его URI
|
||||||
MLflow в TUI и замените Tracking URI на адрес вида `http://mlflow.example:5000`.
|
в настройках WebUI.
|
||||||
|
|
||||||
|
Для каждого завершённого запуска Ultralytics записывает в MLflow параметры,
|
||||||
|
поэпоховые метрики, графики, `results.csv` и checkpoints
|
||||||
|
`weights/best.pt`/`weights/last.pt`. SQLite-файл хранит tracking metadata, а сами
|
||||||
|
файлы находятся в MLflow Artifact Repository (локально — в `./mlruns`). Это
|
||||||
|
артефакты запуска, а не версии MLflow Model Registry: raw YOLO checkpoint не имеет
|
||||||
|
стандартной MLflow `MLmodel`-упаковки.
|
||||||
|
|
||||||
|
Проверка интеграции на минимальных датасетах для всех пяти задач:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run scripts/run_yolo26_smoke_training.py --mlflow
|
||||||
|
uv run scripts/verify_mlflow_smoke.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Второй скрипт завершается с ошибкой, если отсутствует experiment/run, параметры,
|
||||||
|
метрики, `results.csv`, `best.pt` или `last.pt` хотя бы для одной задачи.
|
||||||
|
|
||||||
## Проверка
|
## Проверка
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv run pytest
|
uv run pytest -q
|
||||||
|
node --check src/yolo_webui/static/app.js
|
||||||
|
docker compose config
|
||||||
```
|
```
|
||||||
|
|
||||||
Ultralytics распространяется по лицензии AGPL-3.0; для закрытых коммерческих
|
Ultralytics распространяется по лицензии AGPL-3.0; для закрытых коммерческих
|
||||||
|
|
|
||||||
|
|
@ -2,11 +2,10 @@ services:
|
||||||
webui:
|
webui:
|
||||||
build:
|
build:
|
||||||
context: .
|
context: .
|
||||||
args:
|
|
||||||
- DEVICE=cpu # 'cpu' for Mac, change to 'gpu' on a Linux server with NVIDIA GPU
|
|
||||||
image: yolo-train-webui:latest
|
image: yolo-train-webui:latest
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
# The training API has no built-in user accounts, so expose it locally only.
|
||||||
|
- "127.0.0.1:8000:8000"
|
||||||
volumes:
|
volumes:
|
||||||
- ./datasets:/workspace/datasets
|
- ./datasets:/workspace/datasets
|
||||||
- ./runs:/workspace/runs
|
- ./runs:/workspace/runs
|
||||||
|
|
|
||||||
157
scripts/create_yolo26_smoke_datasets.py
Normal file
157
scripts/create_yolo26_smoke_datasets.py
Normal file
|
|
@ -0,0 +1,157 @@
|
||||||
|
"""Create tiny deterministic datasets for all YOLO tasks supported by the WebUI."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from PIL import Image, ImageDraw
|
||||||
|
|
||||||
|
|
||||||
|
IMAGE_SIZE = 128
|
||||||
|
TRAIN_IMAGES = 6
|
||||||
|
VAL_IMAGES = 2
|
||||||
|
|
||||||
|
|
||||||
|
def image_geometry(index: int) -> tuple[int, tuple[int, int, int, int]]:
|
||||||
|
class_id = index % 2
|
||||||
|
offset = (index % 3) * 5
|
||||||
|
box = (26 + offset, 29, 93 + offset, 98)
|
||||||
|
return class_id, box
|
||||||
|
|
||||||
|
|
||||||
|
def make_image(path: Path, index: int, *, rotated: bool = False) -> None:
|
||||||
|
class_id, box = image_geometry(index)
|
||||||
|
colors = ((225, 72, 72), (55, 145, 225))
|
||||||
|
image = Image.new("RGB", (IMAGE_SIZE, IMAGE_SIZE), (238, 241, 245))
|
||||||
|
draw = ImageDraw.Draw(image)
|
||||||
|
if rotated:
|
||||||
|
cx, cy = 64 + (index % 3) * 3, 64
|
||||||
|
half_w, half_h = 37, 23
|
||||||
|
angle = math.radians(15 if class_id == 0 else -15)
|
||||||
|
points = []
|
||||||
|
for x, y in ((-half_w, -half_h), (half_w, -half_h), (half_w, half_h), (-half_w, half_h)):
|
||||||
|
points.append(
|
||||||
|
(
|
||||||
|
cx + x * math.cos(angle) - y * math.sin(angle),
|
||||||
|
cy + x * math.sin(angle) + y * math.cos(angle),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
draw.polygon(points, fill=colors[class_id], outline=(25, 25, 25), width=2)
|
||||||
|
else:
|
||||||
|
draw.rectangle(box, fill=colors[class_id], outline=(25, 25, 25), width=2)
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
image.save(path)
|
||||||
|
|
||||||
|
|
||||||
|
def normalized_box(box: tuple[int, int, int, int]) -> tuple[float, float, float, float]:
|
||||||
|
left, top, right, bottom = box
|
||||||
|
return (
|
||||||
|
(left + right) / 2 / IMAGE_SIZE,
|
||||||
|
(top + bottom) / 2 / IMAGE_SIZE,
|
||||||
|
(right - left) / IMAGE_SIZE,
|
||||||
|
(bottom - top) / IMAGE_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def write_yaml(root: Path, task: str, extra: str = "") -> None:
|
||||||
|
yaml_text = (
|
||||||
|
f"path: {root.resolve()}\n"
|
||||||
|
"train: images/train\n"
|
||||||
|
"val: images/val\n"
|
||||||
|
"names:\n"
|
||||||
|
" 0: red_shape\n"
|
||||||
|
" 1: blue_shape\n"
|
||||||
|
f"{extra}"
|
||||||
|
)
|
||||||
|
(root / f"{task}.yaml").write_text(yaml_text, encoding="utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def create_detection_style(base: Path, task: str) -> None:
|
||||||
|
root = base / task
|
||||||
|
for split, count in (("train", TRAIN_IMAGES), ("val", VAL_IMAGES)):
|
||||||
|
for index in range(count):
|
||||||
|
sample = index if split == "train" else index + TRAIN_IMAGES
|
||||||
|
image_path = root / "images" / split / f"sample_{sample:02d}.png"
|
||||||
|
label_path = root / "labels" / split / f"sample_{sample:02d}.txt"
|
||||||
|
make_image(image_path, sample, rotated=task == "obb")
|
||||||
|
class_id, box = image_geometry(sample)
|
||||||
|
cx, cy, width, height = normalized_box(box)
|
||||||
|
|
||||||
|
if task == "detect":
|
||||||
|
label = f"{class_id} {cx:.6f} {cy:.6f} {width:.6f} {height:.6f}\n"
|
||||||
|
elif task == "segment":
|
||||||
|
left, top, right, bottom = (value / IMAGE_SIZE for value in box)
|
||||||
|
label = (
|
||||||
|
f"{class_id} {left:.6f} {top:.6f} {right:.6f} {top:.6f} "
|
||||||
|
f"{right:.6f} {bottom:.6f} {left:.6f} {bottom:.6f}\n"
|
||||||
|
)
|
||||||
|
elif task == "pose":
|
||||||
|
class_id = 0
|
||||||
|
points = (
|
||||||
|
(cx, cy - height * 0.25),
|
||||||
|
(cx - width * 0.25, cy),
|
||||||
|
(cx + width * 0.25, cy),
|
||||||
|
(cx, cy + height * 0.25),
|
||||||
|
)
|
||||||
|
keypoints = " ".join(f"{x:.6f} {y:.6f} 2" for x, y in points)
|
||||||
|
label = f"{class_id} {cx:.6f} {cy:.6f} {width:.6f} {height:.6f} {keypoints}\n"
|
||||||
|
elif task == "obb":
|
||||||
|
angle = math.radians(15 if class_id == 0 else -15)
|
||||||
|
center_x, center_y = 64 + (sample % 3) * 3, 64
|
||||||
|
half_w, half_h = 37, 23
|
||||||
|
points = []
|
||||||
|
for x, y in ((-half_w, -half_h), (half_w, -half_h), (half_w, half_h), (-half_w, half_h)):
|
||||||
|
px = center_x + x * math.cos(angle) - y * math.sin(angle)
|
||||||
|
py = center_y + x * math.sin(angle) + y * math.cos(angle)
|
||||||
|
points.extend((px / IMAGE_SIZE, py / IMAGE_SIZE))
|
||||||
|
label = f"{class_id} " + " ".join(f"{value:.6f}" for value in points) + "\n"
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported task: {task}")
|
||||||
|
|
||||||
|
label_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
label_path.write_text(label, encoding="utf-8")
|
||||||
|
|
||||||
|
if task == "pose":
|
||||||
|
write_yaml(
|
||||||
|
root,
|
||||||
|
task,
|
||||||
|
extra=(
|
||||||
|
"kpt_shape: [4, 3]\n"
|
||||||
|
"flip_idx: [0, 2, 1, 3]\n"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
write_yaml(root, task)
|
||||||
|
|
||||||
|
|
||||||
|
def create_classification(base: Path) -> None:
|
||||||
|
root = base / "classify"
|
||||||
|
for split, count in (("train", TRAIN_IMAGES), ("val", 4)):
|
||||||
|
for index in range(count):
|
||||||
|
sample = index if split == "train" else index + TRAIN_IMAGES
|
||||||
|
class_id = sample % 2
|
||||||
|
make_image(root / split / ("red_shape" if class_id == 0 else "blue_shape") / f"sample_{sample:02d}.png", sample)
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--output", type=Path, default=Path("datasets/yolo26_smoke"))
|
||||||
|
parser.add_argument("--force", action="store_true")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if args.output.exists():
|
||||||
|
if not args.force:
|
||||||
|
raise SystemExit(f"Dataset already exists: {args.output}; use --force to recreate it")
|
||||||
|
shutil.rmtree(args.output)
|
||||||
|
|
||||||
|
for task in ("detect", "segment", "pose", "obb"):
|
||||||
|
create_detection_style(args.output, task)
|
||||||
|
create_classification(args.output)
|
||||||
|
print(args.output.resolve())
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
99
scripts/run_yolo26_smoke_training.py
Normal file
99
scripts/run_yolo26_smoke_training.py
Normal file
|
|
@ -0,0 +1,99 @@
|
||||||
|
"""Run one small CPU training epoch for every task supported by the WebUI."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
import traceback
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from yolo_webui import TrainingConfig, TrainingRunner
|
||||||
|
from yolo_webui.config import AugmentationConfig, MlflowConfig
|
||||||
|
|
||||||
|
|
||||||
|
TASKS = {
|
||||||
|
"detect": ("datasets/yolo26_smoke/detect/detect.yaml", "yolo26n.pt"),
|
||||||
|
"segment": ("datasets/yolo26_smoke/segment/segment.yaml", "yolo26n-seg.pt"),
|
||||||
|
"classify": ("datasets/yolo26_smoke/classify", "yolo26n-cls.pt"),
|
||||||
|
"pose": ("datasets/yolo26_smoke/pose/pose.yaml", "yolo26n-pose.pt"),
|
||||||
|
"obb": ("datasets/yolo26_smoke/obb/obb.yaml", "yolo26n-obb.pt"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument(
|
||||||
|
"--mlflow",
|
||||||
|
action="store_true",
|
||||||
|
help="Enable MLflow logging for every smoke-training run.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--tracking-uri",
|
||||||
|
default="sqlite:///runs/yolo26_mlflow_smoke/mlflow.db",
|
||||||
|
)
|
||||||
|
parser.add_argument("--experiment", default="yolo26-mlflow-smoke")
|
||||||
|
parser.add_argument("--project", type=Path)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
results: dict[str, dict[str, object]] = {}
|
||||||
|
default_project = "runs/yolo26_mlflow_smoke/train" if args.mlflow else "runs/yolo26_smoke"
|
||||||
|
project_dir = (args.project or Path(default_project)).resolve()
|
||||||
|
for task, (dataset, model) in TASKS.items():
|
||||||
|
print(f"\n=== {task}: {model} ===", flush=True)
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset=dataset,
|
||||||
|
model=model,
|
||||||
|
task=task, # type: ignore[arg-type]
|
||||||
|
epochs=1,
|
||||||
|
image_size=64,
|
||||||
|
batch_size=2,
|
||||||
|
device="cpu",
|
||||||
|
workers=0,
|
||||||
|
patience=0,
|
||||||
|
project=str(project_dir),
|
||||||
|
run_name=task,
|
||||||
|
augmentation=AugmentationConfig(enabled=False),
|
||||||
|
mlflow=MlflowConfig(
|
||||||
|
enabled=args.mlflow,
|
||||||
|
tracking_uri=args.tracking_uri,
|
||||||
|
experiment_name=args.experiment,
|
||||||
|
run_name=f"{task}-smoke" if args.mlflow else "",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
started = time.monotonic()
|
||||||
|
try:
|
||||||
|
output = TrainingRunner().train(
|
||||||
|
config,
|
||||||
|
lambda event: print(
|
||||||
|
f"[{event.kind}] {event.message}",
|
||||||
|
flush=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
results[task] = {
|
||||||
|
"status": "succeeded",
|
||||||
|
"seconds": round(time.monotonic() - started, 2),
|
||||||
|
"output": str(output) if output else None,
|
||||||
|
}
|
||||||
|
except Exception as exc:
|
||||||
|
traceback.print_exc()
|
||||||
|
results[task] = {
|
||||||
|
"status": "failed",
|
||||||
|
"seconds": round(time.monotonic() - started, 2),
|
||||||
|
"error": f"{type(exc).__name__}: {exc}",
|
||||||
|
}
|
||||||
|
|
||||||
|
summary_path = project_dir.parent / "smoke_summary.json" if args.mlflow else project_dir / "smoke_summary.json"
|
||||||
|
summary_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
summary_path.write_text(
|
||||||
|
json.dumps(results, indent=2, ensure_ascii=False) + "\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
print(f"\nSummary: {summary_path.resolve()}")
|
||||||
|
print(json.dumps(results, indent=2, ensure_ascii=False))
|
||||||
|
if any(result["status"] != "succeeded" for result in results.values()):
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
83
scripts/verify_mlflow_smoke.py
Normal file
83
scripts/verify_mlflow_smoke.py
Normal file
|
|
@ -0,0 +1,83 @@
|
||||||
|
"""Verify that every YOLO smoke task was persisted completely in MLflow."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
|
||||||
|
import mlflow
|
||||||
|
from mlflow.entities import Run
|
||||||
|
from mlflow.tracking import MlflowClient
|
||||||
|
|
||||||
|
|
||||||
|
TASKS = ("detect", "segment", "classify", "pose", "obb")
|
||||||
|
REQUIRED_ARTIFACTS = {"weights/best.pt", "weights/last.pt", "results.csv"}
|
||||||
|
|
||||||
|
|
||||||
|
def artifact_paths(client: MlflowClient, run_id: str, path: str = "") -> set[str]:
|
||||||
|
result: set[str] = set()
|
||||||
|
for artifact in client.list_artifacts(run_id, path):
|
||||||
|
if artifact.is_dir:
|
||||||
|
result.update(artifact_paths(client, run_id, artifact.path))
|
||||||
|
else:
|
||||||
|
result.add(artifact.path)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def latest_task_run(runs: list[Run], task: str) -> Run:
|
||||||
|
expected_name = f"{task}-smoke"
|
||||||
|
for run in runs:
|
||||||
|
if run.data.tags.get("mlflow.runName") == expected_name:
|
||||||
|
return run
|
||||||
|
raise AssertionError(f"MLflow run not found: {expected_name}")
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument(
|
||||||
|
"--tracking-uri",
|
||||||
|
default="sqlite:///runs/yolo26_mlflow_smoke/mlflow.db",
|
||||||
|
)
|
||||||
|
parser.add_argument("--experiment", default="yolo26-mlflow-smoke")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
mlflow.set_tracking_uri(args.tracking_uri)
|
||||||
|
client = MlflowClient()
|
||||||
|
experiment = client.get_experiment_by_name(args.experiment)
|
||||||
|
if experiment is None:
|
||||||
|
raise AssertionError(f"MLflow experiment not found: {args.experiment}")
|
||||||
|
|
||||||
|
runs = client.search_runs(
|
||||||
|
[experiment.experiment_id],
|
||||||
|
order_by=["start_time DESC"],
|
||||||
|
)
|
||||||
|
summary: dict[str, object] = {
|
||||||
|
"tracking_uri": args.tracking_uri,
|
||||||
|
"experiment_id": experiment.experiment_id,
|
||||||
|
"artifact_location": experiment.artifact_location,
|
||||||
|
"tasks": {},
|
||||||
|
}
|
||||||
|
task_summary: dict[str, object] = summary["tasks"] # type: ignore[assignment]
|
||||||
|
|
||||||
|
for task in TASKS:
|
||||||
|
run = latest_task_run(runs, task)
|
||||||
|
artifacts = artifact_paths(client, run.info.run_id)
|
||||||
|
missing = REQUIRED_ARTIFACTS - artifacts
|
||||||
|
assert run.info.status == "FINISHED", (task, run.info.status)
|
||||||
|
assert run.data.params, f"No parameters logged for {task}"
|
||||||
|
assert run.data.metrics, f"No metrics logged for {task}"
|
||||||
|
assert not missing, f"Missing artifacts for {task}: {sorted(missing)}"
|
||||||
|
task_summary[task] = {
|
||||||
|
"run_id": run.info.run_id,
|
||||||
|
"status": run.info.status,
|
||||||
|
"parameters": len(run.data.params),
|
||||||
|
"metrics": len(run.data.metrics),
|
||||||
|
"artifact_uri": run.info.artifact_uri,
|
||||||
|
"required_artifacts": sorted(REQUIRED_ARTIFACTS),
|
||||||
|
}
|
||||||
|
|
||||||
|
print(json.dumps(summary, indent=2, ensure_ascii=False))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|
@ -1,14 +1,16 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
import threading
|
import threading
|
||||||
from dataclasses import asdict, dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
|
@ -23,17 +25,19 @@ from yolo_webui.trainer import TrainingRunner
|
||||||
# Set up logging
|
# Set up logging
|
||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||||
logger = logging.getLogger("yolo_webui")
|
logger = logging.getLogger("yolo_webui")
|
||||||
|
SESSION_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class LiveState:
|
class LiveState:
|
||||||
status: str = "idle" # idle, preparing, training, stopping, finished, failed
|
status: str = "idle" # idle, preparing, training, stopping, succeeded, cancelled, failed
|
||||||
epoch: int = 0
|
epoch: int = 0
|
||||||
total_epochs: int = 0
|
total_epochs: int = 0
|
||||||
logs: list[str] = field(default_factory=list)
|
logs: list[str] = field(default_factory=list)
|
||||||
metrics: list[dict[str, Any]] = field(default_factory=list)
|
metrics: list[dict[str, Any]] = field(default_factory=list)
|
||||||
output_dir: str | None = None
|
output_dir: str | None = None
|
||||||
stop_requested: bool = False
|
stop_requested: bool = False
|
||||||
|
last_event_kind: str | None = None
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
self.status = "idle"
|
self.status = "idle"
|
||||||
|
|
@ -43,6 +47,7 @@ class LiveState:
|
||||||
self.metrics = []
|
self.metrics = []
|
||||||
self.output_dir = None
|
self.output_dir = None
|
||||||
self.stop_requested = False
|
self.stop_requested = False
|
||||||
|
self.last_event_kind = None
|
||||||
|
|
||||||
|
|
||||||
class TrainingManager:
|
class TrainingManager:
|
||||||
|
|
@ -54,10 +59,16 @@ class TrainingManager:
|
||||||
self.active_websockets: set[WebSocket] = set()
|
self.active_websockets: set[WebSocket] = set()
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
self._thread: threading.Thread | None = None
|
self._thread: threading.Thread | None = None
|
||||||
|
self._event_loop: asyncio.AbstractEventLoop | None = None
|
||||||
|
self._broadcast_lock: asyncio.Lock | None = None
|
||||||
|
|
||||||
def add_websocket(self, websocket: WebSocket) -> None:
|
def add_websocket(self, websocket: WebSocket) -> None:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.active_websockets.add(websocket)
|
self.active_websockets.add(websocket)
|
||||||
|
if self._event_loop is not loop:
|
||||||
|
self._event_loop = loop
|
||||||
|
self._broadcast_lock = asyncio.Lock()
|
||||||
|
|
||||||
def remove_websocket(self, websocket: WebSocket) -> None:
|
def remove_websocket(self, websocket: WebSocket) -> None:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
|
|
@ -65,27 +76,64 @@ class TrainingManager:
|
||||||
|
|
||||||
def broadcast(self, data: dict[str, Any]) -> None:
|
def broadcast(self, data: dict[str, Any]) -> None:
|
||||||
payload = json.dumps(data)
|
payload = json.dumps(data)
|
||||||
# Create a copy under lock to avoid modification during traversal
|
with self._lock:
|
||||||
|
loop = self._event_loop
|
||||||
|
has_sockets = bool(self.active_websockets)
|
||||||
|
|
||||||
|
if not has_sockets:
|
||||||
|
return
|
||||||
|
if loop is None or loop.is_closed():
|
||||||
|
logger.warning("WebSocket event loop is unavailable; broadcast was skipped")
|
||||||
|
return
|
||||||
|
|
||||||
|
coroutine = self._send_payload(payload)
|
||||||
|
try:
|
||||||
|
try:
|
||||||
|
running_loop = asyncio.get_running_loop()
|
||||||
|
except RuntimeError:
|
||||||
|
running_loop = None
|
||||||
|
|
||||||
|
if running_loop is loop:
|
||||||
|
future = loop.create_task(coroutine)
|
||||||
|
else:
|
||||||
|
future = asyncio.run_coroutine_threadsafe(coroutine, loop)
|
||||||
|
future.add_done_callback(self._log_broadcast_failure)
|
||||||
|
except Exception:
|
||||||
|
coroutine.close()
|
||||||
|
logger.exception("Failed to schedule WebSocket broadcast")
|
||||||
|
|
||||||
|
async def _send_payload(self, payload: str) -> None:
|
||||||
|
broadcast_lock = self._broadcast_lock
|
||||||
|
if broadcast_lock is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
async with broadcast_lock:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
sockets = list(self.active_websockets)
|
sockets = list(self.active_websockets)
|
||||||
|
|
||||||
# Send outside lock to prevent blocking
|
failed: list[WebSocket] = []
|
||||||
for ws in sockets:
|
for websocket in sockets:
|
||||||
try:
|
try:
|
||||||
import asyncio
|
await websocket.send_text(payload)
|
||||||
# Check if we are in an event loop
|
|
||||||
try:
|
|
||||||
loop = asyncio.get_event_loop()
|
|
||||||
except RuntimeError:
|
|
||||||
loop = asyncio.new_event_loop()
|
|
||||||
asyncio.set_event_loop(loop)
|
|
||||||
|
|
||||||
if loop.is_running():
|
|
||||||
loop.create_task(ws.send_text(payload))
|
|
||||||
else:
|
|
||||||
loop.run_until_complete(ws.send_text(payload))
|
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
failed.append(websocket)
|
||||||
|
logger.warning("Dropping failed WebSocket client", exc_info=True)
|
||||||
|
|
||||||
|
if failed:
|
||||||
|
with self._lock:
|
||||||
|
for websocket in failed:
|
||||||
|
self.active_websockets.discard(websocket)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _log_broadcast_failure(future: Any) -> None:
|
||||||
|
if future.cancelled():
|
||||||
|
return
|
||||||
|
error = future.exception()
|
||||||
|
if error is not None:
|
||||||
|
logger.error(
|
||||||
|
"WebSocket broadcast failed",
|
||||||
|
exc_info=(type(error), error, error.__traceback__),
|
||||||
|
)
|
||||||
|
|
||||||
def add_log(self, text: str, level: str = "info") -> None:
|
def add_log(self, text: str, level: str = "info") -> None:
|
||||||
log_entry = f"__LOG_LEVEL_{level.upper()}__:{text}"
|
log_entry = f"__LOG_LEVEL_{level.upper()}__:{text}"
|
||||||
|
|
@ -98,7 +146,9 @@ class TrainingManager:
|
||||||
|
|
||||||
def start_training(self, config: TrainingConfig) -> None:
|
def start_training(self, config: TrainingConfig) -> None:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
if self.state.status in ("preparing", "training", "stopping"):
|
if self.state.status in ("preparing", "training", "stopping") or (
|
||||||
|
self._thread is not None and self._thread.is_alive()
|
||||||
|
):
|
||||||
raise ValueError("Обучение уже выполняется.")
|
raise ValueError("Обучение уже выполняется.")
|
||||||
|
|
||||||
self.state.reset()
|
self.state.reset()
|
||||||
|
|
@ -149,7 +199,7 @@ class TrainingManager:
|
||||||
parts = message.split(" · ")[1:]
|
parts = message.split(" · ")[1:]
|
||||||
for p in parts:
|
for p in parts:
|
||||||
if "=" in p:
|
if "=" in p:
|
||||||
k, v = p.split("=")
|
k, v = p.split("=", 1)
|
||||||
try:
|
try:
|
||||||
metrics_dict[k.strip()] = float(v.strip())
|
metrics_dict[k.strip()] = float(v.strip())
|
||||||
except ValueError:
|
except ValueError:
|
||||||
|
|
@ -159,10 +209,15 @@ class TrainingManager:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.state.metrics.append(metrics_dict)
|
self.state.metrics.append(metrics_dict)
|
||||||
|
|
||||||
|
status_update = None
|
||||||
with self._lock:
|
with self._lock:
|
||||||
|
self.state.last_event_kind = kind
|
||||||
if kind == "started" and self.state.status == "preparing":
|
if kind == "started" and self.state.status == "preparing":
|
||||||
self.state.status = "training"
|
self.state.status = "training"
|
||||||
self.broadcast({"type": "status", "status": self.state.status})
|
status_update = self.state.status
|
||||||
|
|
||||||
|
if status_update is not None:
|
||||||
|
self.broadcast({"type": "status", "status": status_update})
|
||||||
|
|
||||||
self.add_log(message, "progress" if is_progress else kind)
|
self.add_log(message, "progress" if is_progress else kind)
|
||||||
self.broadcast({
|
self.broadcast({
|
||||||
|
|
@ -217,24 +272,7 @@ class TrainingManager:
|
||||||
process.wait()
|
process.wait()
|
||||||
rc = process.returncode
|
rc = process.returncode
|
||||||
self.runner.clear_subprocess()
|
self.runner.clear_subprocess()
|
||||||
|
self._finalize_process_result(rc)
|
||||||
stopped = False
|
|
||||||
with self._lock:
|
|
||||||
stopped = self.state.stop_requested
|
|
||||||
|
|
||||||
if rc == 0:
|
|
||||||
with self._lock:
|
|
||||||
self.state.status = "finished"
|
|
||||||
self.add_log("Обучение успешно завершено.", "success")
|
|
||||||
else:
|
|
||||||
if stopped:
|
|
||||||
with self._lock:
|
|
||||||
self.state.status = "finished"
|
|
||||||
self.add_log("Обучение остановлено пользователем.", "warning")
|
|
||||||
else:
|
|
||||||
with self._lock:
|
|
||||||
self.state.status = "failed"
|
|
||||||
self.add_log("Процесс обучения завершился с ошибкой. Проверьте логи выше.", "error")
|
|
||||||
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.exception("Error in training process thread:")
|
logger.exception("Error in training process thread:")
|
||||||
|
|
@ -257,7 +295,38 @@ class TrainingManager:
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
self.runner.clear_subprocess()
|
self.runner.clear_subprocess()
|
||||||
self.broadcast({"type": "status", "status": self.state.status, "output_dir": self.state.output_dir})
|
with self._lock:
|
||||||
|
final_status = self.state.status
|
||||||
|
output_dir = self.state.output_dir
|
||||||
|
self.broadcast({"type": "status", "status": final_status, "output_dir": output_dir})
|
||||||
|
|
||||||
|
def _finalize_process_result(self, return_code: int) -> None:
|
||||||
|
with self._lock:
|
||||||
|
stopped = self.state.stop_requested
|
||||||
|
last_event_kind = self.state.last_event_kind
|
||||||
|
|
||||||
|
was_cancelled = stopped and (
|
||||||
|
return_code == 0
|
||||||
|
or last_event_kind == "cancelled"
|
||||||
|
or self.runner.force_stop_triggered
|
||||||
|
)
|
||||||
|
|
||||||
|
if was_cancelled:
|
||||||
|
status = "cancelled"
|
||||||
|
message = "Обучение остановлено пользователем."
|
||||||
|
level = "warning"
|
||||||
|
elif return_code == 0:
|
||||||
|
status = "succeeded"
|
||||||
|
message = "Обучение успешно завершено."
|
||||||
|
level = "success"
|
||||||
|
else:
|
||||||
|
status = "failed"
|
||||||
|
message = "Процесс обучения завершился с ошибкой. Проверьте логи выше."
|
||||||
|
level = "error"
|
||||||
|
|
||||||
|
with self._lock:
|
||||||
|
self.state.status = status
|
||||||
|
self.add_log(message, level)
|
||||||
|
|
||||||
|
|
||||||
manager = TrainingManager()
|
manager = TrainingManager()
|
||||||
|
|
@ -286,6 +355,17 @@ def get_sessions_dir() -> Path:
|
||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def get_session_path(name: str, *, allow_last_run: bool = True) -> Path:
|
||||||
|
if SESSION_NAME_PATTERN.fullmatch(name) is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Имя сессии может содержать только латинские буквы, цифры, '_' и '-'.",
|
||||||
|
)
|
||||||
|
if not allow_last_run and name == "last_run":
|
||||||
|
raise HTTPException(status_code=400, detail="Имя 'last_run' зарезервировано.")
|
||||||
|
return get_sessions_dir() / f"{name}.json"
|
||||||
|
|
||||||
|
|
||||||
@app.get("/api/config/defaults")
|
@app.get("/api/config/defaults")
|
||||||
async def get_defaults():
|
async def get_defaults():
|
||||||
# Return defaults by instantiating with dummy paths and serializing
|
# Return defaults by instantiating with dummy paths and serializing
|
||||||
|
|
@ -350,8 +430,7 @@ async def list_models():
|
||||||
|
|
||||||
@app.get("/api/sessions/{name}")
|
@app.get("/api/sessions/{name}")
|
||||||
async def load_session(name: str):
|
async def load_session(name: str):
|
||||||
sessions_dir = get_sessions_dir()
|
file_path = get_session_path(name)
|
||||||
file_path = sessions_dir / f"{name}.json"
|
|
||||||
if not file_path.exists():
|
if not file_path.exists():
|
||||||
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
||||||
try:
|
try:
|
||||||
|
|
@ -363,8 +442,7 @@ async def load_session(name: str):
|
||||||
|
|
||||||
@app.post("/api/sessions/{name}")
|
@app.post("/api/sessions/{name}")
|
||||||
async def save_session(name: str, config_data: dict[str, Any]):
|
async def save_session(name: str, config_data: dict[str, Any]):
|
||||||
sessions_dir = get_sessions_dir()
|
file_path = get_session_path(name, allow_last_run=False)
|
||||||
file_path = sessions_dir / f"{name}.json"
|
|
||||||
try:
|
try:
|
||||||
with file_path.open("w", encoding="utf-8") as f:
|
with file_path.open("w", encoding="utf-8") as f:
|
||||||
json.dump(config_data, f, ensure_ascii=False, indent=2)
|
json.dump(config_data, f, ensure_ascii=False, indent=2)
|
||||||
|
|
@ -375,8 +453,7 @@ async def save_session(name: str, config_data: dict[str, Any]):
|
||||||
|
|
||||||
@app.delete("/api/sessions/{name}")
|
@app.delete("/api/sessions/{name}")
|
||||||
async def delete_session(name: str):
|
async def delete_session(name: str):
|
||||||
sessions_dir = get_sessions_dir()
|
file_path = get_session_path(name, allow_last_run=False)
|
||||||
file_path = sessions_dir / f"{name}.json"
|
|
||||||
if not file_path.exists():
|
if not file_path.exists():
|
||||||
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
||||||
try:
|
try:
|
||||||
|
|
@ -448,15 +525,16 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
# We format log items for the UI
|
# We format log items for the UI
|
||||||
"logs": [log.split(":", 1) for log in manager.state.logs if ":" in log],
|
"logs": [log.split(":", 1) for log in manager.state.logs if ":" in log],
|
||||||
}
|
}
|
||||||
await websocket.send_text(json.dumps(state_dict))
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
await websocket.send_text(json.dumps(state_dict))
|
||||||
while True:
|
while True:
|
||||||
# Keep connection alive; discard incoming messages
|
# Keep connection alive; discard incoming messages
|
||||||
await websocket.receive_text()
|
await websocket.receive_text()
|
||||||
except WebSocketDisconnect:
|
except WebSocketDisconnect:
|
||||||
manager.remove_websocket(websocket)
|
pass
|
||||||
except Exception:
|
except Exception:
|
||||||
|
logger.warning("WebSocket connection failed", exc_info=True)
|
||||||
|
finally:
|
||||||
manager.remove_websocket(websocket)
|
manager.remove_websocket(websocket)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,10 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Literal, Any
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
import re
|
||||||
|
from typing import Any, Literal
|
||||||
|
|
||||||
|
|
||||||
YoloTask = Literal["detect", "segment", "classify", "pose", "obb"]
|
YoloTask = Literal["detect", "segment", "classify", "pose", "obb"]
|
||||||
|
|
@ -20,6 +23,63 @@ SUPPORTED_AUTO_AUGMENT_POLICIES: tuple[AutoAugmentPolicy, ...] = (
|
||||||
"augmix",
|
"augmix",
|
||||||
)
|
)
|
||||||
SUPPORTED_COPY_PASTE_MODES: tuple[CopyPasteMode, ...] = ("flip", "mixup")
|
SUPPORTED_COPY_PASTE_MODES: tuple[CopyPasteMode, ...] = ("flip", "mixup")
|
||||||
|
SAFE_IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]*$")
|
||||||
|
MODEL_SUFFIXES = {".pt", ".pth", ".yaml", ".yml"}
|
||||||
|
|
||||||
|
|
||||||
|
def _allowed_roots(defaults: tuple[str, ...], environment_name: str) -> tuple[Path, ...]:
|
||||||
|
configured = [
|
||||||
|
item
|
||||||
|
for item in os.environ.get(environment_name, "").split(os.pathsep)
|
||||||
|
if item.strip()
|
||||||
|
]
|
||||||
|
roots = (*defaults, *configured)
|
||||||
|
return tuple(Path(root).expanduser().resolve(strict=False) for root in roots)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_within(path: Path, roots: tuple[Path, ...]) -> bool:
|
||||||
|
resolved = path.expanduser().resolve(strict=False)
|
||||||
|
return any(resolved == root or resolved.is_relative_to(root) for root in roots)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_local_reference(
|
||||||
|
value: str,
|
||||||
|
*,
|
||||||
|
label: str,
|
||||||
|
roots: tuple[Path, ...],
|
||||||
|
allow_identifier: bool = False,
|
||||||
|
allowed_suffixes: set[str] | None = None,
|
||||||
|
) -> None:
|
||||||
|
normalized = value.strip()
|
||||||
|
if "://" in normalized or normalized.startswith("//"):
|
||||||
|
raise ValueError(f"{label} не может быть URL.")
|
||||||
|
|
||||||
|
is_identifier = (
|
||||||
|
"/" not in normalized
|
||||||
|
and "\\" not in normalized
|
||||||
|
and SAFE_IDENTIFIER.fullmatch(normalized) is not None
|
||||||
|
)
|
||||||
|
if allow_identifier and is_identifier:
|
||||||
|
local_candidate = Path.cwd() / normalized
|
||||||
|
if local_candidate.exists() and not _is_within(local_candidate, roots):
|
||||||
|
allowed = ", ".join(str(root) for root in roots)
|
||||||
|
raise ValueError(
|
||||||
|
f"{label} с таким именем найден вне разрешённого каталога: {allowed}."
|
||||||
|
)
|
||||||
|
if allowed_suffixes is not None and Path(normalized).suffix.lower() not in allowed_suffixes:
|
||||||
|
expected = ", ".join(sorted(allowed_suffixes))
|
||||||
|
raise ValueError(f"{label} должен иметь расширение {expected}.")
|
||||||
|
return
|
||||||
|
|
||||||
|
candidate = Path(normalized)
|
||||||
|
if not candidate.is_absolute():
|
||||||
|
candidate = Path.cwd() / candidate
|
||||||
|
if not _is_within(candidate, roots):
|
||||||
|
allowed = ", ".join(str(root) for root in roots)
|
||||||
|
raise ValueError(f"{label} должен находиться в разрешённом каталоге: {allowed}.")
|
||||||
|
if allowed_suffixes is not None and candidate.suffix.lower() not in allowed_suffixes:
|
||||||
|
expected = ", ".join(sorted(allowed_suffixes))
|
||||||
|
raise ValueError(f"{label} должен иметь расширение {expected}.")
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
|
|
@ -155,6 +215,37 @@ class TrainingConfig:
|
||||||
raise ValueError("Укажите путь или имя датасета.")
|
raise ValueError("Укажите путь или имя датасета.")
|
||||||
if not self.model.strip():
|
if not self.model.strip():
|
||||||
raise ValueError("Укажите путь или имя модели.")
|
raise ValueError("Укажите путь или имя модели.")
|
||||||
|
|
||||||
|
data_roots = _allowed_roots(("datasets",), "YOLO_WEBUI_DATA_ROOTS")
|
||||||
|
model_roots = _allowed_roots(
|
||||||
|
("models", "runs"),
|
||||||
|
"YOLO_WEBUI_MODEL_ROOTS",
|
||||||
|
)
|
||||||
|
run_roots = _allowed_roots(("runs",), "YOLO_WEBUI_RUN_ROOTS")
|
||||||
|
_validate_local_reference(
|
||||||
|
self.dataset,
|
||||||
|
label="Датасет",
|
||||||
|
roots=data_roots,
|
||||||
|
allow_identifier=True,
|
||||||
|
)
|
||||||
|
_validate_local_reference(
|
||||||
|
self.model,
|
||||||
|
label="Модель",
|
||||||
|
roots=model_roots,
|
||||||
|
allow_identifier=True,
|
||||||
|
allowed_suffixes=MODEL_SUFFIXES,
|
||||||
|
)
|
||||||
|
_validate_local_reference(
|
||||||
|
self.project.strip() or "runs/train",
|
||||||
|
label="Каталог результатов",
|
||||||
|
roots=run_roots,
|
||||||
|
)
|
||||||
|
if self.split.classes_path.strip():
|
||||||
|
_validate_local_reference(
|
||||||
|
self.split.classes_path,
|
||||||
|
label="Файл классов",
|
||||||
|
roots=data_roots,
|
||||||
|
)
|
||||||
if self.task not in SUPPORTED_TASKS:
|
if self.task not in SUPPORTED_TASKS:
|
||||||
raise ValueError(f"Неизвестный тип задачи: {self.task}.")
|
raise ValueError(f"Неизвестный тип задачи: {self.task}.")
|
||||||
if self.task == "classify" and self.split.enabled:
|
if self.task == "classify" and self.split.enabled:
|
||||||
|
|
@ -198,7 +289,6 @@ class TrainingConfig:
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def resolved_model(self) -> str:
|
def resolved_model(self) -> str:
|
||||||
from pathlib import Path
|
|
||||||
model_path = self.model.strip()
|
model_path = self.model.strip()
|
||||||
if "/" not in model_path and "\\" not in model_path:
|
if "/" not in model_path and "\\" not in model_path:
|
||||||
# Ensure models directory exists inside workspace
|
# Ensure models directory exists inside workspace
|
||||||
|
|
|
||||||
|
|
@ -239,7 +239,8 @@ def split_dataset(
|
||||||
"val": (relative_split_dir / val_txt_path.name).as_posix(),
|
"val": (relative_split_dir / val_txt_path.name).as_posix(),
|
||||||
})
|
})
|
||||||
|
|
||||||
if "names" not in dataset_data:
|
# `read_classes()` has already applied the explicit-path precedence and validated
|
||||||
|
# the IDs. A different root YAML must never replace that authoritative result.
|
||||||
dataset_data["names"] = classes
|
dataset_data["names"] = classes
|
||||||
|
|
||||||
_write_new(
|
_write_new(
|
||||||
|
|
|
||||||
|
|
@ -259,38 +259,35 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
|
|
||||||
function updateChart(epoch, metrics) {
|
function updateChart(epoch, metrics) {
|
||||||
if (!metricsChart) {
|
if (!metricsChart) {
|
||||||
// Generate datasets based on keys in metrics (excluding epoch)
|
initChart();
|
||||||
const datasets = [];
|
}
|
||||||
const colors = ['#f97316', '#10b981', '#3b82f6', '#eab308', '#a855f7'];
|
|
||||||
let colorIdx = 0;
|
|
||||||
|
|
||||||
for (const key in metrics) {
|
let labelIndex = metricsChart.data.labels.indexOf(epoch);
|
||||||
if (key !== 'epoch') {
|
if (labelIndex === -1) {
|
||||||
datasets.push({
|
metricsChart.data.labels.push(epoch);
|
||||||
|
labelIndex = metricsChart.data.labels.length - 1;
|
||||||
|
metricsChart.data.datasets.forEach(dataset => dataset.data.push(null));
|
||||||
|
}
|
||||||
|
|
||||||
|
const colors = ['#f97316', '#10b981', '#3b82f6', '#eab308', '#a855f7'];
|
||||||
|
Object.entries(metrics).forEach(([key, value]) => {
|
||||||
|
if (key === 'epoch') return;
|
||||||
|
|
||||||
|
let dataset = metricsChart.data.datasets.find(item => item.label === key);
|
||||||
|
if (!dataset) {
|
||||||
|
const color = colors[metricsChart.data.datasets.length % colors.length];
|
||||||
|
dataset = {
|
||||||
label: key,
|
label: key,
|
||||||
data: [],
|
data: Array(metricsChart.data.labels.length).fill(null),
|
||||||
borderColor: colors[colorIdx % colors.length],
|
borderColor: color,
|
||||||
backgroundColor: colors[colorIdx % colors.length] + '22',
|
backgroundColor: color + '22',
|
||||||
tension: 0.15,
|
tension: 0.15,
|
||||||
fill: false
|
fill: false
|
||||||
});
|
};
|
||||||
colorIdx++;
|
metricsChart.data.datasets.push(dataset);
|
||||||
}
|
|
||||||
}
|
|
||||||
initChart(datasets);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add label if not present
|
dataset.data[labelIndex] = value;
|
||||||
if (!metricsChart.data.labels.includes(epoch)) {
|
|
||||||
metricsChart.data.labels.push(epoch);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Push data to correct dataset
|
|
||||||
metricsChart.data.datasets.forEach(dataset => {
|
|
||||||
const val = metrics[dataset.label];
|
|
||||||
if (val !== undefined) {
|
|
||||||
dataset.data.push(val);
|
|
||||||
}
|
|
||||||
});
|
});
|
||||||
|
|
||||||
metricsChart.update();
|
metricsChart.update();
|
||||||
|
|
@ -432,7 +429,8 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
startBtn.disabled = true;
|
startBtn.disabled = true;
|
||||||
stopBtn.disabled = true;
|
stopBtn.disabled = true;
|
||||||
break;
|
break;
|
||||||
case 'finished':
|
case 'finished': // Compatibility with sessions created by older versions.
|
||||||
|
case 'succeeded':
|
||||||
statusTitle.textContent = 'ГОТОВО';
|
statusTitle.textContent = 'ГОТОВО';
|
||||||
statusText.textContent = 'Обучение успешно завершено.';
|
statusText.textContent = 'Обучение успешно завершено.';
|
||||||
isTrainingActive = false;
|
isTrainingActive = false;
|
||||||
|
|
@ -440,6 +438,14 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
stopBtn.disabled = true;
|
stopBtn.disabled = true;
|
||||||
stopTimer();
|
stopTimer();
|
||||||
break;
|
break;
|
||||||
|
case 'cancelled':
|
||||||
|
statusTitle.textContent = 'ОСТАНОВЛЕНО';
|
||||||
|
statusText.textContent = 'Обучение остановлено пользователем.';
|
||||||
|
isTrainingActive = false;
|
||||||
|
startBtn.disabled = false;
|
||||||
|
stopBtn.disabled = true;
|
||||||
|
stopTimer();
|
||||||
|
break;
|
||||||
case 'failed':
|
case 'failed':
|
||||||
statusTitle.textContent = 'ОШИБКА';
|
statusTitle.textContent = 'ОШИБКА';
|
||||||
statusText.textContent = 'Процесс завершился с ошибкой. Проверьте логи.';
|
statusText.textContent = 'Процесс завершился с ошибкой. Проверьте логи.';
|
||||||
|
|
@ -462,43 +468,51 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Read/Write Configurations ---
|
// --- Read/Write Configurations ---
|
||||||
|
function readNumber(id, fallback, integer = false) {
|
||||||
|
const rawValue = document.getElementById(id).value;
|
||||||
|
const value = integer
|
||||||
|
? Number.parseInt(rawValue, 10)
|
||||||
|
: Number.parseFloat(rawValue);
|
||||||
|
return Number.isNaN(value) ? fallback : value;
|
||||||
|
}
|
||||||
|
|
||||||
function getFormConfig() {
|
function getFormConfig() {
|
||||||
return {
|
return {
|
||||||
dataset: document.getElementById('dataset').value.trim(),
|
dataset: document.getElementById('dataset').value.trim(),
|
||||||
model: document.getElementById('model').value.trim(),
|
model: document.getElementById('model').value.trim(),
|
||||||
task: taskSelect.value,
|
task: taskSelect.value,
|
||||||
epochs: parseInt(document.getElementById('epochs').value) || 100,
|
epochs: readNumber('epochs', 100, true),
|
||||||
image_size: parseInt(document.getElementById('image-size').value) || 640,
|
image_size: readNumber('image-size', 640, true),
|
||||||
batch_size: parseInt(document.getElementById('batch-size').value) || 16,
|
batch_size: readNumber('batch-size', 16, true),
|
||||||
device: document.getElementById('device').value.trim(),
|
device: document.getElementById('device').value.trim(),
|
||||||
workers: parseInt(document.getElementById('workers').value) || 8,
|
workers: readNumber('workers', 8, true),
|
||||||
patience: parseInt(document.getElementById('patience').value) || 100,
|
patience: readNumber('patience', 100, true),
|
||||||
project: document.getElementById('project').value.trim() || 'runs/train',
|
project: document.getElementById('project').value.trim() || 'runs/train',
|
||||||
run_name: document.getElementById('run-name').value.trim(),
|
run_name: document.getElementById('run-name').value.trim(),
|
||||||
split: {
|
split: {
|
||||||
enabled: splitEnabled.checked,
|
enabled: splitEnabled.checked,
|
||||||
train_ratio: parseFloat(splitRatio.value) || 0.8,
|
train_ratio: readNumber('split-ratio', 0.8),
|
||||||
classes_path: splitClasses.value.trim()
|
classes_path: splitClasses.value.trim()
|
||||||
},
|
},
|
||||||
augmentation: {
|
augmentation: {
|
||||||
enabled: augmentationEnabled.checked,
|
enabled: augmentationEnabled.checked,
|
||||||
hsv_h: parseFloat(document.getElementById('hsv-h').value) || 0,
|
hsv_h: readNumber('hsv-h', 0.015),
|
||||||
hsv_s: parseFloat(document.getElementById('hsv-s').value) || 0,
|
hsv_s: readNumber('hsv-s', 0.7),
|
||||||
hsv_v: parseFloat(document.getElementById('hsv-v').value) || 0,
|
hsv_v: readNumber('hsv-v', 0.4),
|
||||||
degrees: parseFloat(document.getElementById('degrees').value) || 0,
|
degrees: readNumber('degrees', 0),
|
||||||
translate: parseFloat(document.getElementById('translate').value) || 0,
|
translate: readNumber('translate', 0.1),
|
||||||
scale: parseFloat(document.getElementById('scale').value) || 0,
|
scale: readNumber('scale', 0.5),
|
||||||
shear: parseFloat(document.getElementById('shear').value) || 0,
|
shear: readNumber('shear', 0),
|
||||||
perspective: parseFloat(document.getElementById('perspective').value) || 0,
|
perspective: readNumber('perspective', 0),
|
||||||
close_mosaic: parseInt(document.getElementById('close-mosaic').value) || 10,
|
close_mosaic: readNumber('close-mosaic', 10, true),
|
||||||
flipud: parseFloat(document.getElementById('flipud').value) || 0,
|
flipud: readNumber('flipud', 0),
|
||||||
fliplr: parseFloat(document.getElementById('fliplr').value) || 0,
|
fliplr: readNumber('fliplr', 0.5),
|
||||||
bgr: parseFloat(document.getElementById('bgr').value) || 0,
|
bgr: readNumber('bgr', 0),
|
||||||
mosaic: parseFloat(document.getElementById('mosaic').value) || 0,
|
mosaic: readNumber('mosaic', 1),
|
||||||
mixup: parseFloat(document.getElementById('mixup').value) || 0,
|
mixup: readNumber('mixup', 0),
|
||||||
cutmix: parseFloat(document.getElementById('cutmix').value) || 0,
|
cutmix: readNumber('cutmix', 0),
|
||||||
copy_paste: parseFloat(document.getElementById('copy-paste').value) || 0,
|
copy_paste: readNumber('copy-paste', 0),
|
||||||
erasing: parseFloat(document.getElementById('erasing').value) || 0,
|
erasing: readNumber('erasing', 0.4),
|
||||||
copy_paste_mode: document.getElementById('copy-paste-mode').value,
|
copy_paste_mode: document.getElementById('copy-paste-mode').value,
|
||||||
auto_augment: document.getElementById('auto-augment').value
|
auto_augment: document.getElementById('auto-augment').value
|
||||||
},
|
},
|
||||||
|
|
@ -541,17 +555,17 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Split
|
// Split
|
||||||
splitEnabled.checked = data.split?.enabled || false;
|
splitEnabled.checked = data.split?.enabled ?? false;
|
||||||
splitRatio.value = data.split?.train_ratio || 0.8;
|
splitRatio.value = data.split?.train_ratio ?? 0.8;
|
||||||
splitClasses.value = data.split?.classes_path || '';
|
splitClasses.value = data.split?.classes_path || '';
|
||||||
|
|
||||||
// Training params
|
// Training params
|
||||||
document.getElementById('epochs').value = data.epochs || 100;
|
document.getElementById('epochs').value = data.epochs ?? 100;
|
||||||
document.getElementById('image-size').value = data.image_size || 640;
|
document.getElementById('image-size').value = data.image_size ?? 640;
|
||||||
document.getElementById('batch-size').value = data.batch_size || 16;
|
document.getElementById('batch-size').value = data.batch_size ?? 16;
|
||||||
document.getElementById('device').value = data.device || '';
|
document.getElementById('device').value = data.device || '';
|
||||||
document.getElementById('workers').value = data.workers || 8;
|
document.getElementById('workers').value = data.workers ?? 8;
|
||||||
document.getElementById('patience').value = data.patience || 100;
|
document.getElementById('patience').value = data.patience ?? 100;
|
||||||
document.getElementById('project').value = data.project || 'runs/train';
|
document.getElementById('project').value = data.project || 'runs/train';
|
||||||
document.getElementById('run-name').value = data.run_name || '';
|
document.getElementById('run-name').value = data.run_name || '';
|
||||||
|
|
||||||
|
|
@ -796,6 +810,9 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
localStorage.removeItem('draft_config');
|
localStorage.removeItem('draft_config');
|
||||||
await loadSessionsList();
|
await loadSessionsList();
|
||||||
await loadInitialConfig();
|
await loadInitialConfig();
|
||||||
|
} catch (e) {
|
||||||
|
console.error('Delete profile error:', e);
|
||||||
|
showNotification('Не удалось удалить профиль.', 'error');
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|
@ -837,7 +854,6 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const configForm = document.getElementById('config-form');
|
|
||||||
if (configForm) {
|
if (configForm) {
|
||||||
configForm.addEventListener('input', () => {
|
configForm.addEventListener('input', () => {
|
||||||
const config = getFormConfig();
|
const config = getFormConfig();
|
||||||
|
|
|
||||||
|
|
@ -517,7 +517,9 @@ body {
|
||||||
#status-card.status-preparing { border-left-color: var(--warning); animation: pulsingBorder 2s infinite; }
|
#status-card.status-preparing { border-left-color: var(--warning); animation: pulsingBorder 2s infinite; }
|
||||||
#status-card.status-training { border-left-color: var(--success); }
|
#status-card.status-training { border-left-color: var(--success); }
|
||||||
#status-card.status-stopping { border-left-color: var(--warning); }
|
#status-card.status-stopping { border-left-color: var(--warning); }
|
||||||
#status-card.status-finished { border-left-color: var(--success); }
|
#status-card.status-finished,
|
||||||
|
#status-card.status-succeeded { border-left-color: var(--success); }
|
||||||
|
#status-card.status-cancelled { border-left-color: var(--warning); }
|
||||||
#status-card.status-failed { border-left-color: var(--error); }
|
#status-card.status-failed { border-left-color: var(--error); }
|
||||||
|
|
||||||
@keyframes pulsingBorder {
|
@keyframes pulsingBorder {
|
||||||
|
|
@ -550,7 +552,9 @@ body {
|
||||||
#status-card.status-preparing .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); animation: pulseDot 1s infinite; }
|
#status-card.status-preparing .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); animation: pulseDot 1s infinite; }
|
||||||
#status-card.status-training .status-dot { background-color: var(--success); box-shadow: 0 0 8px var(--success); animation: pulseDot 1.5s infinite; }
|
#status-card.status-training .status-dot { background-color: var(--success); box-shadow: 0 0 8px var(--success); animation: pulseDot 1.5s infinite; }
|
||||||
#status-card.status-stopping .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); }
|
#status-card.status-stopping .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); }
|
||||||
#status-card.status-finished .status-dot { background-color: var(--success); box-shadow: 0 0 8px var(--success); }
|
#status-card.status-finished .status-dot,
|
||||||
|
#status-card.status-succeeded .status-dot { background-color: var(--success); box-shadow: 0 0 8px var(--success); }
|
||||||
|
#status-card.status-cancelled .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); }
|
||||||
#status-card.status-failed .status-dot { background-color: var(--error); box-shadow: 0 0 8px var(--error); }
|
#status-card.status-failed .status-dot { background-color: var(--error); box-shadow: 0 0 8px var(--error); }
|
||||||
|
|
||||||
@keyframes pulseDot {
|
@keyframes pulseDot {
|
||||||
|
|
@ -569,7 +573,9 @@ body {
|
||||||
#status-card.status-preparing #status-title { color: var(--warning); }
|
#status-card.status-preparing #status-title { color: var(--warning); }
|
||||||
#status-card.status-training #status-title { color: var(--success); }
|
#status-card.status-training #status-title { color: var(--success); }
|
||||||
#status-card.status-stopping #status-title { color: var(--warning); }
|
#status-card.status-stopping #status-title { color: var(--warning); }
|
||||||
#status-card.status-finished #status-title { color: var(--success); }
|
#status-card.status-finished #status-title,
|
||||||
|
#status-card.status-succeeded #status-title { color: var(--success); }
|
||||||
|
#status-card.status-cancelled #status-title { color: var(--warning); }
|
||||||
#status-card.status-failed #status-title { color: var(--error); }
|
#status-card.status-failed #status-title { color: var(--error); }
|
||||||
|
|
||||||
.status-timer {
|
.status-timer {
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,9 @@ from typing import Any
|
||||||
|
|
||||||
from .config import MlflowConfig, TrainingConfig
|
from .config import MlflowConfig, TrainingConfig
|
||||||
|
|
||||||
|
# Restrict PyTorch checkpoint deserialization to Ultralytics' known model classes.
|
||||||
|
os.environ["ULTRALYTICS_SAFE_LOAD"] = "1"
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class TrainingEvent:
|
class TrainingEvent:
|
||||||
|
|
@ -58,6 +61,7 @@ class TrainingRunner:
|
||||||
self._model: Any | None = None
|
self._model: Any | None = None
|
||||||
self._state_lock = RLock()
|
self._state_lock = RLock()
|
||||||
self._stop_requested = Event()
|
self._stop_requested = Event()
|
||||||
|
self._force_stop_triggered = Event()
|
||||||
self._subprocess: Any | None = None
|
self._subprocess: Any | None = None
|
||||||
self._subprocess_ready = False
|
self._subprocess_ready = False
|
||||||
self._force_stop_timer: Timer | None = None
|
self._force_stop_timer: Timer | None = None
|
||||||
|
|
@ -69,6 +73,7 @@ class TrainingRunner:
|
||||||
self._force_stop_timer = None
|
self._force_stop_timer = None
|
||||||
self._subprocess_ready = False
|
self._subprocess_ready = False
|
||||||
self._stop_requested.clear()
|
self._stop_requested.clear()
|
||||||
|
self._force_stop_triggered.clear()
|
||||||
if timer is not None:
|
if timer is not None:
|
||||||
timer.cancel()
|
timer.cancel()
|
||||||
|
|
||||||
|
|
@ -119,6 +124,10 @@ class TrainingRunner:
|
||||||
def stop_requested(self) -> bool:
|
def stop_requested(self) -> bool:
|
||||||
return self._stop_requested.is_set()
|
return self._stop_requested.is_set()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def force_stop_triggered(self) -> bool:
|
||||||
|
return self._force_stop_triggered.is_set()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _send_cooperative_stop(process: Any) -> None:
|
def _send_cooperative_stop(process: Any) -> None:
|
||||||
try:
|
try:
|
||||||
|
|
@ -149,6 +158,7 @@ class TrainingRunner:
|
||||||
try:
|
try:
|
||||||
if process.poll() is None:
|
if process.poll() is None:
|
||||||
process.kill()
|
process.kill()
|
||||||
|
self._force_stop_triggered.set()
|
||||||
except (AttributeError, OSError, ProcessLookupError):
|
except (AttributeError, OSError, ProcessLookupError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
@ -227,7 +237,7 @@ class TrainingRunner:
|
||||||
def _on_train_end(self, on_event: EventHandler) -> Callable[[Any], None]:
|
def _on_train_end(self, on_event: EventHandler) -> Callable[[Any], None]:
|
||||||
def callback(trainer: Any) -> None:
|
def callback(trainer: Any) -> None:
|
||||||
if self._stop_requested.is_set():
|
if self._stop_requested.is_set():
|
||||||
on_event(TrainingEvent("warning", "Обучение остановлено пользователем."))
|
on_event(TrainingEvent("cancelled", "Обучение остановлено пользователем."))
|
||||||
else:
|
else:
|
||||||
on_event(TrainingEvent("success", "Ultralytics завершил обучение."))
|
on_event(TrainingEvent("success", "Ultralytics завершил обучение."))
|
||||||
|
|
||||||
|
|
|
||||||
172
tests/frontend_smoke.js
Normal file
172
tests/frontend_smoke.js
Normal file
|
|
@ -0,0 +1,172 @@
|
||||||
|
const assert = require('node:assert/strict');
|
||||||
|
const path = require('node:path');
|
||||||
|
|
||||||
|
class FakeElement {
|
||||||
|
constructor(id = '') {
|
||||||
|
this.id = id;
|
||||||
|
this.value = id === 'task' ? 'detect' : '';
|
||||||
|
this.checked = false;
|
||||||
|
this.disabled = false;
|
||||||
|
this.style = {};
|
||||||
|
this.children = [];
|
||||||
|
this.listeners = {};
|
||||||
|
this.className = '';
|
||||||
|
this.textContent = '';
|
||||||
|
this.scrollTop = 0;
|
||||||
|
this.scrollHeight = 0;
|
||||||
|
this.classList = {
|
||||||
|
add() {},
|
||||||
|
remove() {},
|
||||||
|
contains() { return false; }
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
addEventListener(name, handler) {
|
||||||
|
this.listeners[name] = handler;
|
||||||
|
}
|
||||||
|
|
||||||
|
appendChild(child) {
|
||||||
|
this.children.push(child);
|
||||||
|
this.lastElementChild = child;
|
||||||
|
return child;
|
||||||
|
}
|
||||||
|
|
||||||
|
getContext() {
|
||||||
|
return {};
|
||||||
|
}
|
||||||
|
|
||||||
|
remove() {}
|
||||||
|
|
||||||
|
set innerHTML(value) {
|
||||||
|
this._innerHTML = value;
|
||||||
|
this.children = [];
|
||||||
|
}
|
||||||
|
|
||||||
|
get innerHTML() {
|
||||||
|
return this._innerHTML || '';
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const elements = new Map();
|
||||||
|
const element = id => {
|
||||||
|
if (!elements.has(id)) elements.set(id, new FakeElement(id));
|
||||||
|
return elements.get(id);
|
||||||
|
};
|
||||||
|
|
||||||
|
let domReady;
|
||||||
|
global.document = {
|
||||||
|
addEventListener(name, handler) {
|
||||||
|
if (name === 'DOMContentLoaded') domReady = handler;
|
||||||
|
},
|
||||||
|
querySelectorAll() {
|
||||||
|
return [];
|
||||||
|
},
|
||||||
|
getElementById: element,
|
||||||
|
createElement() {
|
||||||
|
return new FakeElement();
|
||||||
|
},
|
||||||
|
head: new FakeElement('head'),
|
||||||
|
body: new FakeElement('body')
|
||||||
|
};
|
||||||
|
|
||||||
|
const storage = new Map();
|
||||||
|
global.localStorage = {
|
||||||
|
getItem(key) { return storage.has(key) ? storage.get(key) : null; },
|
||||||
|
setItem(key, value) { storage.set(key, value); },
|
||||||
|
removeItem(key) { storage.delete(key); }
|
||||||
|
};
|
||||||
|
global.confirm = () => true;
|
||||||
|
global.window = {location: {protocol: 'http:', host: '127.0.0.1:8000'}};
|
||||||
|
|
||||||
|
class FakeChart {
|
||||||
|
static instances = [];
|
||||||
|
|
||||||
|
constructor(_context, config) {
|
||||||
|
this.data = config.data;
|
||||||
|
this.options = config.options;
|
||||||
|
FakeChart.instances.push(this);
|
||||||
|
}
|
||||||
|
|
||||||
|
destroy() {}
|
||||||
|
update() {}
|
||||||
|
}
|
||||||
|
global.Chart = FakeChart;
|
||||||
|
|
||||||
|
class FakeWebSocket {
|
||||||
|
static instances = [];
|
||||||
|
|
||||||
|
constructor(url) {
|
||||||
|
this.url = url;
|
||||||
|
FakeWebSocket.instances.push(this);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
global.WebSocket = FakeWebSocket;
|
||||||
|
|
||||||
|
const response = (ok, data) => ({
|
||||||
|
ok,
|
||||||
|
async json() { return data; }
|
||||||
|
});
|
||||||
|
global.fetch = async url => {
|
||||||
|
if (url === '/api/datasets' || url === '/api/models' || url === '/api/sessions') {
|
||||||
|
return response(true, []);
|
||||||
|
}
|
||||||
|
if (url === '/api/sessions/last_run') return response(false, {});
|
||||||
|
if (url === '/api/config/defaults') {
|
||||||
|
return response(true, {
|
||||||
|
dataset: 'coco8.yaml',
|
||||||
|
model: 'yolo11n.pt',
|
||||||
|
task: 'detect',
|
||||||
|
workers: 8,
|
||||||
|
patience: 100,
|
||||||
|
augmentation: {enabled: true, close_mosaic: 10},
|
||||||
|
mlflow: {enabled: false},
|
||||||
|
split: {enabled: false}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return response(false, {});
|
||||||
|
};
|
||||||
|
|
||||||
|
require(path.resolve(__dirname, '../src/yolo_webui/static/app.js'));
|
||||||
|
|
||||||
|
async function flushPromises() {
|
||||||
|
await new Promise(resolve => setImmediate(resolve));
|
||||||
|
await new Promise(resolve => setImmediate(resolve));
|
||||||
|
}
|
||||||
|
|
||||||
|
(async () => {
|
||||||
|
assert.equal(typeof domReady, 'function');
|
||||||
|
domReady();
|
||||||
|
await flushPromises();
|
||||||
|
|
||||||
|
element('workers').value = '0';
|
||||||
|
element('patience').value = '0';
|
||||||
|
element('close-mosaic').value = '0';
|
||||||
|
element('config-form').listeners.input();
|
||||||
|
const savedConfig = JSON.parse(storage.get('draft_config'));
|
||||||
|
assert.equal(savedConfig.workers, 0);
|
||||||
|
assert.equal(savedConfig.patience, 0);
|
||||||
|
assert.equal(savedConfig.augmentation.close_mosaic, 0);
|
||||||
|
|
||||||
|
assert.equal(FakeWebSocket.instances.length, 1);
|
||||||
|
const socket = FakeWebSocket.instances[0];
|
||||||
|
socket.onmessage({
|
||||||
|
data: JSON.stringify({
|
||||||
|
type: 'init',
|
||||||
|
status: 'idle',
|
||||||
|
logs: [],
|
||||||
|
metrics: [
|
||||||
|
{epoch: 1, mAP50: 0.5},
|
||||||
|
{epoch: 2, loss: 0.2}
|
||||||
|
]
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
const chart = FakeChart.instances.at(-1);
|
||||||
|
assert.deepEqual(chart.data.labels, [1, 2]);
|
||||||
|
assert.deepEqual(chart.data.datasets.map(item => item.label), ['mAP50', 'loss']);
|
||||||
|
assert.deepEqual(chart.data.datasets[0].data, [0.5, null]);
|
||||||
|
assert.deepEqual(chart.data.datasets[1].data, [null, 0.2]);
|
||||||
|
})().catch(error => {
|
||||||
|
console.error(error);
|
||||||
|
process.exitCode = 1;
|
||||||
|
});
|
||||||
|
|
@ -1,8 +1,12 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import threading
|
||||||
|
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
from yolo_webui.app import app
|
from yolo_webui.app import TrainingManager, app
|
||||||
|
|
||||||
|
|
||||||
def test_get_config_defaults() -> None:
|
def test_get_config_defaults() -> None:
|
||||||
|
|
@ -88,3 +92,84 @@ def test_sessions_flow(monkeypatch, tmp_path) -> None:
|
||||||
# 7. Loading nonexistent session should return 404
|
# 7. Loading nonexistent session should return 404
|
||||||
response = client.get("/api/sessions/nonexistent")
|
response = client.get("/api/sessions/nonexistent")
|
||||||
assert response.status_code == 404
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
def test_reserved_session_name_cannot_be_overwritten(monkeypatch, tmp_path) -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
monkeypatch.setattr("yolo_webui.app.get_sessions_dir", lambda: tmp_path)
|
||||||
|
|
||||||
|
response = client.post("/api/sessions/last_run", json={"dataset": "data"})
|
||||||
|
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "зарезервировано" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_started_event_does_not_deadlock() -> None:
|
||||||
|
training_manager = TrainingManager()
|
||||||
|
training_manager.state.status = "preparing"
|
||||||
|
event = {
|
||||||
|
"kind": "started",
|
||||||
|
"message": "Обучение началось.",
|
||||||
|
"epoch": 0,
|
||||||
|
"total_epochs": 3,
|
||||||
|
}
|
||||||
|
worker = threading.Thread(
|
||||||
|
target=training_manager._handle_subprocess_line,
|
||||||
|
args=(f"__YOLO_WEBUI_EVENT__:{json.dumps(event)}",),
|
||||||
|
)
|
||||||
|
|
||||||
|
worker.start()
|
||||||
|
worker.join(timeout=1)
|
||||||
|
|
||||||
|
assert not worker.is_alive()
|
||||||
|
assert training_manager.state.status == "training"
|
||||||
|
|
||||||
|
|
||||||
|
def test_background_broadcast_uses_websocket_event_loop() -> None:
|
||||||
|
async def scenario() -> None:
|
||||||
|
training_manager = TrainingManager()
|
||||||
|
server_thread_id = threading.get_ident()
|
||||||
|
|
||||||
|
class FakeWebSocket:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.messages: list[str] = []
|
||||||
|
self.send_thread_ids: list[int] = []
|
||||||
|
self.sent = asyncio.Event()
|
||||||
|
|
||||||
|
async def send_text(self, payload: str) -> None:
|
||||||
|
self.messages.append(payload)
|
||||||
|
self.send_thread_ids.append(threading.get_ident())
|
||||||
|
self.sent.set()
|
||||||
|
|
||||||
|
websocket = FakeWebSocket()
|
||||||
|
training_manager.add_websocket(websocket) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
worker = threading.Thread(
|
||||||
|
target=training_manager.broadcast,
|
||||||
|
args=({"type": "status", "status": "training"},),
|
||||||
|
)
|
||||||
|
worker.start()
|
||||||
|
worker.join(timeout=1)
|
||||||
|
assert not worker.is_alive()
|
||||||
|
|
||||||
|
await asyncio.wait_for(websocket.sent.wait(), timeout=1)
|
||||||
|
assert json.loads(websocket.messages[0])["status"] == "training"
|
||||||
|
assert websocket.send_thread_ids == [server_thread_id]
|
||||||
|
|
||||||
|
asyncio.run(scenario())
|
||||||
|
|
||||||
|
|
||||||
|
def test_process_result_distinguishes_success_cancellation_and_failure() -> None:
|
||||||
|
succeeded = TrainingManager()
|
||||||
|
succeeded._finalize_process_result(0)
|
||||||
|
assert succeeded.state.status == "succeeded"
|
||||||
|
|
||||||
|
cancelled = TrainingManager()
|
||||||
|
cancelled.state.stop_requested = True
|
||||||
|
cancelled._finalize_process_result(0)
|
||||||
|
assert cancelled.state.status == "cancelled"
|
||||||
|
|
||||||
|
failed_after_stop = TrainingManager()
|
||||||
|
failed_after_stop.state.stop_requested = True
|
||||||
|
failed_after_stop._finalize_process_result(1)
|
||||||
|
assert failed_after_stop.state.status == "failed"
|
||||||
|
|
|
||||||
|
|
@ -108,3 +108,47 @@ def test_classification_rejects_detection_style_auto_split() -> None:
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="classify"):
|
with pytest.raises(ValueError, match="classify"):
|
||||||
config.validate()
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("field", "value"),
|
||||||
|
[
|
||||||
|
("dataset", "https://example.invalid/dataset.yaml"),
|
||||||
|
("model", "https://example.invalid/model.pt"),
|
||||||
|
("project", "https://example.invalid/results"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_training_rejects_remote_references(field: str, value: str) -> None:
|
||||||
|
values = {
|
||||||
|
"dataset": "dataset.yaml",
|
||||||
|
"model": "model.pt",
|
||||||
|
"project": "runs/train",
|
||||||
|
field: value,
|
||||||
|
}
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="не может быть URL"):
|
||||||
|
TrainingConfig(**values).validate()
|
||||||
|
|
||||||
|
|
||||||
|
def test_model_path_must_stay_in_allowed_roots(tmp_path: Path, monkeypatch) -> None:
|
||||||
|
workspace = tmp_path / "workspace"
|
||||||
|
workspace.mkdir()
|
||||||
|
external_model = tmp_path / "external" / "model.pt"
|
||||||
|
monkeypatch.chdir(workspace)
|
||||||
|
config = TrainingConfig(dataset="dataset.yaml", model=str(external_model))
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="разрешённом каталоге"):
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
monkeypatch.setenv("YOLO_WEBUI_MODEL_ROOTS", str(external_model.parent))
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
def test_existing_bare_dataset_cannot_bypass_allowed_roots(
|
||||||
|
tmp_path: Path, monkeypatch
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
(tmp_path / "private.yaml").write_text("secret: value\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="вне разрешённого каталога"):
|
||||||
|
TrainingConfig(dataset="private.yaml", model="model.pt").validate()
|
||||||
|
|
|
||||||
21
tests/test_frontend.py
Normal file
21
tests/test_frontend.py
Normal file
|
|
@ -0,0 +1,21 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
NODE = shutil.which("node")
|
||||||
|
APP_JS = Path("src/yolo_webui/static/app.js")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
||||||
|
def test_frontend_javascript_syntax() -> None:
|
||||||
|
subprocess.run([NODE, "--check", str(APP_JS)], check=True)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
||||||
|
def test_frontend_zero_values_and_dynamic_chart_series() -> None:
|
||||||
|
subprocess.run([NODE, "tests/frontend_smoke.js"], check=True)
|
||||||
|
|
@ -225,3 +225,32 @@ def test_split_dataset_preserves_custom_yaml_keys(tmp_path: Path) -> None:
|
||||||
assert data["kpt_shape"] == [5, 3]
|
assert data["kpt_shape"] == [5, 3]
|
||||||
assert data["flip_idx"] == [0, 2, 1, 4, 3]
|
assert data["flip_idx"] == [0, 2, 1, 4, 3]
|
||||||
assert data["names"] == {0: "person"}
|
assert data["names"] == {0: "person"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_explicit_text_classes_override_root_yaml_names(tmp_path: Path) -> None:
|
||||||
|
images_dir = tmp_path / "images"
|
||||||
|
labels_dir = tmp_path / "labels"
|
||||||
|
images_dir.mkdir()
|
||||||
|
labels_dir.mkdir()
|
||||||
|
for index in range(2):
|
||||||
|
(images_dir / f"image-{index}.jpg").write_bytes(b"")
|
||||||
|
(labels_dir / f"image-{index}.txt").write_text(
|
||||||
|
"0 0.5 0.5 0.2 0.2\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
(tmp_path / "dataset.yaml").write_text(
|
||||||
|
yaml.safe_dump({"names": {0: "old"}}),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
classes_path = tmp_path / "custom.txt"
|
||||||
|
classes_path.write_text("new\n", encoding="utf-8")
|
||||||
|
|
||||||
|
_, _, output_yaml = split_dataset(
|
||||||
|
str(tmp_path),
|
||||||
|
0.5,
|
||||||
|
str(classes_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
data = yaml.safe_load(Path(output_yaml).read_text(encoding="utf-8"))
|
||||||
|
assert data["names"] == {0: "new"}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue