如何让超参数在Amazon SageMaker TensorFlow推理端点中可用
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

