deepspeedai / deepspeedai/DeepSpeed

[BUG]Zero++ quantizer unsupport BFloat16

Open
#3,992 5 comments 4 reactions 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
When using both the Zero++ and BFloat16 features simultaneously. Sometimes the gathered param is Float15 dtype,but the intermediate result are still BFloat16 dtype.

To Reproduce
Steps to reproduce the behavior:

  1. set deepspeed configuration,enable zero++ and bf16.
ds_config = {
    "train_batch_size": GLOBAL_BATCH_SIZE,
    "train_micro_batch_size_per_gpu": MICRO_BATCH_SIZE,
    "steps_per_print": 10,
    "zero_optimization": {
        "stage": 3,
        "offload_param": {
            "device": 'cpu'
        },
        "offload_optimizer": {
            "device": 'cpu'
        },
        "stage3_param_persistence_threshold": 1e4,
        "stage3_max_live_parameters": 3e7,
        "stage3_prefetch_bucket_size": 3e7,
        "memory_efficient_linear": False,

        "zero_quantized_weights": True,
        "zero_hpz_partition_size": 8,
        "zero_quantized_gradients": True,
    },
    "fp16": {
        "enabled": False
    },
    "bf16": {
        "enabled": True
    },
    "gradient_clipping": 1.0,
    "prescale_gradients": False,
    "wall_clock_breakdown": False,
    "hybrid_engine": {
        "enabled": True,
        "inference_tp_size": 8,
        "release_inference_cache": release_inference_cache,
        "pin_parameters": pin_parameters,
        "tp_gather_partition_size": 8,
        "max_out_tokens": 512,
    }
}
  1. Create a Bloom HF model,and initialize model engine by ds_config.
model = AutoModelForCausalLM.from_pretrained(model_name_or_path)
actor_engine, *_ = deepspeed.initialize(model=model, optimizer=optim, config=ds_config)
  1. run inference.
  2. See error.
Traceback (most recent call last):
  File "/opt/conda/lib/python3.8/site-packages/deepspeed/runtime/hybrid_engine.py", line 245, in generate
    generate_ret_vals = self._generate(*inputs, **kwargs)
  File "/opt/conda/lib/python3.8/site-packages/torch/autograd/grad_mode.py", line 27, in decorate_context
    return func(*args, **kwargs)
  File "/opt/conda/lib/python3.8/site-packages/transformers/generation/utils.py", line 1568, in generate
    return self.sample(
  File "/opt/conda/lib/python3.8/site-packages/transformers/generation/utils.py", line 2615, in sample
    outputs = self(
  File "/opt/conda/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1204, in _call_impl
    result = forward_call(*input, **kwargs)
  File "/opt/conda/lib/python3.8/site-packages/transformers/models/bloom/modeling_bloom.py", line 927, in forward
    lm_logits = self.lm_head(hidden_states)
  File "/opt/conda/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1204, in _call_impl
    result = forward_call(*input, **kwargs)
  File "/opt/conda/lib/python3.8/site-packages/deepspeed/module_inject/layers.py", line 52, in forward
    output = torch.matmul(input, self.weight.transpose(-1, -2))
RuntimeError: expected scalar type BFloat16 but found Half

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

Reproduce the failure with the supplied Zero++ and BFloat16 configuration, then inspect runtime/hybrid_engine.py and module_inject/layers.py around inference and the failing matmul. Trace how gathered parameter and intermediate dtypes are selected. Done means the Bloom inference path runs with both features enabled without the BFloat16/Half mismatch.

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
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.