Skip to content

Commit 114382f

Browse files
committed
fixes and docs
1 parent 86b2620 commit 114382f

4 files changed

Lines changed: 36 additions & 35 deletions

File tree

modules/sagemaker/sagemaker-model-monitoring/README.md

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -13,20 +13,7 @@ Available monitoring types:
1313
* Model Bias
1414
* Model Explainability
1515

16-
### Baseline Generation
17-
18-
The module includes an optional automated baseline generation feature that creates baseline statistics and constraints for your monitoring jobs. When you provide training data, the module will:
19-
20-
1. Deploy a Step Functions state machine that orchestrates baseline generation
21-
2. Run SageMaker Processing jobs to analyze your training data
22-
3. Generate baseline statistics and constraints files
23-
4. Store the baseline artifacts in your specified S3 location
24-
5. Schedule automatic baseline regeneration (default: daily at 2 AM UTC)
25-
26-
The baseline generation uses a Lambda function deployed as a Docker container image to handle the SageMaker SDK dependencies efficiently.
27-
28-
Note that updating parameters will require replacing resources. Deployments may be delayed until any
29-
running monitoring jobs complete (and the resources can be destroyed).
16+
The module includes an optional automated baseline generation feature that creates baseline statistics and constraints for your monitoring jobs.
3017

3118
### Architecture
3219

@@ -58,6 +45,22 @@ running monitoring jobs complete (and the resources can be destroyed).
5845
- **CloudWatch Metrics**: Some monitoring types emit metrics (e.g., data quality drift)
5946
- **CloudWatch Alarms**: Can be configured based on emitted metrics for automated alerting
6047

48+
### Baseline Generation
49+
50+
The module includes an optional automated baseline generation feature that creates baseline statistics and constraints for your monitoring jobs.
51+
52+
![SageMaker Model Monitoring Baseline Generation](docs/_static/sagemaker-model-monitoring-baseline.png "SageMaker Model Monitoring Baseline Generation")
53+
54+
When you provide training data, the module will:
55+
56+
1. Deploy a Step Functions state machine that orchestrates baseline generation
57+
2. Run SageMaker Processing jobs to analyze your training data
58+
3. Generate baseline statistics and constraints files
59+
4. Store the baseline artifacts in your specified S3 location
60+
5. Schedule automatic baseline regeneration (default: daily at 2 AM UTC)
61+
62+
## Inputs/Outputs
63+
6164
### Input Parameters
6265

6366
#### Required
142 KB
Loading

modules/sagemaker/sagemaker-model-monitoring/sagemaker_model_monitoring/baselining_construct.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,7 @@ def __init__(
108108
"baseline_output_uri": baseline_output_data_s3_uri,
109109
}
110110
),
111+
output_path="$.Payload",
111112
)
112113

113114
wait_task = sfn.Wait(self, "WaitForBaselining", time=sfn.WaitTime.duration(Duration.minutes(2)))
@@ -116,7 +117,8 @@ def __init__(
116117
self,
117118
"CheckBaselining",
118119
lambda_function=baselining_lambda,
119-
payload=sfn.TaskInput.from_object({"action": "check", "job_name.$": "$.Payload.job_name"}),
120+
payload=sfn.TaskInput.from_object({"action": "check", "job_name.$": "$.job_name"}),
121+
output_path="$.Payload",
120122
)
121123

122124
# State machine
@@ -126,7 +128,7 @@ def __init__(
126128
.next(
127129
sfn.Choice(self, "IsJobComplete")
128130
.when(
129-
sfn.Condition.string_equals("$.Payload.status", "COMPLETED"),
131+
sfn.Condition.string_equals("$.status", "COMPLETED"),
130132
sfn.Succeed(self, "JobCompleted"),
131133
)
132134
.otherwise(wait_task)

modules/sagemaker/sagemaker-model-monitoring/sagemaker_model_monitoring/lambda/baselining_handler.py

Lines changed: 15 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
# mypy: disable-error-code="attr-defined,no-untyped-call,assignment,call-arg"
22
import os
3-
from datetime import datetime
43
from typing import Any, Dict
54

65
import boto3
@@ -36,9 +35,7 @@ def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
3635
return {"statusCode": 500, "error": f"Internal error: {str(e)}"}
3736

3837

39-
def start_data_quality_baseline(
40-
job_name: str, training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]
41-
) -> None:
38+
def start_data_quality_baseline(training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]) -> str:
4239
"""Start data quality baseline job."""
4340
params = DataQualityParams(**event.get("data_quality_params", {}))
4441
monitor = DefaultModelMonitor(
@@ -55,11 +52,10 @@ def start_data_quality_baseline(
5552
wait=False,
5653
logs=False,
5754
)
55+
return monitor.latest_baselining_job.job_name
5856

5957

60-
def start_model_quality_baseline(
61-
job_name: str, training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]
62-
) -> None:
58+
def start_model_quality_baseline(training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]) -> str:
6359
"""Start model quality baseline job."""
6460
params = ModelQualityParams(**event.get("model_quality_params", {}))
6561
monitor = ModelQualityMonitor(
@@ -70,7 +66,6 @@ def start_model_quality_baseline(
7066
max_runtime_in_seconds=params.max_runtime_seconds,
7167
)
7268
monitor.suggest_baseline(
73-
job_name=job_name,
7469
baseline_dataset=training_data_uri,
7570
dataset_format=DatasetFormat.csv(header=True),
7671
output_s3_uri=baseline_output_uri,
@@ -81,11 +76,12 @@ def start_model_quality_baseline(
8176
wait=False,
8277
logs=False,
8378
)
79+
return monitor.latest_baselining_job.job_name
8480

8581

8682
def start_model_bias_baseline(
87-
job_name: str, training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
88-
) -> None:
83+
training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
84+
) -> str:
8985
"""Start model bias baseline job."""
9086
params = ModelBiasParams(**event.get("model_bias_params", {}))
9187
monitor = ModelBiasMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
@@ -118,11 +114,12 @@ def start_model_bias_baseline(
118114
wait=False,
119115
logs=False,
120116
)
117+
return monitor.latest_baselining_job.job_name
121118

122119

123120
def start_model_explainability_baseline(
124-
job_name: str, training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
125-
) -> None:
121+
training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
122+
) -> str:
126123
"""Start model explainability baseline job."""
127124
params = ModelExplainabilityParams(**event.get("model_explainability_params", {}))
128125
monitor = ModelExplainabilityMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
@@ -152,6 +149,7 @@ def start_model_explainability_baseline(
152149
model_config=model_config,
153150
explainability_config=shap_config,
154151
)
152+
return monitor.latest_baselining_job.job_name
155153

156154

157155
def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
@@ -164,18 +162,16 @@ def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
164162
except KeyError as e:
165163
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
166164

167-
timestamp = datetime.utcnow().strftime("%Y%m%d-%H%M%S")
168-
job_name = f"baseline-{monitor_type}-{timestamp}"
169-
170165
try:
166+
job_name = None
171167
if monitor_type == "data_quality":
172-
start_data_quality_baseline(job_name, training_data_uri, baseline_output_uri, event)
168+
job_name = start_data_quality_baseline(training_data_uri, baseline_output_uri, event)
173169
elif monitor_type == "model_quality":
174-
start_model_quality_baseline(job_name, training_data_uri, baseline_output_uri, event)
170+
job_name = start_model_quality_baseline(training_data_uri, baseline_output_uri, event)
175171
elif monitor_type == "model_bias":
176-
start_model_bias_baseline(job_name, training_data_uri, baseline_output_uri, endpoint_name, event)
172+
job_name = start_model_bias_baseline(training_data_uri, baseline_output_uri, endpoint_name, event)
177173
elif monitor_type == "model_explainability":
178-
start_model_explainability_baseline(job_name, training_data_uri, baseline_output_uri, endpoint_name, event)
174+
job_name = start_model_explainability_baseline(training_data_uri, baseline_output_uri, endpoint_name, event)
179175
else:
180176
return {"statusCode": 400, "error": f"Unsupported monitor_type: {monitor_type}"}
181177

0 commit comments

Comments
 (0)