Complete AWS SageMaker Model Monitor Tutorial: ML Model Monitoring in Production
Amazon SageMaker Model Monitor automatically detects data quality issues, model quality degradation, bias drift, and feature attribution drift in ML models deployed to production. It helps maintain model performance over time.
Why Model Monitor?
Key Benefits:- Automated monitoring: Continuous model surveillance
- Drift detection: Data and model quality drift alerts
- Bias detection: Monitor fairness metrics
- Explainability: Feature attribution tracking
- Integration: Native SageMaker integration
- Data Quality Monitor
- Model Quality Monitor
- Bias Drift Monitor
- Feature Attribution Drift Monitor
Prerequisites
pip install sagemaker boto3 pandas numpy
SageMaker SDK >= 2.0
python -c "import sagemaker; print(sagemaker.version)"
Quick Start
1. Setup
import boto3
import sagemaker
from sagemaker import getexecutionrole
from sagemaker.modelmonitor import (
DefaultModelMonitor,
DataCaptureConfig,
CronExpressionGenerator
)
session = sagemaker.Session()
bucket = session.defaultbucket()
role = getexecutionrole()
region = session.botoregionname
Monitor output location
monitoroutput = f"s3://{bucket}/model-monitor"
2. Deploy Model with Data Capture
from sagemaker.model import Model
from sagemaker.predictor import Predictor
Create model
model = Model(
imageuri=xgboostimage,
modeldata=modeldatauri,
role=role
)
Data capture configuration
datacaptureconfig = DataCaptureConfig(
enablecapture=True,
samplingpercentage=100, # Capture all requests
destinations3uri=f"s3://{bucket}/data-capture",
captureoptions=["Input", "Output"],
csvcontenttypes=["text/csv"],
jsoncontenttypes=["application/json"]
)
Deploy with data capture
predictor = model.deploy(
initialinstancecount=1,
instancetype="ml.m5.large",
endpointname="monitored-endpoint",
datacaptureconfig=datacaptureconfig
)
print(f"Endpoint deployed: {predictor.endpointname}")
Data Quality Monitor
1. Create Baseline
from sagemaker.modelmonitor import DefaultModelMonitor
from sagemaker.model
monitor.datasetformat import DatasetFormat
Create monitor
data
qualitymonitor = DefaultModelMonitor(
role=role,
instance
count=1,
instancetype="ml.m5.xlarge",
volumesizeingb=20,
maxruntimeinseconds=3600
)
Create baseline from training data
dataqualitymonitor.suggestbaseline(
baselinedataset=f"s3://{bucket}/training-data/train.csv",
datasetformat=DatasetFormat.csv(header=True),
outputs3uri=f"{monitoroutput}/data-quality/baseline",
wait=True
)
print("Baseline created!")
2. View Baseline Statistics
import json
Get baseline statistics
baselinejob = dataqualitymonitor.latestbaseliningjob
statisticspath = f"{monitoroutput}/data-quality/baseline/statistics.json"
constraintspath = f"{monitoroutput}/data-quality/baseline/constraints.json"
Download and view statistics
s3 = boto3.client("s3")
Parse S3 URI
def parses3uri(uri):
parts = uri.replace("s3://", "").split("/", 1)
return parts[0], parts[1]
bucketname, key = parses3uri(statisticspath)
response = s3.getobject(Bucket=bucketname, Key=key)
statistics = json.loads(response["Body"].read())
print("Baseline Statistics:")
for feature in statistics["features"]:
print(f" {feature['name']}: mean={feature.get('numericalstatistics', {}).get('mean', 'N/A')}")
3. Schedule Monitoring Job
from sagemaker.modelmonitor import CronExpressionGenerator
Create monitoring schedule
dataqualitymonitor.createmonitoringschedule(
monitorschedulename="data-quality-schedule",
endpointinput=predictor.endpointname,
outputs3uri=f"{monitoroutput}/data-quality/reports",
statistics=dataqualitymonitor.baselinestatistics(),
constraints=dataqualitymonitor.suggestedconstraints(),
schedulecronexpression=CronExpressionGenerator.hourly(),
enablecloudwatchmetrics=True
)
print("Monitoring schedule created!")
Model Quality Monitor
1. Setup Ground Truth
from sagemaker.modelmonitor import ModelQualityMonitor
Create model quality monitor
modelqualitymonitor = ModelQualityMonitor(
role=role,
instancecount=1,
instancetype="ml.m5.xlarge",
volumesizeingb=20,
maxruntimeinseconds=3600,
sagemakersession=session
)
Create baseline with ground truth
modelqualitymonitor.suggestbaseline(
baselinedataset=f"s3://{bucket}/ground-truth/baseline.csv",
datasetformat=DatasetFormat.csv(header=True),
outputs3uri=f"{monitoroutput}/model-quality/baseline",
problemtype="BinaryClassification",
inferenceattribute="prediction",
groundtruthattribute="label",
wait=True
)
2. Ground Truth Format
import pandas as pd
from datetime import datetime
Ground truth data format
groundtruthdata = pd.DataFrame({
"inferenceid": ["id-001", "id-002", "id-003"],
"prediction": [1, 0, 1],
"label": [1, 0, 0], # Actual ground truth
"timestamp": [datetime.now().isoformat()] 3
})
Save ground truth
groundtruthdata.tocsv(
f"s3://{bucket}/ground-truth/latest.csv",
index=False
)
3. Schedule Model Quality Monitor
from sagemaker.modelmonitor import EndpointInput
Create endpoint input with ground truth
endpointinput = EndpointInput(
endpointname=predictor.endpointname,
destination="/opt/ml/processing/input/endpoint",
inferenceattribute="prediction"
)
Ground truth input
groundtruthinput = f"s3://{bucket}/ground-truth/"
Create schedule
modelqualitymonitor.createmonitoringschedule(
monitorschedulename="model-quality-schedule",
endpointinput=endpointinput,
groundtruthinput=groundtruthinput,
outputs3uri=f"{monitoroutput}/model-quality/reports",
problemtype="BinaryClassification",
constraints=modelqualitymonitor.suggestedconstraints(),
schedulecronexpression=CronExpressionGenerator.daily()
)
Bias Drift Monitor
1. Create Clarify Monitor
from sagemaker.clarify import (
ModelConfig,
BiasConfig,
DataConfig,
SHAPConfig
)
from sagemaker.modelmonitor import ClarifyModelMonitor
Create Clarify monitor
clarifymonitor = ClarifyModelMonitor(
role=role,
instancecount=1,
instancetype="ml.m5.xlarge",
volumesizeingb=20,
maxruntimeinseconds=3600,
sagemakersession=session
)
Bias configuration
biasconfig = BiasConfig(
labelvaluesorthreshold=[1],
facetname="gender",
facetvaluesorthreshold=[0], # Monitor bias against gender=0
groupname="age"
)
Data configuration
dataconfig = DataConfig(
s3datainputpath=f"s3://{bucket}/training-data/",
s3outputpath=f"{monitoroutput}/bias/baseline",
label="target",
headers=["feature1", "feature2", "gender", "age", "target"],
datasettype="text/csv"
)
Model configuration
modelconfig = ModelConfig(
modelname="my-model",
instancetype="ml.m5.large",
instancecount=1,
accepttype="text/csv"
)
Create baseline
clarifymonitor.suggestbaseline(
dataconfig=dataconfig,
biasconfig=biasconfig,
modelconfig=modelconfig,
wait=True
)
2. Schedule Bias Monitor
# Create bias monitoring schedule
clarifymonitor.createmonitoringschedule(
monitorschedulename="bias-drift-schedule",
endpointinput=predictor.endpointname,
outputs3uri=f"{monitoroutput}/bias/reports",
constraints=clarifymonitor.suggestedconstraints(),
schedulecronexpression=CronExpressionGenerator.daily()
)
Feature Attribution Monitor
1. SHAP Baseline
# SHAP configuration
shapconfig = SHAPConfig(
baseline=[
[0.0] 10 # Baseline values for each feature
],
numsamples=100,
aggmethod="meanabs"
)
Create feature attribution baseline
clarifymonitor.suggestbaseline(
dataconfig=dataconfig,
modelconfig=modelconfig,
explainabilityconfig=shapconfig,
wait=True
)
2. Schedule Attribution Monitor
# Create explainability monitoring schedule
clarifymonitor.createmonitoringschedule(
monitorschedulename="feature-attribution-schedule",
endpointinput=predictor.endpointname,
outputs3uri=f"{monitoroutput}/explainability/reports",
constraints=clarifymonitor.suggestedconstraints(),
schedulecronexpression=CronExpressionGenerator.daily(),
analysistype="explainability"
)
Custom Monitoring
1. Custom Processing Script
# custommonitor.py
import json
import pandas as pd
import os
def preprocesshandler(inputdata):
"""Preprocess captured data."""
# Custom preprocessing logic
return inputdata
def evaluatehandler(baseline, current):
"""Evaluate current data against baseline."""
violations = []
# Check for missing values
if current.isnull().any().any():
violations.append({
"feature": "all",
"violationtype": "missingvalues",
"description": "Missing values detected"
})
# Check for distribution shift
for column in current.columns:
if column in baseline.columns:
baselinemean = baseline[column].mean()
currentmean = current[column].mean()
if abs(currentmean - baselinemean) > baselinemean * 0.2:
violations.append({
"feature": column,
"violationtype": "distributionshift",
"description": f"Mean shifted from {baselinemean} to {currentmean}"
})
return violations
if name == "main":
# Load data
inputpath = "/opt/ml/processing/input"
outputpath = "/opt/ml/processing/output"
# Process and evaluate
# ... implementation
2. Bring Your Own Container
from sagemaker.modelmonitor import ModelMonitor
Custom monitor with BYOC
custommonitor = ModelMonitor(
role=role,
imageuri="your-account.dkr.ecr.region.amazonaws.com/custom-monitor:latest",
instancecount=1,
instancetype="ml.m5.xlarge",
volumesizeingb=20,
maxruntimeinseconds=3600
)
Create schedule with custom monitor
custommonitor.createmonitoringschedule(
monitorschedulename="custom-monitor-schedule",
endpointinput=predictor.endpointname,
outputs3uri=f"{monitoroutput}/custom/reports",
schedulecronexpression=CronExpressionGenerator.hourly()
)
CloudWatch Integration
1. View Metrics
import boto3
from datetime import datetime, timedelta
cloudwatch = boto3.client("cloudwatch")
Get metrics
response = cloudwatch.getmetricstatistics(
Namespace="aws/sagemaker/Endpoints/data-metrics",
MetricName="featurebaselinedrift",
Dimensions=[
{"Name": "Endpoint", "Value": predictor.endpointname},
{"Name": "MonitoringSchedule", "Value": "data-quality-schedule"}
],
StartTime=datetime.utcnow() - timedelta(hours=24),
EndTime=datetime.utcnow(),
Period=3600,
Statistics=["Average"]
)
for datapoint in response["Datapoints"]:
print(f"{datapoint['Timestamp']}: {datapoint['Average']}")
2. Create Alarms
# Create CloudWatch alarm for drift
cloudwatch.putmetricalarm(
AlarmName="ModelDriftAlarm",
MetricName="featurebaselinedrift",
Namespace="aws/sagemaker/Endpoints/data-metrics",
Dimensions=[
{"Name": "Endpoint", "Value": predictor.endpointname}
],
Statistic="Average",
Period=3600,
EvaluationPeriods=2,
Threshold=0.5,
ComparisonOperator="GreaterThanThreshold",
AlarmActions=[
"arn:aws:sns:region:account:alert-topic"
]
)
Analyzing Results
1. Get Monitoring Results
# List monitoring executions
executions = dataqualitymonitor.listexecutions()
for execution in executions:
print(f"Execution: {execution.processingjobname}")
print(f" Status: {execution.status}")
print(f" Start: {execution.creationtime}")
Get latest execution
latest = executions[0]
print(f"\nLatest execution violations:")
Get violation report
if latest.constraintviolationss3uri:
violations = latest.constraintviolations()
for violation in violations.violations:
print(f" {violation.featurename}: {violation.constraintchecktype}")
2. Visualize Drift
import matplotlib.pyplot as plt
def plotdriftovertime(executions):
"""Plot drift metrics over time."""
timestamps = []
driftscores = []
for execution in executions:
if execution.statisticss3uri:
stats = execution.statistics()
timestamps.append(execution.creationtime)
# Extract drift score from stats
driftscores.append(stats.get("driftscore", 0))
plt.figure(figsize=(12, 6))
plt.plot(timestamps, driftscores, marker='o')
plt.xlabel("Time")
plt.ylabel("Drift Score")
plt.title("Model Drift Over Time")
plt.xticks(rotation=45)
plt.tightlayout()
plt.savefig("driftanalysis.png")
plotdriftovertime(executions)
Best Practices
1. Monitoring Strategy
# Comprehensive monitoring setup
def setupcomprehensivemonitoring(endpointname, baselinedatauri):
"""Setup all monitor types for an endpoint."""
monitors = {}
# 1. Data Quality Monitor
monitors["dataquality"] = setupdataqualitymonitor(
endpointname, baselinedatauri
)
# 2. Model Quality Monitor (if ground truth available)
monitors["modelquality"] = setupmodelqualitymonitor(
endpointname
)
# 3. Bias Monitor
monitors["bias"] = setupbiasmonitor(endpointname)
return monitors
2. Alert Configuration
def configuremonitoringalerts(schedulename, snstopicarn):
"""Configure alerts for monitoring violations."""
cloudwatch = boto3.client("cloudwatch")
# Data quality alert
cloudwatch.put
metricalarm(
AlarmName=f"{schedule
name}-data-quality-alert",
MetricName="violationscount",
Namespace="aws/sagemaker/Endpoints/data-metrics",
Threshold=1,
ComparisonOperator="GreaterThanOrEqualToThreshold",
AlarmActions=[snstopic_arn]
)
Conclusion
SageMaker Model Monitor provides:
Key takeaways:
- Create baselines from training data
- Schedule regular monitoring jobs
- Configure CloudWatch alarms
- Analyze violations promptly
- Retrain models when drift detected