Skip to content

Commit fe2f742

Browse files
committed
refactoring
1 parent 30bfdbe commit fe2f742

1 file changed

Lines changed: 174 additions & 157 deletions

File tree

Lines changed: 174 additions & 157 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,6 @@
1+
# mypy: disable-error-code="attr-defined,no-untyped-call,assignment,call-arg"
12
import os
3+
from datetime import datetime
24
from typing import Any, Dict
35

46
import boto3
@@ -19,170 +21,185 @@
1921

2022
def lambda_handler(event: Dict[str, Any], context: Any) -> Dict[str, Any]:
2123
"""Lambda handler for SageMaker Model Monitor baselining jobs."""
22-
23-
action = event.get("action", "start")
24-
25-
if action == "start":
26-
return start_baselining_job(event)
27-
elif action == "check":
28-
return check_baselining_job(event)
29-
30-
return {"statusCode": 200}
24+
try:
25+
action = event.get("action", "start")
26+
27+
if action == "start":
28+
return start_baselining_job(event)
29+
elif action == "check":
30+
return check_baselining_job(event)
31+
32+
return {"statusCode": 200}
33+
except KeyError as e:
34+
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
35+
except Exception as e:
36+
return {"statusCode": 500, "error": f"Internal error: {str(e)}"}
37+
38+
39+
def start_data_quality_baseline(
40+
job_name: str, training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]
41+
) -> None:
42+
"""Start data quality baseline job."""
43+
params = DataQualityParams(**event.get("data_quality_params", {}))
44+
monitor = DefaultModelMonitor(
45+
role=SAGEMAKER_ROLE_ARN,
46+
instance_count=params.instance_count,
47+
instance_type=params.instance_type,
48+
volume_size_in_gb=params.volume_size_gb,
49+
max_runtime_in_seconds=params.max_runtime_seconds,
50+
)
51+
monitor.suggest_baseline(
52+
baseline_dataset=training_data_uri,
53+
dataset_format=DatasetFormat.csv(header=True),
54+
output_s3_uri=baseline_output_uri,
55+
wait=False,
56+
logs=False,
57+
)
58+
59+
60+
def start_model_quality_baseline(
61+
job_name: str, training_data_uri: str, baseline_output_uri: str, event: Dict[str, Any]
62+
) -> None:
63+
"""Start model quality baseline job."""
64+
params = ModelQualityParams(**event.get("model_quality_params", {}))
65+
monitor = ModelQualityMonitor(
66+
role=SAGEMAKER_ROLE_ARN,
67+
instance_count=params.instance_count,
68+
instance_type=params.instance_type,
69+
volume_size_in_gb=params.volume_size_gb,
70+
max_runtime_in_seconds=params.max_runtime_seconds,
71+
)
72+
monitor.suggest_baseline(
73+
job_name=job_name,
74+
baseline_dataset=training_data_uri,
75+
dataset_format=DatasetFormat.csv(header=True),
76+
output_s3_uri=baseline_output_uri,
77+
problem_type=params.problem_type,
78+
inference_attribute=params.inference_attribute,
79+
probability_attribute=params.probability_attribute,
80+
ground_truth_attribute=params.ground_truth_attribute,
81+
wait=False,
82+
logs=False,
83+
)
84+
85+
86+
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:
89+
"""Start model bias baseline job."""
90+
params = ModelBiasParams(**event.get("model_bias_params", {}))
91+
monitor = ModelBiasMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
92+
93+
model_bias_data_config = DataConfig(
94+
s3_data_input_path=training_data_uri,
95+
s3_output_path=baseline_output_uri,
96+
label=params.label_header,
97+
headers=params.get_headers(),
98+
dataset_type=params.dataset_type,
99+
)
100+
model_bias_config = BiasConfig(
101+
label_values_or_threshold=params.get_label_values(),
102+
facet_name=params.facet_name,
103+
facet_values_or_threshold=params.get_facet_values(),
104+
)
105+
model_predicted_label_config = ModelPredictedLabelConfig(probability_threshold=params.probability_threshold)
106+
model_config = ModelConfig(
107+
model_name=params.model_name or endpoint_name,
108+
instance_count=params.instance_count,
109+
instance_type=params.instance_type,
110+
content_type=params.dataset_type,
111+
accept_type=params.dataset_type,
112+
)
113+
monitor.suggest_baseline(
114+
model_config=model_config,
115+
data_config=model_bias_data_config,
116+
bias_config=model_bias_config,
117+
model_predicted_label_config=model_predicted_label_config,
118+
wait=False,
119+
logs=False,
120+
)
121+
122+
123+
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:
126+
"""Start model explainability baseline job."""
127+
params = ModelExplainabilityParams(**event.get("model_explainability_params", {}))
128+
monitor = ModelExplainabilityMonitor(role=SAGEMAKER_ROLE_ARN, max_runtime_in_seconds=params.max_runtime_seconds)
129+
130+
model_explainability_data_config = DataConfig(
131+
s3_data_input_path=training_data_uri,
132+
s3_output_path=baseline_output_uri,
133+
label=params.label_header,
134+
headers=params.get_headers(),
135+
dataset_type=params.dataset_type,
136+
)
137+
shap_config = SHAPConfig(
138+
baseline=params.shap_baseline,
139+
num_samples=params.num_samples,
140+
agg_method=params.agg_method,
141+
save_local_shap_values=params.save_local_shap_values,
142+
)
143+
model_config = ModelConfig(
144+
model_name=params.model_name or endpoint_name,
145+
instance_count=params.instance_count,
146+
instance_type=params.instance_type,
147+
content_type=params.dataset_type,
148+
accept_type=params.dataset_type,
149+
)
150+
monitor.suggest_baseline(
151+
data_config=model_explainability_data_config,
152+
model_config=model_config,
153+
explainability_config=shap_config,
154+
)
31155

32156

33157
def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
34158
"""Start a baselining processing job."""
35-
36-
monitor_type = event["monitor_type"]
37-
endpoint_name = event["endpoint_name"]
38-
training_data_uri = event["training_data_uri"]
39-
baseline_output_uri = event["baseline_output_uri"]
40-
41-
if monitor_type == "data_quality":
42-
job_name = f"{endpoint_name}-data-quality-baseline"
43-
params = DataQualityParams(**event.get("data_quality_params", {}))
44-
45-
monitor = DefaultModelMonitor(
46-
role=SAGEMAKER_ROLE_ARN,
47-
instance_count=params.instance_count,
48-
instance_type=params.instance_type,
49-
volume_size_in_gb=params.volume_size_gb,
50-
max_runtime_in_seconds=params.max_runtime_seconds,
51-
)
52-
53-
monitor.suggest_baseline(
54-
baseline_dataset=training_data_uri,
55-
dataset_format=DatasetFormat.csv(header=True),
56-
output_s3_uri=baseline_output_uri,
57-
wait=False,
58-
logs=False,
59-
)
159+
try:
160+
monitor_type = event["monitor_type"]
161+
endpoint_name = event["endpoint_name"]
162+
training_data_uri = event["training_data_uri"]
163+
baseline_output_uri = event["baseline_output_uri"]
164+
except KeyError as e:
165+
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
166+
167+
timestamp = datetime.utcnow().strftime("%Y%m%d-%H%M%S")
168+
job_name = f"baseline-{monitor_type}-{timestamp}"
169+
170+
try:
171+
if monitor_type == "data_quality":
172+
start_data_quality_baseline(job_name, training_data_uri, baseline_output_uri, event)
173+
elif monitor_type == "model_quality":
174+
start_model_quality_baseline(job_name, training_data_uri, baseline_output_uri, event)
175+
elif monitor_type == "model_bias":
176+
start_model_bias_baseline(job_name, training_data_uri, baseline_output_uri, endpoint_name, event)
177+
elif monitor_type == "model_explainability":
178+
start_model_explainability_baseline(job_name, training_data_uri, baseline_output_uri, endpoint_name, event)
179+
else:
180+
return {"statusCode": 400, "error": f"Unsupported monitor_type: {monitor_type}"}
60181

61182
return {"statusCode": 200, "job_name": job_name, "status": "STARTED"}
62-
63-
elif monitor_type == "model_quality":
64-
job_name = f"{endpoint_name}-model-quality-baseline"
65-
params = ModelQualityParams(**event.get("model_quality_params", {}))
66-
67-
monitor = ModelQualityMonitor(
68-
role=SAGEMAKER_ROLE_ARN,
69-
instance_count=params.instance_count,
70-
instance_type=params.instance_type,
71-
volume_size_in_gb=params.volume_size_gb,
72-
max_runtime_in_seconds=params.max_runtime_seconds,
73-
)
74-
75-
monitor.suggest_baseline(
76-
job_name=job_name,
77-
baseline_dataset=training_data_uri,
78-
dataset_format=DatasetFormat.csv(header=True),
79-
output_s3_uri=baseline_output_uri,
80-
problem_type=params.problem_type,
81-
inference_attribute=params.inference_attribute,
82-
probability_attribute=params.probability_attribute,
83-
ground_truth_attribute=params.ground_truth_attribute,
84-
wait=False,
85-
logs=False,
86-
)
87-
88-
return {"statusCode": 200, "job_name": job_name, "status": "STARTED"}
89-
90-
elif monitor_type == "model_bias":
91-
job_name = f"{endpoint_name}-model-bias-baseline"
92-
params = ModelBiasParams(**event.get("model_bias_params", {}))
93-
94-
monitor = ModelBiasMonitor(
95-
role=SAGEMAKER_ROLE_ARN,
96-
max_runtime_in_seconds=params.max_runtime_seconds,
97-
)
98-
99-
model_bias_data_config = DataConfig(
100-
s3_data_input_path=training_data_uri,
101-
s3_output_path=baseline_output_uri,
102-
label=params.label_header,
103-
headers=params.get_headers(),
104-
dataset_type=params.dataset_type,
105-
)
106-
107-
model_bias_config = BiasConfig(
108-
label_values_or_threshold=params.get_label_values(),
109-
facet_name=params.facet_name,
110-
facet_values_or_threshold=params.get_facet_values(),
111-
)
112-
113-
model_predicted_label_config = ModelPredictedLabelConfig(
114-
probability_threshold=params.probability_threshold,
115-
)
116-
117-
model_config = ModelConfig(
118-
model_name=params.model_name or endpoint_name,
119-
instance_count=params.instance_count,
120-
instance_type=params.instance_type,
121-
content_type=params.dataset_type,
122-
accept_type=params.dataset_type,
123-
)
124-
125-
monitor.suggest_baseline(
126-
model_config=model_config,
127-
data_config=model_bias_data_config,
128-
bias_config=model_bias_config,
129-
model_predicted_label_config=model_predicted_label_config,
130-
wait=False,
131-
logs=False,
132-
)
133-
134-
return {"statusCode": 200, "job_name": job_name, "status": "STARTED"}
135-
136-
elif monitor_type == "model_explainability":
137-
job_name = f"{endpoint_name}-model-explainability-baseline"
138-
params = ModelExplainabilityParams(**event.get("model_explainability_params", {}))
139-
140-
monitor = ModelExplainabilityMonitor(
141-
role=SAGEMAKER_ROLE_ARN,
142-
max_runtime_in_seconds=params.max_runtime_seconds,
143-
)
144-
145-
model_explainability_data_config = DataConfig(
146-
s3_data_input_path=training_data_uri,
147-
s3_output_path=baseline_output_uri,
148-
label=params.label_header,
149-
headers=params.get_headers(),
150-
dataset_type=params.dataset_type,
151-
)
152-
153-
# Use configurable SHAP baseline
154-
shap_config = SHAPConfig(
155-
baseline=params.shap_baseline,
156-
num_samples=params.num_samples,
157-
agg_method=params.agg_method,
158-
save_local_shap_values=params.save_local_shap_values,
159-
)
160-
161-
model_config = ModelConfig(
162-
model_name=params.model_name or endpoint_name,
163-
instance_count=params.instance_count,
164-
instance_type=params.instance_type,
165-
content_type=params.dataset_type,
166-
accept_type=params.dataset_type,
167-
)
168-
169-
monitor.suggest_baseline(
170-
data_config=model_explainability_data_config,
171-
model_config=model_config,
172-
explainability_config=shap_config,
173-
)
174-
175-
return {"statusCode": 200, "job_name": job_name, "status": "STARTED"}
176-
177-
return {"statusCode": 400, "error": f"Unsupported monitor_type: {monitor_type}"}
183+
except Exception as e:
184+
return {"statusCode": 500, "error": f"Failed to start baseline job: {str(e)}"}
178185

179186

180187
def check_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
181188
"""Check the status of a baselining processing job."""
182-
183-
job_name = event["job_name"]
184-
185-
response = sagemaker.describe_processing_job(ProcessingJobName=job_name)
186-
status = response["ProcessingJobStatus"]
187-
188-
return {"statusCode": 200, "job_name": job_name, "status": "COMPLETED" if status == "Completed" else "IN_PROGRESS"}
189+
try:
190+
job_name = event["job_name"]
191+
192+
response = sagemaker.describe_processing_job(ProcessingJobName=job_name)
193+
status = response["ProcessingJobStatus"]
194+
195+
return {
196+
"statusCode": 200,
197+
"job_name": job_name,
198+
"status": "COMPLETED" if status == "Completed" else "IN_PROGRESS",
199+
}
200+
except KeyError as e:
201+
return {"statusCode": 400, "error": f"Missing required parameter: {e}"}
202+
except sagemaker.exceptions.ResourceNotFound:
203+
return {"statusCode": 404, "error": f"Processing job not found: {event.get('job_name', 'unknown')}"}
204+
except Exception as e:
205+
return {"statusCode": 500, "error": f"Failed to check job status: {str(e)}"}

0 commit comments

Comments
 (0)