NVIDIA / NVIDIA/Megatron-LM

Cannot save optimizer when swiglu and FSDP

Open
#5,265 3 comments 0 reactions 1 assignee Claimed by @asolergi-nv View on GitHub
bug community-request module: megatron-fsdp
Dominant language
Python
Stars
17.9k
Forks
4.5k
Avg merge
4d 6h
Merged PRs (30d)
271

Description

**Describe the bug**

Optimizer saving fails when using swiglu and FSDP

@NVIDIA/mcore-oncall

**Steps/Code to reproduce bug**

Run `examples/megatron_fsdp/train_llama3_8b_fsdp_h100_fp8.sh `

I have modified setup to fail fast
```
SEQ_LENGTH=16
--train-samples 1000
--save-interval 1
```

Log
```
[rank3]: Traceback (most recent call last):
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/pretrain_gpt_2.py", line 408, in
[rank3]: pretrain(full_config,
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 1385, in pretrain
[rank3]: iteration, num_floating_point_operations_so_far = train(
[rank3]: ^^^^^^
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 3759, in train
[rank3]: should_exit = checkpoint_and_decide_exit(
[rank3]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2989, in checkpoint_and_decide_exit
[rank3]: save_checkpoint_and_time(
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2816, in save_checkpoint_and_time
[rank3]: save_checkpoint(
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 693, in save_checkpoint
[rank3]: state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0])
[rank3]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 1076, in preprocess_fsdp_dtensor_state_dict
[rank3]: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict(
[rank3]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 378, in handle_swiglu_in_state_dict
[rank3]: weight_w, weight_v = split_swiglu_linear_fc1(
[rank3]: ^^^^^^^^^^^^^^^^^^^^^^^^
[rank3]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 278, in split_swiglu_linear_fc1
[rank3]: assert data.shape[swiglu_shard_axis] % 2 == 0, (
[rank3]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank3]: AssertionError: SWiGLU weights must have an even size along the shard axis 0, got 13307
[rank1]: Traceback (most recent call last):
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/pretrain_gpt_2.py", line 408, in
[rank1]: pretrain(full_config,
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 1385, in pretrain
[rank1]: iteration, num_floating_point_operations_so_far = train(
[rank1]: ^^^^^^
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 3759, in train
[rank1]: should_exit = checkpoint_and_decide_exit(
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2989, in checkpoint_and_decide_exit
[rank1]: save_checkpoint_and_time(
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2816, in save_checkpoint_and_time
[rank1]: save_checkpoint(
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 693, in save_checkpoint
[rank1]: state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0])
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 1076, in preprocess_fsdp_dtensor_state_dict
[rank1]: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict(
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 378, in handle_swiglu_in_state_dict
[rank1]: weight_w, weight_v = split_swiglu_linear_fc1(
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 278, in split_swiglu_linear_fc1
[rank1]: assert data.shape[swiglu_shard_axis] % 2 == 0, (
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: AssertionError: SWiGLU weights must have an even size along the shard axis 0, got 2051
[rank2]: Traceback (most recent call last):
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/pretrain_gpt_2.py", line 408, in
[rank2]: pretrain(full_config,
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 1385, in pretrain
[rank2]: iteration, num_floating_point_operations_so_far = train(
[rank2]: ^^^^^^
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 3759, in train
[rank2]: should_exit = checkpoint_and_decide_exit(
[rank2]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2989, in checkpoint_and_decide_exit
[rank2]: save_checkpoint_and_time(
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2816, in save_checkpoint_and_time
[rank2]: save_checkpoint(
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 693, in save_checkpoint
[rank2]: state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0])
[rank2]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 1076, in preprocess_fsdp_dtensor_state_dict
[rank2]: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict(
[rank2]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 378, in handle_swiglu_in_state_dict
[rank2]: weight_w, weight_v = split_swiglu_linear_fc1(
[rank2]: ^^^^^^^^^^^^^^^^^^^^^^^^
[rank2]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 295, in split_swiglu_linear_fc1
[rank2]: local_tensor = data.to_local()
[rank2]: ^^^^^^^^^^^^^
[rank2]: AttributeError: 'Tensor' object has no attribute 'to_local'
[rank0]: Traceback (most recent call last):
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/pretrain_gpt_2.py", line 408, in
[rank0]: pretrain(full_config,
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 1385, in pretrain
[rank0]: iteration, num_floating_point_operations_so_far = train(
[rank0]: ^^^^^^
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 3759, in train
[rank0]: should_exit = checkpoint_and_decide_exit(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2989, in checkpoint_and_decide_exit
[rank0]: save_checkpoint_and_time(
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/training.py", line 2816, in save_checkpoint_and_time
[rank0]: save_checkpoint(
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 693, in save_checkpoint
[rank0]: state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0])
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/training/checkpointing.py", line 1076, in preprocess_fsdp_dtensor_state_dict
[rank0]: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 378, in handle_swiglu_in_state_dict
[rank0]: weight_w, weight_v = split_swiglu_linear_fc1(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/mnt/p.geyn/tgpt-megatron/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/fsdp_dtensor_checkpoint.py", line 295, in split_swiglu_linear_fc1
[rank0]: local_tensor = data.to_local()
[rank0]: ^^^^^^^^^^^^^
[rank0]: AttributeError: 'Tensor' object has no attribute 'to_local'
[rank0]:[W610 08:07:01.919411572 ProcessGroupNCCL.cpp:1569] Warning: WARNING: destroy_process_group() was not called before program exit, which can leak resources. For more info, please see https://pytorch.org/docs/stable/distributed.html#shutdown (function operator())
```

**Expected behavior**

Checkpoint is correctly saved

**Additional context**

Megatron-Bridge issue: https://github.com/NVIDIA-NeMo/Megatron-Bridge/issues/4236

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.