Skip to content

Commit b1b5b0b

Browse files
committed
feat(test): add modular E2E TPU testing pipeline scripts for GPTOSS-20B
1 parent 0faa9be commit b1b5b0b

5 files changed

Lines changed: 253 additions & 61 deletions

File tree

Lines changed: 46 additions & 61 deletions
Original file line numberDiff line numberDiff line change
@@ -1,73 +1,58 @@
11
#!/bin/bash
22

3-
# This file is documentation for how to get started with gpt-oss-20b on v5p-8.
4-
5-
# The flow of this file is as follows:
6-
# 1. Convert the HuggingFace checkpoint (bf16) to MaxText-compatible checkpoint (bf16):
7-
# Scanned format is better for training; unscanned format is better for decoding.
8-
# 2. Run logit check, pre-training, fine-tuning, and decoding.
9-
10-
# Example Usage: export HF_TOKEN=<huggingface_access_token>; export BASE_OUTPUT_PATH=<GCS_bucket_path>; bash test_gpt_oss.sh
11-
12-
# The golden logit can be generated by:
13-
# python3 -m tests.assets.logits_generation.generate_hf_golden_logits --model-id=openai/gpt-oss-20b --output-path=golden_data_gpt-oss-20b.jsonl --prompts='I love to;Today is a;What is the' --hf-model-path=$local_bf16_path
3+
# Validates the GPTOSS-20B pre-training pipeline starting from converted MaxText checkpoint.
144

155
set -ex
166

7+
run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)}
178
export MODEL_NAME='gpt-oss-20b'
189
export TOKENIZER_PATH='openai/gpt-oss-20b'
1910

2011
if [ -z "${BASE_OUTPUT_PATH}" ]; then
21-
# Non-Googlers please remember to point `BASE_OUTPUT_PATH` to GCS buckets that you own, this script uses internal buckets for testing.
22-
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/$(date +%Y-%m-%d-%H-%M)
23-
echo "BASE_OUTPUT_PATH is not set"
12+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
2413
fi
2514
BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/}
26-
echo using BASE_OUTPUT_PATH = ${BASE_OUTPUT_PATH}
27-
28-
# Installing torch for checkpoint conversion and forward_pass_logit_checker.py
29-
python3 -m pip install torch --index-url https://download.pytorch.org/whl/cpu
30-
31-
# Step 1: Checkpoint conversion
32-
# Assume HF checkpoints are uploaded to GCS bucket at CKPT_BUCKET
33-
# Non-Googlers please remember to point `CKPT_BUCKET` to GCS buckets that you own
34-
# Copying the HF checkpoint into a local directory `/tmp` -- you are free to use a different directory
35-
if [ -z "${CKPT_DISK_LOCATION}" ]; then
36-
export CKPT_BUCKET=gs://maxtext-model-checkpoints/gpt-oss-20b/hf-bf16
37-
gcloud storage cp -r ${CKPT_BUCKET} /tmp
38-
export CKPT_DISK_LOCATION=/tmp/hf-bf16
39-
fi
40-
41-
# 1.1 Convert checkpoint to `scanned` format, more suitable for training
42-
JAX_PLATFORMS=cpu python3 -m maxtext.checkpoint_conversion.standalone_scripts.convert_gpt_oss_ckpt --base-model-path ${CKPT_DISK_LOCATION} --maxtext-model-path ${BASE_OUTPUT_PATH}/scanned --model-size ${MODEL_NAME}
43-
44-
# 1.2 Convert checkpoint to `unscanned` format, more suitable for decoding
45-
JAX_PLATFORMS=cpu python3 -m maxtext.checkpoint_conversion.standalone_scripts.convert_gpt_oss_unscanned_ckpt --base-model-path ${CKPT_DISK_LOCATION} --maxtext-model-path ${BASE_OUTPUT_PATH}/unscanned --model-size ${MODEL_NAME}
46-
47-
# Step 2:
48-
# We define the checkpoint paths. This way it is easier to use these paths in the `train.py` and `decode.py` commands
49-
export SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/0/items
50-
export UNSCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/unscanned/0/items
51-
# Non-Googlers please remember to point `DATASET_PATH` to the GCS bucket where you have your training data
52-
export DATASET_PATH=gs://maxtext-dataset
53-
54-
export LIBTPU_INIT_ARGS='--xla_tpu_scoped_vmem_limit_kib=81920'
55-
56-
# Test whether the forward pass logits match the golden logits
57-
# default golden_logits_path=/deps/tests/assets/golden_logits/golden_data_{MODEL_NAME}.jsonl, copied from gs://maxtext-test-assets/golden_data_${MODEL_NAME}.jsonl
58-
python3 -m tests.utils.forward_pass_logit_checker "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"//base.yml base_output_directory=${BASE_OUTPUT_PATH} run_name=forward_logits_check model_name=${MODEL_NAME} load_parameters_path=${UNSCANNED_CKPT_PATH} scan_layers=false attention=dot_product sparse_matmul=True megablox=True per_device_batch_size=1 max_target_length=4 max_prefill_predict_length=4 dtype=float32 --atol=0.1 --rtol=0.1 --max_kl_div=3e-4
59-
60-
# Run pre-training - megablox implementation
61-
python3 -m maxtext.trainers.pre_train.train "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"//base.yml base_output_directory=${BASE_OUTPUT_PATH} run_name=megablox_pre_training model_name=${MODEL_NAME} tokenizer_type=huggingface tokenizer_path=${TOKENIZER_PATH} dataset_type=synthetic enable_checkpointing=false attention=flash sparse_matmul=True megablox=True dtype=bfloat16 weight_dtype=bfloat16 per_device_batch_size=4 steps=5 max_target_length=1024 ici_fsdp_parallelism=4 gcs_metrics=true
62-
63-
# Run fine-tuning - megablox implementation
64-
# TODO: remove `abort_on_nan_loss=false` after b/497864549
65-
python3 -m maxtext.trainers.pre_train.train "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"//base.yml base_output_directory=${BASE_OUTPUT_PATH} run_name=megablox_fine_tuning model_name=${MODEL_NAME} tokenizer_type=huggingface tokenizer_path=${TOKENIZER_PATH} dataset_path=${DATASET_PATH} enable_checkpointing=true async_checkpointing=false load_parameters_path=${SCANNED_CKPT_PATH} scan_layers=True attention=flash sparse_matmul=True megablox=True dtype=bfloat16 weight_dtype=bfloat16 per_device_batch_size=4 steps=5 max_target_length=1024 ici_fsdp_parallelism=1 ici_expert_parallelism=4 gcs_metrics=true abort_on_nan_loss=false
66-
67-
# Run supervised fine-tuning - megablox implementation
68-
# TODO: remove `abort_on_nan_loss=false` after b/497864549
69-
python3 -m maxtext.trainers.post_train.sft.train_sft_native "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs/post_train}"//sft.yml base_output_directory=${BASE_OUTPUT_PATH} run_name=megablox_supervised_fine_tuning model_name=${MODEL_NAME} tokenizer_type=huggingface tokenizer_path=${TOKENIZER_PATH} dataset_type=hf enable_checkpointing=true async_checkpointing=false load_parameters_path=${SCANNED_CKPT_PATH} scan_layers=True attention=flash sparse_matmul=True megablox=True dtype=bfloat16 weight_dtype=bfloat16 per_device_batch_size=4 steps=5 max_target_length=1024 ici_fsdp_parallelism=1 ici_expert_parallelism=4 gcs_metrics=true abort_on_nan_loss=false
7015

71-
# Run decoding - megablox implementation
72-
# Note decode requires the access token for huggingface tokenizer even if the model is not gated
73-
python3 -m maxtext.inference.decode "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"//base.yml base_output_directory=${BASE_OUTPUT_PATH} run_name=decode model_name=${MODEL_NAME} tokenizer_type=huggingface tokenizer_path=${TOKENIZER_PATH} hf_access_token=${HF_TOKEN} load_parameters_path=${UNSCANNED_CKPT_PATH} scan_layers=False attention=dot_product sparse_matmul=True megablox=True dtype=bfloat16 weight_dtype=bfloat16 per_device_batch_size=1 max_prefill_predict_length=64 max_target_length=128 prompt="I love to" ici_fsdp_parallelism=1 ici_tensor_parallelism=4
16+
export SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/${run_id}/0/items
17+
export UNSCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/unscanned/${run_id}/0/items
18+
19+
export SPARSE_MATMUL="True"
20+
export MEGABLOX="True"
21+
export PRETRAIN_ATTENTION="flash"
22+
23+
# 1. Run Pre-training using synthetic dataset with Megablocks
24+
python3 -m maxtext.trainers.pre_train.train \
25+
"${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"/base.yml \
26+
base_output_directory=${BASE_OUTPUT_PATH}/train \
27+
run_name=${run_id} \
28+
model_name=${MODEL_NAME} \
29+
tokenizer_type=huggingface \
30+
tokenizer_path=${TOKENIZER_PATH} \
31+
dataset_type=synthetic \
32+
enable_checkpointing=true \
33+
async_checkpointing=false \
34+
load_parameters_path=${SCANNED_CKPT_PATH} \
35+
attention=${PRETRAIN_ATTENTION} \
36+
sparse_matmul=${SPARSE_MATMUL} \
37+
megablox=${MEGABLOX} \
38+
dtype=bfloat16 \
39+
weight_dtype=bfloat16 \
40+
per_device_batch_size=4 \
41+
steps=5 \
42+
max_target_length=1024 \
43+
ici_fsdp_parallelism=4 \
44+
gcs_metrics=true
45+
46+
# 2. Run Verification Decoding from the converted checkpoint
47+
python3 -m maxtext.inference.decode \
48+
base_output_directory=${BASE_OUTPUT_PATH} \
49+
run_name=decode \
50+
model_name=${MODEL_NAME} \
51+
tokenizer_path=${TOKENIZER_PATH} \
52+
load_parameters_path=${BASE_OUTPUT_PATH}/train/${run_id}/checkpoints/4/items \
53+
scan_layers=True \
54+
attention=dot_product \
55+
sparse_matmul=${SPARSE_MATMUL} \
56+
megablox=${MEGABLOX} \
57+
prompt="I love to" \
58+
ici_tensor_parallelism=4
Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,55 @@
1+
#!/bin/bash
2+
3+
# Validates the GPTOSS-20B Reinforcement Learning (RL) pipeline using GRPO.
4+
5+
set -ex
6+
7+
run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)}
8+
export MODEL_NAME='gpt-oss-20b'
9+
export TOKENIZER_PATH='openai/gpt-oss-20b'
10+
11+
if [ -z "${BASE_OUTPUT_PATH}" ]; then
12+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
13+
fi
14+
BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/}
15+
16+
export SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/${run_id}/0/items
17+
18+
export SPARSE_MATMUL="True"
19+
export MEGABLOX="True"
20+
export ATTENTION="flash"
21+
export VLLM_ADDITIONAL_CONFIG='{"maxtext_config": {"model_name": "gpt-oss-20b", "log_config": "false"}}'
22+
23+
# 1. Run GRPO Reinforcement Learning
24+
python3 -m maxtext.trainers.post_train.rl.train_rl \
25+
base_output_directory=${BASE_OUTPUT_PATH}/rl \
26+
load_parameters_path=${SCANNED_CKPT_PATH} \
27+
run_name=${run_id} \
28+
rl.loss_algo='grpo' \
29+
scan_layers=true \
30+
num_batches=5 \
31+
batch_size=1 \
32+
num_test_batches=5 \
33+
model_name=${MODEL_NAME} \
34+
checkpoint_storage_use_zarr3=False \
35+
checkpoint_storage_use_ocdbt=False \
36+
rollout_tensor_parallelism=1 \
37+
attention=${ATTENTION} \
38+
sparse_matmul=${SPARSE_MATMUL} \
39+
megablox=${MEGABLOX} \
40+
vllm_hf_overrides='{architectures: ["MaxTextForCausalLM"]}' \
41+
vllm_additional_config="${VLLM_ADDITIONAL_CONFIG}"
42+
43+
# 2. Run Verification Decoding on the newly produced actor checkpoint
44+
python3 -m maxtext.inference.decode \
45+
base_output_directory=${BASE_OUTPUT_PATH} \
46+
run_name=decode_rl \
47+
model_name=${MODEL_NAME} \
48+
tokenizer_path=${TOKENIZER_PATH} \
49+
load_parameters_path=${BASE_OUTPUT_PATH}/rl/${run_id}/checkpoints/actor/4/items \
50+
scan_layers=True \
51+
attention=dot_product \
52+
sparse_matmul=${SPARSE_MATMUL} \
53+
megablox=${MEGABLOX} \
54+
prompt="I love to" \
55+
ici_tensor_parallelism=4
Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
#!/bin/bash
2+
3+
# Validates the GPTOSS-20B Supervised Fine-Tuning (SFT) pipeline.
4+
5+
set -ex
6+
7+
run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)}
8+
export MODEL_NAME='gpt-oss-20b'
9+
export TOKENIZER_PATH='openai/gpt-oss-20b'
10+
11+
if [ -z "${BASE_OUTPUT_PATH}" ]; then
12+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
13+
fi
14+
BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/}
15+
16+
export SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/${run_id}/0/items
17+
18+
export SPARSE_MATMUL="True"
19+
export MEGABLOX="True"
20+
export SFT_ATTENTION="flash"
21+
22+
# 1. Run Supervised Fine-Tuning
23+
python3 -m maxtext.trainers.post_train.sft.train_sft_native \
24+
"${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs/post_train}"/sft.yml \
25+
base_output_directory=${BASE_OUTPUT_PATH}/sft \
26+
run_name=${run_id} \
27+
model_name=${MODEL_NAME} \
28+
tokenizer_type=huggingface \
29+
tokenizer_path=${TOKENIZER_PATH} \
30+
dataset_type=hf \
31+
enable_checkpointing=true \
32+
async_checkpointing=false \
33+
load_parameters_path=${SCANNED_CKPT_PATH} \
34+
scan_layers=True \
35+
attention=${SFT_ATTENTION} \
36+
sparse_matmul=${SPARSE_MATMUL} \
37+
megablox=${MEGABLOX} \
38+
dtype=bfloat16 \
39+
weight_dtype=bfloat16 \
40+
per_device_batch_size=4 \
41+
steps=5 \
42+
max_target_length=1024 \
43+
ici_fsdp_parallelism=1 \
44+
ici_expert_parallelism=4 \
45+
gcs_metrics=true \
46+
abort_on_nan_loss=false
47+
48+
# 2. Run Decoding on the newly produced SFT checkpoint
49+
python3 -m maxtext.inference.decode \
50+
base_output_directory=${BASE_OUTPUT_PATH} \
51+
run_name=decode_sft \
52+
model_name=${MODEL_NAME} \
53+
tokenizer_path=${TOKENIZER_PATH} \
54+
load_parameters_path=${BASE_OUTPUT_PATH}/sft/${run_id}/checkpoints/4/items \
55+
scan_layers=True \
56+
attention=dot_product \
57+
sparse_matmul=${SPARSE_MATMUL} \
58+
megablox=${MEGABLOX} \
59+
prompt="I love to" \
60+
ici_tensor_parallelism=4
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
#!/bin/bash
2+
3+
# Converts a MaxText checkpoint to a Hugging Face model checkpoint for GPTOSS-20B.
4+
5+
set -ex
6+
7+
run_id=$1
8+
CKPT_PATH=$2
9+
SCAN_LAYERS=${3:-false}
10+
11+
export MODEL_NAME='gpt-oss-20b'
12+
13+
if [ -z "${BASE_OUTPUT_PATH}" ]; then
14+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
15+
fi
16+
BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/}
17+
18+
if [ "${SCAN_LAYERS,,}" = "true" ]; then
19+
scan_status="scanned"
20+
else
21+
scan_status="unscanned"
22+
fi
23+
24+
python3 -m maxtext.checkpoint_conversion.to_huggingface \
25+
model_name=${MODEL_NAME} \
26+
tokenizer_type="huggingface" \
27+
load_parameters_path=${CKPT_PATH} \
28+
base_output_directory=${BASE_OUTPUT_PATH}/to_huggingface/${scan_status}/${run_id} \
29+
use_multimodal=false \
30+
scan_layers=$SCAN_LAYERS
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
#!/bin/bash
2+
3+
# Converts GPTOSS-20B HuggingFace checkpoint to MaxText format and validates logit correctness.
4+
5+
set -ex
6+
7+
export PYTHONPATH=src
8+
9+
run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)}
10+
export MODEL_NAME='gpt-oss-20b'
11+
export TOKENIZER_PATH='openai/gpt-oss-20b'
12+
13+
if [ -z "${BASE_OUTPUT_PATH}" ]; then
14+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
15+
fi
16+
BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/}
17+
echo "Using BASE_OUTPUT_PATH = ${BASE_OUTPUT_PATH}"
18+
19+
if [ -z "${CKPT_DISK_LOCATION}" ]; then
20+
export CKPT_BUCKET=gs://maxtext-model-checkpoints/gpt-oss-20b/hf-bf16
21+
gcloud storage cp -r ${CKPT_BUCKET} /tmp
22+
export CKPT_DISK_LOCATION=/tmp/hf-bf16
23+
fi
24+
25+
# 1. Convert to scanned checkpoint (for training)
26+
JAX_PLATFORMS=cpu python3 -m maxtext.checkpoint_conversion.standalone_scripts.convert_gpt_oss_ckpt \
27+
--base-model-path ${CKPT_DISK_LOCATION} \
28+
--maxtext-model-path ${BASE_OUTPUT_PATH}/scanned/${run_id} \
29+
--model-size ${MODEL_NAME} \
30+
--pure-nnx True
31+
32+
SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/${run_id}/0/items
33+
echo "Scanned checkpoint path: ${SCANNED_CKPT_PATH}"
34+
35+
# 2. Convert to unscanned checkpoint (for inference)
36+
JAX_PLATFORMS=cpu python3 -m maxtext.checkpoint_conversion.standalone_scripts.convert_gpt_oss_unscanned_ckpt \
37+
--base-model-path ${CKPT_DISK_LOCATION} \
38+
--maxtext-model-path ${BASE_OUTPUT_PATH}/unscanned/${run_id} \
39+
--model-size ${MODEL_NAME} \
40+
--pure-nnx True
41+
42+
UNSCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/unscanned/${run_id}/0/items
43+
echo "Unscanned checkpoint path: ${UNSCANNED_CKPT_PATH}"
44+
45+
# 3. Logit correctness check
46+
if [ ! -f /tmp/golden_data_gpt-oss-20b.jsonl ]; then
47+
gcloud storage cp gs://maxtext-test-assets/golden_data_gpt-oss-20b.jsonl /tmp/golden_data_gpt-oss-20b.jsonl
48+
fi
49+
50+
SPARSE_MATMUL="True"
51+
MEGABLOX="True"
52+
53+
python3 -m tests.utils.forward_pass_logit_checker \
54+
base_output_directory=${BASE_OUTPUT_PATH} \
55+
model_name=${MODEL_NAME} \
56+
load_parameters_path=${UNSCANNED_CKPT_PATH} \
57+
scan_layers=false \
58+
attention=dot_product \
59+
sparse_matmul=${SPARSE_MATMUL} \
60+
megablox=${MEGABLOX} \
61+
--golden_logits_path=/tmp/golden_data_gpt-oss-20b.jsonl \
62+
--max_kl_div=0.01

0 commit comments

Comments
 (0)