From 98b50d6f7c39937123b04a5316e2e02cba655cc4 Mon Sep 17 00:00:00 2001 From: haozhx23 Date: Wed, 20 Mar 2024 11:25:04 +0000 Subject: [PATCH] testcommitfix --- src/llama_recipes/finetuning.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/llama_recipes/finetuning.py b/src/llama_recipes/finetuning.py index 6b5650b20..2738094ef 100644 --- a/src/llama_recipes/finetuning.py +++ b/src/llama_recipes/finetuning.py @@ -158,9 +158,9 @@ def main(**kwargs): wandb_run.config.update(peft_config) - hsdp_device_mesh = None + hsdp_device_mesh_plan = None if fsdp_config.hsdp and fsdp_config.sharding_strategy == ShardingStrategy.HYBRID_SHARD: - hsdp_device_mesh = hsdp_device_mesh(replica_group_size=fsdp_config.replica_group_size, sharding_group_size=fsdp_config.sharding_group_size) + hsdp_device_mesh_plan = hsdp_device_mesh(replica_group_size=fsdp_config.replica_group_size, sharding_group_size=fsdp_config.sharding_group_size) print("HSDP device mesh is ready") #setting up FSDP if enable_fsdp is enabled @@ -184,7 +184,7 @@ def main(**kwargs): cpu_offload=CPUOffload(offload_params=True) if fsdp_config.fsdp_cpu_offload else None, mixed_precision=mixed_precision_policy if not fsdp_config.pure_bf16 else None, sharding_strategy=fsdp_config.sharding_strategy, - device_mesh=hsdp_device_mesh, + device_mesh=hsdp_device_mesh_plan, device_id=device_id, limit_all_gathers=True, sync_module_states=train_config.low_cpu_fsdp,