本文属于机器翻译版本。若本译文内容与英语原文存在差异,则一律以英文原文为准。
培训作业提交
启动培训作业
部署代理且数据集在 S3 中后,使用以下方法之一创建训练作业。
SageMaker 人工智能工作室
-
在导航窗格中导航到模型,然后选择 JumpStart 基础模型。
-
选择支持多圈 RL 的型号(参见支持的型号表),然后选择 “自定义模型”,然后选择 “使用用户界面自定义”。
-
选择 Multi-Turn 强化学习作为自定义技术。
-
配置您的代理环境 — 选择您的 Bedrock AgentCore 运行时或提供您的 Lambda 转发器 ARN。
-
以 S3 URI 或注册数据集的形式提供您的训练数据集。
-
根据需要调整超参数。
-
查看您的配置并选择提交。
SageMaker AI Python SD
探索支持的型号
from sagemaker.train.multi_turn_rl_trainer import MultiTurnRLTrainer supported_models = MultiTurnRLTrainer.list_supported_models() print(f"Supported MTRL models ({len(supported_models)}):") for m in supported_models: print(f" - {m}")
设置您的代理环境
选项 1:基岩运行 AgentCore 时
# List available runtimes runtimes = MultiTurnRLTrainer.list_bedrock_agentcore_runtimes() for rt in runtimes: print(f" - {rt['name']} ({rt['status']}) → {rt['arn']}")
选项 2:自定义 Lambda 代理
from sagemaker.train.agent_lambda import AgentLambda # Create from inline code adapter = AgentLambda.create( source=''' import json def handler(event, context): prompt = event.get("prompt", "") return {"statusCode": 200, "body": json.dumps({"status": "ok", "agentResponse": prompt})} ''', role="arn:aws:iam::123456789012:role/AgentLambdaRole", ) # Create from a local file adapter = AgentLambda.create( source="~/my_agent_handler.py", role="arn:aws:iam::123456789012:role/AgentLambdaRole", ) # Create from S3 adapter = AgentLambda.create( source="s3://my-bucket/agent_handler.py", role="arn:aws:iam::123456789012:role/AgentLambdaRole", ) # Wrap an existing Lambda adapter = AgentLambda.get("arn:aws:lambda:us-west-2:123456789012:function:my-agent")
注册您的数据集(可选)
from sagemaker.ai_registry.dataset import DataSet dataset = DataSet.create( name="my-mtrl-dataset", source="s3://my-bucket/prompts/training_prompts.parquet" ) print(f"Dataset ARN: {dataset.arn}")
为 Nova 创建受限型号 Package 组(可选)
如果您选择 Nova 模型 (nova-textgeneration-lite-v2),则可以在提交训练作业之前创建受限模型包组(下一步)。如果您跳过此步骤,SDK 会自动为您创建一个。
受限型号封装组 (RMPG) 是一个模型包组,包含 ManagedStorageType:受限。对于像 Nova 这样的闭源模型,这是必需的,在这种模型中,模型权重由 Amazon 客户管理,客户无法直接访问。
RFT 作业架构需要两个单独的受限 MPG:
-
输出 MPG — 存储最终经过微调的模型包
-
中级检查点 MPG — 保留给中级训练检查点(必须与输出 MPG 不同)
from sagemaker.core.resources import Job, ModelPackageGroup from sagemaker.core.shapes import ManagedConfiguration model_name = "nova-textgeneration-lite-v2" # Restricted configuration managed_config = ManagedConfiguration(managed_storage_type="Restricted") # Output Model package group output_mpg_name = f"{model_name}-mtrl-output-mpg" create_kwargs = { "model_package_group_name": output_mpg_name, "region": "us-east-1", "managed_configuration": managed_config } output_mpg = ModelPackageGroup.create(**create_kwargs) # Intermediate Model package group intermediate_mpg_name = f"{model_name}-mtrl-inter-mpg" create_kwargs = { "model_package_group_name": intermediate_mpg_name, "region": "us-east-1", "managed_configuration": managed_config } intermediate_mpg = ModelPackageGroup.create(**create_kwargs)
创建模型包群组后,在下一步提交训练作业时传递群组。
使用 Bedrock 提交训练作业 AgentCore
from sagemaker.train.multi_turn_rl_trainer import MultiTurnRLTrainer trainer = MultiTurnRLTrainer( model="openai-reasoning-gpt-oss-20b", agent_env="arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/my-agent-runtime", training_dataset="s3://my-bucket/prompts/prompts.parquet", mlflow_app_arn="arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id", s3_output_path="s3://my-bucket/output/", role="arn:aws:iam::123456789012:role/SageMakerRole", accept_eula=True, ) # View and adjust hyperparameters trainer.hyperparameters.get_info() trainer.hyperparameters.max_epochs = 1 trainer.hyperparameters.global_batch_size = 32 trainer.hyperparameters.max_steps = 12 job = trainer.train(wait=True) print(f"Job: {job.job_name}") print(f"Status: {job.job_status}") print(f"Output Model Package: {job.output_model_package_arn}")
使用自定义 Lambda 代理提交训练作业
trainer = MultiTurnRLTrainer( model="openai-reasoning-gpt-oss-20b", agent_env=adapter, # AgentLambda object or Lambda ARN string training_dataset="s3://my-bucket/prompts/prompts.parquet", mlflow_app_arn="arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id", s3_output_path="s3://my-bucket/output/", role="arn:aws:iam::123456789012:role/SageMakerRole", accept_eula=True, ) trainer.hyperparameters.max_epochs = 1 trainer.hyperparameters.global_batch_size = 32 trainer.hyperparameters.max_steps = 12 job = trainer.train(wait=True) print(f"Job: {job.job_name}") print(f"Status: {job.job_status}") print(f"Output Model Package: {job.output_model_package_arn}")
使用 Nova 的受限模型包群组提交训练作业
有关如何创建受限模型包组的信息,请参阅上述步骤(为 Nova 创建受限模型包组)。
from sagemaker.train.multi_turn_rl_trainer import MultiTurnRLTrainer trainer = MultiTurnRLTrainer( model="nova-textgeneration-lite-v2", agent_env="arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/my-agent-runtime", training_dataset="s3://my-bucket/prompts/prompts.parquet", mlflow_app_arn="arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id", s3_output_path="s3://my-bucket/output/", role="arn:aws:iam::123456789012:role/SageMakerRole", accept_eula=True, output_model_package_group=output_mpg, intermediate_checkpoint_model_package_group=intermediate_mpg ) # View and adjust hyperparameters trainer.hyperparameters.get_info() trainer.hyperparameters.max_epochs = 1 trainer.hyperparameters.global_batch_size = 32 trainer.hyperparameters.max_steps = 12 job = trainer.train(wait=True) print(f"Job: {job.job_name}") print(f"Status: {job.job_status}") print(f"Output Model Package: {job.output_model_package_arn}")
Amazon CLI
使用 CreateJob API 创建训练作业。您可以在中指定代理配置、训练数据位置、基础模型和输出设置JobConfigDocument。
要检索完整 JobConfigDocument 架构,请执行以下操作:
aws sagemaker list-job-schema-versions --job-category AgentRFT aws sagemaker describe-job-schema-version --job-category AgentRFT --version "1.0.0"
使用 Bedrock 创造就业机会 AgentCore
aws sagemaker create-job \ --job-category AgentRFT \ --job-name "my-agent-rft-job" \ --role-arn "arn:aws:iam::123456789012:role/SageMakerFineTuningJobRole" \ --job-config-schema-version "1.0.0" \ --job-config-document '{ "AgentConfig": { "BedrockAgentCoreConfig": { "AgentRuntimeArn": "arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/my-agent" } }, "InputDataConfig": [...], "OutputDataConfig": {...}, "ModelPackageConfig": {...}, "TrainingConfig": {...} }' \ --region us-west-2
使用自定义 Lambda 代理创建任务
aws sagemaker create-job \ --job-category AgentRFT \ --job-name "my-custom-agent-rft-job" \ --role-arn "arn:aws:iam::account-id:role/SageMakerFineTuningJobRole" \ --job-config-schema-version "1.0.0" \ --job-config-document '{ "AgentConfig": { "CustomAgentLambdaConfig": { "LambdaArn": "arn:aws:lambda:us-west-2:account-id:function:rft-agent-forwarder" } }, "InputDataConfig": [...], "OutputDataConfig": {...}, "ModelPackageConfig": {...}, "TrainingConfig": {...} }' \ --region us-west-2
boto3
使用 Bedrock 创造就业机会 AgentCore
import json import boto3 sm = boto3.client("sagemaker") response = sm.create_job( JobName="my-agent-rft-job", RoleArn="arn:aws:iam::123456789012:role/SageMakerFineTuningJobRole", JobCategory="AgentRFT", JobConfigSchemaVersion="1.0.0", JobConfigDocument=json.dumps({ "AgentConfig": { "BedrockAgentCoreConfig": { "AgentRuntimeArn": "arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/my-agent" } }, "InputDataConfig": [...], "OutputDataConfig": {...}, "ModelPackageConfig": {...}, "TrainingConfig": {...} }) ) print(f"Job ARN: {response['JobArn']}")
使用自定义 Lambda 代理创建任务
import json import boto3 sm = boto3.client("sagemaker") response = sm.create_job( JobName="my-custom-agent-rft-job", RoleArn="arn:aws:iam::account-id:role/SageMakerFineTuningJobRole", JobCategory="AgentRFT", JobConfigSchemaVersion="1.0.0", JobConfigDocument=json.dumps({ "AgentConfig": { "CustomAgentLambdaConfig": { "LambdaArn": "arn:aws:lambda:us-west-2:account-id:function:rft-agent-forwarder" } }, "InputDataConfig": [...], "OutputDataConfig": {...}, "ModelPackageConfig": {...}, "TrainingConfig": {...} }) ) print(f"Job ARN: {response['JobArn']}")
监控训练
监控你的 Training Job
使用 DescribeJob API 随时查看任务的当前状态。任务状态会通过InProgress,然后转换到Completed、Failed或Stopped。
aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2
使用软件开发工具包:
# Run without blocking job = trainer.train(wait=False) job.wait(poll=5, timeout=3000, max_log_lines=10) # Check status job.refresh() print(f"Status: {job.job_status}") print(f"Secondary Status: {job.secondary_status}") print(f"Output Model Package: {job.output_model_package_arn}") print(f"MLflow Details: {job.mlflow_details}") print(f"Billable Tokens: {job.billable_token_usage}") # Open MLflow tracking URL job.get_mlflow_url() # Stop a running job job.stop() # Attach to an existing job from a different session existing_job = MultiTurnRLTrainer.attach(job_name="my-existing-job-name") print(f"Status: {existing_job.job_status}") print(f"Output Model: {existing_job.output_model_package_arn}") # List all completed jobs from sagemaker.train.agent_rft_job import AgentRFTJob for j in AgentRFTJob.get_all(status_equals="Completed"): print(f"{j.job_name}: {j.job_status}")
在 mlFlow 中监控训练
SageMaker AI 会自动与托管 mlFlow 集成,以跟踪训练作业的进度、指标和工件。要启用 mlFlow 跟踪,请在你的任务MlflowConfig中加上:OutputDataConfig
"OutputDataConfig": { "S3OutputPath": "s3://your-bucket/output/", "MlflowConfig": { "MlflowResourceArn": "arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/my-rft-mlflow-app" } }
先决条件
-
在您的账户中创建托管 mlFlow 应用程序。有关设置说明,请参阅 mlFlow 应用程序设置。
-
确保您的 SageMaker AI 执行角色有权写入 mlFlow 应用程序(
sagemaker-mlflow:*操作)。 -
MlflowResourceArn在您的任务配置中包括。
记录了什么
| # | 类别 | 记录了什么 | mlFlow 用户界面中的位置 |
|---|---|---|---|
| 1 | 训练指标 | Per-step 计数器、吞吐量、数据和代币核算、步骤每个阶段的全时钟持续时间、轨迹奖励汇总推出批次以及每个轨迹的回合计数分布 | “指标” 选项卡(时间序列图表) |
| 2 | 轨迹轨迹 | 完整的多回合对话,包括工具调用和奖励 | 追踪选项卡 |
详细的训练指标参考
在每个训练步骤中都会记录以下指标。
步数计数器和吞吐量 (training/)
| 指标 | 说明 |
|---|---|
training/epoch |
当前纪元数 |
training/global_step |
全球训练计步器 |
training/num_groups |
此步骤中的轨迹组 |
training/num_trajectories |
此步骤中处理的总轨迹 |
training/total_tokens |
在此步骤中,所有微批次的代币总和 |
training/num_datums |
由轨迹形成的训练基准 |
training/datums_per_trajectory |
每条轨迹发射的平均基准 |
training/action_tokens_mean |
每个轨迹的平均动作(响应)标记 |
training/obs_tokens_mean |
每个轨迹的平均观察(提示)标记 |
training/trainable_token_positions |
此步骤中可训练的目标位置总数 |
training/nontrainable_token_positions |
此步骤中不可训练的目标位置总数 |
training/trainable_token_ratio |
比率:trainable / (trainable + nontrainable)代币头寸 |
阶段持续时间 () timing_s/
| 指标 | 说明 |
|---|---|
timing_s/step |
整个步骤的总时间 |
timing_s/training |
是时候 forward/backward 通过和优化器步骤了 |
timing_s/policy_update |
为采样器更新权重节省时间 |
timing_s/save_checkpoint |
保存检查点的时间(仅限于检查点步骤) |
timing_s/eval |
时间运行评估(仅适用于评估步骤) |
奖励分配 (rollout/reward/)
| 指标 | 说明 |
|---|---|
rollout/reward/mean |
所有组别的平均轨迹奖励 |
rollout/reward/valid_mean |
仅对有效(非零优势)组的平均奖励;未进行过滤mean时等于 |
rollout/reward/std |
轨迹奖励的标准差 |
rollout/reward/min |
最低弹道奖励 |
rollout/reward/max |
最大弹道奖励 |
rollout/reward/zero_frac |
总奖励正好为 0.0 的轨迹分数 |
回合计数 (rollout/turns/)
| 指标 | 说明 |
|---|---|
rollout/turns/mean |
每个轨迹的平均转弯(过渡) |
rollout/turns/min |
穿过轨迹的最小转弯次数 |
rollout/turns/max |
穿过轨迹的最大转弯次数 |
代币长度 (rollout/tokens/)
| 指标 | 说明 |
|---|---|
rollout/tokens/prompt_mean |
每次过渡的平均提示令牌数量 |
rollout/tokens/response_mean |
每次过渡的平均响应令牌数量 |
rollout/tokens/response_std |
响应令牌计数的标准差 |
rollout/tokens/response_min |
最低响应令牌 |
rollout/tokens/response_max |
最大响应令牌(请注意聚类sampling_max_tokens) |
Log-probability 健康 (rollout/logprob/)
| 指标 | 说明 |
|---|---|
rollout/logprob/zero_count |
零日志探测代币总数 |
rollout/logprob/zero_frac |
所有日志探测值中精确为 0.0 的比例 |
rollout/logprob/zero_per_group |
每个轨迹组的平均对数探测值为零 |
rollout/logprob/nz_mean |
非零对数概率的平均值 |
rollout/logprob/nz_std |
非零对数探数的标准差 |
rollout/logprob/nz_min |
最小非零对数概率 |
rollout/logprob/nz_max |
最大非零对数概率 |
优势分布 (rollout/advantage/)
| 指标 | 说明 |
|---|---|
rollout/advantage/mean |
所有过渡的平均优势值 |
rollout/advantage/std |
优势的标准差 |
rollout/advantage/min |
最低优势 |
rollout/advantage/max |
最大优势 |
rollout/advantage/n_positive |
具有积极优势的过渡 |
rollout/advantage/n_negative |
具有负面优势的过渡 |
Batch-quality 分类 (analysis/)
| 指标 | 说明 |
|---|---|
analysis/batch_completion_ratio |
total_completed / batch_size— 预计抵达的团体所占比例 |
analysis/batch_valid_ratio |
valid_count / batch_size— 相对于整批次的非零优势组 |
analysis/zero_adv_groups |
所有过渡都具有接近零优势的群体 |
analysis/zero_adv_nonzero_reward |
Zero-advantage 至少有一次过渡的奖励不等于 0 的群组(二进制奖励的情况完全正确) |
analysis/zero_adv_zero_reward |
Zero-advantage 所有奖励都为 0 的群组(全是错误的情况) |
analysis/reward_variance_across_groups |
每组平均奖励的方差(高 = 不同批次) |
analysis/mean_group_reward_spread |
组内平均奖励点差 max - min |
评估奖励并通过 @k (val/reward/)
在基线(步骤 0)、每个val_every间隔和最后一步发射。包括与按提示汇总的群rollout/reward组奖励指标相同的分发指标。
分布:
| 指标 | 说明 |
|---|---|
val/reward/mean |
平均奖励超过评估套装 |
val/reward/std |
奖励标准开发者 |
val/reward/min |
最低奖励 |
val/reward/max |
最高奖励 |
val/reward/zero_frac |
零奖励轨迹的比例 |
Group-reward (每个提示聚合):
| 指标 | 说明 |
|---|---|
val/reward/min_within_groups |
每次提示的平均最低奖励 |
val/reward/mean_within_groups |
每次提示的平均奖励均值 |
val/reward/max_within_groups |
每次提示的平均最高奖励 |
val/reward/std_within_groups |
每次提示的平均奖励 std(一致性) |
val/reward/rollouts_per_prompt |
跨提示的平均推出次数 (n) |
val/reward/num_prompts |
评估了不同的提示 |
pass @k 和成功会计:
| 指标 | 说明 |
|---|---|
val/reward/succeeded_rollouts |
有奖励的投放总数 ≥ success_threshold |
val/reward/failed_rollouts |
带奖励的总推出次数 < success_threshold |
val/reward/success_threshold |
使用的阈值(为了清晰起见,进行了回声) |
val/reward/pass_at_{k} |
k 个样本中有 ≥ 1 个通过概率 |
val/reward/pass_power_{k} |
所有 k 个样本通过的概率(可靠性) |
评估回合计数 (val/turns/)
| 指标 | 说明 |
|---|---|
val/turns/mean |
每个评估轨迹的平均转弯数 |
val/turns/min |
最小转弯数 |
val/turns/max |
最大回合数 |
评估令牌长度 (val/tokens/)
| 指标 | 说明 |
|---|---|
val/tokens/prompt_mean |
每次过渡的平均提示令牌 |
val/tokens/response_mean |
每次过渡的平均响应令牌 |
val/tokens/response_std |
响应令牌的标准差 |
val/tokens/response_min |
最低响应令牌 |
val/tokens/response_max |
最大响应令牌 |
评估日志概率生命值 () val/logprob/
| 指标 | 说明 |
|---|---|
val/logprob/zero_count |
零日志探测代币总数 |
val/logprob/zero_frac |
零日志探测器的分数 |
val/logprob/zero_per_group |
每组零日志探测数 |
val/logprob/nz_mean |
非零对数概率的平均值 |
val/logprob/nz_std |
非零对数探数的标准差 |
val/logprob/nz_min |
最小非零对数概率 |
val/logprob/nz_max |
最大非零对数概率 |
访问 mlFlow 用户界面
通过预签名 URL 访问 mlFlow 用户界面:
aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2
将输出AuthorizedUrl中的内容复制到浏览器中。
特工轨迹和踪迹
在训练期间, SageMaker AI 会将您的代理与策略模型之间的每一次互动记录为轨迹,即一次部署的完整记录。每个轨迹都会捕获发送到模型的每个提示、生成的每个响应、发出的每个工具调用以及最终的奖励。轨迹以结构化轨迹的形式发布到您的 mlFlow 实验中。
追踪内容
-
来自训练数据集的输入提示
-
每个模型推理回合(提示、响应和令牌级数据)
-
工具调用及其结果(如果您的代理使用工具)
-
最终奖励分数
-
每回合的时间信息
在 mlFlow 用户界面中查看轨迹
通过预签名 URL 访问 mlFlow 用户界面:
aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2
将输出AuthorizedUrl中的内容复制到浏览器中。
使用上面的预签名 URL 打开 mlFlow 用户界面。导航到您的实验运行并选择 Traces 选项卡。每条跟踪记录代表一个已完成的部署,并显示:
-
系统提示符和用户提示符
-
每个助手的回应( thinking/reasoning 如果适用,还有)
-
工具使用跨度显示了哪些工具被调用及其输出
-
分配给轨迹的奖励分数
使用轨迹调试低奖励分数
| 症状 | 要查找的内容 |
|---|---|
| 大多数推出的奖励都很低 | 模型响应是否一致? 提示格式正确吗? |
| Tool-related 失败 | 工具调用成功了吗? 输入和输出是否格式正确? |
| 代理循环 | 代理是否在没有取得进展的情况下重复了同样的操作? |
| 截断的响应 | maxTokens 限制是否会切断回复? |
获取训练结果
训练作业完成后,训练后的模型权重将存储为 SageMaker AI Model Package。本节介绍如何查找结果、了解训练期间生成的检查点类型,以及如何使用它们进行部署或继续训练。
结果是如何存储的
SageMaker AI 将训练输出作为版本化、不可变的模型包存储在 Model Package 组中。 Multi-turn RL 使用两个单独的组,您在创建作业时指定这两个组:
| Group | 用途 | 内容 |
|---|---|---|
| 输出模型 Package 组 | 最终训练好的模型 | HuggingFace-compatible LoRa 适配器权重(adapter_config.json + adapter_model.safetensors) |
| 中间检查点模型 Package Group | 可恢复训练状态 | LoRa 适配器权重 + 优化器状态 + 训练步骤元数据 |
在您的:中配置这两个群组 ModelPackageConfig:
"ModelPackageConfig": { "OutputModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-final-models", "IntermediateCheckpointModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-intermediate-checkpoints" }
检查点类型
训练会生成两种类型的检查点,每个训练步骤都会保存:
模型检查点(仅限权重)
-
存储在输出模型 Package 组中
-
包含格式化 HuggingFace-compatible 的 LoRa 适配器 SafeTensors 权重
-
用于推理、部署或作为新训练作业的起点
-
在每个步骤、任务完成时和作业停止时创建
可恢复检查点(完整状态)
-
存储在中间检查点模型 Package 组中
-
包含 LoRa 适配器权重、优化器状态和每个 GPU 的训练步骤元数据
-
用于从中断的作业停止的确切步骤开始恢复
-
内部格式 — 不能直接用于推理
检查点生命周期
Step 1 → Intermediate Checkpoint (resumable) Step 1 → Intermediate Checkpoint (HF-compatible) ... Step N-1 → Intermediate Checkpoint (resumable) Step N-1 → Intermediate Checkpoint (HF-compatible) ... Step N (final) → Model Checkpoint (HuggingFace LoRA) → Output Model Package Group
检索经过训练的模型
作业成功完成后,最终模型将作为模型包保存在输出模型包组中。工作记录上的OutputModelPackageArn字段包含 ARN。
检查任务完成情况并检索输出模型 ARN:
aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2
在回复OutputModelPackageArn中查找。用它来描述 Model Package 并获取权重的 S3 位置:
aws sagemaker describe-model-package \ --model-package-name "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-final-models/5"
如果作业失败或在完成之前停止,则会尽最大努力将最后一个中间检查点提升到 Output Model Package 组。OutputModelPackageArn用同样的方法检查。
要在训练期间监控检查点的创建情况,请查看 DescribeJob 输出中的ResumableCheckpoint和ModelCheckpoint字段。
恢复中断的作业
如果作业失败或在训练中停止,则可以开始一项新作业,该作业从上次停止的确切步骤继续进行。该平台从可恢复的检查点恢复完整的训练状态(权重、优化器动量和计步器)。
将中间检查点模型 Package 组中的可恢复检查点指定为InputModelPackageArn:
"ModelPackageConfig": { "OutputModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-final-models", "IntermediateCheckpointModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-intermediate-checkpoints", "InputModelPackageArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-intermediate-checkpoints/5" }
InputModelPackageArn必须指向可恢复的检查点(其 Model Package 元数据IsCheckpoint=true中包含的检查点)。训练将从检查点之后的步骤中继续,例如,如果在步骤 4 中保存了检查点,则从步骤 5 开始训练。
在原始作业和恢复的作业之间,以下内容必须保持不变:
-
基础模型
-
LoRa 配置(等级和 alpha)
-
超参数(学习率、批量大小等)
-
数据集
继续就新工作进行培训(迭代训练)
迭代训练允许您在先前训练过的模型的基础上进行构建,该模型具有不同的数据集、不同的超参数或精细的奖励函数。与恢复不同,这会开始新的训练——优化器重置,计步器重置为 0,只有经过训练的 LoRa 权重才会延续。
将输出模型 Package 组中的模型检查点指定为InputModelPackageArn:
"ModelPackageConfig": { "OutputModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-final-models", "IntermediateCheckpointModelPackageGroupArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package-group/my-intermediate-checkpoints", "InputModelPackageArn": "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-final-models/3" }
在迭代之间可以更改的内容:
-
超参数(学习率、批次大小、最大步数、组大小等)
-
数据集(不同的提示或数据分布)
-
奖励函数
-
代理配置
什么必须保持不变:
-
基本型号 — LoRa 适配器绑定到基本模型架构
迭代训练的常见模式:
-
课程学习 — 先针对较简单的问题进行训练,然后继续处理更难的问题
-
奖励细化 — 从一个简单的奖励函数开始,然后用一个更细致入微的功能进行迭代
-
超参数调整 — 在观察初始训练动态后,增加批次大小或调整学习率
检查点最佳实践
-
监控检查点的创建。在训练期间 DescribeJob 用于
ModelCheckpoint田径ResumableCheckpoint运动,以便在需要恢复时知道有哪些可用的内容。 -
为长时间作业的失败做好准备。如果作业有许多步骤,请将工作流程设计为从检查点恢复,而不是从头开始。