翻訳は機械翻訳により提供されています。提供された翻訳内容と英語版の間で齟齬、不一致または矛盾がある場合、英語版が優先します。
トレーニングジョブの送信
トレーニングジョブの起動
エージェントがデプロイされ、データセットが 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) を選択する場合は、トレーニングジョブを送信する前に、オプションで制限付きモデルパッケージグループを作成します (次のステップ)。このステップをスキップすると、SDK によって自動的に作成されます。
制限付きモデルパッケージグループ (RMPG) は、ManagedStorageType: 制限付きモデルパッケージグループです。これは、モデルの重みが によって管理され AWS 、顧客が直接アクセスできない Nova などのクローズドソースモデルに必要です。
RFT ジョブスキーマには、2 つの個別の制限付き 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 の制限付きモデルパッケージグループを使用してトレーニングジョブを送信する
制限付きモデルパッケージグループの作成方法については、上記のステップ (Create Restricted Model Package Group for 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 |
現在のエポック番号 |
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 |
ゼロログプローブトークンの合計 |
rollout/logprob/zero_frac |
正確に 0.0 であるすべての logprob の割合 |
rollout/logprob/zero_per_group |
軌道グループあたりの平均ゼロログ確率 |
rollout/logprob/nz_mean |
ゼロ以外の logprobs の平均 |
rollout/logprob/nz_std |
ゼロ以外の logprobs の標準偏差 |
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 |
少なくとも 1 つの移行で報酬が 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 |
報酬 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} |
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 |
ゼロ以外の logprobs の平均 |
val/logprob/nz_std |
ゼロ以外の logprobs の標準偏差 |
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 はエージェントとポリシーモデル間のすべてのやり取りを、1 つのロールアウトの完全なレコードである軌跡として記録します。各軌道は、モデルに送信されたすべてのプロンプト、生成されたすべてのレスポンス、行われたすべてのツール呼び出し、最終的な報酬をキャプチャします。軌道は構造化トレースとして 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 を開きます。実験の実行に移動し、トレースタブを選択します。各トレースは、完了したロールアウトを 1 つ表し、以下を示します。
-
システムプロンプトとユーザープロンプト
-
各アシスタントの応答 (該当する場合は思考/理由付き)
-
呼び出されたツールとその出力を示すツールの使用スパン
-
軌道に割り当てられた報酬スコア
軌跡を使用して低報酬スコアをデバッグする
| 症状 | 検索対象 |
|---|---|
| ほとんどのロールアウトで報酬が低い | モデルレスポンスは一貫性がありますか? プロンプト形式は正しいですか? |
| ツール関連の障害 | ツール呼び出しは成功していますか? 入力と出力は適切に形成されていますか? |
| エージェントループ | エージェントは進行せずに同じアクションを繰り返していますか? |
| 切り捨てられたレスポンス | レスポンスは maxTokens の制限でカットされていますか? |
トレーニング結果の取得
トレーニングジョブが完了すると、トレーニングされたモデルの重みは SageMaker AI モデルパッケージとして保存されます。このセクションでは、結果を検索し、トレーニング中に生成されたチェックポイントタイプを理解し、デプロイまたは継続的なトレーニングに使用する方法について説明します。
結果の保存方法
SageMaker AI は、トレーニング出力をバージョン管理されたイミュータブルなモデルパッケージとしてモデルパッケージグループ内に保存します。マルチターン RL は、ジョブの作成時に指定する 2 つの個別のグループを使用します。
| 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" }
チェックポイントタイプ
トレーニングでは、トレーニングステップごとに 2 種類のチェックポイントが生成されます。
モデルチェックポイント (重みのみ)
-
出力モデルパッケージグループに保存されます
-
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フィールドを監視します。 DescribeJob
中断されたジョブを再開する
トレーニング中にジョブが失敗または停止した場合、中断した正確なステップから取得する新しいジョブを開始できます。プラットフォームは、再開可能なチェックポイントから、重み、オプティマイザの勢い、ステップカウンターなど、完全なトレーニング状態を復元します。
中間チェックポイントモデルパッケージグループから再開可能なチェックポイントを として指定します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" }
は、再開可能なチェックポイント (モデルパッケージメタデータIsCheckpoint=true内の を持つチェックポイント) を指すInputModelPackageArn必要があります。トレーニングはチェックポイントの後のステップから再開します。例えば、チェックポイントがステップ 4 で保存された場合、トレーニングはステップ 5 から続行されます。
次の内容は、元のジョブと再開されたジョブで同じである必要があります。
-
ベースモデル
-
LoRA 設定 (ランクとアルファ)
-
ハイパーパラメータ (学習レート、バッチサイズなど)
-
データセット
新しいジョブでトレーニングを継続する (反復トレーニング)
反復トレーニングを使用すると、異なるデータセット、異なるハイパーパラメータ、または洗練された報酬関数を使用して、以前にトレーニングしたモデルを構築できます。再開とは異なり、これにより新しいトレーニング実行が開始されます。オプティマイザがリセットされ、ステップカウンターが 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フィールドを追跡し、再開する必要がある場合に利用できる内容を把握します。 -
長いジョブで障害が発生した場合の計画を立てます。ジョブに多くのステップがある場合は、ゼロから再起動するのではなく、チェックポイントから再開するようにワークフローを設計します。