Menjalankan pelatihan reinforcement learning di VM TPU menggunakan MaxText

Tutorial ini memberikan panduan langkah demi langkah untuk menjalankan pelatihan reinforcement learning (RL) pada satu instance virtual machine (VM) Tensor Processing Unit (TPU) di Google Cloud dengan menggunakan MaxText, sebuah stack pelatihan berbasis JAX berperforma tinggi untuk model bahasa besar (LLM).v6e-8

Tujuan

  • Siapkan instance VM Cloud TPU.
  • Instal MaxText dan dependensinya.
  • Mengonversi model Hugging Face ke format MaxText.
  • Jalankan workload Pengoptimalan Kebijakan Relatif Grup (GRPO) RL di TPU.
  • Konversi model terlatih kembali ke format Hugging Face untuk inferensi.

Biaya

Dalam dokumen ini, Anda akan menggunakan komponen Google Cloudyang dapat ditagih berikut:

Untuk membuat perkiraan biaya berdasarkan proyeksi penggunaan Anda, gunakan kalkulator harga.

Pengguna Google Cloud baru mungkin memenuhi syarat untuk mendapatkan uji coba gratis.

Setelah menyelesaikan tugas yang dijelaskan dalam dokumen ini, Anda dapat menghindari penagihan berkelanjutan dengan menghapus resource yang Anda buat. Untuk mengetahui informasi selengkapnya, lihat Pembersihan.

Sebelum memulai

  • Anda memerlukan token akses Hugging Face untuk menggunakan tutorial ini. Anda dapat mendaftar untuk mendapatkan akun gratis di Hugging Face. Setelah Anda memiliki akun, buat token akses:

    1. Di halaman Welcome to Hugging Face, klik avatar akun Anda, lalu pilih Access tokens.
    2. Di halaman Access tokens, klik Create new token.
    3. Pilih jenis token Baca dan masukkan nama untuk token Anda.
    4. Token akses Anda akan ditampilkan. Simpan token di tempat yang aman.

  • Di situs Hugging Face, setujui perjanjian lisensi untuk model yang ingin Anda latih. Tutorial ini menggunakan model llama3.1-8b-Instruct.

Untuk mendapatkan izin yang Anda perlukan untuk menyelesaikan tutorial ini, minta administrator Anda untuk memberi Anda peran IAM berikut di project Anda:

Untuk mengetahui informasi selengkapnya tentang pemberian peran, lihat Mengelola akses ke project, folder, dan organisasi.

Anda mungkin juga bisa mendapatkan izin yang diperlukan melalui peran khusus atau peran bawaan lainnya.

Menyiapkan lingkungan

Siapkan variabel lingkungan Anda dengan menjalankan skrip berikut:

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

Ganti kode berikut:

  • YOUR_PROJECT_ID: Project ID Google Cloud Anda
  • ZONE_NAME: zona yang ingin Anda gunakan
  • RESERVATION_NAME: reservasi kapasitas Anda
  • TPU_MACHINE_NAME: nama instance VM Cloud TPU Anda

Lakukan autentikasi dengan Google Cloud menjalankan perintah berikut:

gcloud auth login

Buat VM Cloud TPU

Buat instance VM Cloud TPU dengan 8 chip TPU v6e, yang terikat dengan reservasi kapasitas Anda.

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

Setelah instance VM dibuat, hubungkan ke instance tersebut menggunakan SSH.

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

Selesaikan langkah-langkah berikut dalam instance VM TPU Anda.

Menginstal MaxText

Update paket sistem dalam instance VM TPU.

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

Instal Python 3.12, yang diperlukan MaxText, dan paket lingkungan virtualnya.

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

Gunakan uv untuk mempercepat penginstalan paket Python.

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

Buat lingkungan virtual bernama maxtext_venv dan aktifkan.

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

Instal MaxText dan dependensi yang diperlukan untuk tugas pasca-pelatihan.

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

Instal dependensi wajib yang tersisa dengan menjalankan perintah berikut:

install_tpu_post_train_extra_deps

Mengonversi model ke format MaxText

Untuk melatih model dalam format MaxText, Anda harus mengonversinya dari format Hugging Face ke format MaxText.

Berikan nilai berikut:

  • Token akses Hugging Face Anda
  • Nama model yang ingin Anda gunakan
  • Direktori tempat Anda ingin menyimpan model dalam format MaxText
  • Opsi untuk pemuatan dan penyimpanan
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.

Ganti YOUR_HF_TOKEN dengan token akses Hugging Face yang Anda buat sebelumnya.

Untuk mengonversi model dari format Hugging Face ke format MaxText, jalankan skrip berikut. Konversi ini membutuhkan waktu sekitar lima menit hingga selesai.

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?}

Mulai workload pelatihan

Setelah proses konversi selesai, Anda dapat memulai workload RL.

  1. Konfigurasi parameter pelatihan workload 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. Mulai tugas pelatihan. Proses ini memerlukan waktu sekitar 10 menit di instance 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

Mengonversi kembali model terlatih ke format Hugging Face

Setelah beban kerja pelatihan selesai, konversi model kembali ke format Hugging Face.

  1. Tetapkan jalur untuk ekspor dan parameter terlatih.

    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. Jalankan konversi kembali ke 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

Setelah konversi selesai, model yang disesuaikan dan disimpan di /dev/shm/$MODEL_NAME/hf-trained siap digunakan. Karena Anda kehilangan akses ke isi folder /dev/shm saat VM di-reboot, Anda harus memindahkan model yang disesuaikan ke penyimpanan persisten atau menguploadnya ke Hugging Face Hub.

Pembersihan

Agar tidak menimbulkan biaya tambahan, hapus resource yang dibuat selama tutorial ini.

Menghapus instance VM TPU

Hapus instance VM Cloud TPU Anda.

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

Langkah berikutnya

  • Untuk mengetahui informasi selengkapnya tentang Cloud TPU, lihat Pengantar Cloud TPU.
  • Untuk mengetahui detail arsitektur dan konfigurasi TPU v6e-8, lihat TPU v6e.