diff --git a/model_manager/workflows/train_model.py b/model_manager/workflows/train_model.py index 3f41f5e..69aa7c5 100644 --- a/model_manager/workflows/train_model.py +++ b/model_manager/workflows/train_model.py @@ -16,7 +16,7 @@ with workflow.unsafe.imports_passed_through(): from datetime import timedelta from typing import Any - from sientia_do.temporal.policies import retry_policy + from temporalio.common import RetryPolicy from model_manager.activities.activities import Activities from model_manager.activities.experiment_tracking import UpdateType @@ -34,6 +34,44 @@ with workflow.unsafe.imports_passed_through(): TIMEOUT_DELETE_FILE = int(os.getenv('TIMEOUT_DELETE_FILE', '60')) TIMEOUT_UPDATE_DATABASE = int(os.getenv('TIMEOUT_UPDATE_DATABASE', '30')) + # Retry Policies - Granular strategies for different operation types + # Fast retry for transient network errors (MinIO operations) + network_retry_policy = RetryPolicy( + initial_interval=timedelta(seconds=1), + maximum_interval=timedelta(seconds=10), + backoff_coefficient=2.0, + maximum_attempts=5, + ) + + # No retry for training - data errors are permanent + no_retry_policy = RetryPolicy( + maximum_attempts=1, + ) + + # Moderate retry with backoff for MLFlow operations + mlflow_retry_policy = RetryPolicy( + initial_interval=timedelta(seconds=5), + maximum_interval=timedelta(seconds=30), + backoff_coefficient=2.0, + maximum_attempts=3, + ) + + # Database retry with exponential backoff + database_retry_policy = RetryPolicy( + initial_interval=timedelta(seconds=2), + maximum_interval=timedelta(seconds=20), + backoff_coefficient=2.0, + maximum_attempts=5, + ) + + # Filesystem retry for cleanup operations + filesystem_retry_policy = RetryPolicy( + initial_interval=timedelta(seconds=2), + maximum_interval=timedelta(seconds=10), + backoff_coefficient=1.5, + maximum_attempts=3, + ) + @workflow.defn(name='train_model') class TrainModel: @@ -183,7 +221,7 @@ class TrainModel: train_params = await workflow.execute_activity_method( Activities.validate_train_params, validation_input, - retry_policy=retry_policy, + retry_policy=no_retry_policy, # Validation errors are permanent start_to_close_timeout=timedelta(seconds=TIMEOUT_VALIDATE_PARAMS), ) @@ -253,7 +291,7 @@ class TrainModel: uploaded_file = await workflow.execute_activity_method( Activities.fetch_file_from_minio, download_input, - retry_policy=retry_policy, + retry_policy=network_retry_policy, # Fast retry for network issues start_to_close_timeout=timedelta(seconds=TIMEOUT_DOWNLOAD_FILE), ) @@ -267,7 +305,7 @@ class TrainModel: train_result = await workflow.execute_activity_method( Activities.train_model, train_input, - retry_policy=retry_policy, + retry_policy=no_retry_policy, # Training errors are permanent (bad data) start_to_close_timeout=timedelta(seconds=TIMEOUT_TRAIN_MODEL), ) @@ -340,7 +378,7 @@ class TrainModel: saved_result = await workflow.execute_activity_method( Activities.save_model, save_input, - retry_policy=retry_policy, + retry_policy=mlflow_retry_policy, # Retry MLFlow with backoff start_to_close_timeout=timedelta(seconds=TIMEOUT_SAVE_MODEL), ) @@ -406,7 +444,7 @@ class TrainModel: await workflow.execute_activity_method( Activities.cleanup_run_directory, cleanup_input, - retry_policy=retry_policy, + retry_policy=filesystem_retry_policy, # Retry filesystem operations start_to_close_timeout=timedelta(seconds=TIMEOUT_CLEANUP_DIRECTORY), ) @@ -424,7 +462,7 @@ class TrainModel: await workflow.execute_activity_method( Activities.delete_file_from_minio, delete_input, - retry_policy=retry_policy, + retry_policy=network_retry_policy, # Fast retry for network issues start_to_close_timeout=timedelta(seconds=TIMEOUT_DELETE_FILE), ) @@ -496,6 +534,6 @@ class TrainModel: await workflow.execute_activity_method( Activities.update_experiment_run, update_input, - retry_policy=retry_policy, + retry_policy=database_retry_policy, # Retry DB with exponential backoff start_to_close_timeout=timedelta(seconds=TIMEOUT_UPDATE_DATABASE), )