View a markdown version of this page

培训作业提交 - 亚马逊 SageMaker AI
Amazon Web Services 文档中描述的 Amazon Web Services 服务或功能可能因区域而异。要查看适用于中国区域的差异,请参阅 中国的 Amazon Web Services 服务入门 (PDF)

本文属于机器翻译版本。若本译文内容与英语原文存在差异,则一律以英文原文为准。

培训作业提交

启动培训作业

部署代理且数据集在 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,然后转换到CompletedFailedStopped

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 输出中的ResumableCheckpointModelCheckpoint字段。

恢复中断的作业

如果作业失败或在训练中停止,则可以开始一项新作业,该作业从上次停止的确切步骤继续进行。该平台从可恢复的检查点恢复完整的训练状态(权重、优化器动量和计步器)。

将中间检查点模型 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运动,以便在需要恢复时知道有哪些可用的内容。

  • 为长时间作业的失败做好准备。如果作业有许多步骤,请将工作流程设计为从检查点恢复,而不是从头开始。