diff --git a/docs/guides/checkpointing_solutions/convert_checkpoint.md b/docs/guides/checkpointing_solutions/convert_checkpoint.md index bece6db18b..d0dba429ed 100644 --- a/docs/guides/checkpointing_solutions/convert_checkpoint.md +++ b/docs/guides/checkpointing_solutions/convert_checkpoint.md @@ -12,6 +12,8 @@ The following models are supported: | :---------------------- | :--------------------- | :-------------------: | :---------------------: | :-------------------: | :---------------------: | | **Gemma2** | 2B, 9B, 27B | √ | √ | √ | √ | | **Gemma3** (Multimodal) | 4B, 12B, 27B | √ | √ | √ | √ | +| **Gemma4** | 26B (MoE), 31B | √ | √ | √ | √ | +| **Gemma4 Small** | E2B, E4B | - | √ | - | √ | | **Llama3.1** | 8B, 70B, 450B | √ | √ | √ | √ | | **Qwen2.5** | 1.5B, 7B, 14B | √ | √ | √ | √ | | **Qwen3** | 0.6B, 4B, 8B, 14B, 32B | √ | √ | √ | √ | @@ -23,6 +25,8 @@ The following models are supported: | **DeepSeek3.2** | 671B | √ | √ | - | - | | **Qwen3 Next** | 80B | √ | √ | √ | √ | +> **Note:** *Gemma4 Small* (E2B / E4B) uses the `gemma4_small` decoder block, which has per-layer KV sharing and is incompatible with `nn.scan`. Only unscanned (`scan_layers=False`) conversion is supported for these variants. + ## Prerequisites - MaxText must be installed in a Python virtual environment using the `maxtext[tpu]` option. For instructions on installing MaxText on your VM, please refer to the official [installation documentation](install-from-source). @@ -37,6 +41,8 @@ Use the `to_maxtext.py` script to convert a Hugging Face model checkpoint into a > **Note:** For more information, checkout [qwen3-4b example script](https://github.com/AI-Hypercomputer/maxtext/blob/main/tests/end_to_end/tpu/qwen3/4b/test_qwen3_to_mt.sh) and [gemma3-4b example script](https://github.com/AI-Hypercomputer/maxtext/blob/main/tests/end_to_end/tpu/gemma3/4b/test_gemma3_to_mt.sh). +> **Note (Gemma 4 E2B / E4B):** The `gemma4-e2b` and `gemma4-e4b` variants use the `gemma4_small` decoder block, which has per-layer KV sharing and is **incompatible with `nn.scan`**. Conversion must therefore be run with `scan_layers=False` (unscanned), and any downstream training/inference run must also use `scan_layers=False`. Additionally, if you want to convert the **base** (pretrained) model rather than the instruction-tuned default, pass `--hf_model_path=google/gemma-4-E4B` (or `google/gemma-4-E2B`) explicitly — `HF_IDS` defaults to the `-it` repository. See [convert_gemma4_base.sh](https://github.com/AI-Hypercomputer/maxtext/blob/main/tests/end_to_end/tpu/gemma4/e4b/convert_gemma4_base.sh) for a full example. + ### Setup Environment ```bash @@ -94,6 +100,8 @@ Use the `to_huggingface.py` script to convert a MaxText checkpoint into the Hugg > **Note:** For more information, checkout [qwen3-4b example script](https://github.com/AI-Hypercomputer/maxtext/blob/main/tests/end_to_end/tpu/qwen3/4b/test_qwen3_to_hf.sh). +> **Note (Gemma 4 E2B / E4B):** The `gemma4-e2b` and `gemma4-e4b` variants use the `gemma4_small` decoder block and must be converted with `scan_layers=False`. If your MaxText checkpoint was fine-tuned from the base model, also pass `--hf_model_path=google/gemma-4-E4B` (or `google/gemma-4-E2B`) so the exported checkpoint bundles the base tokenizer — otherwise `to_huggingface` sources the tokenizer from `HF_IDS[gemma4-e4b]` = `google/gemma-4-E4B-it`. + ### Setup Environment ```bash diff --git a/docs/tutorials/posttraining/rl_gemma4_e4b.md b/docs/tutorials/posttraining/rl_gemma4_e4b.md new file mode 100644 index 0000000000..37cfd2f23b --- /dev/null +++ b/docs/tutorials/posttraining/rl_gemma4_e4b.md @@ -0,0 +1,152 @@ + + +# Reinforcement Learning with gemma4-e4b on Multi-Host TPUs + +This tutorial provides step-by-step instructions for setting up the environment +and training the gemma4-e4b model with GRPO on the [OpenMathInstruct-2 dataset](https://huggingface.co/datasets/nvidia/OpenMathInstruct-2) on a Cloud TPU v6e (Trillium) GKE cluster using a `v6e-32` (4x8) slice. + +## Prerequisites + +Before starting, ensure you have: + +- Access to a Google Cloud Project with TPU quotas. +- A Hugging Face account with an access token for downloading models (the `google/gemma-4-E4B` and `google/gemma-4-E4B-it` repositories are gated; request access before proceeding). +- Permissions for Google Artifact Registry (Artifact Registry Writer role). +- Prerequisites for XPK installed (follow [official documentation](https://github.com/AI-Hypercomputer/xpk/blob/main/docs/installation.md#1-prerequisites)). +- A Pathways-ready GKE cluster (see [create GKE cluster](https://docs.cloud.google.com/ai-hypercomputer/docs/workloads/pathways-on-cloud/create-gke-cluster)). +- **Docker** installed and configured for sudoless use. Follow the steps to [configure sudoless Docker](https://docs.docker.com/engine/install/linux-postinstall/). + +## Setup Environment Variables + +Set up the following environment variables to configure your training run. Replace +placeholders with your actual values. + +```bash +# Your GCP project ID. +# If you've already set it in your local config, you can retrieve it via: +# gcloud config get-value project +export PROJECT_ID= + +# The name of your GKE cluster. +export CLUSTER_NAME= + +# The GCP location of your GKE cluster. +export ZONE= # e.g., 'us-central1' or 'us-central1-a' + +# Use a GCS bucket you own to store logs and checkpoints. +export BASE_OUTPUT_DIRECTORY= # e.g., gs://my-bucket/maxtext-runs +``` + +## Authenticate with Hugging Face + +To download the `gemma4-e4b` model checkpoint from Hugging Face, you need to authenticate using your Hugging Face account credentials. Run the following command and follow the prompts to log in: + +```bash +hf auth login +``` + +## Get Your MaxText Compatible Model Checkpoint + +### Option 1: Using an existing MaxText checkpoint + +If you already have a MaxText-compatible model checkpoint, simply set the +following environment variable and move on to the next section. + +```bash +export MAXTEXT_CKPT_PATH= # e.g., gs://my-bucket/my-model-checkpoint/0/items +``` + +### Option 2: Converting from a Hugging Face checkpoint + +Refer to [Hugging Face to MaxText](hf-to-maxtext) to convert a Hugging Face checkpoint to MaxText format. After conversion finishes, set `MAXTEXT_CKPT_PATH` to the converted MaxText checkpoint path. + +```bash +export MAXTEXT_CKPT_PATH= # e.g., gs://my-bucket/my-model-checkpoint/0/items +``` + +> **Note (Gemma 4 E4B specifics):** +> +> - For the `gemma4-e4b` model, you must run the conversion with `scan_layers=False` (the `gemma4_small` decoder block is incompatible with `nn.scan`); the resulting unscanned checkpoint matches the `scan_layers=False` setting used for the RL run below. +> - This RL recipe fine-tunes the **base** model (`google/gemma-4-E4B`), not the instruction-tuned default. Pass `--hf_model_path=google/gemma-4-E4B` explicitly when running `to_maxtext` — otherwise MaxText defaults to `HF_IDS[gemma4-e4b]` = `google/gemma-4-E4B-it`. + +## Chat Template Configuration + +Unlike an instruction-tuned tokenizer, the `google/gemma-4-E4B` base tokenizer does **not** ship with a chat template, so the RL run must supply one explicitly. The [run_gemma4_e4b_rl.sh](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/trainers/post_train/rl/scripts/run_gemma4_e4b_rl.sh) script points at two files bundled with the repo (and therefore baked into your Docker image): + +- `data_template_path=maxtext/examples/chat_templates/openmathinstruct2_rl.json` — a stripped-down data template that injects only the system prompt and question (no turn markers). This avoids the doubled `` delimiters that would occur if the default `gsm8k_rl.json` (which bakes literal Gemma turn markers into the message content) were combined with the tokenizer's Jinja chat template. +- `chat_template_path=maxtext/examples/chat_templates/gemma-3-27b-chat_template.json` — the Gemma 3 chat template, wrapped in a JSON object with a `chat_template` key. It is applied by the tokenizer as the single source of truth for turn-marker formatting. + +Both files are already included under `src/maxtext/examples/chat_templates/`, so no additional setup is required. If you want to customize the prompt formatting, edit these files (or point the config at your own) before building the Docker image. + +## Run RL Workload + +### Build and Upload MaxText Docker Image + +For instructions on building and uploading the MaxText Docker image with post-training dependencies, please refer to the [official documentation](../../build_maxtext.md). + +### Submit your workload + +```bash +# The Docker image you pushed in the previous step +export CLOUD_IMAGE_NAME= +export DOCKER_IMAGE="gcr.io/${PROJECT_ID?}/${CLOUD_IMAGE_NAME?}" + +# Run the RL training script on your cluster +run_tutorial maxtext/trainers/post_train/rl/scripts/run_gemma4_e4b_rl.sh +``` + +> **Note:** The `run_gemma4_e4b_rl.sh` script pins the Pathways component images to specific versions via the xpk `--server-image` and `--proxy-server-image` flags (set through the `PATHWAYS_SERVER_IMAGE` and `PATHWAYS_PROXY_SERVER_IMAGE` variables at the top of the script). The `--server-image` is used for both the Pathways resource-manager server and the workers (the reference config uses the same image for both). Update these variables if you need a different Pathways release. + +### Monitor your workload + +To monitor your job's progress, you can use `kubectl` to check the `Jobset` status and stream logs directly from the pods. + +```bash +kubectl get jobset -n default ${WORKLOAD_NAME} + +# List pods to find the specific name +kubectl get pods | grep ${WORKLOAD_NAME} + +# stream the logs from the running pod (replace with the name you found) +kubectl logs -f +``` + +Alternatively, after running the bash script, you will also get a link to the Google Cloud Console to view your workload logs. Follow the link to view logs and monitor your workload's progress in the Cloud Console. + +### Monitor RL Metrics + +During RL training, you can monitor key metrics to track model convergence, reward trends, and hardware performance. + +To enable Tunix-managed metrics measurement, set `enable_tunix_perf_metrics` to `true` in RL configurations. Note that this flag is already set to `True` by default for this tutorial workload. When enabled, Tunix automatically collects and uploads these metrics to TensorBoard. + +For a complete list of collected metrics, see the [Tunix Metrics Documentation](https://tunix.readthedocs.io/en/latest/metrics.html). Key metrics to monitor include: + +- **Model Quality & Reward Metrics:** + - `rewards/mean`: The average reward across the batch (crucial for tracking learning progress). + - `score/mean`: The average raw score from the reward model before applying the KL penalty. +- **Rollout & Generation Metrics:** + - `rollout_time`: How long each rollout step takes. + - `completions/mean_length`: The average token length of generated completions. + - `actor_dequeue_time`: The time spent waiting for data from the rollout workers (relevant when async rollout is enabled). +- **Performance & Efficiency Metrics:** + - `step_time_sec`: The execution time for a single training step. + +## Convert Checkpoint to Hugging Face Format + +Refer to [MaxText to Hugging Face](maxtext-to-hf) to convert a MaxText checkpoint back to Hugging Face format. + +> **Note (Gemma 4 E4B specifics):** Because this recipe fine-tunes the **base** model, pass `--hf_model_path=google/gemma-4-E4B` to `to_huggingface` so the exported checkpoint bundles the base tokenizer. Without it, `to_huggingface` sources the tokenizer from `HF_IDS[gemma4-e4b]` = `google/gemma-4-E4B-it` (the instruction-tuned model). Also keep `scan_layers=False`, since the `gemma4_small` decoder block is not compatible with scanned layers. diff --git a/src/maxtext/checkpoint_conversion/to_huggingface.py b/src/maxtext/checkpoint_conversion/to_huggingface.py index f0e94d4829..e23c7f2309 100644 --- a/src/maxtext/checkpoint_conversion/to_huggingface.py +++ b/src/maxtext/checkpoint_conversion/to_huggingface.py @@ -108,6 +108,15 @@ "with values from the MaxText config. If False, raises a ValueError on mismatch.", ) +flags.DEFINE_string( + "hf_model_path", + None, + "(Optional) Customized remote HF repo (or local path) to source the tokenizer/processor " + "from. If not specified, defaults to maxtext.utils.globals.HF_IDS[model_name]. Use this to " + "bundle a different tokenizer than the default (e.g., the base model's tokenizer instead of " + "the instruction-tuned one).", +) + FLAGS = flags.FLAGS @@ -493,7 +502,7 @@ def main(argv: Sequence[str]) -> None: if model_key not in HF_IDS: raise ValueError(f"HF Tokenizer ID not found for model key: {model_key}") hf_token = config.hf_access_token - hf_tokenizer_id = HF_IDS[model_key] + hf_tokenizer_id = FLAGS.hf_model_path or HF_IDS[model_key] tokenizer = AutoTokenizer.from_pretrained(hf_tokenizer_id, token=hf_token) # For multi-modal case: diff --git a/src/maxtext/examples/chat_templates/gemma-3-27b-chat_template.json b/src/maxtext/examples/chat_templates/gemma-3-27b-chat_template.json new file mode 100644 index 0000000000..719b0cd0d7 --- /dev/null +++ b/src/maxtext/examples/chat_templates/gemma-3-27b-chat_template.json @@ -0,0 +1,3 @@ +{ + "chat_template": "{{ bos_token }}\n{%- if messages[0]['role'] == 'system' -%}\n {%- if messages[0]['content'] is string -%}\n {%- set first_user_prefix = messages[0]['content'] + '\n\n' -%}\n {%- else -%}\n {%- set first_user_prefix = messages[0]['content'][0]['text'] + '\n\n' -%}\n {%- endif -%}\n {%- set loop_messages = messages[1:] -%}\n{%- else -%}\n {%- set first_user_prefix = \"\" -%}\n {%- set loop_messages = messages -%}\n{%- endif -%}\n{%- for message in loop_messages -%}\n {%- if (message['role'] == 'user') != (loop.index0 % 2 == 0) -%}\n {{ raise_exception(\"Conversation roles must alternate user/assistant/user/assistant/...\") }}\n {%- endif -%}\n {%- if (message['role'] == 'assistant') -%}\n {%- set role = \"model\" -%}\n {%- else -%}\n {%- set role = message['role'] -%}\n {%- endif -%}\n {{ '' + role + '\n' + (first_user_prefix if loop.first else \"\") }}\n {%- if message['content'] is string -%}\n {{ message['content'] | trim }}\n {%- elif message['content'] is iterable -%}\n {%- for item in message['content'] -%}\n {%- if item['type'] == 'image' -%}\n {{ '' }}\n {%- elif item['type'] == 'text' -%}\n {{ item['text'] | trim }}\n {%- endif -%}\n {%- endfor -%}\n {%- else -%}\n {{ raise_exception(\"Invalid content type\") }}\n {%- endif -%}\n {{ '\n' }}\n{%- endfor -%}\n{%- if add_generation_prompt -%}\n {{'model\n'}}\n{%- endif -%}\n" +} diff --git a/src/maxtext/examples/chat_templates/openmathinstruct2_rl.json b/src/maxtext/examples/chat_templates/openmathinstruct2_rl.json new file mode 100644 index 0000000000..3962f96deb --- /dev/null +++ b/src/maxtext/examples/chat_templates/openmathinstruct2_rl.json @@ -0,0 +1,4 @@ +{ + "SYSTEM_PROMPT": "You are given a problem. Think about the problem and provide your reasoning. Place it between {reasoning_start_token} and {reasoning_end_token}. Then, provide the final answer (i.e., just one numerical value) between {solution_start_token} and {solution_end_token}.", + "TEMPLATE": "{system_prompt}\n\n{question}" +} diff --git a/src/maxtext/trainers/post_train/rl/scripts/run_gemma4_e4b_rl.sh b/src/maxtext/trainers/post_train/rl/scripts/run_gemma4_e4b_rl.sh new file mode 100644 index 0000000000..4ec84d7b46 --- /dev/null +++ b/src/maxtext/trainers/post_train/rl/scripts/run_gemma4_e4b_rl.sh @@ -0,0 +1,180 @@ +#!/bin/bash + +# This script launches a Reinforcement Learning (RL) training workload for the +# gemma4-e4b model on a GKE cluster using XPK. + +set -e + +# --- Environment Setup --- +if ! pip show xpk &> /dev/null; then + echo "xpk not found in the environment. Please install it by running:" + echo "uv pip install -e .[runner] --resolution=lowest" + exit 1 +fi + +# --- Environment Variables --- +export PROJECT_ID="${PROJECT_ID:-}" # GCP project ID where the v6e cluster is deployed +export CLUSTER_NAME="${CLUSTER_NAME:-}" # Name of your v6e (Trillium) cluster +export ZONE="${ZONE:-}" # Zone where your v6e cluster is deployed +export BASE_OUTPUT_DIRECTORY="${BASE_OUTPUT_DIRECTORY:-}" # GCS bucket path for outputs (e.g., gs://my-bucket/outputs) +export DOCKER_IMAGE="${DOCKER_IMAGE:-}" # Full path to the Docker image you pushed (e.g., gcr.io/my-project/my-image:tag) +export MAXTEXT_CKPT_PATH="${MAXTEXT_CKPT_PATH:-}" # GCS path of the MaxText checkpoint you want to fine-tune from (e.g., gs://my-bucket/checkpoints/maxtext-ckpt) +export TPU_TYPE="v6e-32" +export WORKLOAD_NAME="rl-$(date +%Y%m%d-%H%M)" + +# Pathways component images, pinned to specific versions (mirrors the reference config). +# Note: --server-image is used for BOTH the Pathways resource-manager server and the +# workers; in the reference config the worker and pathwaysServer images are identical. +export PATHWAYS_SERVER_IMAGE="us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server@sha256:d03068b39a8a2fab0621086ccb7c9445ce17ad34f1520159ca4ecd395346a162" +export PATHWAYS_PROXY_SERVER_IMAGE="us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server@sha256:6342a3cae2818d2f887d396e9ae4156b3316f79fef186d71ff16d535b44e1724" + +# --- Variable Validation --- +if [ -z "$PROJECT_ID" ]; then + echo "Error: PROJECT_ID is not set. Please set it in the script or as an environment variable." + exit 1 +fi +if [ -z "$CLUSTER_NAME" ]; then + echo "Error: CLUSTER_NAME is not set. Please set it in the script or as an environment variable." + exit 1 +fi +if [ -z "$ZONE" ]; then + echo "Error: ZONE is not set. Please set it in the script or as an environment variable." + exit 1 +fi +if [ -z "$BASE_OUTPUT_DIRECTORY" ]; then + echo "Error: BASE_OUTPUT_DIRECTORY is not set. Please set it in the script or as an environment variable." + exit 1 +fi +if [ -z "$DOCKER_IMAGE" ]; then + echo "Error: DOCKER_IMAGE is not set. Please set it in the script or as an environment variable." + exit 1 +fi + +if [ -z "$MAXTEXT_CKPT_PATH" ]; then + echo "MAXTEXT_CKPT_PATH is not set. Please set it in the script or as an environment variable." + exit 1 +fi + +# XLA Flags (tuned for v6e / Trillium) +XLA_FLAGS="--xla_tpu_scoped_vmem_limit_kib=65536 \ +--xla_tpu_bf16_emission_mode=NATIVE_EMISSION \ +--xla_tpu_enable_sparse_core_reduce_scatter_v2=true \ +--xla_tpu_enable_sparse_core_collective_offload_all_gather=true \ +--xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true \ +--xla_tpu_use_single_sparse_core_for_all_gather_offload=true \ +--xla_tpu_enable_sparse_core_collective_offload_all_reduce=true \ +--xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true \ +--xla_tpu_enable_sparse_core_collective_offload_3d_all_gather=true \ +--xla_tpu_enable_all_gather_offload_tracing=true \ +--xla_tpu_use_tc_device_shape_on_sc=True \ +--xla_sc_disable_megacore_partitioning=True \ +--xla_enable_async_all_gather=true \ +--xla_tpu_prefer_async_allgather_to_allreduce=true \ +--xla_tpu_enable_latency_hiding_layer_scheduler=true \ +--xla_tpu_scheduler_percent_shared_memory_limit=150 \ +--xla_tpu_enable_layer_scheduler_for_dependent_collectives=true \ +--xla_tpu_enable_sparse_core_collective_offload_nd_reduce_scatter=true \ +--xla_tpu_enable_sparse_core_collective_aggregator=true" + +# MaxText command +MAXTEXT_COMMAND="JAX_RANDOM_WEIGHTS=1 \ +VLLM_ENABLE_V1_MULTIPROCESSING=0 \ +NEW_MODEL_DESIGN=1 \ +TF_CPP_MIN_LOG_LEVEL=0 \ +ENABLE_PJRT_COMPATIBILITY=true \ +ENABLE_PATHWAYS_PERSISTENCE=1 \ +TPU_BACKEND_TYPE=jax \ +VLLM_LOGGING_LEVEL=WARNING \ +JAX_PLATFORMS=proxy,cpu \ +JAX_BACKEND_TARGET=grpc://127.0.0.1:29000 \ +NUM_SLICES=1 \ +VLLM_ENGINE_READY_TIMEOUT_S=7200 \ +RPA_D_BLOCK_SIZES=1,1024,1,512 \ +TOKENIZERS_PARALLELISM=False \ +python3 -m maxtext.trainers.post_train.rl.train_rl \ +model_name=gemma4-e4b \ +tokenizer_type=huggingface \ +tokenizer_path=google/gemma-4-E4B \ +run_name=$WORKLOAD_NAME \ +async_scheduling=True \ +base_output_directory=$BASE_OUTPUT_DIRECTORY \ +chips_per_vm=4 \ +data_template_path=maxtext/examples/chat_templates/openmathinstruct2_rl.json \ +chat_template_path=maxtext/examples/chat_templates/gemma-3-27b-chat_template.json \ +remat_policy=full \ +decoder_layer_input=offload \ +megablox=True \ +sparse_matmul=True \ +use_custom_sort_vjp=True \ +attention=flash \ +use_tokamax_splash=True \ +sa_use_fused_bwd_kernel=True \ +sa_block_q=1024 \ +sa_block_kv=1024 \ +sa_block_kv_compute=512 \ +sa_block_q_dkv=1024 \ +sa_block_kv_dkv=1024 \ +sa_block_kv_dkv_compute=256 \ +mu_dtype=bfloat16 \ +grad_dtype=bfloat16 \ +use_iota_embed=True \ +batch_size=128 \ +num_batches=500 \ +learning_rate_schedule_steps=500 \ +num_test_batches=5 \ +eval_interval=100 \ +rl.num_generations=8 \ +rl.num_iterations=1 \ +rl.grpo_beta=0.05 \ +rl.grpo_epsilon=0.2 \ +debug.rl=True \ +gradient_clipping_threshold=1.0 \ +decode_sampling_temperature=0.8 \ +decode_sampling_top_k=50 \ +decode_sampling_nucleus_p=0.95 \ +hf_name=default \ +dataset_name=nvidia/OpenMathInstruct-2 \ +hf_train_files=hf://datasets/nvidia/OpenMathInstruct-2/data/train_1M-*.parquet \ +train_split=train_1M \ +eval_dataset_name=nvidia/OpenMathInstruct-2 \ +eval_split=test \ +learning_rate=1e-6 \ +max_prefill_predict_length=512 \ +max_target_length=4096 \ +kv_cache_buffer=512 \ +max_num_seqs=16 \ +max_num_batched_tokens=2048 \ +rollout_data_parallelism=8 \ +rollout_tensor_parallelism=2 \ +rollout_expert_parallelism=1 \ +rl.reshard_chunk_size=420 \ +enable_dp_attention=True \ +hbm_utilization_vllm=0.5 \ +scan_layers=False \ +allow_split_physical_axes=True \ +enable_tunix_perf_metrics=True \ +checkpoint_period=25 \ +max_num_checkpoints_to_keep=30 \ +enable_checkpointing=True \ +train_micro_batch_size=1 \ +rollout_micro_batch_size=16 \ +ici_tensor_parallelism=2 \ +load_parameters_path=$MAXTEXT_CKPT_PATH \ +vllm_hf_overrides='{architectures: [\"MaxTextForCausalLM\"]}' \ +vllm_additional_config='{\"maxtext_config\": {\"model_name\": \"gemma4-e4b\", \"model_call_mode\": \"inference\", \"enable_dp_attention\": true, \"allow_split_physical_axes\": true, \"log_config\": false, \"weight_dtype\": \"bfloat16\", \"prefuse_moe_weights\": true}}'" + +# Workload Creation +xpk workload create-pathways \ + --cluster=$CLUSTER_NAME \ + --project=$PROJECT_ID \ + --zone=$ZONE \ + --priority=medium \ + --max-restarts=0 \ + --tpu-type=$TPU_TYPE \ + --num-slices=1 \ + --docker-image="${DOCKER_IMAGE}" \ + --workload="${WORKLOAD_NAME}" \ + --server-image="${PATHWAYS_SERVER_IMAGE}" \ + --proxy-server-image="${PATHWAYS_PROXY_SERVER_IMAGE}" \ + --custom-pathways-proxy-server-args='${XLA_FLAGS}' \ + --command="${MAXTEXT_COMMAND}"