从源码到应用:tensor_parallel关键函数tensor_parallel()深度解析
从源码到应用:tensor_parallel关键函数tensor_parallel()深度解析
【免费下载链接】tensor_parallelAutomatically split your PyTorch models on multiple GPUs for training & inference项目地址: https://gitcode.com/gh_mirrors/te/tensor_parallel
tensor_parallel是一个能自动将PyTorch模型在多个GPU上拆分以进行训练和推理的工具,其核心函数tensor_parallel()在实现这一功能中发挥着关键作用。本文将深度解析该函数,帮助新手和普通用户理解其工作原理与应用方法。
一、tensor_parallel()函数基本介绍
tensor_parallel()函数位于src/tensor_parallel/factory.py文件中,它的主要作用是为现有的PyTorch模块添加张量并行功能,并返回等效的张量并行模块。通过该函数,用户可以轻松实现模型在多个设备上的并行处理,提升训练和推理效率。
1.1 函数定义与参数说明
函数的定义如下:
def tensor_parallel( module: nn.Module, device_ids: Optional[Sequence[Union[torch.device, str]]] = None, tensor_parallel_config: Optional[Config] = None, distributed: Optional[bool] = None, sharded: Optional[bool] = None, sharded_param_names: Optional[Collection[str]] = None, **kwargs, ) -> nn.Module:主要参数说明:
- module:原始的PyTorch模块,建议将输入模块存储在CPU上以最小化GPU内存占用。
- device_ids:模型将在设备列表(如GPU)之间拆分,默认是所有可用的CUDA设备。
- tensor_parallel_config:用于描述模型如何并行化的自定义配置,默认为自动配置。
- distributed:若为True,使用torch.distributed而非线程,默认在torch.distributed初始化时为True,否则为False。
- sharded:若为True,任何非张量并行参数(如layernorm权重)仍将被分片,并在每次前向传播时手动重新组装,相当于PyTorch的FullyShardedDataParallel。
- sharded_param_names:当sharded=True时,这是ZeRO-3应用的所有参数名称列表,默认情况下,ZeRO-3适用于所有未使用张量并行拆分的参数。
1.2 简单使用示例
以下是一个简单的使用示例,展示了如何使用tensor_parallel()函数对模型进行并行化处理:
import torch, transformers import tensor_parallel as tp model = transformers.AutoModel.from_pretrained("t5-11b") model = tp.tensor_parallel(model, device_ids=['cuda:0', 'cuda:1']) outputs_as_usual = model(**inputs_as_usual) # 反向传播也适用!二、tensor_parallel()函数工作流程
tensor_parallel()函数的工作流程主要包括分布式模式判断、设备处理以及模块包装等步骤,下面将详细介绍。
2.1 分布式模式判断
函数首先会判断是否采用分布式模式,代码如下:
distributed = distributed if distributed is not None else torch.distributed.is_initialized()这里根据用户传入的distributed参数或当前torch.distributed是否初始化来确定是否使用分布式模式。
2.2 分布式模式下的处理
如果处于分布式模式,函数会对设备进行处理,确保只指定一个当前设备,并返回分布式分片模型,代码如下:
if distributed: if device_ids is None: device_ids = [torch.device("cuda" if torch.cuda.is_available() else "cpu")] assert len(device_ids) == 1, "if distributed=True, please specify a single (current) device" assert not sharded, "distributed + sharded mode is not implemented, please keep one" return make_distributed_shard(module, device=torch.device(device_ids[0]), **kwargs)2.3 非分布式模式下的模块包装
在非分布式模式下,函数会根据模块类型进行不同的包装。如果是PreTrainedModel类型,使用TensorParallelPreTrainedModel进行包装;否则使用TensorParallel进行包装,代码如下:
else: if isinstance(module, PreTrainedModel): return TensorParallelPreTrainedModel( module, device_ids=device_ids, tensor_parallel_config=tensor_parallel_config, distributed=distributed, sharded=sharded, sharded_param_names=sharded_param_names,** kwargs, ) else: return TensorParallel( module, device_ids=device_ids, tensor_parallel_config=tensor_parallel_config, distributed=distributed, sharded=sharded, sharded_param_names=sharded_param_names, **kwargs, )三、关键参数深入解析
为了更好地理解和使用tensor_parallel()函数,下面对一些关键参数进行深入解析。
3.1 device_ids参数
device_ids参数用于指定模型拆分的设备列表。在src/tensor_parallel/tensor_parallel.py中,有对device_ids的检查和处理函数check_device_ids(),它确保设备列表的有效性。如果用户未指定device_ids,函数会默认使用所有可用的CUDA设备或CPU设备。在实际应用中,用户可以根据自己的硬件情况灵活指定设备,例如device_ids=['cuda:0', 'cuda:1']表示将模型拆分到0号和1号GPU上。
3.2 sharded参数
sharded参数决定是否对非张量并行参数进行分片处理。当sharded=True时,会对相关参数进行分片,在src/tensor_parallel/tensor_parallel.py中,apply_sharding()方法会实现这一功能。通过分片处理,可以进一步优化内存使用,提高模型并行效率。但需要注意的是,在分布式模式下,sharded模式暂未实现,不能同时使用。
四、实际应用场景与注意事项
4.1 应用场景
tensor_parallel()函数适用于需要在多个GPU上进行模型训练和推理的场景。例如,当处理大型语言模型(如t5-11b)时,单个GPU的内存可能无法满足需求,此时使用tensor_parallel()函数将模型拆分到多个GPU上,可以有效解决内存不足的问题,同时加快训练和推理速度。
4.2 注意事项
- 在使用分布式模式时,需要确保只指定一个当前设备,并且不能与sharded模式同时使用。
- 对于PreTrainedModel类型的模块和普通nn.Module类型的模块,函数会进行不同的包装处理,用户在使用时无需额外区分,函数会自动判断。
- 在指定device_ids时,要根据实际可用的设备情况进行设置,避免出现设备不存在或不可用的情况。
通过对tensor_parallel()函数的深度解析,相信大家对其工作原理和使用方法有了更清晰的认识。在实际应用中,合理使用该函数可以充分利用多GPU资源,提升模型训练和推理的效率,为处理大型模型提供有力支持。
要使用该项目,可通过以下命令克隆仓库:git clone https://gitcode.com/gh_mirrors/te/tensor_parallel
【免费下载链接】tensor_parallelAutomatically split your PyTorch models on multiple GPUs for training & inference项目地址: https://gitcode.com/gh_mirrors/te/tensor_parallel
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考