diff --git a/src/viam/app/ml_training_client.py b/src/viam/app/ml_training_client.py index c62f29d8b..a8dbd141a 100644 --- a/src/viam/app/ml_training_client.py +++ b/src/viam/app/ml_training_client.py @@ -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. @@ -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. @@ -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 diff --git a/tests/mocks/services.py b/tests/mocks/services.py index 49b2dbd9d..2b160b0ef 100644 --- a/tests/mocks/services.py +++ b/tests/mocks/services.py @@ -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: diff --git a/tests/test_ml_training_client.py b/tests/test_ml_training_client.py index 232bc5794..741b98899 100644 --- a/tests/test_ml_training_client.py +++ b/tests/test_ml_training_client.py @@ -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: