本文為英文版的機器翻譯版本,如內容有任何歧義或不一致之處,概以英文版為準。
訓練任務提交
啟動訓練任務
部署代理程式且資料集位於 S3 之後,請使用下列其中一種方法建立訓練任務。
SageMaker AI Studio
-
在導覽窗格中導覽至模型,然後選取 JumpStart 基礎模型。
-
選取支援多迴轉 RL 的模型 (請參閱支援的模型表格),然後選擇自訂模型,然後選擇使用 UI 自訂。
-
選取多轉強化學習作為自訂技術。
-
設定您的代理程式環境:選取 Bedrock AgentCore 執行時間或提供 Lambda 轉送器 ARN。
-
將您的訓練資料集提供為 S3 URI 或已註冊的資料集。
-
視需要調整超參數。
-
檢閱您的組態,然後選擇提交。
SageMaker AI Python SDK
探索支援的模型
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:Bedrock 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 建立受限制的模型套件群組 (選用)
如果您選擇的是 Nova 模型 (nova-textgeneration-lite-v2),則選擇性地在提交訓練任務 (下一個步驟) 之前建立受限制模型套件群組。如果您略過此步驟,軟體開發套件會自動為您建立一個。
受限模型套件群組 (RMPG) 是具有 ManagedStorageType: Restricted 的模型套件群組。這是 Nova 等封閉來源模型的必要項目,其中模型權重由 管理 AWS ,且無法直接供客戶存取。
RFT 任務結構描述需要兩個單獨的受限 MPGs:
-
輸出 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}")
AWS 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']}")
監控訓練
監控您的訓練任務
使用 DescribeJob API 隨時檢查任務的目前狀態。任務狀態會透過 轉換InProgress,然後轉換至 Completed、 Failed或 Stopped。
aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2
使用 SDK:
# 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中包含 。
記錄的內容
| # | Category | 記錄的內容 | MLflow UI 中的位置 |
|---|---|---|---|
| 1 | 訓練指標 | 每個步驟計數器、輸送量、基準和權杖會計、步驟每個階段的全天候持續時間、軌跡獎勵摘要推展批次,以及每個軌跡的周轉計數分佈 | 指標索引標籤 (時間序列圖表) |
| 2 | 軌跡追蹤 | 使用工具呼叫和獎勵進行完整多迴轉對話 | 追蹤索引標籤 |
詳細訓練指標參考
每個訓練步驟都會記錄下列指標。
步驟計數器和輸送量 (training/)
| 指標 | 說明 |
|---|---|
training/epoch |
目前的 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 |
向前/向後通過和最佳化工具步驟的時間 |
timing_s/policy_update |
節省取樣器更新權重的時間 |
timing_s/save_checkpoint |
儲存檢查點的時間 (僅適用於檢查點步驟) |
timing_s/eval |
時間執行評估 (僅適用於評估步驟) |
獎勵分佈 (rollout/reward/)
| 指標 | 說明 |
|---|---|
rollout/reward/mean |
所有群組的平均軌跡獎勵 |
rollout/reward/valid_mean |
僅對有效 (non-zero-advantage群組的平均獎勵;在未進行篩選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) |
日誌機率運作狀態 (rollout/logprob/)
| 指標 | 說明 |
|---|---|
rollout/logprob/zero_count |
零 logprob 權杖總數 |
rollout/logprob/zero_frac |
剛好為 0.0 的所有 logprob 分數 |
rollout/logprob/zero_per_group |
每個軌跡群組的平均零 logprob |
rollout/logprob/nz_mean |
非零 logprob 的平均值 |
rollout/logprob/nz_std |
非零 logprob 的標準差 |
rollout/logprob/nz_min |
最小非零 logprob |
rollout/logprob/nz_max |
最大非零 logprob |
優勢分佈 (rollout/advantage/)
| 指標 | 說明 |
|---|---|
rollout/advantage/mean |
所有轉換的平均優勢值 |
rollout/advantage/std |
優點的標準差 |
rollout/advantage/min |
最低優勢 |
rollout/advantage/max |
最大優勢 |
rollout/advantage/n_positive |
具有正面優勢的轉換 |
rollout/advantage/n_negative |
具有負面優勢的轉換 |
批次品質分類 (analysis/)
| 指標 | 說明 |
|---|---|
analysis/batch_completion_ratio |
total_completed / batch_size — 到達的預期群組部分 |
analysis/batch_valid_ratio |
valid_count / batch_size — non-zero-advantage群組 |
analysis/zero_adv_groups |
所有轉換都有近乎零優點的群組 |
analysis/zero_adv_nonzero_reward |
零優勢群組,其中至少一個轉換具有不等於 0 的獎勵 (二進位獎勵的所有正確案例) |
analysis/zero_adv_zero_reward |
零優勢群組,其中所有獎勵都是 0 (全錯誤案例) |
analysis/reward_variance_across_groups |
每個群組平均獎勵的差異 (高 = 多樣化批次) |
analysis/mean_group_reward_spread |
平均群組內獎勵分散 max - min |
評估獎勵和 pass@k (val/reward/)
在基準 (步驟 0)、每個val_every間隔和最後一個步驟發出。包含與依提示彙總rollout/reward的群組獎勵指標相同的分佈指標。
分佈:
| 指標 | 說明 |
|---|---|
val/reward/mean |
評估集的平均獎勵 |
val/reward/std |
獎勵標準開發 |
val/reward/min |
最低獎勵 |
val/reward/max |
獎勵上限 |
val/reward/zero_frac |
零獎勵軌跡的分數 |
群組獎勵 (每次提示彙總):
| 指標 | 說明 |
|---|---|
val/reward/min_within_groups |
每個提示的平均最低獎勵 |
val/reward/mean_within_groups |
每個提示的平均獎勵 |
val/reward/max_within_groups |
每個提示的平均最高獎勵 |
val/reward/std_within_groups |
每個提示的平均獎勵標準 (一致性) |
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} |
機率 ≥1 k 樣本通過 |
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 |
零 logprob 權杖總數 |
val/logprob/zero_frac |
零 logprob 的分數 |
val/logprob/zero_per_group |
每個群組零 logprob |
val/logprob/nz_mean |
非零 logprob 的平均值 |
val/logprob/nz_std |
非零 logprob 的標準差 |
val/logprob/nz_min |
最小非零 logprob |
val/logprob/nz_max |
最大非零 logprob |
存取 MLflow UI
透過預先簽章的 URL 存取 MLflow UI:
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 UI 中檢視軌跡
透過預先簽章的 URL 存取 MLflow UI:
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 UI。導覽至實驗的執行,然後選取追蹤索引標籤。每個追蹤代表一個已完成的推展,並顯示:
-
系統提示和使用者提示
-
每個助理回應 (如適用,使用思考/理性)
-
工具使用範圍顯示呼叫了哪些工具及其輸出
-
指派給軌跡的獎勵分數
使用軌跡來偵錯低獎勵分數
| 徵狀 | 要尋找的內容 |
|---|---|
| 大多數推展的低獎勵 | 模型回應是否一致? 提示格式是否正確? |
| 工具相關的故障 | 工具呼叫是否成功? 輸入和輸出格式是否正確? |
| 代理程式迴圈 | 客服人員是否在不進行的情況下重複相同的動作? |
| 截斷的回應 | 回應是否被 maxTokens 限制截斷? |
取得訓練結果
當訓練任務完成時,您訓練過的模型權重會儲存為 SageMaker AI 模型套件。本節說明如何尋找結果、了解訓練期間產生的檢查點類型,以及將其用於部署或持續訓練。
如何儲存結果
SageMaker AI 會將訓練輸出儲存為模型套件群組中的版本控制、不可變模型套件。多迴轉 RL 使用兩個不同的群組,您在建立任務時指定:
| Group | 用途 | 目錄 |
|---|---|---|
| 輸出模型套件群組 | 最終訓練模型 | HuggingFace 相容的 LoRA 轉接器權重 (adapter_config.json + adapter_model.safetensors) |
| 中繼檢查點模型套件群組 | 可繼續的訓練狀態 | 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" }
檢查點類型
訓練會產生兩種檢查點類型,並在每個訓練步驟中儲存:
模型檢查點 (僅限權重)
-
存放在輸出模型套件群組中
-
包含 SafeTensors 格式的 HuggingFace 相容 LoRA 轉接器權重
-
用於推論、部署或作為新訓練任務的起點
-
在每個步驟、任務完成和任務停止時建立
可繼續檢查點 (完整狀態)
-
存放在中繼檢查點模型套件群組
-
包含 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中尋找 。使用它來描述模型套件,並取得權重的 S3 位置:
aws sagemaker describe-model-package \ --model-package-name "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-final-models/5"
如果任務在完成之前失敗或停止,則會盡力將最後一個中繼檢查點提升為輸出模型套件群組。檢查OutputModelPackageArn的方式相同。
若要在訓練期間監控檢查點建立,請監看 DescribeJob 輸出中的 ResumableCheckpoint和 ModelCheckpoint 欄位。
繼續中斷的任務
如果任務失敗或已在訓練中停止,您可以從停止的確切步驟開始新的任務。平台會從可繼續的檢查點還原完整訓練狀態:權重、最佳化工具動量和步進計數器。
從中繼檢查點模型套件群組將可恢復檢查點指定為 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指向可恢復的檢查點 (模型套件中繼資料IsCheckpoint=true中具有 的檢查點)。檢查點之後的步驟會繼續訓練,例如,如果檢查點是在步驟 4 儲存,則訓練會從步驟 5 繼續。
原始任務和繼續任務之間的下列項目必須保持不變:
-
基礎模型
-
LoRA 組態 (排名和 Alpha)
-
超參數 (學習率、批次大小等)
-
資料集
繼續訓練新任務 (反覆訓練)
反覆訓練可讓您以先前訓練的模型為基礎,使用不同的資料集、不同的超參數或精細的獎勵函數。與繼續不同,這會啟動新的訓練執行:最佳化工具重設、步驟計數器重設為 0,而且只有訓練過的 LoRA 權重會轉移。
從輸出模型套件群組將模型檢查點指定為 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" }
您可以在反覆運算之間變更的內容:
-
超參數 (學習速率、批次大小、Max_steps、group_size 等)
-
資料集 (不同的提示或資料分佈)
-
獎勵函數
-
代理程式組態
哪些項目必須保持不變:
-
基本模型 — LoRA 轉接器繫結至基本模型架構
反覆訓練的常見模式:
-
課程學習 - 首先針對較容易的問題進行訓練,然後繼續較困難的問題
-
獎勵精簡 – 從簡單的獎勵函數開始,然後使用更細微的獎勵函數反覆運算
-
超參數調整 — 在觀察初始訓練動態之後增加批次大小或調整學習率
檢查點最佳實務
-
監控檢查點建立。使用 DescribeJob 在訓練期間追蹤
ResumableCheckpoint和ModelCheckpoint欄位,以便在您需要繼續時了解可用的項目。 -
規劃長時間工作的失敗。如果任務有許多步驟,請將您的工作流程設計為從檢查點繼續,而不是從頭開始重新啟動。