train_utility/tests/test_config.py
2026-08-05 09:32:44 +04:00

273 lines
8.4 KiB
Python
Executable file

from __future__ import annotations
import os
import pytest
from yolo_webui.config import AugmentationConfig, DatasetSplitConfig, MlflowConfig, TrainingConfig
from yolo_webui.trainer import TrainingRunner, mlflow_environment
def test_train_kwargs_omit_optional_empty_values() -> None:
config = TrainingConfig(
dataset=" dataset.yaml ",
model=" model.pt ",
augmentation=AugmentationConfig(enabled=False),
)
from pathlib import Path
assert config.train_kwargs() == {
"data": "dataset.yaml",
"epochs": 100,
"imgsz": 640,
"batch": 16,
"workers": 8,
"patience": 100,
"project": str(Path("runs/train").resolve()),
"verbose": True,
}
def test_augmentation_kwargs_are_passed_to_ultralytics() -> None:
config = TrainingConfig(
dataset="dataset.yaml",
model="model.pt",
augmentation=AugmentationConfig(mosaic=0.5, mixup=0.2, close_mosaic=5),
)
kwargs = config.train_kwargs()
assert kwargs["mosaic"] == 0.5
assert kwargs["mixup"] == 0.2
assert kwargs["close_mosaic"] == 5
assert kwargs["auto_augment"] == "randaugment"
def test_auto_augment_none_policy() -> None:
config = TrainingConfig(
dataset="dataset.yaml",
model="model.pt",
augmentation=AugmentationConfig(auto_augment="none"),
)
kwargs = config.train_kwargs()
assert kwargs["auto_augment"] is None
def test_invalid_auto_augment_policy() -> None:
config = TrainingConfig(
dataset="dataset.yaml",
model="model.pt",
augmentation=AugmentationConfig(auto_augment="invalid_policy"), # type: ignore[arg-type]
)
with pytest.raises(ValueError, match="Неизвестная политика AutoAugment"):
config.validate()
@pytest.mark.parametrize("field", ["mosaic", "fliplr", "erasing", "perspective"])
def test_augmentation_probabilities_are_validated(field: str) -> None:
augmentation = AugmentationConfig(**{field: 1.1})
with pytest.raises(ValueError, match="от 0 до 1"):
augmentation.validate()
@pytest.mark.parametrize("batch", [0, -2, -0.5, 1.5])
def test_invalid_batch_is_rejected(batch: int | float) -> None:
config = TrainingConfig(dataset="dataset.yaml", model="model.pt", batch_size=batch)
with pytest.raises(ValueError, match="Batch"):
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:
monkeypatch.setenv("MLFLOW_TRACKING_URI", "previous")
config = MlflowConfig(
enabled=True,
tracking_uri="http://127.0.0.1:5000",
experiment_name="test-experiment",
run_name="test-run",
)
with mlflow_environment(config):
assert os.environ["MLFLOW_TRACKING_URI"] == config.tracking_uri
assert os.environ["MLFLOW_EXPERIMENT_NAME"] == config.experiment_name
assert os.environ["MLFLOW_RUN"] == config.run_name
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_RUN" not in os.environ
def test_metrics_summary_skips_non_numeric_values() -> None:
summary = TrainingRunner._metrics_summary(
{
"metrics/mAP50": 0.81234,
"label": "invalid",
"val/loss": 0.12345,
"metrics/precision": 0.95678,
"another": 12.34,
}
)
assert summary == "mAP50=0.8123 · loss=0.1235 · precision=0.9568"
@pytest.mark.parametrize("ratio", [0.05, 0.98, float("nan")])
def test_dataset_split_ratio_is_validated(ratio: float) -> None:
config = DatasetSplitConfig(enabled=True, train_ratio=ratio)
with pytest.raises(ValueError, match="Доля обучающей выборки"):
config.validate()
def test_classification_rejects_detection_style_auto_split() -> None:
config = TrainingConfig(
dataset="classification-dataset",
model="model.pt",
task="classify",
split=DatasetSplitConfig(enabled=True),
)
with pytest.raises(ValueError, match="classify"):
config.validate()
@pytest.mark.parametrize(
("field", "value"),
[
("dataset", "https://example.invalid/dataset.yaml"),
("model", "https://example.invalid/model.pt"),
("project", "https://example.invalid/results"),
],
)
def test_training_rejects_remote_references(field: str, value: str) -> None:
values = {
"dataset": "dataset.yaml",
"model": "model.pt",
"project": "runs/train",
field: value,
}
with pytest.raises(ValueError, match="не может быть URL"):
TrainingConfig(**values).validate()
def test_model_path_must_stay_in_allowed_roots(tmp_path: Path, monkeypatch) -> None:
workspace = tmp_path / "workspace"
workspace.mkdir()
external_model = tmp_path / "external" / "model.pt"
monkeypatch.chdir(workspace)
config = TrainingConfig(dataset="dataset.yaml", model=str(external_model))
with pytest.raises(ValueError, match="разрешённом каталоге"):
config.validate()
monkeypatch.setenv("YOLO_WEBUI_MODEL_ROOTS", str(external_model.parent))
config.validate()
def test_existing_bare_dataset_cannot_bypass_allowed_roots(
tmp_path: Path, monkeypatch
) -> None:
monkeypatch.chdir(tmp_path)
(tmp_path / "private.yaml").write_text("secret: value\n", encoding="utf-8")
with pytest.raises(ValueError, match="вне разрешённого каталога"):
TrainingConfig(dataset="private.yaml", model="model.pt").validate()