יצירת פלח TPU עם כמה מארחים

במאמר הזה מוסבר איך ליצור פלח TPU מרובה מארחים באמצעות קבוצת מופעי מכונה מנוהלים (MIG), איך להתחבר לפלח ואיך להריץ חישוב. במדריך הזה נשתמש באפשרות של צריכה על פי דרישה. מריצים את הפקודות במדריך למתחילים הזה בטרמינל המקומי או ב-Cloud Shell.

לפני שמתחילים

  1. נכנסים לחשבון Google Cloud . אם אתם משתמשים חדשים ב- Google Cloud, צרו חשבון כדי שתוכלו להעריך את הביצועים של המוצרים שלנו בתרחישים מהעולם האמיתי. לקוחות חדשים מקבלים בחינם גם קרדיט בשווי 300$ להרצה, לבדיקה ולפריסה של עומסי העבודה.
  2. התקינו את ה-CLI של Google Cloud.

  3. אם אתם משתמשים בספק זהויות חיצוני (IdP), קודם אתם צריכים להיכנס ל-CLI של gcloud באמצעות המאגר המאוחד לניהול זהויות.

  4. כדי לאתחל את ה-CLI של gcloud, הריצו את הפקודה הבאה:

    gcloud init
  5. יוצרים או בוחרים Google Cloud פרויקט.

    תפקידים שנדרשים כדי לבחור או ליצור פרויקט

    • Select a project: כדי לבחור פרויקט לא צריך תפקיד IAM ספציפי – אפשר לבחור כל פרויקט שקיבלתם בו תפקיד.
    • יצירת פרויקט: כדי ליצור פרויקט, צריך את התפקיד Project Creator (יצירת פרויקטים) (roles/resourcemanager.projectCreator), שכולל את ההרשאה resourcemanager.projects.create. איך מקצים תפקידים
    • יוצרים Google Cloud פרויקט:

      gcloud projects create PROJECT_ID

      מחליפים את PROJECT_ID בשם של פרויקט Google Cloud שיוצרים.

    • בוחרים את הפרויקט שיצרתם: Google Cloud

      gcloud config set project PROJECT_ID

      מחליפים את PROJECT_ID בשם הפרויקט ב- Google Cloud .

  6. אם משתמשים בפרויקט קיים, מוודאים שיש את ההרשאות הנדרשות כדי להשלים את ההדרכה. אם משתמשים בפרויקט חדש, לא צריך לוודא כי כבר יש את ההרשאות הנדרשות.

  7. מוודאים שהחיוב מופעל בפרויקט Google Cloud .

  8. מפעילים את Compute Engine API:

    תפקידים שנדרשים להפעלת ממשקי API

    כדי להפעיל ממשקי API, נדרשת ההרשאה serviceusage.services.enable. אם יצרתם את הפרויקט, סביר להניח שכבר יש לכם את ההרשאה הזו דרך התפקיד 'בעלים' (roles/owner). אחרת, תוכלו לקבל את ההרשאה הזו דרך התפקיד 'אדמין של Service Usage' (roles/serviceusage.serviceUsageAdmin). איך מקצים תפקידים

    gcloud services enable compute.googleapis.com
  9. התקינו את ה-CLI של Google Cloud.

  10. אם אתם משתמשים בספק זהויות חיצוני (IdP), קודם אתם צריכים להיכנס ל-CLI של gcloud באמצעות המאגר המאוחד לניהול זהויות.

  11. כדי לאתחל את ה-CLI של gcloud, הריצו את הפקודה הבאה:

    gcloud init
  12. יוצרים או בוחרים Google Cloud פרויקט.

    תפקידים שנדרשים כדי לבחור או ליצור פרויקט

    • Select a project: כדי לבחור פרויקט לא צריך תפקיד IAM ספציפי – אפשר לבחור כל פרויקט שקיבלתם בו תפקיד.
    • יצירת פרויקט: כדי ליצור פרויקט, צריך את התפקיד Project Creator (יצירת פרויקטים) (roles/resourcemanager.projectCreator), שכולל את ההרשאה resourcemanager.projects.create. איך מקצים תפקידים
    • יוצרים Google Cloud פרויקט:

      gcloud projects create PROJECT_ID

      מחליפים את PROJECT_ID בשם של פרויקט Google Cloud שיוצרים.

    • בוחרים את הפרויקט שיצרתם: Google Cloud

      gcloud config set project PROJECT_ID

      מחליפים את PROJECT_ID בשם הפרויקט ב- Google Cloud .

  13. אם משתמשים בפרויקט קיים, מוודאים שיש את ההרשאות הנדרשות כדי להשלים את ההדרכה. אם משתמשים בפרויקט חדש, לא צריך לוודא כי כבר יש את ההרשאות הנדרשות.

  14. מוודאים שהחיוב מופעל בפרויקט Google Cloud .

  15. מפעילים את Compute Engine API:

    תפקידים שנדרשים להפעלת ממשקי API

    כדי להפעיל ממשקי API, נדרשת ההרשאה serviceusage.services.enable. אם יצרתם את הפרויקט, סביר להניח שכבר יש לכם את ההרשאה הזו דרך התפקיד 'בעלים' (roles/owner). אחרת, תוכלו לקבל את ההרשאה הזו דרך התפקיד 'אדמין של Service Usage' (roles/serviceusage.serviceUsageAdmin). איך מקצים תפקידים

    gcloud services enable compute.googleapis.com

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

כדי לקבל את ההרשאות שדרושות ליצירת קבוצת מכונות מנוהלת (MIG) שיוצרת פרוסת TPU מרובת מארחים, צריך להתחבר לכל מכונה וירטואלית בקבוצת המכונות המנוהלת באמצעות SSH ולהריץ פקודות. לשם כך, צריך לבקש מהאדמין להקצות לכם את תפקידי ה-IAM הבאים בפרויקט:

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

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

יצירת תבנית מכונה

כדי ליצור תבנית של הגדרות מכונה ל-TPU v6e, משתמשים בפקודה gcloud compute instance-templates create:

gcloud compute instance-templates create quickstart-tpu-instance-template \
    --machine-type=ct6e-standard-4t \
    --maintenance-policy=TERMINATE \
    --image-family=ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6e \
    --image-project=ubuntu-os-accelerator-images \
    --region=us-east5

יצירת מדיניות של עומס עבודה

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

כדי ליצור מדיניות של עומס עבודה עבור חלוקת TPU מרובת-מארחים, משתמשים בפקודה gcloud compute resource-policies create workload-policy עם הדגל --accelerator-topology. הפקודה הבאה יוצרת מדיניות של עומס עבודה עם טופולוגיה של 2x4:

gcloud compute resource-policies create workload-policy quickstart-tpu-workload-policy \
    --type=high-throughput \
    --accelerator-topology=2x4 \
    --region=us-east5

יצירת קבוצת מופעים מנוהלת (MIG)

מריצים את הפקודות הבאות כדי ליצור קבוצת מופעים מנוהלת (MIG) שיוצרת פלח TPU מרובה-מארחים.

  1. כדי ליצור קבוצת מופעים מנוהלת (MIG) שיוצרת פרוסת TPU מרובת-מארחים, משתמשים בפקודה gcloud compute instance-groups managed create:

    gcloud compute instance-groups managed create quickstart-tpu-mig \
        --size=2 \
        --target-size-policy-mode=bulk \
        --template=quickstart-tpu-instance-template \
        --region=us-east5 \
        --target-distribution-shape=any-single-zone \
        --instance-redistribution-type=none \
        --default-action-on-vm-failure=do-nothing \
        --workload-policy=projects/PROJECT_ID/regions/us-east5/resourcePolicies/quickstart-tpu-workload-policy
    

    מחליפים את PROJECT_ID במזהה הפרויקט ב- Google Cloud .

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

התקנה של JAX

מתקינים יחסי תלות ואת מסגרת JAX בסביבה וירטואלית בכל מופעי TPU VM ב-MIG. אם במכונות הווירטואליות של TPU מותקנת גרסת Python מוקדמת יותר מ-3.11, צריך להתקין את Python 3.11 כדי להריץ את הגרסה העדכנית של JAX.

  1. בודקים איזו גרסה של Python פועלת במכונות הווירטואליות של TPU:

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='python3 --version'
    

    אם הגרסה קודמת ל-Python 3.11, צריך להתקין את Python 3.11:

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='sudo apt update && \
        sudo apt install -y software-properties-common && \
        sudo add-apt-repository -y ppa:deadsnakes/ppa && \
        sudo apt update && \
        sudo apt install -y python3.11 python3.11-dev'
    
  2. יוצרים סביבה וירטואלית:

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='sudo apt install -y python3.11-venv && \
        python3.11 -m venv ~/jax_venv'
    
  3. מתקינים את JAX בסביבה הווירטואלית:

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='source ~/jax_venv/bin/activate && \
        pip install --upgrade pip -q && \
        pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html -q'
    

הרצת קוד JAX בפרוסה

כדי להריץ קוד JAX בפלח TPU, צריך להריץ את הקוד בכל מארח בפלח ה-TPU. הפונקציה jax.device_count() call מפסיקה להגיב עד שהיא נקראת בכל מארח בפרוסה. בדוגמה הבאה אפשר לראות איך מריצים חישוב JAX על TPU slice.

הכנת הקוד

יוצרים קובץ בשם example.py בכל מכונה:

gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
    --region=us-east5 \
    --uri \
| xargs -I {} -P 0 gcloud compute ssh {} \
    --command="cat << 'EOF' > ~/example.py
import jax

# Initialize the slice
jax.distributed.initialize()

# The total number of TPU cores in the slice
device_count = jax.device_count()

# The number of TPU cores attached to this host
local_device_count = jax.local_device_count()

# The psum is performed over all mapped devices across the slice
xs = jax.numpy.ones(jax.local_device_count())
r = jax.pmap(lambda x: jax.lax.psum(x, 'i'), axis_name='i')(xs)

# Print from a single host to avoid duplicated output
if jax.process_index() == 0:
    print('global device count:', jax.device_count())
    print('local device count:', jax.local_device_count())
    print('pmap result:', r)
EOF"

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

מריצים את התוכנית example.py בכל TPU VM בפרוסת ה-TPU:

gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
    --region=us-east5 \
    --uri \
| xargs -I {} -P 0 gcloud compute ssh {} \
    --command='source ~/jax_venv/bin/activate && python3 ~/example.py'

הפלט אמור להיראות כך:

global device count: 8
local device count: 4
pmap result: [8. 8. 8. 8.]

הסרת המשאבים

כדי לא לצבור חיובים בחשבון על המשאבים שבהם השתמשתם בדף הזה, אתם צריכים למחוק את הפרויקט יחד עם המשאבים. Google Cloud Google Cloud

לחלופין, אם אתם רוצים לשמור את הפרויקט, אתם יכולים למחוק רק את ה-MIG ואת כל מכונות ה-VM בקבוצה באמצעות הפקודה gcloud compute instance-groups managed delete:

gcloud compute instance-groups managed delete quickstart-tpu-mig --region=us-east5

המאמרים הבאים