Le traduzioni sono generate tramite traduzione automatica. In caso di conflitto tra il contenuto di una traduzione e la versione originale in Inglese, quest'ultima prevarrà.
Invio di un lavoro di formazione
Avvio di lavori di formazione
Dopo aver distribuito l'agente e aver installato il set di dati in S3, crea un processo di formazione utilizzando uno dei seguenti metodi.
SageMaker AI Studio
-
Vai a Modelli nel pannello di navigazione e seleziona Modelli JumpStart base.
-
Seleziona un modello che supporti l'RL multigiro (consulta la tabella dei modelli supportati) e scegli Personalizza modello, quindi Personalizza con interfaccia utente.
-
Seleziona Multi-Turn Reinforcement Learning come tecnica di personalizzazione.
-
Configura l'ambiente del tuo agente: seleziona il AgentCore runtime Bedrock o fornisci l'ARN del tuo spedizioniere Lambda.
-
Fornisci il tuo set di dati di addestramento come URI S3 o set di dati registrato.
-
Regola gli iperparametri secondo necessità.
-
Controlla la configurazione e scegli Invia.
SageMaker SDK per AI per Python
Scopri i modelli supportati
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}")
Configura l'ambiente del tuo agente
Opzione 1: Bedrock runtime AgentCore
# List available runtimes runtimes = MultiTurnRLTrainer.list_bedrock_agentcore_runtimes() for rt in runtimes: print(f" - {rt['name']} ({rt['status']}) → {rt['arn']}")
Opzione 2: agente Lambda personalizzato
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")
Registra il tuo set di dati (opzionale)
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}")
Crea un gruppo di pacchetti di modelli con restrizioni per Nova (opzionale)
Se scegli Nova model (nova-textgeneration-lite-v2), crea facoltativamente Restricted Model Package Group prima di inviare un lavoro di formazione (passaggio successivo). Se salti questo passaggio, l'SDK ne crea automaticamente uno per te.
Restricted Model Package Group (RMPG) è un Model Package Group con ManagedStorageType: Restricted. È necessario per i modelli closed-source come Nova, in cui i pesi dei modelli sono gestiti dal cliente AWS e non sono direttamente accessibili al cliente.
Lo schema di lavoro RFT richiede due MPG con restrizioni separati:
-
Output MPG: memorizza il pacchetto modello finale ottimizzato
-
Intermediate Checkpoint MPG — riservato ai checkpoint di allenamento intermedi (deve essere diverso dall'Output 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)
Una volta creato il gruppo di pacchetti Model, passate i gruppi nel passaggio successivo quando inviate un lavoro di formazione.
Invia un lavoro di formazione con 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}")
Invia un lavoro di formazione con un agente Lambda personalizzato
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}")
Invia un lavoro di formazione con il gruppo di pacchetti Restricted Model per Nova
Fai riferimento al passaggio precedente (Create Restricted Model Package Group for Nova) su come creare un gruppo 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
Crea un lavoro di formazione utilizzando l'CreateJobAPI. È possibile specificare la configurazione dell'agente, la posizione dei dati di addestramento, il modello di base e le impostazioni di output inJobConfigDocument.
Per recuperare lo JobConfigDocument schema completo:
aws sagemaker list-job-schema-versions --job-category AgentRFT aws sagemaker describe-job-schema-version --job-category AgentRFT --version "1.0.0"
Crea lavoro con 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
Crea un lavoro con un agente Lambda personalizzato
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
Crea lavoro con 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']}")
Crea un lavoro con un agente Lambda personalizzato
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']}")
Formazione sul monitoraggio
Monitora il tuo Training Job
Utilizza l'DescribeJobAPI per controllare lo stato attuale del tuo lavoro in qualsiasi momento. Lo stato del lavoro passa daInProgress, e poi aCompleted, Failed oStopped.
aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2
Usa l'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}")
Monitora la formazione in MLFlow
SageMaker L'intelligenza artificiale si integra automaticamente con MLFlow gestito per tenere traccia dei progressi, delle metriche e degli artefatti del processo di formazione. Per abilitare il tracciamento MLFlow, includi nel tuo lavoro: MlflowConfig OutputDataConfig
"OutputDataConfig": { "S3OutputPath": "s3://your-bucket/output/", "MlflowConfig": { "MlflowResourceArn": "arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/my-rft-mlflow-app" } }
Prerequisiti
-
Crea un'app MLFlow gestita nel tuo account. Per istruzioni di configurazione, consulta Configurazione dell'app MLFlow.
-
Assicurati che il tuo ruolo di esecuzione SageMaker AI disponga delle autorizzazioni per scrivere sull'app MLFlow (
sagemaker-mlflow:*azioni). -
Includili
MlflowResourceArnnella configurazione del tuo lavoro.
Cosa viene registrato
| # | Categoria | Cosa viene registrato | Dove nell'interfaccia utente MLFlow |
|---|---|---|---|
| 1 | Metriche di addestramento | Per-step contatori, velocità effettiva, contabilizzazione dei dati e dei token, durata continua di ogni fase di un passaggio, batch riepilogativo di implementazione tra traiettoria e ricompensa e distribuzione del conteggio dei turni per traiettoria | Scheda Metriche (grafici delle serie temporali) |
| 2 | Tracce di traiettoria | Conversazioni complete a più turni con chiamate e premi | Scheda Tracce |
Riferimento dettagliato alle metriche di allenamento
Le seguenti metriche vengono registrate in ogni fase dell'allenamento.
Contatori di passi e produttività () training/
| Metrica | Description |
|---|---|
training/epoch |
Numero dell'epoca attuale |
training/global_step |
Contapassi di formazione globale |
training/num_groups |
Gruppi di traiettorie in questa fase |
training/num_trajectories |
Traiettorie totali elaborate in questa fase |
training/total_tokens |
I token sono stati sommati in tutti i microbatch in questa fase |
training/num_datums |
Datum di addestramento formati da traiettorie |
training/datums_per_trajectory |
Dati medi emessi per traiettoria |
training/action_tokens_mean |
Token medi di azione (risposta) per traiettoria |
training/obs_tokens_mean |
Token medi di osservazione (prompt) per traiettoria |
training/trainable_token_positions |
Posizioni target totali addestrabili in questa fase |
training/nontrainable_token_positions |
Posizioni bersaglio totali non allenabili in questa fase |
training/trainable_token_ratio |
Rapporto: trainable / (trainable + nontrainable) posizioni dei token |
Durate delle fasi () timing_s/
| Metrica | Description |
|---|---|
timing_s/step |
Tempo totale per l'intero passaggio |
timing_s/training |
Tempo di esecuzione dei forward/backward passaggi e fase di ottimizzazione |
timing_s/policy_update |
Risparmio di tempo sui pesi aggiornati per il campionatore |
timing_s/save_checkpoint |
Risparmio di tempo in un checkpoint (solo nelle fasi del checkpoint) |
timing_s/eval |
Valutazione del tempo di esecuzione (solo nelle fasi di valutazione) |
Distribuzione dei premi () rollout/reward/
| Metrica | Description |
|---|---|
rollout/reward/mean |
Traiettoria media (ricompensa) in tutti i gruppi |
rollout/reward/valid_mean |
Ricompensa media relativa solo ai gruppi validi (con vantaggio diverso da zero); uguale quando non si è verificato alcun filtraggio mean |
rollout/reward/std |
Deviazione standard delle ricompense della traiettoria |
rollout/reward/min |
Ricompensa minima sulla traiettoria |
rollout/reward/max |
Ricompensa massima sulla traiettoria |
rollout/reward/zero_frac |
Frazione di traiettorie con ricompensa totale esattamente 0,0 |
I turni contano () rollout/turns/
| Metrica | Description |
|---|---|
rollout/turns/mean |
Svolte medie (transizioni) per traiettoria |
rollout/turns/min |
Numero minimo di giri tra le traiettorie |
rollout/turns/max |
Numero massimo di giri tra le traiettorie |
Lunghezze dei token () rollout/tokens/
| Metrica | Description |
|---|---|
rollout/tokens/prompt_mean |
Numero medio di token rapidi per transizione |
rollout/tokens/response_mean |
Numero medio di token di risposta per transizione |
rollout/tokens/response_std |
Deviazione standard del conteggio dei token di risposta |
rollout/tokens/response_min |
Token di risposta minimi |
rollout/tokens/response_max |
Token di risposta massimi (attenzione al clustering su) sampling_max_tokens |
Log-probability salute () rollout/logprob/
| Metrica | Description |
|---|---|
rollout/logprob/zero_count |
Token zero-logprob totali |
rollout/logprob/zero_frac |
Frazione di tutti i logprob che è esattamente 0,0 |
rollout/logprob/zero_per_group |
In media zero logprob per gruppo di traiettorie |
rollout/logprob/nz_mean |
Media dei logprob diversi da zero |
rollout/logprob/nz_std |
Deviazione standard di logprob diversi da zero |
rollout/logprob/nz_min |
Logprob minimo diverso da zero |
rollout/logprob/nz_max |
Logprob massimo diverso da zero |
Distribuzione dei vantaggi () rollout/advantage/
| Metrica | Description |
|---|---|
rollout/advantage/mean |
Valore medio del vantaggio in tutte le transizioni |
rollout/advantage/std |
Deviazione standard dei vantaggi |
rollout/advantage/min |
Vantaggio minimo |
rollout/advantage/max |
Vantaggio massimo |
rollout/advantage/n_positive |
Transizioni con vantaggio positivo |
rollout/advantage/n_negative |
Transizioni con vantaggio negativo |
Batch-quality classificazione () analysis/
| Metrica | Description |
|---|---|
analysis/batch_completion_ratio |
total_completed / batch_size— frazione dei gruppi previsti arrivati |
analysis/batch_valid_ratio |
valid_count / batch_size— gruppi con vantaggio diverso da zero rispetto all'intero lotto |
analysis/zero_adv_groups |
Gruppi in cui tutte le transizioni hanno un vantaggio vicino allo zero |
analysis/zero_adv_nonzero_reward |
Zero-advantage gruppi in cui almeno una transizione ha una ricompensa diversa da 0 (caso corretto per le ricompense binarie) |
analysis/zero_adv_zero_reward |
Zero-advantage gruppi in cui tutte le ricompense sono 0 (caso completamente sbagliato) |
analysis/reward_variance_across_groups |
Varianza delle ricompense medie per gruppo (alta = lotto diverso) |
analysis/mean_group_reward_spread |
Diffusione media delle ricompense all'interno del gruppo max - min |
Premio di valutazione e pass @k () val/reward/
Emesso alla linea di base (fase 0), ad ogni val_every intervallo e nella fase finale. Include le stesse metriche di distribuzione rollout/reward più le metriche di ricompensa di gruppo aggregate per prompt.
Distribuzione:
| Metrica | Description |
|---|---|
val/reward/mean |
Ricompensa media rispetto al set di valutazione |
val/reward/std |
Ricompensa std dev |
val/reward/min |
Ricompensa minima |
val/reward/max |
Ricompensa massima |
val/reward/zero_frac |
Frazione di traiettorie a ricompensa zero |
Group-reward (aggregazione per prompt):
| Metrica | Description |
|---|---|
val/reward/min_within_groups |
Ricompensa minima media per richiesta |
val/reward/mean_within_groups |
Ricompensa media media per richiesta |
val/reward/max_within_groups |
Ricompensa massima media per richiesta |
val/reward/std_within_groups |
Ricompensa media per richiesta std (coerenza) |
val/reward/rollouts_per_prompt |
Implementazioni medie (n) tra i prompt |
val/reward/num_prompts |
Sono stati valutati prompt distinti |
Pass @k e contabilità di successo:
| Metrica | Description |
|---|---|
val/reward/succeeded_rollouts |
Implementazioni totali con ricompensa ≥ success_threshold |
val/reward/failed_rollouts |
Implementazioni totali con ricompensa < success_threshold |
val/reward/success_threshold |
Soglia utilizzata (riportata per maggiore chiarezza) |
val/reward/pass_at_{k} |
Probabilità ≥ 1% di k campioni superati |
val/reward/pass_power_{k} |
Probabilità che tutti i k campioni vengano superati (affidabilità) |
I turni di valutazione contano (val/turns/)
| Metrica | Description |
|---|---|
val/turns/mean |
Giri medi per traiettoria nominale |
val/turns/min |
Numero minimo di giri |
val/turns/max |
Numero massimo di giri |
Lunghezze dei token di valutazione () val/tokens/
| Metrica | Description |
|---|---|
val/tokens/prompt_mean |
Token prompt medi per transizione |
val/tokens/response_mean |
Token di risposta medi per transizione |
val/tokens/response_std |
Deviazione standard dei token di risposta |
val/tokens/response_min |
Token di risposta minimi |
val/tokens/response_max |
Token di risposta massimi |
Salute log-probabilistica di valutazione () val/logprob/
| Metrica | Description |
|---|---|
val/logprob/zero_count |
Token zero-logprob totali |
val/logprob/zero_frac |
Frazione di zero logprob |
val/logprob/zero_per_group |
Zero logprobs per gruppo |
val/logprob/nz_mean |
Media dei logprob diversi da zero |
val/logprob/nz_std |
Deviazione standard di logprob diversi da zero |
val/logprob/nz_min |
Logprob minimo diverso da zero |
val/logprob/nz_max |
Logprob massimo diverso da zero |
Accesso all'interfaccia utente MLFlow
Accedi all'interfaccia utente MLFlow tramite un URL predefinito:
aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2
Copia il file AuthorizedUrl dall'output nel tuo browser.
Tracce e traiettorie degli agenti
Durante la formazione, l' SageMaker intelligenza artificiale registra ogni interazione tra l'agente e il modello di policy come traiettoria, il record completo di un'unica implementazione. Ogni traiettoria registra ogni richiesta inviata al modello, ogni risposta generata, ogni intervento effettuato e la ricompensa finale. Le traiettorie vengono pubblicate nell'esperimento MLFlow come tracce strutturate.
Contenuti delle tracce
-
La richiesta di input del set di dati di allenamento
-
Ogni turno di inferenza del modello (richiesta, risposta e dati a livello di token)
-
Chiamate agli strumenti e relativi risultati, se l'agente utilizza strumenti
-
Il punteggio finale della ricompensa
-
Informazioni sulla tempistica per ogni turno
Visualizzazione delle traiettorie nell'interfaccia utente MLFlow
Accedi all'interfaccia utente MLFlow tramite un URL predefinito:
aws sagemaker create-presigned-mlflow-app-url \ --arn arn:aws:sagemaker:us-west-2:123456789012:mlflow-app/mlflow-app-id \ --region us-west-2
Copia il file AuthorizedUrl dall'output nel tuo browser.
Apri l'interfaccia utente MLFlow utilizzando l'URL predefinito riportato sopra. Vai all'esecuzione dell'esperimento e seleziona la scheda Traces. Ogni traccia rappresenta un'implementazione completata e mostra:
-
Il prompt di sistema e il prompt dell'utente
-
Ogni risposta dell'assistente (con thinking/reasoning se applicabile)
-
Intervalli di utilizzo degli strumenti che mostrano quali strumenti sono stati chiamati e i relativi risultati
-
Il punteggio di ricompensa assegnato alla traiettoria
Usa le traiettorie per eseguire il debug dei punteggi di ricompensa bassi
| Caratteristiche | Cosa cercare |
|---|---|
| Ricompense basse nella maggior parte delle implementazioni | Le risposte dei modelli sono coerenti? Il formato del prompt è corretto? |
| Tool-related fallimenti | Le chiamate agli strumenti hanno esito positivo? Gli input e gli output sono ben formati? |
| Agente in loop | L'agente ripete le stesse azioni senza fare progressi? |
| Risposte troncate | Le risposte vengono interrotte dal limite MaxTokens? |
Ottieni risultati di formazione
Al termine di un processo di formazione, i pesi del modello addestrato vengono archiviati come SageMaker AI Model Package. Questa sezione spiega come trovare i risultati, comprendere i tipi di checkpoint prodotti durante la formazione e utilizzarli per l'implementazione o la formazione continua.
Come vengono archiviati i risultati
SageMaker L'intelligenza artificiale memorizza i risultati della formazione come pacchetti modello con versioni e immutabili all'interno dei Model Package Groups. Multi-turn RL utilizza due gruppi separati, che vengono specificati durante la creazione di un lavoro:
| Gruppo | Scopo | Indice |
|---|---|---|
| Output Model Package Group | Modello finale addestrato | HuggingFace-compatible Pesi degli adattatori LoRa (adapter_config.json + adapter_model.safetensors) |
| Intermediate Checkpoint Model Package Group | Stato di allenamento riprendibile | Pesi dell'adattatore LoRa + stati dell'ottimizzatore + metadati delle fasi di allenamento |
Configura entrambi i gruppi nel tuo: 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" }
Tipi di checkpoint
La formazione produce due tipi di checkpoint, salvati in ogni fase dell'allenamento:
Modello di checkpoint (solo pesi)
-
Memorizzato nell'Output Model Package Group
-
Contiene i pesi degli adattatori HuggingFace-compatible LoRa in formato SafeTensors
-
Utilizzalo per inferenza, implementazione o come punto di partenza per un nuovo lavoro di formazione
-
Creato in ogni fase, al completamento del lavoro e quando un lavoro viene interrotto
Checkpoint ripristinabile (stato completo)
-
Memorizzato nell'Intermediate Checkpoint Model Package Group
-
Contiene i pesi degli adattatori LoRa, gli stati dell'ottimizzatore e i metadati delle fasi di addestramento per GPU
-
Si usa per riprendere un lavoro interrotto dalla fase esatta in cui è stato interrotto
-
Formato interno: non utilizzabile direttamente per l'inferenza
Ciclo di vita di Checkpoint
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
Recupera il tuo modello addestrato
Quando un lavoro viene completato correttamente, il modello finale viene salvato come Model Package nell'Output Model Package Group. Il OutputModelPackageArn campo nel record del lavoro contiene l'ARN.
Controlla il completamento del lavoro e recupera il modello di output ARN:
aws sagemaker describe-job \ --job-name "my-agent-rft-job" \ --job-category AgentRFT \ --region us-west-2
Cerca OutputModelPackageArn nella risposta. Usalo per descrivere il Model Package e ottenere la posizione S3 dei pesi:
aws sagemaker describe-model-package \ --model-package-name "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-final-models/5"
Se un job fallisce o viene interrotto prima del completamento, l'ultimo checkpoint intermedio viene promosso all'Output Model Package Group con la massima diligenza possibile. Controlla OutputModelPackageArn allo stesso modo.
Per monitorare la creazione dei checkpoint durante l'allenamento, guarda i ModelCheckpoint campi ResumableCheckpoint and in DescribeJob output.
Riprendere un lavoro interrotto
Se un lavoro fallisce o viene interrotto durante la fase di formazione, puoi iniziare un nuovo lavoro che riprende esattamente dalla fase in cui era stato interrotto. La piattaforma ripristina l'intero stato di allenamento (pesi, ottimizzatore dello slancio e contapassi) dal checkpoint riutilizzabile.
Specificare un checkpoint ripristinabile dall'Intermediate Checkpoint Model Package Group come: 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" }
InputModelPackageArnDeve puntare a un checkpoint ripristinabile (uno con IsCheckpoint=true i metadati del Model Package). L'allenamento riprende dalla fase successiva al checkpoint: ad esempio, se il checkpoint è stato salvato nella fase 4, la formazione continua dalla fase 5.
Quanto segue deve rimanere lo stesso tra il lavoro originale e il lavoro ripreso:
-
Modelli base
-
Configurazione LoRa (rank e alpha)
-
Iperparametri (tasso di apprendimento, dimensione del batch, ecc.)
-
Set di dati
Continuare la formazione su un nuovo lavoro (formazione iterativa)
La formazione iterativa consente di basarsi su un modello precedentemente addestrato con un set di dati diverso, diversi iperparametri o una funzione di ricompensa perfezionata. A differenza della ripresa, questo avvia una nuova sessione di allenamento: l'ottimizzatore si resetta, il contapassi si resetta a 0 e vengono trasferiti solo i pesi LoRa allenati.
Specificate un checkpoint del modello dall'Output Model Package Group comeInputModelPackageArn:
"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" }
Cosa puoi cambiare tra le iterazioni:
-
Iperparametri (tasso di apprendimento, dimensione del batch, max_steps, group_size, ecc.)
-
Set di dati (diverse richieste o distribuzione dei dati)
-
Funzione di ricompensa
-
Configurazione dell'agente
Cosa deve rimanere uguale:
-
Modello base: l'adattatore LoRa è legato all'architettura del modello base
Modelli comuni per la formazione iterativa:
-
Apprendimento mediante programmi di studio: allenati prima sui problemi più semplici, poi prosegui su quelli più difficili
-
Perfezionamento delle ricompense: inizia con una semplice funzione di ricompensa, poi ripeti con una più sfumata
-
Regolazione degli iperparametri: aumenta la dimensione del batch o regola la velocità di apprendimento dopo aver osservato le dinamiche di allenamento iniziali
Le migliori pratiche di Checkpoint
-
Monitora la creazione di checkpoint. Utilizzatelo DescribeJob per tenere traccia
ResumableCheckpointduranteModelCheckpointl'allenamento, in modo da sapere cosa c'è a disposizione in caso di necessità di riprendere il corso. -
Pianifica gli insuccessi in caso di lavori di lunga durata. Se un lavoro prevede molti passaggi, progetta il flusso di lavoro in modo che riprenda dai checkpoint anziché riavviarlo da zero.