关于SageMaker TensorFlow自定义模型JSON调用预测端点的技术问询
Hey there! Let’s walk through how to get your JavaScript JSON prediction requests playing nicely with your custom SageMaker TensorFlow Estimator. I’ll break down the key details based on your setup:
Key Context from Your Setup
First, let’s recap your serving_input_fn code since it dictates how your model accepts input:
def serving_input_fn(params): tensor = tf.placeholder(tf.float32, shape=[1, NUM_FEATURES]) return build_raw_serving_input_receiver_fn({INPUT_TENSOR_NAME: tensor})()
This tells SageMaker your model expects a 2-dimensional tensor with shape [1, NUM_FEATURES]—meaning a single batch (batch size 1) of NUM_FEATURES features.
JSON Request Format Requirements
To match your model’s input shape, your JSON payload needs to mirror this tensor structure exactly:
- The top-level key must match your
INPUT_TENSOR_NAME(e.g., if it’s"inputs", use that as the key) - The value must be a 2-dimensional array (even for a single sample)
Example payload (assuming NUM_FEATURES = 8 and INPUT_TENSOR_NAME = "my_features"):
{"my_features": [[1.1, 2.2, 3.3, 4.4, 5.5, 6.6, 7.7, 8.8]]}
If you try sending a 1-dimensional array like [1.1, 2.2, ...], you’ll get a shape mismatch error—your model is explicitly expecting that extra layer of nesting for the batch dimension.
JavaScript Implementation Tips
Here’s a clean, working example using the AWS SDK for JavaScript v3 (the latest recommended version):
import { SageMakerRuntimeClient, InvokeEndpointCommand } from "@aws-sdk/client-sagemaker-runtime"; // Initialize the SageMaker Runtime client with your region const sageMakerClient = new SageMakerRuntimeClient({ region: "us-east-1" }); async function runPrediction() { const endpointName = "your-endpoint-name-here"; const inputFeatures = [1.1, 2.2, 3.3, 4.4, 5.5, 6.6, 7.7, 8.8]; // Format payload to match model's expected tensor shape const payload = { "INPUT_TENSOR_NAME": [inputFeatures] // Wrap single sample in an array for batch dimension }; try { const command = new InvokeEndpointCommand({ EndpointName: endpointName, ContentType: "application/json", Body: JSON.stringify(payload) }); const response = await sageMakerClient.send(command); // Decode the Uint8Array response body to JSON const predictionResult = JSON.parse(new TextDecoder().decode(response.Body)); console.log("Prediction output:", predictionResult.predictions); } catch (error) { console.error("Error calling endpoint:", error); } } // Run the prediction runPrediction();
Common Fixes for Issues You Might Hit
- Shape Mismatch Errors: Double-check that your JSON payload is a 2D array. If you want to support batch predictions (multiple samples at once), update your
serving_input_fntensor shape to[None, NUM_FEATURES]—then you can send a payload like{"INPUT_TENSOR_NAME": [[sample1], [sample2], ...]}. - Mismatched Tensor Name: Ensure the key in your JSON matches exactly with
INPUT_TENSOR_NAMEin your Python code (case-sensitive!). - Response Parsing: SageMaker returns the response body as a
Uint8Array, so you need to decode it withTextDecoderbefore parsing as JSON.
内容的提问来源于stack exchange,提问作者JonL

