Improve train utility functionality
This commit is contained in:
parent
53f758cc07
commit
c87450393e
22 changed files with 1907 additions and 262 deletions
17
.dockerignore
Normal file
17
.dockerignore
Normal file
|
|
@ -0,0 +1,17 @@
|
||||||
|
.git
|
||||||
|
.agents
|
||||||
|
.venv
|
||||||
|
.pytest_cache
|
||||||
|
**/__pycache__
|
||||||
|
*.py[cod]
|
||||||
|
.coverage
|
||||||
|
build
|
||||||
|
dist
|
||||||
|
datasets
|
||||||
|
models
|
||||||
|
runs
|
||||||
|
mlflow
|
||||||
|
mlruns
|
||||||
|
mlflow.db*
|
||||||
|
*.pt
|
||||||
|
.DS_Store
|
||||||
17
Dockerfile
17
Dockerfile
|
|
@ -9,14 +9,19 @@ RUN pip install --no-cache-dir uv==0.10.6
|
||||||
# Set working directory
|
# Set working directory
|
||||||
WORKDIR /workspace
|
WORKDIR /workspace
|
||||||
|
|
||||||
# Install our dependencies into the system python without resolving/syncing
|
ENV ULTRALYTICS_SAFE_LOAD=1
|
||||||
# which would remove packages installed by ultralytics base image (like tensorrt).
|
|
||||||
COPY pyproject.toml README.md ./
|
|
||||||
RUN uv pip install --system fastapi mlflow uvicorn websockets
|
|
||||||
|
|
||||||
# Copy source code and install the project
|
# Keep the training/GPU stack supplied by the Ultralytics image, while installing
|
||||||
|
# the remaining application dependencies at the exact versions recorded in uv.lock.
|
||||||
|
COPY pyproject.toml uv.lock README.md ./
|
||||||
|
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||||
|
uv export --locked --no-dev --group build --no-emit-project --prune ultralytics \
|
||||||
|
--output-file /tmp/requirements.lock && \
|
||||||
|
uv pip install --system --requirement /tmp/requirements.lock
|
||||||
|
|
||||||
|
# Copy source code and install only the project itself without re-resolving deps.
|
||||||
COPY src ./src
|
COPY src ./src
|
||||||
RUN uv pip install --system -e .
|
RUN uv pip install --system --no-deps --no-build-isolation .
|
||||||
|
|
||||||
# Expose Web UI port and MLflow port
|
# Expose Web UI port and MLflow port
|
||||||
EXPOSE 8000
|
EXPOSE 8000
|
||||||
|
|
|
||||||
16
README.md
16
README.md
|
|
@ -42,14 +42,17 @@ Ultralytics скачает веса. Пользовательские модел
|
||||||
docker compose up --build
|
docker compose up --build
|
||||||
```
|
```
|
||||||
|
|
||||||
WebUI будет доступен по `http://127.0.0.1:8000`. Compose намеренно публикует порт
|
WebUI будет доступен по `http://127.0.0.1:8000`, а MLflow — по
|
||||||
только на loopback. Не заменяйте адрес на `0.0.0.0` без аутентифицирующего reverse
|
`http://127.0.0.1:5000`. Compose намеренно публикует оба порта только на loopback и
|
||||||
proxy: API позволяет запускать и останавливать ресурсоёмкие задачи.
|
хранит состояние MLflow в `./mlflow`. Не заменяйте адреса на `0.0.0.0` без
|
||||||
|
аутентифицирующего reverse proxy: API позволяет запускать и останавливать
|
||||||
|
ресурсоёмкие задачи.
|
||||||
|
|
||||||
Для NVIDIA GPU раскомментируйте секцию `deploy.resources.reservations.devices` в
|
Для NVIDIA GPU раскомментируйте секцию `deploy.resources.reservations.devices` в
|
||||||
`docker-compose.yml`. Образ устанавливает зафиксированные в `uv.lock` зависимости;
|
`docker-compose.yml`. Web/MLflow-зависимости устанавливаются в версиях из
|
||||||
для другого варианта PyTorch используйте отдельно сгенерированный и проверенный
|
`uv.lock`, а PyTorch, Ultralytics и GPU runtime предоставляет базовый образ
|
||||||
lock-файл.
|
Ultralytics. Для полностью воспроизводимой production-сборки дополнительно
|
||||||
|
зафиксируйте базовый образ по digest.
|
||||||
|
|
||||||
## Разрешённые пути
|
## Разрешённые пути
|
||||||
|
|
||||||
|
|
@ -124,6 +127,7 @@ trainer и сохраняется при запуске Ultralytics DDP на н
|
||||||
Проверка интеграции на минимальных датасетах для всех пяти задач:
|
Проверка интеграции на минимальных датасетах для всех пяти задач:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
uv run scripts/create_yolo26_smoke_datasets.py
|
||||||
uv run scripts/run_yolo26_smoke_training.py --mlflow
|
uv run scripts/run_yolo26_smoke_training.py --mlflow
|
||||||
uv run scripts/verify_mlflow_smoke.py
|
uv run scripts/verify_mlflow_smoke.py
|
||||||
```
|
```
|
||||||
|
|
|
||||||
|
|
@ -5,19 +5,49 @@ services:
|
||||||
image: yolo-train-webui:latest
|
image: yolo-train-webui:latest
|
||||||
ports:
|
ports:
|
||||||
# The training API has no built-in user accounts, so expose it locally only.
|
# The training API has no built-in user accounts, so expose it locally only.
|
||||||
- "0.0.0.0:8000:8000"
|
- "127.0.0.1:8000:8000"
|
||||||
volumes:
|
volumes:
|
||||||
- ./datasets:/workspace/datasets
|
- ./datasets:/workspace/datasets
|
||||||
- ./runs:/workspace/runs
|
- ./runs:/workspace/runs
|
||||||
- ./models:/workspace/models
|
- ./models:/workspace/models
|
||||||
- ./models/.config:/root/.config/Ultralytics
|
- ./models/.config:/root/.config/Ultralytics
|
||||||
|
environment:
|
||||||
|
# Use the Compose MLflow service so metadata and artifacts survive container recreation.
|
||||||
|
YOLO_WEBUI_MLFLOW_TRACKING_URI: http://mlflow:5000
|
||||||
|
depends_on:
|
||||||
|
- mlflow
|
||||||
# Uncomment the block below on Linux with NVIDIA GPU to pass the graphics card into the container:
|
# Uncomment the block below on Linux with NVIDIA GPU to pass the graphics card into the container:
|
||||||
deploy:
|
# deploy:
|
||||||
resources:
|
# resources:
|
||||||
reservations:
|
# reservations:
|
||||||
devices:
|
# devices:
|
||||||
- driver: nvidia
|
# - driver: nvidia
|
||||||
count: all
|
# count: all
|
||||||
capabilities: [gpu]
|
# capabilities: [gpu]
|
||||||
ipc: host
|
ipc: host
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
|
|
||||||
|
mlflow:
|
||||||
|
build:
|
||||||
|
context: .
|
||||||
|
image: yolo-train-webui:latest
|
||||||
|
command:
|
||||||
|
- mlflow
|
||||||
|
- server
|
||||||
|
- --backend-store-uri
|
||||||
|
- sqlite:////mlflow/mlflow.db
|
||||||
|
- --artifacts-destination
|
||||||
|
- /mlflow/artifacts
|
||||||
|
- --host
|
||||||
|
- 0.0.0.0
|
||||||
|
- --port
|
||||||
|
- "5000"
|
||||||
|
- --workers
|
||||||
|
- "1"
|
||||||
|
- --allowed-hosts
|
||||||
|
- localhost:5000,127.0.0.1:5000,mlflow:5000
|
||||||
|
ports:
|
||||||
|
- "127.0.0.1:5000:5000"
|
||||||
|
volumes:
|
||||||
|
- ./mlflow:/mlflow
|
||||||
|
restart: unless-stopped
|
||||||
|
|
|
||||||
|
|
@ -16,13 +16,16 @@ dependencies = [
|
||||||
yolo-train-webui = "yolo_webui.app:main"
|
yolo-train-webui = "yolo_webui.app:main"
|
||||||
|
|
||||||
[dependency-groups]
|
[dependency-groups]
|
||||||
|
build = [
|
||||||
|
"hatchling==1.27.0",
|
||||||
|
]
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=8.3",
|
"pytest>=8.3",
|
||||||
"httpx",
|
"httpx",
|
||||||
]
|
]
|
||||||
|
|
||||||
[build-system]
|
[build-system]
|
||||||
requires = ["hatchling"]
|
requires = ["hatchling==1.27.0"]
|
||||||
build-backend = "hatchling.build"
|
build-backend = "hatchling.build"
|
||||||
|
|
||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
|
|
@ -31,4 +34,3 @@ packages = ["src/yolo_webui"]
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
addopts = "-q"
|
addopts = "-q"
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,7 @@ from fastapi.staticfiles import StaticFiles
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
|
||||||
from yolo_webui.config import TrainingConfig
|
from yolo_webui.config import TrainingConfig
|
||||||
|
from yolo_webui.export_runner import ExportConfig
|
||||||
from yolo_webui.trainer import TrainingRunner
|
from yolo_webui.trainer import TrainingRunner
|
||||||
|
|
||||||
# Set up logging
|
# Set up logging
|
||||||
|
|
@ -49,6 +50,17 @@ class LiveState:
|
||||||
self.stop_requested = False
|
self.stop_requested = False
|
||||||
self.last_event_kind = None
|
self.last_event_kind = None
|
||||||
|
|
||||||
|
def snapshot(self) -> dict[str, Any]:
|
||||||
|
"""Return a detached snapshot safe to serialize after releasing a lock."""
|
||||||
|
return {
|
||||||
|
"status": self.status,
|
||||||
|
"epoch": self.epoch,
|
||||||
|
"total_epochs": self.total_epochs,
|
||||||
|
"logs": list(self.logs),
|
||||||
|
"metrics": [dict(metric) for metric in self.metrics],
|
||||||
|
"output_dir": self.output_dir,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class TrainingManager:
|
class TrainingManager:
|
||||||
"""Manages the background training subprocess and WebSocket clients."""
|
"""Manages the background training subprocess and WebSocket clients."""
|
||||||
|
|
@ -355,6 +367,51 @@ def get_sessions_dir() -> Path:
|
||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def write_json_atomic(path: Path, data: dict[str, Any]) -> None:
|
||||||
|
"""Replace a JSON file atomically so concurrent readers never see partial data."""
|
||||||
|
temp_path: Path | None = None
|
||||||
|
try:
|
||||||
|
with tempfile.NamedTemporaryFile(
|
||||||
|
"w",
|
||||||
|
prefix=f".{path.name}.",
|
||||||
|
suffix=".tmp",
|
||||||
|
dir=path.parent,
|
||||||
|
delete=False,
|
||||||
|
encoding="utf-8",
|
||||||
|
) as temp_file:
|
||||||
|
temp_path = Path(temp_file.name)
|
||||||
|
json.dump(data, temp_file, ensure_ascii=False, indent=2)
|
||||||
|
os.replace(temp_path, path)
|
||||||
|
temp_path = None
|
||||||
|
finally:
|
||||||
|
if temp_path is not None:
|
||||||
|
try:
|
||||||
|
temp_path.unlink()
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def parse_training_config(config_data: dict[str, Any]) -> TrainingConfig:
|
||||||
|
config = TrainingConfig.from_dict(config_data)
|
||||||
|
config.validate()
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def parse_export_config(config_data: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
config = ExportConfig.from_mapping(config_data)
|
||||||
|
return {
|
||||||
|
"model": str(config.model),
|
||||||
|
"format": config.export_format,
|
||||||
|
"imgsz": config.imgsz,
|
||||||
|
"half": config.half,
|
||||||
|
"int8": config.int8,
|
||||||
|
"dynamic": config.dynamic,
|
||||||
|
"simplify": config.simplify,
|
||||||
|
"batch": config.batch,
|
||||||
|
"workspace": config.workspace,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def get_session_path(name: str, *, allow_last_run: bool = True) -> Path:
|
def get_session_path(name: str, *, allow_last_run: bool = True) -> Path:
|
||||||
if SESSION_NAME_PATTERN.fullmatch(name) is None:
|
if SESSION_NAME_PATTERN.fullmatch(name) is None:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|
@ -432,13 +489,13 @@ async def list_models():
|
||||||
try:
|
try:
|
||||||
for path in runs_dir.rglob("*.pt"):
|
for path in runs_dir.rglob("*.pt"):
|
||||||
items.append({
|
items.append({
|
||||||
"name": f"{path.parent.parent.parent.name}/{path.parent.parent.name}/{path.name}",
|
"name": path.relative_to(runs_dir).as_posix(),
|
||||||
"path": str(path.absolute()),
|
"path": str(path.absolute()),
|
||||||
"source": "runs"
|
"source": "runs"
|
||||||
})
|
})
|
||||||
for path in runs_dir.rglob("*.pth"):
|
for path in runs_dir.rglob("*.pth"):
|
||||||
items.append({
|
items.append({
|
||||||
"name": f"{path.parent.parent.parent.name}/{path.parent.parent.name}/{path.name}",
|
"name": path.relative_to(runs_dir).as_posix(),
|
||||||
"path": str(path.absolute()),
|
"path": str(path.absolute()),
|
||||||
"source": "runs"
|
"source": "runs"
|
||||||
})
|
})
|
||||||
|
|
@ -464,8 +521,13 @@ async def load_session(name: str):
|
||||||
async def save_session(name: str, config_data: dict[str, Any]):
|
async def save_session(name: str, config_data: dict[str, Any]):
|
||||||
file_path = get_session_path(name, allow_last_run=False)
|
file_path = get_session_path(name, allow_last_run=False)
|
||||||
try:
|
try:
|
||||||
with file_path.open("w", encoding="utf-8") as f:
|
config = parse_training_config(config_data)
|
||||||
json.dump(config_data, f, ensure_ascii=False, indent=2)
|
except ValueError as exc:
|
||||||
|
raise HTTPException(status_code=400, detail=str(exc))
|
||||||
|
except Exception as exc:
|
||||||
|
raise HTTPException(status_code=400, detail=f"Некорректная конфигурация: {exc}")
|
||||||
|
try:
|
||||||
|
write_json_atomic(file_path, config.to_dict())
|
||||||
return {"message": "Сессия успешно сохранена."}
|
return {"message": "Сессия успешно сохранена."}
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
raise HTTPException(status_code=500, detail=f"Не удалось сохранить сессию: {exc}")
|
raise HTTPException(status_code=500, detail=f"Не удалось сохранить сессию: {exc}")
|
||||||
|
|
@ -486,40 +548,32 @@ async def delete_session(name: str):
|
||||||
@app.get("/api/train/status")
|
@app.get("/api/train/status")
|
||||||
async def get_status():
|
async def get_status():
|
||||||
with manager._lock:
|
with manager._lock:
|
||||||
return {
|
return manager.state.snapshot()
|
||||||
"status": manager.state.status,
|
|
||||||
"epoch": manager.state.epoch,
|
|
||||||
"total_epochs": manager.state.total_epochs,
|
|
||||||
"output_dir": manager.state.output_dir,
|
|
||||||
"metrics": manager.state.metrics,
|
|
||||||
"logs": manager.state.logs,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/train/start")
|
@app.post("/api/train/start")
|
||||||
async def start_training(config_data: dict[str, Any]):
|
async def start_training(config_data: dict[str, Any]):
|
||||||
try:
|
try:
|
||||||
config = TrainingConfig.from_dict(config_data)
|
config = parse_training_config(config_data)
|
||||||
config.validate()
|
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
raise HTTPException(status_code=400, detail=str(exc))
|
raise HTTPException(status_code=400, detail=str(exc))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
raise HTTPException(status_code=400, detail=f"Некорректная конфигурация: {exc}")
|
raise HTTPException(status_code=400, detail=f"Некорректная конфигурация: {exc}")
|
||||||
|
|
||||||
# Auto-save last configuration on start
|
try:
|
||||||
|
manager.start_training(config)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise HTTPException(status_code=400, detail=str(exc))
|
||||||
|
|
||||||
|
# Only a configuration accepted for execution becomes the last run.
|
||||||
try:
|
try:
|
||||||
sessions_dir = get_sessions_dir()
|
sessions_dir = get_sessions_dir()
|
||||||
last_run_path = sessions_dir / "last_run.json"
|
last_run_path = sessions_dir / "last_run.json"
|
||||||
with last_run_path.open("w", encoding="utf-8") as f:
|
write_json_atomic(last_run_path, config.to_dict())
|
||||||
json.dump(config_data, f, ensure_ascii=False, indent=2)
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(f"Failed to auto-save last run: {exc}")
|
logger.error(f"Failed to auto-save last run: {exc}")
|
||||||
|
|
||||||
try:
|
return {"message": "Обучение запущено."}
|
||||||
manager.start_training(config)
|
|
||||||
return {"message": "Обучение запущено."}
|
|
||||||
except ValueError as exc:
|
|
||||||
raise HTTPException(status_code=400, detail=str(exc))
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/train/stop")
|
@app.post("/api/train/stop")
|
||||||
|
|
@ -572,8 +626,10 @@ class ExportManager:
|
||||||
future = loop.create_task(coroutine)
|
future = loop.create_task(coroutine)
|
||||||
else:
|
else:
|
||||||
future = asyncio.run_coroutine_threadsafe(coroutine, loop)
|
future = asyncio.run_coroutine_threadsafe(coroutine, loop)
|
||||||
|
future.add_done_callback(TrainingManager._log_broadcast_failure)
|
||||||
except Exception:
|
except Exception:
|
||||||
coroutine.close()
|
coroutine.close()
|
||||||
|
logger.exception("Failed to schedule export WebSocket broadcast")
|
||||||
|
|
||||||
async def _send_payload(self, payload: str) -> None:
|
async def _send_payload(self, payload: str) -> None:
|
||||||
broadcast_lock = self._broadcast_lock
|
broadcast_lock = self._broadcast_lock
|
||||||
|
|
@ -621,19 +677,28 @@ class ExportManager:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
if self.state.status not in ("preparing", "exporting"):
|
if self.state.status not in ("preparing", "exporting"):
|
||||||
return
|
return
|
||||||
|
|
||||||
|
process = self._process
|
||||||
|
if process is not None:
|
||||||
|
if process.poll() is not None:
|
||||||
|
# The subprocess has already completed; let its real return code
|
||||||
|
# determine the final status instead of relabeling it cancelled.
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
process.terminate()
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Failed to terminate export subprocess", exc_info=True)
|
||||||
|
return
|
||||||
|
|
||||||
self.state.status = "stopping"
|
self.state.status = "stopping"
|
||||||
self.state.stop_requested = True
|
self.state.stop_requested = True
|
||||||
if self._process is not None:
|
|
||||||
try:
|
|
||||||
self._process.terminate()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
self.broadcast({"type": "status", "status": self.state.status})
|
self.broadcast({"type": "status", "status": self.state.status})
|
||||||
self.add_log("Запрошена остановка экспорта...", "warning")
|
self.add_log("Запрошена остановка экспорта...", "warning")
|
||||||
|
|
||||||
def _run_subprocess(self, config_data: dict[str, Any]) -> None:
|
def _run_subprocess(self, config_data: dict[str, Any]) -> None:
|
||||||
temp_config_path = None
|
temp_config_path = None
|
||||||
|
process: subprocess.Popen[str] | None = None
|
||||||
try:
|
try:
|
||||||
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False, encoding="utf-8") as f:
|
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False, encoding="utf-8") as f:
|
||||||
json.dump(config_data, f)
|
json.dump(config_data, f)
|
||||||
|
|
@ -641,17 +706,29 @@ class ExportManager:
|
||||||
|
|
||||||
cmd = [sys.executable, "-u", "-m", "yolo_webui.export_runner", temp_config_path]
|
cmd = [sys.executable, "-u", "-m", "yolo_webui.export_runner", temp_config_path]
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._process = subprocess.Popen(
|
if self.state.stop_requested:
|
||||||
cmd,
|
cancelled_before_start = True
|
||||||
stdout=subprocess.PIPE,
|
else:
|
||||||
stderr=subprocess.STDOUT,
|
cancelled_before_start = False
|
||||||
text=True,
|
process = subprocess.Popen(
|
||||||
bufsize=1,
|
cmd,
|
||||||
)
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.STDOUT,
|
||||||
|
text=True,
|
||||||
|
bufsize=1,
|
||||||
|
)
|
||||||
|
self._process = process
|
||||||
|
|
||||||
|
if cancelled_before_start:
|
||||||
|
self._finalize_process_result(0)
|
||||||
|
return
|
||||||
|
|
||||||
|
if process is None or process.stdout is None:
|
||||||
|
raise RuntimeError("Не удалось открыть stdout процесса экспорта.")
|
||||||
|
|
||||||
buffer = ""
|
buffer = ""
|
||||||
while True:
|
while True:
|
||||||
char = self._process.stdout.read(1)
|
char = process.stdout.read(1)
|
||||||
if not char:
|
if not char:
|
||||||
if buffer:
|
if buffer:
|
||||||
self._handle_subprocess_line(buffer)
|
self._handle_subprocess_line(buffer)
|
||||||
|
|
@ -664,26 +741,28 @@ class ExportManager:
|
||||||
else:
|
else:
|
||||||
buffer += char
|
buffer += char
|
||||||
|
|
||||||
self._process.wait()
|
process.wait()
|
||||||
rc = self._process.returncode
|
rc = process.returncode
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._process = None
|
if self._process is process:
|
||||||
|
self._process = None
|
||||||
|
|
||||||
self._finalize_process_result(rc)
|
self._finalize_process_result(rc)
|
||||||
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.exception("Error in export process thread:")
|
logger.exception("Error in export process thread:")
|
||||||
if self._process is not None:
|
if process is not None:
|
||||||
try:
|
try:
|
||||||
if self._process.poll() is None:
|
if process.poll() is None:
|
||||||
self._process.kill()
|
process.kill()
|
||||||
self._process.wait(timeout=5)
|
process.wait(timeout=5)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._process = None
|
if self._process is process:
|
||||||
|
self._process = None
|
||||||
self.state.status = "failed"
|
self.state.status = "failed"
|
||||||
self.add_log(f"Внутренняя ошибка менеджера: {exc}", "error")
|
self.add_log(f"Внутренняя ошибка менеджера: {exc}", "error")
|
||||||
finally:
|
finally:
|
||||||
|
|
@ -703,9 +782,13 @@ class ExportManager:
|
||||||
return
|
return
|
||||||
|
|
||||||
if line_str == "__YOLO_WEBUI_READY__":
|
if line_str == "__YOLO_WEBUI_READY__":
|
||||||
|
status_changed = False
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.state.status = "exporting"
|
if self.state.status == "preparing":
|
||||||
self.broadcast({"type": "status", "status": "exporting"})
|
self.state.status = "exporting"
|
||||||
|
status_changed = True
|
||||||
|
if status_changed:
|
||||||
|
self.broadcast({"type": "status", "status": "exporting"})
|
||||||
elif line_str.startswith("__YOLO_WEBUI_RESULT__:"):
|
elif line_str.startswith("__YOLO_WEBUI_RESULT__:"):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.state.output_dir = line_str[len("__YOLO_WEBUI_RESULT__:") :]
|
self.state.output_dir = line_str[len("__YOLO_WEBUI_RESULT__:") :]
|
||||||
|
|
@ -738,16 +821,18 @@ export_manager = ExportManager()
|
||||||
@app.get("/api/export/status")
|
@app.get("/api/export/status")
|
||||||
async def get_export_status():
|
async def get_export_status():
|
||||||
with export_manager._lock:
|
with export_manager._lock:
|
||||||
|
snapshot = export_manager.state.snapshot()
|
||||||
return {
|
return {
|
||||||
"status": export_manager.state.status,
|
"status": snapshot["status"],
|
||||||
"output_dir": export_manager.state.output_dir,
|
"output_dir": snapshot["output_dir"],
|
||||||
"logs": export_manager.state.logs,
|
"logs": snapshot["logs"],
|
||||||
}
|
}
|
||||||
|
|
||||||
@app.post("/api/export/start")
|
@app.post("/api/export/start")
|
||||||
async def start_export(config_data: dict[str, Any]):
|
async def start_export(config_data: dict[str, Any]):
|
||||||
try:
|
try:
|
||||||
export_manager.start_export(config_data)
|
validated_data = parse_export_config(config_data)
|
||||||
|
export_manager.start_export(validated_data)
|
||||||
return {"message": "Экспорт запущен."}
|
return {"message": "Экспорт запущен."}
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
raise HTTPException(status_code=400, detail=str(exc))
|
raise HTTPException(status_code=400, detail=str(exc))
|
||||||
|
|
@ -763,11 +848,12 @@ async def export_websocket_endpoint(websocket: WebSocket):
|
||||||
export_manager.add_websocket(websocket)
|
export_manager.add_websocket(websocket)
|
||||||
|
|
||||||
with export_manager._lock:
|
with export_manager._lock:
|
||||||
|
snapshot = export_manager.state.snapshot()
|
||||||
state_dict = {
|
state_dict = {
|
||||||
"type": "init",
|
"type": "init",
|
||||||
"status": export_manager.state.status,
|
"status": snapshot["status"],
|
||||||
"output_dir": export_manager.state.output_dir,
|
"output_dir": snapshot["output_dir"],
|
||||||
"logs": [log.split(":", 1) for log in export_manager.state.logs if ":" in log],
|
"logs": [log.split(":", 1) for log in snapshot["logs"] if ":" in log],
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
await websocket.send_text(json.dumps(state_dict))
|
await websocket.send_text(json.dumps(state_dict))
|
||||||
|
|
@ -788,15 +874,16 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||||
|
|
||||||
# Send current state upon connection
|
# Send current state upon connection
|
||||||
with manager._lock:
|
with manager._lock:
|
||||||
|
snapshot = manager.state.snapshot()
|
||||||
state_dict = {
|
state_dict = {
|
||||||
"type": "init",
|
"type": "init",
|
||||||
"status": manager.state.status,
|
"status": snapshot["status"],
|
||||||
"epoch": manager.state.epoch,
|
"epoch": snapshot["epoch"],
|
||||||
"total_epochs": manager.state.total_epochs,
|
"total_epochs": snapshot["total_epochs"],
|
||||||
"output_dir": manager.state.output_dir,
|
"output_dir": snapshot["output_dir"],
|
||||||
"metrics": manager.state.metrics,
|
"metrics": snapshot["metrics"],
|
||||||
# We format log items for the UI
|
# We format log items for the UI
|
||||||
"logs": [log.split(":", 1) for log in manager.state.logs if ":" in log],
|
"logs": [log.split(":", 1) for log in snapshot["logs"] if ":" in log],
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
await websocket.send_text(json.dumps(state_dict))
|
await websocket.send_text(json.dumps(state_dict))
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,9 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
import math
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path, PureWindowsPath
|
||||||
import re
|
import re
|
||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
|
|
@ -28,6 +29,58 @@ SAFE_IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]*$")
|
||||||
MODEL_SUFFIXES = {".pt", ".pth", ".yaml", ".yml"}
|
MODEL_SUFFIXES = {".pt", ".pth", ".yaml", ".yml"}
|
||||||
|
|
||||||
|
|
||||||
|
def _default_mlflow_tracking_uri() -> str:
|
||||||
|
return os.environ.get(
|
||||||
|
"YOLO_WEBUI_MLFLOW_TRACKING_URI",
|
||||||
|
"sqlite:///mlflow.db",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _require_string(value: Any, label: str) -> str:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
raise ValueError(f"{label} должен быть строкой.")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _require_bool(value: Any, label: str) -> bool:
|
||||||
|
if not isinstance(value, bool):
|
||||||
|
raise ValueError(f"{label} должен быть логическим значением.")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _require_finite_number(value: Any, label: str) -> float:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
|
raise ValueError(f"{label} должен быть числом.")
|
||||||
|
try:
|
||||||
|
number = float(value)
|
||||||
|
except OverflowError as exc:
|
||||||
|
raise ValueError(f"{label} должен быть конечным числом.") from exc
|
||||||
|
if not math.isfinite(number):
|
||||||
|
raise ValueError(f"{label} должен быть конечным числом.")
|
||||||
|
return number
|
||||||
|
|
||||||
|
|
||||||
|
def _require_integer(value: Any, label: str) -> int:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int):
|
||||||
|
raise ValueError(f"{label} должен быть целым числом.")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_run_name(value: str) -> None:
|
||||||
|
normalized = value.strip()
|
||||||
|
if not normalized:
|
||||||
|
return
|
||||||
|
if (
|
||||||
|
normalized in {".", ".."}
|
||||||
|
or "/" in normalized
|
||||||
|
or "\\" in normalized
|
||||||
|
or PureWindowsPath(normalized).drive
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"Имя запуска должно быть одним именем каталога без '/' или '\\'."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _allowed_roots(defaults: tuple[str, ...], environment_name: str) -> tuple[Path, ...]:
|
def _allowed_roots(defaults: tuple[str, ...], environment_name: str) -> tuple[Path, ...]:
|
||||||
configured = [
|
configured = [
|
||||||
item
|
item
|
||||||
|
|
@ -109,7 +162,7 @@ class AugmentationConfig:
|
||||||
close_mosaic: int = 10
|
close_mosaic: int = 10
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
if not self.enabled:
|
if not _require_bool(self.enabled, "Флаг аугментации"):
|
||||||
return
|
return
|
||||||
|
|
||||||
fractions = {
|
fractions = {
|
||||||
|
|
@ -129,20 +182,23 @@ class AugmentationConfig:
|
||||||
"Erasing": self.erasing,
|
"Erasing": self.erasing,
|
||||||
}
|
}
|
||||||
for label, value in fractions.items():
|
for label, value in fractions.items():
|
||||||
if not 0.0 <= value <= 1.0:
|
number = _require_finite_number(value, f"Параметр «{label}»")
|
||||||
|
if not 0.0 <= number <= 1.0:
|
||||||
raise ValueError(f"Параметр «{label}» должен быть от 0 до 1.")
|
raise ValueError(f"Параметр «{label}» должен быть от 0 до 1.")
|
||||||
if self.degrees < 0:
|
if _require_finite_number(self.degrees, "Угол поворота") < 0:
|
||||||
raise ValueError("Угол поворота не может быть отрицательным.")
|
raise ValueError("Угол поворота не может быть отрицательным.")
|
||||||
if self.shear < 0:
|
if _require_finite_number(self.shear, "Угол сдвига") < 0:
|
||||||
raise ValueError("Угол сдвига не может быть отрицательным.")
|
raise ValueError("Угол сдвига не может быть отрицательным.")
|
||||||
if self.close_mosaic < 0:
|
if _require_integer(self.close_mosaic, "Close mosaic") < 0:
|
||||||
raise ValueError("Close mosaic не может быть отрицательным.")
|
raise ValueError("Close mosaic не может быть отрицательным.")
|
||||||
|
_require_string(self.copy_paste_mode, "Режим copy-paste")
|
||||||
if self.copy_paste_mode not in SUPPORTED_COPY_PASTE_MODES:
|
if self.copy_paste_mode not in SUPPORTED_COPY_PASTE_MODES:
|
||||||
raise ValueError(f"Неизвестный режим copy-paste: {self.copy_paste_mode}.")
|
raise ValueError(f"Неизвестный режим copy-paste: {self.copy_paste_mode}.")
|
||||||
|
_require_string(self.auto_augment, "Политика AutoAugment")
|
||||||
if self.auto_augment not in SUPPORTED_AUTO_AUGMENT_POLICIES:
|
if self.auto_augment not in SUPPORTED_AUTO_AUGMENT_POLICIES:
|
||||||
raise ValueError(f"Неизвестная политика AutoAugment: {self.auto_augment}.")
|
raise ValueError(f"Неизвестная политика AutoAugment: {self.auto_augment}.")
|
||||||
|
|
||||||
def train_kwargs(self) -> dict[str, str | int | float]:
|
def train_kwargs(self) -> dict[str, str | int | float | None]:
|
||||||
if not self.enabled:
|
if not self.enabled:
|
||||||
return {}
|
return {}
|
||||||
return {
|
return {
|
||||||
|
|
@ -171,11 +227,15 @@ class AugmentationConfig:
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class MlflowConfig:
|
class MlflowConfig:
|
||||||
enabled: bool = True
|
enabled: bool = True
|
||||||
tracking_uri: str = "sqlite:///mlflow.db"
|
tracking_uri: str = field(default_factory=_default_mlflow_tracking_uri)
|
||||||
experiment_name: str = "yolo-webui"
|
experiment_name: str = "yolo-webui"
|
||||||
run_name: str = ""
|
run_name: str = ""
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
|
_require_bool(self.enabled, "Флаг MLflow")
|
||||||
|
_require_string(self.tracking_uri, "URI хранилища MLflow")
|
||||||
|
_require_string(self.experiment_name, "Название эксперимента MLflow")
|
||||||
|
_require_string(self.run_name, "Название запуска MLflow")
|
||||||
if self.enabled and not self.tracking_uri.strip():
|
if self.enabled and not self.tracking_uri.strip():
|
||||||
raise ValueError("Укажите URI хранилища MLflow.")
|
raise ValueError("Укажите URI хранилища MLflow.")
|
||||||
if self.enabled and not self.experiment_name.strip():
|
if self.enabled and not self.experiment_name.strip():
|
||||||
|
|
@ -189,8 +249,14 @@ class DatasetSplitConfig:
|
||||||
classes_path: str = ""
|
classes_path: str = ""
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
|
_require_bool(self.enabled, "Флаг разделения датасета")
|
||||||
|
ratio = _require_finite_number(
|
||||||
|
self.train_ratio,
|
||||||
|
"Доля обучающей выборки (Train)",
|
||||||
|
)
|
||||||
|
_require_string(self.classes_path, "Путь к файлу классов")
|
||||||
if self.enabled:
|
if self.enabled:
|
||||||
if not 0.1 <= self.train_ratio <= 0.95:
|
if not 0.1 <= ratio <= 0.95:
|
||||||
raise ValueError("Доля обучающей выборки (Train) должна быть от 0.1 до 0.95.")
|
raise ValueError("Доля обучающей выборки (Train) должна быть от 0.1 до 0.95.")
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -212,10 +278,27 @@ class TrainingConfig:
|
||||||
split: DatasetSplitConfig = field(default_factory=DatasetSplitConfig)
|
split: DatasetSplitConfig = field(default_factory=DatasetSplitConfig)
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
if not self.dataset.strip():
|
dataset = _require_string(self.dataset, "Датасет")
|
||||||
|
model = _require_string(self.model, "Модель")
|
||||||
|
_require_string(self.task, "Тип задачи")
|
||||||
|
_require_string(self.device, "Устройство")
|
||||||
|
project = _require_string(self.project, "Каталог результатов")
|
||||||
|
run_name = _require_string(self.run_name, "Имя запуска")
|
||||||
|
if not isinstance(self.augmentation, AugmentationConfig):
|
||||||
|
raise ValueError("Параметры аугментации должны быть объектом.")
|
||||||
|
if not isinstance(self.mlflow, MlflowConfig):
|
||||||
|
raise ValueError("Параметры MLflow должны быть объектом.")
|
||||||
|
if not isinstance(self.split, DatasetSplitConfig):
|
||||||
|
raise ValueError("Параметры разделения датасета должны быть объектом.")
|
||||||
|
self.augmentation.validate()
|
||||||
|
self.mlflow.validate()
|
||||||
|
self.split.validate()
|
||||||
|
|
||||||
|
if not dataset.strip():
|
||||||
raise ValueError("Укажите путь или имя датасета.")
|
raise ValueError("Укажите путь или имя датасета.")
|
||||||
if not self.model.strip():
|
if not model.strip():
|
||||||
raise ValueError("Укажите путь или имя модели.")
|
raise ValueError("Укажите путь или имя модели.")
|
||||||
|
_validate_run_name(run_name)
|
||||||
|
|
||||||
data_roots = _allowed_roots(("datasets",), "YOLO_WEBUI_DATA_ROOTS")
|
data_roots = _allowed_roots(("datasets",), "YOLO_WEBUI_DATA_ROOTS")
|
||||||
model_roots = _allowed_roots(
|
model_roots = _allowed_roots(
|
||||||
|
|
@ -224,23 +307,29 @@ class TrainingConfig:
|
||||||
)
|
)
|
||||||
run_roots = _allowed_roots(("runs",), "YOLO_WEBUI_RUN_ROOTS")
|
run_roots = _allowed_roots(("runs",), "YOLO_WEBUI_RUN_ROOTS")
|
||||||
_validate_local_reference(
|
_validate_local_reference(
|
||||||
self.dataset,
|
dataset,
|
||||||
label="Датасет",
|
label="Датасет",
|
||||||
roots=data_roots,
|
roots=data_roots,
|
||||||
allow_identifier=True,
|
allow_identifier=True,
|
||||||
)
|
)
|
||||||
_validate_local_reference(
|
_validate_local_reference(
|
||||||
self.model,
|
model,
|
||||||
label="Модель",
|
label="Модель",
|
||||||
roots=model_roots,
|
roots=model_roots,
|
||||||
allow_identifier=True,
|
allow_identifier=True,
|
||||||
allowed_suffixes=MODEL_SUFFIXES,
|
allowed_suffixes=MODEL_SUFFIXES,
|
||||||
)
|
)
|
||||||
_validate_local_reference(
|
_validate_local_reference(
|
||||||
self.project.strip() or "runs/train",
|
project.strip() or "runs/train",
|
||||||
label="Каталог результатов",
|
label="Каталог результатов",
|
||||||
roots=run_roots,
|
roots=run_roots,
|
||||||
)
|
)
|
||||||
|
if run_name.strip():
|
||||||
|
_validate_local_reference(
|
||||||
|
str(Path(project.strip() or "runs/train") / run_name.strip()),
|
||||||
|
label="Каталог запуска",
|
||||||
|
roots=run_roots,
|
||||||
|
)
|
||||||
if self.split.classes_path.strip():
|
if self.split.classes_path.strip():
|
||||||
_validate_local_reference(
|
_validate_local_reference(
|
||||||
self.split.classes_path,
|
self.split.classes_path,
|
||||||
|
|
@ -254,19 +343,24 @@ class TrainingConfig:
|
||||||
"Автоматическое разделение доступно только для YOLO-датасетов "
|
"Автоматическое разделение доступно только для YOLO-датасетов "
|
||||||
"с папками images/labels и не поддерживает задачу classify."
|
"с папками images/labels и не поддерживает задачу classify."
|
||||||
)
|
)
|
||||||
if self.epochs < 1:
|
if _require_integer(self.epochs, "Количество эпох") < 1:
|
||||||
raise ValueError("Количество эпох должно быть не меньше 1.")
|
raise ValueError("Количество эпох должно быть не меньше 1.")
|
||||||
if self.image_size < 32:
|
if _require_integer(self.image_size, "Размер изображения") < 32:
|
||||||
raise ValueError("Размер изображения должен быть не меньше 32.")
|
raise ValueError("Размер изображения должен быть не меньше 32.")
|
||||||
if self.batch_size == 0 or self.batch_size < -1:
|
batch_size = _require_finite_number(self.batch_size, "Batch")
|
||||||
raise ValueError("Batch должен быть положительным числом или -1 для автоподбора.")
|
if isinstance(self.batch_size, int):
|
||||||
if self.workers < 0:
|
valid_batch = self.batch_size == -1 or self.batch_size >= 1
|
||||||
|
else:
|
||||||
|
valid_batch = 0 < batch_size < 1
|
||||||
|
if not valid_batch:
|
||||||
|
raise ValueError(
|
||||||
|
"Batch должен быть целым числом от 1, значением -1 для "
|
||||||
|
"автоподбора либо дробью от 0 до 1."
|
||||||
|
)
|
||||||
|
if _require_integer(self.workers, "Количество workers") < 0:
|
||||||
raise ValueError("Количество workers не может быть отрицательным.")
|
raise ValueError("Количество workers не может быть отрицательным.")
|
||||||
if self.patience < 0:
|
if _require_integer(self.patience, "Patience") < 0:
|
||||||
raise ValueError("Patience не может быть отрицательным.")
|
raise ValueError("Patience не может быть отрицательным.")
|
||||||
self.augmentation.validate()
|
|
||||||
self.mlflow.validate()
|
|
||||||
self.split.validate()
|
|
||||||
|
|
||||||
def train_kwargs(self) -> dict[str, str | int | float | bool]:
|
def train_kwargs(self) -> dict[str, str | int | float | bool]:
|
||||||
"""Convert the form values to arguments accepted by YOLO.train()."""
|
"""Convert the form values to arguments accepted by YOLO.train()."""
|
||||||
|
|
@ -303,12 +397,30 @@ class TrainingConfig:
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(cls, data: dict[str, Any]) -> TrainingConfig:
|
def from_dict(cls, data: dict[str, Any]) -> TrainingConfig:
|
||||||
aug_data = data.get("augmentation", {})
|
if not isinstance(data, dict):
|
||||||
mlflow_data = data.get("mlflow", {})
|
raise ValueError("Конфигурация должна быть JSON-объектом.")
|
||||||
split_data = data.get("split", {})
|
|
||||||
|
sections: dict[str, dict[str, Any]] = {}
|
||||||
|
for key, label in (
|
||||||
|
("augmentation", "аугментации"),
|
||||||
|
("mlflow", "MLflow"),
|
||||||
|
("split", "разделения датасета"),
|
||||||
|
):
|
||||||
|
raw_section = data.get(key, {})
|
||||||
|
if not isinstance(raw_section, dict):
|
||||||
|
raise ValueError(f"Параметры {label} должны быть JSON-объектом.")
|
||||||
|
sections[key] = raw_section
|
||||||
|
|
||||||
|
try:
|
||||||
|
augmentation = AugmentationConfig(**sections["augmentation"])
|
||||||
|
mlflow = MlflowConfig(**sections["mlflow"])
|
||||||
|
split = DatasetSplitConfig(**sections["split"])
|
||||||
|
except TypeError as exc:
|
||||||
|
raise ValueError(f"Конфигурация содержит неизвестные параметры: {exc}") from exc
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
dataset=data["dataset"],
|
dataset=data.get("dataset", ""),
|
||||||
model=data["model"],
|
model=data.get("model", ""),
|
||||||
task=data.get("task", "detect"),
|
task=data.get("task", "detect"),
|
||||||
epochs=data.get("epochs", 100),
|
epochs=data.get("epochs", 100),
|
||||||
image_size=data.get("image_size", 640),
|
image_size=data.get("image_size", 640),
|
||||||
|
|
@ -318,7 +430,7 @@ class TrainingConfig:
|
||||||
patience=data.get("patience", 100),
|
patience=data.get("patience", 100),
|
||||||
project=data.get("project", "runs/train"),
|
project=data.get("project", "runs/train"),
|
||||||
run_name=data.get("run_name", ""),
|
run_name=data.get("run_name", ""),
|
||||||
augmentation=AugmentationConfig(**aug_data) if aug_data else AugmentationConfig(),
|
augmentation=augmentation,
|
||||||
mlflow=MlflowConfig(**mlflow_data) if mlflow_data else MlflowConfig(),
|
mlflow=mlflow,
|
||||||
split=DatasetSplitConfig(**split_data) if split_data else DatasetSplitConfig(),
|
split=split,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
import random
|
import random
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
@ -8,6 +9,25 @@ from uuid import uuid4
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_class_key(raw_key: Any, source: Path) -> int:
|
||||||
|
if isinstance(raw_key, bool):
|
||||||
|
raise ValueError(f"Некорректный ID класса в '{source}': {raw_key!r}.")
|
||||||
|
if isinstance(raw_key, int):
|
||||||
|
class_id = raw_key
|
||||||
|
elif isinstance(raw_key, str):
|
||||||
|
try:
|
||||||
|
class_id = int(raw_key.strip())
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Некорректный ID класса в '{source}': {raw_key!r}."
|
||||||
|
) from exc
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Некорректный ID класса в '{source}': {raw_key!r}.")
|
||||||
|
if class_id < 0:
|
||||||
|
raise ValueError(f"ID класса в '{source}' не может быть отрицательным.")
|
||||||
|
return class_id
|
||||||
|
|
||||||
|
|
||||||
def _normalize_names(names: Any, source: Path) -> dict[int, str]:
|
def _normalize_names(names: Any, source: Path) -> dict[int, str]:
|
||||||
if isinstance(names, list):
|
if isinstance(names, list):
|
||||||
items = enumerate(names)
|
items = enumerate(names)
|
||||||
|
|
@ -15,14 +35,7 @@ def _normalize_names(names: Any, source: Path) -> dict[int, str]:
|
||||||
normalized_items: list[tuple[int, Any]] = []
|
normalized_items: list[tuple[int, Any]] = []
|
||||||
seen: set[int] = set()
|
seen: set[int] = set()
|
||||||
for raw_key, value in names.items():
|
for raw_key, value in names.items():
|
||||||
if isinstance(raw_key, bool):
|
class_id = _normalize_class_key(raw_key, source)
|
||||||
raise ValueError(f"Некорректный ID класса в '{source}': {raw_key!r}.")
|
|
||||||
try:
|
|
||||||
class_id = int(raw_key)
|
|
||||||
except (TypeError, ValueError) as exc:
|
|
||||||
raise ValueError(
|
|
||||||
f"Некорректный ID класса в '{source}': {raw_key!r}."
|
|
||||||
) from exc
|
|
||||||
if class_id in seen:
|
if class_id in seen:
|
||||||
raise ValueError(f"Повторяющийся ID класса {class_id} в '{source}'.")
|
raise ValueError(f"Повторяющийся ID класса {class_id} в '{source}'.")
|
||||||
seen.add(class_id)
|
seen.add(class_id)
|
||||||
|
|
@ -86,6 +99,36 @@ def _parse_classes_file(path: Path) -> dict[int, str]:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_label_class_ids(path: Path) -> set[int]:
|
||||||
|
class_ids: set[int] = set()
|
||||||
|
try:
|
||||||
|
with path.open("r", encoding="utf-8") as label_file:
|
||||||
|
for line_number, line in enumerate(label_file, start=1):
|
||||||
|
parts = line.strip().split()
|
||||||
|
if not parts:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
numeric_id = float(parts[0])
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Некорректный ID класса в '{path}', строка {line_number}: "
|
||||||
|
f"{parts[0]!r}."
|
||||||
|
) from exc
|
||||||
|
if (
|
||||||
|
not math.isfinite(numeric_id)
|
||||||
|
or not numeric_id.is_integer()
|
||||||
|
or numeric_id < 0
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"Некорректный ID класса в '{path}', строка {line_number}: "
|
||||||
|
f"{parts[0]!r}."
|
||||||
|
)
|
||||||
|
class_ids.add(int(numeric_id))
|
||||||
|
except (OSError, UnicodeError) as exc:
|
||||||
|
raise ValueError(f"Не удалось прочитать файл разметки '{path}': {exc}.") from exc
|
||||||
|
return class_ids
|
||||||
|
|
||||||
|
|
||||||
def read_classes(dataset_dir: Path, custom_classes_path: str) -> dict[int, str]:
|
def read_classes(dataset_dir: Path, custom_classes_path: str) -> dict[int, str]:
|
||||||
# An explicit path is authoritative: typos and malformed files must not fall back.
|
# An explicit path is authoritative: typos and malformed files must not fall back.
|
||||||
if custom_classes_path.strip():
|
if custom_classes_path.strip():
|
||||||
|
|
@ -119,17 +162,10 @@ def read_classes(dataset_dir: Path, custom_classes_path: str) -> dict[int, str]:
|
||||||
class_ids: set[int] = set()
|
class_ids: set[int] = set()
|
||||||
labels_dir = dataset_dir / "labels"
|
labels_dir = dataset_dir / "labels"
|
||||||
if labels_dir.exists():
|
if labels_dir.exists():
|
||||||
for txt_file in labels_dir.rglob("*.txt"):
|
for txt_file in sorted(labels_dir.rglob("*.txt")):
|
||||||
if txt_file.name == "classes.txt":
|
if txt_file.name == "classes.txt":
|
||||||
continue
|
continue
|
||||||
try:
|
class_ids.update(_read_label_class_ids(txt_file))
|
||||||
with txt_file.open("r", encoding="utf-8") as label_file:
|
|
||||||
for line in label_file:
|
|
||||||
parts = line.strip().split()
|
|
||||||
if parts:
|
|
||||||
class_ids.add(int(parts[0]))
|
|
||||||
except (OSError, ValueError):
|
|
||||||
continue
|
|
||||||
|
|
||||||
if class_ids:
|
if class_ids:
|
||||||
max_id = max(class_ids)
|
max_id = max(class_ids)
|
||||||
|
|
@ -146,15 +182,52 @@ def _write_new(path: Path, content: str) -> None:
|
||||||
output_file.write(content)
|
output_file.write(content)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_image_labels(
|
||||||
|
images_dir: Path,
|
||||||
|
labels_dir: Path,
|
||||||
|
image_files: list[Path],
|
||||||
|
classes: dict[int, str],
|
||||||
|
) -> None:
|
||||||
|
checked_paths: set[Path] = set()
|
||||||
|
for image_path in image_files:
|
||||||
|
label_path = labels_dir / image_path.relative_to(images_dir).with_suffix(".txt")
|
||||||
|
if label_path in checked_paths:
|
||||||
|
continue
|
||||||
|
checked_paths.add(label_path)
|
||||||
|
if not label_path.exists():
|
||||||
|
# YOLO treats an image without a label file as a background image.
|
||||||
|
continue
|
||||||
|
if not label_path.is_file():
|
||||||
|
raise ValueError(f"Путь к разметке '{label_path}' не является файлом.")
|
||||||
|
for class_id in _read_label_class_ids(label_path):
|
||||||
|
if class_id not in classes:
|
||||||
|
raise ValueError(
|
||||||
|
f"В файле '{label_path}' указан класс {class_id}, "
|
||||||
|
f"но список классов содержит ID от 0 до {len(classes) - 1}."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def split_dataset(
|
def split_dataset(
|
||||||
dataset_dir: str, train_ratio: float, classes_path: str
|
dataset_dir: str, train_ratio: float, classes_path: str
|
||||||
) -> tuple[int, int, str]:
|
) -> tuple[int, int, str]:
|
||||||
|
if not isinstance(dataset_dir, str):
|
||||||
|
raise ValueError("Каталог датасета должен быть строкой.")
|
||||||
if not dataset_dir.strip():
|
if not dataset_dir.strip():
|
||||||
raise ValueError("Укажите каталог датасета.")
|
raise ValueError("Укажите каталог датасета.")
|
||||||
if not 0.1 <= train_ratio <= 0.95:
|
if not isinstance(classes_path, str):
|
||||||
|
raise ValueError("Путь к файлу классов должен быть строкой.")
|
||||||
|
if isinstance(train_ratio, bool) or not isinstance(train_ratio, (int, float)):
|
||||||
|
raise ValueError("Доля обучающей выборки должна быть числом.")
|
||||||
|
try:
|
||||||
|
ratio = float(train_ratio)
|
||||||
|
except OverflowError as exc:
|
||||||
|
raise ValueError("Доля обучающей выборки должна быть конечным числом.") from exc
|
||||||
|
if not math.isfinite(ratio):
|
||||||
|
raise ValueError("Доля обучающей выборки должна быть конечным числом.")
|
||||||
|
if not 0.1 <= ratio <= 0.95:
|
||||||
raise ValueError("Доля обучающей выборки должна быть от 0.1 до 0.95.")
|
raise ValueError("Доля обучающей выборки должна быть от 0.1 до 0.95.")
|
||||||
|
|
||||||
base_dir = Path(dataset_dir.strip()).expanduser().absolute()
|
base_dir = Path(dataset_dir.strip()).expanduser().resolve(strict=False)
|
||||||
images_dir = base_dir / "images"
|
images_dir = base_dir / "images"
|
||||||
labels_dir = base_dir / "labels"
|
labels_dir = base_dir / "labels"
|
||||||
|
|
||||||
|
|
@ -185,13 +258,14 @@ def split_dataset(
|
||||||
rng = random.Random(42)
|
rng = random.Random(42)
|
||||||
rng.shuffle(image_files)
|
rng.shuffle(image_files)
|
||||||
|
|
||||||
split_idx = int(len(image_files) * train_ratio)
|
split_idx = int(len(image_files) * ratio)
|
||||||
split_idx = max(1, min(split_idx, len(image_files) - 1))
|
split_idx = max(1, min(split_idx, len(image_files) - 1))
|
||||||
train_images = image_files[:split_idx]
|
train_images = image_files[:split_idx]
|
||||||
val_images = image_files[split_idx:]
|
val_images = image_files[split_idx:]
|
||||||
|
|
||||||
# Resolve classes before creating output so invalid input leaves no partial split.
|
# Resolve classes before creating output so invalid input leaves no partial split.
|
||||||
classes = read_classes(base_dir, classes_path)
|
classes = read_classes(base_dir, classes_path)
|
||||||
|
_validate_image_labels(images_dir, labels_dir, image_files, classes)
|
||||||
|
|
||||||
relative_split_dir = Path(".yolo-webui") / "splits" / uuid4().hex
|
relative_split_dir = Path(".yolo-webui") / "splits" / uuid4().hex
|
||||||
split_dir = base_dir / relative_split_dir
|
split_dir = base_dir / relative_split_dir
|
||||||
|
|
|
||||||
|
|
@ -1,14 +1,213 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import math
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import traceback
|
import traceback
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
# Force headless Matplotlib to avoid any thread/process GUI issues
|
from yolo_webui.config import MODEL_SUFFIXES, _allowed_roots, _is_within
|
||||||
|
|
||||||
|
# Force headless Matplotlib to avoid any thread/process GUI issues.
|
||||||
os.environ["MPLBACKEND"] = "Agg"
|
os.environ["MPLBACKEND"] = "Agg"
|
||||||
|
# Restrict PyTorch checkpoint deserialization for direct CLI invocation too.
|
||||||
|
os.environ["ULTRALYTICS_SAFE_LOAD"] = "1"
|
||||||
|
|
||||||
def main() -> int:
|
SUPPORTED_EXPORT_FORMATS = frozenset(
|
||||||
args = sys.argv[1:]
|
{
|
||||||
|
"coreml",
|
||||||
|
"engine",
|
||||||
|
"onnx",
|
||||||
|
"openvino",
|
||||||
|
"pb",
|
||||||
|
"saved_model",
|
||||||
|
"tflite",
|
||||||
|
"torchscript",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
EXPORT_CONFIG_FIELDS = frozenset(
|
||||||
|
{
|
||||||
|
"model",
|
||||||
|
"format",
|
||||||
|
"imgsz",
|
||||||
|
"half",
|
||||||
|
"int8",
|
||||||
|
"dynamic",
|
||||||
|
"simplify",
|
||||||
|
"batch",
|
||||||
|
"workspace",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _bounded_int(
|
||||||
|
data: Mapping[str, Any],
|
||||||
|
key: str,
|
||||||
|
default: int,
|
||||||
|
*,
|
||||||
|
minimum: int,
|
||||||
|
maximum: int,
|
||||||
|
) -> int:
|
||||||
|
value = data.get(key, default)
|
||||||
|
if type(value) is not int or not minimum <= value <= maximum:
|
||||||
|
raise ValueError(
|
||||||
|
f"Параметр {key} должен быть целым числом от {minimum} до {maximum}."
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _boolean(data: Mapping[str, Any], key: str, default: bool = False) -> bool:
|
||||||
|
value = data.get(key, default)
|
||||||
|
if type(value) is not bool:
|
||||||
|
raise ValueError(f"Параметр {key} должен быть логическим значением.")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _bounded_number(
|
||||||
|
data: Mapping[str, Any],
|
||||||
|
key: str,
|
||||||
|
default: float,
|
||||||
|
*,
|
||||||
|
minimum: float,
|
||||||
|
maximum: float,
|
||||||
|
) -> float:
|
||||||
|
value = data.get(key, default)
|
||||||
|
if type(value) not in (int, float):
|
||||||
|
raise ValueError(
|
||||||
|
f"Параметр {key} должен быть числом от {minimum:g} до {maximum:g}."
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
number = float(value)
|
||||||
|
except OverflowError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"Параметр {key} должен быть числом от {minimum:g} до {maximum:g}."
|
||||||
|
) from exc
|
||||||
|
if not math.isfinite(number) or not minimum <= number <= maximum:
|
||||||
|
raise ValueError(
|
||||||
|
f"Параметр {key} должен быть числом от {minimum:g} до {maximum:g}."
|
||||||
|
)
|
||||||
|
return number
|
||||||
|
|
||||||
|
|
||||||
|
def _model_path(value: Any) -> Path:
|
||||||
|
if not isinstance(value, str) or not value.strip():
|
||||||
|
raise ValueError("Укажите путь к модели.")
|
||||||
|
|
||||||
|
normalized = value.strip()
|
||||||
|
if any(character in normalized for character in ("\0", "\r", "\n")):
|
||||||
|
raise ValueError("Путь к модели содержит недопустимые управляющие символы.")
|
||||||
|
|
||||||
|
candidate = Path(normalized).expanduser()
|
||||||
|
if not candidate.is_absolute():
|
||||||
|
if len(candidate.parts) == 1:
|
||||||
|
candidate = Path("models") / candidate
|
||||||
|
candidate = Path.cwd() / candidate
|
||||||
|
|
||||||
|
try:
|
||||||
|
resolved = candidate.resolve(strict=False)
|
||||||
|
except (OSError, RuntimeError, ValueError) as exc:
|
||||||
|
raise ValueError(f"Некорректный путь к модели: {normalized}") from exc
|
||||||
|
|
||||||
|
roots = _allowed_roots(("models", "runs"), "YOLO_WEBUI_MODEL_ROOTS")
|
||||||
|
if not _is_within(resolved, roots):
|
||||||
|
allowed = ", ".join(str(root) for root in roots)
|
||||||
|
raise ValueError(
|
||||||
|
f"Модель должна находиться в разрешённом каталоге: {allowed}."
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
is_file = resolved.is_file()
|
||||||
|
except OSError as exc:
|
||||||
|
raise ValueError(f"Не удалось проверить файл модели: {normalized}") from exc
|
||||||
|
if not is_file:
|
||||||
|
raise ValueError(f"Файл модели не найден: {normalized}")
|
||||||
|
if resolved.suffix.lower() not in MODEL_SUFFIXES:
|
||||||
|
expected = ", ".join(sorted(MODEL_SUFFIXES))
|
||||||
|
raise ValueError(f"Модель должна иметь расширение {expected}.")
|
||||||
|
return resolved
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ExportConfig:
|
||||||
|
model: Path
|
||||||
|
export_format: str
|
||||||
|
imgsz: int
|
||||||
|
half: bool
|
||||||
|
int8: bool
|
||||||
|
dynamic: bool
|
||||||
|
simplify: bool
|
||||||
|
batch: int
|
||||||
|
workspace: float
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_mapping(cls, data: Mapping[str, Any]) -> ExportConfig:
|
||||||
|
unknown = set(data) - EXPORT_CONFIG_FIELDS
|
||||||
|
if unknown:
|
||||||
|
fields = ", ".join(sorted(map(str, unknown)))
|
||||||
|
raise ValueError(f"Неизвестные параметры экспорта: {fields}.")
|
||||||
|
|
||||||
|
raw_format = data.get("format", "onnx")
|
||||||
|
if not isinstance(raw_format, str):
|
||||||
|
raise ValueError("Параметр format должен быть строкой.")
|
||||||
|
export_format = raw_format.strip().lower()
|
||||||
|
if export_format not in SUPPORTED_EXPORT_FORMATS:
|
||||||
|
supported = ", ".join(sorted(SUPPORTED_EXPORT_FORMATS))
|
||||||
|
raise ValueError(
|
||||||
|
f"Неподдерживаемый формат экспорта: {raw_format}. "
|
||||||
|
f"Доступны: {supported}."
|
||||||
|
)
|
||||||
|
|
||||||
|
half = _boolean(data, "half")
|
||||||
|
int8 = _boolean(data, "int8")
|
||||||
|
if half and int8:
|
||||||
|
raise ValueError("Параметры half и int8 нельзя включать одновременно.")
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
model=_model_path(data.get("model")),
|
||||||
|
export_format=export_format,
|
||||||
|
imgsz=_bounded_int(data, "imgsz", 640, minimum=32, maximum=8192),
|
||||||
|
half=half,
|
||||||
|
int8=int8,
|
||||||
|
dynamic=_boolean(data, "dynamic"),
|
||||||
|
simplify=_boolean(data, "simplify"),
|
||||||
|
batch=_bounded_int(data, "batch", 1, minimum=1, maximum=1024),
|
||||||
|
workspace=_bounded_number(
|
||||||
|
data,
|
||||||
|
"workspace",
|
||||||
|
4,
|
||||||
|
minimum=1,
|
||||||
|
maximum=64,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _result_path(value: Any) -> Path:
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
if not value:
|
||||||
|
raise RuntimeError("Ultralytics не вернул путь к экспортированной модели.")
|
||||||
|
value = value[0]
|
||||||
|
if not isinstance(value, (str, os.PathLike)):
|
||||||
|
raise RuntimeError("Ultralytics не вернул путь к экспортированной модели.")
|
||||||
|
|
||||||
|
raw_path = os.fspath(value)
|
||||||
|
if not raw_path or any(character in raw_path for character in ("\0", "\r", "\n")):
|
||||||
|
raise RuntimeError("Ultralytics вернул некорректный путь результата.")
|
||||||
|
path = Path(raw_path).expanduser()
|
||||||
|
if not path.is_absolute():
|
||||||
|
path = Path.cwd() / path
|
||||||
|
try:
|
||||||
|
return path.resolve(strict=True)
|
||||||
|
except (FileNotFoundError, OSError, RuntimeError, ValueError) as exc:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Экспорт завершился без ожидаемого артефакта: {raw_path}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
|
args = list(sys.argv[1:] if argv is None else argv)
|
||||||
if not args:
|
if not args:
|
||||||
print(
|
print(
|
||||||
"Usage: python -m yolo_webui.export_runner <config_json_path>",
|
"Usage: python -m yolo_webui.export_runner <config_json_path>",
|
||||||
|
|
@ -16,63 +215,58 @@ def main() -> int:
|
||||||
)
|
)
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
config_path = Path(args[0])
|
# The parent waits for this marker before treating the process as started.
|
||||||
|
|
||||||
# The parent waits for this marker
|
|
||||||
print("__YOLO_WEBUI_READY__", flush=True)
|
print("__YOLO_WEBUI_READY__", flush=True)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with config_path.open("r", encoding="utf-8") as config_file:
|
with Path(args[0]).open("r", encoding="utf-8") as config_file:
|
||||||
config_dict = json.load(config_file)
|
raw_config = json.load(config_file)
|
||||||
|
if not isinstance(raw_config, Mapping):
|
||||||
|
raise ValueError("Конфигурация экспорта должна быть JSON-объектом.")
|
||||||
|
config = ExportConfig.from_mapping(raw_config)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
print(f"Error loading config: {exc}", file=sys.stderr)
|
print(f"Ошибка конфигурации экспорта: {exc}", file=sys.stderr)
|
||||||
return 1
|
|
||||||
|
|
||||||
model_path = config_dict.get("model")
|
|
||||||
export_format = config_dict.get("format", "onnx")
|
|
||||||
imgsz = config_dict.get("imgsz", 640)
|
|
||||||
half = config_dict.get("half", False)
|
|
||||||
int8 = config_dict.get("int8", False)
|
|
||||||
dynamic = config_dict.get("dynamic", False)
|
|
||||||
simplify = config_dict.get("simplify", False)
|
|
||||||
batch = config_dict.get("batch", 1)
|
|
||||||
workspace = config_dict.get("workspace", 4)
|
|
||||||
|
|
||||||
if not model_path or not Path(model_path).exists():
|
|
||||||
print(f"Model file not found: {model_path}")
|
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
os.environ["ULTRALYTICS_SAFE_LOAD"] = "1"
|
||||||
from ultralytics import YOLO
|
from ultralytics import YOLO
|
||||||
|
|
||||||
print(f"Загрузка модели {model_path}...", flush=True)
|
print(f"Загрузка модели {config.model}...", flush=True)
|
||||||
model = YOLO(model_path)
|
model = YOLO(str(config.model))
|
||||||
|
|
||||||
print(f"Запуск экспорта в формат {export_format}...", flush=True)
|
print(
|
||||||
print(f"Параметры: imgsz={imgsz}, half={half}, int8={int8}, dynamic={dynamic}, simplify={simplify}, batch={batch}, workspace={workspace}", flush=True)
|
f"Запуск экспорта в формат {config.export_format}...",
|
||||||
|
flush=True,
|
||||||
exported_path = model.export(
|
)
|
||||||
format=export_format,
|
print(
|
||||||
imgsz=imgsz,
|
"Параметры: "
|
||||||
half=half,
|
f"imgsz={config.imgsz}, half={config.half}, int8={config.int8}, "
|
||||||
int8=int8,
|
f"dynamic={config.dynamic}, simplify={config.simplify}, "
|
||||||
dynamic=dynamic,
|
f"batch={config.batch}, workspace={config.workspace}",
|
||||||
simplify=simplify,
|
flush=True,
|
||||||
batch=batch,
|
|
||||||
workspace=workspace
|
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"Экспорт завершен успешно.", flush=True)
|
exported_path = model.export(
|
||||||
if exported_path:
|
format=config.export_format,
|
||||||
# ultralytics returns either a single path (string) or list of paths depending on the format.
|
imgsz=config.imgsz,
|
||||||
if isinstance(exported_path, list):
|
half=config.half,
|
||||||
exported_path = exported_path[0]
|
int8=config.int8,
|
||||||
print(f"__YOLO_WEBUI_RESULT__:{exported_path}", flush=True)
|
dynamic=config.dynamic,
|
||||||
|
simplify=config.simplify,
|
||||||
|
batch=config.batch,
|
||||||
|
workspace=config.workspace,
|
||||||
|
)
|
||||||
|
result_path = _result_path(exported_path)
|
||||||
|
|
||||||
|
print("Экспорт завершен успешно.", flush=True)
|
||||||
|
print(f"__YOLO_WEBUI_RESULT__:{result_path}", flush=True)
|
||||||
return 0
|
return 0
|
||||||
except Exception as e:
|
except Exception as exc:
|
||||||
print(f"Ошибка при экспорте модели: {e}", flush=True)
|
print(f"Ошибка при экспорте модели: {exc}", flush=True)
|
||||||
traceback.print_exc()
|
traceback.print_exc()
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
raise SystemExit(main())
|
raise SystemExit(main())
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,10 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
};
|
};
|
||||||
let discoveredModels = [];
|
let discoveredModels = [];
|
||||||
|
|
||||||
|
function getDiscoveredModelValue(model) {
|
||||||
|
return model.path || model.name;
|
||||||
|
}
|
||||||
|
|
||||||
function updateModelOptions() {
|
function updateModelOptions() {
|
||||||
const task = taskSelect.value;
|
const task = taskSelect.value;
|
||||||
const stdModels = standardModels[task] || [];
|
const stdModels = standardModels[task] || [];
|
||||||
|
|
@ -54,13 +58,13 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
modelSelect.appendChild(stdGroup);
|
modelSelect.appendChild(stdGroup);
|
||||||
|
|
||||||
// Group 2: Discovered Models
|
// Group 2: Discovered Models
|
||||||
const localModels = discoveredModels.filter(m => !stdModels.includes(m.name));
|
const localModels = discoveredModels.filter(m => getDiscoveredModelValue(m));
|
||||||
if (localModels.length > 0) {
|
if (localModels.length > 0) {
|
||||||
const localGroup = document.createElement('optgroup');
|
const localGroup = document.createElement('optgroup');
|
||||||
localGroup.label = 'Локальные/скачанные модели';
|
localGroup.label = 'Локальные/скачанные модели';
|
||||||
localModels.forEach(m => {
|
localModels.forEach(m => {
|
||||||
const opt = document.createElement('option');
|
const opt = document.createElement('option');
|
||||||
opt.value = m.name;
|
opt.value = getDiscoveredModelValue(m);
|
||||||
opt.textContent = m.name;
|
opt.textContent = m.name;
|
||||||
localGroup.appendChild(opt);
|
localGroup.appendChild(opt);
|
||||||
});
|
});
|
||||||
|
|
@ -74,8 +78,8 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
modelSelect.appendChild(customOpt);
|
modelSelect.appendChild(customOpt);
|
||||||
|
|
||||||
// Match selection if valid
|
// Match selection if valid
|
||||||
const allAvailable = [...stdModels, ...localModels.map(m => m.name)];
|
const allAvailable = [...stdModels, ...localModels.map(getDiscoveredModelValue)];
|
||||||
if (allAvailable.includes(currentSelectVal)) {
|
if (currentSelectVal === '__custom__' || allAvailable.includes(currentSelectVal)) {
|
||||||
modelSelect.value = currentSelectVal;
|
modelSelect.value = currentSelectVal;
|
||||||
} else {
|
} else {
|
||||||
modelSelect.value = stdModels[0] || '__custom__';
|
modelSelect.value = stdModels[0] || '__custom__';
|
||||||
|
|
@ -129,7 +133,11 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
let trainingTimer = null;
|
let trainingTimer = null;
|
||||||
let secondsElapsed = 0;
|
let secondsElapsed = 0;
|
||||||
let socket = null;
|
let socket = null;
|
||||||
|
let socketReconnectTimer = null;
|
||||||
let isTrainingActive = false;
|
let isTrainingActive = false;
|
||||||
|
let isTrainingStartPending = false;
|
||||||
|
let trainingStatusRevision = 0;
|
||||||
|
let currentTrainingStatus = 'idle';
|
||||||
|
|
||||||
// --- View Tab Switching ---
|
// --- View Tab Switching ---
|
||||||
const viewBtns = document.querySelectorAll('.view-btn');
|
const viewBtns = document.querySelectorAll('.view-btn');
|
||||||
|
|
@ -215,7 +223,17 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
function updateMlflowHeaderLink() {
|
function updateMlflowHeaderLink() {
|
||||||
const uri = trackingUriInput.value.trim();
|
const uri = trackingUriInput.value.trim();
|
||||||
if (uri.startsWith('http://') || uri.startsWith('https://')) {
|
if (uri.startsWith('http://') || uri.startsWith('https://')) {
|
||||||
mlflowHeaderLink.href = uri;
|
let browserUri = uri;
|
||||||
|
try {
|
||||||
|
const parsedUri = new URL(uri);
|
||||||
|
if (parsedUri.hostname === 'mlflow') {
|
||||||
|
parsedUri.hostname = window.location.hostname || 'localhost';
|
||||||
|
browserUri = parsedUri.href;
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Invalid MLflow tracking URI:', error);
|
||||||
|
}
|
||||||
|
mlflowHeaderLink.href = browserUri;
|
||||||
mlflowHeaderLink.style.opacity = '1';
|
mlflowHeaderLink.style.opacity = '1';
|
||||||
mlflowHeaderLink.style.pointerEvents = 'auto';
|
mlflowHeaderLink.style.pointerEvents = 'auto';
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -343,41 +361,75 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- WebSocket Sync ---
|
// --- WebSocket Sync ---
|
||||||
|
function scheduleWebSocketReconnect() {
|
||||||
|
if (socketReconnectTimer !== null) return;
|
||||||
|
socketReconnectTimer = setTimeout(() => {
|
||||||
|
socketReconnectTimer = null;
|
||||||
|
connectWebSocket();
|
||||||
|
}, 5000);
|
||||||
|
}
|
||||||
|
|
||||||
|
function parseWebSocketMessage(event, label, logLine) {
|
||||||
|
try {
|
||||||
|
const data = JSON.parse(event.data);
|
||||||
|
if (!data || typeof data !== 'object' || Array.isArray(data)) {
|
||||||
|
throw new Error('WebSocket payload must be an object');
|
||||||
|
}
|
||||||
|
return data;
|
||||||
|
} catch (error) {
|
||||||
|
console.error(`${label} message error:`, error);
|
||||||
|
logLine('Получено некорректное сообщение от сервера.', 'warning');
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
function connectWebSocket() {
|
function connectWebSocket() {
|
||||||
const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
|
const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
|
||||||
const wsUrl = `${protocol}//${window.location.host}/api/ws`;
|
const wsUrl = `${protocol}//${window.location.host}/api/ws`;
|
||||||
|
|
||||||
socket = new WebSocket(wsUrl);
|
let nextSocket;
|
||||||
|
try {
|
||||||
|
nextSocket = new WebSocket(wsUrl);
|
||||||
|
} catch (error) {
|
||||||
|
console.error('WS connection error:', error);
|
||||||
|
addLogLine('Не удалось подключиться к серверу. Повторная попытка через 5 секунд...', 'warning');
|
||||||
|
scheduleWebSocketReconnect();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
socket = nextSocket;
|
||||||
|
|
||||||
socket.onopen = () => {
|
nextSocket.onopen = () => {
|
||||||
addLogLine('Соединение с сервером установлено.', 'info');
|
addLogLine('Соединение с сервером установлено.', 'info');
|
||||||
};
|
};
|
||||||
|
|
||||||
socket.onclose = () => {
|
nextSocket.onclose = () => {
|
||||||
|
if (socket !== nextSocket) return;
|
||||||
|
socket = null;
|
||||||
addLogLine('Соединение потеряно. Повторная попытка через 5 секунд...', 'warning');
|
addLogLine('Соединение потеряно. Повторная попытка через 5 секунд...', 'warning');
|
||||||
setTimeout(connectWebSocket, 5000);
|
scheduleWebSocketReconnect();
|
||||||
};
|
};
|
||||||
|
|
||||||
socket.onerror = (err) => {
|
nextSocket.onerror = (err) => {
|
||||||
console.error('WS Error:', err);
|
console.error('WS Error:', err);
|
||||||
};
|
};
|
||||||
|
|
||||||
socket.onmessage = (event) => {
|
nextSocket.onmessage = (event) => {
|
||||||
const data = JSON.parse(event.data);
|
const data = parseWebSocketMessage(event, 'Training WebSocket', addLogLine);
|
||||||
|
if (!data) return;
|
||||||
|
|
||||||
if (data.type === 'init') {
|
if (data.type === 'init') {
|
||||||
updateUIStatus(data.status);
|
updateUIStatus(data.status);
|
||||||
|
|
||||||
// Load logs
|
// Load logs
|
||||||
logContainer.innerHTML = '';
|
logContainer.innerHTML = '';
|
||||||
data.logs.forEach(([levelCode, msg]) => {
|
(Array.isArray(data.logs) ? data.logs : []).forEach(([levelCode, msg]) => {
|
||||||
const level = levelCode.replace('__LOG_LEVEL_', '').replace('__', '').toLowerCase();
|
const level = levelCode.replace('__LOG_LEVEL_', '').replace('__', '').toLowerCase();
|
||||||
addLogLine(msg, level);
|
addLogLine(msg, level);
|
||||||
});
|
});
|
||||||
|
|
||||||
// Draw initial chart points
|
// Draw initial chart points
|
||||||
initChart();
|
initChart();
|
||||||
if (data.metrics && data.metrics.length > 0) {
|
if (Array.isArray(data.metrics) && data.metrics.length > 0) {
|
||||||
data.metrics.forEach(m => {
|
data.metrics.forEach(m => {
|
||||||
updateChart(m.epoch, m);
|
updateChart(m.epoch, m);
|
||||||
});
|
});
|
||||||
|
|
@ -424,6 +476,8 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
function updateUIStatus(status) {
|
function updateUIStatus(status) {
|
||||||
|
trainingStatusRevision++;
|
||||||
|
currentTrainingStatus = status;
|
||||||
statusCard.className = `status-${status}`;
|
statusCard.className = `status-${status}`;
|
||||||
|
|
||||||
switch (status) {
|
switch (status) {
|
||||||
|
|
@ -505,6 +559,25 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
return Number.isNaN(value) ? fallback : value;
|
return Number.isNaN(value) ? fallback : value;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function readValidatedNumber(id, label, {integer = false, min, max} = {}) {
|
||||||
|
const rawValue = document.getElementById(id).value.trim();
|
||||||
|
const value = Number(rawValue);
|
||||||
|
const outOfRange = (min !== undefined && value < min)
|
||||||
|
|| (max !== undefined && value > max);
|
||||||
|
if (
|
||||||
|
rawValue === ''
|
||||||
|
|| !Number.isFinite(value)
|
||||||
|
|| (integer && !Number.isInteger(value))
|
||||||
|
|| outOfRange
|
||||||
|
) {
|
||||||
|
const range = max === undefined ? `не меньше ${min}` : `от ${min} до ${max}`;
|
||||||
|
const integerHint = integer ? 'целым числом ' : '';
|
||||||
|
showNotification(`${label} должен быть ${integerHint}${range}.`, 'warning');
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
|
||||||
function getFormConfig() {
|
function getFormConfig() {
|
||||||
return {
|
return {
|
||||||
dataset: document.getElementById('dataset').value.trim(),
|
dataset: document.getElementById('dataset').value.trim(),
|
||||||
|
|
@ -561,11 +634,18 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
const modelVal = data.model || 'yolo11n.pt';
|
const modelVal = data.model || 'yolo11n.pt';
|
||||||
const task = data.task || 'detect';
|
const task = data.task || 'detect';
|
||||||
const stdModels = standardModels[task] || [];
|
const stdModels = standardModels[task] || [];
|
||||||
const allAvailable = [...stdModels, ...discoveredModels.map(m => m.name)];
|
const matchedLocalModel = discoveredModels.find(
|
||||||
if (allAvailable.includes(modelVal)) {
|
model => model.path === modelVal || model.name === modelVal
|
||||||
|
);
|
||||||
|
if (stdModels.includes(modelVal)) {
|
||||||
modelSelect.value = modelVal;
|
modelSelect.value = modelVal;
|
||||||
modelCustomWrapper.style.display = 'none';
|
modelCustomWrapper.style.display = 'none';
|
||||||
document.getElementById('model').value = modelVal;
|
document.getElementById('model').value = modelVal;
|
||||||
|
} else if (matchedLocalModel) {
|
||||||
|
const discoveredValue = getDiscoveredModelValue(matchedLocalModel);
|
||||||
|
modelSelect.value = discoveredValue;
|
||||||
|
modelCustomWrapper.style.display = 'none';
|
||||||
|
document.getElementById('model').value = discoveredValue;
|
||||||
} else {
|
} else {
|
||||||
modelSelect.value = '__custom__';
|
modelSelect.value = '__custom__';
|
||||||
modelCustomWrapper.style.display = 'block';
|
modelCustomWrapper.style.display = 'block';
|
||||||
|
|
@ -831,9 +911,9 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
});
|
});
|
||||||
|
|
||||||
sessionSaveBtn.addEventListener('click', async () => {
|
sessionSaveBtn.addEventListener('click', async () => {
|
||||||
const name = sessionNameInput.value.trim().replace(/[^a-zA-Z0-9_\-]/g, "");
|
const name = sessionNameInput.value.trim();
|
||||||
if (!name) {
|
if (!/^[a-zA-Z0-9_-]+$/.test(name)) {
|
||||||
showNotification('Введите корректное имя профиля (латиница, цифры, дефисы).', 'warning');
|
showNotification('Введите корректное имя профиля (латиница, цифры, дефисы и подчёркивания).', 'warning');
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
if (name === "last_run") {
|
if (name === "last_run") {
|
||||||
|
|
@ -888,8 +968,12 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
|
|
||||||
// --- Form submit ---
|
// --- Form submit ---
|
||||||
async function startTraining() {
|
async function startTraining() {
|
||||||
if (isTrainingActive) return;
|
if (isTrainingActive || isTrainingStartPending) return;
|
||||||
|
isTrainingStartPending = true;
|
||||||
|
startBtn.disabled = true;
|
||||||
const config = getFormConfig();
|
const config = getFormConfig();
|
||||||
|
const statusRevisionAtStart = trainingStatusRevision;
|
||||||
|
let started = false;
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const res = await fetch('/api/train/start', {
|
const res = await fetch('/api/train/start', {
|
||||||
|
|
@ -902,10 +986,21 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
if (!res.ok) {
|
if (!res.ok) {
|
||||||
throw new Error(data.detail || 'Failed to start training');
|
throw new Error(data.detail || 'Failed to start training');
|
||||||
}
|
}
|
||||||
|
started = true;
|
||||||
|
const terminalStatusReceived = trainingStatusRevision !== statusRevisionAtStart
|
||||||
|
&& ['finished', 'succeeded', 'cancelled', 'failed'].includes(currentTrainingStatus);
|
||||||
|
if (!isTrainingActive && !terminalStatusReceived) {
|
||||||
|
updateUIStatus('preparing');
|
||||||
|
}
|
||||||
showNotification('Обучение успешно запущено!', 'success');
|
showNotification('Обучение успешно запущено!', 'success');
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
console.error('Start error:', err);
|
console.error('Start error:', err);
|
||||||
showNotification(err.message, 'error');
|
showNotification(err.message, 'error');
|
||||||
|
} finally {
|
||||||
|
isTrainingStartPending = false;
|
||||||
|
if (!started && !isTrainingActive) {
|
||||||
|
startBtn.disabled = false;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1002,7 +1097,11 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
const exportAutoscrollCheck = document.getElementById('export-autoscroll');
|
const exportAutoscrollCheck = document.getElementById('export-autoscroll');
|
||||||
|
|
||||||
let isExportActive = false;
|
let isExportActive = false;
|
||||||
|
let isExportStartPending = false;
|
||||||
let exportSocket = null;
|
let exportSocket = null;
|
||||||
|
let exportSocketReconnectTimer = null;
|
||||||
|
let exportStatusRevision = 0;
|
||||||
|
let currentExportStatus = 'idle';
|
||||||
|
|
||||||
function addExportLogLine(message, level = 'info') {
|
function addExportLogLine(message, level = 'info') {
|
||||||
const line = document.createElement('div');
|
const line = document.createElement('div');
|
||||||
|
|
@ -1019,32 +1118,52 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
exportLogContainer.innerHTML = '';
|
exportLogContainer.innerHTML = '';
|
||||||
});
|
});
|
||||||
|
|
||||||
|
function scheduleExportWebSocketReconnect() {
|
||||||
|
if (exportSocketReconnectTimer !== null) return;
|
||||||
|
exportSocketReconnectTimer = setTimeout(() => {
|
||||||
|
exportSocketReconnectTimer = null;
|
||||||
|
connectExportWebSocket();
|
||||||
|
}, 5000);
|
||||||
|
}
|
||||||
|
|
||||||
function connectExportWebSocket() {
|
function connectExportWebSocket() {
|
||||||
const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
|
const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
|
||||||
const wsUrl = `${protocol}//${window.location.host}/api/export/ws`;
|
const wsUrl = `${protocol}//${window.location.host}/api/export/ws`;
|
||||||
|
|
||||||
exportSocket = new WebSocket(wsUrl);
|
let nextSocket;
|
||||||
|
try {
|
||||||
|
nextSocket = new WebSocket(wsUrl);
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Export WS connection error:', error);
|
||||||
|
addExportLogLine('Не удалось подключиться к серверу. Повторная попытка через 5 секунд...', 'warning');
|
||||||
|
scheduleExportWebSocketReconnect();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
exportSocket = nextSocket;
|
||||||
|
|
||||||
exportSocket.onopen = () => {
|
nextSocket.onopen = () => {
|
||||||
addExportLogLine('Соединение с сервером установлено.', 'info');
|
addExportLogLine('Соединение с сервером установлено.', 'info');
|
||||||
};
|
};
|
||||||
|
|
||||||
exportSocket.onclose = () => {
|
nextSocket.onclose = () => {
|
||||||
|
if (exportSocket !== nextSocket) return;
|
||||||
|
exportSocket = null;
|
||||||
addExportLogLine('Соединение потеряно. Повторная попытка через 5 секунд...', 'warning');
|
addExportLogLine('Соединение потеряно. Повторная попытка через 5 секунд...', 'warning');
|
||||||
setTimeout(connectExportWebSocket, 5000);
|
scheduleExportWebSocketReconnect();
|
||||||
};
|
};
|
||||||
|
|
||||||
exportSocket.onerror = (err) => {
|
nextSocket.onerror = (err) => {
|
||||||
console.error('Export WS Error:', err);
|
console.error('Export WS Error:', err);
|
||||||
};
|
};
|
||||||
|
|
||||||
exportSocket.onmessage = (event) => {
|
nextSocket.onmessage = (event) => {
|
||||||
const data = JSON.parse(event.data);
|
const data = parseWebSocketMessage(event, 'Export WebSocket', addExportLogLine);
|
||||||
|
if (!data) return;
|
||||||
|
|
||||||
if (data.type === 'init') {
|
if (data.type === 'init') {
|
||||||
updateExportUIStatus(data.status);
|
updateExportUIStatus(data.status);
|
||||||
exportLogContainer.innerHTML = '';
|
exportLogContainer.innerHTML = '';
|
||||||
data.logs.forEach(([levelCode, msg]) => {
|
(Array.isArray(data.logs) ? data.logs : []).forEach(([levelCode, msg]) => {
|
||||||
const level = levelCode.replace('__LOG_LEVEL_', '').replace('__', '').toLowerCase();
|
const level = levelCode.replace('__LOG_LEVEL_', '').replace('__', '').toLowerCase();
|
||||||
addExportLogLine(msg, level);
|
addExportLogLine(msg, level);
|
||||||
});
|
});
|
||||||
|
|
@ -1060,6 +1179,8 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
function updateExportUIStatus(status) {
|
function updateExportUIStatus(status) {
|
||||||
|
exportStatusRevision++;
|
||||||
|
currentExportStatus = status;
|
||||||
exportStatusCard.className = `status-${status}`;
|
exportStatusCard.className = `status-${status}`;
|
||||||
|
|
||||||
switch (status) {
|
switch (status) {
|
||||||
|
|
@ -1110,7 +1231,7 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
}
|
}
|
||||||
|
|
||||||
async function startExport() {
|
async function startExport() {
|
||||||
if (isExportActive) return;
|
if (isExportActive || isExportStartPending) return;
|
||||||
|
|
||||||
const modelVal = document.getElementById('export-model').value.trim();
|
const modelVal = document.getElementById('export-model').value.trim();
|
||||||
if (!modelVal) {
|
if (!modelVal) {
|
||||||
|
|
@ -1118,17 +1239,47 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const imgsz = readValidatedNumber(
|
||||||
|
'export-imgsz',
|
||||||
|
'Размер изображения',
|
||||||
|
{integer: true, min: 32, max: 8192}
|
||||||
|
);
|
||||||
|
if (imgsz === null) return;
|
||||||
|
const batch = readValidatedNumber(
|
||||||
|
'export-batch',
|
||||||
|
'Размер батча',
|
||||||
|
{integer: true, min: 1, max: 1024}
|
||||||
|
);
|
||||||
|
if (batch === null) return;
|
||||||
|
const workspace = readValidatedNumber(
|
||||||
|
'export-workspace',
|
||||||
|
'Workspace',
|
||||||
|
{min: 1, max: 64}
|
||||||
|
);
|
||||||
|
if (workspace === null) return;
|
||||||
|
|
||||||
|
const half = document.getElementById('export-half').checked;
|
||||||
|
const int8 = document.getElementById('export-int8').checked;
|
||||||
|
if (half && int8) {
|
||||||
|
showNotification('FP16 и INT8 нельзя включать одновременно.', 'warning');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
isExportStartPending = true;
|
||||||
|
exportStartBtn.disabled = true;
|
||||||
const config = {
|
const config = {
|
||||||
model: modelVal,
|
model: modelVal,
|
||||||
format: document.getElementById('export-format').value,
|
format: document.getElementById('export-format').value,
|
||||||
imgsz: readNumber('export-imgsz', 640, true),
|
imgsz,
|
||||||
half: document.getElementById('export-half').checked,
|
half,
|
||||||
int8: document.getElementById('export-int8').checked,
|
int8,
|
||||||
dynamic: document.getElementById('export-dynamic').checked,
|
dynamic: document.getElementById('export-dynamic').checked,
|
||||||
simplify: document.getElementById('export-simplify').checked,
|
simplify: document.getElementById('export-simplify').checked,
|
||||||
batch: readNumber('export-batch', 1, true),
|
batch,
|
||||||
workspace: readNumber('export-workspace', 4, true)
|
workspace
|
||||||
};
|
};
|
||||||
|
const statusRevisionAtStart = exportStatusRevision;
|
||||||
|
let started = false;
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const res = await fetch('/api/export/start', {
|
const res = await fetch('/api/export/start', {
|
||||||
|
|
@ -1141,10 +1292,21 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||||
if (!res.ok) {
|
if (!res.ok) {
|
||||||
throw new Error(data.detail || 'Failed to start export');
|
throw new Error(data.detail || 'Failed to start export');
|
||||||
}
|
}
|
||||||
|
started = true;
|
||||||
|
const terminalStatusReceived = exportStatusRevision !== statusRevisionAtStart
|
||||||
|
&& ['succeeded', 'cancelled', 'failed'].includes(currentExportStatus);
|
||||||
|
if (!isExportActive && !terminalStatusReceived) {
|
||||||
|
updateExportUIStatus('preparing');
|
||||||
|
}
|
||||||
showNotification('Экспорт успешно запущен!', 'success');
|
showNotification('Экспорт успешно запущен!', 'success');
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
console.error('Export start error:', err);
|
console.error('Export start error:', err);
|
||||||
showNotification(err.message, 'error');
|
showNotification(err.message, 'error');
|
||||||
|
} finally {
|
||||||
|
isExportStartPending = false;
|
||||||
|
if (!started && !isExportActive) {
|
||||||
|
exportStartBtn.disabled = false;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -30,7 +30,7 @@
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="header-actions">
|
<div class="header-actions">
|
||||||
<a href="http://localhost:5000" target="_blank" class="mlflow-link" id="mlflow-header-link">
|
<a href="http://localhost:5000" target="_blank" rel="noopener noreferrer" class="mlflow-link" id="mlflow-header-link">
|
||||||
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" class="link-icon">
|
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" class="link-icon">
|
||||||
<path d="M18 13v6a2 2 0 0 1-2 2H5a2 2 0 0 1-2-2V8a2 2 0 0 1 2-2h6"></path>
|
<path d="M18 13v6a2 2 0 0 1-2 2H5a2 2 0 0 1-2-2V8a2 2 0 0 1 2-2h6"></path>
|
||||||
<polyline points="15 3 21 3 21 9"></polyline>
|
<polyline points="15 3 21 3 21 9"></polyline>
|
||||||
|
|
@ -70,7 +70,7 @@
|
||||||
<div class="session-row">
|
<div class="session-row">
|
||||||
<div class="session-field">
|
<div class="session-field">
|
||||||
<label for="session-name">Сохранить текущие настройки как профиль</label>
|
<label for="session-name">Сохранить текущие настройки как профиль</label>
|
||||||
<input type="text" id="session-name" placeholder="Введите имя профиля...">
|
<input type="text" id="session-name" pattern="[A-Za-z0-9_-]+" autocomplete="off" placeholder="Введите имя профиля...">
|
||||||
</div>
|
</div>
|
||||||
<button type="button" id="session-save-btn" class="btn-action btn-save" title="Сохранить настройки">
|
<button type="button" id="session-save-btn" class="btn-action btn-save" title="Сохранить настройки">
|
||||||
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" class="btn-icon-small"><path d="M19 21H5a2 2 0 0 1-2-2V5a2 2 0 0 1 2-2h11l5 5v11a2 2 0 0 1-2 2z"></path><polyline points="17 21 17 13 7 13 7 21"></polyline></svg>
|
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" class="btn-icon-small"><path d="M19 21H5a2 2 0 0 1-2-2V5a2 2 0 0 1 2-2h11l5 5v11a2 2 0 0 1-2 2z"></path><polyline points="17 21 17 13 7 13 7 21"></polyline></svg>
|
||||||
|
|
@ -433,11 +433,11 @@
|
||||||
<option value="onnx">ONNX (.onnx)</option>
|
<option value="onnx">ONNX (.onnx)</option>
|
||||||
<option value="engine">TensorRT (.engine)</option>
|
<option value="engine">TensorRT (.engine)</option>
|
||||||
<option value="openvino">OpenVINO (_openvino_model/)</option>
|
<option value="openvino">OpenVINO (_openvino_model/)</option>
|
||||||
<option value="triton">Triton Inference Server (triton/)</option>
|
|
||||||
<option value="torchscript">TorchScript (.torchscript)</option>
|
<option value="torchscript">TorchScript (.torchscript)</option>
|
||||||
<option value="coreml">CoreML (.mlpackage)</option>
|
<option value="coreml">CoreML (.mlpackage)</option>
|
||||||
<option value="tflite">TFLite (.tflite)</option>
|
<option value="tflite">TFLite (.tflite)</option>
|
||||||
<option value="pb">TensorFlow SavedModel (_saved_model/)</option>
|
<option value="saved_model">TensorFlow SavedModel (_saved_model/)</option>
|
||||||
|
<option value="pb">TensorFlow GraphDef (.pb)</option>
|
||||||
</select>
|
</select>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
|
@ -445,15 +445,15 @@
|
||||||
<div class="row">
|
<div class="row">
|
||||||
<div class="field">
|
<div class="field">
|
||||||
<label for="export-imgsz">Размер изображения (imgsz)</label>
|
<label for="export-imgsz">Размер изображения (imgsz)</label>
|
||||||
<input type="number" id="export-imgsz" name="export-imgsz" value="640">
|
<input type="number" id="export-imgsz" name="export-imgsz" value="640" min="32" max="8192" step="1" required>
|
||||||
</div>
|
</div>
|
||||||
<div class="field">
|
<div class="field">
|
||||||
<label for="export-batch" title="Максимальный размер батча. Важен при Dynamic=True для TensorRT и Triton.">Макс. батч (batch)</label>
|
<label for="export-batch" title="Максимальный размер батча. Важен при Dynamic=True для TensorRT.">Макс. батч (batch)</label>
|
||||||
<input type="number" id="export-batch" name="export-batch" value="1" min="1">
|
<input type="number" id="export-batch" name="export-batch" value="1" min="1" max="1024" step="1" required>
|
||||||
</div>
|
</div>
|
||||||
<div class="field">
|
<div class="field">
|
||||||
<label for="export-workspace" title="Размер памяти для сборки TensorRT (ГБ).">Workspace (ГБ)</label>
|
<label for="export-workspace" title="Размер памяти для сборки TensorRT (ГБ).">Workspace (ГБ)</label>
|
||||||
<input type="number" id="export-workspace" name="export-workspace" value="4" min="1" max="64">
|
<input type="number" id="export-workspace" name="export-workspace" value="4" min="1" max="64" step="0.5" required>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="row">
|
<div class="row">
|
||||||
|
|
|
||||||
|
|
@ -171,12 +171,14 @@ body {
|
||||||
height: 100%;
|
height: 100%;
|
||||||
}
|
}
|
||||||
|
|
||||||
#config-pane {
|
#config-pane,
|
||||||
|
#export-config-pane {
|
||||||
padding: 1.25rem;
|
padding: 1.25rem;
|
||||||
max-height: 100%;
|
max-height: 100%;
|
||||||
}
|
}
|
||||||
|
|
||||||
#run-pane {
|
#run-pane,
|
||||||
|
#export-run-pane {
|
||||||
padding: 1.25rem;
|
padding: 1.25rem;
|
||||||
display: flex;
|
display: flex;
|
||||||
flex-direction: column;
|
flex-direction: column;
|
||||||
|
|
@ -187,28 +189,36 @@ body {
|
||||||
|
|
||||||
/* Custom Scrollbars */
|
/* Custom Scrollbars */
|
||||||
#run-pane::-webkit-scrollbar,
|
#run-pane::-webkit-scrollbar,
|
||||||
|
#export-run-pane::-webkit-scrollbar,
|
||||||
.form-container::-webkit-scrollbar,
|
.form-container::-webkit-scrollbar,
|
||||||
#log-container::-webkit-scrollbar {
|
#log-container::-webkit-scrollbar,
|
||||||
|
#export-log-container::-webkit-scrollbar {
|
||||||
width: 6px;
|
width: 6px;
|
||||||
height: 6px;
|
height: 6px;
|
||||||
}
|
}
|
||||||
|
|
||||||
#run-pane::-webkit-scrollbar-track,
|
#run-pane::-webkit-scrollbar-track,
|
||||||
|
#export-run-pane::-webkit-scrollbar-track,
|
||||||
.form-container::-webkit-scrollbar-track,
|
.form-container::-webkit-scrollbar-track,
|
||||||
#log-container::-webkit-scrollbar-track {
|
#log-container::-webkit-scrollbar-track,
|
||||||
|
#export-log-container::-webkit-scrollbar-track {
|
||||||
background: transparent;
|
background: transparent;
|
||||||
}
|
}
|
||||||
|
|
||||||
#run-pane::-webkit-scrollbar-thumb,
|
#run-pane::-webkit-scrollbar-thumb,
|
||||||
|
#export-run-pane::-webkit-scrollbar-thumb,
|
||||||
.form-container::-webkit-scrollbar-thumb,
|
.form-container::-webkit-scrollbar-thumb,
|
||||||
#log-container::-webkit-scrollbar-thumb {
|
#log-container::-webkit-scrollbar-thumb,
|
||||||
|
#export-log-container::-webkit-scrollbar-thumb {
|
||||||
background: var(--border-color);
|
background: var(--border-color);
|
||||||
border-radius: 3px;
|
border-radius: 3px;
|
||||||
}
|
}
|
||||||
|
|
||||||
#run-pane::-webkit-scrollbar-thumb:hover,
|
#run-pane::-webkit-scrollbar-thumb:hover,
|
||||||
|
#export-run-pane::-webkit-scrollbar-thumb:hover,
|
||||||
.form-container::-webkit-scrollbar-thumb:hover,
|
.form-container::-webkit-scrollbar-thumb:hover,
|
||||||
#log-container::-webkit-scrollbar-thumb:hover {
|
#log-container::-webkit-scrollbar-thumb:hover,
|
||||||
|
#export-log-container::-webkit-scrollbar-thumb:hover {
|
||||||
background: var(--border-hover);
|
background: var(--border-hover);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -543,7 +553,8 @@ body {
|
||||||
}
|
}
|
||||||
|
|
||||||
/* Status Cards & Themes */
|
/* Status Cards & Themes */
|
||||||
#status-card {
|
#status-card,
|
||||||
|
#export-status-card {
|
||||||
border-radius: 8px;
|
border-radius: 8px;
|
||||||
padding: 1.25rem;
|
padding: 1.25rem;
|
||||||
border-left: 5px solid var(--text-dim);
|
border-left: 5px solid var(--text-dim);
|
||||||
|
|
@ -602,7 +613,8 @@ body {
|
||||||
100% { transform: scale(0.9); opacity: 0.6; }
|
100% { transform: scale(0.9); opacity: 0.6; }
|
||||||
}
|
}
|
||||||
|
|
||||||
#status-title {
|
#status-title,
|
||||||
|
#export-status-title {
|
||||||
font-size: 1rem;
|
font-size: 1rem;
|
||||||
font-weight: 700;
|
font-weight: 700;
|
||||||
letter-spacing: 0.05em;
|
letter-spacing: 0.05em;
|
||||||
|
|
@ -624,12 +636,33 @@ body {
|
||||||
color: var(--text-muted);
|
color: var(--text-muted);
|
||||||
}
|
}
|
||||||
|
|
||||||
#status-text {
|
#status-text,
|
||||||
|
#export-status-text {
|
||||||
font-size: 0.9rem;
|
font-size: 0.9rem;
|
||||||
color: var(--text-main);
|
color: var(--text-main);
|
||||||
margin-bottom: 1rem;
|
margin-bottom: 1rem;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#export-status-card.status-preparing,
|
||||||
|
#export-status-card.status-stopping { border-left-color: var(--warning); }
|
||||||
|
#export-status-card.status-exporting,
|
||||||
|
#export-status-card.status-succeeded { border-left-color: var(--success); }
|
||||||
|
#export-status-card.status-cancelled { border-left-color: var(--warning); }
|
||||||
|
#export-status-card.status-failed { border-left-color: var(--error); }
|
||||||
|
#export-status-card.status-preparing .status-dot,
|
||||||
|
#export-status-card.status-stopping .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); }
|
||||||
|
#export-status-card.status-exporting .status-dot,
|
||||||
|
#export-status-card.status-succeeded .status-dot { background-color: var(--success); box-shadow: 0 0 8px var(--success); }
|
||||||
|
#export-status-card.status-cancelled .status-dot { background-color: var(--warning); box-shadow: 0 0 8px var(--warning); }
|
||||||
|
#export-status-card.status-failed .status-dot { background-color: var(--error); box-shadow: 0 0 8px var(--error); }
|
||||||
|
#export-status-card.status-preparing #export-status-title,
|
||||||
|
#export-status-card.status-stopping #export-status-title { color: var(--warning); }
|
||||||
|
#export-status-card.status-exporting #export-status-title,
|
||||||
|
#export-status-card.status-succeeded #export-status-title { color: var(--success); }
|
||||||
|
#export-status-card.status-cancelled #export-status-title { color: var(--warning); }
|
||||||
|
#export-status-card.status-failed #export-status-title { color: var(--error); }
|
||||||
|
|
||||||
|
|
||||||
/* Progress bar inside status card */
|
/* Progress bar inside status card */
|
||||||
.progress-container {
|
.progress-container {
|
||||||
display: flex;
|
display: flex;
|
||||||
|
|
|
||||||
|
|
@ -191,25 +191,31 @@ class TrainingRunner:
|
||||||
from ultralytics import YOLO, settings
|
from ultralytics import YOLO, settings
|
||||||
from .ultralytics_trainers import trainer_for_task
|
from .ultralytics_trainers import trainer_for_task
|
||||||
|
|
||||||
settings.update({"mlflow": config.mlflow.enabled})
|
previous_mlflow_setting = settings["mlflow"]
|
||||||
|
setting_changed = previous_mlflow_setting is not config.mlflow.enabled
|
||||||
with mlflow_environment(config.mlflow):
|
if setting_changed:
|
||||||
model = YOLO(config.resolved_model, task=config.task)
|
settings.update({"mlflow": config.mlflow.enabled})
|
||||||
with self._state_lock:
|
try:
|
||||||
self._model = model
|
with mlflow_environment(config.mlflow):
|
||||||
|
model = YOLO(config.resolved_model, task=config.task)
|
||||||
model.add_callback("on_train_start", self._on_train_start(on_event))
|
|
||||||
model.add_callback("on_fit_epoch_end", self._on_epoch_end(on_event))
|
|
||||||
model.add_callback("on_train_end", self._on_train_end(on_event))
|
|
||||||
|
|
||||||
try:
|
|
||||||
model.train(trainer=trainer_for_task(config.task), **train_args)
|
|
||||||
trainer = getattr(model, "trainer", None)
|
|
||||||
save_dir = getattr(trainer, "save_dir", None)
|
|
||||||
return Path(save_dir) if save_dir else None
|
|
||||||
finally:
|
|
||||||
with self._state_lock:
|
with self._state_lock:
|
||||||
self._model = None
|
self._model = model
|
||||||
|
|
||||||
|
model.add_callback("on_train_start", self._on_train_start(on_event))
|
||||||
|
model.add_callback("on_fit_epoch_end", self._on_epoch_end(on_event))
|
||||||
|
model.add_callback("on_train_end", self._on_train_end(on_event))
|
||||||
|
|
||||||
|
try:
|
||||||
|
model.train(trainer=trainer_for_task(config.task), **train_args)
|
||||||
|
trainer = getattr(model, "trainer", None)
|
||||||
|
save_dir = getattr(trainer, "save_dir", None)
|
||||||
|
return Path(save_dir) if save_dir else None
|
||||||
|
finally:
|
||||||
|
with self._state_lock:
|
||||||
|
self._model = None
|
||||||
|
finally:
|
||||||
|
if setting_changed:
|
||||||
|
settings.update({"mlflow": previous_mlflow_setting})
|
||||||
|
|
||||||
def _on_train_start(self, on_event: EventHandler) -> Callable[[Any], None]:
|
def _on_train_start(self, on_event: EventHandler) -> Callable[[Any], None]:
|
||||||
def callback(trainer: Any) -> None:
|
def callback(trainer: Any) -> None:
|
||||||
|
|
|
||||||
|
|
@ -76,7 +76,25 @@ global.localStorage = {
|
||||||
removeItem(key) { storage.delete(key); }
|
removeItem(key) { storage.delete(key); }
|
||||||
};
|
};
|
||||||
global.confirm = () => true;
|
global.confirm = () => true;
|
||||||
global.window = {location: {protocol: 'http:', host: '127.0.0.1:8000'}};
|
global.window = {
|
||||||
|
location: {
|
||||||
|
protocol: 'http:',
|
||||||
|
host: '127.0.0.1:8000',
|
||||||
|
hostname: '127.0.0.1'
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let nextTimerId = 1;
|
||||||
|
const scheduledTimers = [];
|
||||||
|
global.setTimeout = (handler, delay) => {
|
||||||
|
const timer = {id: nextTimerId++, handler, delay, cancelled: false, fired: false};
|
||||||
|
scheduledTimers.push(timer);
|
||||||
|
return timer.id;
|
||||||
|
};
|
||||||
|
global.clearTimeout = id => {
|
||||||
|
const timer = scheduledTimers.find(item => item.id === id);
|
||||||
|
if (timer) timer.cancelled = true;
|
||||||
|
};
|
||||||
|
|
||||||
class FakeChart {
|
class FakeChart {
|
||||||
static instances = [];
|
static instances = [];
|
||||||
|
|
@ -106,8 +124,25 @@ const response = (ok, data) => ({
|
||||||
ok,
|
ok,
|
||||||
async json() { return data; }
|
async json() { return data; }
|
||||||
});
|
});
|
||||||
global.fetch = async url => {
|
const fetchCalls = [];
|
||||||
|
let holdTrainingStart = false;
|
||||||
|
let holdExportStart = false;
|
||||||
|
global.fetch = async (url, options = {}) => {
|
||||||
|
fetchCalls.push({url, options});
|
||||||
|
if (url === '/api/train/start' && holdTrainingStart) {
|
||||||
|
return new Promise(() => {});
|
||||||
|
}
|
||||||
|
if (url === '/api/export/start' && holdExportStart) {
|
||||||
|
return new Promise(() => {});
|
||||||
|
}
|
||||||
if (url === '/api/datasets' || url === '/api/models' || url === '/api/sessions') {
|
if (url === '/api/datasets' || url === '/api/models' || url === '/api/sessions') {
|
||||||
|
if (url === '/api/models') {
|
||||||
|
return response(true, [{
|
||||||
|
name: 'detect/experiment/best.pt',
|
||||||
|
path: '/workspace/runs/detect/experiment/weights/best.pt',
|
||||||
|
source: 'runs'
|
||||||
|
}]);
|
||||||
|
}
|
||||||
return response(true, []);
|
return response(true, []);
|
||||||
}
|
}
|
||||||
if (url === '/api/sessions/last_run') return response(false, {});
|
if (url === '/api/sessions/last_run') return response(false, {});
|
||||||
|
|
@ -119,7 +154,7 @@ global.fetch = async url => {
|
||||||
workers: 8,
|
workers: 8,
|
||||||
patience: 100,
|
patience: 100,
|
||||||
augmentation: {enabled: true, close_mosaic: 10},
|
augmentation: {enabled: true, close_mosaic: 10},
|
||||||
mlflow: {enabled: false},
|
mlflow: {enabled: false, tracking_uri: 'http://mlflow:5000'},
|
||||||
split: {enabled: false}
|
split: {enabled: false}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
@ -133,11 +168,30 @@ async function flushPromises() {
|
||||||
await new Promise(resolve => setImmediate(resolve));
|
await new Promise(resolve => setImmediate(resolve));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function descendants(node) {
|
||||||
|
return node.children.flatMap(child => [child, ...descendants(child)]);
|
||||||
|
}
|
||||||
|
|
||||||
|
function runTimer(timer) {
|
||||||
|
timer.fired = true;
|
||||||
|
timer.handler();
|
||||||
|
}
|
||||||
|
|
||||||
(async () => {
|
(async () => {
|
||||||
assert.equal(typeof domReady, 'function');
|
assert.equal(typeof domReady, 'function');
|
||||||
domReady();
|
domReady();
|
||||||
await flushPromises();
|
await flushPromises();
|
||||||
|
|
||||||
|
assert.equal(element('mlflow-header-link').href, 'http://127.0.0.1:5000/');
|
||||||
|
|
||||||
|
const localModelPath = '/workspace/runs/detect/experiment/weights/best.pt';
|
||||||
|
const localModelOption = descendants(element('model-select'))
|
||||||
|
.find(item => item.textContent === 'detect/experiment/best.pt');
|
||||||
|
assert.ok(localModelOption);
|
||||||
|
assert.equal(localModelOption.value, localModelPath);
|
||||||
|
|
||||||
|
element('model-select').value = localModelOption.value;
|
||||||
|
element('model-select').listeners.change();
|
||||||
element('workers').value = '0';
|
element('workers').value = '0';
|
||||||
element('patience').value = '0';
|
element('patience').value = '0';
|
||||||
element('close-mosaic').value = '0';
|
element('close-mosaic').value = '0';
|
||||||
|
|
@ -148,9 +202,28 @@ async function flushPromises() {
|
||||||
assert.equal(savedConfig.patience, 0);
|
assert.equal(savedConfig.patience, 0);
|
||||||
assert.equal(savedConfig.augmentation.close_mosaic, 0);
|
assert.equal(savedConfig.augmentation.close_mosaic, 0);
|
||||||
assert.equal(savedConfig.augmentation.auto_augment, 'none');
|
assert.equal(savedConfig.augmentation.auto_augment, 'none');
|
||||||
|
assert.equal(savedConfig.model, localModelPath);
|
||||||
|
|
||||||
|
element('model-select').value = '__custom__';
|
||||||
|
element('model-select').listeners.change();
|
||||||
|
element('model').value = '/workspace/models/custom-segment.pt';
|
||||||
|
element('task').value = 'segment';
|
||||||
|
element('task').listeners.change();
|
||||||
|
assert.equal(element('model-select').value, '__custom__');
|
||||||
|
assert.equal(element('model').value, '/workspace/models/custom-segment.pt');
|
||||||
|
|
||||||
assert.equal(FakeWebSocket.instances.length, 2);
|
assert.equal(FakeWebSocket.instances.length, 2);
|
||||||
const socket = FakeWebSocket.instances[0];
|
const socket = FakeWebSocket.instances[0];
|
||||||
|
const exportSocket = FakeWebSocket.instances[1];
|
||||||
|
const originalConsoleError = console.error;
|
||||||
|
console.error = () => {};
|
||||||
|
try {
|
||||||
|
assert.doesNotThrow(() => socket.onmessage({data: '{broken'}));
|
||||||
|
assert.doesNotThrow(() => exportSocket.onmessage({data: 'null'}));
|
||||||
|
} finally {
|
||||||
|
console.error = originalConsoleError;
|
||||||
|
}
|
||||||
|
|
||||||
socket.onmessage({
|
socket.onmessage({
|
||||||
data: JSON.stringify({
|
data: JSON.stringify({
|
||||||
type: 'init',
|
type: 'init',
|
||||||
|
|
@ -168,6 +241,77 @@ async function flushPromises() {
|
||||||
assert.deepEqual(chart.data.datasets.map(item => item.label), ['mAP50', 'loss']);
|
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[0].data, [0.5, null]);
|
||||||
assert.deepEqual(chart.data.datasets[1].data, [null, 0.2]);
|
assert.deepEqual(chart.data.datasets[1].data, [null, 0.2]);
|
||||||
|
|
||||||
|
socket.onclose();
|
||||||
|
socket.onclose();
|
||||||
|
let activeReconnectTimers = scheduledTimers.filter(
|
||||||
|
timer => timer.delay === 5000 && !timer.cancelled && !timer.fired
|
||||||
|
);
|
||||||
|
assert.equal(activeReconnectTimers.length, 1);
|
||||||
|
runTimer(activeReconnectTimers[0]);
|
||||||
|
assert.equal(FakeWebSocket.instances.length, 3);
|
||||||
|
assert.equal(FakeWebSocket.instances[2].url, 'ws://127.0.0.1:8000/api/ws');
|
||||||
|
|
||||||
|
exportSocket.onclose();
|
||||||
|
exportSocket.onclose();
|
||||||
|
activeReconnectTimers = scheduledTimers.filter(
|
||||||
|
timer => timer.delay === 5000 && !timer.cancelled && !timer.fired
|
||||||
|
);
|
||||||
|
assert.equal(activeReconnectTimers.length, 1);
|
||||||
|
runTimer(activeReconnectTimers[0]);
|
||||||
|
assert.equal(FakeWebSocket.instances.length, 4);
|
||||||
|
assert.equal(FakeWebSocket.instances[3].url, 'ws://127.0.0.1:8000/api/export/ws');
|
||||||
|
|
||||||
|
element('session-name').value = 'profile with spaces';
|
||||||
|
await element('session-save-btn').listeners.click();
|
||||||
|
assert.equal(
|
||||||
|
fetchCalls.filter(call => call.options.method === 'POST' && call.url.startsWith('/api/sessions/')).length,
|
||||||
|
0
|
||||||
|
);
|
||||||
|
|
||||||
|
element('export-model').value = localModelPath;
|
||||||
|
element('export-format').value = 'pb';
|
||||||
|
element('export-imgsz').value = '31';
|
||||||
|
element('export-batch').value = '1';
|
||||||
|
element('export-workspace').value = '4';
|
||||||
|
await element('export-start-btn').listeners.click();
|
||||||
|
assert.equal(fetchCalls.filter(call => call.url === '/api/export/start').length, 0);
|
||||||
|
|
||||||
|
element('export-imgsz').value = '640';
|
||||||
|
element('export-workspace').value = '1.5';
|
||||||
|
element('export-half').checked = true;
|
||||||
|
element('export-int8').checked = true;
|
||||||
|
await element('export-start-btn').listeners.click();
|
||||||
|
assert.equal(fetchCalls.filter(call => call.url === '/api/export/start').length, 0);
|
||||||
|
|
||||||
|
element('export-half').checked = false;
|
||||||
|
element('export-int8').checked = false;
|
||||||
|
holdExportStart = true;
|
||||||
|
element('export-start-btn').listeners.click();
|
||||||
|
element('export-start-btn').listeners.click();
|
||||||
|
const exportStartCalls = fetchCalls.filter(call => call.url === '/api/export/start');
|
||||||
|
assert.equal(exportStartCalls.length, 1);
|
||||||
|
assert.equal(element('export-start-btn').disabled, true);
|
||||||
|
assert.deepEqual(
|
||||||
|
JSON.parse(exportStartCalls[0].options.body),
|
||||||
|
{
|
||||||
|
model: localModelPath,
|
||||||
|
format: 'pb',
|
||||||
|
imgsz: 640,
|
||||||
|
half: false,
|
||||||
|
int8: false,
|
||||||
|
dynamic: false,
|
||||||
|
simplify: false,
|
||||||
|
batch: 1,
|
||||||
|
workspace: 1.5
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
holdTrainingStart = true;
|
||||||
|
element('start-btn').listeners.click();
|
||||||
|
element('start-btn').listeners.click();
|
||||||
|
assert.equal(fetchCalls.filter(call => call.url === '/api/train/start').length, 1);
|
||||||
|
assert.equal(element('start-btn').disabled, true);
|
||||||
})().catch(error => {
|
})().catch(error => {
|
||||||
console.error(error);
|
console.error(error);
|
||||||
process.exitCode = 1;
|
process.exitCode = 1;
|
||||||
|
|
|
||||||
|
|
@ -5,8 +5,16 @@ import json
|
||||||
import threading
|
import threading
|
||||||
|
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
import pytest
|
||||||
|
|
||||||
from yolo_webui.app import TrainingManager, app
|
import yolo_webui.app as app_module
|
||||||
|
from yolo_webui.app import (
|
||||||
|
ExportManager,
|
||||||
|
LiveState,
|
||||||
|
TrainingManager,
|
||||||
|
app,
|
||||||
|
write_json_atomic,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_get_config_defaults() -> None:
|
def test_get_config_defaults() -> None:
|
||||||
|
|
@ -94,6 +102,28 @@ def test_sessions_flow(monkeypatch, tmp_path) -> None:
|
||||||
assert response.status_code == 404
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
def test_invalid_session_config_is_not_saved(monkeypatch, tmp_path) -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
monkeypatch.setattr("yolo_webui.app.get_sessions_dir", lambda: tmp_path)
|
||||||
|
|
||||||
|
response = client.post("/api/sessions/broken", json={"unexpected": True})
|
||||||
|
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "датасета" in response.json()["detail"].lower()
|
||||||
|
assert not (tmp_path / "broken.json").exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_atomic_json_write_cleans_temp_file_after_serialization_error(tmp_path) -> None:
|
||||||
|
destination = tmp_path / "session.json"
|
||||||
|
destination.write_text('{"old": true}', encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(TypeError):
|
||||||
|
write_json_atomic(destination, {"invalid": object()})
|
||||||
|
|
||||||
|
assert destination.read_text(encoding="utf-8") == '{"old": true}'
|
||||||
|
assert list(tmp_path.iterdir()) == [destination]
|
||||||
|
|
||||||
|
|
||||||
def test_reserved_session_name_cannot_be_overwritten(monkeypatch, tmp_path) -> None:
|
def test_reserved_session_name_cannot_be_overwritten(monkeypatch, tmp_path) -> None:
|
||||||
client = TestClient(app)
|
client = TestClient(app)
|
||||||
monkeypatch.setattr("yolo_webui.app.get_sessions_dir", lambda: tmp_path)
|
monkeypatch.setattr("yolo_webui.app.get_sessions_dir", lambda: tmp_path)
|
||||||
|
|
@ -104,6 +134,107 @@ def test_reserved_session_name_cannot_be_overwritten(monkeypatch, tmp_path) -> N
|
||||||
assert "зарезервировано" in response.json()["detail"]
|
assert "зарезервировано" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_rejected_busy_start_does_not_overwrite_last_run(
|
||||||
|
monkeypatch,
|
||||||
|
tmp_path,
|
||||||
|
) -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
previous = {"dataset": "previous.yaml", "model": "previous.pt"}
|
||||||
|
last_run_path = tmp_path / "last_run.json"
|
||||||
|
last_run_path.write_text(json.dumps(previous), encoding="utf-8")
|
||||||
|
monkeypatch.setattr("yolo_webui.app.get_sessions_dir", lambda: tmp_path)
|
||||||
|
|
||||||
|
def reject_start(_config) -> None:
|
||||||
|
raise ValueError("Обучение уже выполняется.")
|
||||||
|
|
||||||
|
monkeypatch.setattr(app_module.manager, "start_training", reject_start)
|
||||||
|
|
||||||
|
response = client.post(
|
||||||
|
"/api/train/start",
|
||||||
|
json={"dataset": "coco8.yaml", "model": "yolo11n.pt"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert json.loads(last_run_path.read_text(encoding="utf-8")) == previous
|
||||||
|
|
||||||
|
|
||||||
|
def test_start_api_rejects_non_finite_numeric_value() -> None:
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
response = client.post(
|
||||||
|
"/api/train/start",
|
||||||
|
content='{"dataset":"coco8.yaml","model":"yolo11n.pt","epochs":NaN}',
|
||||||
|
headers={"content-type": "application/json"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "целым числом" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_models_uses_path_relative_to_runs(monkeypatch, tmp_path) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
direct_model = tmp_path / "runs" / "direct.pt"
|
||||||
|
nested_model = tmp_path / "runs" / "detect" / "train" / "weights" / "best.pt"
|
||||||
|
direct_model.parent.mkdir()
|
||||||
|
nested_model.parent.mkdir(parents=True)
|
||||||
|
direct_model.write_bytes(b"")
|
||||||
|
nested_model.write_bytes(b"")
|
||||||
|
|
||||||
|
response = TestClient(app).get("/api/models")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
run_names = {
|
||||||
|
item["name"] for item in response.json() if item["source"] == "runs"
|
||||||
|
}
|
||||||
|
assert run_names == {"direct.pt", "detect/train/weights/best.pt"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_start_rejects_invalid_config_synchronously() -> None:
|
||||||
|
response = TestClient(app).post(
|
||||||
|
"/api/export/start",
|
||||||
|
json={"model": "missing.pt", "format": "not-a-format"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 400
|
||||||
|
assert "Неподдерживаемый формат" in response.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_start_passes_canonical_validated_config(
|
||||||
|
monkeypatch,
|
||||||
|
tmp_path,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
model_path = tmp_path / "models" / "model.pt"
|
||||||
|
model_path.parent.mkdir()
|
||||||
|
model_path.write_bytes(b"checkpoint")
|
||||||
|
received: list[dict[str, object]] = []
|
||||||
|
monkeypatch.setattr(app_module.export_manager, "start_export", received.append)
|
||||||
|
|
||||||
|
response = TestClient(app).post(
|
||||||
|
"/api/export/start",
|
||||||
|
json={
|
||||||
|
"model": "model.pt",
|
||||||
|
"format": "ONNX",
|
||||||
|
"workspace": 2.5,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert received == [
|
||||||
|
{
|
||||||
|
"model": str(model_path),
|
||||||
|
"format": "onnx",
|
||||||
|
"imgsz": 640,
|
||||||
|
"half": False,
|
||||||
|
"int8": False,
|
||||||
|
"dynamic": False,
|
||||||
|
"simplify": False,
|
||||||
|
"batch": 1,
|
||||||
|
"workspace": 2.5,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_started_event_does_not_deadlock() -> None:
|
def test_started_event_does_not_deadlock() -> None:
|
||||||
training_manager = TrainingManager()
|
training_manager = TrainingManager()
|
||||||
training_manager.state.status = "preparing"
|
training_manager.state.status = "preparing"
|
||||||
|
|
@ -173,3 +304,55 @@ def test_process_result_distinguishes_success_cancellation_and_failure() -> None
|
||||||
failed_after_stop.state.stop_requested = True
|
failed_after_stop.state.stop_requested = True
|
||||||
failed_after_stop._finalize_process_result(1)
|
failed_after_stop._finalize_process_result(1)
|
||||||
assert failed_after_stop.state.status == "failed"
|
assert failed_after_stop.state.status == "failed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_live_state_snapshot_is_detached() -> None:
|
||||||
|
state = LiveState(
|
||||||
|
status="training",
|
||||||
|
logs=["first"],
|
||||||
|
metrics=[{"epoch": 1, "loss": 0.5}],
|
||||||
|
)
|
||||||
|
|
||||||
|
snapshot = state.snapshot()
|
||||||
|
state.logs.append("second")
|
||||||
|
state.metrics[0]["loss"] = 0.25
|
||||||
|
|
||||||
|
assert snapshot["logs"] == ["first"]
|
||||||
|
assert snapshot["metrics"] == [{"epoch": 1, "loss": 0.5}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_stop_does_not_relabel_completed_failure_as_cancelled() -> None:
|
||||||
|
class CompletedProcess:
|
||||||
|
returncode = 1
|
||||||
|
|
||||||
|
def poll(self) -> int:
|
||||||
|
return self.returncode
|
||||||
|
|
||||||
|
def terminate(self) -> None:
|
||||||
|
raise AssertionError("completed process must not be terminated")
|
||||||
|
|
||||||
|
export = ExportManager()
|
||||||
|
export.state.status = "exporting"
|
||||||
|
export._process = CompletedProcess() # type: ignore[assignment]
|
||||||
|
|
||||||
|
export.stop_export()
|
||||||
|
export._finalize_process_result(1)
|
||||||
|
|
||||||
|
assert export.state.stop_requested is False
|
||||||
|
assert export.state.status == "failed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_stopped_while_preparing_never_starts_process(monkeypatch) -> None:
|
||||||
|
export = ExportManager()
|
||||||
|
export.state.status = "preparing"
|
||||||
|
export.stop_export()
|
||||||
|
|
||||||
|
def unexpected_popen(*_args, **_kwargs):
|
||||||
|
raise AssertionError("subprocess must not start after cancellation")
|
||||||
|
|
||||||
|
monkeypatch.setattr(app_module.subprocess, "Popen", unexpected_popen)
|
||||||
|
|
||||||
|
export._run_subprocess({"model": "unused.pt", "format": "onnx"})
|
||||||
|
|
||||||
|
assert export.state.status == "cancelled"
|
||||||
|
assert export._process is None
|
||||||
|
|
|
||||||
|
|
@ -72,14 +72,102 @@ def test_augmentation_probabilities_are_validated(field: str) -> None:
|
||||||
augmentation.validate()
|
augmentation.validate()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("batch", [0, -2])
|
@pytest.mark.parametrize("batch", [0, -2, -0.5, 1.5])
|
||||||
def test_invalid_batch_is_rejected(batch: int) -> None:
|
def test_invalid_batch_is_rejected(batch: int | float) -> None:
|
||||||
config = TrainingConfig(dataset="dataset.yaml", model="model.pt", batch_size=batch)
|
config = TrainingConfig(dataset="dataset.yaml", model="model.pt", batch_size=batch)
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="Batch"):
|
with pytest.raises(ValueError, match="Batch"):
|
||||||
config.validate()
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("batch", [-1, 1, 16, 0.25])
|
||||||
|
def test_supported_batch_modes_are_accepted(batch: int | float) -> None:
|
||||||
|
TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
batch_size=batch,
|
||||||
|
).validate()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("field", "value", "message"),
|
||||||
|
[
|
||||||
|
("epochs", float("nan"), "целым числом"),
|
||||||
|
("epochs", True, "целым числом"),
|
||||||
|
("workers", 1.5, "целым числом"),
|
||||||
|
("patience", float("inf"), "целым числом"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_training_numeric_fields_reject_wrong_types_and_non_finite_values(
|
||||||
|
field: str,
|
||||||
|
value: object,
|
||||||
|
message: str,
|
||||||
|
) -> None:
|
||||||
|
values = {"dataset": "dataset.yaml", "model": "model.pt", field: value}
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match=message):
|
||||||
|
TrainingConfig(**values).validate()
|
||||||
|
|
||||||
|
|
||||||
|
def test_augmentation_rejects_non_finite_unbounded_values() -> None:
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
augmentation=AugmentationConfig(degrees=float("inf")),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="конечным числом"):
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
def test_augmentation_rejects_integer_too_large_for_float() -> None:
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
augmentation=AugmentationConfig(degrees=10**10_000),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="конечным числом"):
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"run_name",
|
||||||
|
["../escape", "/tmp/escape", r"..\\escape", "C:escape"],
|
||||||
|
)
|
||||||
|
def test_run_name_cannot_escape_project(run_name: str) -> None:
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
run_name=run_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Имя запуска"):
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
|
def test_nested_config_sections_must_be_objects() -> None:
|
||||||
|
with pytest.raises(ValueError, match="разделения датасета.*JSON-объектом"):
|
||||||
|
TrainingConfig.from_dict(
|
||||||
|
{
|
||||||
|
"dataset": "dataset.yaml",
|
||||||
|
"model": "model.pt",
|
||||||
|
"split": None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_classes_path_type_is_validated_before_path_operations() -> None:
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
split=DatasetSplitConfig(classes_path=123), # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Путь к файлу классов должен быть строкой"):
|
||||||
|
config.validate()
|
||||||
|
|
||||||
|
|
||||||
def test_mlflow_environment_is_restored(monkeypatch: pytest.MonkeyPatch) -> None:
|
def test_mlflow_environment_is_restored(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
monkeypatch.setenv("MLFLOW_TRACKING_URI", "previous")
|
monkeypatch.setenv("MLFLOW_TRACKING_URI", "previous")
|
||||||
config = MlflowConfig(
|
config = MlflowConfig(
|
||||||
|
|
@ -95,6 +183,14 @@ def test_mlflow_environment_is_restored(monkeypatch: pytest.MonkeyPatch) -> None
|
||||||
assert os.environ["MLFLOW_RUN"] == config.run_name
|
assert os.environ["MLFLOW_RUN"] == config.run_name
|
||||||
|
|
||||||
assert os.environ["MLFLOW_TRACKING_URI"] == "previous"
|
assert os.environ["MLFLOW_TRACKING_URI"] == "previous"
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlflow_default_tracking_uri_can_be_overridden_by_environment(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setenv("YOLO_WEBUI_MLFLOW_TRACKING_URI", "http://mlflow:5000")
|
||||||
|
|
||||||
|
assert MlflowConfig().tracking_uri == "http://mlflow:5000"
|
||||||
assert "MLFLOW_EXPERIMENT_NAME" not in os.environ
|
assert "MLFLOW_EXPERIMENT_NAME" not in os.environ
|
||||||
assert "MLFLOW_RUN" not in os.environ
|
assert "MLFLOW_RUN" not in os.environ
|
||||||
|
|
||||||
|
|
@ -113,7 +209,7 @@ def test_metrics_summary_skips_non_numeric_values() -> None:
|
||||||
assert summary == "mAP50=0.8123 · loss=0.1235 · precision=0.9568"
|
assert summary == "mAP50=0.8123 · loss=0.1235 · precision=0.9568"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("ratio", [0.05, 0.98])
|
@pytest.mark.parametrize("ratio", [0.05, 0.98, float("nan")])
|
||||||
def test_dataset_split_ratio_is_validated(ratio: float) -> None:
|
def test_dataset_split_ratio_is_validated(ratio: float) -> None:
|
||||||
config = DatasetSplitConfig(enabled=True, train_ratio=ratio)
|
config = DatasetSplitConfig(enabled=True, train_ratio=ratio)
|
||||||
with pytest.raises(ValueError, match="Доля обучающей выборки"):
|
with pytest.raises(ValueError, match="Доля обучающей выборки"):
|
||||||
|
|
|
||||||
45
tests/test_deployment.py
Normal file
45
tests/test_deployment.py
Normal file
|
|
@ -0,0 +1,45 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
|
||||||
|
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_compose_is_local_and_does_not_require_nvidia() -> None:
|
||||||
|
compose = yaml.safe_load(
|
||||||
|
(PROJECT_ROOT / "docker-compose.yml").read_text(encoding="utf-8")
|
||||||
|
)
|
||||||
|
webui = compose["services"]["webui"]
|
||||||
|
|
||||||
|
assert webui["ports"] == ["127.0.0.1:8000:8000"]
|
||||||
|
assert "deploy" not in webui
|
||||||
|
|
||||||
|
mlflow = compose["services"]["mlflow"]
|
||||||
|
assert mlflow["ports"] == ["127.0.0.1:5000:5000"]
|
||||||
|
assert mlflow["volumes"] == ["./mlflow:/mlflow"]
|
||||||
|
assert "server" in mlflow["command"]
|
||||||
|
assert webui["environment"]["YOLO_WEBUI_MLFLOW_TRACKING_URI"] == "http://mlflow:5000"
|
||||||
|
|
||||||
|
|
||||||
|
def test_docker_dependencies_are_exported_from_lock_file() -> None:
|
||||||
|
dockerfile = (PROJECT_ROOT / "Dockerfile").read_text(encoding="utf-8")
|
||||||
|
|
||||||
|
assert "COPY pyproject.toml uv.lock README.md ./" in dockerfile
|
||||||
|
assert "uv export --locked --no-dev" in dockerfile
|
||||||
|
assert "--group build" in dockerfile
|
||||||
|
assert "--prune ultralytics" in dockerfile
|
||||||
|
assert "uv pip install --system --no-deps --no-build-isolation ." in dockerfile
|
||||||
|
assert "ULTRALYTICS_SAFE_LOAD=1" in dockerfile
|
||||||
|
|
||||||
|
|
||||||
|
def test_docker_context_excludes_local_training_artifacts() -> None:
|
||||||
|
ignored = {
|
||||||
|
line.strip()
|
||||||
|
for line in (PROJECT_ROOT / ".dockerignore").read_text(encoding="utf-8").splitlines()
|
||||||
|
if line.strip() and not line.lstrip().startswith("#")
|
||||||
|
}
|
||||||
|
|
||||||
|
assert {".git", ".venv", "datasets", "models", "runs", "mlruns"} <= ignored
|
||||||
244
tests/test_export_runner.py
Normal file
244
tests/test_export_runner.py
Normal file
|
|
@ -0,0 +1,244 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from types import ModuleType
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from yolo_webui import export_runner
|
||||||
|
|
||||||
|
|
||||||
|
def _write_config(tmp_path: Path, data: Any) -> Path:
|
||||||
|
path = tmp_path / "export.json"
|
||||||
|
path.write_text(json.dumps(data), encoding="utf-8")
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def _write_model(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
models = tmp_path / "models"
|
||||||
|
models.mkdir()
|
||||||
|
model = models / "model.pt"
|
||||||
|
model.write_bytes(b"checkpoint")
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def _install_fake_ultralytics(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
result: Any,
|
||||||
|
) -> tuple[list[str], list[dict[str, Any]]]:
|
||||||
|
constructed: list[str] = []
|
||||||
|
export_calls: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
class FakeYOLO:
|
||||||
|
def __init__(self, model: str) -> None:
|
||||||
|
assert os.environ["ULTRALYTICS_SAFE_LOAD"] == "1"
|
||||||
|
constructed.append(model)
|
||||||
|
|
||||||
|
def export(self, **kwargs: Any) -> Any:
|
||||||
|
export_calls.append(kwargs)
|
||||||
|
return result
|
||||||
|
|
||||||
|
fake_ultralytics = ModuleType("ultralytics")
|
||||||
|
fake_ultralytics.YOLO = FakeYOLO # type: ignore[attr-defined]
|
||||||
|
monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics)
|
||||||
|
return constructed, export_calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_main_validates_config_and_reports_existing_export(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
model = _write_model(tmp_path, monkeypatch)
|
||||||
|
result = tmp_path / "models" / "model.onnx"
|
||||||
|
result.write_bytes(b"onnx")
|
||||||
|
constructed, export_calls = _install_fake_ultralytics(monkeypatch, result)
|
||||||
|
monkeypatch.setenv("ULTRALYTICS_SAFE_LOAD", "0")
|
||||||
|
config = _write_config(
|
||||||
|
tmp_path,
|
||||||
|
{
|
||||||
|
"model": model.name,
|
||||||
|
"format": "ONNX",
|
||||||
|
"imgsz": 320,
|
||||||
|
"half": False,
|
||||||
|
"int8": False,
|
||||||
|
"dynamic": True,
|
||||||
|
"simplify": True,
|
||||||
|
"batch": 2,
|
||||||
|
"workspace": 2.5,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return_code = export_runner.main([str(config)])
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 0
|
||||||
|
assert constructed == [str(model)]
|
||||||
|
assert export_calls == [
|
||||||
|
{
|
||||||
|
"format": "onnx",
|
||||||
|
"imgsz": 320,
|
||||||
|
"half": False,
|
||||||
|
"int8": False,
|
||||||
|
"dynamic": True,
|
||||||
|
"simplify": True,
|
||||||
|
"batch": 2,
|
||||||
|
"workspace": 2.5,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
assert f"__YOLO_WEBUI_RESULT__:{result}" in output.out
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"override",
|
||||||
|
[
|
||||||
|
{"format": "onnxx"},
|
||||||
|
{"format": ["onnx"]},
|
||||||
|
{"imgsz": "640"},
|
||||||
|
{"imgsz": True},
|
||||||
|
{"half": 1},
|
||||||
|
{"batch": 0},
|
||||||
|
{"workspace": float("inf")},
|
||||||
|
{"workspace": 10**400},
|
||||||
|
{"workspace": 64.5},
|
||||||
|
{"unexpected": "value"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_invalid_export_values_are_rejected_before_loading_ultralytics(
|
||||||
|
override: dict[str, Any],
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
model = _write_model(tmp_path, monkeypatch)
|
||||||
|
constructed, _ = _install_fake_ultralytics(monkeypatch, model)
|
||||||
|
config = _write_config(
|
||||||
|
tmp_path,
|
||||||
|
{
|
||||||
|
"model": str(model),
|
||||||
|
"format": "onnx",
|
||||||
|
**override,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return_code = export_runner.main([str(config)])
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert constructed == []
|
||||||
|
assert "Ошибка конфигурации экспорта:" in output.err
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_object_json_is_reported_as_invalid_config(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
|
||||||
|
return_code = export_runner.main([str(_write_config(tmp_path, ["model.pt"]))])
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert "JSON-объектом" in output.err
|
||||||
|
|
||||||
|
|
||||||
|
def test_model_must_be_a_regular_file_inside_an_allowed_root(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
(tmp_path / "models").mkdir()
|
||||||
|
outside = tmp_path / "outside.pt"
|
||||||
|
outside.write_bytes(b"checkpoint")
|
||||||
|
|
||||||
|
return_code = export_runner.main(
|
||||||
|
[str(_write_config(tmp_path, {"model": str(outside), "format": "onnx"}))]
|
||||||
|
)
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert "разрешённом каталоге" in output.err
|
||||||
|
|
||||||
|
|
||||||
|
def test_symlink_cannot_escape_an_allowed_model_root(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
models = tmp_path / "models"
|
||||||
|
models.mkdir()
|
||||||
|
outside = tmp_path / "outside.pt"
|
||||||
|
outside.write_bytes(b"checkpoint")
|
||||||
|
link = models / "linked.pt"
|
||||||
|
link.symlink_to(outside)
|
||||||
|
|
||||||
|
return_code = export_runner.main(
|
||||||
|
[str(_write_config(tmp_path, {"model": str(link), "format": "onnx"}))]
|
||||||
|
)
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert "разрешённом каталоге" in output.err
|
||||||
|
|
||||||
|
|
||||||
|
def test_half_and_int8_cannot_be_enabled_together(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
model = _write_model(tmp_path, monkeypatch)
|
||||||
|
|
||||||
|
return_code = export_runner.main(
|
||||||
|
[
|
||||||
|
str(
|
||||||
|
_write_config(
|
||||||
|
tmp_path,
|
||||||
|
{
|
||||||
|
"model": str(model),
|
||||||
|
"format": "onnx",
|
||||||
|
"half": True,
|
||||||
|
"int8": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert "нельзя включать одновременно" in output.err
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("result", [None, [], "/missing/export.onnx"])
|
||||||
|
def test_missing_export_artifact_is_a_failure(
|
||||||
|
result: Any,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
capsys: pytest.CaptureFixture[str],
|
||||||
|
) -> None:
|
||||||
|
model = _write_model(tmp_path, monkeypatch)
|
||||||
|
_install_fake_ultralytics(monkeypatch, result)
|
||||||
|
|
||||||
|
return_code = export_runner.main(
|
||||||
|
[
|
||||||
|
str(
|
||||||
|
_write_config(
|
||||||
|
tmp_path,
|
||||||
|
{"model": str(model), "format": "onnx"},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
output = capsys.readouterr()
|
||||||
|
assert return_code == 1
|
||||||
|
assert "__YOLO_WEBUI_RESULT__:" not in output.out
|
||||||
|
assert "Ошибка при экспорте модели:" in output.out
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from html.parser import HTMLParser
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
|
|
@ -9,6 +10,37 @@ import pytest
|
||||||
|
|
||||||
NODE = shutil.which("node")
|
NODE = shutil.which("node")
|
||||||
APP_JS = Path("src/yolo_webui/static/app.js")
|
APP_JS = Path("src/yolo_webui/static/app.js")
|
||||||
|
INDEX_HTML = Path("src/yolo_webui/static/index.html")
|
||||||
|
STYLE_CSS = Path("src/yolo_webui/static/style.css")
|
||||||
|
|
||||||
|
|
||||||
|
class _MarkupCollector(HTMLParser):
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.by_id: dict[str, dict[str, str | None]] = {}
|
||||||
|
self.options: dict[str, str] = {}
|
||||||
|
self._option_value: str | None = None
|
||||||
|
|
||||||
|
def handle_starttag(
|
||||||
|
self,
|
||||||
|
tag: str,
|
||||||
|
attrs: list[tuple[str, str | None]],
|
||||||
|
) -> None:
|
||||||
|
attributes = dict(attrs)
|
||||||
|
if element_id := attributes.get("id"):
|
||||||
|
self.by_id[element_id] = attributes
|
||||||
|
if tag == "option":
|
||||||
|
self._option_value = attributes.get("value")
|
||||||
|
|
||||||
|
def handle_data(self, data: str) -> None:
|
||||||
|
if self._option_value is not None:
|
||||||
|
self.options[self._option_value] = (
|
||||||
|
self.options.get(self._option_value, "") + data
|
||||||
|
).strip()
|
||||||
|
|
||||||
|
def handle_endtag(self, tag: str) -> None:
|
||||||
|
if tag == "option":
|
||||||
|
self._option_value = None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
||||||
|
|
@ -17,5 +49,30 @@ def test_frontend_javascript_syntax() -> None:
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
@pytest.mark.skipif(NODE is None, reason="Node.js is required for frontend checks")
|
||||||
def test_frontend_zero_values_and_dynamic_chart_series() -> None:
|
def test_frontend_interactions_smoke() -> None:
|
||||||
subprocess.run([NODE, "tests/frontend_smoke.js"], check=True)
|
subprocess.run([NODE, "tests/frontend_smoke.js"], check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_export_controls_have_correct_formats_constraints_and_styles() -> None:
|
||||||
|
parser = _MarkupCollector()
|
||||||
|
parser.feed(INDEX_HTML.read_text(encoding="utf-8"))
|
||||||
|
|
||||||
|
assert parser.options["saved_model"].startswith("TensorFlow SavedModel")
|
||||||
|
assert parser.options["pb"] == "TensorFlow GraphDef (.pb)"
|
||||||
|
assert "triton" not in parser.options
|
||||||
|
assert parser.by_id["export-imgsz"]["min"] == "32"
|
||||||
|
assert parser.by_id["export-imgsz"]["max"] == "8192"
|
||||||
|
assert parser.by_id["export-imgsz"]["step"] == "1"
|
||||||
|
assert parser.by_id["export-batch"]["min"] == "1"
|
||||||
|
assert parser.by_id["export-batch"]["max"] == "1024"
|
||||||
|
assert parser.by_id["export-batch"]["step"] == "1"
|
||||||
|
assert parser.by_id["export-workspace"]["min"] == "1"
|
||||||
|
assert parser.by_id["export-workspace"]["max"] == "64"
|
||||||
|
assert parser.by_id["export-workspace"]["step"] == "0.5"
|
||||||
|
|
||||||
|
css = STYLE_CSS.read_text(encoding="utf-8")
|
||||||
|
assert "#export-config-pane" in css
|
||||||
|
assert "#export-run-pane" in css
|
||||||
|
assert "#export-log-container::-webkit-scrollbar" in css
|
||||||
|
assert "#export-status-card.status-exporting" in css
|
||||||
|
assert "#export-status-card.status-failed" in css
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,14 @@ def test_read_classes_custom_yaml(
|
||||||
assert read_classes(tmp_path, str(custom_file)) == expected
|
assert read_classes(tmp_path, str(custom_file)) == expected
|
||||||
|
|
||||||
|
|
||||||
|
def test_read_classes_rejects_fractional_yaml_class_id(tmp_path: Path) -> None:
|
||||||
|
custom_file = tmp_path / "custom_classes.yaml"
|
||||||
|
custom_file.write_text("names:\n 0.5: person\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Некорректный ID класса"):
|
||||||
|
read_classes(tmp_path, str(custom_file))
|
||||||
|
|
||||||
|
|
||||||
def test_missing_custom_classes_path_does_not_fall_back(tmp_path: Path) -> None:
|
def test_missing_custom_classes_path_does_not_fall_back(tmp_path: Path) -> None:
|
||||||
(tmp_path / "classes.txt").write_text("fallback\n", encoding="utf-8")
|
(tmp_path / "classes.txt").write_text("fallback\n", encoding="utf-8")
|
||||||
|
|
||||||
|
|
@ -145,6 +153,59 @@ def test_split_dataset_rejects_single_image_without_writing_output(
|
||||||
assert not (tmp_path / ".yolo-webui").exists()
|
assert not (tmp_path / ".yolo-webui").exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_split_dataset_rejects_label_id_missing_from_classes(
|
||||||
|
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(
|
||||||
|
"1 0.5 0.5 0.2 0.2\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
(tmp_path / "classes.txt").write_text("only-class\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="указан класс 1"):
|
||||||
|
split_dataset(str(tmp_path), 0.5, "")
|
||||||
|
|
||||||
|
assert not (tmp_path / ".yolo-webui").exists()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("invalid_id", ["-1", "0.5", "nan", "class-a"])
|
||||||
|
def test_split_dataset_rejects_invalid_label_class_id(
|
||||||
|
tmp_path: Path,
|
||||||
|
invalid_id: str,
|
||||||
|
) -> 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(
|
||||||
|
f"{invalid_id} 0.5 0.5 0.2 0.2\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
(tmp_path / "classes.txt").write_text("item\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Некорректный ID класса"):
|
||||||
|
split_dataset(str(tmp_path), 0.5, "")
|
||||||
|
|
||||||
|
assert not (tmp_path / ".yolo-webui").exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_split_dataset_rejects_non_finite_ratio_before_writing(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
with pytest.raises(ValueError, match="конечным числом"):
|
||||||
|
split_dataset(str(tmp_path), float("nan"), "")
|
||||||
|
|
||||||
|
assert not (tmp_path / ".yolo-webui").exists()
|
||||||
|
|
||||||
|
|
||||||
def test_split_dataset_finds_nested_images_and_labels(tmp_path: Path) -> None:
|
def test_split_dataset_finds_nested_images_and_labels(tmp_path: Path) -> None:
|
||||||
images_dir = tmp_path / "images" / "day"
|
images_dir = tmp_path / "images" / "day"
|
||||||
labels_dir = tmp_path / "labels" / "day"
|
labels_dir = tmp_path / "labels" / "day"
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,8 @@ from pathlib import Path
|
||||||
from types import ModuleType, SimpleNamespace
|
from types import ModuleType, SimpleNamespace
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from yolo_webui.config import MlflowConfig, TrainingConfig
|
from yolo_webui.config import MlflowConfig, TrainingConfig
|
||||||
from yolo_webui.trainer import TrainingEvent, TrainingRunner
|
from yolo_webui.trainer import TrainingEvent, TrainingRunner
|
||||||
|
|
||||||
|
|
@ -37,8 +39,15 @@ def test_runner_wires_yolo_callbacks_and_returns_output(
|
||||||
train_arguments: list[dict[str, Any]] = []
|
train_arguments: list[dict[str, Any]] = []
|
||||||
|
|
||||||
class FakeSettings:
|
class FakeSettings:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.values = {"mlflow": True}
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> bool:
|
||||||
|
return self.values[key]
|
||||||
|
|
||||||
def update(self, values: dict[str, bool]) -> None:
|
def update(self, values: dict[str, bool]) -> None:
|
||||||
settings_updates.append(values)
|
settings_updates.append(values)
|
||||||
|
self.values.update(values)
|
||||||
|
|
||||||
class FakeYOLO:
|
class FakeYOLO:
|
||||||
def __init__(self, model: str, task: str) -> None:
|
def __init__(self, model: str, task: str) -> None:
|
||||||
|
|
@ -97,7 +106,7 @@ def test_runner_wires_yolo_callbacks_and_returns_output(
|
||||||
|
|
||||||
assert output == tmp_path / "run"
|
assert output == tmp_path / "run"
|
||||||
assert constructed == [("models/model.pt", "pose")]
|
assert constructed == [("models/model.pt", "pose")]
|
||||||
assert settings_updates == [{"mlflow": False}]
|
assert settings_updates == [{"mlflow": False}, {"mlflow": True}]
|
||||||
assert train_arguments[0]["data"] == "dataset.yaml"
|
assert train_arguments[0]["data"] == "dataset.yaml"
|
||||||
assert train_arguments[0]["verbose"] is True
|
assert train_arguments[0]["verbose"] is True
|
||||||
assert train_arguments[0]["trainer"] is FakeTrainer
|
assert train_arguments[0]["trainer"] is FakeTrainer
|
||||||
|
|
@ -106,6 +115,48 @@ def test_runner_wires_yolo_callbacks_and_returns_output(
|
||||||
assert "mAP50=0.6" in events[3].message
|
assert "mAP50=0.6" in events[3].message
|
||||||
|
|
||||||
|
|
||||||
|
def test_ultralytics_mlflow_setting_is_restored_when_model_loading_fails(
|
||||||
|
monkeypatch: Any,
|
||||||
|
) -> None:
|
||||||
|
settings_updates: list[bool] = []
|
||||||
|
|
||||||
|
class FakeSettings:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.mlflow = True
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> bool:
|
||||||
|
assert key == "mlflow"
|
||||||
|
return self.mlflow
|
||||||
|
|
||||||
|
def update(self, values: dict[str, bool]) -> None:
|
||||||
|
self.mlflow = values["mlflow"]
|
||||||
|
settings_updates.append(self.mlflow)
|
||||||
|
|
||||||
|
class FailingYOLO:
|
||||||
|
def __init__(self, _model: str, task: str) -> None:
|
||||||
|
assert task == "detect"
|
||||||
|
raise RuntimeError("model load failed")
|
||||||
|
|
||||||
|
fake_ultralytics = ModuleType("ultralytics")
|
||||||
|
fake_ultralytics.YOLO = FailingYOLO # type: ignore[attr-defined]
|
||||||
|
fake_ultralytics.settings = FakeSettings() # type: ignore[attr-defined]
|
||||||
|
fake_trainers = ModuleType("yolo_webui.ultralytics_trainers")
|
||||||
|
fake_trainers.trainer_for_task = lambda _task: object # type: ignore[attr-defined]
|
||||||
|
monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics)
|
||||||
|
monkeypatch.setitem(sys.modules, "yolo_webui.ultralytics_trainers", fake_trainers)
|
||||||
|
config = TrainingConfig(
|
||||||
|
dataset="dataset.yaml",
|
||||||
|
model="model.pt",
|
||||||
|
mlflow=MlflowConfig(enabled=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="model load failed"):
|
||||||
|
TrainingRunner().train(config, lambda _event: None)
|
||||||
|
|
||||||
|
assert settings_updates == [False, True]
|
||||||
|
assert fake_ultralytics.settings["mlflow"] is True # type: ignore[index]
|
||||||
|
|
||||||
|
|
||||||
def test_final_checkpoint_validation_is_not_reported_as_an_extra_epoch() -> None:
|
def test_final_checkpoint_validation_is_not_reported_as_an_extra_epoch() -> None:
|
||||||
events: list[TrainingEvent] = []
|
events: list[TrainingEvent] = []
|
||||||
trainer = SimpleNamespace(
|
trainer = SimpleNamespace(
|
||||||
|
|
|
||||||
42
uv.lock
generated
42
uv.lock
generated
|
|
@ -834,7 +834,7 @@ name = "exceptiongroup"
|
||||||
version = "1.3.1"
|
version = "1.3.1"
|
||||||
source = { registry = "https://pypi.org/simple" }
|
source = { registry = "https://pypi.org/simple" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
|
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
|
||||||
]
|
]
|
||||||
sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" }
|
sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" }
|
||||||
wheels = [
|
wheels = [
|
||||||
|
|
@ -1232,7 +1232,7 @@ name = "gunicorn"
|
||||||
version = "26.0.0"
|
version = "26.0.0"
|
||||||
source = { registry = "https://pypi.org/simple" }
|
source = { registry = "https://pypi.org/simple" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "packaging" },
|
{ name = "packaging", marker = "sys_platform != 'win32'" },
|
||||||
]
|
]
|
||||||
sdist = { url = "https://files.pythonhosted.org/packages/6d/b7/a4a3f632f823e432ce6bc65f62961b7980c898c77f075a2f7118cb3846fe/gunicorn-26.0.0.tar.gz", hash = "sha256:ca9346f85e3a4aeeb64d491045c16b9a35647abd37ea15efe53080eb8b090baf", size = 727286, upload-time = "2026-05-05T06:38:25.529Z" }
|
sdist = { url = "https://files.pythonhosted.org/packages/6d/b7/a4a3f632f823e432ce6bc65f62961b7980c898c77f075a2f7118cb3846fe/gunicorn-26.0.0.tar.gz", hash = "sha256:ca9346f85e3a4aeeb64d491045c16b9a35647abd37ea15efe53080eb8b090baf", size = 727286, upload-time = "2026-05-05T06:38:25.529Z" }
|
||||||
wheels = [
|
wheels = [
|
||||||
|
|
@ -1248,6 +1248,22 @@ wheels = [
|
||||||
{ url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" },
|
{ url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "hatchling"
|
||||||
|
version = "1.27.0"
|
||||||
|
source = { registry = "https://pypi.org/simple" }
|
||||||
|
dependencies = [
|
||||||
|
{ name = "packaging" },
|
||||||
|
{ name = "pathspec" },
|
||||||
|
{ name = "pluggy" },
|
||||||
|
{ name = "tomli", marker = "python_full_version < '3.11'" },
|
||||||
|
{ name = "trove-classifiers" },
|
||||||
|
]
|
||||||
|
sdist = { url = "https://files.pythonhosted.org/packages/8f/8a/cc1debe3514da292094f1c3a700e4ca25442489731ef7c0814358816bb03/hatchling-1.27.0.tar.gz", hash = "sha256:971c296d9819abb3811112fc52c7a9751c8d381898f36533bb16f9791e941fd6", size = 54983, upload-time = "2024-12-15T17:08:11.894Z" }
|
||||||
|
wheels = [
|
||||||
|
{ url = "https://files.pythonhosted.org/packages/08/e7/ae38d7a6dfba0533684e0b2136817d667588ae3ec984c1a4e5df5eb88482/hatchling-1.27.0-py3-none-any.whl", hash = "sha256:d3a2f3567c4f926ea39849cdf924c7e99e6686c9c8e288ae1037c8fa2a5d937b", size = 75794, upload-time = "2024-12-15T17:08:10.364Z" },
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "httpcore"
|
name = "httpcore"
|
||||||
version = "1.0.9"
|
version = "1.0.9"
|
||||||
|
|
@ -2500,6 +2516,15 @@ wheels = [
|
||||||
{ url = "https://files.pythonhosted.org/packages/70/44/5191d2e4026f86a2a109053e194d3ba7a31a2d10a9c2348368c63ed4e85a/pandas-2.3.3-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:3869faf4bd07b3b66a9f462417d0ca3a9df29a9f6abd5d0d0dbab15dac7abe87", size = 13202175, upload-time = "2025-09-29T23:31:59.173Z" },
|
{ url = "https://files.pythonhosted.org/packages/70/44/5191d2e4026f86a2a109053e194d3ba7a31a2d10a9c2348368c63ed4e85a/pandas-2.3.3-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:3869faf4bd07b3b66a9f462417d0ca3a9df29a9f6abd5d0d0dbab15dac7abe87", size = 13202175, upload-time = "2025-09-29T23:31:59.173Z" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "pathspec"
|
||||||
|
version = "1.1.1"
|
||||||
|
source = { registry = "https://pypi.org/simple" }
|
||||||
|
sdist = { url = "https://files.pythonhosted.org/packages/5a/82/42f767fc1c1143d6fd36efb827202a2d997a375e160a71eb2888a925aac1/pathspec-1.1.1.tar.gz", hash = "sha256:17db5ecd524104a120e173814c90367a96a98d07c45b2e10c2f3919fff91bf5a", size = 135180, upload-time = "2026-04-27T01:46:08.907Z" }
|
||||||
|
wheels = [
|
||||||
|
{ url = "https://files.pythonhosted.org/packages/f1/d9/7fb5aa316bc299258e68c73ba3bddbc499654a07f151cba08f6153988714/pathspec-1.1.1-py3-none-any.whl", hash = "sha256:a00ce642f577bf7f473932318056212bc4f8bfdf53128c78bbd5af0b9b20b189", size = 57328, upload-time = "2026-04-27T01:46:07.06Z" },
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "pillow"
|
name = "pillow"
|
||||||
version = "12.3.0"
|
version = "12.3.0"
|
||||||
|
|
@ -3802,6 +3827,15 @@ wheels = [
|
||||||
{ url = "https://files.pythonhosted.org/packages/f0/ac/229b7d4589d2e5937310e72c6d46e89599d16a4a12b479ffa1499fee8eb8/triton-3.7.1-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:10ba85fa2cca4a2fbdeb36bf1cb082f2c252bda55bf9fccd74f65ec5bc647e68", size = 197824404, upload-time = "2026-06-17T19:53:42.772Z" },
|
{ url = "https://files.pythonhosted.org/packages/f0/ac/229b7d4589d2e5937310e72c6d46e89599d16a4a12b479ffa1499fee8eb8/triton-3.7.1-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:10ba85fa2cca4a2fbdeb36bf1cb082f2c252bda55bf9fccd74f65ec5bc647e68", size = 197824404, upload-time = "2026-06-17T19:53:42.772Z" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "trove-classifiers"
|
||||||
|
version = "2026.6.1.19"
|
||||||
|
source = { registry = "https://pypi.org/simple" }
|
||||||
|
sdist = { url = "https://files.pythonhosted.org/packages/c2/e3/7ca82ee24c82d344584abd5b8637b3bd056f2900226e8d82fc22f1184b92/trove_classifiers-2026.6.1.19.tar.gz", hash = "sha256:c5132b4b61a829d11cfbd2d72e97f20a45ed6edb95e45c5efdeb5e00836b2745", size = 17059, upload-time = "2026-06-01T19:41:34.649Z" }
|
||||||
|
wheels = [
|
||||||
|
{ url = "https://files.pythonhosted.org/packages/7c/a4/81502f486f01db95bc8320646a8a12511f5e556cb63d5e224d91816605c4/trove_classifiers-2026.6.1.19-py3-none-any.whl", hash = "sha256:ab4c4ec93cc4a4e7815fa759906e05e6bb3f2fbd92ea0f897288c6a43efd15b3", size = 14211, upload-time = "2026-06-01T19:41:33.434Z" },
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "typing-extensions"
|
name = "typing-extensions"
|
||||||
version = "4.16.0"
|
version = "4.16.0"
|
||||||
|
|
@ -4171,6 +4205,9 @@ dependencies = [
|
||||||
]
|
]
|
||||||
|
|
||||||
[package.dev-dependencies]
|
[package.dev-dependencies]
|
||||||
|
build = [
|
||||||
|
{ name = "hatchling" },
|
||||||
|
]
|
||||||
dev = [
|
dev = [
|
||||||
{ name = "httpx" },
|
{ name = "httpx" },
|
||||||
{ name = "pytest" },
|
{ name = "pytest" },
|
||||||
|
|
@ -4186,6 +4223,7 @@ requires-dist = [
|
||||||
]
|
]
|
||||||
|
|
||||||
[package.metadata.requires-dev]
|
[package.metadata.requires-dev]
|
||||||
|
build = [{ name = "hatchling", specifier = "==1.27.0" }]
|
||||||
dev = [
|
dev = [
|
||||||
{ name = "httpx" },
|
{ name = "httpx" },
|
||||||
{ name = "pytest", specifier = ">=8.3" },
|
{ name = "pytest", specifier = ">=8.3" },
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue