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 | Приоритет | Статус | Исправление |
|
||||
|---|---|---|---|
|
||||
| BUG-001 | Критический | Исправлено | `subprocess_runner.main()` возвращает код, а `SystemExit` создаётся только снаружи обрабатывающего блока |
|
||||
| BUG-002 | Высокий | Исправлено | Перед каждым запуском `prepare_run()` сбрасывает состояние остановки |
|
||||
| BUG-003 | Высокий | Исправлено | Родитель отправляет кооперативный сигнал; принудительный `kill()` используется только после таймаута |
|
||||
| BUG-004 | Высокий | Исправлено | Запрос, сделанный до готовности subprocess, сохраняется и доставляется после маркера `READY` |
|
||||
| BUG-005 | Высокий | Исправлено | Detection-style авторазбиение запрещено для `classify` в UI и конфигурации |
|
||||
| BUG-006 | Средний | Исправлено | `.yaml`/`.yml` разбираются через `yaml.safe_load()`, поле `names` валидируется |
|
||||
| BUG-007 | Средний | Исправлено | Датасет с одним изображением отклоняется с понятной ошибкой |
|
||||
| BUG-008 | Средний | Исправлено | Изображения и метки ищутся рекурсивно с сохранением вложенных путей |
|
||||
| BUG-009 | Средний | Исправлено | Каждый результат создаётся в уникальном `.yolo-tui/splits/<id>` без перезаписи пользовательского `split/` |
|
||||
| BUG-010 | Низкий | Исправлено | Traceback выводится в журнал TUI; абсолютный путь другого пользователя удалён |
|
||||
| BUG-011 | Низкий | Исправлено | Явно указанный отсутствующий или некорректный файл классов вызывает точную ошибку без fallback |
|
||||
| BUG-001 | Критический | Исправлено | Закрыт `try/catch`, удалено повторное объявление `configForm`, добавлен `node --check` в тесты |
|
||||
| BUG-002 | Критический | Исправлено | Status broadcast вынесен за пределы `threading.Lock`; добавлен тест на отсутствие deadlock |
|
||||
| SEC-001 | Критический при сетевой публикации | Исправлено | Compose публикует loopback, URL запрещены, пути ограничены доверенными корнями, restricted checkpoint loading включён |
|
||||
| BUG-003 | Высокий | Исправлено | Все WebSocket send выполняются в ASGI loop через `run_coroutine_threadsafe`; ошибки логируются, сломанные сокеты удаляются |
|
||||
| BUG-004 | Высокий | Исправлено | Введены состояния `succeeded`, `cancelled`, `failed`; ошибка после stop больше не маскируется как отмена |
|
||||
| DOC-001 | Высокий | Исправлено | README полностью обновлён для WebUI, актуальных CLI-команд, Docker и модели безопасности |
|
||||
| BUG-005 | Средний | Исправлено | Провалидированные классы всегда записываются в итоговый YAML и имеют приоритет над случайным корневым YAML |
|
||||
| BUG-006 | Средний | Исправлено | Числа разбираются с проверкой `Number.isNaN`; нули сохраняются при чтении и восстановлении формы |
|
||||
| BUG-007 | Средний | Исправлено | Серии графика добавляются динамически и выравниваются по эпохам, включая новые ключи метрик |
|
||||
| BUILD-001 | Средний | Исправлено | Docker устанавливает frozen-набор из `uv.lock`; версия `uv` также зафиксирована |
|
||||
|
||||
## Жизненный цикл обучения
|
||||
## Жизненный цикл и WebSocket
|
||||
|
||||
- Дочерний процесс устанавливает обработчики остановки и только затем печатает
|
||||
`__YOLO_TUI_READY__`.
|
||||
- Если пользователь нажал «Остановить» раньше, родитель запоминает запрос и
|
||||
отправляет его после получения маркера готовности.
|
||||
- Дочерний `TrainingRunner` устанавливает `trainer.stop = True`; Ultralytics
|
||||
останавливается между пакетами данных, затем выполняет штатную финализацию и
|
||||
завершающие callbacks.
|
||||
- Если процесс не завершился за 30 секунд, используется принудительный fallback.
|
||||
- После завершения ссылка на subprocess и таймер очищаются; перед следующим
|
||||
запуском флаг остановки сбрасывается.
|
||||
- Событие `started` меняет состояние под lock, но отправляет статус только после
|
||||
освобождения lock.
|
||||
- Event loop запоминается при подключении WebSocket. Вызовы из фонового потока
|
||||
передаются в него через `asyncio.run_coroutine_threadsafe()`.
|
||||
- Отправки сериализуются `asyncio.Lock`, поэтому сообщения одного запуска сохраняют
|
||||
порядок. Ошибка доставки попадает в журнал, а нерабочий клиент удаляется.
|
||||
- Финальная классификация учитывает return code, stop-флаг, последнее
|
||||
структурированное событие и факт принудительной остановки.
|
||||
- Штатная кооперативная остановка даёт `cancelled`; ненулевой код после stop без
|
||||
подтверждённой отмены даёт `failed`.
|
||||
|
||||
## Работа с датасетами
|
||||
## Безопасность
|
||||
|
||||
- Текущий splitter предназначен для задач `detect`, `segment`, `pose` и `obb`
|
||||
со структурой `images/` + `labels/`.
|
||||
- Для `classify` требуется готовый каталог с `train`/`val` и подкаталогами
|
||||
классов; несовместимый переключатель в UI отключён.
|
||||
- Поддерживаются `classes.txt`, `.yaml` и `.yml`; YAML может хранить `names` как
|
||||
список или словарь с последовательными ID от 0.
|
||||
- Явный путь к классам считается обязательным и не заменяется автопоиском при
|
||||
опечатке или ошибке формата.
|
||||
- Split требует минимум два изображения и рекурсивно обрабатывает вложенные
|
||||
каталоги.
|
||||
- Файлы каждого запуска создаются эксклюзивно в отдельном управляемом каталоге.
|
||||
- `docker-compose.yml` публикует `127.0.0.1:8000:8000`.
|
||||
- Dataset, model и project не принимают URL.
|
||||
- Локальные пути ограничены `datasets`, `models` и `runs`; дополнительные доверенные
|
||||
корни задаются переменными `YOLO_WEBUI_DATA_ROOTS`,
|
||||
`YOLO_WEBUI_MODEL_ROOTS`, `YOLO_WEBUI_RUN_ROOTS`.
|
||||
- Проверка использует разрешённые абсолютные пути после `resolve()`, поэтому
|
||||
symlink/`..` не позволяют выйти из доверенного корня.
|
||||
- Имена профилей валидируются на сервере, а `last_run` нельзя перезаписать через
|
||||
публичный 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
|
||||
uv run pytest -q
|
||||
uv run python -m compileall -q src tests
|
||||
git diff --check
|
||||
uv run pytest -q -> 49 passed, 1 warning
|
||||
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 -> успешно, host_ip=127.0.0.1
|
||||
git diff --check -> успешно
|
||||
```
|
||||
|
||||
Результат: `38 passed`; ошибок компиляции и форматирования diff нет.
|
||||
Полная сборка Docker-образа локально не запускалась: Docker daemon недоступен.
|
||||
Конфигурация Compose проверена отдельно, а соответствие lock-файла — через
|
||||
`uv lock --check`.
|
||||
|
||||
Регрессионные тесты проверяют:
|
||||
|
||||
1. успешный и ошибочный коды `subprocess_runner.main()`;
|
||||
2. сброс остановки между запусками;
|
||||
3. кооперативный сигнал вместо немедленного `terminate()`;
|
||||
4. доставку раннего запроса после готовности subprocess;
|
||||
5. запрет авторазбиения для `classify`;
|
||||
6. пользовательские YAML-файлы классов и ошибочный явный путь;
|
||||
7. датасеты из одного и двух изображений;
|
||||
8. вложенные изображения и метки;
|
||||
9. сохранность пользовательского каталога `split/` и уникальность результатов.
|
||||
Оставшееся предупреждение pytest относится к deprecated-связке
|
||||
`fastapi.testclient`/`starlette.testclient` с `httpx`; оно не связано с исправленными
|
||||
дефектами и не ломает тесты.
|
||||
|
|
|
|||
35
Dockerfile
35
Dockerfile
|
|
@ -1,8 +1,5 @@
|
|||
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
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
|
|
@ -12,33 +9,29 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
|||
git \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install uv for fast dependency resolution using pip (avoids ghcr.io network issues)
|
||||
RUN pip install --no-cache-dir uv
|
||||
# Pin the installer as well as application dependencies.
|
||||
RUN pip install --no-cache-dir uv==0.10.6
|
||||
|
||||
# Set working directory
|
||||
WORKDIR /workspace
|
||||
|
||||
# Copy dependency definition
|
||||
COPY pyproject.toml ./
|
||||
ENV UV_COMPILE_BYTECODE=1 \
|
||||
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
|
||||
# and fetch the correct PyTorch package based on the target DEVICE (CPU or GPU)
|
||||
# Install the exact dependency set recorded in uv.lock. Keeping the project out of
|
||||
# this layer allows dependency caching while source files change.
|
||||
COPY pyproject.toml uv.lock README.md ./
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
if [ "$DEVICE" = "cpu" ]; then \
|
||||
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
|
||||
uv sync --locked --no-dev --no-install-project
|
||||
|
||||
# Copy source code and files
|
||||
# Copy source code and install the project without re-resolving dependencies.
|
||||
COPY src ./src
|
||||
COPY README.md ./
|
||||
|
||||
# Install the project itself without re-installing dependencies
|
||||
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 8000
|
||||
|
|
|
|||
110
README.md
110
README.md
|
|
@ -1,59 +1,117 @@
|
|||
# YOLO Train TUI
|
||||
# YOLO Train WebUI
|
||||
|
||||
Терминальный интерфейс для обучения моделей Ultralytics YOLO с автоматической
|
||||
регистрацией параметров, метрик и артефактов в MLflow.
|
||||
Локальный веб-интерфейс для обучения моделей Ultralytics YOLO с журналом,
|
||||
графиками метрик, мягкой остановкой и интеграцией MLflow.
|
||||
|
||||
## Возможности
|
||||
|
||||
- задачи `detect`, `segment`, `classify`, `pose` и `obb`;
|
||||
- локальные пути, YAML-конфигурации и официальные имена моделей/датасетов;
|
||||
- локальные датасеты и официальные имена моделей Ultralytics;
|
||||
- настройка эпох, размера изображения, batch, устройства, workers и patience;
|
||||
- настройка цветовых и геометрических аугментаций, flip, Mosaic, MixUp,
|
||||
CutMix, copy-paste, erasing и AutoAugment;
|
||||
- обучение в фоновом потоке, прогресс по эпохам, журнал и мягкая остановка;
|
||||
- встроенная интеграция Ultralytics ↔ MLflow;
|
||||
- локальное MLflow-хранилище по умолчанию или внешний tracking server.
|
||||
- цветовые и геометрические аугментации, Mosaic, MixUp, CutMix, copy-paste,
|
||||
erasing и AutoAugment;
|
||||
- live-прогресс, журнал, графики метрик и восстановление состояния после
|
||||
переподключения браузера;
|
||||
- сохранение профилей запуска и кооперативная остановка обучения;
|
||||
- локальное MLflow-хранилище или внешний tracking server.
|
||||
|
||||
## Установка и запуск
|
||||
## Локальная установка и запуск
|
||||
|
||||
Нужны Python 3.11+ и [uv](https://docs.astral.sh/uv/).
|
||||
|
||||
```bash
|
||||
uv sync
|
||||
uv run yolo-train-tui
|
||||
uv sync --locked
|
||||
uv run yolo-train-webui
|
||||
```
|
||||
|
||||
Также приложение можно запустить как модуль:
|
||||
Альтернативный запуск как Python-модуля:
|
||||
|
||||
```bash
|
||||
uv run -m yolo_tui
|
||||
uv run -m yolo_webui
|
||||
```
|
||||
|
||||
При первом использовании официального имени модели (например, `yolo11n.pt`)
|
||||
Ultralytics автоматически скачает веса. Для полностью локальной работы укажите
|
||||
путь к уже загруженному `.pt` или `.yaml` файлу.
|
||||
Откройте `http://127.0.0.1:8000`. Сервер по умолчанию слушает только loopback.
|
||||
|
||||
При первом использовании официального имени модели, например `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-файлу датасета.
|
||||
Для `classify` укажите каталог с подкаталогами `train`, `test`/`val`, внутри
|
||||
которых изображения разложены по классам.
|
||||
Для `detect`, `segment`, `pose` и `obb` укажите YAML-файл либо каталог со структурой
|
||||
`images/` + `labels/`. WebUI может детерминированно разделить такой каталог на
|
||||
train/val. Для `classify` нужен готовый каталог с `train` и `val`/`test`, внутри
|
||||
которых изображения разложены по классам; автоматическое detection-style разбиение
|
||||
для этой задачи отключено.
|
||||
|
||||
## MLflow
|
||||
|
||||
По умолчанию метаданные записываются в локальную SQLite-базу `./mlflow.db`,
|
||||
сервер для обучения не требуется. Артефакты сохраняются локально средствами MLflow.
|
||||
Открыть интерфейс просмотра:
|
||||
По умолчанию метаданные записываются в `./mlflow.db`. Открыть интерфейс просмотра:
|
||||
|
||||
```bash
|
||||
uv run mlflow ui --backend-store-uri sqlite:///mlflow.db
|
||||
```
|
||||
|
||||
Затем откройте `http://127.0.0.1:5000`. Для удаленного MLflow-сервера включите
|
||||
MLflow в TUI и замените Tracking URI на адрес вида `http://mlflow.example:5000`.
|
||||
Затем откройте `http://127.0.0.1:5000`. Для внешнего tracking server укажите его URI
|
||||
в настройках 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
|
||||
uv run pytest
|
||||
uv run pytest -q
|
||||
node --check src/yolo_webui/static/app.js
|
||||
docker compose config
|
||||
```
|
||||
|
||||
Ultralytics распространяется по лицензии AGPL-3.0; для закрытых коммерческих
|
||||
|
|
|
|||
|
|
@ -2,11 +2,10 @@ services:
|
|||
webui:
|
||||
build:
|
||||
context: .
|
||||
args:
|
||||
- DEVICE=cpu # 'cpu' for Mac, change to 'gpu' on a Linux server with NVIDIA GPU
|
||||
image: yolo-train-webui:latest
|
||||
ports:
|
||||
- "8000:8000"
|
||||
# The training API has no built-in user accounts, so expose it locally only.
|
||||
- "127.0.0.1:8000:8000"
|
||||
volumes:
|
||||
- ./datasets:/workspace/datasets
|
||||
- ./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
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
|
@ -23,17 +25,19 @@ from yolo_webui.trainer import TrainingRunner
|
|||
# Set up logging
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||
logger = logging.getLogger("yolo_webui")
|
||||
SESSION_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LiveState:
|
||||
status: str = "idle" # idle, preparing, training, stopping, finished, failed
|
||||
status: str = "idle" # idle, preparing, training, stopping, succeeded, cancelled, failed
|
||||
epoch: int = 0
|
||||
total_epochs: int = 0
|
||||
logs: list[str] = field(default_factory=list)
|
||||
metrics: list[dict[str, Any]] = field(default_factory=list)
|
||||
output_dir: str | None = None
|
||||
stop_requested: bool = False
|
||||
last_event_kind: str | None = None
|
||||
|
||||
def reset(self) -> None:
|
||||
self.status = "idle"
|
||||
|
|
@ -43,6 +47,7 @@ class LiveState:
|
|||
self.metrics = []
|
||||
self.output_dir = None
|
||||
self.stop_requested = False
|
||||
self.last_event_kind = None
|
||||
|
||||
|
||||
class TrainingManager:
|
||||
|
|
@ -54,10 +59,16 @@ class TrainingManager:
|
|||
self.active_websockets: set[WebSocket] = set()
|
||||
self._lock = threading.Lock()
|
||||
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:
|
||||
loop = asyncio.get_running_loop()
|
||||
with self._lock:
|
||||
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:
|
||||
with self._lock:
|
||||
|
|
@ -65,27 +76,64 @@ class TrainingManager:
|
|||
|
||||
def broadcast(self, data: dict[str, Any]) -> None:
|
||||
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:
|
||||
sockets = list(self.active_websockets)
|
||||
|
||||
# Send outside lock to prevent blocking
|
||||
for ws in sockets:
|
||||
failed: list[WebSocket] = []
|
||||
for websocket in sockets:
|
||||
try:
|
||||
import asyncio
|
||||
# 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))
|
||||
await websocket.send_text(payload)
|
||||
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:
|
||||
log_entry = f"__LOG_LEVEL_{level.upper()}__:{text}"
|
||||
|
|
@ -98,7 +146,9 @@ class TrainingManager:
|
|||
|
||||
def start_training(self, config: TrainingConfig) -> None:
|
||||
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("Обучение уже выполняется.")
|
||||
|
||||
self.state.reset()
|
||||
|
|
@ -149,7 +199,7 @@ class TrainingManager:
|
|||
parts = message.split(" · ")[1:]
|
||||
for p in parts:
|
||||
if "=" in p:
|
||||
k, v = p.split("=")
|
||||
k, v = p.split("=", 1)
|
||||
try:
|
||||
metrics_dict[k.strip()] = float(v.strip())
|
||||
except ValueError:
|
||||
|
|
@ -159,10 +209,15 @@ class TrainingManager:
|
|||
with self._lock:
|
||||
self.state.metrics.append(metrics_dict)
|
||||
|
||||
status_update = None
|
||||
with self._lock:
|
||||
self.state.last_event_kind = kind
|
||||
if kind == "started" and self.state.status == "preparing":
|
||||
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.broadcast({
|
||||
|
|
@ -217,24 +272,7 @@ class TrainingManager:
|
|||
process.wait()
|
||||
rc = process.returncode
|
||||
self.runner.clear_subprocess()
|
||||
|
||||
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")
|
||||
self._finalize_process_result(rc)
|
||||
|
||||
except Exception as exc:
|
||||
logger.exception("Error in training process thread:")
|
||||
|
|
@ -257,7 +295,38 @@ class TrainingManager:
|
|||
except Exception:
|
||||
pass
|
||||
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()
|
||||
|
|
@ -286,6 +355,17 @@ def get_sessions_dir() -> 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")
|
||||
async def get_defaults():
|
||||
# Return defaults by instantiating with dummy paths and serializing
|
||||
|
|
@ -350,8 +430,7 @@ async def list_models():
|
|||
|
||||
@app.get("/api/sessions/{name}")
|
||||
async def load_session(name: str):
|
||||
sessions_dir = get_sessions_dir()
|
||||
file_path = sessions_dir / f"{name}.json"
|
||||
file_path = get_session_path(name)
|
||||
if not file_path.exists():
|
||||
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
||||
try:
|
||||
|
|
@ -363,8 +442,7 @@ async def load_session(name: str):
|
|||
|
||||
@app.post("/api/sessions/{name}")
|
||||
async def save_session(name: str, config_data: dict[str, Any]):
|
||||
sessions_dir = get_sessions_dir()
|
||||
file_path = sessions_dir / f"{name}.json"
|
||||
file_path = get_session_path(name, allow_last_run=False)
|
||||
try:
|
||||
with file_path.open("w", encoding="utf-8") as f:
|
||||
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}")
|
||||
async def delete_session(name: str):
|
||||
sessions_dir = get_sessions_dir()
|
||||
file_path = sessions_dir / f"{name}.json"
|
||||
file_path = get_session_path(name, allow_last_run=False)
|
||||
if not file_path.exists():
|
||||
raise HTTPException(status_code=404, detail="Сессия не найдена.")
|
||||
try:
|
||||
|
|
@ -448,15 +525,16 @@ async def websocket_endpoint(websocket: WebSocket):
|
|||
# We format log items for the UI
|
||||
"logs": [log.split(":", 1) for log in manager.state.logs if ":" in log],
|
||||
}
|
||||
await websocket.send_text(json.dumps(state_dict))
|
||||
|
||||
try:
|
||||
await websocket.send_text(json.dumps(state_dict))
|
||||
while True:
|
||||
# Keep connection alive; discard incoming messages
|
||||
await websocket.receive_text()
|
||||
except WebSocketDisconnect:
|
||||
manager.remove_websocket(websocket)
|
||||
pass
|
||||
except Exception:
|
||||
logger.warning("WebSocket connection failed", exc_info=True)
|
||||
finally:
|
||||
manager.remove_websocket(websocket)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,10 @@
|
|||
from __future__ import annotations
|
||||
|
||||
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"]
|
||||
|
|
@ -20,6 +23,63 @@ SUPPORTED_AUTO_AUGMENT_POLICIES: tuple[AutoAugmentPolicy, ...] = (
|
|||
"augmix",
|
||||
)
|
||||
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)
|
||||
|
|
@ -155,6 +215,37 @@ class TrainingConfig:
|
|||
raise ValueError("Укажите путь или имя датасета.")
|
||||
if not self.model.strip():
|
||||
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:
|
||||
raise ValueError(f"Неизвестный тип задачи: {self.task}.")
|
||||
if self.task == "classify" and self.split.enabled:
|
||||
|
|
@ -198,7 +289,6 @@ class TrainingConfig:
|
|||
|
||||
@property
|
||||
def resolved_model(self) -> str:
|
||||
from pathlib import Path
|
||||
model_path = self.model.strip()
|
||||
if "/" not in model_path and "\\" not in model_path:
|
||||
# Ensure models directory exists inside workspace
|
||||
|
|
|
|||
|
|
@ -239,7 +239,8 @@ def split_dataset(
|
|||
"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
|
||||
|
||||
_write_new(
|
||||
|
|
|
|||
|
|
@ -259,38 +259,35 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
|
||||
function updateChart(epoch, metrics) {
|
||||
if (!metricsChart) {
|
||||
// Generate datasets based on keys in metrics (excluding epoch)
|
||||
const datasets = [];
|
||||
const colors = ['#f97316', '#10b981', '#3b82f6', '#eab308', '#a855f7'];
|
||||
let colorIdx = 0;
|
||||
initChart();
|
||||
}
|
||||
|
||||
for (const key in metrics) {
|
||||
if (key !== 'epoch') {
|
||||
datasets.push({
|
||||
let labelIndex = metricsChart.data.labels.indexOf(epoch);
|
||||
if (labelIndex === -1) {
|
||||
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,
|
||||
data: [],
|
||||
borderColor: colors[colorIdx % colors.length],
|
||||
backgroundColor: colors[colorIdx % colors.length] + '22',
|
||||
data: Array(metricsChart.data.labels.length).fill(null),
|
||||
borderColor: color,
|
||||
backgroundColor: color + '22',
|
||||
tension: 0.15,
|
||||
fill: false
|
||||
});
|
||||
colorIdx++;
|
||||
}
|
||||
}
|
||||
initChart(datasets);
|
||||
};
|
||||
metricsChart.data.datasets.push(dataset);
|
||||
}
|
||||
|
||||
// Add label if not present
|
||||
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);
|
||||
}
|
||||
dataset.data[labelIndex] = value;
|
||||
});
|
||||
|
||||
metricsChart.update();
|
||||
|
|
@ -432,7 +429,8 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
startBtn.disabled = true;
|
||||
stopBtn.disabled = true;
|
||||
break;
|
||||
case 'finished':
|
||||
case 'finished': // Compatibility with sessions created by older versions.
|
||||
case 'succeeded':
|
||||
statusTitle.textContent = 'ГОТОВО';
|
||||
statusText.textContent = 'Обучение успешно завершено.';
|
||||
isTrainingActive = false;
|
||||
|
|
@ -440,6 +438,14 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
stopBtn.disabled = true;
|
||||
stopTimer();
|
||||
break;
|
||||
case 'cancelled':
|
||||
statusTitle.textContent = 'ОСТАНОВЛЕНО';
|
||||
statusText.textContent = 'Обучение остановлено пользователем.';
|
||||
isTrainingActive = false;
|
||||
startBtn.disabled = false;
|
||||
stopBtn.disabled = true;
|
||||
stopTimer();
|
||||
break;
|
||||
case 'failed':
|
||||
statusTitle.textContent = 'ОШИБКА';
|
||||
statusText.textContent = 'Процесс завершился с ошибкой. Проверьте логи.';
|
||||
|
|
@ -462,43 +468,51 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
}
|
||||
|
||||
// --- 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() {
|
||||
return {
|
||||
dataset: document.getElementById('dataset').value.trim(),
|
||||
model: document.getElementById('model').value.trim(),
|
||||
task: taskSelect.value,
|
||||
epochs: parseInt(document.getElementById('epochs').value) || 100,
|
||||
image_size: parseInt(document.getElementById('image-size').value) || 640,
|
||||
batch_size: parseInt(document.getElementById('batch-size').value) || 16,
|
||||
epochs: readNumber('epochs', 100, true),
|
||||
image_size: readNumber('image-size', 640, true),
|
||||
batch_size: readNumber('batch-size', 16, true),
|
||||
device: document.getElementById('device').value.trim(),
|
||||
workers: parseInt(document.getElementById('workers').value) || 8,
|
||||
patience: parseInt(document.getElementById('patience').value) || 100,
|
||||
workers: readNumber('workers', 8, true),
|
||||
patience: readNumber('patience', 100, true),
|
||||
project: document.getElementById('project').value.trim() || 'runs/train',
|
||||
run_name: document.getElementById('run-name').value.trim(),
|
||||
split: {
|
||||
enabled: splitEnabled.checked,
|
||||
train_ratio: parseFloat(splitRatio.value) || 0.8,
|
||||
train_ratio: readNumber('split-ratio', 0.8),
|
||||
classes_path: splitClasses.value.trim()
|
||||
},
|
||||
augmentation: {
|
||||
enabled: augmentationEnabled.checked,
|
||||
hsv_h: parseFloat(document.getElementById('hsv-h').value) || 0,
|
||||
hsv_s: parseFloat(document.getElementById('hsv-s').value) || 0,
|
||||
hsv_v: parseFloat(document.getElementById('hsv-v').value) || 0,
|
||||
degrees: parseFloat(document.getElementById('degrees').value) || 0,
|
||||
translate: parseFloat(document.getElementById('translate').value) || 0,
|
||||
scale: parseFloat(document.getElementById('scale').value) || 0,
|
||||
shear: parseFloat(document.getElementById('shear').value) || 0,
|
||||
perspective: parseFloat(document.getElementById('perspective').value) || 0,
|
||||
close_mosaic: parseInt(document.getElementById('close-mosaic').value) || 10,
|
||||
flipud: parseFloat(document.getElementById('flipud').value) || 0,
|
||||
fliplr: parseFloat(document.getElementById('fliplr').value) || 0,
|
||||
bgr: parseFloat(document.getElementById('bgr').value) || 0,
|
||||
mosaic: parseFloat(document.getElementById('mosaic').value) || 0,
|
||||
mixup: parseFloat(document.getElementById('mixup').value) || 0,
|
||||
cutmix: parseFloat(document.getElementById('cutmix').value) || 0,
|
||||
copy_paste: parseFloat(document.getElementById('copy-paste').value) || 0,
|
||||
erasing: parseFloat(document.getElementById('erasing').value) || 0,
|
||||
hsv_h: readNumber('hsv-h', 0.015),
|
||||
hsv_s: readNumber('hsv-s', 0.7),
|
||||
hsv_v: readNumber('hsv-v', 0.4),
|
||||
degrees: readNumber('degrees', 0),
|
||||
translate: readNumber('translate', 0.1),
|
||||
scale: readNumber('scale', 0.5),
|
||||
shear: readNumber('shear', 0),
|
||||
perspective: readNumber('perspective', 0),
|
||||
close_mosaic: readNumber('close-mosaic', 10, true),
|
||||
flipud: readNumber('flipud', 0),
|
||||
fliplr: readNumber('fliplr', 0.5),
|
||||
bgr: readNumber('bgr', 0),
|
||||
mosaic: readNumber('mosaic', 1),
|
||||
mixup: readNumber('mixup', 0),
|
||||
cutmix: readNumber('cutmix', 0),
|
||||
copy_paste: readNumber('copy-paste', 0),
|
||||
erasing: readNumber('erasing', 0.4),
|
||||
copy_paste_mode: document.getElementById('copy-paste-mode').value,
|
||||
auto_augment: document.getElementById('auto-augment').value
|
||||
},
|
||||
|
|
@ -541,17 +555,17 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
}
|
||||
|
||||
// Split
|
||||
splitEnabled.checked = data.split?.enabled || false;
|
||||
splitRatio.value = data.split?.train_ratio || 0.8;
|
||||
splitEnabled.checked = data.split?.enabled ?? false;
|
||||
splitRatio.value = data.split?.train_ratio ?? 0.8;
|
||||
splitClasses.value = data.split?.classes_path || '';
|
||||
|
||||
// Training params
|
||||
document.getElementById('epochs').value = data.epochs || 100;
|
||||
document.getElementById('image-size').value = data.image_size || 640;
|
||||
document.getElementById('batch-size').value = data.batch_size || 16;
|
||||
document.getElementById('epochs').value = data.epochs ?? 100;
|
||||
document.getElementById('image-size').value = data.image_size ?? 640;
|
||||
document.getElementById('batch-size').value = data.batch_size ?? 16;
|
||||
document.getElementById('device').value = data.device || '';
|
||||
document.getElementById('workers').value = data.workers || 8;
|
||||
document.getElementById('patience').value = data.patience || 100;
|
||||
document.getElementById('workers').value = data.workers ?? 8;
|
||||
document.getElementById('patience').value = data.patience ?? 100;
|
||||
document.getElementById('project').value = data.project || 'runs/train';
|
||||
document.getElementById('run-name').value = data.run_name || '';
|
||||
|
||||
|
|
@ -796,6 +810,9 @@ document.addEventListener('DOMContentLoaded', () => {
|
|||
localStorage.removeItem('draft_config');
|
||||
await loadSessionsList();
|
||||
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) {
|
||||
configForm.addEventListener('input', () => {
|
||||
const config = getFormConfig();
|
||||
|
|
|
|||
|
|
@ -517,7 +517,9 @@ body {
|
|||
#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-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); }
|
||||
|
||||
@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-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-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); }
|
||||
|
||||
@keyframes pulseDot {
|
||||
|
|
@ -569,7 +573,9 @@ body {
|
|||
#status-card.status-preparing #status-title { color: var(--warning); }
|
||||
#status-card.status-training #status-title { color: var(--success); }
|
||||
#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-timer {
|
||||
|
|
|
|||
|
|
@ -11,6 +11,9 @@ from typing import Any
|
|||
|
||||
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)
|
||||
class TrainingEvent:
|
||||
|
|
@ -58,6 +61,7 @@ class TrainingRunner:
|
|||
self._model: Any | None = None
|
||||
self._state_lock = RLock()
|
||||
self._stop_requested = Event()
|
||||
self._force_stop_triggered = Event()
|
||||
self._subprocess: Any | None = None
|
||||
self._subprocess_ready = False
|
||||
self._force_stop_timer: Timer | None = None
|
||||
|
|
@ -69,6 +73,7 @@ class TrainingRunner:
|
|||
self._force_stop_timer = None
|
||||
self._subprocess_ready = False
|
||||
self._stop_requested.clear()
|
||||
self._force_stop_triggered.clear()
|
||||
if timer is not None:
|
||||
timer.cancel()
|
||||
|
||||
|
|
@ -119,6 +124,10 @@ class TrainingRunner:
|
|||
def stop_requested(self) -> bool:
|
||||
return self._stop_requested.is_set()
|
||||
|
||||
@property
|
||||
def force_stop_triggered(self) -> bool:
|
||||
return self._force_stop_triggered.is_set()
|
||||
|
||||
@staticmethod
|
||||
def _send_cooperative_stop(process: Any) -> None:
|
||||
try:
|
||||
|
|
@ -149,6 +158,7 @@ class TrainingRunner:
|
|||
try:
|
||||
if process.poll() is None:
|
||||
process.kill()
|
||||
self._force_stop_triggered.set()
|
||||
except (AttributeError, OSError, ProcessLookupError):
|
||||
pass
|
||||
|
||||
|
|
@ -227,7 +237,7 @@ class TrainingRunner:
|
|||
def _on_train_end(self, on_event: EventHandler) -> Callable[[Any], None]:
|
||||
def callback(trainer: Any) -> None:
|
||||
if self._stop_requested.is_set():
|
||||
on_event(TrainingEvent("warning", "Обучение остановлено пользователем."))
|
||||
on_event(TrainingEvent("cancelled", "Обучение остановлено пользователем."))
|
||||
else:
|
||||
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
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from yolo_webui.app import app
|
||||
from yolo_webui.app import TrainingManager, app
|
||||
|
||||
|
||||
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
|
||||
response = client.get("/api/sessions/nonexistent")
|
||||
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"):
|
||||
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["flip_idx"] == [0, 2, 1, 4, 3]
|
||||
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