Skip to content

Commit cebbc97

Browse files
authored
Merge branch 'main' into release/2.1.0
2 parents d50d80d + 5ea9b34 commit cebbc97

4 files changed

Lines changed: 134 additions & 28 deletions

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
2020
- updated `sagemaker-model-monitor` module and tested 4 types of monitoring end-to-end
2121
- added baseline generation step function to `sagemaker-model-monitor` module
2222
- updated lambda runtime and depdencies in `sagemaker-templates` module
23+
- fixed baseline error handling in `sagemaker-model-monitor` module
2324

2425
## v2.0.0
2526

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

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ def __init__(
5050
memory_size=2048,
5151
environment={
5252
"SAGEMAKER_ROLE_ARN": sagemaker_role_arn,
53+
"LOG_LEVEL": "INFO",
5354
},
5455
)
5556

@@ -131,6 +132,15 @@ def __init__(
131132
sfn.Condition.string_equals("$.status", "COMPLETED"),
132133
sfn.Succeed(self, "JobCompleted"),
133134
)
135+
.when(
136+
sfn.Condition.string_equals("$.status", "FAILED"),
137+
sfn.Fail(
138+
self,
139+
"JobFailed",
140+
cause="Baselining job failed",
141+
error="ProcessingJobFailed",
142+
),
143+
)
134144
.otherwise(wait_task)
135145
)
136146
)

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

Lines changed: 62 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,6 @@
11
# mypy: disable-error-code="attr-defined,no-untyped-call,assignment,call-arg,no-any-return,union-attr"
2+
import json
3+
import logging
24
import os
35
from typing import Any, Dict
46

@@ -13,15 +15,22 @@
1315
)
1416
from sagemaker.model_monitor.dataset_format import DatasetFormat
1517

18+
logger = logging.getLogger()
19+
logger.setLevel(os.environ.get("LOG_LEVEL", "INFO"))
20+
log = logging.getLogger(__name__)
21+
1622
SAGEMAKER_ROLE_ARN = os.environ["SAGEMAKER_ROLE_ARN"]
1723

1824
sagemaker = boto3.client("sagemaker")
1925

2026

2127
def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
2228
"""Lambda handler for SageMaker Model Monitor baselining jobs."""
29+
logger.info(f"Received event: {json.dumps(event)}")
30+
2331
try:
2432
action = event.get("action", "start")
33+
logger.info(f"Action: {action}")
2534

2635
if action == "start":
2736
return start_baselining_job(event)
@@ -30,14 +39,18 @@ def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
3039

3140
return {"statusCode": 200}
3241
except KeyError as e:
42+
logger.error(f"Missing required parameter: {e}")
3343
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
3444
except Exception as e:
45+
logger.exception(f"Internal error: {str(e)}")
3546
return {"statusCode": 500, "error": f"Internal error: {str(e)}"}
3647

3748

3849
def start_data_quality_baseline(training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]) -> str:
3950
"""Start data quality baseline job."""
4051
params = DataQualityParams(**event.get("data_quality_params", {}))
52+
logger.info(f"Data quality baseline params: {params.model_dump_json()}")
53+
4154
monitor = DefaultModelMonitor(
4255
role=SAGEMAKER_ROLE_ARN,
4356
instance_count=params.instance_count,
@@ -52,12 +65,16 @@ def start_data_quality_baseline(training_data_uri: str, baseline_output_uri: str
5265
wait=False,
5366
logs=False,
5467
)
55-
return monitor.latest_baselining_job.job_name
68+
job_name = monitor.latest_baselining_job.job_name
69+
logger.info(f"Started data quality baseline job: {job_name}")
70+
return job_name
5671

5772

5873
def start_model_quality_baseline(training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]) -> str:
5974
"""Start model quality baseline job."""
6075
params = ModelQualityParams(**event.get("model_quality_params", {}))
76+
logger.info(f"Model quality baseline params: {params.model_dump_json()}")
77+
6178
monitor = ModelQualityMonitor(
6279
role=SAGEMAKER_ROLE_ARN,
6380
instance_count=params.instance_count,
@@ -76,14 +93,18 @@ def start_model_quality_baseline(training_data_uri: str, baseline_output_uri: st
7693
wait=False,
7794
logs=False,
7895
)
79-
return monitor.latest_baselining_job.job_name
96+
job_name = monitor.latest_baselining_job.job_name
97+
logger.info(f"Started model quality baseline job: {job_name}")
98+
return job_name
8099

81100

82101
def start_model_bias_baseline(
83102
training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
84103
) -> str:
85104
"""Start model bias baseline job."""
86105
params = ModelBiasParams(**event.get("model_bias_params", {}))
106+
logger.info(f"Model bias baseline params: {params.model_dump_json()}")
107+
87108
monitor = ModelBiasMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
88109

89110
model_bias_data_config = DataConfig(
@@ -114,14 +135,18 @@ def start_model_bias_baseline(
114135
wait=False,
115136
logs=False,
116137
)
117-
return monitor.latest_baselining_job.job_name
138+
job_name = monitor.latest_baselining_job.job_name
139+
logger.info(f"Started model bias baseline job: {job_name}")
140+
return job_name
118141

119142

120143
def start_model_explainability_baseline(
121144
training_data_uri: str, baseline_output_uri: str, endpoint_name: str, event: Dict[str, Any]
122145
) -> str:
123146
"""Start model explainability baseline job."""
124147
params = ModelExplainabilityParams(**event.get("model_explainability_params", {}))
148+
logger.info(f"Model explainability baseline params: {params.model_dump_json()}")
149+
125150
monitor = ModelExplainabilityMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
126151

127152
model_explainability_data_config = DataConfig(
@@ -149,7 +174,9 @@ def start_model_explainability_baseline(
149174
model_config=model_config,
150175
explainability_config=shap_config,
151176
)
152-
return monitor.latest_baselining_job.job_name
177+
job_name = monitor.latest_baselining_job.job_name
178+
logger.info(f"Started model explainability baseline job: {job_name}")
179+
return job_name
153180

154181

155182
def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
@@ -159,7 +186,13 @@ def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
159186
endpoint_name = event["endpoint_name"]
160187
training_data_uri = event["training_data_uri"]
161188
baseline_output_uri = event["baseline_output_uri"]
189+
190+
logger.info(
191+
f"Starting baselining job - Monitor type: {monitor_type}, Endpoint: {endpoint_name}, "
192+
f"Training URI: {training_data_uri}, Output URI: {baseline_output_uri}"
193+
)
162194
except KeyError as e:
195+
logger.error(f"Missing required parameter: {e}")
163196
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
164197

165198
try:
@@ -173,29 +206,48 @@ def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
173206
elif monitor_type == "model_explainability":
174207
job_name = start_model_explainability_baseline(training_data_uri, baseline_output_uri, endpoint_name, event)
175208
else:
209+
logger.error(f"Unsupported monitor_type: {monitor_type}")
176210
return {"statusCode": 400, "error": f"Unsupported monitor_type: {monitor_type}"}
177211

212+
logger.info(f"Successfully started baselining job: {job_name}")
178213
return {"statusCode": 200, "job_name": job_name, "status": "STARTED"}
179214
except Exception as e:
215+
logger.exception(f"Failed to start baseline job: {str(e)}")
180216
return {"statusCode": 500, "error": f"Failed to start baseline job: {str(e)}"}
181217

182218

183219
def check_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
184220
"""Check the status of a baselining processing job."""
185221
try:
186222
job_name = event["job_name"]
223+
logger.info(f"Checking status for job: {job_name}")
187224

188225
response = sagemaker.describe_processing_job(ProcessingJobName=job_name)
189226
status = response["ProcessingJobStatus"]
190-
191-
return {
192-
"statusCode": 200,
193-
"job_name": job_name,
194-
"status": "COMPLETED" if status == "Completed" else "IN_PROGRESS",
195-
}
227+
logger.info(f"Job {job_name} status: {status}")
228+
229+
# Map SageMaker processing job statuses to simplified states
230+
if status == "Completed":
231+
logger.info(f"Job {job_name} completed successfully")
232+
return {"statusCode": 200, "job_name": job_name, "status": "COMPLETED"}
233+
elif status in ["Failed", "Stopped", "Stopping"]:
234+
failure_reason = response.get("FailureReason", "Unknown failure reason")
235+
logger.error(f"Job {job_name} failed: {failure_reason}")
236+
return {
237+
"statusCode": 200,
238+
"job_name": job_name,
239+
"status": "FAILED",
240+
"failure_reason": failure_reason,
241+
}
242+
else: # InProgress
243+
logger.info(f"Job {job_name} still in progress")
244+
return {"statusCode": 200, "job_name": job_name, "status": "IN_PROGRESS"}
196245
except KeyError as e:
246+
logger.error(f"Missing required parameter: {e}")
197247
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
198248
except sagemaker.exceptions.ResourceNotFound:
249+
logger.error(f"Processing job not found: {event.get('job_name', 'unknown')}")
199250
return {"statusCode": 404, "error": f"Processing job not found: {event.get('job_name', 'unknown')}"}
200251
except Exception as e:
252+
logger.exception(f"Failed to check job status: {str(e)}")
201253
return {"statusCode": 500, "error": f"Failed to check job status: {str(e)}"}

modules/sagemaker/sagemaker-model-monitoring/tests/test_baselining_handler.py

Lines changed: 61 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -84,31 +84,29 @@ def test_lambda_handler_check_action(mock_sagemaker_modules, mock_env, mock_sage
8484
assert result["job_name"] == "test-job"
8585

8686

87-
def test_check_baselining_job_completed(mock_sagemaker_modules, mock_env, mock_sagemaker_client):
88-
"""Test check_baselining_job with completed job."""
87+
@pytest.mark.parametrize(
88+
"processing_job_status,expected_status",
89+
[
90+
("Completed", "COMPLETED"),
91+
("Failed", "FAILED"),
92+
("InProgress", "IN_PROGRESS"),
93+
("Stopping", "FAILED"),
94+
("Stopped", "FAILED"),
95+
],
96+
)
97+
def test_check_baselining_job_status(
98+
mock_sagemaker_modules, mock_env, mock_sagemaker_client, processing_job_status, expected_status
99+
):
100+
"""Test check_baselining_job with various job statuses."""
89101
check_baselining_job, lambda_handler = mock_sagemaker_modules
90102
event = {"job_name": "test-job"}
91103

92-
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": "Completed"}
93-
94-
result = check_baselining_job(event)
95-
96-
assert result["statusCode"] == 200
97-
assert result["status"] == "COMPLETED"
98-
assert result["job_name"] == "test-job"
99-
100-
101-
def test_check_baselining_job_in_progress(mock_sagemaker_modules, mock_env, mock_sagemaker_client):
102-
"""Test check_baselining_job with in-progress job."""
103-
check_baselining_job, lambda_handler = mock_sagemaker_modules
104-
event = {"job_name": "test-job"}
105-
106-
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": "InProgress"}
104+
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": processing_job_status}
107105

108106
result = check_baselining_job(event)
109107

110108
assert result["statusCode"] == 200
111-
assert result["status"] == "IN_PROGRESS"
109+
assert result["status"] == expected_status
112110
assert result["job_name"] == "test-job"
113111

114112

@@ -122,3 +120,48 @@ def test_check_baselining_job_missing_job_name(mock_sagemaker_modules, mock_env)
122120
assert result["statusCode"] == 400
123121
assert "error" in result
124122
assert "Missing required parameter" in result["error"]
123+
124+
125+
@pytest.mark.parametrize(
126+
"monitor_type,monitor_class_path,expected_job_name",
127+
[
128+
("data_quality", "baselining_handler.DefaultModelMonitor", "data-quality-job-123"),
129+
("model_quality", "baselining_handler.ModelQualityMonitor", "model-quality-job-456"),
130+
("model_bias", "baselining_handler.ModelBiasMonitor", "model-bias-job-789"),
131+
("model_explainability", "baselining_handler.ModelExplainabilityMonitor", "model-explainability-job-012"),
132+
],
133+
)
134+
def test_start_baselining_job_types(
135+
mock_sagemaker_modules, mock_env, monitor_type, monitor_class_path, expected_job_name
136+
):
137+
"""Test starting different types of baselining jobs."""
138+
check_baselining_job, lambda_handler = mock_sagemaker_modules
139+
140+
event = {
141+
"action": "start",
142+
"monitor_type": monitor_type,
143+
"endpoint_name": "test-endpoint",
144+
"training_data_uri": "s3://bucket/training-data",
145+
"baseline_output_uri": "s3://bucket/baseline-output",
146+
}
147+
148+
with patch(monitor_class_path) as mock_monitor_class:
149+
mock_monitor_instance = MagicMock()
150+
mock_monitor_class.return_value = mock_monitor_instance
151+
152+
# Mock the baselining job
153+
mock_baselining_job = MagicMock()
154+
mock_baselining_job.job_name = expected_job_name
155+
mock_monitor_instance.latest_baselining_job = mock_baselining_job
156+
157+
result = lambda_handler(event, None)
158+
159+
assert result["statusCode"] == 200
160+
assert result["status"] == "STARTED"
161+
assert result["job_name"] == expected_job_name
162+
163+
# Verify the monitor was instantiated
164+
mock_monitor_class.assert_called_once()
165+
166+
# Verify suggest_baseline was called
167+
mock_monitor_instance.suggest_baseline.assert_called_once()

0 commit comments

Comments
 (0)