הרצת קוד PyTorch בפרוסות TPU

לפני שמריצים את הפקודות במאמר הזה, צריך לפעול לפי ההוראות במאמר הגדרת חשבון ופרויקט Cloud TPU.

אחרי שקוד PyTorch פועל במכונה וירטואלית אחת של TPU, אפשר להגדיל את קנה המידה של הקוד על ידי הפעלתו בפרוסת TPU. ‫TPU slices הם כמה לוחות TPU שמחוברים זה לזה באמצעות חיבורי רשת ייעודיים במהירות גבוהה. במאמר הזה מוסבר איך להריץ קוד PyTorch בפרוסות TPU.

התפקידים הנדרשים

כדי לקבל את ההרשאות שדרושות ליצירת TPU ולהתחבר אליו באמצעות SSH, צריך לבקש מהאדמין להקצות לכם בפרויקט את תפקידי ה-IAM הבאים:

להסבר על מתן תפקידים, ראו איך מנהלים את הגישה ברמת הפרויקט, התיקייה והארגון.

יכול להיות שאפשר לקבל את ההרשאות הנדרשות גם באמצעות תפקידים בהתאמה אישית או תפקידים מוגדרים מראש.

יצירת פרוסת Cloud TPU

  1. הגדרת משתני סביבה לפרמטרים של פקודות Google Cloud CLI.

    export PROJECT_ID=your-project-id
    export TPU_NAME=your-tpu-name
    export ZONE=us-central1-b
    export ACCELERATOR_TYPE=v6e-8
    export RUNTIME_VERSION=v2-alpha-tpuv6e

    תיאורים של משתני סביבה

    • PROJECT_ID: מזהה הפרויקט ב- Google Cloud . משתמשים בפרויקט קיים או יוצרים פרויקט חדש.
    • TPU_NAME: השם של ה-TPU.
    • ZONE: האזור שבו יוצרים את ה-TPU VM. מידע נוסף על אזורים נתמכים זמין במאמר אזורים ותחומים של TPU.
    • ACCELERATOR_TYPE: סוג המאיץ מציין את הגרסה והגודל של Cloud TPU שרוצים ליצור. מידע נוסף על סוגי המאיצים הנתמכים בכל גרסת TPU זמין במאמר בנושא גרסאות TPU.
    • RUNTIME_VERSION: גרסת התוכנה של Cloud TPU.

  2. מריצים את הפקודה הבאה כדי ליצור TPU VM:

    gcloud compute tpus tpu-vm create $TPU_NAME \
         --zone=$ZONE \
         --project=$PROJECT_ID \
         --accelerator-type=$ACCELERATOR_TYPE \
         --version=$RUNTIME_VERSION
    

התקנת PyTorch/XLA בפרוסת ה-TPU

אחרי שיוצרים את פרוסת ה-TPU, צריך להתקין את PyTorch בכל המארחים בפרוסת ה-TPU. אפשר לעשות את זה באמצעות הפקודה gcloud compute tpus tpu-vm ssh עם הפרמטרים --worker=all ו---command.

  1. יוצרים קובץ בשם requirements.txt עם התוכן הבא:

    --find-links https://storage.googleapis.com/libtpu-releases/index.html
    --find-links https://storage.googleapis.com/libtpu-wheels/index.html
    torch~=2.6.0
    torch_xla[tpu]~=2.6.0
    torchvision
    ray[default]==2.40.0
    
  2. מעתיקים את requirements.txt לכל מכונה וירטואלית בפרוסה:

    gcloud compute tpus tpu-vm tpu-vm scp ./requirements.txt \
         $TPU_NAME:~/ \
         --zone=$ZONE \
         --project=$PROJECT_ID \
         --worker=all
    
  3. מתקינים את יחסי התלות בכל מכונת VM בפרוסה:

    gcloud compute tpus tpu-vm ssh $TPU_NAME \
         --zone=$ZONE \
         --project=$PROJECT_ID \
         --worker=all \
         --command="pip3 install -r requirements.txt"
    
  4. משכפלים את XLA בכל העובדים של TPU VM:

    gcloud compute tpus tpu-vm ssh $TPU_NAME \
         --zone=$ZONE \
         --project=$PROJECT_ID \
         --worker=all \
         --command="git clone https://github.com/pytorch/xla.git"
    

הרצת סקריפט אימון ב-TPU slice

מריצים את סקריפט האימון על כל העובדים. סקריפט האימון משתמש בשיטת שרדינג של Single Program Multiple Data (SPMD). מידע נוסף על SPMD זמין במדריך למשתמש של PyTorch/XLA SPMD.

gcloud compute tpus tpu-vm ssh $TPU_NAME \
   --zone=$ZONE \
   --project=$PROJECT_ID \
   --worker=all \
   --command="PJRT_DEVICE=TPU python3 ~/xla/test/spmd/test_train_spmd_imagenet.py  \
   --fake_data \
   --model=resnet50  \
   --num_epochs=1 2>&1 | tee ~/logs.txt"

ההדרכה נמשכת כ-15 דקות. בסיום התהליך אמורה להופיע הודעה דומה לזו שכאן למטה:

   Epoch 1 test end 23:49:15, Accuracy=100.00
   10.164.0.11 [0] Max Accuracy: 100.00%

הסרת המשאבים

כשמסיימים להשתמש ב-TPU VM, פועלים לפי השלבים הבאים כדי לנקות את המשאבים.

  1. אם עדיין לא עשיתם זאת, צריך להתנתק ממכונת Cloud TPU:

    exit
    

    ההנחיה אמורה להיות עכשיו username@projectname, כדי שתדעו שאתם ב-Cloud Shell.

  2. מוחקים את משאבי Cloud TPU.

    gcloud compute tpus tpu-vm delete $TPU_NAME --zone=$ZONE
    
  3. כדי לוודא שהמשאבים נמחקו, מריצים את הפקודה gcloud compute tpus tpu-vm list. יכול להיות שיחלפו כמה דקות עד שהפריט יימחק. הפלט מהפקודה הבאה לא אמור לכלול אף אחד מהמשאבים שנוצרו במדריך הזה:

    gcloud compute tpus tpu-vm list --zone=$ZONE