איך מבצעים כוונון עדין של מודל LLM באמצעות TPUs ב-GKE עם JAX

במדריך הזה נסביר איך לכוונן מודל גדול של שפה (LLM) באמצעות יחידות לעיבוד טנסורים (TPU) ב-Google Kubernetes Engine ‏ (GKE) עם JAX. באמצעות כוונון עדין אפשר להתאים מודל בסיס כמו Gemma 3 לדומיין או למשימה ספציפיים. התהליך הזה משפר את הדיוק של המודל על ידי עדכון הפרמטרים שלו באמצעות מערך נתונים מיוחד משלכם.

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

רקע

באמצעות שימוש ב-TPU ב-GKE עם Jax כדי לבצע כוונון עדין של LLM, אתם יכולים ליצור פתרון כוונון עדין חזק ומוכן לייצור עם כל היתרונות של Kubernetes מנוהל.

‏Gemma

Gemma היא סדרה של מודלים קלים ורב-מודאליים של AI גנרטיבי/ML שזמינים לשימוש חופשי ומופצים ברישיון קוד פתוח. מודלים ה-AI האלה זמינים להפעלה באפליקציות, בחומרה, במכשירים ניידים או בשירותים מתארחים. ‫Gemma 3 מציג יכולות מולטימודאליות, והוא תומך בקלט של שפה ויזואלית ובתפוקות של טקסט. הוא יכול לטפל בחלונות הקשר של עד 128,000 טוקנים ותומך ביותר מ-140 שפות. בנוסף, מודל Gemma 3 מציע יכולות משופרות במתמטיקה, בחשיבה רציונלית ובצ'אט, כולל פלט מובנה וקריאה להפעלת פונקציות.

אתם יכולים להשתמש במודלים של Gemma ליצירת טקסט, או לכוונן אותם למשימות מיוחדות.

מידע נוסף זמין במסמכי התיעוד של Gemma.

TPUs

יחידות TPU הן מעגלים משולבים לאפליקציות ספציפיות (ASIC) שפותחו על ידי Google כדי להאיץ מודלים של למידת מכונה ו-AI שנבנים באמצעות frameworks כמו TensorFlow,‏ PyTorch ו-JAX.

לפני שמשתמשים ב-TPU ב-GKE, מומלץ להשלים את תוכנית הלימודים הבאה:

  1. מידע על הזמינות של גרסאות TPU עדכניות מופיע בארכיטקטורת המערכת של Cloud TPU.
  2. מידע על TPU ב-GKE

JAX

JAX היא מסגרת ללמידת מכונה עם ביצועים גבוהים, שנועדה לשימוש עם TPU ו-GPU. ‫JAX מספקת API ליצירה ולאימון של מודלים של למידת מכונה.

מידע נוסף זמין במאגר JAX.

מטרות

במדריך הזה מוסבר איך:

  1. יוצרים אשכול GKE במצב Autopilot או אשכול רגיל עם טופולוגיית TPU מומלצת, על סמך מאפייני המודל. במהלך המדריך הזה, תבצעו כוונון עדין במאגרי צמתים של מארח יחיד.
  2. מוסיפים נתונים לקטגוריה של Cloud Storage וטוענים אותה למאגר באמצעות Cloud Storage FUSE.
  3. פורסים את משימת הכוונון העדין של ה-LLM ב-GKE.
  4. עוקבים אחרי משימת הכוונון העדין ומציגים את היומנים.

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

  • נכנסים לחשבון Google Cloud . אם אתם משתמשים חדשים ב- Google Cloud, צרו חשבון כדי שתוכלו להעריך את הביצועים של המוצרים שלנו בתרחישים מהעולם האמיתי. לקוחות חדשים מקבלים בחינם גם קרדיט בשווי 300$ להרצה, לבדיקה ולפריסה של עומסי העבודה.
  • In the Google Cloud console, on the project selector page, select or create a Google Cloud project.

    Roles required to select or create a project

    • Select a project: Selecting a project doesn't require a specific IAM role—you can select any project that you've been granted a role on.
    • Create a project: To create a project, you need the Project Creator role (roles/resourcemanager.projectCreator), which contains the resourcemanager.projects.create permission. Learn how to grant roles.

    Go to project selector

  • Verify that billing is enabled for your Google Cloud project.

  • Enable the required API.

    Roles required to enable APIs

    To enable APIs, you need the Service Usage Admin IAM role (roles/serviceusage.serviceUsageAdmin), which contains the serviceusage.services.enable permission. Learn how to grant roles.

    Enable the API

  • In the Google Cloud console, on the project selector page, select or create a Google Cloud project.

    Roles required to select or create a project

    • Select a project: Selecting a project doesn't require a specific IAM role—you can select any project that you've been granted a role on.
    • Create a project: To create a project, you need the Project Creator role (roles/resourcemanager.projectCreator), which contains the resourcemanager.projects.create permission. Learn how to grant roles.

    Go to project selector

  • Verify that billing is enabled for your Google Cloud project.

  • Enable the required API.

    Roles required to enable APIs

    To enable APIs, you need the Service Usage Admin IAM role (roles/serviceusage.serviceUsageAdmin), which contains the serviceusage.services.enable permission. Learn how to grant roles.

    Enable the API

  • צריך לוודא שיש לכם בפרויקט את התפקיד או התפקידים הבאים: roles/container.admin,roles/iam.serviceAccountAdmin,roles/storage.admin

    בדיקת התפקידים

    1. נכנסים לדף IAM במסוף Google Cloud .

      כניסה לדף IAM
    2. בוחרים את הפרויקט.
    3. בעמודה Principal (חשבון המשתמש), מוצאים את כל השורות שבהן מופיע השם שלכם או של קבוצה שאתם נכללים בה. כדי לברר באילו קבוצות אתם נכללים, פנו לאדמין.

    4. בודקים את העמודה Role בכל השורות שבהן מצוין או מופיע השם שלכם, כדי לראות אם רשימת התפקידים כוללת את התפקידים הנדרשים.

    מתן התפקידים

    1. נכנסים לדף IAM במסוף Google Cloud .

      כניסה לדף IAM
    2. בוחרים את הפרויקט.
    3. לוחצים על Grant access.
    4. בשדה New principals, מזינים את מזהה המשתמש. ‫ בדרך כלל מזהה המשתמש הוא כתובת האימייל של חשבון Google.

    5. לוחצים על Select a role ומחפשים את התפקיד.
    6. כדי לתת עוד תפקידים, לוחצים על Add another role ומוסיפים אותם.
    7. לוחצים על Save.
  • מוודאים שיש לכם מספיק מכסת שימוש בשבבי TPU Trillium ‏ (v6e)‎. במדריך הזה משתמשים בהגדרת מאגר צמתים שדורשת 16 שבבים ומופעים על פי דרישה.
  • ודאו שיש לכם מאגר Docker. אם אין לכם מאגר, אתם צריכים ליצור מאגר רגיל ב-Artifact Registry.

הכנת הסביבה

במדריך הזה תשתמשו ב-Cloud Shell כדי לנהל משאבים שמתארחים ב- Google Cloud. ב-Cloud Shell מותקנת מראש התוכנה שדרושה למדריך הזה, כולל kubectl ו-Google Cloud CLI.

כדי להגדיר את הסביבה באמצעות Cloud Shell:

  1. במסוף Google Cloud , מפעילים סשן של Cloud Shell ולוחצים על Activate Cloud Shell.סמל ההפעלה של Cloud Shell הפעולה הזו מפעילה סשן בחלונית התחתונה של מסוף Google Cloud .

  2. מגדירים את משתני הסביבה שמוגדרים כברירת מחדל:

    gcloud config set project PROJECT_ID
    gcloud config set billing/quota_project PROJECT_ID
    export PROJECT_ID=$(gcloud config get project)
    export CLUSTER_NAME=CLUSTER_NAME
    export REGION=CONTROL_PLANE_LOCATION
    export ZONE=ZONE
    export GCS_BUCKET_NAME=BUCKET_NAME
    

    מחליפים את הערכים הבאים:

    • PROJECT_ID: מזהה הפרויקט ב- Google Cloud .
    • CLUSTER_NAME: השם של אשכול GKE.
    • CONTROL_PLANE_LOCATION: האזור ב-Compute Engine שבו נמצאים אשכול GKE וצומתי TPU. האזור צריך להכיל אזורים שבהם זמינים סוגי מכונות TPU Trillium‏ (v6e).
    • ZONE: אזור בתוך אזור CONTROL_PLANE_LOCATION שבחרתם, שבו זמינים סוגי מכונות TPU Trillium ‏ (v6e). כדי להציג את רשימת האזורים שבהם זמינים מכשירי TPU Trillium ‏ (v6e), מריצים את הפקודה הבאה:

        gcloud compute accelerator-types list --filter="name~ct6e" --format="value(zone)"
      
    • BUCKET_NAME: השם של הקטגוריה של Cloud Storage שמכילה את נתוני האימון.

  3. משכפלים את המאגר לדוגמה:

    git clone https://github.com/GoogleCloudPlatform/kubernetes-engine-samples.git
    cd kubernetes-engine-samples
    
  4. עוברים לספריית העבודה:

    cd ai-ml/llm-training-jax-tpu-gemma3
    

יצירה והגדרה של Google Cloud משאבים

בקטע הזה יוצרים ומגדירים Google Cloud משאבים.

יצירת אשכול GKE

אפשר לכוונן מודל שפה גדול (LLM) ב-TPU באשכול GKE במצב Autopilot או באשכול רגיל. מומלץ להשתמש באשכול Autopilot כדי ליהנות מחוויית Kubernetes מנוהלת באופן מלא. כדי לבחור את מצב הפעולה של GKE שהכי מתאים לעומסי העבודה שלכם, אפשר לעיין במאמר בחירת מצב פעולה של GKE.

טייס אוטומטי

יוצרים אשכול GKE Autopilot שמשתמש ב איחוד זהויות של עומסי עבודה ל-GKE ומופעל בו Cloud Storage FUSE.

gcloud container clusters create-auto ${CLUSTER_NAME} \
    --location=${REGION}

יצירת האשכול עשויה להימשך כמה דקות.

רגילה

  1. יוצרים אשכול GKE Standard אזורי שמשתמש באיחוד זהויות של עומסי עבודה ל-GKE ומופעל בו Cloud Storage FUSE.

    gcloud container clusters create ${CLUSTER_NAME} \
        --enable-ip-alias \
        --addons GcsFuseCsiDriver \
        --machine-type=n2-standard-4 \
        --num-nodes=2 \
        --workload-pool=${PROJECT_ID}.svc.id.goog \
        --location=${REGION}
    

    יצירת האשכול עשויה להימשך כמה דקות.

  2. יוצרים מאגר צמתים עם מארח יחיד:

    gcloud container node-pools create jax-tpu-nodepool \
        --cluster=${CLUSTER_NAME} \
        --machine-type=ct6e-standard-1t \
        --num-nodes=1 \
        --location=${REGION} \
        --node-locations=${ZONE} \
        --workload-metadata=GKE_METADATA
    

‫GKE יוצר מאגר צמתים של TPU Trillium עם טופולוגיה 1x1 וצומת אחד. הדגל --workload-metadata=GKE_METADATA מגדיר את מאגר הצמתים לשימוש בשרת המטא-נתונים של GKE.

התקנת JobSet

  1. מגדירים את kubectl לתקשורת עם האשכול:

    gcloud container clusters get-credentials ${CLUSTER_NAME} --location=${REGION}
    
  2. מתקינים את הגרסה האחרונה שפורסמה של JobSet:

    kubectl apply --server-side -f https://github.com/kubernetes-sigs/jobset/releases/download/JOBSET_VERSION/manifests.yaml
    

    מחליפים את JOBSET_VERSION בגרסה האחרונה של JobSet. לדוגמה, v0.11.0.

  3. מאמתים את ההתקנה של JobSet:

    kubectl get pods -n jobset-system
    

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

    NAME                                         READY   STATUS    RESTARTS   AGE
    jobset-controller-manager-6c56668494-l4dhc   1/1     Running   0          4m45s
    

    יכול להיות שתצטרכו להוסיף עוד צמתים אם JobSet ממתין למשאבים.

הגדרת Cloud Storage FUSE

כדי לבצע התאמה עדינה של מודל שפה גדול, צריך לספק נתונים לאימון. במדריך הזה השתמשנו במערך הנתונים TinyStories מ-Hugging Face. קבוצת הנתונים הזו מכילה סיפורים קצרים שנוצרו באופן סינתטי על ידי GPT-3.5 ו-GPT-4, עם אוצר מילים מוגבל.

בקטע הזה מוסבר איך מגדירים את Cloud Storage FUSE לקריאת נתונים מקטגוריה של Cloud Storage.

  1. מורידים את מערך הנתונים:

    wget https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStories-train.txt?download=true -O TinyStories-train.txt
    
  2. מעלים את הנתונים לקטגוריה חדשה ב-Cloud Storage:

    gcloud storage buckets create gs://${GCS_BUCKET_NAME} \
        --location=${REGION} \
        --enable-hierarchical-namespace \
        --uniform-bucket-level-access
    gcloud storage cp TinyStories-train.txt gs://${GCS_BUCKET_NAME}
    
  3. כדי לאפשר לעומס העבודה לקרוא נתונים דרך Cloud Storage FUSE, צריך ליצור חשבון שירות של Kubernetes (KSA) ולהוסיף את ההרשאות הנדרשות. מריצים את הסקריפט permissionsetup.sh:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    #!/bin/bash
    
    # --- Configuration Variables ---
    # Kubernetes Service Account details
    export KSA_NAME="jaxserviceaccout"
    export NAMESPACE="default"
    
    # Google Cloud IAM Service Account details
    export GSA_NAME="<GSA_NAME>"
    # Automatically get the current project ID
    export PROJECT_ID=$(gcloud config get-value project)
    export  GSA_DESCRIPTION="GKE Service Account to read GCS bucket for ${KSA_NAME}"
    
    # GCS Bucket details
    export GCS_BUCKET_NAME="<GCS_BUCKET_NAME>" # <--- IMPORTANT: Update this to your bucket name
    
    # Derived Variables
    export GSA_EMAIL="${GSA_NAME}@${PROJECT_ID}.iam.gserviceaccount.com"
    export WI_MEMBER="serviceAccount:${PROJECT_ID}.svc.id.goog[${NAMESPACE}/${KSA_NAME}]"
    
    # --- Check if PROJECT_ID is set ---
    if [ -z "${PROJECT_ID}" ]; then
      echo "Error: PROJECT_ID is not set. Please set it using 'gcloud config set project YOUR_PROJECT_ID'"
      exit 1
    fi
    
    echo "--- Configuration ---"
    echo "KSA_NAME:      ${KSA_NAME}"
    echo "NAMESPACE:     ${NAMESPACE}"
    echo "GSA_NAME:      ${GSA_NAME}"
    echo "PROJECT_ID:    ${PROJECT_ID}"
    echo "GSA_EMAIL:     ${GSA_EMAIL}"
    echo "GCS_BUCKET_NAME:   ${GCS_BUCKET_NAME}"
    echo "WI_MEMBER:     ${WI_MEMBER}"
    echo "--------------------"
    read -p "Press enter to continue..."
    
    # --- Command Execution ---
    
    echo "[1/5] Creating Google Cloud IAM Service Account (GSA): ${GSA_NAME}"
    gcloud iam service-accounts create "${GSA_NAME}" \
        --project="${PROJECT_ID}" \
        --description="${GSA_DESCRIPTION}" \
        --display-name="${GSA_NAME}"
    
    echo "[2/5] Granting GSA '${GSA_EMAIL}' read access (roles/storage.objectViewer) to bucket 'gs://${GCS_BUCKET_NAME}'"
    gcloud storage buckets add-iam-policy-binding "gs://${GCS_BUCKET_NAME}" \
        --member="serviceAccount:${GSA_EMAIL}" \
        --role="roles/storage.objectViewer" \
        --project="${PROJECT_ID}"
    
    echo "[3/5] Creating Kubernetes Service Account (KSA): ${KSA_NAME} in namespace ${NAMESPACE}"
    kubectl create serviceaccount "${KSA_NAME}" --namespace "${NAMESPACE}"
    
    echo "[4/5] Allowing KSA to impersonate GSA (Workload Identity Binding): ${GSA_EMAIL}"
    gcloud iam service-accounts add-iam-policy-binding "${GSA_EMAIL}" \
        --role roles/iam.workloadIdentityUser \
        --member "${WI_MEMBER}" \
        --project="${PROJECT_ID}"
    
    echo "[5/5] Annotating KSA '${KSA_NAME}' to link with GSA '${GSA_EMAIL}'"
    kubectl annotate serviceaccount "${KSA_NAME}" \
        --namespace "${NAMESPACE}" \
        iam.gke.io/gcp-service-account="${GSA_EMAIL}"
    
    echo "--- Setup Complete ---"
    echo "Pods in namespace '${NAMESPACE}' using serviceAccount '${KSA_NAME}' can now authenticate as '${GSA_EMAIL}' and have read access to 'gs://${GCS_BUCKET_NAME}'."
    

    אחרי שמריצים את הסקריפט הזה, המשאבים הבאים מוגדרים בGoogle Cloud פרויקט ובאשכול GKE:

    • נוצר חשבון שירות חדש ב-IAM בשם gcs-fuse-sa בפרויקט.
    • לחשבון השירות (GSA) שנוצר Google Cloud (gcs-fuse-sa) מוקצה התפקיד roles/storage.objectViewer בקטגוריה של Cloud Storage שצוינה על ידי ${GCS_BUCKET_NAME}. ההרשאה הזו מאפשרת ל-GSA לקרוא אובייקטים מהקטגוריה.
    • במרחב השמות default באשכול GKE נוצר KSA חדש בשם jaxserviceaccount.
    • מדיניות ה-IAM של GSA מתעדכנת כדי להעניק את התפקיד roles/iam.workloadIdentityUser ל-KSA. ההרשאה הזו מאפשרת ל-KSA להתחזות ל-GSA.
    • ה-KSA מתויג כדי לקשר אותו ל-GSA. ההערה הזו מציינת ל-GKE לאיזה חשבון שירות של Google (GSA) חשבון השירות של Kubernetes (KSA) צריך להתחזות באמצעות Workload Identity.

      כל פוד שפועל במרחב השמות default של אשכול GKE ומשתמש בחשבון השירות jaxserviceaccount יוכל עכשיו לבצע אימות בתור חשבון שירות של Google‏ (GSA) מספר gcs-fuse-sa. ל-Pods האלה תהיה גישת קריאה לאובייקטים שמאוחסנים בקטגוריה gs://${GCS_BUCKET_NAME}, וזה חיוני כדי שה-Job של הכוונון העדין יוכל לגשת למערך הנתונים באמצעות Cloud Storage FUSE.

יצירת סקריפט לכוונון עדין

בקטע הזה נסביר על סקריפט ההדרכה שמבצע פעולת כוונון עדין במודל Gemma 3. הסקריפט הזה משתמש ב-Gemma3Tokenizer.

כדאי לעיין בסקריפט הבא של כוונון עדין:Gemma3LLMTrain.py

# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import grain.python as pygrain
import jax
import jax.numpy as jnp
import optax
import pandas as pd
import time
import argparse

from dataclasses import dataclass
from functools import partial
from gemma import gm
from flax.training import train_state
from jax.sharding import Mesh, PartitionSpec, NamedSharding

jax.distributed.initialize()
print("Global device count:", jax.device_count())
print("jax version:", jax.__version__)

tokenizer = gm.text.Gemma3Tokenizer()
num_epochs = 1
learning_rate = 2e-5

@dataclass
class TextDataset:
    data: list
    maxlen: int
    tokenizer: gm.text.Gemma3Tokenizer

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx: int):
        encoding = self.tokenizer.encode(self.data[idx])[:self.maxlen]  # Tokenize and truncate
        return encoding + [0] * (self.maxlen - len(encoding))  # Pad to maxlen

def load_and_preprocess_data(file_path, batch_size, maxlen, datacount, tokenizer):

    with open(file_path, 'r') as f:
      text = f.read()

    stories = text.split('<|endoftext|>')
    stories = [story for story in stories if story.strip()][:datacount]
    df = pd.DataFrame({'text': stories})
    data = df['text'].dropna().tolist()
    dataset = TextDataset(data, maxlen, tokenizer)

    sampler = pygrain.IndexSampler(
        len(dataset),
        shuffle=False,
        seed=42,
        shard_options=pygrain.NoSharding(),
        num_epochs=num_epochs,
    )

    dataloader = pygrain.DataLoader(
        data_source=dataset,
        sampler=sampler,
        operations=[pygrain.Batch(batch_size=batch_size, drop_remainder=True)],
    )

    return dataloader

def generate_text(model, params, tokenizer, prompt):
    sampler = gm.text.Sampler(
        model=model,
        params=params,
        tokenizer=tokenizer,
    )
    print("Generating response for: " + prompt)
    out = sampler.sample(prompt, max_new_tokens=32)
    print("Reponse: \n" + out + "\n")
    return out

prep_target_batch = jax.vmap(lambda tokens: jnp.concatenate((tokens[1:], jnp.array([0]))))

@partial(jax.jit, donate_argnums=(0,))
def train_step(state, batch):
    """Performs one supervised fine-tuning step."""

    def loss_fn(params):
        # Run the forward pass. The model returns logits.
        logits = state.apply_fn({'params': params}, batch[0]).logits

        # Calculate the cross-entropy loss.
        loss = optax.softmax_cross_entropy_with_integer_labels(
            logits=logits, labels=batch[1]
        ).mean()

        return loss

    # Compute gradients
    grad_fn = jax.value_and_grad(loss_fn)
    loss, grads = grad_fn(state.params)

    # Update the model state
    state = state.apply_gradients(grads=grads)

    metrics = {'loss': loss}
    return state, metrics

def train_model(state, text_dl, num_epochs, sharding):
    batchCount = 0
    start_time = time.time()
    for epoch in range(num_epochs):
        start_time = time.time()
        for batch in text_dl:
            if len(batch) % len(jax.devices()) != 0:
              continue  # skip the remaining elements
            input_batch = jnp.array(jnp.array(batch).T)
            target_batch = prep_target_batch(input_batch)
            state, metrics = train_step(state, jax.device_put((input_batch, target_batch), sharding))

            if batchCount % 10 == 0:
                print(f"Loss after batch {batchCount}: {metrics['loss']}")
            batchCount += 1

    end_time = time.time()
    print(f"Completed training model. Total time for training {end_time - start_time} seconds \n")
    return state

def run_training(maxlen, batch_size, datacount):
    print(f"Batch size: {batch_size}, Max length: {maxlen}, Data count: {datacount}")
    #Load the training data
    tiny_stories_dl = load_and_preprocess_data('/data/TinyStories-train.txt', batch_size, maxlen, datacount, tokenizer)
    # Get the Gemma3 model
    model = gm.nn.Gemma3_270M()
    # Load the pretrained parameters
    params = gm.ckpts.load_params(gm.ckpts.CheckpointPath.GEMMA3_270M_PT)
    # Create an optimizer
    optimizer = optax.adamw(learning_rate=learning_rate)
    # Define sharding for data parallel training
    mesh = Mesh(jax.devices(), ('batch',))
    sharding = NamedSharding(mesh, PartitionSpec('batch', None))

    # Testing out current state of the model
    test_prompt = "Once upon a time, there was a girl named Amy."
    generate_text(model, params, tokenizer, test_prompt)

    state = train_state.TrainState.create(
        apply_fn=model.apply,
        params=params,
        tx=optimizer
    )

    # Perform post training
    print("Start training model")
    state = train_model(state, tiny_stories_dl, num_epochs, sharding)

    # Final text generation
    generate_text(model, state.params, tokenizer, test_prompt)

if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='Train Gemma model with custom parameters.')
    parser.add_argument('--maxlen', type=int, default=256, help='Maximum sequence length')
    parser.add_argument('--batch_size', type=int, default=128, help='Batch size')
    parser.add_argument('--datacount', type=int, default=296000, help='Number of data samples to use')
    args = parser.parse_args()

    run_training(maxlen=args.maxlen, batch_size=args.batch_size, datacount=args.datacount)

בסקריפט הזה, התנאים הבאים חלים:

  • Gemma3Tokenizer ממיר נתוני טקסט לטוקנים שהמודל יכול לעבד.
  • הפונקציה load_and_preprocess_data קוראת את נתוני האימון מקובץ, מפצלת אותם לסיפורים נפרדים ומשתמשת ב-tokenizer כדי להמיר את הטקסט לרצפים מרופדים של טוקנים.
  • הפונקציה generate_text מקבלת את המודל, את הפרמטרים שלו ואת ההנחיה ליצירת טקסט.
  • הפונקציה train_step מגדירה איטרציה אחת של אימון שכוללת העברה קדימה, חישוב הפסד (באמצעות אנטרופיה צולבת), חישוב גרדיאנט ועדכוני פרמטרים.
  • הפונקציה train_model מבצעת איטרציה במערך הנתונים למשך מספר מוגדר של תקופות, ובמהלכה היא קוראת לפונקציה train_step לכל אצווה.
  • הפונקציה run_training מתזמנת את כל התהליך של טעינת הנתונים, אתחול מודל Gemma 3 ‏ (Gemma3_270M) ואופטימיזציה, טעינת פרמטרים שאומנו מראש, הגדרת חלוקת נתונים לעיבוד מקביל, הפעלת יצירת טקסט לבדיקה, ביצוע לולאת האימון וביצוע יצירת טקסט סופית כדי להדגים את ההשפעה של כוונון עדין.
  • הסקריפט משתמש בספריית argparse כדי לקבל ארגומנטים בשורת הפקודה עבור הפרמטרים maxlen,‏ batch_size ו-datacount.

אחרי שבדקתם את סקריפט הכוונון העדין, אתם יכולים להוסיף אותו לקונטיינר כדי להריץ אותו ב-GKE.

העברה של סקריפט הכוונון העדין למאגר

לפני שמריצים את סקריפט הכוונון העדין באשכול GKE, צריך להוסיף אותו לקונטיינר. במדריך הזה נעשה שימוש בתמונה מ-AI של JAX כתמונת הבסיס.

  1. פותחים את Dockerfile באותה ספרייה שבה נמצא הקובץ Gemma3LLMTrain.py:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    FROM us-docker.pkg.dev/cloud-tpu-images/jax-ai-image/tpu:jax0.7.2-rev1
    RUN apt-get update && apt-get install -y wget && rm -rf /var/lib/apt/lists/*
    
    RUN pip install --upgrade pip
    RUN pip install gemma grain
    
    WORKDIR /app
    
    # Copy your training script into the container
    COPY Gemma3LLMTrain.py .
    

    קובץ ה-Dockerfile הזה מתקין את הרכיבים התלויים הנדרשים ומעתיק את הקובץ Gemma3LLMTrain.py לקונטיינר.

  2. יוצרים את קובץ האימג' של Docker ומעבירים אותו בדחיפה למאגר אימג'ים:

    export REPOSITORY=REPOSITORY_NAME
    export IMAGE_NAME="jax-gemma3-training"
    export IMAGE_TAG="latest"
    export DOCKERFILE_PATH="./Dockerfile"
    export IMAGE_URI="${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPOSITORY}/${IMAGE_NAME}:${IMAGE_TAG}"
    
    docker build -t "${IMAGE_URI}" -f "${DOCKERFILE_PATH}" .
    gcloud auth configure-docker "${REGION}-docker.pkg.dev" -q
    docker push "${IMAGE_URI}"
    

    מחליפים את REPOSITORY_NAME בשם המאגר שלכם ב-Artifact Registry.

  3. מוסיפים קישורי תפקידים לחשבון השירות:

    export PROJECT_NUMBER=$(gcloud projects describe $PROJECT_ID --format 'get(projectNumber)')
    gcloud artifacts repositories add-iam-policy-binding ${REPOSITORY} \
        --project=${PROJECT_ID} \
        --location=${REGION} \
        --member="serviceAccount:${PROJECT_NUMBER}-compute@developer.gserviceaccount.com" \
        --role="roles/artifactregistry.reader"
    

אחרי שהתמונה נמצאת במאגר, אפשר לפרוס את משימת הכוונון העדין באשכול GKE.

פריסת משימת הכוונון העדין של מודל LLM

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

  1. פותחים את קובץ המניפסט training_singlehost.yaml:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    apiVersion: batch/v1
    kind: Job
    metadata:
      name: jax-gemma3-train-singlehost
    spec:
      template:
        metadata:
          annotations:
            gke-gcsfuse/volumes: "true"
        spec:
          serviceAccountName: jaxserviceaccout
          containers:
          - name: training-container
            image: ${IMAGE_URI}
            imagePullPolicy: "Always"
            command: ["python", "Gemma3LLMTrain.py", "--maxlen", "256", "--batch_size", "64", "--datacount", "355120"]
            resources:
              limits:
                google.com/tpu: 1
            volumeMounts:
            - name: gcs-fuse-csi-ephemeral
              mountPath: /data
          nodeSelector:
            cloud.google.com/gke-tpu-accelerator: tpu-v6e-slice
            cloud.google.com/gke-tpu-topology: 1x1
          restartPolicy: Never
          volumes:
          - name: gcs-fuse-csi-ephemeral
            csi:
              driver: gcsfuse.csi.storage.gke.io
              volumeAttributes:
                bucketName: ${GCS_BUCKET_NAME}
                mountOptions: "implicit-dirs,file-cache:enable-parallel-downloads:true,file-cache:parallel-downloads-per-file:100,file-cache:max-parallel-downloads:-1,file-cache:download-chunk-size-mb:10,file-cache:max-size-mb:-1"
      backoffLimit: 1
  2. החלת המניפסט:

    envsubst < training_singlehost.yaml | kubectl apply -f -
    

‫GKE יוצר Job שמפעיל Pod בצומת TPU Trillium ‏ (v6e). ה-Pod הזה מריץ את סקריפט הכוונון העדין של Python, שמקבל גישה לנתוני הכוונון העדין מקטגוריית Cloud Storage שצוינה, שנטענה בנתיב /data באמצעות Cloud Storage FUSE. הסקריפט מבצע כוונון עדין של מודל Gemma.

מעקב אחרי עבודת ההדרכה

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

איך בודקים את התקדמות הכוונון העדין

  1. מציגים את רשימת ה-Pods:

    # Find the Pods
    kubectl get pods
    
  2. עוקבים אחרי פלט היומן:

    kubectl logs -f pods/POD_NAME
    

    מחליפים את POD_NAME בשם ה-Pod.

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

    Global device count: 1
    Batch size: 128, Max length: 256, Data count: 96000
    I1028 00:12:55.925999 1387 google_auth_provider.cc:181] Running on GCE, using service account ...
    Generating response for: Once upon a time, there was a girl named Amy.
    Response:
    Amy lived in a small house. The house was in a big field. Amy liked to play in the big field. She
    Start training model
    Loss after batch 0: 10.25
    Loss after batch 10: 4.3125
    .
    .
    .
    Loss after batch 740: 1.41406
    Completed training model. Total time for training 294.6791355609894 seconds
    Generating response for: Once upon a time, there was a girl named Amy.
    Response:
    She loved to play with her toys. One day, Amy's mom told her that she had to go to the store to
    
  3. ניתוח הפלט:

    • הקו Global device count: 1 מציין את ליבות ה-TPU שהיו בשימוש.
    • המודל יוצר טקסט סביר לפני הרצת הכוונון העדין, כי הוא נטען מנקודת ביקורת שאומנה מראש.
    • הפלט שנוצר אחרי כוונון עדין דומה יותר להתחלה של סיפור קצר, מה שמצביע על כך שהמודל לומד ממערך הנתונים החדש.
    • כוונון עדין של המודל על מערך הנתונים המלא אמור להניב תוצאות מדויקות עוד יותר.

התבוננות במדדים

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

הגדרות חלופיות לכוונון עדין

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

בחירת מודל

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

רשימה מלאה של המודלים הזמינים מופיעה בתיעוד של Gemma.

הגדרות של מאגר צמתים

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

בכרטיסיות הבאות מוצגות הוראות ליצירת מאגרי צמתים עם כמה מארחים וכמה פרוסות:

כמה מארחים

  1. ב-Cloud Shell, מריצים את הפקודה הבאה:

    gcloud container node-pools create jax-tpu-multihost1 \
        --cluster=${CLUSTER_NAME} \
        --machine-type=ct6e-standard-4t \
        --num-nodes=2 \
        --tpu-topology=2x4 \
        --location=${REGION} \
        --node-locations=${ZONE}
    

    ‫GKE יוצר מאגר צמתים של TPU Trillium עם טופולוגיה של 2x4 ושני צמתים.

  2. פותחים את הגדרת המשרה training_multihost_jobset.yaml:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    apiVersion: jobset.x-k8s.io/v1alpha2
    kind: JobSet
    metadata:
      name: jax-gemma3-train-multihost
    spec:
      replicatedJobs:
        - name: trainers
          replicas: 1
          template:
            spec:
              parallelism: 2
              completions: 2
              backoffLimit: 1
              template:
                metadata:
                  annotations:
                    gke-gcsfuse/volumes: "true"
                spec:
                  serviceAccountName: jaxserviceaccout
                  nodeSelector:
                    cloud.google.com/gke-tpu-accelerator: tpu-v6e-slice
                    cloud.google.com/gke-tpu-topology: 2x4
                    cloud.google.com/gke-nodepool: jax-tpu-multihost1
                  containers:
                  - name: training-container
                    image: ${IMAGE_URI} 
                    imagePullPolicy: "Always"
                    ports:
                      - containerPort: 8471
                    command: ["python", "Gemma3LLMTrain.py", "--maxlen", "256", "--batch_size", "64", "--datacount", "5120"]
                    resources:
                      limits:
                        google.com/tpu: 4
                    volumeMounts:
                    - name: gcs-fuse-csi-ephemeral
                      mountPath: /data
                  volumes:
                    - name: gcs-fuse-csi-ephemeral
                      csi:
                        driver: gcsfuse.csi.storage.gke.io
                        volumeAttributes:
                          bucketName: ${GCS_BUCKET_NAME}
                          mountOptions: "implicit-dirs,file-cache:enable-parallel-downloads:true,file-cache:parallel-downloads-per-file:100,file-cache:max-parallel-downloads:-1,file-cache:download-chunk-size-mb:10,file-cache:max-size-mb:-1"
    
  3. פורסים את משימת הכוונון העדין:

    envsubst < training_multihost_jobset.yaml | kubectl apply -f -
    

Multislice

  1. ב-Cloud Shell, מריצים את הפקודה הבאה:

    gcloud container node-pools create jax-tpu-multihost1 \
      --cluster=${CLUSTER_NAME} \
      --machine-type=ct6e-standard-4t \
      --num-nodes=2 \
      --tpu-topology=2x4 \
      --location=${REGION} \
      --node-locations=${ZONE}
    
    gcloud container node-pools create jax-tpu-multihost2 \
      --cluster=${CLUSTER_NAME} \
      --machine-type=ct6e-standard-4t \
      --num-nodes=2 \
      --tpu-topology=2x4 \
      --location=${REGION} \
      --node-locations=${ZONE}
    

    ‫GKE יוצר שני מאגרי צמתים של TPU Trillium. לכל מאגר צמתים יש 2x4 טופולוגיה ושני צמתים.

  2. פותחים את הגדרת המשרה training_multislice_jobset.yaml:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    apiVersion: jobset.x-k8s.io/v1alpha2
    kind: JobSet
    metadata:
      name: jax-gemma3-train-multislice
    spec:
      replicatedJobs:
        - name: trainers
          replicas: 2
          template:
            spec:
              parallelism: 2
              completions: 2
              backoffLimit: 1
              template:
                metadata:
                  annotations:
                    gke-gcsfuse/volumes: "true"
                spec:
                  serviceAccountName: jaxserviceaccout
                  nodeSelector:
                    cloud.google.com/gke-tpu-accelerator: tpu-v6e-slice
                    cloud.google.com/gke-tpu-topology: 2x4
                  containers:
                  - name: training-container
                    image: ${IMAGE_URI}
                    imagePullPolicy: "Always"
                    ports:
                      - containerPort: 8471
                    command: ["python", "Gemma3LLMTrain.py", "--maxlen", "256", "--batch_size", "64", "--datacount", "5120"]
                    resources:
                      limits:
                        google.com/tpu: 4
                    volumeMounts:
                    - name: gcs-fuse-csi-ephemeral
                      mountPath: /data
                  volumes:
                    - name: gcs-fuse-csi-ephemeral
                      csi:
                        driver: gcsfuse.csi.storage.gke.io
                        volumeAttributes:
                          bucketName: ${GCS_BUCKET_NAME}
                          mountOptions: "implicit-dirs,file-cache:enable-parallel-downloads:true,file-cache:parallel-downloads-per-file:100,file-cache:max-parallel-downloads:-1,file-cache:download-chunk-size-mb:10,file-cache:max-size-mb:100"
    
  3. פורסים את משימת הכוונון העדין:

    envsubst < training_multislice_jobset.yaml | kubectl apply -f -
    

ניתוח ביצועים ואופטימיזציה

כדי לנתח את הביצועים של כוונון עדין של למידת מכונה ולבצע אופטימיזציה שלהם, אפשר להשתמש ב-XProf. ‫XProf הוא חבילת כלים שיוצרת פרופילים של עומסי עבודה של ML שנבנו באמצעות JAX, ‏ TensorFlow או PyTorch/XLA, ובודקת אותם. הכלי XProf מציג עקבות של ביצוע, שימוש בזיכרון ונתונים אחרים, וכך מאפשר לכם לשפר את המודלים ואת הגדרות האימון כדי להשיג יעילות טובה יותר ואימון מהיר יותר.

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

  • מתקינים את חבילת xprof. משנים את סקריפט ההכשרה כדי להפעיל את שרת XProf.
  • משנים את מניפסט העבודה של Kubernetes כך שיכלול נקודת טעינה של נפח ליומני XProf.
  • נותנים לחשבון השירות הרשאות לכתוב יומני XProf לקטגוריה של Cloud Storage.
  • מריצים את XProf בתוך ה-Pod ומגדירים העברת יציאות כדי לגשת ללוח הבקרה של XProf.

התקנת חבילת XProf

  1. עוברים לספרייה שמכילה את הדוגמאות של XProf:

      cd ai-ml/llm-training-jax-tpu-gemma3/xprof-enabled
    
  2. יוצרים את קובץ האימג' של Docker ומעבירים אותו בדחיפה למאגר אימג'ים:

    export REPOSITORY=REPOSITORY_NAME
    export IMAGE_NAME="jax-gemma3-training-xp"
    export IMAGE_TAG="latest"
    export DOCKERFILE_PATH="./Dockerfile"
    export IMAGE_URI="${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPOSITORY}/${IMAGE_NAME}:${IMAGE_TAG}"
    
    docker build -t "${IMAGE_URI}" -f "${DOCKERFILE_PATH}" .
    gcloud auth configure-docker "${REGION}-docker.pkg.dev" -q
    docker push "${IMAGE_URI}"
    

    מחליפים את REPOSITORY_NAME בשם המאגר שלכם ב-Artifact Registry.

  3. מריצים את הסקריפט Dockerfile:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    FROM us-docker.pkg.dev/cloud-tpu-images/jax-ai-image/tpu:jax0.7.2-rev1
    RUN apt-get update && apt-get install -y wget && rm -rf /var/lib/apt/lists/*
    
    RUN pip install --upgrade pip
    RUN pip install gemma grain equinox
    RUN pip install xprof
    
    WORKDIR /app
    
    # Copy your training script into the container
    COPY Gemma3LLMTrain.py .
    

    קובץ ה-Dockerfile הזה מתקין יחסי תלות של XProf.

מעתיקים את סקריפט הכוונון העדין אל הקונטיינר.

בקטע הזה, יוצרים ומחילים מניפסט של Kubernetes Job שכולל את נקודות הטעינה (mount) של נפח האחסון שנדרשות ליומני XProf.

  1. פותחים את הגדרת המשרה training_singlehost.yaml:

    # Copyright 2026 Google LLC
    #
    # Licensed under the Apache License, Version 2.0 (the "License");
    # you may not use this file except in compliance with the License.
    # You may obtain a copy of the License at
    #
    #     http://www.apache.org/licenses/LICENSE-2.0
    #
    # Unless required by applicable law or agreed to in writing, software
    # distributed under the License is distributed on an "AS IS" BASIS,
    # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    # See the License for the specific language governing permissions and
    # limitations under the License.
    
    apiVersion: batch/v1
    kind: Job
    metadata:
      name: jax-gemma3-train-singlehost
    spec:
      template:
        metadata:
          annotations:
            gke-gcsfuse/volumes: "true"
        spec:
          serviceAccountName: jaxserviceaccout
          containers:
          - name: training-container
            image: ${IMAGE_URI}
            imagePullPolicy: "Always"
            command: ["python", "Gemma3LLMTrain.py", "--maxlen", "256", "--batch_size", "64", "--datacount", "851200"]
            resources:
              limits:
                google.com/tpu: 1
            volumeMounts:
            - name: gcs-fuse-csi-ephemeral
              mountPath: /data
            - name: gcs-fuse-csi-ephemeral2
              mountPath: /xprof
          nodeSelector:
            cloud.google.com/gke-tpu-accelerator: tpu-v6e-slice
            cloud.google.com/gke-tpu-topology: 1x1
          restartPolicy: Never
          volumes:
          - name: gcs-fuse-csi-ephemeral
            csi:
              driver: gcsfuse.csi.storage.gke.io
              volumeAttributes:
                bucketName: ${GCS_BUCKET_NAME}
                mountOptions: "implicit-dirs,file-cache:enable-parallel-downloads:true,file-cache:parallel-downloads-per-file:100,file-cache:max-parallel-downloads:-1,file-cache:download-chunk-size-mb:10,file-cache:max-size-mb:-1"
          - name: gcs-fuse-csi-ephemeral2
            csi:
              driver: gcsfuse.csi.storage.gke.io
              volumeAttributes:
                bucketName: ${XPROF_GCS_BUCKET_NAME}
                mountOptions: "implicit-dirs,file-cache:enable-parallel-downloads:true,file-cache:parallel-downloads-per-file:100,file-cache:max-parallel-downloads:-1,file-cache:download-chunk-size-mb:10,file-cache:max-size-mb:-1"
      backoffLimit: 1
  2. החלת המניפסט:

    envsubst < training_singlehost.yaml | kubectl apply -f -
    

מתן הרשאות לחשבון השירות לכתוב יומנים של XProf

  1. כדי לאפשר לחשבון השירות לכתוב ולקרוא, מוסיפים את התפקיד "roles/storage.objectUser":

    export GSA_NAME="GSA_NAME" # Same as used in initial setup
    
    # Automatically get the current project ID
    export PROJECT_ID=$(gcloud config get-value project)
    
    # Cloud Storage Bucket details
    export XPROF_GCS_BUCKET_NAME="XPROF_GCS_BUCKET_NAME"
    
    # Derived Variables
    export GSA_EMAIL="${GSA_NAME}@${PROJECT_ID}.iam.gserviceaccount.com"
    
    gcloud storage buckets add-iam-policy-binding "gs://${XPROF_GCS_BUCKET_NAME}" \
        --member="serviceAccount:${GSA_EMAIL}" \
        --role="roles/storage.objectUser" \
        --project="${PROJECT_ID}"
    

    מחליפים את מה שכתוב בשדות הבאים:

    • GSA_NAME: השם של חשבון השירות של Google שרוצים להעניק לו את התפקיד.
    • XPROF_GCS_BUCKET_NAME: שם הקטגוריה שרוצים להעניק לה את התפקיד.
  2. מריצים את XProf בתוך ה-Pod:

    kubectl exec POD_NAME -c training-container -it -- bash # exec into the container
    xprof --port 9001 --logdir /xprof # start xprof
    

    מחליפים את POD_NAME בשם ה-Pod.

גישה למרכז הבקרה של XProf

  1. מגדירים העברה ליציאה אחרת לשרת XProf ב-Pod:

    kubectl port-forward POD_NAME 9001:9001
    
  2. בסרגל הכתובות של הדפדפן, מזינים את הטקסט הבא:

    http://localhost:9001/
    

    הכלי XProf Trace Viewer ייפתח.

  3. בחלון TensorBoard, לוחצים על Capture profile (לכידת פרופיל).

  4. בשדה Profile Service URL(s) or TPU name (כתובות URL של שירות פרופילים או שם TPU), מזינים localhost:9002.

  5. כדי ללכוד פרטים נוספים, בקטע Host Trace (TraceMe) Level (רמת מעקב אחר המארח (TraceMe)), בוחרים באפשרות verbose (מפורט) ומפעילים את האפשרות Python trace logging (רישום מעקב של Python).

  6. כדי להציג את לוח הבקרה, לוחצים על Capture (צילום).

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

דוגמה לכלי XProf trace viewer שבו מוצג גרף של מטריצת ביצועים

אפשרויות נוספות ליצירת פרופילים לניתוח הביצועים של עומס העבודה של האימון מפורטות במאמר בנושא יצירת פרופילים של חישובים במאמרי העזרה של JAX.

כוונון עדין בסביבות ייצור

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

כדי למזער את אובדן ההתקדמות במהלך תקלה, מומלץ להגדיר יצירת נקודות ביקורת (checkpointing) של עומסי עבודה ארוכי טווח של אימון או כוונון עדין בסביבת ייצור. מידע נוסף על הגדרת נקודות ביקורת רב-שכבתיות זמין במאמר אימון מודלים של למידת מכונה בקנה מידה גדול ב-GKE באמצעות נקודות ביקורת רב-שכבתיות.

הסרת המשאבים

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

מחיקת המשאבים הבודדים

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

  1. מוחקים את המשאבים שיצרתם במדריך הזה:

    gcloud container clusters delete ${CLUSTER_NAME} --location=${REGION}
    gcloud storage rm --recursive gs://${GCS_BUCKET_NAME}
    gcloud artifacts docker images delete ${IMAGE_URI} --delete-tags
    
  2. אם אתם לא צריכים את הנתונים שנוצרו על ידי XProf, אתם יכולים להסיר את קטגוריית Cloud Storage שבה נעשה שימוש על ידי XProf:

    gcloud storage rm --recursive gs://${XPROF_GCS_BUCKET_NAME}
    

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