AWS SageMaker是否支持gRPC预测请求?本地TFS部署遇连接问题
问题描述
我通过SageMaker的estimator在本地模式部署了TensorFlow模型,尝试通过gRPC调用TensorFlow Serving(TFS)预测端点时,出现以下错误:
grpc._channel._InactiveRpcError: <_InactiveRpcError of RPC that terminated with: status = StatusCode.UNAVAILABLE details = "failed to connect to all addresses"
我的gRPC请求代码如下:
import grpc from tensorflow.compat.v1 import make_tensor_proto from tensorflow_serving.apis import predict_pb2 from tensorflow_serving.apis import prediction_service_pb2_grpc grpc_port = 9000 # 也尝试过8500等其他端口 request = predict_pb2.PredictRequest() request.model_spec.name = 'model' request.model_spec.signature_name = 'serving_default' request.inputs['input_tensor'].CopyFrom(make_tensor_proto(instance)) options = [ ('grpc.enable_http_proxy', 0), ('grpc.max_send_message_length', MAX_GRPC_MESSAGE_LENGTH), ('grpc.max_receive_message_length', MAX_GRPC_MESSAGE_LENGTH) ] channel = grpc.insecure_channel(f'0.0.0.0:{grpc_port}', options=options) stub = prediction_service_pb2_grpc.PredictionServiceStub(channel) result_future = stub.Predict.future(request, 30) output_tensor_proto = result_future.result().outputs['predictions'] output_shape = [dim.size for dim in output_tensor_proto.tensor_shape.dim] output_np = np.array(output_tensor_proto.float_val).reshape(output_shape) prediction_json = {'predictions': output_np.tolist()}
查看运行TFS的SageMaker Docker容器日志,发现REST端点已暴露,但gRPC端点虽显示运行却无法连接:
tensorflow_serving/model_servers/server.cc:417] Running gRPC ModelServer at 0.0.0.0:9000 ...
REST端点日志:
tensorflow_serving/model_servers/server.cc:438] Exporting HTTP/REST API at:localhost:8501 ...
请问SageMaker TFS容器是否支持gRPC?如何在SageMaker中发起gRPC TFS预测请求?
解决方案
1. SageMaker TFS容器是否支持gRPC?
是,SageMaker官方TensorFlow Serving容器完全支持gRPC协议。本地模式下出现连接问题,核心原因是默认情况下SageMaker本地模式只会映射REST端口(8501),不会自动暴露gRPC的9000端口。
2. 本地模式下修复gRPC连接问题
(1)部署时显式映射gRPC端口
使用SageMaker estimator部署本地端点时,必须通过container_port参数指定容器内的gRPC端口,让本地端口与容器端口做映射:
from sagemaker.tensorflow.serving import Model model = Model(model_data='path/to/your/model', role='SageMakerExecutionRole', framework_version='2.8') predictor = model.deploy( initial_instance_count=1, instance_type='local', container_port=9000, # 指定容器内的gRPC端口 endpoint_name='local-tfs-grpc-endpoint' )
如果使用LocalSession,需同步配置:
from sagemaker.local import LocalSession sagemaker_session = LocalSession() sagemaker_session.config = {'local': {'local_code': True}} predictor = model.deploy( initial_instance_count=1, instance_type='local', container_port=9000, sagemaker_session=sagemaker_session )
(2)验证端口映射是否生效
部署完成后,执行docker ps命令查看容器信息,确认存在0.0.0.0:9000->9000/tcp的端口映射项:
docker ps
如果未看到该映射,说明部署时未正确配置container_port,需重新部署。
(3)调整请求代码细节
- 确保
MAX_GRPC_MESSAGE_LENGTH已定义,例如设置为1024*1024*100(100MB) - 确认
request.model_spec.name与部署时的模型名称一致(默认是model,若自定义过需对应修改) - 宿主机测试时用
0.0.0.0:9000,容器内部测试用localhost:9000
3. 云端SageMaker端点的gRPC调用方式
部署到云端SageMaker端点时,无需手动映射端口,直接通过官方域名调用即可:
- 云端gRPC端口默认是
443,需使用grpc.secure_channel而非insecure_channel - 端点地址格式:
{endpoint-name}.{region}.sagemaker.aws:443 - 需要配置AWS SigV4签名认证,可通过
grpcio-auth库实现:
import grpc from grpc_auth import SigV4Auth # 替换为你的端点名称和区域 endpoint_name = 'your-cloud-endpoint' region = 'us-east-1' channel = grpc.secure_channel(f'{endpoint_name}.{region}.sagemaker.aws:443', grpc.ssl_channel_credentials()) auth = SigV4Auth(region=region, service='sagemaker') stub = prediction_service_pb2_grpc.PredictionServiceStub(channel) # 发起请求时携带认证元数据 result_future = stub.Predict.future(request, 30, metadata=auth.metadata())
内容的提问来源于stack exchange,提问作者theeDude
相关产品推荐
相关产品推荐

