44 lines
1.4 KiB
Python
44 lines
1.4 KiB
Python
"""Tests for mock and Azure ML stub training backends."""
|
|
|
|
from uuid import uuid4
|
|
|
|
import pytest
|
|
|
|
from forja_core.config import Settings
|
|
from forja_core.domain.training import TrainingHyperparameters, TrainingJobSpec
|
|
from forja_core.training.azure_ml import AzureMLTrainingBackend
|
|
from forja_core.training.mock import MockTrainingBackend
|
|
|
|
|
|
@pytest.fixture
|
|
def spec() -> TrainingJobSpec:
|
|
return TrainingJobSpec(
|
|
dataset_name="incident_sft",
|
|
dataset_version="v1",
|
|
model_name="gpt4o_lora_base",
|
|
model_version="v1",
|
|
hyperparameters=TrainingHyperparameters(epochs=2),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mock_submit_and_status_succeeds(spec: TrainingJobSpec) -> None:
|
|
backend = MockTrainingBackend()
|
|
run_id = uuid4()
|
|
job_id = await backend.submit(spec, run_id)
|
|
assert job_id == f"mock:{run_id}"
|
|
status = await backend.status(job_id)
|
|
assert status.status == "succeeded"
|
|
assert len(status.artifacts) == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_azure_stub_submit_and_status_succeeds(spec: TrainingJobSpec) -> None:
|
|
backend = AzureMLTrainingBackend(settings=Settings())
|
|
run_id = uuid4()
|
|
job_id = await backend.submit(spec, run_id)
|
|
assert job_id.startswith("azureml-stub:")
|
|
status = await backend.status(job_id)
|
|
assert status.status == "succeeded"
|
|
assert "azureml://" in status.artifacts[0].uri
|