Skip to content

Commit ac2bd7d

Browse files
committed
update handler; package dependencies
1 parent 562b73f commit ac2bd7d

4 files changed

Lines changed: 11 additions & 12 deletions

File tree

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

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,12 @@ def __init__(
4747
command=[
4848
"bash",
4949
"-c",
50-
"pip install -r requirements.txt -t /asset-output && cp -au . /asset-output",
50+
"pip install pydantic==2.10.3 -t /asset-output --no-cache-dir && "
51+
"pip install sagemaker==2.232.2 -t /asset-output --no-deps --no-cache-dir && "
52+
"rm -rf /asset-output/{boto*,urllib3*,certifi*,six*,python_dateutil*,jmespath*,s3transfer*} && "
53+
"find /asset-output \\( -name '*.pyc' -o -name '*.so' \\) -delete && "
54+
"find /asset-output -type d \\( -name '__pycache__' -o -name 'test*' \\) -exec rm -rf {} + && "
55+
"cp -au . /asset-output",
5156
],
5257
),
5358
),
@@ -87,8 +92,8 @@ def __init__(
8792
"s3:PutObject",
8893
],
8994
resources=[
90-
f"{baseline_training_data_s3_uri}*",
91-
f"{baseline_output_data_s3_uri}*",
95+
f"arn:aws:s3:::{baseline_training_data_s3_uri.replace('s3://', '')}*",
96+
f"arn:aws:s3:::{baseline_output_data_s3_uri.replace('s3://', '')}*",
9297
],
9398
)
9499
)

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

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
from typing import Any, Dict
33

44
import boto3
5-
import pandas as pd
65
from models import DataQualityParams, ModelBiasParams, ModelExplainabilityParams, ModelQualityParams
76
from sagemaker.clarify import BiasConfig, DataConfig, ModelConfig, ModelPredictedLabelConfig, SHAPConfig
87
from sagemaker.model_monitor import (
@@ -151,12 +150,9 @@ def start_baselining_job(event: Dict[str, Any]) -> Dict[str, Any]:
151150
dataset_type=params.dataset_type,
152151
)
153152

154-
# Use mean value of test dataset as SHAP baseline
155-
test_dataframe = pd.read_csv(training_data_uri, header=None)
156-
shap_baseline = [list(test_dataframe.mean())]
157-
153+
# Use configurable SHAP baseline
158154
shap_config = SHAPConfig(
159-
baseline=shap_baseline,
155+
baseline=params.shap_baseline,
160156
num_samples=params.num_samples,
161157
agg_method=params.agg_method,
162158
save_local_shap_values=params.save_local_shap_values,

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

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,7 @@ class ModelExplainabilityParams(BaseParams):
4949
num_samples: int = Field(default=100, ge=1, description="Number of samples for SHAP")
5050
agg_method: str = Field(default="mean_abs", description="SHAP aggregation method")
5151
save_local_shap_values: bool = Field(default=False, description="Save local SHAP values")
52+
shap_baseline: List[List[float]] = Field(default=[[0.0]], description="SHAP baseline values")
5253

5354
def get_headers(self) -> Optional[List[str]]:
5455
return [h.strip() for h in self.headers.split(",")] if self.headers else None

modules/sagemaker/sagemaker-model-monitoring/sagemaker_model_monitoring/lambda/requirements.txt

Lines changed: 0 additions & 3 deletions
This file was deleted.

0 commit comments

Comments
 (0)