已开启
【RFC】hyper_parallel.hsdp 兼容 torch.distributed.fsdp.fully_shard - HSDPParam #14
cuiyushi创建于 1月30日
1月30日 关联了pull request:[feat/hsdp_refactor] add TorchHSDPParamV2 for HSDP parameter management
1月30日 关联了pull request:[feat/hsdp_refactor] add TorchHSDPParamV2 for HSDP parameter management
2月9日 关联了pull request:add_fully_shard_module_testcase
2月13日 关联了pull request:bugfix_fully_shard_post_backward
背景
使用数据并行进行训练时,各节点存储的模型参数是相同的,模型参数更新过程使用的优化器状态也是相同的。从存储的角度看,集群中模型参数和优化器状态存在冗余存储。从计算角度看,节点间更新模型参数的优化器计算完全一致,存在冗余计算。如果将模型参数和优化器状态按节点个数切分,每个节点存储不同的模型参数切片和优化器状态切片,不仅消除了冗余存储,由于通过优化器参与模型参数更新的是切片数据,还能消除优化器的冗余计算。这种并行方式,我们一般称为优化器并行,又称ZeRO(Zero Redundancy Optimizer)。
目前,业界中主要使用的是Torch的fully_shard接口,其本质上就是ZeRO-3的实现。
方案
基本原理
以业界常见的优化器并行方案ZeRO为例,有ZeRO-1,ZeRO-2,ZeRO-3三种优化级别:
ZeRO1
ZeRO2
ZeRO3
接口设计
在接口设计上,我们选择兼容Torch的fully_shard接口,设计如下
2.1 fully_shard 接口入参分析
cellmeshreshard_after_forwardshard_placement_fnmp_policyoffload_policyignored_params2.2 HSDPCell 核心方法分析
HSDPCell提供以下方法实现对模型运行时状态的精细化调控。set_requires_gradient_syncset_forward_prefetch_cellsset_backward_prefetch_cellsreshardunshardset_is_last_backwardset_reshard_after_forwardset_reshard_after_backwardset_requires_all_reduce可以通过上述接口搭配,达到ZeRO-1, ZeRO-2,ZeRO-3级别的优化器并行切分策略:
ZeRO-1=rehsard_after_forward=True+set_require_gradient_sync=FalseZeRO-2=rehsard_after_forward=True+set_require_gradient_sync=TrueZeRO-3=rehsard_after_forward=True+set_require_gradient_sync=True2.3 HSDPParam
HSDPParam 的主要作用是维护参数状态,进行参数分片计算,对张量生命周期进行管理。
_init_sharded_param_init_sharded_post_forward_param_metadatainit_all_gather_outputsinit_unsharded_paramto_shardedto_sharded_post_forwardto_unshardedto_sharded_dtensorto_sharded_post_forward_dtensoralloc_all_gather_outputsfree_unsharded_paramall_gather_inputsto_accumulated_grad_if_neededaccumulate_unsharded_grad_if_needed