Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions src/viam/app/ml_training_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,7 @@ async def submit_custom_training_job(
model_name: str,
model_version: str,
container_id: str = "",
refresh_dataset_cache: bool = False,
) -> str:
"""Submit a custom training job.

Expand All @@ -158,6 +159,8 @@ async def submit_custom_training_job(
model_version (str): the model version.
container_id (str): the ID of the custom training container to run the job in. If unspecified, the training script's
default container is used.
refresh_dataset_cache (bool): whether to export the dataset fresh instead of reusing a cached export, and replace the cached
export with it. Defaults to False.

Returns:
str: the ID of the training job.
Expand All @@ -173,6 +176,7 @@ async def submit_custom_training_job(
model_name=model_name,
model_version=model_version,
container_id=container_id,
refresh_dataset_cache=refresh_dataset_cache,
)
response: SubmitCustomTrainingJobResponse = await self._ml_training_client.SubmitCustomTrainingJob(request, metadata=self._metadata)
return response.id
Expand Down
1 change: 1 addition & 0 deletions tests/mocks/services.py
Original file line number Diff line number Diff line change
Expand Up @@ -1483,6 +1483,7 @@ async def SubmitCustomTrainingJob(self, stream: Stream[SubmitCustomTrainingJobRe
self.model_name = request.model_name
self.model_version = request.model_version
self.container_id = request.container_id
self.refresh_dataset_cache = request.refresh_dataset_cache
await stream.send_message(SubmitCustomTrainingJobResponse(id=self.job_id))

async def GetTrainingJob(self, stream: Stream[GetTrainingJobRequest, GetTrainingJobResponse]) -> None:
Expand Down
2 changes: 2 additions & 0 deletions tests/test_ml_training_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,9 +89,11 @@ async def test_custom_submit_training_job(self, service: MockMLTraining):
model_name=MODEL_NAME,
model_version=MODEL_VERSION,
container_id=CONTAINER_ID,
refresh_dataset_cache=True,
)
assert id == JOB_ID
assert service.container_id == CONTAINER_ID
assert service.refresh_dataset_cache is True

async def test_get_training_job(self, service: MockMLTraining):
async with ChannelFor([service]) as channel:
Expand Down
Loading