Skip to content

Commit cbe2520

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

5 files changed

Lines changed: 269 additions & 59 deletions

File tree

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

3-
# This file is documentation for how to get started with gpt-oss-20b on v5p-8.
3+
# Validates the GPTOSS-20B pre-training pipeline starting from converted MaxText checkpoint.
44

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.
5+
set -ex
96

10-
# Example Usage: export HF_TOKEN=<huggingface_access_token>; export BASE_OUTPUT_PATH=<GCS_bucket_path>; bash test_gpt_oss.sh
117

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
148

15-
set -ex
169

10+
11+
run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)}
1712
export MODEL_NAME='gpt-oss-20b'
1813
export TOKENIZER_PATH='openai/gpt-oss-20b'
1914

2015
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"
16+
export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME}
2417
fi
2518
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
7019

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

0 commit comments

Comments
 (0)