Exécuter un entraînement par renforcement sur une VM TPU à l'aide de MaxText

Ce tutoriel fournit un guide pas à pas pour exécuter l'entraînement par apprentissage par renforcement (RL) sur une seule v6e-8 instance de machine virtuelle (VM) Tensor Processing Unit (TPU) su Google Cloud en utilisant MaxText, une pile d'entraînement haute performance basée sur JAX pour les grands modèles de langage (LLM).

Objectifs

  • Configurer une instance de VM Cloud TPU.
  • Installer MaxText et ses dépendances.
  • Convertir un modèle Hugging Face au format MaxText.
  • Exécuter une charge de travail d'optimisation de la stratégie relative au groupe (GRPO) RL sur le TPU.
  • Reconvertir le modèle entraîné au format Hugging Face pour la diffusion.

Coûts

Dans ce document, vous utilisez les composants facturables suivants de Google Cloud:

Obtenez une estimation des coûts en fonction de votre utilisation prévue, utilisez le simulateur de coût.

Les nouveaux utilisateurs de peuvent bénéficier d'un essai sans frais. Google Cloud

Une fois que vous avez terminé les tâches décrites dans ce document, supprimez les ressources que vous avez créées pour éviter que des frais vous soient facturés. Pour en savoir plus, consultez la section Effectuer un nettoyage.

Avant de commencer

  • Vous avez besoin d'un jeton d'accès Hugging Face pour suivre ce tutoriel. Vous pouvez créer un compte sans frais sur Hugging Face. Une fois que vous avez un compte, générez un jeton d'accès :

    1. Sur la page "Welcome to Hugging Face" (Bienvenue sur Hugging Face), cliquez sur l'avatar de votre compte, puis sélectionnez Access tokens (Jetons d'accès).
    2. Sur la page Access tokens (Jetons d'accès), cliquez sur Create new token (Créer un jeton).
    3. Sélectionnez le type de jeton Read (Lecture), puis saisissez un nom pour votre jeton.
    4. Votre jeton d'accès s'affiche. Enregistrez le jeton dans un endroit sûr.

  • Sur le site Web Hugging Face, acceptez le contrat de licence du modèle que vous prévoyez d'entraîner. Ce tutoriel utilise le modèle llama3.1-8b-Instruct.

Pour obtenir les autorisations nécessaires pour suivre ce tutoriel, demandez à votre administrateur de vous accorder les rôles IAM suivants sur votre projet :

Pour en savoir plus sur l'attribution de rôles, consultez Gérer l'accès aux projets, aux dossiers et aux organisations.

Vous pouvez également obtenir les autorisations requises via des rôles personnalisés ou d'autres rôles prédéfinis.

Configurer l'environnement

Configurez vos variables d'environnement en exécutant le script suivant :

export PROJECT="YOUR_PROJECT_ID"
export ZONE="ZONE_NAME"
export RESERVATION="RESERVATION_NAME"
export TPU_NAME="TPU_MACHINE_NAME"

Remplacez les éléments suivants :

  • YOUR_PROJECT_ID: ID de votre Google Cloud projet
  • ZONE_NAME : zone que vous souhaitez utiliser
  • RESERVATION_NAME : réservation de capacité
  • TPU_MACHINE_NAME : nom de votre instance de VM Cloud TPU

Authentifiez-vous auprès de Google Cloud en exécutant la commande suivante :

gcloud auth login

Créer votre VM Cloud TPU

Créez une instance de VM Cloud TPU avec huit puces TPU v6e, liées à votre réservation de capacité.

gcloud alpha compute tpus tpu-vm create $TPU_NAME \
    --zone=$ZONE \
    --project=$PROJECT \
    --accelerator-type=v6e-8 \
    --version=v2-alpha-tpuv6e \
    --provisioning-model=reservation-bound \
    --reservation=$RESERVATION

Une fois l'instance de VM créée, connectez-vous à celle-ci à l'aide de SSH.

gcloud compute tpus tpu-vm ssh $TPU_NAME --zone $ZONE --project $PROJECT

Suivez les étapes ci-dessous dans votre instance de VM TPU.

Installer MaxText

Mettez à jour les packages système dans l'instance de VM TPU.

sudo apt update && sudo apt upgrade -y --fix-missing

Installez Python 3.12, qui est requis par MaxText, et son package d'environnement virtuel.

sudo apt install -y python3.12 python3.12-venv

Utilisez uv pour accélérer l'installation du package Python.

curl -LsSf https://astral.sh/uv/install.sh | sh
source $HOME/.local/bin/env

Créez un environnement virtuel nommé maxtext_venv, puis activez-le.

uv venv --python 3.12 --seed maxtext_venv
source maxtext_venv/bin/activate

Installez MaxText et les dépendances requises pour les tâches post-entraînement.

uv pip install maxtext[tpu-post-train]==0.2.2 --resolution=lowest

Installez les dépendances requises restantes en exécutant la commande suivante :

install_tpu_post_train_extra_deps

Convertir le modèle au format MaxText

Pour entraîner le modèle au format MaxText, vous devez le convertir du format Hugging Face au format MaxText.

Indiquez les valeurs suivantes :

  • Votre jeton d'accès Hugging Face
  • Nom du modèle que vous souhaitez utiliser
  • Répertoire dans lequel vous souhaitez enregistrer le modèle au format MaxText
  • Options de chargement et de stockage
export HF_TOKEN="YOUR_HF_TOKEN"
export MODEL_NAME='llama3.1-8b-Instruct'
export MODEL_CHECKPOINT_DIRECTORY=/dev/shm/$MODEL_NAME/mt-format/
export USE_PATHWAYS=0 # Set to 1 for Pathways, 0 for McJAX
export LAZY_LOAD_TENSORS=False # True to use lazy load, False to use eager load.

Remplacez YOUR_HF_TOKEN par le jeton d'accès Hugging Face que vous avez créé précédemment.

Pour convertir le modèle du format Hugging Face au format MaxText, exécutez le script suivant. Cette conversion prend environ cinq minutes.

python3 -m maxtext.checkpoint_conversion.to_maxtext \
    model_name=${MODEL_NAME?} \
    hf_access_token=${HF_TOKEN?} \
    base_output_directory=${MODEL_CHECKPOINT_DIRECTORY?} \
    scan_layers=True \
    use_multimodal=False \
    hardware=cpu \
    skip_jax_distributed_system=true \
    checkpoint_storage_use_zarr3=$((1 - USE_PATHWAYS)) \
    checkpoint_storage_use_ocdbt=$((1 - USE_PATHWAYS)) \
    --lazy_load_tensors=${LAZY_LOAD_TENSORS?}

Démarrer la charge de travail d'entraînement

Une fois le processus de conversion terminé, vous pouvez démarrer la charge de travail RL.

  1. Configurez les paramètres d'entraînement de la charge de travail RL.

    # -- MaxText configuration --
    export BASE_OUTPUT_DIRECTORY=/dev/shm/$MODEL_NAME/post-train/
    export RUN_NAME=$(date +%Y-%m-%d-%H-%M-%S)
    export CHIPS_PER_VM=8
    export NUM_BATCHES=50
    export MAXTEXT_CKPT_PATH=$MODEL_CHECKPOINT_DIRECTORY/0/items
  2. Démarrez la tâche d'entraînement. Cette opération prend environ 10 minutes sur une instance de VM v6e-8.

    python3 -m maxtext.trainers.post_train.rl.train_rl \
        model_name=${MODEL_NAME?} \
        load_parameters_path=${MAXTEXT_CKPT_PATH?} \
        run_name=${RUN_NAME?} \
        base_output_directory=${BASE_OUTPUT_DIRECTORY?} \
        chips_per_vm=${CHIPS_PER_VM?} \
        num_batches=${NUM_BATCHES?} \
        num_test_batches=10 \
        rollout_data_parallelism=1 \
        rollout_tensor_parallelism=-1

Reconvertir le modèle entraîné au format Hugging Face

Une fois la charge de travail d'entraînement terminée, reconvertissez le modèle au format Hugging Face.

  1. Définissez les chemins d'exportation et les paramètres entraînés.

    export HF_EXPORT=/dev/shm/$MODEL_NAME/hf-trained/
    export HF_MODEL_NAME=llama3.1-8b
    export POST_TRAIN_PATH=$BASE_OUTPUT_DIRECTORY/$RUN_NAME/checkpoints/actor/$NUM_BATCHES/model_params
  2. Exécutez la reconversion au format Hugging Face.

    python3 -m maxtext.checkpoint_conversion.to_huggingface \
        model_name=${HF_MODEL_NAME?} \
        load_parameters_path=${POST_TRAIN_PATH?} \
        base_output_directory=${HF_EXPORT?} \
        scan_layers=True \
        use_multimodal=False \
        weight_dtype=bfloat16

Une fois la conversion terminée, votre modèle réglé stocké dans /dev/shm/$MODEL_NAME/hf-trained est prêt à être utilisé. Étant donné que vous perdez l'accès au contenu du dossier /dev/shm lorsque la VM redémarre, vous devez déplacer le modèle réglé vers un stockage persistant ou le importer dans le hub Hugging Face.

Effectuer un nettoyage

Pour éviter des frais supplémentaires, supprimez les ressources créées lors de ce tutoriel.

Supprimer votre instance de VM TPU

Supprimez votre instance de VM Cloud TPU.

gcloud alpha compute tpus tpu-vm delete $TPU_NAME --zone=$ZONE --project=$PROJECT --quiet

Étape suivante

  • Pour en savoir plus sur Cloud TPU, consultez la Présentation de Cloud TPU.
  • Pour en savoir plus sur l'architecture et la configuration du TPU v6e-8, consultez TPU v6e.