View a markdown version of this page

Presentación de trabajos de formación - Amazon SageMaker AI

Las traducciones son generadas a través de traducción automática. En caso de conflicto entre la traducción y la version original de inglés, prevalecerá la version en inglés.

Presentación de trabajos de formación

Lanzamiento de trabajos de formación

Una vez que el agente esté desplegado y el conjunto de datos esté en S3, cree un trabajo de formación mediante uno de los siguientes métodos.

SageMaker AI Studio

  • Navegue hasta Modelos en el panel de navegación y seleccione Modelos JumpStart base.

  • Seleccione un modelo que admita el RL de varias vueltas (consulte la tabla de modelos compatibles) y elija Personalizar el modelo y, a continuación, Personalizar con la interfaz de usuario.

  • Seleccione el aprendizaje por Multi-Turn refuerzo como técnica de personalización.

  • Configure su entorno de agentes: seleccione su entorno de AgentCore ejecución de Bedrock o proporcione el ARN de su reenviador Lambda.

  • Proporcione su conjunto de datos de entrenamiento como un URI de S3 o un conjunto de datos registrado.

  • Ajusta los hiperparámetros según sea necesario.

  • Revise la configuración y seleccione Enviar.

SageMaker SDK de Python para IA

Descubra los modelos compatibles

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 su entorno de agentes

Opción 1: tiempo de ejecución de Bedrock AgentCore

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

Opción 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 su conjunto de datos (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}")

Crear un grupo de paquetes de modelos restringidos para Nova (opcional)

Si elige el modelo Nova (nova-textgeneration-lite-v2), puede crear opcionalmente un grupo de paquetes de modelos restringidos antes de enviar un trabajo de formación (siguiente paso). Si te saltas este paso, el SDK creará uno automáticamente para ti.

Restricted Model Package Group (RMPG) es un grupo de paquetes modelo con ManagedStorageType: Restricted. Es obligatorio para los modelos de código cerrado, como Nova, donde los pesos de los modelos son gestionados por el cliente AWS y el cliente no puede acceder directamente a ellos.

El esquema de tareas de RFT requiere dos MPG restringidos independientes:

  • MPG de salida: almacena el paquete final del modelo ajustado

  • MPG de punto de control intermedio: reservado para los puntos de control de entrenamiento de nivel intermedio (debe diferir del MPG de salida)

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 vez creado el grupo de paquetes modelo, pase los grupos al siguiente paso al enviar un trabajo de formación.

Envíe un trabajo de formación 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}")

Envíe un trabajo de formación con un 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}")

Envíe un trabajo de formación con el grupo de paquetes de modelos restringidos para Nova

Consulte el paso anterior (Crear un grupo de paquetes de modelos restringidos para Nova) sobre cómo crear un grupo de paquetes de modelos restringidos.

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

Cree un trabajo de formación mediante la CreateJob API. La configuración del agente, la ubicación de los datos de entrenamiento, el modelo base y los ajustes de salida se especifican enJobConfigDocument.

Para recuperar el 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"

Cree un trabajo 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

Cree un trabajo con un 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

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

Cree un trabajo con un 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']}")

Capacitación sobre monitoreo

Supervise su trabajo de formación

Usa la DescribeJob API para comprobar el estado actual de tu trabajo en cualquier momento. El estado del trabajo pasa de un lado a otro yInProgress, despuésCompleted, a Failed oStopped.

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

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

Supervise la formación en MLflow

SageMaker La IA se integra automáticamente con el MLflow gestionado para realizar un seguimiento del progreso, las métricas y los artefactos de su trabajo de formación. Para habilitar el seguimiento de MLflow, incluye uno de los siguientes MlflowConfig elementos en tu trabajo: OutputDataConfig

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

Requisitos previos

  • Cree una aplicación MLflow gestionada en su cuenta. Para obtener instrucciones de configuración, consulte Configuración de la aplicación MLflow.

  • Asegúrese de que su función de ejecución de SageMaker IA tenga permisos para escribir en la aplicación MLflow (sagemaker-mlflow:*acciones).

  • Inclúyalo MlflowResourceArn en la configuración de su trabajo.

¿Qué se registra

# Categoría ¿Qué se registra En qué parte de la interfaz de usuario de MLflow
1 Métricas de entrenamiento Per-step contadores, rendimiento, contabilidad de datos y símbolos, duración ininterrumpida de cada fase de un paso, lote de despliegue resumido con recompensa por trayectoria y distribución del recuento de turnos por trayectoria Pestaña de métricas (gráficos de series temporales)
2 Trazos de trayectoria Conversaciones completas en varios turnos con llamadas a herramientas y recompensas Pestaña Traces

Referencia detallada de las métricas de entrenamiento

Las siguientes métricas se registran en cada paso del entrenamiento.

Contadores de pasos y rendimiento () training/

Métrica Description (Descripción)
training/epoch Número de época actual
training/global_step Contador de pasos de entrenamiento global
training/num_groups La trayectoria se agrupa en este paso
training/num_trajectories Trayectorias totales procesadas en este paso
training/total_tokens Los tokens se han sumado en todos los microlotes de este paso
training/num_datums Los datos de entrenamiento se forman a partir de trayectorias
training/datums_per_trajectory Datos medios emitidos por trayectoria
training/action_tokens_mean Número medio de símbolos de acción (respuesta) por trayectoria
training/obs_tokens_mean Indicadores medios de observación (rápidos) por trayectoria
training/trainable_token_positions Total de posiciones objetivo entrenables en este paso
training/nontrainable_token_positions Número total de posiciones objetivo no entrenables en este paso
training/trainable_token_ratio Proporción: posiciones simbólicas trainable / (trainable + nontrainable)

Duraciones de fase () timing_s/

Métrica Description (Descripción)
timing_s/step Tiempo total del paso completo
timing_s/training Tiempo de forward/backward pases y paso de optimización
timing_s/policy_update Los pesos actualizados ahorran tiempo para el muestreador
timing_s/save_checkpoint Ahorra tiempo en un punto de control (solo en los escalones del punto de control)
timing_s/eval Duración de la evaluación (solo en las etapas de evaluación)

Distribución de recompensas (rollout/reward/)

Métrica Description (Descripción)
rollout/reward/mean Recompensa media por trayectoria en todos los grupos
rollout/reward/valid_mean La recompensa media solo se aplica a los grupos válidos (con ventajas distintas de cero); es igual mean cuando no se ha realizado ningún filtrado
rollout/reward/std Desviación estándar de las recompensas por trayectoria
rollout/reward/min Recompensa de trayectoria mínima
rollout/reward/max Recompensa máxima de trayectoria
rollout/reward/zero_frac Fracción de trayectorias con una recompensa total exacta de 0.0

El turno cuenta () rollout/turns/

Métrica Description (Descripción)
rollout/turns/mean Media de giros (transiciones) por trayectoria
rollout/turns/min Giros mínimos entre trayectorias
rollout/turns/max Máximos giros entre trayectorias

Longitudes de los tokens () rollout/tokens/

Métrica Description (Descripción)
rollout/tokens/prompt_mean Recuento medio de fichas rápidas por transición
rollout/tokens/response_mean Recuento medio de tokens de respuesta por transición
rollout/tokens/response_std Desviación estándar de los recuentos de tokens de respuesta
rollout/tokens/response_min Tokens de respuesta mínima
rollout/tokens/response_max Tokens de respuesta máxima (observa la agrupación ensampling_max_tokens)

Log-probability salud () rollout/logprob/

Métrica Description (Descripción)
rollout/logprob/zero_count Total de fichas con cero probabilidades de registro
rollout/logprob/zero_frac Fracción de todos los logprobs que son exactamente 0.0
rollout/logprob/zero_per_group Promedio de cero logprobs por grupo de trayectorias
rollout/logprob/nz_mean Media de los sondeos logarítmicos distintos de cero
rollout/logprob/nz_std Desviación estándar de los probos logarítmicos distintos de cero
rollout/logprob/nz_min Probabilidad logarítmica mínima distinta de cero
rollout/logprob/nz_max Problema de registro máximo distinto de cero

Distribución de ventajas () rollout/advantage/

Métrica Description (Descripción)
rollout/advantage/mean Valor medio de la ventaja en todas las transiciones
rollout/advantage/std Desviación estándar de las ventajas
rollout/advantage/min Ventaja mínima
rollout/advantage/max Máxima ventaja
rollout/advantage/n_positive Transiciones con ventajas positivas
rollout/advantage/n_negative Transiciones con ventajas negativas

Batch-quality clasificación (analysis/)

Métrica Description (Descripción)
analysis/batch_completion_ratio total_completed / batch_size— fracción de los grupos esperados que llegaron
analysis/batch_valid_ratio valid_count / batch_size— grupos con ventajas distintas de cero en relación con el lote completo
analysis/zero_adv_groups Grupos en los que todas las transiciones tienen una ventaja cercana a cero
analysis/zero_adv_nonzero_reward Zero-advantage grupos en los que al menos una transición tiene una recompensa distinta de 0 (todos los casos son correctos para las recompensas binarias)
analysis/zero_adv_zero_reward Zero-advantage grupos en los que todas las recompensas son 0 (mayúsculas y minúsculas)
analysis/reward_variance_across_groups Variación de las recompensas medias por grupo (alta = lote diverso)
analysis/mean_group_reward_spread Diferencia media de recompensas dentro del grupo max - min

Recompensa y pase de evaluación @k () val/reward/

Se emite al inicio (paso 0), en cada val_every intervalo y en el paso final. Incluye las mismas métricas de distribución, rollout/reward además de las métricas de recompensas grupales agregadas por solicitud.

Distribución:

Métrica Description (Descripción)
val/reward/mean Recompensa media sobre el conjunto de evaluación
val/reward/std Recompensa al desarrollador estándar
val/reward/min Recompensa mínima
val/reward/max Recompensa máxima
val/reward/zero_frac Fracción de trayectorias sin recompensa

Group-reward (agregación por mensaje):

Métrica Description (Descripción)
val/reward/min_within_groups Recompensa mínima media por mensaje
val/reward/mean_within_groups Recompensa media media por mensaje
val/reward/max_within_groups Recompensa máxima media por mensaje
val/reward/std_within_groups Estándar de recompensa promedio por mensaje (consistencia)
val/reward/rollouts_per_prompt Número medio de despliegues (n) en todas las solicitudes
val/reward/num_prompts Se evaluaron distintas solicitudes

Pass @k y la contabilidad del éxito:

Métrica Description (Descripción)
val/reward/succeeded_rollouts Número total de lanzamientos con recompensa ≥ success_threshold
val/reward/failed_rollouts Número total de lanzamientos con recompensa < success_threshold
val/reward/success_threshold Umbral utilizado (se repite para mayor claridad)
val/reward/pass_at_{k} Probabilidad de aprobación ≥ 1 de k muestras
val/reward/pass_power_{k} Probabilidad de que pasen todas las k muestras (fiabilidad)

El número de turnos de evaluación cuenta (val/turns/)

Métrica Description (Descripción)
val/turns/mean Número medio de vueltas por trayectoria de evaluación
val/turns/min Número mínimo de vueltas
val/turns/max Número máximo de vueltas

Longitudes de los tokens de evaluación (val/tokens/)

Métrica Description (Descripción)
val/tokens/prompt_mean Promedio de fichas rápidas por transición
val/tokens/response_mean Número medio de indicadores de respuesta por transición
val/tokens/response_std Desviación estándar de los indicadores de respuesta
val/tokens/response_min Tokens de respuesta mínima
val/tokens/response_max Tokens de respuesta máxima

Estado de probabilidad logarítmica de evaluación () val/logprob/

Métrica Description (Descripción)
val/logprob/zero_count Número total de fichas con cero valores logarítmicos
val/logprob/zero_frac Fracción de cero logprobs
val/logprob/zero_per_group Cero sondeos logarítmicos por grupo
val/logprob/nz_mean Media de problemas logarítmicos distintos de cero
val/logprob/nz_std Desviación estándar de los probos logarítmicos distintos de cero
val/logprob/nz_min Probabilidad logarítmica mínima distinta de cero
val/logprob/nz_max Problema de registro máximo distinto de cero

Acceso a la interfaz de usuario de MLflow

Acceda a la interfaz de usuario de MLflow a través de una URL prefirmada:

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 el AuthorizedUrl archivo de salida a su navegador.

Trayectorias y trazas de los agentes

Durante la formación, la SageMaker IA registra todas las interacciones entre el agente y el modelo de política como una trayectoria, es decir, el registro completo de una implementación. Cada trayectoria captura cada solicitud enviada al modelo, cada respuesta generada, cada llamada de herramienta realizada y la recompensa final. Las trayectorias se publican en su experimento de MLflow como trazas estructuradas.

Rastrea el contenido

  • El mensaje de entrada de tu conjunto de datos de entrenamiento

  • Cada turno de inferencia del modelo (datos de aviso, respuesta y nivel de token)

  • Las llamadas de herramientas y sus resultados, si su agente utiliza herramientas

  • La puntuación final de la recompensa

  • Información sobre el tiempo de cada turno

Visualización de las trayectorias en la interfaz de usuario de MLflow

Acceda a la interfaz de usuario de MLflow a través de una URL prefirmada:

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 el AuthorizedUrl archivo de salida a su navegador.

Abra la interfaz de usuario de MLflow con la URL prefirmada de arriba. Navega hasta la ejecución del experimento y selecciona la pestaña Rastros. Cada traza representa una implementación completada y muestra lo siguiente:

  • El mensaje del sistema y el mensaje del usuario

  • La respuesta de cada asistente (con, thinking/reasoning si corresponde)

  • Los intervalos de uso de las herramientas muestran a qué herramientas se recurrió y sus resultados

  • La puntuación de recompensa asignada a la trayectoria

Usa las trayectorias para depurar las puntuaciones de recompensa bajas

Síntoma Qué buscar
La recompensa es baja en la mayoría de los lanzamientos ¿Las respuestas de los modelos son coherentes? ¿El formato del mensaje es correcto?
Tool-related fallas ¿Las llamadas a las herramientas se realizan correctamente? ¿Están bien formadas las entradas y salidas?
Agente en bucle ¿El agente repite las mismas acciones sin avanzar?
Respuestas truncadas ¿Las respuestas están limitadas por el límite de MaxTokens?

Obtenga los resultados de la capacitación

Cuando se completa un trabajo de entrenamiento, las pesas del modelo entrenado se almacenan como un SageMaker AI Model Package. En esta sección se explica cómo encontrar los resultados, comprender los tipos de puntos de control que se generan durante el entrenamiento y utilizarlos para el despliegue o la formación continua.

¿Cómo se almacenan los resultados

SageMaker La IA almacena los resultados del entrenamiento como paquetes de modelos versionados e inmutables dentro de los grupos de paquetes de modelos. Multi-turn RL usa dos grupos separados, que se especifican al crear un trabajo:

Group Finalidad Contenido
Grupo de paquetes de modelos de salida Modelo entrenado final HuggingFace-compatible Pesos del adaptador LoRa (adapter_config.json + adapter_model.safetensors)
Grupo de paquetes modelo de punto de control intermedio Estado de entrenamiento reanudable El adaptador LoRa pesa más los estados del optimizador y los metadatos de los pasos de entrenamiento

Configura ambos grupos en tu: 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 puntos de control

La formación produce dos tipos de puntos de control, que se guardan en cada paso de la formación:

Punto de control modelo (solo pesos)

  • Almacenado en el grupo de paquetes del modelo de salida

  • Contiene los pesos del adaptador HuggingFace-compatible LoRa en formato SafeTensors

  • Úselo para inferencias, despliegues o como punto de partida para un nuevo trabajo de formación

  • Se crea en cada paso, al finalizar el trabajo y cuando se detiene un trabajo

Punto de control reanudable (estado completo)

  • Almacenado en el grupo de paquetes del modelo Intermediate Checkpoint

  • Contiene los pesos de los adaptadores LoRa, los estados del optimizador y los metadatos de los pasos de entrenamiento por GPU

  • Se utiliza para reanudar un trabajo interrumpido desde el mismo paso en el que se detuvo

  • Formato interno: no se puede utilizar directamente para realizar inferencias

Ciclo de vida de los puntos

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 su modelo entrenado

Cuando un trabajo se completa correctamente, el modelo final se guarda como paquete de modelos en el grupo de paquetes de modelos de salida. El OutputModelPackageArn campo del registro de trabajo contiene el ARN.

Compruebe la finalización del trabajo y recupere el ARN del modelo de salida:

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

Busque OutputModelPackageArn en la respuesta. Úselo para describir el Model Package y obtener la ubicación S3 de las pesas:

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

Si un trabajo falla o se detiene antes de completarse, el último punto de control intermedio pasa al Output Model Package Group haciendo todo lo posible. Compruébelo OutputModelPackageArn de la misma manera.

Para supervisar la creación de puntos de control durante el entrenamiento, observa los ModelCheckpoint campos ResumableCheckpoint y del DescribeJob resultado.

Reanude un trabajo interrumpido

Si un trabajo fracasa o se interrumpe a mitad del entrenamiento, puedes empezar un nuevo trabajo retomando el mismo paso en el que lo dejaste. La plataforma restaura el estado completo del entrenamiento (pesas, optimizador de impulso y contador de pasos) desde el punto de control que se puede reanudar.

Especifique un punto de control reanudable del grupo de paquetes del modelo de punto de control intermedio 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" }

InputModelPackageArnDebe apuntar a un punto de control reanudable (uno que figure IsCheckpoint=true en sus metadatos del Model Package). El entrenamiento se reanuda desde el paso posterior al punto de control; por ejemplo, si el punto de control se guardó en el paso 4, el entrenamiento continúa desde el paso 5.

Lo siguiente debe permanecer igual entre el trabajo original y el trabajo reanudado:

  • Modelo básico

  • Configuración LoRa (rango y alfa)

  • Hiperparámetros (tasa de aprendizaje, tamaño del lote, etc.)

  • Conjunto de datos

Continuar formándose para un nuevo trabajo (formación iterativa)

El entrenamiento iterativo te permite basarte en un modelo previamente entrenado con un conjunto de datos diferente, diferentes hiperparámetros o una función de recompensa refinada. A diferencia de la reanudación, esto inicia una nueva sesión de entrenamiento: el optimizador se restablece, el contador de pasos se restablece a 0 y solo se transfieren las pesas LoRa entrenadas.

Especifique un punto de control de modelo del grupo 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" }

Qué puede cambiar entre iteraciones:

  • Hiperparámetros (tasa de aprendizaje, tamaño del lote, max_steps, group_size, etc.)

  • Conjunto de datos (diferentes indicaciones o distribución de datos)

  • Función de recompensa

  • Configuración del agente

Qué debe permanecer igual:

  • Modelo base: el adaptador LoRa está vinculado a la arquitectura del modelo base

Patrones comunes para el entrenamiento iterativo:

  • Aprendizaje curricular: entrénese primero en los problemas más fáciles y luego continúe con los más difíciles

  • Refinamiento de recompensas: comience con una función de recompensa simple y luego repita con una más matizada

  • Ajuste de hiperparámetros: aumente el tamaño del lote o ajuste la velocidad de aprendizaje después de observar la dinámica inicial del entrenamiento

Mejores prácticas de Checkpoint

  • Supervise la creación de puntos de control. DescribeJob Úsalo para realizar un seguimiento ResumableCheckpoint durante el entrenamiento para saber qué hay disponible en caso de que necesites reanudarlo. ModelCheckpoint

  • Planifique los fallos en los trabajos largos. Si un trabajo consta de muchos pasos, diseñe el flujo de trabajo para que se reanude desde los puntos de control en lugar de reiniciarlo desde cero.