אימונים גמישים עם Pathways

התכונה 'תוכניות לימודים' מספקת יתרונות של עמידות בדרכים הבאות:

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

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

ודאו שיש לכם:

השהיה והמשך

בדרך כלל, GKE שולח הודעת קדימות ל-pod של מאיץ, לפני שה-pod נדחק. הסבילות להקדמה של Pathways מופעלת כברירת מחדל בכל פריסות הענן, והמשימות של מאיץ Pathways מאזינות להודעות האלה.

כשמתקבלת הודעה על קדימות, המערכת של Pathways קודם כל בודקת אם אפשר לשחזר את עומס העבודה הנוכחי – כלומר, אם אפשר לשמור ולשחזר את עומס העבודה באופן שקוף. אם כן, המערכת מנסה להשהות את עומס העבודה של ה-ML באופן שקוף על ידי כתיבת המצב הנוכחי שלו לאחסון קבוע, כמו Cloud Storage, לפני ש-GKE מפנה את משימות ההאצה. כש-GKE מתזמן מחדש את העבודות שלכם מאוחר יותר, Pathways מפעיל מחדש את עומס העבודה של ה-ML על ידי קריאה חוזרת של המצב שנשמר.

אם אי אפשר לשחזר את עומס העבודה, Pathways משבית את משימת ההאצה ומעביר את הכשל למשימה שלכם אם Elastic training מוגדר. אם לא מוגדרת הדרכה גמישה, ‏ GKE מפעיל מחדש את כל עומס העבודה על סמך מדיניות ההפעלה מחדש של JobSet.

עומסי עבודה אופייניים של למידת מכונה שמוגדרים באמצעות JAX מסתמכים על רכיבי Pathways XLA בלי שמירת מצב, שאפשר לשחזר באמצעות תמונת מצב של זיכרון ברוחב פס גבוה (HBM). עומסי עבודה מסוימים של ML, כמו אלה שמוגדרים באמצעות JAX colocated python API, מסתמכים על רכיבי Pathways עם שמירת מצב. אי אפשר לשחזר אותם.

אימון עם גומיות

אימון גמיש מאפשר לעבודת האימון שלכם להימשך גם כשמתרחשים כשלים בחומרה. השילוב הזה מתבצע באמצעות היכולות של מערכת ה-Pathways והלוגיקה של שחזור המודל שהוגדרה על ידי המשתמש:

  • זיהוי של כשל: כשמתרחש כשל בחומרה (לדוגמה, קורס TPU worker), מערכת Pathways מזהה את זה ומודיעה על כך למשימת האימון של המשתמש באמצעות חריגה בפעם הבאה שמתבצעת גישה לנתונים שנמצאים בחומרה הזו. ההתראה הזו לא גורמת לקריסה של עומס העבודה. היא מאפשרת לקוד לטפל בהתראה ולהגדיר מחדש את המשאבים כדי להמשיך את העיבוד או לצאת בצורה תקינה.
  • מטפל בגמישות שהוגדר על ידי המשתמש: קוד המודל של המשתמש צריך להיות מסוגל לטפל בחריגה הזו. זו הסיבה לכך שהשיטה נקראת 'שחזור ספציפי למודל'.
    • יצירת תמונת מצב: הגישה הנפוצה ביותר היא לשמור מעת לעת תמונות מצב של מצב המודל. אם מתרחש כשל, אפשר לטעון את התמונה העדכנית ביותר כדי להמשיך את האימון.
    • הגדרה מחדש: סביר להניח שתצטרכו להגדיר מחדש את משימת האימון כדי להתאים אותה למספר הפרוסות הזמינות. לדוגמה, אם אחד מהפלחים מפסיק לפעול, יכול להיות שתצטרכו להקטין את מספר הפלחים הפעילים באחד עד שיהיה פלח חלופי. מידע נוסף מופיע במאמר בנושא Elastic Handler.
    • עדכונים של גרף הנתונים או החישוב: הקוד צריך לטפל בכל שינוי במספר המכשירים שזמינים לחישוב, על ידי יצירה מחדש של גרף החישוב לפי הצורך. יכול להיות שתצטרכו לחלק מחדש את הנתונים או לקמפל מחדש את המודל.
  • התפקיד של Pathways בתהליך השחזור: Pathways מספק את הפרימיטיבים לתמיכה בהגדרה מחדש שהוגדרה על ידי המשתמש:
    • החלפת פרוסה: אם פרוסה שנכשלה מוחלפת, אפשר לעדכן את הלקוח כשהפרוסה החדשה זמינה. אחרי כן, הקוד יכול להגדיר מחדש את השימוש בפלח החדש הזה.
    • שחזור שקוף: ‏Pathways מטפל בפרטים ברמה הנמוכה של השחזור, כמו יצירה מחדש של חיבורים לחלקים תקינים של האשכול.
  • ‫Utilities in pathwaysutils: קבוצה של כלי Pathways שמוגדרים ב-pathways-utils.

הטמעה של handler גמיש

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

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

def elastic_handler(elastic_utils, *args, **kwargs):
  mesh = initialize_mesh(**kwargs["mesh_kwargs"])
  initial_state, initial_step, jitted_train_step, other_variables =
      initialize_training_loop(mesh, **kwargs["initialize_training_loop_kwargs"])

  step, snapshot = elastic_utils.get_next_snapshot()
  state = initial_state.replace(**snapshot)

  return state, step, mesh, jitted_train_step, other_variables

עדכון לולאת האימון

צריך לבצע את השינויים הבאים בלולאת ההכשרה:

  1. יצירת חשבון ניהול גמיש
  2. עוטפים את לולאת האימון בבלוקים try-except שמטפלים ב-jax.errors.JaxRuntimeError
  3. בתוך הפונקציה jax.errors.JaxRuntimeError handler, קוראים לפונקציה maybe_reshard_down. אם השגיאה קשורה לאירוע גמיש, המנהל הגמיש יבצע שוב את חלוקת הנתונים, אחרת הוא יציג את השגיאה מחדש.
  4. התקשרות אל maybe_snapshot ו-maybe_reshard_up בסוף לולאת האימון
import pathwaysutils
from pathwaysutils.elastic import manager

pathwaysutils.initialize()

def initialize_mesh(**kwargs):
  ...


def initialize_training_loop(**kwargs):
  ...


def train_loop(
    final_step,
    elastic_manager,
    mesh_kwargs,
    initialize_training_loop_kwargs,
):
  mesh = initialize_mesh(**mesh_kwargs)
  initial_state, initial_step, jitted_train_step, other_variables =
      initialize_training_loop(mesh, **initialize_training_loop_kwargs)

  step = initial_step
  while step < final_step:
    try:
      state = jitted_train_step(state)

      elastic_manager.maybe_snapshot(step=step, snapshot=state)
      handler_returns = elastic_manager.maybe_reshard_up(
          step=step,
          snapshot=state,
          elastic_handler=elastic_handler,
          handler_args=(),
          handler_kwargs=dict(
              mesh_kwargs=mesh_kwargs,
              initialize_training_loop_kwargs=initialize_training_loop_kwargs,
          ),
      )
      if handler_returns:
        state, step, mesh, jitted_train_step, other_variables = handler_returns
      step += 1
    except jax.errors.JaxRuntimeError as error:
      handler_returns = elastic_manager.maybe_reshard_down(
          error=error,
          elastic_handler=elastic_handler,
          handler_args=(),
          handler_kwargs=dict(
              mesh_kwargs=mesh_kwargs,
              initialize_training_loop_kwargs=initialize_training_loop_kwargs,
          ),
      )
      if handler_returns:
        state, step, mesh, jitted_train_step, other_variables = handler_returns

  return state


def main():
  elastic_manager = manager.Manager(
      devices=jax.devices(),
      snapshot_period=10,
      snapshot_buffer_size=1,
      reshard_check_period=5,
      max_elastic_down_event_count=10,
      max_reshard_retry_count=3,
  )

  train_loop(100, elastic_manager, {}, {})

הגדרת מנהל הגמישות

אפשר להגדיר את המנהל הגמיש בכמה דרכים שונות. התדירות של יצירת תמונת מצב נקבעת לפי תקופת תמונת המצב. התקופה של התמונה המייצגת משפיעה על המספר הממוצע של שלבים שאבדו בגלל אירוע אלסטי. התקופה שבה מתבצעת בדיקת הפיצול מחדש קובעת באיזו תדירות לולאת האימון תבדוק אם הפרוסות זמינות. הפרמטר max_elastic_down_event_count מאפשר להגדיר את מספר האירועים הגמישים שהלולאה של האימון תתמוך בהם בגלל אובדן של פרוסת נתונים. התג max_reshard_retry_count מציין את מספר הפעמים שמנהל ה-Elastic צריך לנסות שוב את חלוקת הנתונים מחדש. האובייקט manager הוא אובייקט יחידני וצריך ליצור אותו רק פעם אחת.

תמונות מצב

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

הפחתת הפיצול

אחרי שמתרחשת שגיאה jax.errors.JaxRuntimeError, התכונה 'תוכניות לימודים' בודקת אם השגיאה נובעת מאירוע גמיש בגלל פרוסה שאבדה. אם כן, הפונקציה תקרא ל-elastic handler בלולאה עד שהפעולה תצליח או עד שיגיע למספר המקסימלי של ניסיונות חוזרים. אם השגיאה לא נובעת מאירוע גמיש, השגיאה תופעל שוב. הערכים שמוחזרים מה-handler הגמיש מועברים אל הפונקציה שקוראת לו.

הגדלת מספר הרסיסים

בהתאם להגדרות של מנהל המשאבים הגמיש, אם יש פרוסות לא זמינות, מערכת Pathways תבדוק אם יש פרוסות נוספות שהפכו לזמינות. אם כן, המערכת תשמור מיד תמונת מצב (אם לא צולמה כבר תמונת מצב קיימת לשלב הנוכחי) ותקרא ל-elastic handler בלולאה עד להצלחה או עד להגעה למספר המקסימלי של ניסיונות חוזרים. אם מתבצעת חלוקה מחדש של הנתונים, ערכי ההחזרה של ה-handler הגמיש מועברים אל הפונקציה שקוראת לו. אחרת, מוחזר הערך None.

החלפה חמה

החלפה חמה היא תכונה ב-GKE JobSet API שבה עבודה עם עדיפות גבוהה יותר יכולה להשתלט במהירות על משאבים מעבודה עם עדיפות נמוכה יותר, וכך למזער את זמן ההשבתה ולהבטיח התאוששות מהירה יותר.

כשיוצרים JobSet, ‏ GKE מתזמן את עומס העבודה בכמה פרוסות, בהתאם להגדרות של JobSet. אם מתרחש כשל בחומרה באחד או יותר מהפלחים, ה-Pods המושפעים מסומנים ככשל. כשמזמנים מחדש את Jobset, אם בחרתם להשאיר פרוסת זמן פנויה באשכול GKE שאפשר להשתמש בה עבור עבודה בעדיפות נמוכה יותר, מערכת JobSet תמפה מחדש את עומס העבודה של פרוסת הזמן שנכשלה בעבודה בעדיפות גבוהה יותר לפרוסת הזמן הפנויה שמשמשת לעבודה בעדיפות נמוכה יותר באותו אשכול GKE. המיפוי מחדש בדרך כלל נמשך פחות מדקה.

אחרי הפעלה מחדש של JobSet, החלפה בזמן ריצה יכולה להתרחש במצבים הבאים:

  1. מצב ברירת מחדל: אם יש פרוסות TPU פנויות במצב המתנה באותו אשכול, מתזמן Kubernetes ייתן עדיפות לתזמון של המשימות שהופעלו מחדש בפרוסות האלה, במקום לחכות לתיקון של הפרוסות שנכשלו. כך אפשר לשחזר את החשבון מהר יותר.
  2. עומסי עבודה הטרוגניים: באשכולות שבהם פועלים כמה עומסי עבודה עם PriorityClass מוגדר של Kubernetes, הפעלה מחדש של JobSet יכולה להפעיל החלפה מהירה. אם ההתאמה של האפיניות של המשימה שהופעלה מחדש תואמת למשאבים של משימה עם עדיפות נמוכה יותר, מערכת Kubernetes תבצע קדימה (preemption) של המשימה עם העדיפות הנמוכה יותר, ותאפשר למשימה עם העדיפות הגבוהה יותר להתחיל מיד. לדוגמה, אפשר להגדיר את ה-pods של העובדים ב-Pathways עם עדיפויות שונות באמצעות PriorityClass.

כדי להשתמש בעדיפויות באשכול, מגדירים מחלקת עדיפות, למשל:

kind: PriorityClass
metadata:
  name: high-prior-job
value: 2000
globalDefault: false
description: "This priority class should be used for high priority job."

מחילים את קובץ ה-YAML הזה על אשכול GKE:

kubectl apply -f high-prior-job.yaml

לאחר מכן, מצרפים את מחלקת העדיפות החדשה לעבודת ה-worker של Pathways על ידי הוספת הטקסט הבא ל-podspec של pathways-worker Pod.

priorityClassName: high-prior-job

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