Skip to content

Commit 07205e2

Browse files
Merge pull request #2486 from AI-Hypercomputer:aireen/fix_logits_xlml
PiperOrigin-RevId: 818822679
2 parents d6193a5 + f7236a5 commit 07205e2

6 files changed

Lines changed: 23 additions & 17 deletions

File tree

end_to_end/tpu/llama2/7b/test_llama2_7b.sh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ export NEW_CKPT_PATH=${BASE_OUTPUT_DIRECTORY}/${PARAMETER_CHECKPOINT_RUN}/checkp
7373
python3 -m MaxText.decode "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml load_parameters_path=${NEW_CKPT_PATH} run_name=runner_decode_finetuned_${idx} base_output_directory=${BASE_OUTPUT_DIRECTORY} per_device_batch_size=1 model_name='llama2-7b' ici_autoregressive_parallelism=4 max_prefill_predict_length=4 max_target_length=16 prompt="I love to" attention=dot_product scan_layers=false
7474

7575
# We also test whether the forward pass logits match the golden logits for Llama2-7b
76-
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=llama2-7b ici_tensor_parallelism=4 max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 scan_layers=false
76+
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=llama2-7b ici_tensor_parallelism=4 max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 scan_layers=false --rtol=0.1 --atol=0.1
7777

7878
# Converting MaxText orbax checkpoint to HF
7979
JAX_PLATFORMS=cpu python3 -m MaxText.llama_mistral_mixtral_orbax_to_hf "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml base_output_directory=gs://runner-maxtext-logs load_parameters_path=${CONVERTED_CHECKPOINT} run_name=convert_to_hf model_name=llama2-7b hf_model_path=/tmp/hf_llama2

end_to_end/tpu/llama3/70b/2_test_llama3_70b.sh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -64,4 +64,4 @@ python3 -m MaxText.generate_param_only_checkpoint "${MAXTEXT_PKG_DIR:-${MAXTEXT_
6464
python3 -m MaxText.decode "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${BASE_OUTPUT_PATH}/${PARAM_RUN_NAME}/checkpoints/0/items per_device_batch_size=1 run_name=runner_$(date +%Y-%m-%d-%H-%M) max_prefill_predict_length=4 max_target_length=16 dataset_type=synthetic steps=10 async_checkpointing=false scan_layers=false model_name=${MODEL_VARIATION} attention=dot_product prompt="I love to"
6565

6666
# We also test whether the forward pass logits match the golden logits for Llama3-70B
67-
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml base_output_directory=${BASE_OUTPUT_PATH} tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 async_checkpointing=false scan_layers=false
67+
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml base_output_directory=${BASE_OUTPUT_PATH} tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 async_checkpointing=false scan_layers=false --atol=0.1 --rtol=0.1

end_to_end/tpu/llama3/8b/2_test_llama3_8b.sh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -64,4 +64,4 @@ python3 -m MaxText.generate_param_only_checkpoint "${MAXTEXT_PKG_DIR:-${MAXTEXT_
6464
python3 -m MaxText.decode "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${BASE_OUTPUT_PATH}/${PARAM_RUN_NAME}/checkpoints/0/items per_device_batch_size=1 run_name=runner_$(date +%Y-%m-%d-%H-%M) max_prefill_predict_length=4 max_target_length=16 dataset_type=synthetic steps=10 async_checkpointing=false scan_layers=false model_name=${MODEL_VARIATION} attention=dot_product prompt="I love to"
6565

6666
# We also test whether the forward pass logits match the golden logits for Llama3-8B
67-
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml base_output_directory=${BASE_OUTPUT_PATH} tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 async_checkpointing=false scan_layers=false
67+
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml base_output_directory=${BASE_OUTPUT_PATH} tokenizer_path="${MAXTEXT_ASSETS_ROOT:-${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText/assets}}"/tokenizer_llama3.tiktoken load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 dataset_type=synthetic dtype=float32 async_checkpointing=false scan_layers=false --atol=0.1 --rtol=0.1

end_to_end/tpu/llama4/2_test_llama4.sh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,4 +35,4 @@ export MODEL_BUCKET=gs://maxtext-llama/${MODEL_VARIATION}
3535
export UNSCANNED_CKPT_PATH=${MODEL_BUCKET}/${idx}/unscanned/0/items
3636

3737
# Step 2: run logit checking
38-
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml tokenizer_path=${TOKENIZER_PATH} load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test_${MODEL_VARIATION} attention=dot_product per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 scan_layers=false --atol=0.5 --rtol=0.5 async_checkpointing=false sparse_matmul=false weight_dtype=float32 dtype=float32
38+
python3 -m tests.forward_pass_logit_checker "${MAXTEXT_PKG_DIR:-${MAXTEXT_REPO_ROOT:-$PWD}/src/MaxText}/"configs/base.yml tokenizer_path=${TOKENIZER_PATH} load_parameters_path=${UNSCANNED_CKPT_PATH} run_name=forward_pass_test_${MODEL_VARIATION} attention=dot_product per_device_batch_size=1 model_name=${MODEL_VARIATION} max_prefill_predict_length=4 max_target_length=4 scan_layers=false --atol=0.01 --rtol=0.01 async_checkpointing=false sparse_matmul=false weight_dtype=float32 dtype=float32 activations_in_float32=true matmul_precision=float32 float32_logits=true float32_qk_product=true

src/MaxText/layers/decoders.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -485,7 +485,9 @@ def scan_decoder_layers(self, cfg, decoder_layer, length, metadata_axis_name, me
485485
length=length,
486486
metadata_params={nn.PARTITION_NAME: metadata_axis_name},
487487
)
488-
return scan_fn(config=cfg, mesh=mesh, name=metadata_axis_name, quant=self.quant, **kwargs) # pytype: disable=wrong-keyword-args
488+
return scan_fn(
489+
config=cfg, mesh=mesh, name=metadata_axis_name, quant=self.quant, **kwargs
490+
) # pytype: disable=wrong-keyword-args
489491

490492
def get_pipeline_stage_module(self, decoder_blocks):
491493
"""get pipeline stage module"""

tests/forward_pass_logit_checker.py

Lines changed: 16 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -329,16 +329,17 @@ def main(config, test_args): # pylint: disable=W0621
329329
}
330330
all_data_to_save.append(data_to_save)
331331

332-
max_logging.log("\n[test criteria]")
333-
max_logging.log(
334-
f"Checking Numerical Differences between train logits and golden logits against "
335-
f"atol={test_args.rtol} rtol={test_args.atol}."
336-
)
337-
rtol_val = float(test_args.rtol)
338-
atol_val = float(test_args.atol)
339-
assert jax.numpy.allclose(
340-
train_logits_slice, golden_logits_slice, rtol=rtol_val, atol=atol_val, equal_nan=False
341-
), f"Logits do not match closely enough. Required rtol={test_args.rtol}, atol={test_args.atol}."
332+
if test_args.atol is not None:
333+
max_logging.log("\n[test criteria]")
334+
max_logging.log(
335+
f"Checking Numerical Differences between train logits and golden logits against "
336+
f"atol={test_args.rtol} rtol={test_args.atol}."
337+
)
338+
rtol_val = float(test_args.rtol)
339+
atol_val = float(test_args.atol)
340+
assert jax.numpy.allclose(
341+
train_logits_slice, golden_logits_slice, rtol=rtol_val, atol=atol_val, equal_nan=False
342+
), f"Logits do not match closely enough. Required rtol={test_args.rtol}, atol={test_args.atol}."
342343

343344
if test_args.max_kl_div is not None:
344345
max_logging.log(
@@ -451,8 +452,8 @@ def main(config, test_args): # pylint: disable=W0621
451452
os.environ["TF_CPP_MIN_LOG_LEVEL"] = "0"
452453

453454
parser = argparse.ArgumentParser()
454-
parser.add_argument("--atol", type=float, required=False, default=0.1)
455-
parser.add_argument("--rtol", type=float, required=False, default=0.1)
455+
parser.add_argument("--atol", type=float, required=False, default=None)
456+
parser.add_argument("--rtol", type=float, required=False, default=1e-05) # default from jnp.allclose
456457
parser.add_argument("--token_size", type=int, required=False)
457458
parser.add_argument("--max_kl_div", type=float, required=False, default=None)
458459
parser.add_argument("--golden_logits_path", type=str, required=False, default="")
@@ -479,6 +480,9 @@ def main(config, test_args): # pylint: disable=W0621
479480
model_args = [s for s in model_args if not s.startswith(arg)]
480481

481482
cfg = pyconfig.initialize(model_args)
483+
assert (
484+
test_args.atol is not None or test_args.max_kl_div is not None
485+
), "At least one of --atol or --max_kl_div must be specified to define the test criteria."
482486
if cfg.use_multimodal:
483487
assert not test_args.run_hf_model, (
484488
"Multimodal does not support running hf model on-the-fly, please generate hf golden logits "

0 commit comments

Comments
 (0)