View a markdown version of this page

Envio de trabalhos de treinamento - SageMaker IA da Amazon

As traduções são geradas por tradução automática. Em caso de conflito entre o conteúdo da tradução e da versão original em inglês, a versão em inglês prevalecerá.

Envio de trabalhos de treinamento

Lançamento de trabalhos de treinamento

Depois que seu agente for implantado e seu conjunto de dados estiver no S3, crie um trabalho de treinamento usando um dos métodos a seguir.

SageMaker Estúdio AI

  • Navegue até Modelos no painel de navegação e selecione Modelos JumpStart básicos.

  • Selecione um modelo que suporte RL de várias voltas (consulte a tabela de modelos compatíveis) e escolha Personalizar modelo e, em seguida, Personalizar com interface do usuário.

  • Selecione Multi-Turn Aprendizado por Reforço como técnica de personalização.

  • Configure seu ambiente de agente — selecione seu tempo de AgentCore execução do Bedrock ou forneça o ARN do encaminhador Lambda.

  • Forneça seu conjunto de dados de treinamento como um URI do S3 ou conjunto de dados registrado.

  • Ajuste os hiperparâmetros conforme necessário.

  • Revise sua configuração e escolha Enviar.

SageMaker SDK AI para Python

Descubra os modelos compatíveis

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

Configure seu ambiente de agente

Opção 1: tempo de execução do Bedrock AgentCore

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

Opção 2: agente Lambda personalizado

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

Registre seu conjunto de dados (opcional)

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

Crie um grupo restrito de pacotes de modelos para Nova (opcional)

Se você estiver escolhendo o modelo Nova (nova-textgeneration-lite-v2), crie opcionalmente o Restricted Model Package Group antes de enviar um trabalho de treinamento (próxima etapa). Se você pular essa etapa, o SDK criará automaticamente uma para você.

O Restricted Model Package Group (RMPG) é um grupo de pacotes de modelos com ManagedStorageType: Restricted. É necessário para modelos de código fechado, como o Nova, em que os pesos do modelo são gerenciados pelo cliente AWS e não estão diretamente acessíveis ao cliente.

O esquema de trabalho do RFT requer dois MPGs restritos separados:

  • Output MPG — armazena o pacote final do modelo ajustado

  • Ponto de verificação intermediário MPG — reservado para pontos de verificação de treinamento intermediário (deve ser diferente do MPG de saída)

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)

Depois que o grupo de pacotes do modelo for criado, passe os grupos na próxima etapa ao enviar um trabalho de treinamento.

Envie um trabalho de treinamento com a 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}")

Envie um trabalho de treinamento com um agente Lambda personalizado

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

Envie um trabalho de treinamento com o grupo de pacotes Restricted Model para Nova

Consulte a etapa acima (Create Restricted Model Package Group for Nova) para saber como criar um grupo de 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

Crie um trabalho de treinamento usando a CreateJob API. Você especifica a configuração do agente, a localização dos dados de treinamento, o modelo básico e as configurações de saída noJobConfigDocument.

Para recuperar o JobConfigDocument esquema completo:

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

Crie emprego com o 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

Crie um trabalho com um agente Lambda personalizado

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

Crie emprego com o 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']}")

Crie um trabalho com um agente Lambda personalizado

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

Treinamento de monitoramento

Monitore seu Training Job

Use a DescribeJob API para verificar o status atual do seu trabalho a qualquer momento. O status do trabalho passa porInProgress, e depois paraCompleted, Failed ouStopped.

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

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

Monitore o treinamento no MLflow

SageMaker A IA se integra automaticamente ao MLflow gerenciado para monitorar o progresso, as métricas e os artefatos do seu trabalho de treinamento. Para habilitar o rastreamento do MLflow, inclua um MlflowConfig em seu trabalho: OutputDataConfig

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

Pré-requisitos

  • Crie um aplicativo MLflow gerenciado em sua conta. Para obter instruções de configuração, consulte Configuração do aplicativo MLflow.

  • Certifique-se de que sua função de execução de SageMaker IA tenha permissões para gravar no aplicativo MLflow (sagemaker-mlflow:*ações).

  • Inclua o MlflowResourceArn em sua configuração de trabalho.

O que é registrado

# Categoria O que está registrado Onde está o MLFlow UI
1 Métricas de treinamento Per-step contadores, produtividade, contabilização de dados e tokens, duração total de cada fase de uma etapa, resumo da trajetória-recompensa, lote de lançamento e distribuição da contagem de turnos por trajetória Aba Métricas (gráficos de séries temporais)
2 Traços de trajetória Conversas completas em vários turnos com chamadas de ferramentas e recompensas Aba Traços

Referência detalhada de métricas de treinamento

As métricas a seguir são registradas em cada etapa do treinamento.

Contadores de etapas e taxa de transferência () training/

Métrica Description
training/epoch Número da época atual
training/global_step Contador global de etapas de treinamento
training/num_groups Grupos de trajetória nesta etapa
training/num_trajectories Total de trajetórias processadas nesta etapa
training/total_tokens Tokens somados em todos os microlotes nesta etapa
training/num_datums Dados de treinamento formados a partir de trajetórias
training/datums_per_trajectory Dados médios emitidos por trajetória
training/action_tokens_mean Tokens médios de ação (resposta) por trajetória
training/obs_tokens_mean Tokens médios de observação (imediata) por trajetória
training/trainable_token_positions Total de posições-alvo treináveis nesta etapa
training/nontrainable_token_positions Total de posições-alvo não treináveis nesta etapa
training/trainable_token_ratio Proporção: posições de trainable / (trainable + nontrainable) token

Durações de fase () timing_s/

Métrica Description
timing_s/step Tempo total para a etapa completa
timing_s/training Hora dos forward/backward passes e etapa do otimizador
timing_s/policy_update Pesos atualizados que economizam tempo para o amostrador
timing_s/save_checkpoint Economizando tempo em um ponto de verificação (somente nas etapas do ponto de verificação)
timing_s/eval Avaliação do tempo de execução (somente nas etapas de avaliação)

Distribuição de recompensas (rollout/reward/)

Métrica Description
rollout/reward/mean Recompensa média da trajetória em todos os grupos
rollout/reward/valid_mean Recompensa média somente sobre os grupos válidos (sem vantagem zero); é igual mean quando nenhuma filtragem ocorreu
rollout/reward/std Desvio padrão das recompensas da trajetória
rollout/reward/min Recompensa mínima de trajetória
rollout/reward/max Recompensa máxima de trajetória
rollout/reward/zero_frac Fração de trajetórias com recompensa total de exatamente 0,0

Contagens de turnos (rollout/turns/)

Métrica Description
rollout/turns/mean Média de curvas (transições) por trajetória
rollout/turns/min Curvas mínimas em trajetórias
rollout/turns/max Voltas máximas em trajetórias

Comprimentos de token () rollout/tokens/

Métrica Description
rollout/tokens/prompt_mean Contagem média de tokens imediatos por transição
rollout/tokens/response_mean Contagem média de tokens de resposta por transição
rollout/tokens/response_std Desvio padrão da contagem de tokens de resposta
rollout/tokens/response_min Tokens de resposta mínimos
rollout/tokens/response_max Tokens de resposta máxima (observe o agrupamento em) sampling_max_tokens

Log-probability saúde (rollout/logprob/)

Métrica Description
rollout/logprob/zero_count Total de tokens zero-logprob
rollout/logprob/zero_frac Fração de todos os logprobs que são exatamente 0,0
rollout/logprob/zero_per_group Média de zero logprobs por grupo de trajetória
rollout/logprob/nz_mean Média de sondas logarítmicas diferentes de zero
rollout/logprob/nz_std Desvio padrão de logprobs diferentes de zero
rollout/logprob/nz_min Logprob mínimo diferente de zero
rollout/logprob/nz_max Logprob máximo diferente de zero

Distribuição de vantagens (rollout/advantage/)

Métrica Description
rollout/advantage/mean Valor médio da vantagem em todas as transições
rollout/advantage/std Desvio padrão das vantagens
rollout/advantage/min Vantagem mínima
rollout/advantage/max Vantagem máxima
rollout/advantage/n_positive Transições com vantagem positiva
rollout/advantage/n_negative Transições com vantagem negativa

Batch-quality classificação (analysis/)

Métrica Description
analysis/batch_completion_ratio total_completed / batch_size— fração dos grupos esperados que chegaram
analysis/batch_valid_ratio valid_count / batch_size— grupos de vantagem sem zero em relação ao lote completo
analysis/zero_adv_groups Grupos em que todas as transições têm vantagens quase nulas
analysis/zero_adv_nonzero_reward Zero-advantage grupos em que pelo menos uma transição tem recompensa diferente de 0 (caso correto para recompensas binárias)
analysis/zero_adv_zero_reward Zero-advantage grupos em que todas as recompensas são 0 (caso totalmente errado)
analysis/reward_variance_across_groups Variação das recompensas médias por grupo (alta = lote diverso)
analysis/mean_group_reward_spread Distribuição média de recompensas dentro do grupo max - min

Recompensa de avaliação e aprovação @k (val/reward/)

Emitido na linha de base (etapa 0), em cada val_every intervalo e na etapa final. Inclui as mesmas métricas de distribuição, rollout/reward além das métricas de recompensa de grupo agregadas por prompt.

Distribuição:

Métrica Description
val/reward/mean Recompensa média em relação ao conjunto de avaliação
val/reward/std Recompensa std dev
val/reward/min Recompensa mínima
val/reward/max Recompensa máxima
val/reward/zero_frac Fração de trajetórias de recompensa zero

Group-reward (agregação por prompt):

Métrica Description
val/reward/min_within_groups Recompensa mínima média por solicitação
val/reward/mean_within_groups Recompensa média por solicitação
val/reward/max_within_groups Recompensa máxima média por solicitação
val/reward/std_within_groups Padrão médio de recompensa por alerta (consistência)
val/reward/rollouts_per_prompt Implantações médias (n) em todos os prompts
val/reward/num_prompts Solicitações distintas avaliadas

Pass @k e a contabilidade do sucesso:

Métrica Description
val/reward/succeeded_rollouts Total de lançamentos com recompensa ≥ success_threshold
val/reward/failed_rollouts Total de lançamentos com recompensa < success_threshold
val/reward/success_threshold Limite usado (ecoado para maior clareza)
val/reward/pass_at_{k} Probabilidade ≥1 de k amostras passadas
val/reward/pass_power_{k} Probabilidade de todas as k amostras passarem (confiabilidade)

Contagens de turnos de avaliação (val/turns/)

Métrica Description
val/turns/mean Média de voltas por trajetória de avaliação
val/turns/min Voltas mínimas
val/turns/max Voltas máximas

Comprimentos do token de avaliação () val/tokens/

Métrica Description
val/tokens/prompt_mean Média de tokens imediatos por transição
val/tokens/response_mean Média de tokens de resposta por transição
val/tokens/response_std Desvio padrão dos tokens de resposta
val/tokens/response_min Tokens de resposta mínimos
val/tokens/response_max Tokens de resposta máxima

Avaliação da saúde log-probabilística () val/logprob/

Métrica Description
val/logprob/zero_count Total de tokens zero-logprob
val/logprob/zero_frac Fração de zero logprobs
val/logprob/zero_per_group Zero problemas de registro por grupo
val/logprob/nz_mean Média de sondas logarítmicas diferentes de zero
val/logprob/nz_std Desvio padrão de logprobs diferentes de zero
val/logprob/nz_min Logprob mínimo diferente de zero
val/logprob/nz_max Logprob máximo diferente de zero

Acessando a interface do usuário do MLflow

Acesse a interface do usuário do MLflow por meio de um URL pré-assinado:

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

Copie o AuthorizedUrl da saída para o seu navegador.

Trajetórias e traços do agente

Durante o treinamento, a SageMaker IA registra cada interação entre seu agente e o modelo de política como uma trajetória — o registro completo de uma implementação. Cada trajetória captura cada solicitação enviada ao modelo, cada resposta gerada, cada chamada de ferramenta feita e a recompensa final. As trajetórias são publicadas em seu experimento MLflow como traços estruturados.

Rastrear conteúdo

  • O prompt de entrada do seu conjunto de dados de treinamento

  • Cada turno de inferência do modelo (dados de solicitação, resposta e nível de token)

  • Chamadas de ferramentas e seus resultados, se seu agente usar ferramentas

  • A pontuação final da recompensa

  • Informações de tempo para cada turno

Visualizando trajetórias na interface do usuário do MLflow

Acesse a interface do usuário do MLflow por meio de um URL pré-assinado:

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

Copie o AuthorizedUrl da saída para o seu navegador.

Abra a interface do usuário do MLflow usando o URL pré-assinado acima. Navegue até a execução do seu experimento e selecione a guia Traços. Cada rastreamento representa uma implantação concluída e mostra:

  • O prompt do sistema e o prompt do usuário

  • Cada resposta do assistente (com, thinking/reasoning se aplicável)

  • Extensões de uso de ferramentas mostrando quais ferramentas foram chamadas e suas saídas

  • A pontuação da recompensa atribuída à trajetória

Use trajetórias para depurar baixas pontuações de recompensa

Sintomas O que procurar
Baixa recompensa na maioria dos lançamentos As respostas do modelo são coerentes? O formato do prompt está correto?
Tool-related fracassos As chamadas de ferramentas estão sendo bem-sucedidas? As entradas e saídas estão bem formadas?
Agente em loop O agente está repetindo as mesmas ações sem progredir?
Respostas truncadas As respostas estão sendo cortadas pelo limite de maxTokens?

Obtenha resultados de treinamento

Quando um trabalho de treinamento é concluído, os pesos do modelo treinado são armazenados como um SageMaker AI Model Package. Esta seção explica como encontrar seus resultados, entender os tipos de pontos de verificação produzidos durante o treinamento e usá-los para implantação ou treinamento contínuo.

Como os resultados são armazenados

SageMaker A IA armazena os resultados do treinamento como Pacotes de Modelos versionados e imutáveis dentro dos Model Package Groups. Multi-turn O RL usa dois grupos separados, que você especifica ao criar um trabalho:

Group (Grupo) Finalidade Conteúdo
Grupo de pacotes do modelo de saída Modelo final treinado HuggingFace-compatible Pesos do adaptador LoRa (adapter_config.json + adapter_model.safetensors)
Grupo de pacotes de modelos de ponto de verificação intermediário Estado de treinamento retomável Pesos do adaptador LoRa + estados do otimizador + metadados da etapa de treinamento

Configure os dois grupos em seu 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" }

Tipos de pontos de verificação

O treinamento produz dois tipos de pontos de verificação, salvos em cada etapa do treinamento:

Ponto de verificação do modelo (somente pesos)

  • Armazenado no Output Model Package Group

  • Contém pesos do adaptador HuggingFace-compatible LoRa em formato SafeTensors

  • Use para inferência, implantação ou como ponto de partida para um novo trabalho de treinamento

  • Criado em cada etapa, na conclusão do trabalho e quando um trabalho é interrompido

Ponto de verificação retomável (estado completo)

  • Armazenado no Intermediate Checkpoint Model Package Group

  • Contém pesos do adaptador LoRa, estados do otimizador e metadados da etapa de treinamento por GPU

  • Use para retomar um trabalho interrompido a partir da etapa exata em que foi interrompido

  • Formato interno — não pode ser usado diretamente para inferência

Ciclo de vida do ponto de verificação

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

Recupere seu modelo treinado

Quando um trabalho é concluído com êxito, o modelo final é salvo como um Model Package no Output Model Package Group. O OutputModelPackageArn campo no registro do trabalho contém o ARN.

Verifique a conclusão do trabalho e recupere o ARN do modelo de saída:

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

Procure OutputModelPackageArn na resposta. Use-o para descrever o Model Package e obter a localização S3 dos pesos:

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

Se um trabalho falhar ou for interrompido antes da conclusão, o último ponto de verificação intermediário será promovido para o Output Model Package Group com base no melhor esforço. Verifique OutputModelPackageArn da mesma forma.

Para monitorar a criação de pontos de verificação durante o treinamento, observe os ModelCheckpoint campos ResumableCheckpoint e na DescribeJob saída.

Retomar um trabalho interrompido

Se um trabalho falhar ou for interrompido no meio do treinamento, você pode começar um novo trabalho que continue exatamente na etapa em que parou. A plataforma restaura todo o estado de treinamento — pesos, impulso do otimizador e contador de passos — a partir do ponto de verificação retomável.

Especifique um ponto de verificação retomável do Intermediate Checkpoint Model Package Group como: 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" }

Eles InputModelPackageArn devem apontar para um ponto de verificação retomável (um com IsCheckpoint=true metadados do Model Package). O treinamento é retomado a partir da etapa após o ponto de verificação — por exemplo, se o ponto de verificação foi salvo na etapa 4, o treinamento continua a partir da etapa 5.

O seguinte deve permanecer o mesmo entre o trabalho original e o trabalho retomado:

  • Modelo de base

  • Configuração LoRa (classificação e alfa)

  • Hiperparâmetros (taxa de aprendizado, tamanho do lote etc.)

  • Conjunto de dados

Continue treinando em um novo emprego (treinamento iterativo)

O treinamento iterativo permite que você desenvolva um modelo previamente treinado com um conjunto de dados diferente, hiperparâmetros diferentes ou uma função de recompensa refinada. Ao contrário da retomada, isso inicia uma nova corrida de treinamento - o otimizador é reiniciado, o contador de passos é redefinido para 0 e apenas os pesos LoRa treinados são transferidos.

Especifique um ponto de verificação do modelo do Output Model Package Group comoInputModelPackageArn:

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

O que você pode mudar entre as iterações:

  • Hiperparâmetros (taxa de aprendizado, tamanho do lote, max_steps, group_size, etc.)

  • Conjunto de dados (solicitações ou distribuição de dados diferentes)

  • Função de recompensa

  • Configuração do agente

O que deve permanecer o mesmo:

  • Modelo básico — o adaptador LoRa está vinculado à arquitetura do modelo básico

Padrões comuns para treinamento iterativo:

  • Aprendizagem curricular — treine primeiro nos problemas mais fáceis e depois continue nos mais difíceis

  • Refinamento da recompensa — comece com uma função de recompensa simples e, em seguida, repita com uma mais sutil

  • Ajuste de hiperparâmetros — aumente o tamanho do lote ou ajuste a taxa de aprendizado após observar a dinâmica inicial do treinamento

Melhores práticas do Checkpoint

  • Monitore a criação de pontos de verificação. Use DescribeJob para monitorar ResumableCheckpoint e treinar em ModelCheckpoint campo durante o treinamento para saber o que está disponível se precisar continuar.

  • Planeje falhas em trabalhos longos. Se um trabalho tiver muitas etapas, projete seu fluxo de trabalho para ser retomado a partir dos pontos de verificação em vez de reiniciar do zero.