Refinamiento por refuerzo (RFT) en Nova 2.0 en SageMaker HyperPod
En esta sección se describe una fórmula de ejemplo, cómo iniciar un trabajo de refinamiento, la guía de hiperparámetros y la monitorización del entrenamiento de RFT en Nova 2.0 Lite en SageMaker HyperPod. Para obtener información sobre el formato de datos, las características compatibles, las restricciones y las prácticas recomendadas para preparar los datos de entrenamiento de RFT, consulte Preparación de datos para el RFT en Amazon Nova 2.
Para determinar si el RFT es adecuado para su caso de uso, consulte Refinamiento por refuerzo (RFT).
Temas
# Note: # This recipe can run on p5.48xlarge, p5e.48xlarge, and p5en.48xlarge instance types. run: name: "my-rft-run" # Unique run name (appears in logs and artifacts). model_type: amazon.nova-2-lite-v1:0:256k model_name_or_path: nova-lite-2/prod data_s3_path: s3://<bucket>/<data-file> # Training dataset in JSONL format. replicas: 4 # Number of total training instances. generation_replicas: 2 # Number of total instances dedicated to response generation. reward_lambda_arn: arn:aws:lambda:<region>:<account-id>:function:<function-name> ## MLFlow configs mlflow_tracking_uri: "" # Required for MLFlow mlflow_experiment_name: "my-rft-experiment" # Optional for MLFlow. Note: leave this field non-empty mlflow_run_name: "my-rft-run" # Optional for MLFlow. Note: leave this field non-empty ## SMHP RFT training configs training_config: max_length: 8192 # Context window (tokens) for inputs and prompt. global_batch_size: 32 # Total samples per optimizer step across all replicas (16/32/64/128/256). reasoning_effort: high # Reasoning mode: high, low, or null for non-reasoning. data: shuffle: true # Shuffle training data each epoch. rollout: # Controls how responses are generated for advantage calculation. rollout_strategy: type: off_policy_async # Asynchronous rollout for higher throughput. age_tolerance: 2 # Maximum policy age before regeneration. advantage_strategy: number_generation: 4 # Samples per prompt to estimate advantages (higher = lower variance but higher cost). generator: max_new_tokens: 6000 # Cap on tokens generated per sample. set_random_seed: true # Seed generation for reproducibility across runs. temperature: 1 # Softmax temperature for sampling. top_k: 1 # Sample only from top-K logits. rewards: preset_reward_function: null # Preset reward functions: exact_match or null for custom. api_endpoint: lambda_arn: arn:aws:lambda:<region>:<account-id>:function:<function-name> lambda_concurrency_limit: 12 # Max concurrent Lambda invocations (throughput vs. throttling). lambda_batch_size: 128 # Number of samples per Lambda invocation. trainer: max_steps: 2 # Steps to train for. One step = global_batch_size samples. save_steps: 5 # Save a checkpoint every N steps. test_steps: 1 # Run validation every N reference model updates. refit_freq: 4 # Frequency of reference model updates. clip_ratio_high: 0.2 # PPO clip ratio for policy updates. loss_scale: 1.0 # Scaling factor for the policy loss. # RL parameters ent_coeff: 0.0 # Entropy bonus added to the policy loss (higher = more exploration). kl_loss_coef: 0.0 # Weight on the KL penalty between the current and reference policy. optim_config: # Optimizer settings. lr: 1e-6 # Learning rate. weight_decay: 0.0 # L2 regularization strength (0.0 to 1.0). adam_beta1: 0.9 adam_beta2: 0.95 peft: # Parameter-efficient fine-tuning (LoRA). peft_scheme: "lora" # Enable LoRA for PEFT. lora_tuning: alpha: 64 # LoRA scaling factor. lora_plus_lr_ratio: 64.0 # LoRA+ learning rate scaling factor (0.0 to 100.0).
Inicio de un trabajo de refinamiento en SageMaker HyperPod
Preparación de los datos de entrada
Para obtener información sobre el formato de datos, las características compatibles, las restricciones y las prácticas recomendadas para preparar los datos de entrenamiento de RFT, consulte Preparación de datos para el RFT en Amazon Nova 2.
Cargar los datos
Cargue su conjunto de datos de entrenamiento en un bucket de S3. Especifique su ubicación en el bloque run de la fórmula:
## Run config run: ... data_s3_path: "s3://<bucket-name>/<training-directory>/<training-file>.jsonl"
nota
Sustituya <bucket-name>, <training-directory> y <training-file> por las rutas de S3 reales.
Definición de su configuración
Defina el modelo base mediante los campos model_type y model_name_or_path del bloque run:
## Run config run: ... model_type: amazon.nova-2-lite-v1:0:256k model_name_or_path: nova-lite-2/prod ...
Guía de hiperparámetros
Utilice los siguientes hiperparámetros recomendados en función de su enfoque de entrenamiento:
General:
-
Épocas: 1
-
Tasa de aprendizaje (lr): 1e-7
-
Número de generaciones: 8
-
Máximo de tokens nuevos: 8192
-
Tamaño del lote: 256
LoRA (adaptación de rango bajo):
-
Rango de LoRA: 32
nota
Ajuste estos valores en función del tamaño del conjunto de datos y del rendimiento de la validación. Supervise las métricas de entrenamiento para evitar el sobreajuste.