Skip to content

Commit 4344b39

Browse files
committed
Fix: Fixed Jetstream_benchmarking_servig DAG
- Point v5e tests to new reservation in CIENET-CMCS project - Point v6e tests to new reservation in tpu-prod-env-automated
1 parent d4c530a commit 4344b39

4 files changed

Lines changed: 118 additions & 78 deletions

File tree

dags/common/vm_resource.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,12 @@
1818
import enum
1919
from xlml.apis.xpk_cluster_config import XpkClusterConfig
2020

21+
# Cienet Reservation for V5e and V6e
22+
CIENET_NETWORK_PREFIX = f"projects/cienet-cmcs"
23+
CIENET_NETWORKS = f"{CIENET_NETWORK_PREFIX}/global/networks/mas-test"
24+
CIENET_V5E_SUBNETWORKS = (
25+
f"{CIENET_NETWORK_PREFIX}/regions/us-west4/subnetworks/mas-test"
26+
)
2127

2228
V5_NETWORKS_PREFIX = "projects/tpu-prod-env-automated"
2329
V5_NETWORKS = f"{V5_NETWORKS_PREFIX}/global/networks/mas-test"
@@ -27,7 +33,7 @@
2733
f"{V5_NETWORKS_PREFIX}/regions/europe-west4/subnetworks/mas-test-v2"
2834
)
2935
V6E_SUBNETWORKS = (
30-
f"{V5_NETWORKS_PREFIX}/regions/us-central2/subnetworks/mas-test"
36+
f"{V5_NETWORKS_PREFIX}/regions/southamerica-west1-a/subnetworks/mas-test"
3137
)
3238
# TODO: Figure V6E_GCE_NETWORK and V6E_GCE_SUBNETWORK
3339
V6E_GCE_NETWORK = "default"
@@ -72,6 +78,7 @@ class Project(enum.Enum):
7278
TPU_PROD_ENV_LARGE_ADHOC = "tpu-prod-env-large-adhoc"
7379
TPU_PROD_ENV_ONE_VM = "tpu-prod-env-one-vm"
7480
TPU_PROD_ENV_LARGE_CONT = "tpu-prod-env-large-cont"
81+
CIENET_CMCS = "cienet-cmcs"
7582

7683

7784
class ImageProject(enum.Enum):
@@ -124,6 +131,10 @@ class Zone(enum.Enum):
124131
US_EAST5_C = "us-east5-c"
125132
# reserved v5e in tpu-prod-env-multipod
126133
US_WEST4_B = "us-west4-b"
134+
# reserved v5e in cienet-cmcs
135+
US_WEST4_A = "us-west4-a"
136+
# reserved v6e in cienet-cmcs
137+
US_EAST4_B = "us-east4-b"
127138
# reserved v5e in cloud-tpu-inference-test
128139
US_WEST1_C = "us-west1-c"
129140
# reserved a3+ cluster in supercomputer-testing

dags/inference/configs/jetstream_benchmark_serving_gce_config.py

Lines changed: 54 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
PROJECT_NAME = Project.CLOUD_ML_AUTO_SOLUTIONS.value
2626
RUNTIME_IMAGE = RuntimeVersion.TPU_UBUNTU2204_BASE.value
2727
GCS_SUBFOLDER_PREFIX = test_owner.Team.INFERENCE.value
28+
VENV_DIR = "venv-312"
2829

2930

3031
def get_config(
@@ -55,8 +56,20 @@ def get_config(
5556
# Download jetstream and maxtext
5657
f"if [ ! -d maxtext ]; then git clone {maxtext_branch} https://github.com/google/maxtext.git; fi",
5758
f"if [ ! -d JetStream ]; then git clone {jetstream_branch} https://github.com/google/JetStream.git; fi",
58-
"sudo apt-get -y update",
59-
"sudo apt-get -y install jq",
59+
"sudo systemctl stop unattended-upgrades",
60+
"sudo systemctl disable unattended-upgrades",
61+
# 2. Aggressively kill any lingering apt, dpkg, or unattended-upgrades processes
62+
"sudo pkill -9 unattended-upgr || true",
63+
"sudo pkill -9 apt || true",
64+
"sudo pkill -9 dpkg || true",
65+
# 3. Remove any remaining lock files
66+
"sudo rm -f /var/lib/apt/lists/lock",
67+
"sudo rm -f /var/cache/apt/archives/lock",
68+
"sudo rm -f /var/lib/dpkg/lock*",
69+
# 4. Repair any interrupted package configurations
70+
"sudo dpkg --configure -a",
71+
"sudo apt-get -o Dpkg::Lock::Timeout=300 -y update",
72+
"sudo apt-get -o Dpkg::Lock::Timeout=300 -y install jq",
6073
"cd JetStream && pip install -e . && cd benchmarks && pip install -r requirements.in",
6174
"pip install torch --index-url https://download.pytorch.org/whl/cpu",
6275
"cd ..",
@@ -70,11 +83,14 @@ def get_config(
7083
# Make the PATH change permanent for subsequent sessions (optional, but good practice)
7184
"echo 'export PATH=\"$HOME/.local/bin:$PATH\"' >> ~/.bashrc",
7285
"source ~/.bashrc",
73-
"uv venv --python 3.12 venv-312 --seed",
74-
"source venv-312/bin/activate",
86+
f"uv venv --python 3.12 {VENV_DIR} --seed --clear",
87+
f"source {VENV_DIR}/bin/activate",
7588
"pip install uv",
7689
"uv pip install maxtext --resolution=lowest",
7790
"install_maxtext_github_deps",
91+
"uv pip install rouge-score",
92+
"sudo apt-get -o Dpkg::Lock::Timeout=300 install -y git-lfs",
93+
"git lfs install",
7894
)
7995

8096
set_up_cmds += setup_maxtext_cmds
@@ -113,15 +129,19 @@ def get_config(
113129

114130
# Let gcs path be directly used, else use maxtext/assets dir
115131
if not model_configs["tokenizer"].startswith("gs://"):
116-
tokenizer_path = f"assets/{model_configs['tokenizer']}"
117-
full_tokenizer_path = f"maxtext/assets/{model_configs['tokenizer']}"
132+
tokenizer_path = (
133+
f"src/maxtext/assets/tokenizers/{model_configs['tokenizer']}"
134+
)
135+
full_tokenizer_path = (
136+
f"maxtext/src/maxtext/assets/tokenizers/{model_configs['tokenizer']}"
137+
)
118138
else:
119139
tokenizer_path = model_configs["tokenizer"]
120140
full_tokenizer_path = model_configs["tokenizer"]
121141

122142
run_model_cmds = (
123143
# Start virtual environment
124-
"source .env/bin/activate",
144+
f"source {VENV_DIR}/bin/activate",
125145
"wget https://huggingface.co/datasets/anon8231489123/ShareGPT_Vicuna_unfiltered/resolve/main/ShareGPT_V3_unfiltered_cleaned_split.json > /dev/null 2>&1",
126146
# Get commit hash of the maxtext and jetstream repos
127147
f"export METADATA_DICT='{json.dumps(additional_metadata_dict)}'",
@@ -131,58 +151,37 @@ def get_config(
131151
'export METADATA_DICT=$(jq -c \'. + { "jetstream_commit_hash": $newVal}\' --arg newVal ${JETSTREAM_COMMIT_HASH} <<<"$METADATA_DICT")',
132152
### Benchmark
133153
"cd maxtext",
134-
# Configure flags
135-
f"export MODEL_NAME={model_configs['model_name']}",
136-
f"export TOKENIZER_PATH={tokenizer_path}",
137-
f"export WEIGHT_DTYPE={model_configs['weight_dtype']}",
138-
f"export SCAN_LAYERS={model_configs['scan_layers']}",
139-
f"export MAX_PREFILL_PREDICT_LENGTH={model_configs['max_prefill_predict_length']}",
140-
f"export MAX_TARGET_LENGTH={model_configs['max_target_length']}",
141-
f"export ATTENTION={model_configs['attention']}",
142-
f"export ICI_FSDP_PARALLELISM={model_configs['ici_fsdp_parallelism']}",
143-
f"export ICI_AUTOREGRESSIVE_PARALLELISM={model_configs['ici_autoregressive_parallelism']}",
144-
f"export ICI_TENSOR_PARALLELISM={model_configs['ici_tensor_parallelism']}",
145-
f"export UNSCANNED_CKPT_PATH={model_configs['checkpoint']}",
146-
"export LOAD_PARAMETERS_PATH=${UNSCANNED_CKPT_PATH}",
147-
f"export QUANTIZATION={model_configs['quantization']}",
148-
f"export QUANTIZE_KVCACHE={model_configs['quantize_kvcache']}",
149-
f"export KV_QUANT_DTYPE={model_configs['kv_quant_dtype']}",
150-
f"export PER_DEVICE_BATCH_SIZE={model_configs['per_device_batch_size']}",
151-
f"export PREFILL_CACHE_AXIS_ORDER={model_configs['prefill_cache_axis_order']}",
152-
f"export AR_CACHE_AXIS_ORDER={model_configs['ar_cache_axis_order']}",
153-
f"export COMPUTE_AXIS_ORDER={model_configs['compute_axis_order']}",
154-
f"export RESHAPE_Q={model_configs['reshape_q']}",
155-
f"export KV_QUANT_AXIS={model_configs['kv_quant_axis']}",
156154
# Start JetStream MaxText server in the background
157-
"""python3 -m MaxText.maxengine_server \
158-
src/maxtext/configs/inference/inference_jetstream.yml \
159-
model_name=${MODEL_NAME} \
160-
tokenizer_path=${TOKENIZER_PATH} \
161-
weight_dtype=${WEIGHT_DTYPE} \
162-
scan_layers=${SCAN_LAYERS} \
163-
max_prefill_predict_length=${MAX_PREFILL_PREDICT_LENGTH} \
164-
max_target_length=${MAX_TARGET_LENGTH} \
165-
attention=${ATTENTION} \
166-
ici_fsdp_parallelism=${ICI_FSDP_PARALLELISM} \
167-
ici_autoregressive_parallelism=${ICI_AUTOREGRESSIVE_PARALLELISM} \
168-
ici_tensor_parallelism=${ICI_TENSOR_PARALLELISM} \
169-
load_parameters_path=${LOAD_PARAMETERS_PATH} \
170-
quantization=${QUANTIZATION} \
171-
quantize_kvcache=${QUANTIZE_KVCACHE} \\"""
172-
+ (
173-
"""kv_quant_dtype=${KV_QUANT_DTYPE} \\"""
174-
if model_configs["kv_quant_dtype"]
175-
else ""
176-
)
177-
+ """per_device_batch_size=${PER_DEVICE_BATCH_SIZE} \
178-
prefill_cache_axis_order=${PREFILL_CACHE_AXIS_ORDER} \
179-
ar_cache_axis_order=${AR_CACHE_AXIS_ORDER} \
180-
compute_axis_order=${COMPUTE_AXIS_ORDER} \
181-
reshape_q=${RESHAPE_Q} \
182-
kv_quant_axis=${KV_QUANT_AXIS} &""",
155+
f"""python3 -m MaxText.maxengine_server \\
156+
src/maxtext/configs/inference/inference_jetstream.yml \\
157+
model_name='{model_configs['model_name']}' \\
158+
tokenizer_path='{tokenizer_path}' \\
159+
weight_dtype='{model_configs['weight_dtype']}' \\
160+
scan_layers='{model_configs['scan_layers']}' \\
161+
max_prefill_predict_length='{model_configs['max_prefill_predict_length']}' \\
162+
max_target_length='{model_configs['max_target_length']}' \\
163+
attention='{model_configs['attention']}' \\
164+
ici_fsdp_parallelism='{model_configs['ici_fsdp_parallelism']}' \\
165+
ici_autoregressive_parallelism='{model_configs['ici_autoregressive_parallelism']}' \\
166+
ici_tensor_parallelism='{model_configs['ici_tensor_parallelism']}' \\
167+
load_parameters_path='{model_configs['checkpoint']}' \\
168+
quantization='"{model_configs.get('quantization') or ""}"' \\
169+
quantize_kvcache='{model_configs['quantize_kvcache']}' \\
170+
kv_quant_dtype='"{model_configs.get('kv_quant_dtype') or ""}"' \\
171+
per_device_batch_size='{model_configs['per_device_batch_size']}' \\
172+
prefill_cache_axis_order='{model_configs['prefill_cache_axis_order']}' \\
173+
ar_cache_axis_order='{model_configs['ar_cache_axis_order']}' \\
174+
compute_axis_order='{model_configs['compute_axis_order']}' \\
175+
reshape_q='{model_configs['reshape_q']}' \\
176+
kv_quant_axis='"{model_configs.get('kv_quant_axis') or ""}"' &""",
183177
"cd ..",
184178
# Give server time to start
185179
f"sleep {model_configs['sleep_time']}",
180+
"cd JetStream",
181+
# Since we change to the Jetstream dir as root, we need to change the ownership of the dir to the user running the script
182+
"sudo chown -R $USER:$USER /home/sa_112632397993248756658/JetStream",
183+
"git lfs pull",
184+
"cd ..",
186185
# Run benchmark, run eval, save benchmark and eval results, and save predictions to /tmp/request-outputs.json
187186
f"""python JetStream/benchmarks/benchmark_serving.py \
188187
--tokenizer {full_tokenizer_path} \
@@ -233,7 +232,6 @@ def get_config(
233232
json_lines=metric_config.JSONLinesConfig("metric_report.jsonl"),
234233
use_runtime_generated_gcs_folder=True,
235234
)
236-
237235
return task.run_queued_resource_test(
238236
task_test_config=job_test_config,
239237
task_gcp_config=job_gcp_config,

dags/inference/maxtext_inference.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -576,7 +576,6 @@
576576
},
577577
}
578578
)
579-
580579
# run_configs = [
581580
# f"{LLAMA2_7B}-{BASE_MODE}-{W_BF16_KV_BF16}",
582581
# f"{LLAMA2_7B}-{BASE_MODE}-{W_INT8_KV_INT8}",
@@ -616,6 +615,7 @@
616615
request_rate=request_rate,
617616
tpu_version=tpu_version,
618617
tpu_cores=tpu_cores,
618+
backup_tpu_resources=True,
619619
)
620620
)
621621
dags.append(jetstream_benchmark_serving_kv_cache_layout)

dags/inference/maxtext_model_config_generator.py

Lines changed: 51 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
"""A helper to generate maxtext model configs."""
1616

17-
from dags.common.vm_resource import TpuVersion, Zone, Project, V5_NETWORKS, V5E_SUBNETWORKS, V5P_SUBNETWORKS, RuntimeVersion, V6E_GCE_NETWORK, V6E_GCE_SUBNETWORK
17+
from dags.common.vm_resource import TpuVersion, Zone, Project, V5_NETWORKS, V5E_SUBNETWORKS, V5P_SUBNETWORKS, RuntimeVersion, V6E_GCE_NETWORK, V6E_GCE_SUBNETWORK, V6E_SUBNETWORKS, CIENET_NETWORKS, CIENET_V5E_SUBNETWORKS
1818
from dags.inference.configs import jetstream_benchmark_serving_gce_config
1919
from dags.multipod.configs.common import SetupMode
2020

@@ -28,6 +28,7 @@ def generate_model_configs(
2828
request_rate,
2929
tpu_version,
3030
tpu_cores,
31+
backup_tpu_resources=False,
3132
):
3233
model_configs = {}
3334
model_configs["model_config_name"] = model_config_name
@@ -99,25 +100,55 @@ def generate_model_configs(
99100

100101
test_name = f"{test_name_prefix}-{test_run_tag}"
101102

102-
if tpu_version == TpuVersion.V5E:
103-
# v5e benchmarks
104-
project_name = Project.TPU_PROD_ENV_AUTOMATED.value
105-
zone = Zone.US_EAST1_C.value
106-
network = V5_NETWORKS
107-
subnetwork = V5E_SUBNETWORKS
108-
runtime_version = RuntimeVersion.V2_ALPHA_TPUV5_LITE.value
109-
elif tpu_version == TpuVersion.V5P:
110-
zone = Zone.US_EAST5_A.value
111-
runtime_version = RuntimeVersion.V2_ALPHA_TPUV5.value
112-
project_name = Project.TPU_PROD_ENV_AUTOMATED.value
113-
network = V5_NETWORKS
114-
subnetwork = V5P_SUBNETWORKS
115-
elif tpu_version == TpuVersion.TRILLIUM:
116-
zone = Zone.US_EAST5_A.value
117-
runtime_version = RuntimeVersion.V2_ALPHA_TPUV6.value
118-
project_name = Project.TPU_PROD_ENV_AUTOMATED.value
119-
network = V6E_GCE_NETWORK
120-
subnetwork = V6E_GCE_SUBNETWORK
103+
# Default configurations for each TPU version
104+
tpu_configs = {
105+
TpuVersion.V5E: {
106+
"runtime_version": RuntimeVersion.V2_ALPHA_TPUV5_LITE.value,
107+
"project_name": Project.TPU_PROD_ENV_AUTOMATED.value,
108+
"zone": Zone.US_EAST1_C.value,
109+
"network": V5_NETWORKS,
110+
"subnetwork": V5E_SUBNETWORKS,
111+
},
112+
TpuVersion.V5P: {
113+
"runtime_version": RuntimeVersion.V2_ALPHA_TPUV5.value,
114+
"project_name": Project.TPU_PROD_ENV_AUTOMATED.value,
115+
"zone": Zone.US_EAST5_A.value,
116+
"network": V5_NETWORKS,
117+
"subnetwork": V5P_SUBNETWORKS,
118+
},
119+
TpuVersion.TRILLIUM: {
120+
"runtime_version": RuntimeVersion.V2_ALPHA_TPUV6.value,
121+
"project_name": Project.TPU_PROD_ENV_AUTOMATED.value,
122+
"zone": Zone.US_EAST5_A.value,
123+
"network": V6E_GCE_NETWORK,
124+
"subnetwork": V6E_GCE_SUBNETWORK,
125+
},
126+
}
127+
128+
# Overrides when backup TPU resources are requested
129+
backup_overrides = {
130+
TpuVersion.V5E: {
131+
"project_name": Project.CIENET_CMCS.value,
132+
"zone": Zone.US_WEST4_A.value,
133+
"network": CIENET_NETWORKS,
134+
"subnetwork": CIENET_V5E_SUBNETWORKS,
135+
},
136+
TpuVersion.TRILLIUM: {
137+
"zone": Zone.SOUTHAMERICA_WEST1_A.value,
138+
"network": V5_NETWORKS,
139+
"subnetwork": V6E_SUBNETWORKS,
140+
},
141+
}
142+
143+
config = tpu_configs[tpu_version]
144+
if backup_tpu_resources and tpu_version in backup_overrides:
145+
config.update(backup_overrides[tpu_version])
146+
147+
runtime_version = config["runtime_version"]
148+
project_name = config["project_name"]
149+
zone = config["zone"]
150+
network = config["network"]
151+
subnetwork = config["subnetwork"]
121152
jetstream_benchmark_serving = (
122153
jetstream_benchmark_serving_gce_config.get_config(
123154
tpu_version=tpu_version,

0 commit comments

Comments
 (0)