|
1 | 1 | #!/bin/bash |
2 | 2 |
|
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. |
4 | 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 |
| 5 | +set -ex |
11 | 6 |
|
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 |
| 7 | +# Activate virtual environment if available |
| 8 | +if [ -f "/home/jackyf_google_com/maxtext/.venv/bin/activate" ]; then |
| 9 | + source /home/jackyf_google_com/maxtext/.venv/bin/activate |
| 10 | +elif [ -f "${PWD}/.venv/bin/activate" ]; then |
| 11 | + source ${PWD}/.venv/bin/activate |
| 12 | +fi |
14 | 13 |
|
15 | | -set -ex |
| 14 | +# Ensure src is in PYTHONPATH so maxtext module is found |
| 15 | +export PYTHONPATH=${PWD}/src:${PYTHONPATH} |
16 | 16 |
|
| 17 | +run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)} |
17 | 18 | export MODEL_NAME='gpt-oss-20b' |
18 | 19 | export TOKENIZER_PATH='openai/gpt-oss-20b' |
19 | 20 |
|
20 | 21 | 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" |
| 22 | + export BASE_OUTPUT_PATH=gs://runner-maxtext-logs/${MODEL_NAME} |
24 | 23 | fi |
25 | 24 | BASE_OUTPUT_PATH=${BASE_OUTPUT_PATH%/} |
26 | | -echo using BASE_OUTPUT_PATH = ${BASE_OUTPUT_PATH} |
27 | 25 |
|
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 |
| 26 | +export SCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/scanned/${run_id}/0/items |
| 27 | +export UNSCANNED_CKPT_PATH=${BASE_OUTPUT_PATH}/unscanned/${run_id}/0/items |
| 28 | + |
| 29 | +# Dynamically configure for CPU or TPU environment |
| 30 | +if python3 -c "import jax; print(any(d.platform == 'tpu' for d in jax.devices()))" 2>/dev/null | grep -q "True"; then |
| 31 | + echo "TPU detected. Running with Megablox and Flash/Pallas kernels." |
| 32 | + export SPARSE_MATMUL="True" |
| 33 | + export MEGABLOX="True" |
| 34 | + export PRETRAIN_ATTENTION="flash" |
| 35 | +else |
| 36 | + echo "No TPU detected. Running on CPU with dense layers and dot_product attention." |
| 37 | + export SPARSE_MATMUL="False" |
| 38 | + export MEGABLOX="False" |
| 39 | + export PRETRAIN_ATTENTION="dot_product" |
| 40 | + export JAX_PLATFORMS=cpu |
| 41 | + export XLA_FLAGS="--xla_force_host_platform_device_count=4" |
39 | 42 | fi |
40 | 43 |
|
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 |
70 | | - |
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 |
| 44 | +# 1. Run Pre-training using synthetic dataset with Megablocks |
| 45 | +python3 -m maxtext.trainers.pre_train.train \ |
| 46 | + "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"/base.yml \ |
| 47 | + base_output_directory=${BASE_OUTPUT_PATH}/train \ |
| 48 | + run_name=${run_id} \ |
| 49 | + model_name=${MODEL_NAME} \ |
| 50 | + tokenizer_type=huggingface \ |
| 51 | + tokenizer_path=${TOKENIZER_PATH} \ |
| 52 | + dataset_type=synthetic \ |
| 53 | + enable_checkpointing=true \ |
| 54 | + async_checkpointing=false \ |
| 55 | + load_parameters_path=${SCANNED_CKPT_PATH} \ |
| 56 | + attention=${PRETRAIN_ATTENTION} \ |
| 57 | + sparse_matmul=${SPARSE_MATMUL} \ |
| 58 | + megablox=${MEGABLOX} \ |
| 59 | + dtype=bfloat16 \ |
| 60 | + weight_dtype=bfloat16 \ |
| 61 | + per_device_batch_size=4 \ |
| 62 | + steps=5 \ |
| 63 | + max_target_length=1024 \ |
| 64 | + ici_fsdp_parallelism=4 \ |
| 65 | + gcs_metrics=true |
| 66 | + |
| 67 | +# 2. Run Verification Decoding from the converted checkpoint |
| 68 | +python3 -m maxtext.inference.decode \ |
| 69 | + "${MAXTEXT_CONFIGS_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/maxtext/configs}"/base.yml \ |
| 70 | + base_output_directory=${BASE_OUTPUT_PATH} \ |
| 71 | + run_name=decode \ |
| 72 | + model_name=${MODEL_NAME} \ |
| 73 | + tokenizer_type=huggingface \ |
| 74 | + tokenizer_path=${TOKENIZER_PATH} \ |
| 75 | + load_parameters_path=${UNSCANNED_CKPT_PATH} \ |
| 76 | + scan_layers=False \ |
| 77 | + pure_nnx=False \ |
| 78 | + pure_nnx_decoder=False \ |
| 79 | + attention=dot_product \ |
| 80 | + sparse_matmul=${SPARSE_MATMUL} \ |
| 81 | + megablox=${MEGABLOX} \ |
| 82 | + dtype=bfloat16 \ |
| 83 | + weight_dtype=bfloat16 \ |
| 84 | + per_device_batch_size=1 \ |
| 85 | + max_prefill_predict_length=64 \ |
| 86 | + max_target_length=128 \ |
| 87 | + prompt="I love to" \ |
| 88 | + ici_fsdp_parallelism=1 \ |
| 89 | + ici_tensor_parallelism=4 |
0 commit comments