You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让超参数在Amazon SageMaker TensorFlow推理端点中可用

Solution: Making Training Hyperparameters Accessible in SageMaker Inference Endpoints

Got it, let's tackle this problem—getting your training hyperparameters accessible in a SageMaker inference endpoint is a common need, and there are a few reliable approaches to make this work:

1. Save Hyperparameters with the Model During Training

This is the most straightforward method: persist your hyperparameters alongside your trained model files, then load them when the inference endpoint starts up.

Step 1: Modify your training script (autocat.py) to save hyperparameters

SageMaker injects your hyperparameters into the training container as a JSON string in the SM_HPS environment variable. You can capture this and save it to the model directory:

import json
import os

def train():
    # Fetch hyperparameters from the environment
    hyperparams = json.loads(os.environ.get('SM_HPS', '{}'))
    
    # ... your existing training logic here ...
    
    # Save hyperparameters to the model directory (will be packaged with the model)
    model_dir = os.environ.get('SM_MODEL_DIR')
    with open(os.path.join(model_dir, 'hyperparameters.json'), 'w') as f:
        json.dump(hyperparams, f)

if __name__ == '__main__':
    train()

Step 2: Load hyperparameters in your inference script

When deploying, your inference script (either the same autocat.py if you're using a single script for train/infer, or a dedicated inference.py) can load the saved hyperparameters:

import json
import os

def model_fn(model_dir):
    # Load the saved hyperparameters
    with open(os.path.join(model_dir, 'hyperparameters.json'), 'r') as f:
        hyperparams = json.load(f)
    
    # Load your trained model
    model = load_your_trained_model(model_dir)
    
    # Use hyperparameters to configure the model for inference (e.g., set thresholds, scaling factors)
    model.configure(hyperparams)
    
    return model

# ... implement input_fn, predict_fn, output_fn as needed ...

2. Inject Hyperparameters via Environment Variables at Deployment

If you don't want to modify the training script, you can pass the hyperparameters directly to the inference container using environment variables when creating the SageMaker Model object.

Step 1: Create the Model with Hyperparameters in Environment

import json
from sagemaker.tensorflow import TensorFlowModel

# Assume you still have access to the original `params` dictionary used for training
trained_hyperparams = params

# Create the Model, injecting hyperparameters as an environment variable
model = TensorFlowModel(
    model_data=estimator.model_data,
    role=role,
    entry_point='inference.py',
    environment={
        'TRAINING_HYPERPARAMS': json.dumps(trained_hyperparams)
    }
)

# Deploy the endpoint
predictor = model.deploy(
    initial_instance_count=1,
    instance_type='ml.c4.xlarge'
)

Step 2: Read Environment Variables in Inference Script

In your inference script, pull the hyperparameters from the environment:

import json
import os

def model_fn(model_dir):
    # Load hyperparameters from environment variable
    hyperparams = json.loads(os.environ.get('TRAINING_HYPERPARAMS', '{}'))
    
    # Load and configure your model
    model = load_your_trained_model(model_dir)
    model.set_hyperparameters(hyperparams)
    
    return model

3. Use SageMaker Model Registry (For Managed Model Versions)

If you're using SageMaker Model Registry to track model versions, you can attach hyperparameters as metadata to the model version during training, then retrieve them at deployment.

Step 1: Associate Hyperparameters with Model Registry

After training, when registering your model, add the hyperparameters as properties:

model_package = estimator.register(
    model_package_group_name="your-model-group",
    properties={
        "hyperparameters": json.dumps(params)
    }
)

Step 2: Retrieve Hyperparameters When Deploying

When deploying from the Model Registry, fetch the metadata and pass it to the inference container (similar to method 2, but pulling from the registry instead of the original params):

from sagemaker.model import ModelPackage

model_package_arn = "your-model-package-arn"
model = ModelPackage(
    role=role,
    model_package_arn=model_package_arn,
    environment={
        'TRAINING_HYPERPARAMS': model_package.properties["hyperparameters"]
    }
)

model.deploy(initial_instance_count=1, instance_type='ml.c4.xlarge')

内容的提问来源于stack exchange,提问作者tlanigan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.20 12:35:00