|
1 | 1 | #!/bin/bash |
2 | 2 |
|
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. |
14 | 4 |
|
15 | 5 | set -ex |
16 | 6 |
|
| 7 | +run_id=${1:-$(date +%Y-%m-%d-%H-%M-%S)} |
17 | 8 | export MODEL_NAME='gpt-oss-20b' |
18 | 9 | export TOKENIZER_PATH='openai/gpt-oss-20b' |
19 | 10 |
|
20 | 11 | 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} |
24 | 13 | fi |
25 | 14 | 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 |
70 | 15 |
|
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 |
0 commit comments