facebookresearch / facebookresearch/fairscale

Non shared params is empty when large translation model is traind by FSDP.

Open
#925 13 comments 0 reactions 1 assignee Claimed by @anj-s View on GitHub
FSDP help wanted
Dominant language
Python
Stars
3.4k
Forks
293
PR merge metrics
No merged PRs in 30d

Description

I'm trying to train translation model which is using transformer encoder-decoder model with FSDP.
But the same problem as https://github.com/facebookresearch/fairscale/issues/883 occures even if fairscale 0.4.5 fairseq 1.0.0a0+4a7835b.
When I run the following command, I got the same error.

``` bash
fairseq-train data-bin/iwslt14.tokenized.de-en --arch transformer_iwslt_de_en --share-decoder-input-output-embed --optimizer cpu_adam --adam-betas '(0.9, 0.98)' --clip-norm 0.0 --lr 5e-4 --lr-scheduler inverse_sqrt --warmup-updates 4 --dropout 0.3 --weight-decay 0.0001 --criterion label_smoothed_cross_entropy --label-smoothing 0.1 --max-tokens 4096 --max-update 4 --ddp-backend fully_sharded --fp16 --fp16-init-scale 4 --cpu-offload --log-format json --log-interval 1 --encoder-layers 6 --decoder-layers 6 --encoder-embed-dim 5120 --decoder-embed-dim 5120 --encoder-ffn-embed-dim 5120 --decoder-ffn-embed-dim 5120 --encoder-attention-heads 40 --decoder-attention-heads 40
```

```
Traceback (most recent call last):
File "/home/christopher/gitrepos/fairseq/env/bin/fairseq-train", line 33, in
sys.exit(load_entry_point('fairseq', 'console_scripts', 'fairseq-train')())
File "/home/christopher/gitrepos/fairseq/fairseq_cli/train.py", line 528, in cli_main
distributed_utils.call_main(cfg, main)
File "/home/christopher/gitrepos/fairseq/fairseq/distributed/utils.py", line 344, in call_main
torch.multiprocessing.spawn(
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/torch/multiprocessing/spawn.py", line 230, in spawn
return start_processes(fn, args, nprocs, join, daemon, start_method='spawn')
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/torch/multiprocessing/spawn.py", line 188, in start_processes
while not context.join():
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/torch/multiprocessing/spawn.py", line 150, in join
raise ProcessRaisedException(msg, error_index, failed_process.pid)
torch.multiprocessing.spawn.ProcessRaisedException:

-- Process 1 terminated with the following error:
Traceback (most recent call last):
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/torch/multiprocessing/spawn.py", line 59, in _wrap
fn(i, *args)
File "/home/christopher/gitrepos/fairseq/fairseq/distributed/utils.py", line 328, in distributed_main
main(cfg, **kwargs)
File "/home/christopher/gitrepos/fairseq/fairseq_cli/train.py", line 188, in main
valid_losses, should_stop = train(cfg, trainer, task, epoch_itr)
File "/usr/lib/python3.9/contextlib.py", line 79, in inner
return func(*args, **kwds)
File "/home/christopher/gitrepos/fairseq/fairseq_cli/train.py", line 317, in train
valid_losses, should_stop = validate_and_save(
File "/home/christopher/gitrepos/fairseq/fairseq_cli/train.py", line 414, in validate_and_save
checkpoint_utils.save_checkpoint(
File "/home/christopher/gitrepos/fairseq/fairseq/checkpoint_utils.py", line 50, in save_checkpoint
trainer.consolidate_optimizer() # TODO(SS): do we need this if no_save_optimizer_state
File "/home/christopher/gitrepos/fairseq/fairseq/trainer.py", line 380, in consolidate_optimizer
st = self.model.gather_full_optim_state_dict(
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/fairscale/nn/data_parallel/fully_sharded_data_parallel.py", line 2267, in gather_full_optim_state_dict
state, singleton_state = self._gather_optim_state(sd.pop("state"))
File "/home/christopher/gitrepos/fairseq/env/lib/python3.9/site-packages/fairscale/nn/data_parallel/fully_sharded_data_parallel.py", line 2211, in _gather_optim_state
assert (
AssertionError: Only flatten param or a single non-shared param is supported: len=0
```

My environment is below.
- PyTorch: 1.10.2
- Fairseq: 1.0.0a0+4a7835b
- Fairscale: 0.4.5
- GeForce Titan RTX x2

`params` is empty because `encoder` and `decoder` the transformer model of are both sharded,
https://github.com/facebookresearch/fairscale/blob/main/fairscale/nn/data_parallel/fully_sharded_data_parallel.py#L414

Then `to_be_flatten` is empty, and therefore self._fsdp_wrapped_module is empty.
https://github.com/facebookresearch/fairscale/blob/main/fairscale/nn/data_parallel/fully_sharded_data_parallel.py#L442

As a result, non_shared_params is empty, which seems to satisfy the assert condition.

Is there any workaround?

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.