I am now trying to use FSDP in Huggingface transformers Trainer. The training script is something like
train_dataset = Mydataset(...)
args = TrainingArguments(...)
model = LlamaForCausalLM.from_pretrained(model_path, attn_implementation="flex_attention")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
trainer = Trainer(model=model, train_dataset=train_dataset, args=args, processing_class=tokenizer)
trainer.train()
In Trainer, I modified training_step function for using flex attention. And it worked just fine when I was using DDP.
The command I use for training is like
accelerate launch --config_file default_config.yaml train.py
Here the default_config.yaml is
compute_environment: LOCAL_MACHINE
debug: true
distributed_type: FSDP
downcast_bf16: 'no'
dynamo_config:
dynamo_backend: EAGER
dynamo_mode: default
dynamo_use_dynamic: true
dynamo_use_fullgraph: true
dynamo_use_regional_compilation: true
enable_cpu_affinity: false
fsdp_config:
fsdp_activation_checkpointing: true
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_transformer_layer_cls_to_wrap: LlamaDecoderLayer
fsdp_backward_prefetch: BACKWARD_PRE
fsdp_cpu_ram_efficient_loading: false
fsdp_forward_prefetch: false
fsdp_offload_params: true
fsdp_reshard_after_forward: FULL_SHARD
fsdp_state_dict_type: FULL_STATE_DICT
fsdp_sync_module_states: true
fsdp_use_orig_params: true
fsdp_version: 1
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 2
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
And I came up with the following error
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/transformers/trainer.py", line 2328, in train
[rank1]: return inner_training_loop(
[rank1]: ^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/transformers/trainer.py", line 2469, in _inner_training_loop
[rank1]: self.model = self.accelerator.prepare(self.model)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/accelerator.py", line 1559, in prepare
[rank1]: result = tuple(
[rank1]: ^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/accelerator.py", line 1560, in
[rank1]: self._prepare_one(obj, first_pass=True, device_placement=d) for obj, d in zip(args, device_placement)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/accelerator.py", line 1402, in _prepare_one
[rank1]: return self.prepare_model(obj, device_placement=device_placement)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/accelerator.py", line 1995, in prepare_model
[rank1]: model = compile_regions(model, **self.state.dynamo_plugin.to_kwargs())
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/utils/other.py", line 165, in compile_regions
[rank1]: new_module = _compile_regions(module, **compile_kwargs)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/accelerate/utils/other.py", line 159, in _compile_regions
[rank1]: new_module.add_module(name, _compile_regions(submodule, **compile_kwargs))
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/nn/modules/module.py", line 640, in add_module
[rank1]: elif hasattr(self, name) and name not in self._modules:
[rank1]: ^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/distributed/fsdp/fully_sharded_data_parallel.py", line 540, in __getattr__
[rank1]: return getattr(self._fsdp_wrapped_module, name)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/distributed/fsdp/fully_sharded_data_parallel.py", line 540, in __getattr__
[rank1]: return getattr(self._fsdp_wrapped_module, name)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/distributed/fsdp/fully_sharded_data_parallel.py", line 540, in __getattr__
[rank1]: return getattr(self._fsdp_wrapped_module, name)
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: [Previous line repeated 984 more times]
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/distributed/fsdp/fully_sharded_data_parallel.py", line 538, in __getattr__
[rank1]: return super().__getattr__(name) # defer to nn.Module's logic
[rank1]: ^^^^^^^^^^^^^^^^^^^^^^^^^
[rank1]: File "/home/xvehao/.conda/envs/tree/lib/python3.11/site-packages/torch/nn/modules/module.py", line 1940, in __getattr__
[rank1]: raise AttributeError(
[rank1]: ^^^^^^^^^^^^^^^
The package version
# Name Version Build Channel
torch 2.7.1+cu118 pypi_0 pypi
torchaudio 2.7.1+cu118 pypi_0 pypi
torchvision 0.22.1+cu118 pypi_0 pypi
transformers 4.56.0 pypi_0 pypi
accelerate 1.11.0 pypi_0 pypi
I wonder what could possibly cause such error and how can I fix it?