View a markdown version of this page

Pengajuan pekerjaan pelatihan - Amazon SageMaker AI

Terjemahan disediakan oleh mesin penerjemah. Jika konten terjemahan yang diberikan bertentangan dengan versi bahasa Inggris aslinya, utamakan versi bahasa Inggris.

Pengajuan pekerjaan pelatihan

Peluncuran Pekerjaan Pelatihan

Setelah agen Anda digunakan dan kumpulan data Anda ada di S3, buat pekerjaan pelatihan menggunakan salah satu metode berikut.

SageMaker Studio AI

  • Arahkan ke Model di panel navigasi dan pilih Model JumpStart Dasar.

  • Pilih model yang mendukung RL multi-putaran (lihat tabel model yang didukung) dan pilih Sesuaikan model, lalu Sesuaikan dengan UI.

  • Pilih Multi-Turn Reinforcement Learning sebagai teknik kustomisasi.

  • Konfigurasikan lingkungan agen Anda — pilih AgentCore runtime Bedrock Anda atau berikan ARN forwarder Lambda Anda.

  • Berikan kumpulan data pelatihan Anda sebagai URI S3 atau kumpulan data terdaftar.

  • Sesuaikan hyperparameters sesuai kebutuhan.

  • Tinjau konfigurasi Anda dan pilih Kirim.

SageMaker SDK Python AI

Temukan model yang didukung

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}")

Siapkan lingkungan agen Anda

Opsi 1: runtime batuan dasar AgentCore

# List available runtimes runtimes = MultiTurnRLTrainer.list_bedrock_agentcore_runtimes() for rt in runtimes: print(f" - {rt['name']} ({rt['status']}) → {rt['arn']}")

Opsi 2: Agen Lambda Kustom

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")

Daftarkan dataset Anda (opsional)

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}")

Buat Grup Paket Model Terbatas untuk Nova (opsional)

Jika Anda memilih Nova model (nova-textgeneration-lite-v2), maka secara opsional buat Grup Paket Model Terbatas sebelum mengirimkan pekerjaan pelatihan (langkah selanjutnya). Jika Anda melewati langkah ini, SDK secara otomatis membuatnya untuk Anda.

Restricted Model Package Group (RMPG) adalah Model Package Group dengan ManagedStorageType: Restricted. Ini diperlukan untuk model sumber tertutup seperti Nova di mana bobot model dikelola oleh AWS dan tidak dapat diakses secara langsung oleh pelanggan.

Skema pekerjaan RFT membutuhkan dua MPG Terbatas terpisah:

  • Output MPG - menyimpan paket model fine-tuned akhir

  • Pos Pemeriksaan Menengah MPG — disediakan untuk pos pemeriksaan pelatihan menengah (harus berbeda dari MPG Output)

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)

Setelah grup paket Model dibuat, lewati grup pada langkah berikutnya saat mengirimkan pekerjaan pelatihan.

Kirim pekerjaan pelatihan dengan 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}")

Kirim pekerjaan pelatihan dengan agen Lambda khusus

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}")

Kirim pekerjaan pelatihan dengan grup paket Model Terbatas untuk Nova

Lihat langkah di atas (Create Restricted Model Package Group for Nova) tentang cara membuat grup Restricted Model Package.

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

Buat pekerjaan pelatihan menggunakan CreateJob API. Anda menentukan konfigurasi agen, lokasi data pelatihan, model dasar, dan pengaturan output diJobConfigDocument.

Untuk mengambil JobConfigDocument skema lengkap:

aws sagemaker list-job-schema-versions --job-category AgentRFT aws sagemaker describe-job-schema-version --job-category AgentRFT --version "1.0.0"

Buat pekerjaan dengan 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

Buat pekerjaan dengan agen Lambda khusus

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

Buat pekerjaan dengan 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']}")

Buat pekerjaan dengan agen Lambda khusus

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']}")

Pelatihan Pemantauan

Pantau Training Job Anda

Gunakan DescribeJob API untuk memeriksa status pekerjaan Anda saat ini kapan saja. Transisi status pekerjaan melaluiInProgress, dan kemudian keCompleted, Failed atauStopped.

aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2

Gunakan 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}")

Memantau Pelatihan di MLFlow

SageMaker AI secara otomatis terintegrasi dengan MLFlow terkelola untuk melacak kemajuan, metrik, dan artefak pekerjaan pelatihan Anda. Untuk mengaktifkan pelacakan MLFlow, sertakan MlflowConfig dalam pekerjaan Anda: OutputDataConfig

"OutputDataConfig": { "S3OutputPath": "s3://your-bucket/output/", "MlflowConfig": { "MlflowResourceArn": "arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/my-rft-mlflow-app" } }

Prasyarat

  • Buat Aplikasi MLFlow terkelola di akun Anda. Untuk petunjuk penyiapan, lihat Pengaturan Aplikasi MLFlow.

  • Pastikan peran eksekusi SageMaker AI Anda memiliki izin untuk menulis ke Aplikasi MLFlow (sagemaker-mlflow:*tindakan).

  • Sertakan MlflowResourceArn dalam konfigurasi pekerjaan Anda.

Apa yang akan dicatat

# Kategori Apa yang dicatat Dimana di MLFlow UI
1 Metrik pelatihan Per-step penghitung, throughput, datum dan akuntansi token, durasi sepanjang jam dari setiap fase langkah, batch peluncuran ringkasan lintasan-hadiah, dan distribusi turn-count per lintasan Tab metrik (grafik deret waktu)
2 Jejak lintasan Percakapan multi-putaran penuh dengan panggilan alat dan hadiah Tab Jejak

Referensi metrik pelatihan terperinci

Metrik berikut dicatat pada setiap langkah pelatihan.

Penghitung langkah dan throughput () training/

Metrik Deskripsi
training/epoch Nomor zaman saat ini
training/global_step Penghitung langkah pelatihan global
training/num_groups Kelompok lintasan dalam langkah ini
training/num_trajectories Total lintasan yang diproses pada langkah ini
training/total_tokens Token dijumlahkan di semua batch mikro dalam langkah ini
training/num_datums Data pelatihan terbentuk dari lintasan
training/datums_per_trajectory Datum rata-rata yang dipancarkan per lintasan
training/action_tokens_mean Token tindakan (respons) rata-rata per lintasan
training/obs_tokens_mean Token observasi rata-rata (prompt) per lintasan
training/trainable_token_positions Total posisi target yang dapat dilatih dalam langkah ini
training/nontrainable_token_positions Total posisi target yang tidak dapat dilatih dalam langkah ini
training/trainable_token_ratio Rasio: posisi trainable / (trainable + nontrainable) token

Durasi fase () timing_s/

Metrik Deskripsi
timing_s/step Total waktu untuk langkah penuh
timing_s/training Waktu untuk forward/backward operan dan langkah pengoptimal
timing_s/policy_update Menghemat waktu bobot yang diperbarui untuk sampler
timing_s/save_checkpoint Menghemat waktu pos pemeriksaan (hanya pada langkah-langkah pos pemeriksaan)
timing_s/eval Evaluasi waktu berjalan (hanya pada langkah-langkah evaluasi)

Distribusi hadiah (rollout/reward/)

Metrik Deskripsi
rollout/reward/mean Berarti hadiah lintasan di semua kelompok
rollout/reward/valid_mean Hadiah rata-rata hanya atas grup yang valid (non-zero-advantage); sama dengan saat tidak ada mean penyaringan yang terjadi
rollout/reward/std Standar deviasi imbalan lintasan
rollout/reward/min Hadiah lintasan minimum
rollout/reward/max Hadiah lintasan maksimum
rollout/reward/zero_frac Fraksi lintasan dengan total hadiah tepat 0,0

Putar hitungan () rollout/turns/

Metrik Deskripsi
rollout/turns/mean Rata-rata belokan (transisi) per lintasan
rollout/turns/min Belokan minimum melintasi lintasan
rollout/turns/max Belokan maksimum melintasi lintasan

Panjang token () rollout/tokens/

Metrik Deskripsi
rollout/tokens/prompt_mean Rata-rata jumlah token prompt per transisi
rollout/tokens/response_mean Rata-rata jumlah token respons per transisi
rollout/tokens/response_std Deviasi standar jumlah token respons
rollout/tokens/response_min Token respons minimum
rollout/tokens/response_max Token respons maksimum (perhatikan pengelompokan disampling_max_tokens)

Log-probability kesehatan (rollout/logprob/)

Metrik Deskripsi
rollout/logprob/zero_count Total token zero-logprob
rollout/logprob/zero_frac Fraksi dari semua logprob yang persis 0,0
rollout/logprob/zero_per_group Rata-rata nol logprobs per kelompok lintasan
rollout/logprob/nz_mean Rata-rata logprob bukan nol
rollout/logprob/nz_std Standar deviasi logprobs bukan nol
rollout/logprob/nz_min Minimum logprob bukan nol
rollout/logprob/nz_max Maksimum logprob bukan nol

Distribusi keuntungan (rollout/advantage/)

Metrik Deskripsi
rollout/advantage/mean Nilai keuntungan rata-rata di semua transisi
rollout/advantage/std Standar deviasi keuntungan
rollout/advantage/min Keuntungan minimum
rollout/advantage/max Keuntungan maksimal
rollout/advantage/n_positive Transisi dengan keuntungan positif
rollout/advantage/n_negative Transisi dengan keuntungan negatif

Batch-quality klasifikasi (analysis/)

Metrik Deskripsi
analysis/batch_completion_ratio total_completed / batch_size— fraksi kelompok yang diharapkan yang tiba
analysis/batch_valid_ratio valid_count / batch_size— kelompok non-nol-keuntungan relatif terhadap batch penuh
analysis/zero_adv_groups Grup di mana semua transisi memiliki keuntungan mendekati nol
analysis/zero_adv_nonzero_reward Zero-advantage grup di mana setidaknya satu transisi memiliki hadiah yang tidak sama dengan 0 (kasus yang benar untuk hadiah biner)
analysis/zero_adv_zero_reward Zero-advantage grup di mana semua hadiah adalah 0 (kasus semua salah)
analysis/reward_variance_across_groups Varians imbalan rata-rata per kelompok (tinggi = batch beragam)
analysis/mean_group_reward_spread Rata-rata spread hadiah dalam kelompok max - min

Hadiah evaluasi dan lulus @k (val/reward/)

Dipancarkan pada baseline (langkah 0), pada setiap val_every interval, dan pada langkah terakhir. Termasuk metrik distribusi yang sama dengan rollout/reward ditambah metrik reward grup yang dikumpulkan berdasarkan prompt.

Distribusi:

Metrik Deskripsi
val/reward/mean Berarti hadiah atas set eval
val/reward/std Hadiah std dev
val/reward/min Hadiah minimum
val/reward/max Hadiah maksimum
val/reward/zero_frac Fraksi lintasan hadiah nol

Group-reward (agregasi per prompt):

Metrik Deskripsi
val/reward/min_within_groups Hadiah minimum rata-rata per prompt
val/reward/mean_within_groups Rata-rata hadiah rata-rata per prompt
val/reward/max_within_groups Hadiah maksimum rata-rata per prompt
val/reward/std_within_groups Rata-rata hadiah per prompt std (konsistensi)
val/reward/rollouts_per_prompt Peluncuran rata-rata (n) di seluruh prompt
val/reward/num_prompts Permintaan berbeda dievaluasi

Lulus @k dan akuntansi sukses:

Metrik Deskripsi
val/reward/succeeded_rollouts Total peluncuran dengan hadiah ≥ success_threshold
val/reward/failed_rollouts Total peluncuran dengan hadiah < success_threshold
val/reward/success_threshold Ambang batas yang digunakan (digaungkan untuk kejelasan)
val/reward/pass_at_{k} Probabilitas ≥1 dari k sampel lewat
val/reward/pass_power_{k} Probabilitas semua k sampel lulus (keandalan)

Jumlah giliran evaluasi () val/turns/

Metrik Deskripsi
val/turns/mean Rata-rata putaran per lintasan eval
val/turns/min Belokan minimum
val/turns/max Giliran maksimum

Panjang token evaluasi () val/tokens/

Metrik Deskripsi
val/tokens/prompt_mean Berarti token prompt per transisi
val/tokens/response_mean Token respons rata-rata per transisi
val/tokens/response_std Standar deviasi token respons
val/tokens/response_min Token respons minimum
val/tokens/response_max Token respons maksimum

Evaluasi log-probabilitas kesehatan () val/logprob/

Metrik Deskripsi
val/logprob/zero_count Total token zero-logprob
val/logprob/zero_frac Fraksi dari nol logprobs
val/logprob/zero_per_group Nol logprobs per grup
val/logprob/nz_mean Rata-rata logprob bukan nol
val/logprob/nz_std Standar deviasi logprobs bukan nol
val/logprob/nz_min Minimum logprob bukan nol
val/logprob/nz_max Maksimum logprob bukan nol

Mengakses UI MLFlow

Akses UI MLFlow melalui URL yang telah ditetapkan sebelumnya:

aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2

Salin AuthorizedUrl dari output ke browser Anda.

Lintasan dan jejak agen

Selama pelatihan, SageMaker AI mencatat setiap interaksi antara agen Anda dan model kebijakan sebagai lintasan - catatan lengkap dari satu peluncuran. Setiap lintasan menangkap setiap prompt yang dikirim ke model, setiap respons yang dihasilkan, setiap panggilan alat yang dilakukan, dan hadiah akhir. Lintasan dipublikasikan ke eksperimen MLFlow Anda sebagai jejak terstruktur.

Melacak konten

  • Permintaan input dari kumpulan data pelatihan Anda

  • Setiap giliran inferensi model (prompt, respons, dan data tingkat token)

  • Panggilan alat dan hasilnya, jika agen Anda menggunakan alat

  • Skor hadiah akhir

  • Informasi waktu untuk setiap belokan

Melihat lintasan di UI MLFlow

Akses UI MLFlow melalui URL yang telah ditetapkan sebelumnya:

aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2

Salin AuthorizedUrl dari output ke browser Anda.

Buka UI MLFlow menggunakan URL presigned di atas. Arahkan ke proses eksperimen Anda dan pilih tab Jejak. Setiap jejak mewakili satu peluncuran yang telah selesai dan menunjukkan:

  • Prompt sistem dan prompt pengguna

  • Setiap respons asisten (dengan thinking/reasoning jika ada)

  • Rentang penggunaan alat yang menunjukkan alat mana yang dipanggil dan keluarannya

  • Skor hadiah diberikan ke lintasan

Gunakan lintasan untuk men-debug skor hadiah rendah

Gejala Apa yang harus dicari
Hadiah rendah di sebagian besar peluncuran Apakah respons model koheren? Apakah format prompt benar?
Tool-related kegagalan Apakah panggilan alat berhasil? Apakah input dan output terbentuk dengan baik?
Perulangan agen Apakah agen mengulangi tindakan yang sama tanpa membuat kemajuan?
Tanggapan terpotong Apakah tanggapan terputus oleh batas maxTokens?

Dapatkan Hasil Pelatihan

Saat pekerjaan pelatihan selesai, bobot model terlatih Anda disimpan sebagai Paket Model SageMaker AI. Bagian ini menjelaskan cara menemukan hasil Anda, memahami jenis pos pemeriksaan yang dihasilkan selama pelatihan, dan menggunakannya untuk penerapan atau pelatihan lanjutan.

Bagaimana hasil disimpan

SageMaker AI menyimpan output pelatihan sebagai Paket Model berversi dan tidak dapat diubah di dalam Grup Paket Model. Multi-turn RL menggunakan dua grup terpisah, yang Anda tentukan saat membuat pekerjaan:

Kelompok Tujuan Daftar Isi
Grup Package Model Output Model terlatih terakhir HuggingFace-compatible Bobot adaptor LoRa (adapter_config.json+adapter_model.safetensors)
Grup Paket Model Pos Pemeriksaan Menengah Status pelatihan yang dapat dilanjutkan Bobot adaptor LoRa+status pengoptimal+metadata langkah pelatihan

Konfigurasikan kedua grup di 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" }

Jenis pos pemeriksaan

Pelatihan menghasilkan dua jenis pos pemeriksaan, disimpan di setiap langkah pelatihan:

Pos pemeriksaan model (hanya bobot)

  • Disimpan di Output Model Package Group

  • Berisi bobot HuggingFace-compatible adaptor LoRa dalam format SafeTensors

  • Gunakan untuk inferensi, penyebaran, atau sebagai titik awal untuk pekerjaan pelatihan baru

  • Dibuat di setiap langkah, saat penyelesaian pekerjaan, dan saat pekerjaan dihentikan

Pos pemeriksaan yang dapat dilanjutkan (status penuh)

  • Disimpan di Grup Package Model Checkpoint Menengah

  • Berisi bobot adaptor LoRa, status pengoptimal, dan metadata langkah pelatihan per GPU

  • Gunakan untuk melanjutkan pekerjaan yang terputus dari langkah yang tepat yang dihentikan

  • Format internal — tidak dapat digunakan secara langsung untuk inferensi

Siklus hidup pos pemeriksaan

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

Ambil model terlatih Anda

Ketika pekerjaan selesai dengan sukses, model akhir disimpan sebagai Model Package di Output Model Package Group. OutputModelPackageArnBidang pada catatan pekerjaan berisi ARN.

Periksa penyelesaian pekerjaan dan ambil model keluaran ARN:

aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2

Cari OutputModelPackageArn dalam tanggapannya. Gunakan untuk mendeskripsikan Model Package dan dapatkan lokasi S3 dari bobot:

aws sagemaker describe-model-package \ --model-package-name "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-final-models/5"

Jika pekerjaan gagal atau dihentikan sebelum selesai, pos pemeriksaan perantara terakhir dipromosikan ke Output Model Package Group berdasarkan upaya terbaik. Periksa OutputModelPackageArn dengan cara yang sama.

Untuk memantau pembuatan pos pemeriksaan selama pelatihan, perhatikan ResumableCheckpoint dan ModelCheckpoint bidang dalam DescribeJob output.

Lanjutkan pekerjaan yang terputus

Jika suatu pekerjaan gagal atau dihentikan di tengah pelatihan, Anda dapat memulai pekerjaan baru yang diambil dari langkah yang tepat di mana ia tinggalkan. Platform ini mengembalikan status pelatihan penuh — bobot, momentum pengoptimal, dan penghitung langkah — dari pos pemeriksaan yang dapat dilanjutkan.

Tentukan pos pemeriksaan yang dapat dilanjutkan dari Grup Paket Model Checkpoint Menengah sebagai: 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" }

InputModelPackageArnHarus menunjuk ke pos pemeriksaan yang dapat dilanjutkan (satu dengan IsCheckpoint=true metadata Model Package). Pelatihan dilanjutkan dari langkah setelah pos pemeriksaan - misalnya, jika pos pemeriksaan disimpan pada langkah 4, pelatihan berlanjut dari langkah 5.

Berikut ini harus tetap sama antara pekerjaan asli dan pekerjaan yang dilanjutkan:

  • Model dasar

  • Konfigurasi LoRa (peringkat dan alfa)

  • Hyperparameter (tingkat pembelajaran, ukuran batch, dll.)

  • Set data

Lanjutkan pelatihan pada pekerjaan baru (pelatihan berulang)

Pelatihan berulang memungkinkan Anda membangun model yang dilatih sebelumnya dengan kumpulan data yang berbeda, hiperparameter yang berbeda, atau fungsi hadiah yang disempurnakan. Tidak seperti melanjutkan, ini memulai latihan baru — pengoptimal mengatur ulang, penghitung langkah disetel ulang ke 0, dan hanya bobot LoRa yang terlatih yang terbawa.

Tentukan pos pemeriksaan model dari Output Model Package Group sebagaiInputModelPackageArn:

"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" }

Apa yang dapat Anda ubah di antara iterasi:

  • Hyperparameters (tingkat pembelajaran, ukuran batch, max_steps, group_size, dll.)

  • Dataset (petunjuk atau distribusi data yang berbeda)

  • Fungsi skor

  • Konfigurasi agen

Apa yang harus tetap sama:

  • Model dasar — adaptor LoRa terikat pada arsitektur model dasar

Pola umum untuk pelatihan berulang:

  • Pembelajaran kurikulum — latih masalah yang lebih mudah terlebih dahulu, kemudian lanjutkan masalah yang lebih sulit

  • Penyempurnaan hadiah — mulai dengan fungsi hadiah sederhana, lalu ulangi dengan yang lebih bernuansa

  • Penyesuaian hyperparameter — tingkatkan ukuran batch atau tune learning rate setelah mengamati dinamika pelatihan awal

Praktik terbaik pos pemeriksaan

  • Pantau pembuatan pos pemeriksaan. Gunakan DescribeJob untuk melacak ResumableCheckpoint dan ModelCheckpoint bidang selama pelatihan sehingga Anda tahu apa yang tersedia jika Anda perlu melanjutkan.

  • Rencanakan kegagalan pada pekerjaan panjang. Jika pekerjaan memiliki banyak langkah, rancang alur kerja Anda untuk dilanjutkan dari pos pemeriksaan daripada memulai ulang dari awal.