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
MlflowResourceArnem 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
ResumableCheckpointe treinar emModelCheckpointcampo 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.