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
MlflowResourceArnen 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
ResumableCheckpointdurante 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.