2525PROJECT_NAME = Project .CLOUD_ML_AUTO_SOLUTIONS .value
2626RUNTIME_IMAGE = RuntimeVersion .TPU_UBUNTU2204_BASE .value
2727GCS_SUBFOLDER_PREFIX = test_owner .Team .INFERENCE .value
28+ VENV_DIR = "venv-312"
2829
2930
3031def get_config (
@@ -70,11 +71,14 @@ def get_config(
7071 # Make the PATH change permanent for subsequent sessions (optional, but good practice)
7172 "echo 'export PATH=\" $HOME/.local/bin:$PATH\" ' >> ~/.bashrc" ,
7273 "source ~/.bashrc" ,
73- "uv venv --python 3.12 venv-312 --seed" ,
74- "source venv-312 /bin/activate" ,
74+ f "uv venv --python 3.12 { VENV_DIR } --seed --clear " ,
75+ f "source { VENV_DIR } /bin/activate" ,
7576 "pip install uv" ,
7677 "uv pip install maxtext --resolution=lowest" ,
7778 "install_maxtext_github_deps" ,
79+ "uv pip install rouge-score" ,
80+ "sudo apt-get -o Dpkg::Lock::Timeout=300 install -y git-lfs" ,
81+ "git lfs install" ,
7882 )
7983
8084 set_up_cmds += setup_maxtext_cmds
@@ -113,15 +117,19 @@ def get_config(
113117
114118 # Let gcs path be directly used, else use maxtext/assets dir
115119 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' ]} "
120+ tokenizer_path = (
121+ f"src/maxtext/assets/tokenizers/{ model_configs ['tokenizer' ]} "
122+ )
123+ full_tokenizer_path = (
124+ f"maxtext/src/maxtext/assets/tokenizers/{ model_configs ['tokenizer' ]} "
125+ )
118126 else :
119127 tokenizer_path = model_configs ["tokenizer" ]
120128 full_tokenizer_path = model_configs ["tokenizer" ]
121129
122130 run_model_cmds = (
123131 # Start virtual environment
124- "source .env /bin/activate" ,
132+ f "source { VENV_DIR } /bin/activate" ,
125133 "wget https://huggingface.co/datasets/anon8231489123/ShareGPT_Vicuna_unfiltered/resolve/main/ShareGPT_V3_unfiltered_cleaned_split.json > /dev/null 2>&1" ,
126134 # Get commit hash of the maxtext and jetstream repos
127135 f"export METADATA_DICT='{ json .dumps (additional_metadata_dict )} '" ,
@@ -131,58 +139,37 @@ def get_config(
131139 'export METADATA_DICT=$(jq -c \' . + { "jetstream_commit_hash": $newVal}\' --arg newVal ${JETSTREAM_COMMIT_HASH} <<<"$METADATA_DICT")' ,
132140 ### Benchmark
133141 "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' ]} " ,
156142 # 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} &""" ,
143+ f"""python3 -m MaxText.maxengine_server \\
144+ src/maxtext/configs/inference/inference_jetstream.yml \\
145+ model_name='{ model_configs ['model_name' ]} ' \\
146+ tokenizer_path='{ tokenizer_path } ' \\
147+ weight_dtype='{ model_configs ['weight_dtype' ]} ' \\
148+ scan_layers='{ model_configs ['scan_layers' ]} ' \\
149+ max_prefill_predict_length='{ model_configs ['max_prefill_predict_length' ]} ' \\
150+ max_target_length='{ model_configs ['max_target_length' ]} ' \\
151+ attention='{ model_configs ['attention' ]} ' \\
152+ ici_fsdp_parallelism='{ model_configs ['ici_fsdp_parallelism' ]} ' \\
153+ ici_autoregressive_parallelism='{ model_configs ['ici_autoregressive_parallelism' ]} ' \\
154+ ici_tensor_parallelism='{ model_configs ['ici_tensor_parallelism' ]} ' \\
155+ load_parameters_path='{ model_configs ['checkpoint' ]} ' \\
156+ quantization='"{ model_configs .get ('quantization' ) or "" } "' \\
157+ quantize_kvcache='{ model_configs ['quantize_kvcache' ]} ' \\
158+ kv_quant_dtype='"{ model_configs .get ('kv_quant_dtype' ) or "" } "' \\
159+ per_device_batch_size='{ model_configs ['per_device_batch_size' ]} ' \\
160+ prefill_cache_axis_order='{ model_configs ['prefill_cache_axis_order' ]} ' \\
161+ ar_cache_axis_order='{ model_configs ['ar_cache_axis_order' ]} ' \\
162+ compute_axis_order='{ model_configs ['compute_axis_order' ]} ' \\
163+ reshape_q='{ model_configs ['reshape_q' ]} ' \\
164+ kv_quant_axis='"{ model_configs .get ('kv_quant_axis' ) or "" } "' &""" ,
183165 "cd .." ,
184166 # Give server time to start
185167 f"sleep { model_configs ['sleep_time' ]} " ,
168+ "cd JetStream" ,
169+ # Since we change to the Jetstream dir as root, we need to change the ownership of the dir to the user running the script
170+ "sudo chown -R $USER:$USER $HOME/JetStream" ,
171+ "git lfs pull" ,
172+ "cd .." ,
186173 # Run benchmark, run eval, save benchmark and eval results, and save predictions to /tmp/request-outputs.json
187174 f"""python JetStream/benchmarks/benchmark_serving.py \
188175 --tokenizer { full_tokenizer_path } \
@@ -233,7 +220,6 @@ def get_config(
233220 json_lines = metric_config .JSONLinesConfig ("metric_report.jsonl" ),
234221 use_runtime_generated_gcs_folder = True ,
235222 )
236-
237223 return task .run_queued_resource_test (
238224 task_test_config = job_test_config ,
239225 task_gcp_config = job_gcp_config ,
0 commit comments