feat: enhance configuration and error handling in project setup

- Added new ignore rule for Ruff to allow temporary paths in tests.
- Introduced MyPy overrides for specific modules to ignore errors.
- Refactored `Cleanup` and `ExperimentTracking` classes to remove async keywords from methods, improving consistency in method signatures.
- Updated `Training` class methods to handle synchronous operations, enhancing performance and clarity.
- Adjusted `requirements.txt` to remove unnecessary Git dependency, streamlining project setup.
This commit is contained in:
vitor-aignosi
2026-04-07 10:25:17 -03:00
parent 6b1df7c3a7
commit 09ee92f100
21 changed files with 500 additions and 309 deletions

View File

@@ -67,7 +67,6 @@ class Activities(ExperimentTracking, Training, Cleanup):
Exception: If any parent class initialization fails
"""
ExperimentTracking.__init__(
self,
host=postgres_config['host'],

View File

@@ -62,7 +62,7 @@ class Cleanup(SientiaMonitoring):
) # name_YYYYMMDD_HHMMSS_microseconds
@activity.defn(name='cleanup_temp_directories')
async def cleanup_temp_directories(self, input_data: dict[str, Any]) -> None:
def cleanup_temp_directories(self, input_data: dict[str, Any]) -> None:
"""
Clean up stale temporary directories based on timestamp in directory name.
@@ -178,14 +178,14 @@ class Cleanup(SientiaMonitoring):
raise
finally:
await self._emit_metrics(
self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='cleanup_temp_directories',
emit_workflow_metric=True,
)
async def _emit_metrics(
def _emit_metrics(
self,
metadata: dict[str, Any],
metrics_status: str,
@@ -201,7 +201,7 @@ class Cleanup(SientiaMonitoring):
activity_name: Name of the activity being executed
"""
if emit_workflow_metric:
await self.emit_metric(
self.emit_metric_sync(
metric_object=WORKFLOW_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),
@@ -210,7 +210,7 @@ class Cleanup(SientiaMonitoring):
},
)
await self.emit_metric(
self.emit_metric_sync(
metric_object=ACTIVITY_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),

View File

@@ -11,7 +11,6 @@ import enum
from temporalio import activity, workflow
with workflow.unsafe.imports_passed_through():
import asyncio
import traceback
from collections.abc import Mapping
from datetime import UTC, datetime
@@ -110,9 +109,9 @@ class ExperimentTracking(Postgres):
# Silently ignore errors during garbage collection
pass
async def _execute_update(self, query: str, params: Mapping[str, Any]) -> dict[str, Any]:
def _execute_update(self, query: str, params: Mapping[str, Any]) -> dict[str, Any]:
"""
Execute an UPDATE SQL statement asynchronously.
Execute an UPDATE SQL statement.
Args:
query: Parameterized SQL string to execute.
@@ -122,12 +121,9 @@ class ExperimentTracking(Postgres):
dict: A dictionary containing the affected row count: {'rowcount': int}.
"""
def _run() -> dict[str, Any]:
with self.engine.begin() as connection:
result = connection.execute(text(query), params)
return {'rowcount': result.rowcount}
return await asyncio.to_thread(_run)
with self.engine.begin() as connection:
result = connection.execute(text(query), params)
return {'rowcount': result.rowcount}
def _build_status_update_query(
self, status: str | None, experiment_run_id: int
@@ -223,7 +219,7 @@ class ExperimentTracking(Postgres):
raise ValueError(f'Invalid update_type: {update_type}')
@activity.defn(name='update_experiment_run')
async def update_experiment_run(self, input_data: dict[str, Any]) -> None:
def update_experiment_run(self, input_data: dict[str, Any]) -> None:
"""
Update experiment run with status, errors, or model information.
@@ -256,7 +252,7 @@ class ExperimentTracking(Postgres):
update_type, experiment_run_id, input_data
)
result = await self._execute_update(sql_query, query_params)
result = self._execute_update(sql_query, query_params)
if result.get('rowcount', 0) == 0:
error_msg = (
@@ -273,7 +269,7 @@ class ExperimentTracking(Postgres):
error_msg = f'Error updating experiment run - ID: {experiment_run_id}, Status: {status}, Error: {str(e)}'
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata or {},
notification_id='UPDATE_EXPERIMENT_RUN_ERROR',
message=error_msg,

View File

@@ -12,6 +12,7 @@ with workflow.unsafe.imports_passed_through():
import traceback
from typing import Any
import mlflow
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
@@ -22,10 +23,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_model.model_repository.plugin_store import PluginStore
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
from model_manager.utils.models.train_model_result import TrainModelResult
import mlflow
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
class Training(SientiaMonitoring):
@@ -61,7 +60,7 @@ class Training(SientiaMonitoring):
self.minio_repository = minio_repository
@activity.defn(name='load_model_metadata')
async def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]:
def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Load model metadata/schemas from the model store.
@@ -90,7 +89,7 @@ class Training(SientiaMonitoring):
return train_params.to_dict()
except Exception as exc:
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata,
notification_id='LOAD_MODEL_METADATA_ERROR',
message=f'Error loading model metadata: {str(exc)}',
@@ -99,9 +98,9 @@ class Training(SientiaMonitoring):
attachment_content=trace,
)
raise
@activity.defn(name='validate_train_params')
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams:
def validate_train_params(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Validate and convert training parameters from dict to TrainModelParams.
@@ -115,7 +114,7 @@ class Training(SientiaMonitoring):
- All TrainModelParams fields (experiment_run_id, target_variable, etc.)
Returns:
TrainModelParams: Validated and converted training parameters
dict[str, Any]: Validated and converted training parameters as dictionary
Raises:
Exception: If validation fails (after sending notification)
@@ -123,7 +122,7 @@ class Training(SientiaMonitoring):
metadata = input_data.get('metadata', {})
try:
train_params = TrainModelParams.from_dict(input_data)
train_params.validate_business_rules()
self.info(
@@ -133,12 +132,12 @@ class Training(SientiaMonitoring):
metadata,
)
return train_params
return train_params.to_dict()
except Exception as e:
error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
message=error_msg,
@@ -149,7 +148,7 @@ class Training(SientiaMonitoring):
raise
@activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Train a machine learning model.
@@ -162,7 +161,7 @@ class Training(SientiaMonitoring):
input_data: Training configuration containing:
- metadata (dict): Workflow execution metadata.
- uploaded_file (BytesIO): Training data already downloaded from MinIO.
- train_params (TrainModelParams | dict): Training parameters.
- train_params (dict): Training parameters.
Returns:
dict[str, Any]: Serializable summary (run identifiers, run_dir for cleanup, regression metrics).
@@ -172,14 +171,11 @@ class Training(SientiaMonitoring):
Exception: If training fails (after sending notification).
"""
metadata = input_data.get('metadata')
train_params = input_data['train_params']
if isinstance(train_params, dict):
train_params = TrainModelParams.from_dict(train_params)
train_params = TrainModelParams.from_dict(input_data['train_params'])
try:
# Download training file bytes from MinIO
train_bytes = await self.minio_repository.download_file(
train_bytes = self.minio_repository.download_file_sync(
object_name=train_params.file_name,
bucket=train_params.bucket_name,
metadata=metadata,
@@ -189,7 +185,7 @@ class Training(SientiaMonitoring):
val_bytes: bytes | None = None
validation_name = train_params.val_file_name
if validation_name is not None:
val_bytes = await self.minio_repository.download_file(
val_bytes = self.minio_repository.download_file_sync(
object_name=validation_name,
bucket=train_params.bucket_name,
metadata=metadata,
@@ -202,13 +198,13 @@ class Training(SientiaMonitoring):
metadata=metadata,
)
wrapper = await self.plugin_store.get_model(
wrapper = self.plugin_store.get_model(
model_name=train_params.model_name,
force_download=False,
opt_params=train_params.opt_params or {},
model_kwargs=train_params.model_kwargs or {},
data_model_kwargs=train_params.data_model_kwargs or {},
metadata=metadata
metadata=metadata,
)
train_data = train_result.train_data
@@ -238,7 +234,7 @@ class Training(SientiaMonitoring):
wrapper,
)
async with self.mlflow_repository.start_run(
with self.mlflow_repository.start_run(
model_name=train_params.model_name,
run_name=None,
experiment_name=f'{train_params.model_name}_experiment',
@@ -247,33 +243,19 @@ class Training(SientiaMonitoring):
) as run_info:
train_result.run_name = run_info.run_name
train_result.run_id = run_info.run_id
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if train_result.report_path is None or train_result.train_data_path is None or train_result.test_data_path is None:
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
return {
'run_name': train_result.run_name,
'run_id': train_result.run_id,
'run_dir': train_result.run_dir
'run_dir': train_result.run_dir,
}
except Exception as e: # noqa: BLE001
error_msg = f'Error training model - error: {str(e)}'
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata or {},
notification_id='TRAIN_MODEL_ERROR',
message=error_msg,
@@ -284,8 +266,32 @@ class Training(SientiaMonitoring):
raise e
def _persist_training_artifacts(
self,
train_result: TrainModelResult,
train_params: TrainModelParams,
wrapper: Any,
metadata: dict[str, Any] | None,
) -> None:
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if (
train_result.report_path is None
or train_result.train_data_path is None
or train_result.test_data_path is None
):
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
@activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
def cleanup_resources(self, input_data: dict[str, Any]) -> None:
"""
Cleanup temporary resources created during training.
@@ -317,4 +323,3 @@ class Training(SientiaMonitoring):
)
raise

View File

@@ -62,8 +62,8 @@ class TrainModelParams:
# New Parameters
val_file_name: str | None
data_model_kwargs: dict | None # Removed params used in DataPreprocessor here
model_kwargs: dict | None # Removed params used in Linear Regression Model here
data_model_kwargs: dict | None # Removed params used in DataPreprocessor here
model_kwargs: dict | None # Removed params used in Linear Regression Model here
opt_params: dict | None
model_type: str
model_id: str | None
@@ -102,12 +102,16 @@ class TrainModelParams:
model_name = cls._check_none(data.get('model_name'), str, 'model_name')
return cls(
variable_columns=cls._check_none(data.get('variable_columns'), list, 'variable_columns'),
variable_columns=cls._check_none(
data.get('variable_columns'), list, 'variable_columns'
),
target_variable=cls._check_none(data.get('target_variable'), str, 'target_variable'),
bucket_name=cls._check_none(data.get('bucket_name'), str, 'bucket_name'),
file_name=cls._check_none(data.get('file_name'), str, 'file_name'),
line_separator=cls._check_none(data.get('line_separator'), str, 'line_separator'),
decimal_separator=cls._check_none(data.get('decimal_separator'), str, 'decimal_separator'),
decimal_separator=cls._check_none(
data.get('decimal_separator'), str, 'decimal_separator'
),
date_column=data.get('date_column'),
date_format=data.get('date_format'),
train_size=cls._check_none(data.get('train_size'), int, 'train_size'),
@@ -117,12 +121,13 @@ class TrainModelParams:
model_name=model_name,
experiment_name=model_name + '_experiment',
val_file_name=data.get('val_file_name'),
data_model_kwargs=cls._check_none(data.get('data_model_kwargs'), dict, 'data_model_kwargs'),
data_model_kwargs=cls._check_none(
data.get('data_model_kwargs'), dict, 'data_model_kwargs'
),
model_kwargs=cls._check_none(data.get('model_kwargs'), dict, 'model_kwargs'),
opt_params=cls._check_none(data.get('opt_params'), dict, 'opt_params'),
model_type=cls._check_none(data.get('model_type'), str, 'model_type'),
model_id=data.get('model_id'),
model_metadata=cls._parse_optional_model_metadata(data.get('model_metadata')),
)
@@ -233,9 +238,7 @@ class TrainModelParams:
return None
if isinstance(value, dict):
return value
raise TypeError(
f'model_metadata must be a dict or None, but got {type(value).__name__}.'
)
raise TypeError(f'model_metadata must be a dict or None, but got {type(value).__name__}.')
def validate_business_rules(self) -> None:
"""
@@ -261,21 +264,20 @@ class TrainModelParams:
if not self.variable_columns:
raise ValueError('variable_columns cannot be empty')
def _validate_model_params(self) -> None:
"""Validate model-related parameters."""
if not self.model_metadata:
raise ValueError('model_metadata is required')
schemas = self.model_metadata.get('schemas', {}).get("components", {}).get("schemas")
schemas = self.model_metadata.get('schemas', {}).get('components', {}).get('schemas')
if not schemas:
return
data_model_schema = schemas.get("data_model")
model_schema = schemas.get("model")
opt_params_schema = schemas.get("opt_params")
data_model_schema = schemas.get('data_model')
model_schema = schemas.get('model')
opt_params_schema = schemas.get('opt_params')
if data_model_schema:
self._validate_model_param(data_model_schema, self.data_model_kwargs)
@@ -283,8 +285,6 @@ class TrainModelParams:
self._validate_model_param(model_schema, self.model_kwargs)
if opt_params_schema:
self._validate_model_param(opt_params_schema, self.opt_params)
def _validate_model_param(self, schema: dict[str, Any], value: Any) -> None:
"""Validate model parameter against schema."""
@@ -292,9 +292,7 @@ class TrainModelParams:
validator = Draft202012Validator(schema)
validator.validate(value)
except ValidationError as e:
raise ValueError(f'Model parameters validation failed: {e.message}')
except Exception as e:
raise ValueError(f'Unexpected error: {e}')
raise ValueError(f'Model parameters validation failed: {e.message}') from e
def _validate_required_strings(self) -> None:
"""Validate required string fields are not empty."""
@@ -313,4 +311,4 @@ class TrainModelParams:
def _validate_date_format(self) -> None:
"""Validate date_format is one of the allowed frontend formats when set."""
if self.date_format:
validate_frontend_date_format(self.date_format)
validate_frontend_date_format(self.date_format)

View File

@@ -1,5 +1,4 @@
from dataclasses import dataclass
from typing import Any
import pandas as pd

View File

@@ -13,12 +13,12 @@ integration. Models are trained elsewhere (e.g., via SientiaModel wrappers),
and this repository focuses solely on preparing data structures for them.
"""
import json
from datetime import datetime
from io import BytesIO
import json
from os import makedirs, path
from typing import Any
from shutil import rmtree
from typing import Any
import numpy as np
import pandas as pd
@@ -28,35 +28,42 @@ from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
from sientia_model.wrappers.sientia_model import SientiaModel
from model_manager.sientia.metrics import mae, mse, r2
from model_manager.sientia.reports import Reports # type: ignore[import-untyped]
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.sientia.reports import Reports # type: ignore[import-untyped]
def train_test_split(data: pd.DataFrame | pd.Series, train_size: float, random_state: int | None = None, shuffle: bool = True) -> tuple[pd.DataFrame, pd.DataFrame]:
def train_test_split(
data: pd.DataFrame | pd.Series,
train_size: float,
random_state: int | None = None,
shuffle: bool = True,
) -> tuple[pd.DataFrame, pd.DataFrame]:
# 1. Definir a semente (seed) para reprodutibilidade
if random_state is not None:
np.random.seed(random_state)
# 2. Gerar índices e embaralhar se necessário
indices = np.arange(len(data))
if shuffle:
np.random.shuffle(indices)
# 3. Calcular o ponto de corte (split point)
# Cálculo: N_treino = tamanho_total * proporcao_treino
n_train = int(len(data) * train_size)
# 4. Dividir os índices
train_indices = indices[:n_train]
test_indices = indices[n_train:]
# 5. Retornar os dados fatiados (funciona para DataFrame ou Series)
if isinstance(data, (pd.DataFrame, pd.Series)):
return data.iloc[train_indices], data.iloc[test_indices]
return data[train_indices], data[test_indices]
def _ensure_date_column_parsed(data: pd.DataFrame, params: TrainModelParams) -> pd.DataFrame:
"""
If date_column is set, parse the column as timezone-aware
@@ -178,7 +185,6 @@ class DataManagerRepository(SientiaMonitoring):
if len(val_df) <= 0:
raise ValueError('Validation data view is empty after transformation')
val_data = pd.DataFrame(val_df[params.variable_columns + [params.target_variable]])
else:
# Fallback path: derive validation via train/test split from a single dataset.
@@ -194,11 +200,7 @@ class DataManagerRepository(SientiaMonitoring):
metadata,
)
return TrainModelResult(
params=params,
train_data=train_data,
val_data=val_data
)
return TrainModelResult(params=params, train_data=train_data, val_data=val_data)
def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series:
if isinstance(pred, pd.Series):
@@ -207,10 +209,8 @@ class DataManagerRepository(SientiaMonitoring):
if pred.shape[1] == 1:
return pred.iloc[:, 0]
raise ValueError('y_pred/y_train_pred must be a Series or single-column DataFrame')
def _extract_model_equation(
self, regr: Any, params: TrainModelParams
) -> dict:
def _extract_model_equation(self, regr: Any, params: TrainModelParams) -> dict:
"""
Extract the linear regression equation coefficients and create equation metadata.
@@ -237,7 +237,7 @@ class DataManagerRepository(SientiaMonitoring):
model_kwargs = params.model_kwargs or {}
degree = model_kwargs.get('degree', 1)
poly_feature_names = model_kwargs.get('poly_feature_names', None)
if degree > 1 and poly_feature_names:
feature_names = poly_feature_names
else:
@@ -271,7 +271,6 @@ class DataManagerRepository(SientiaMonitoring):
'original_features': feature_names,
}
def compute_regression_metrics(
self,
tmr: TrainModelResult,
@@ -301,7 +300,7 @@ class DataManagerRepository(SientiaMonitoring):
# True values are expected to come from val_data.
y_true_val = tmr.val_data[target]
y_pred_val = self._as_series(tmr.y_pred).sort_index()
y_true_val = y_true_val.sort_index()
@@ -407,7 +406,9 @@ class DataManagerRepository(SientiaMonitoring):
reports_dir = path.join(model_manager_dir, 'reports')
return reports_dir
def _create_run_directory(self, base_path: str, run_name: str, metadata: dict[str, Any] | None = None) -> str:
def _create_run_directory(
self, base_path: str, run_name: str, metadata: dict[str, Any] | None = None
) -> str:
"""
Creates a directory inside the 'reports' folder with the run name and a timestamp.
@@ -471,7 +472,6 @@ class DataManagerRepository(SientiaMonitoring):
reference_data_float = reference_data.astype(np.float64)
current_data_float = current_data.astype(np.float64)
# Initialize report generator
data.run_dir = self._create_run_directory(self._get_reports_directory(), data.run_name)
report = Reports(

View File

@@ -0,0 +1,106 @@
import os
import re
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor
from typing import Any
from sientia_do.observability.logger import Logger
from temporalio.client import Client
from temporalio.worker import PollerBehaviorAutoscaling, Worker
# Worker configuration parameters with default values.
parameters = [
('MAX_CONCURRENT_WORKFLOW_TASKS', '200'),
('MAX_CONCURRENT_ACTIVITIES', '200'),
('MAX_CONCURRENT_LOCAL_ACTIVITIES', '200'),
('MAX_CACHED_WORKFLOWS', '200'),
('WORKFLOW_POLLER_BEHAVIOUR_MINIMUM', '10'),
('WORKFLOW_POLLER_BEHAVIOUR_INITIAL', '100'),
('WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM', '200'),
('ACTIVITY_POLLER_BEHAVIOUR_MINIMUM', '10'),
('ACTIVITY_POLLER_BEHAVIOUR_INITIAL', '100'),
('ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM', '200'),
('ACTIVITY_EXECUTOR_MAX_WORKERS', '32'),
]
def camel_to_snake(text: str) -> str:
"""
Convert a CamelCase or camelCase string into snake_case.
Args:
- text: str, original string in CamelCase or camelCase format
Return:
str: converted string in snake_case format
"""
text = re.sub('(.)([A-Z][a-z]+)', r'\1_\2', text)
text = re.sub('([a-z0-9])([A-Z])', r'\1_\2', text)
return text.lower()
def prepare_worker(
main_workflow: type,
other_workflows: Sequence[type],
activities: Sequence[Any],
temporal_client: Client,
logger: Logger,
runtime: str | None = None,
) -> Worker:
"""
Build and configure a Temporal worker for the given workflow and activities.
Args:
- main_workflow: type, main workflow class used as worker entry point
- other_workflows: Sequence[type], additional workflows in the same worker
- activities: Sequence[Any], activity callables registered in this worker
- temporal_client: Client, Temporal client used by the worker
- logger: Logger, logger instance used during worker preparation
- runtime: str | None, runtime suffix appended to queue name when present
Return:
Worker: fully configured Temporal worker instance ready to run
"""
main_workflow_name = main_workflow.__name__.upper()
queue_name = (
f'{camel_to_snake(main_workflow.__name__)}-{runtime}-queue'
if runtime
else f'{camel_to_snake(main_workflow.__name__)}-queue'
)
local_workflow_parameters: dict[str, int] = {}
for parameter_name, default_value in parameters:
local_workflow_parameters[parameter_name] = int(
os.getenv(f'{main_workflow_name}_{parameter_name}', default_value)
)
logger.info(f'Preparing worker for {main_workflow_name} with queue {queue_name}')
activity_executor = ThreadPoolExecutor(
max_workers=local_workflow_parameters['ACTIVITY_EXECUTOR_MAX_WORKERS'],
thread_name_prefix=f'{queue_name}-activity',
)
return Worker(
temporal_client,
task_queue=queue_name,
workflows=[main_workflow, *other_workflows],
activities=[*activities],
activity_executor=activity_executor,
max_concurrent_workflow_tasks=local_workflow_parameters['MAX_CONCURRENT_WORKFLOW_TASKS'],
max_concurrent_activities=local_workflow_parameters['MAX_CONCURRENT_ACTIVITIES'],
max_concurrent_local_activities=local_workflow_parameters[
'MAX_CONCURRENT_LOCAL_ACTIVITIES'
],
max_cached_workflows=local_workflow_parameters['MAX_CACHED_WORKFLOWS'],
workflow_task_poller_behavior=PollerBehaviorAutoscaling(
minimum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MINIMUM'],
initial=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_INITIAL'],
maximum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM'],
),
activity_task_poller_behavior=PollerBehaviorAutoscaling(
minimum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MINIMUM'],
initial=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_INITIAL'],
maximum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM'],
),
)

View File

@@ -30,7 +30,6 @@ Environment Variables:
from temporalio import client, workflow
from temporalio.runtime import PrometheusConfig, Runtime, TelemetryConfig
from temporalio.worker import PollerBehaviorAutoscaling, Worker
with workflow.unsafe.imports_passed_through():
import asyncio
@@ -40,9 +39,8 @@ with workflow.unsafe.imports_passed_through():
from prometheus_client import start_http_server
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger as SientiaLogger
from sientia_model.model_repository.plugin_store import PluginStore
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.temporal.worker.prepare_worker import prepare_worker
from sientia_model.model_repository.plugin_store import PluginStore
from model_manager import metrics
from model_manager.activities.activities import Activities
@@ -55,6 +53,7 @@ with workflow.unsafe.imports_passed_through():
build_postgres_config,
)
from model_manager.utils.logger_helper import get_logger
from model_manager.worker.prepare_worker import prepare_worker
from model_manager.workflows.cleanup_files import CleanupFiles
from model_manager.workflows.train_model import TrainModel
@@ -96,7 +95,6 @@ async def main():
'pod_id': POD_ID,
'runtime': RUNTIME,
}
start_prometheus_server(logger, metadata)
mongo_config = build_mongodb_config()
@@ -110,11 +108,9 @@ async def main():
logger.custom_info(f'MongoDB client initialized at {mongo_config["uri"]}', metadata)
logger.custom_info(f'Initializing metrics controller', metadata)
logger.custom_info('Initializing metrics controller', metadata)
metrics_controller = MetricsController(
logger=logger
)
metrics_controller = MetricsController(logger=logger)
logger.custom_info(f'Installing runtime {RUNTIME}', metadata)
@@ -185,6 +181,7 @@ async def main():
],
temporal_client=temporal_client,
logger=logger,
runtime=RUNTIME,
),
prepare_worker(
main_workflow=CleanupFiles,
@@ -194,12 +191,11 @@ async def main():
],
temporal_client=temporal_client,
logger=logger,
runtime=RUNTIME,
),
]
handlers = [
w.run() for w in workers
]
handlers = [w.run() for w in workers]
logger.custom_info('Model manager workers initialized', metadata)

View File

@@ -21,7 +21,6 @@ with workflow.unsafe.imports_passed_through():
from model_manager.activities.activities import Activities
from model_manager.activities.experiment_tracking import UpdateType
from model_manager.utils.models.experiment_status import ExperimentStatus
from model_manager.utils.models.train_model_params import TrainModelParams
# Activity timeouts (seconds). Tune per environment (large uploads, long training).
# Training uses no_retry_policy: extend TIMEOUT_TRAIN_MODEL instead of adding retries
@@ -54,7 +53,6 @@ with workflow.unsafe.imports_passed_through():
)
@workflow.defn(name='train_model')
class TrainModel:
"""
@@ -133,7 +131,7 @@ class TrainModel:
)
else:
pass
except Exception:
except Exception: # noqa: BLE001
# If cleanup fails after training failed, there is nothing extra to log (DB not committed).
if training_succeeded: # pragma: no branch
workflow.logger.warning(
@@ -179,7 +177,7 @@ class TrainModel:
input_data: dict[str, Any],
experiment_run_id: int,
metadata: dict[str, Any],
) -> TrainModelParams:
) -> dict[str, Any]:
"""
Validate and convert training parameters from dict to TrainModelParams.
@@ -193,7 +191,7 @@ class TrainModel:
metadata: Workflow execution metadata
Returns:
TrainModelParams: Validated training parameters object
dict[str, Any]: Validated training parameters
Raises:
Exception: If validation fails (after updating DB status)
@@ -220,7 +218,7 @@ class TrainModel:
)
await self._update_experiment_run(
metadata=metadata,
metadata=metadata,
experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS,
status=ExperimentStatus.ORCHESTRATOR_WAITING_PROC,
@@ -236,7 +234,7 @@ class TrainModel:
status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR,
error_message=self._extract_error_message(e),
)
except Exception as secondary:
except Exception as secondary: # noqa: BLE001
workflow.logger.warning(
'Failed to persist ORCHESTRATOR_VALIDATION_ERROR to experiment_run: %s',
secondary,
@@ -245,7 +243,7 @@ class TrainModel:
async def _train_model(
self,
train_params: TrainModelParams,
train_params: dict[str, Any],
experiment_run_id: int,
metadata: dict[str, Any],
) -> dict[str, Any]:
@@ -288,7 +286,6 @@ class TrainModel:
return train_result
except Exception as e:
try:
await self._update_experiment_run(
metadata=metadata,
@@ -297,7 +294,7 @@ class TrainModel:
status=ExperimentStatus.TRAINING_ERROR,
error_message=self._extract_error_message(e),
)
except Exception as secondary:
except Exception as secondary: # noqa: BLE001
workflow.logger.warning(
'Failed to persist TRAINING_ERROR status to experiment_run: %s',
secondary,