View a markdown version of this page

訓練任務提交 - Amazon SageMaker AI

本文為英文版的機器翻譯版本,如內容有任何歧義或不一致之處,概以英文版為準。

訓練任務提交

啟動訓練任務

部署代理程式且資料集位於 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,然後轉換至 CompletedFailedStopped

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 輸出中的 ResumableCheckpointModelCheckpoint 欄位。

繼續中斷的任務

如果任務失敗或已在訓練中停止,您可以從停止的確切步驟開始新的任務。平台會從可繼續的檢查點還原完整訓練狀態:權重、最佳化工具動量和步進計數器。

從中繼檢查點模型套件群組將可恢復檢查點指定為 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 在訓練期間追蹤 ResumableCheckpointModelCheckpoint 欄位,以便在您需要繼續時了解可用的項目。

  • 規劃長時間工作的失敗。如果任務有許多步驟,請將您的工作流程設計為從檢查點繼續,而不是從頭開始重新啟動。