pagina unica
This commit is contained in:
@@ -0,0 +1,43 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user