feat: improve error handling and resource cleanup in training workflow

- Updated the `Training` class to raise `ModelTrainingError` on training failures for better error management.
- Enhanced the `run` method in `TrainModel` to return training results and ensure proper resource cleanup, including validation files.
- Refactored exception handling to prevent silent failures during resource cleanup and experiment run updates.
- Adjusted type hints for improved clarity and consistency in method signatures.
This commit is contained in:
vitor-aignosi
2026-04-06 08:26:42 -03:00
parent 342a02d6f7
commit bf2b6b3888
2 changed files with 77 additions and 47 deletions

View File

@@ -72,7 +72,7 @@ class TrainModel:
"""
@workflow.run
async def run(self, input_data: dict[str, Any]) -> None:
async def run(self, input_data: dict[str, Any]) -> dict[str, str | None]:
"""
Execute the complete model training workflow.
@@ -112,18 +112,29 @@ class TrainModel:
input_data, experiment_run_id, metadata
)
train_result = await self._train_model(
train_params=train_params,
experiment_run_id=experiment_run_id,
metadata=metadata,
)
training_succeeded = False
train_result: dict[str, str | None]
await self._cleanup_resources(
experiment_run_id=experiment_run_id,
bucket_name=train_params.bucket_name,
file_name=train_params.file_name,
metadata=metadata,
)
try:
train_result = await self._train_model(
train_params=train_params,
experiment_run_id=experiment_run_id,
metadata=metadata,
)
training_succeeded = True
finally:
try:
await self._cleanup_resources(
experiment_run_id=experiment_run_id,
bucket_name=train_params.bucket_name,
file_name=train_params.file_name,
val_file_name=train_params.val_file_name,
metadata=metadata,
)
except Exception:
if training_succeeded:
raise
return train_result
def _validate_experiment_run_id(self, input_data: dict[str, Any]) -> int:
"""
@@ -147,12 +158,15 @@ class TrainModel:
if experiment_run_id is None:
raise ValueError('experiment_run_id is required but was not provided')
if not isinstance(experiment_run_id, int):
raise ValueError(
f'experiment_run_id must be an integer, got {type(experiment_run_id).__name__}'
)
if isinstance(experiment_run_id, int):
return experiment_run_id
return experiment_run_id
if isinstance(experiment_run_id, str) and experiment_run_id.strip().isdigit():
return int(experiment_run_id.strip())
raise ValueError(
f'experiment_run_id must be an integer or numeric string, got {type(experiment_run_id).__name__}'
)
async def _validate_training_parameters(
self,
@@ -208,14 +222,16 @@ class TrainModel:
return train_params
except Exception as e:
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR,
error_message=self._extract_error_message(e),
)
try:
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR,
error_message=self._extract_error_message(e),
)
except Exception:
pass
raise
async def _train_model(
@@ -273,14 +289,16 @@ class TrainModel:
if isinstance(e, ModelTrainingError) and (e.model_trained and not e.model_saved):
status = ExperimentStatus.TRACKING_SEND_ERROR
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=status,
error_message=self._extract_error_message(e),
)
try:
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=status,
error_message=self._extract_error_message(e),
)
except Exception:
pass
raise
async def _cleanup_resources(
@@ -288,6 +306,7 @@ class TrainModel:
experiment_run_id: int,
bucket_name: str,
file_name: str,
val_file_name: str | None,
metadata: dict[str, Any],
) -> None:
"""
@@ -311,6 +330,7 @@ class TrainModel:
**metadata,
'bucket_name': bucket_name,
'file_name': file_name,
'val_file_name': val_file_name,
},
retry_policy=network_retry_policy,
start_to_close_timeout=timedelta(seconds=TIMEOUT_DELETE_FILE),
@@ -323,14 +343,16 @@ class TrainModel:
status=ExperimentStatus.FILE_DELETED,
)
except Exception as e:
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.FILE_DELETE_ERROR,
error_message=self._extract_error_message(e),
)
try:
await self._update_experiment_run(
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS_WITH_ERROR,
status=ExperimentStatus.FILE_DELETE_ERROR,
error_message=self._extract_error_message(e),
)
except Exception:
pass
raise
async def _update_experiment_run(
@@ -388,7 +410,9 @@ class TrainModel:
if text and text not in message_parts:
message_parts.append(text)
current = getattr(current, 'cause', None)
cause = getattr(current, '__cause__', None)
context = getattr(current, '__context__', None)
current = cause if isinstance(cause, Exception) else context
if not message_parts:
return repr(exc)