deepspeedai / deepspeedai/DeepSpeed

[BUG] Deepspeed zero3 error in loading pretrained model by transformers.LlamaForCausalLM.from_pretrained function.

Open
#5,250 2 comments 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

bug training
Dominant language
Python
Stars
43.1k
Forks
5k
Avg merge
4d 15h
Merged PRs (30d)
112

Description

Describe the bug
A clear and concise description of what the bug is.
"Unable to load weights from pytorch checkpoint file for '../llama/hug_llama_2_7b_chat/pytorch_model-00001-of-00003.bin"
When I try to pretrain llava that uses Llama-2-7b-chat as language model using deepspeed zero3 on 2 v100 gpus on 1 node, it shows that rank 0 process and rank 1 process both try to load '../llama/hug_llama_2_7b_chat/pytorch_model-00001-of-00003.bin" simultaneously and finally an error was reported as "OSError: Unable to load weights from pytorch checkpoint file for '../llama/hug_llama_2_7b_chat/pytorch_model-00001-of-00003.bin' at '../llama/hug_llama_2_7b_chat/pytorch_model-00001-of-00003.bin'. If you tried to load a PyTorch model from a TF 2.0 checkpoint, please set from_tf=True.".
However, when I trained with deepspeed zero2, it works successfully. How can I fix the error with zero3?
To Reproduce
Steps to reproduce the behavior:

  1. zero2 config file
    {
    "fp16": {
    "enabled": "auto",
    "loss_scale": 0,
    "loss_scale_window": 1000,
    "initial_scale_power": 16,
    "hysteresis": 2,
    "min_loss_scale": 1
    },
    "bf16": {
    "enabled": "auto"
    },
    "train_micro_batch_size_per_gpu": "auto",
    "train_batch_size": "auto",
    "gradient_accumulation_steps": "auto",
    "zero_optimization": {
    "stage": 2,
    "overlap_comm": true,
    "contiguous_gradients": true,
    "sub_group_size": 1e9,
    "reduce_bucket_size": "auto"
    }
    }
  2. zero3 config file
    {
    "fp16": {
    "enabled": "auto",
    "loss_scale": 0,
    "loss_scale_window": 1000,
    "initial_scale_power": 16,
    "hysteresis": 2,
    "min_loss_scale": 1
    },
    "bf16": {
    "enabled": "auto"
    },
    "train_micro_batch_size_per_gpu": "auto",
    "train_batch_size": "auto",
    "gradient_accumulation_steps": "auto",
    "zero_optimization": {
    "stage": 3,
    "overlap_comm": true,
    "contiguous_gradients": true,
    "sub_group_size": 1e9,
    "reduce_bucket_size": "auto",
    "stage3_prefetch_bucket_size": "auto",
    "stage3_param_persistence_threshold": "auto",
    "stage3_max_live_parameters": 1e9,
    "stage3_max_reuse_distance": 1e9,
    "stage3_gather_16bit_weights_on_model_save": true
    }
    }
  3. transformers.LlamaForCausalLM.from_pretrained to load the sharded weight in huggingface format of LLama-2-7b-chat

Screenshots
image
image

System info (please complete the following information):

  • OS: [Ubuntu 18.04]
  • GPU count and types [one machine with x2 v100s]
  • Interconnects (if applicable) [e.g., two machines connected with 100 Gbps IB]
  • Python 3.9.1
  • transformers 4.37.2
  • deepspeed 0.12.6
  • pytorch 1.12.1 cuda 11.6
### Tasks

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start at transformers.LlamaForCausalLM.from_pretrained and compare the provided ZeRO-2 and ZeRO-3 configurations, focusing on loading the sharded Hugging Face weights. Reproduce with the listed DeepSpeed, Transformers, PyTorch, CUDA, Python, and two-V100 setup. Done means the Llama-2-7b-chat weights load successfully under ZeRO-3 without the checkpoint error.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.