如何在SageMaker端点的inference.py中获取CustomAttributes参数?
在SageMaker推理脚本中获取CustomAttributes参数
完全可行,你可以通过context参数在transform_fn中获取CustomAttributes,具体实现步骤如下:
1. 修改transform_fn的函数签名
SageMaker的推理框架支持为transform_fn传入第五个参数context,这个参数包含了请求的元数据,其中就包括你传递的CustomAttributes。
2. 从context中提取CustomAttributes
context.request_attributes是一个字典,直接通过键CustomAttributes就能拿到对应的值。如果你的CustomAttributes是逗号分隔的键值对格式(比如key1=value1,key2=value2),可以进一步解析成字典方便使用。
修改后的完整inference.py示例
import json import torch def get_device(): device = 'cuda:0' if torch.cuda.is_available() else 'cpu' return device def model_fn(model_dir): # 加载你的模型(保留原有逻辑) device = get_device() model = ... # 替换为你的模型加载代码 model.to(device) return model def transform_fn(model, request_body, content_type, accept, context): # 获取CustomAttributes参数 custom_attrs = context.request_attributes.get('CustomAttributes') # 解析键值对(可选,根据你的参数格式调整) attr_dict = {} if custom_attrs: for item in custom_attrs.split(','): if '=' in item: key, value = item.split('=', 1) attr_dict[key.strip()] = value.strip() # 原有推理逻辑 input_data = json.loads(request_body) # ... 你的模型推理代码 ... predictions = model(input_data) # 示例:根据CustomAttributes调整返回结果 if attr_dict.get('return_full_data') == 'yes': response = {"predictions": predictions.tolist(), "received_attributes": attr_dict} else: response = {"predictions": predictions.tolist()} return json.dumps(response), accept
调用端点时传递CustomAttributes
用boto3调用端点时,直接在invoke_endpoint中指定CustomAttributes参数即可:
import boto3 import json client = boto3.client('sagemaker-runtime') response = client.invoke_endpoint( EndpointName='your-endpoint-name', ContentType='application/json', Body=json.dumps({"input": "your-data"}), CustomAttributes="return_full_data=yes,request_id=abc123" ) # 解析返回结果 result = json.loads(response['Body'].read().decode()) print(result)
注意:context参数是SageMaker推理容器自动注入的,不需要你手动传递,只要在transform_fn的函数签名中声明就能使用。
内容的提问来源于stack exchange,提问作者Abdeelrahman Hassanein
相关产品推荐
相关产品推荐

