View a markdown version of this page

Invio di un lavoro di formazione - Amazon SageMaker AI

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 MlflowResourceArn nella 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 ResumableCheckpoint durante ModelCheckpoint l'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.