第1章 探索 TorchRec 分片机制
探索 TorchRec 的分片 #创建时间:2022年5月10日 | 最后更新:2022年5月13日 | 最后验证:2024年11月5日
本教程主要介绍通过 EmbeddingPlanner 和 DistributedModelParallel API 实现嵌入表(embedding tables)的分片方案,并通过显式配置来探索不同分片方案对嵌入表带来的收益。
安装 # 要求: - python >= 3.7
强烈建议在使用 TorchRec 时搭配 CUDA 使用。若使用 CUDA: - cuda >= 11.0
# 安装 conda,以便更轻松地安装带有 cudatoolkit 11.3 的 PyTorch。 !sudo rm Miniconda3-py37_4.9.2-Linux-x86_64.sh Miniconda3-py37_4.9.2-Linux-x86_64.sh.* !sudo wget https://repo.anaconda.com/miniconda/Miniconda3-py37_4.9.2-Linux-x86_64.sh !sudo chmod +x Miniconda3-py37_4.9.2-Linux-x86_64.sh !sudo bash ./Miniconda3-py37_4.9.2-Linux-x86_64.sh -b -f -p /usr/local
# 安装带有 cudatoolkit 11.3 的 PyTorch !sudo conda install pytorch cudatoolkit=11.3 -c pytorch-nightly -y
安装 TorchRec 会自动安装 FBGEMM,这是一组包含 CUDA kernel 和 GPU 加速操作的集合。 # 安装 torchrec !pip3 install torchrec-nightly
安装 multiprocess,以便在 Colab 中结合 ipython 进行多进程编程。 !pip3 install multiprocess
以下步骤用于让 Colab 运行时环境检测新添加的共享库。运行时会在 /usr/lib 中搜索共享库,因此我们将安装在 /usr/local/lib/ 中的库复制过去。这是仅适用于 Colab 运行时的必要步骤。 !sudo cp /usr/local/lib/lib* /usr/lib/
在此处重启运行时环境,以便识别新安装的包。重启后立即执行以下步骤,让 Python 知道去哪里查找包。每次重启运行时后都要执行此步骤。 import sys sys.path = ['', '/env/python', '/usr/local/lib/python37.zip', '/usr/local/lib/python3.7', '/usr/local/lib/python3.7/lib-dynload', '/usr/local/lib/python3.7/site-packages', './.local/lib/python3.7/site-packages']
分布式配置 # 受限于 Notebook 环境,我们无法在此运行 SPMD 程序,但可以在 Notebook 内部使用多进程来模拟该设置。用户在使用 TorchRec 时应自行负责设置 SPMD 启动器。我们配置环境以确保基于 torch distributed 的通信后端能够正常工作。 import os import torch import torchrec
os.environ["MASTER_ADDR"] = "localhost" os.environ["MASTER_PORT"] = "29500"
构建嵌入模型 # 这里我们使用 TorchRec 提供的 EmbeddingBagCollection 来构建包含嵌入表的嵌入 bag 模型。 这里,我们创建一个包含四个嵌入 bag 的 EmbeddingBagCollection (EBC)。表分为两种类型:大表和小表,通过行数差异来区分:4096 对比 1024。每个表仍由 64 维嵌入表示。 我们为表配置 ParameterConstraints 数据结构,为模型并行 API 提供提示,帮助决定表的分片和放置策略。在 TorchRec 中,我们支持: * table-wise:将整个表放在一个设备上;* row-wise:按行维度均匀分片表,并将每个分片放在通信域中的每个设备上;* column-wise: 按嵌入维度均匀分片表,并将每个分片放在通信域中的每个设备上;* table-row-wise: 针对可用快速机内设备互连(如 NVLink)优化的特殊分片方式,用于主机内通信;* data_parallel: 在每个设备上复制表; 注意我们最初在 device “meta” 上分配 EBC。这将告诉 EBC 暂时不分配内存。 from torchrec.distributed.planner.types import ParameterConstraints from torchrec.distributed.embedding_types import EmbeddingComputeKernel from torchrec.distributed.types import ShardingType from typing import Dict
large_table_cnt = 2 small_table_cnt = 2 large_tables=[ torchrec.EmbeddingBagConfig( name="large_table_" + str(i), embedding_dim=64, num_embeddings=4096, feature_names=["large_table_feature_" + str(i)], pooling=torchrec.PoolingType.SUM, ) for i in range(large_table_cnt) ] small_tables=[ torchrec.EmbeddingBagConfig( name="small_table_" + str(i), embedding_dim=64, num_embeddings=1024, feature_names=["small_table_feature_" + str(i)], pooling=torchrec.PoolingType.SUM, ) for i in range(small_table_cnt) ]
def gen_constraints(sharding_type: ShardingType = ShardingType.TABLE_WISE) -> Dict[str, ParameterConstraints]: large_table_constraints = { "large_table_" + str(i): ParameterConstraints( sharding_types=[sharding_type.value], ) for i in range(large_table_cnt) } small_table_constraints = { "small_table_" + str(i): ParameterConstraints( sharding_types=[sharding_type.value], ) for i in range(small_table_cnt) } constraints = {large_table_constraints, small_table_constraints} return constraints
ebc = torchrec.EmbeddingBagCollection( device="cuda", tables=large_tables + small_tables ]
多进程中的 DistributedModelParallel # 现在,我们有一个单进程执行函数,用于模拟 SPMD 执行过程中单个 rank 的工作。 这段代码将与其他进程协同分片模型,并相应地分配内存。它首先设置进程组,使用 planner 进行嵌入表放置,并使用 DistributedModelParallel 生成分片模型。 def single_rank_execution( rank: int, world_size: int, constraints: Dict[str, ParameterConstraints], module: torch.nn.Module, backend: str, ) -> None: import os import torch import torch.distributed as dist from torchrec.distributed.embeddingbag import EmbeddingBagCollectionSharder from torchrec.distributed.model_parallel import DistributedModelParallel from torchrec.distributed.planner import EmbeddingShardingPlanner, Topology from torchrec.distributed.types import ModuleSharder, ShardingEnv from typing import cast
def init_distributed_single_host( rank: int, world_size: int, backend: str, # pyre-fixme[11]: Annotation `ProcessGroup` is not defined as a type. ) -> dist.ProcessGroup: os.environ["RANK"] = f"{rank}" os.environ["WORLD_SIZE"] = f"{world_size}" dist.init_process_group(rank=rank, world_size=world_size, backend=backend) return dist.group.WORLD
if backend == "nccl": device = torch.device(f"cuda:{rank}") torch.cuda.set_device(device) else: device = torch.device("cpu") topology = Topology(world_size=world_size, compute_device="cuda") pg = init_distributed_single_host(rank, world_size, backend) planner = EmbeddingShardingPlanner( topology=topology, constraints=constraints, ) sharders = [cast(ModuleSharder[torch.nn.Module], EmbeddingBagCollectionSharder())] plan: ShardingPlan = planner.collective_plan(module, sharders, pg)
sharded_model = DistributedModelParallel( module, env=ShardingEnv.from_process_group(pg), plan=plan, sharders=sharders, device=device, ) print(f"rank:{rank},sharding plan: {plan}") return sharded_model
多进程执行 # 现在让我们在代表多个 GPU rank 的多进程中执行代码。 import multiprocess
def spmd_sharing_simulation( sharding_type: ShardingType = ShardingType.TABLE_WISE, world_size = 2, ): ctx = multiprocess.get_context("spawn") processes = [] for rank in range(world_size): p = ctx.Process( target=single_rank_execution, args=( rank, world_size, gen_constraints(sharding_type), ebc, "nccl" ), ) p.start() processes.append(p)
for p in processes: p.join() assert 0 == p.exitcode
表级分片 # 现在让我们执行代码,使用两个进程对应 2 个 GPU。我们可以从打印的计划中看到表如何在 GPU 之间分片。每个节点有一个大表和一个小表,这表明我们的 planner 试图对嵌入表进行负载平衡。对于许多中小尺寸的表,Table-wise 是默认的、通用的分片方案,用于在设备间进行负载平衡。 spmd_sharing_simulation(ShardingType.TABLE_WISE)
rank:1,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[0], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 64], placement=rank:0/cuda:0)])), 'large_table_1': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 64], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[0], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 64], placement=rank:0/cuda:0)])), 'small_table_1': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 64], placement=rank:1/cuda:1)]))}} rank:0,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[0], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 64], placement=rank:0/cuda:0)])), 'large_table_1': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 64], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[0], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 64], placement=rank:0/cuda:0)])), 'small_table_1': ParameterSharding(sharding_type='table_wise', compute_kernel='batched_fused', ranks=[1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 64], placement=rank:1/cuda:1)]))}}
探索其他分片模式 # 我们最初探索了 table-wise 分片的样子以及它如何平衡表的放置。现在我们探索对负载平衡关注更细致的分片模式:row-wise。Row-wise 专门针对那些因嵌入行数增加导致内存占用过大而单个设备无法承载的大表。它可以解决模型中超大表的放置问题。用户可以在打印的计划日志的 shard_sizes 部分看到,表按行维度减半以分布到两个 GPU 上。 spmd_sharing_simulation(ShardingType.ROW_WISE)
rank:1,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[2048, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[2048, 0], shard_sizes=[2048, 64], placement=rank:1/cuda:1)])), 'large_table_1': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[2048, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[2048, 0], shard_sizes=[2048, 64], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[512, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[512, 0], shard_sizes=[512, 64], placement=rank:1/cuda:1)])), 'small_table_1': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[512, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[512, 0], shard_sizes=[512, 64], placement=rank:1/cuda:1)]))}} rank:0,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[2048, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[2048, 0], shard_sizes=[2048, 64], placement=rank:1/cuda:1)])), 'large_table_1': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[2048, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[2048, 0], shard_sizes=[2048, 64], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[512, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[512, 0], shard_sizes=[512, 64], placement=rank:1/cuda:1)])), 'small_table_1': ParameterSharding(sharding_type='row_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[512, 64], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[512, 0], shard_sizes=[512, 64], placement=rank:1/cuda:1)]))}}
另一方面,Column-wise 解决的是具有较大嵌入维度的表的负载不平衡问题。我们将垂直切分表。用户可以在打印的计划日志的 shard_sizes 部分看到,表按嵌入维度减半以分布到两个 GPU 上。 spmd_sharing_simulation(ShardingType.COLUMN_WISE)
rank:0,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[4096, 32], placement=rank:1/cuda:1)])), 'large_table_1': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[4096, 32], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[1024, 32], placement=rank:1/cuda:1)])), 'small_table_1': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[1024, 32], placement=rank:1/cuda:1)]))}} rank:1,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[4096, 32], placement=rank:1/cuda:1)])), 'large_table_1': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[4096, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[4096, 32], placement=rank:1/cuda:1)])), 'small_table_0': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[1024, 32], placement=rank:1/cuda:1)])), 'small_table_1': ParameterSharding(sharding_type='column_wise', compute_kernel='batched_fused', ranks=[0, 1], sharding_spec=EnumerableShardingSpec(shards=[ShardMetadata(shard_offsets=[0, 0], shard_sizes=[1024, 32], placement=rank:0/cuda:0), ShardMetadata(shard_offsets=[0, 32], shard_sizes=[1024, 32], placement=rank:1/cuda:1)]))}}
对于 table-row-wise,不幸的是我们无法模拟它,因为其性质决定了它需要在多主机设置下运行。未来我们将提供一个 Python SPMD 示例来训练使用 table-row-wise 的模型。 使用 data parallel,我们会在所有设备上复制表。 spmd_sharing_simulation(ShardingType.DATA_PARALLEL)
rank:0,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'large_table_1': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'small_table_0': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'small_table_1': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None)}} rank:1,sharding plan: {'': {'large_table_0': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'large_table_1': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'small_table_0': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None), 'small_table_1': ParameterSharding(sharding_type='data_parallel', compute_kernel='batched_dense', ranks=[0, 1], sharding_spec=None)}}