PyTorch多GPU大矩阵乘法实现问题及优化方案咨询
多GPU矩阵乘法内存优化与张量并行问题
需求
使用多GPU执行矩阵乘法(如torch.mm(a, b)),降低单GPU的内存占用。
单GPU运行代码
import torch a = torch.randn(30000, 30000).cuda(1) b = torch.randn(30000, 30000).cuda(1) c = torch.mm(a, b) # 此过程中最大内存占用为10491 MB。
双GPU手动拆分实现(存在OOM问题)
手动拆分大矩阵到双GPU计算,结果拼接时触发内存不足:
import torch # 假设`a1`和`a2`是大矩阵的拆分部分 a1 = torch.randn(15000, 30000).cuda(0) a2 = torch.randn(15000, 30000).cuda(1) b1 = torch.randn(30000, 30000).cuda(0) b2 = b1.cuda(1) c1 = torch.mm(a1,b1) c2 = torch.mm(a2,b2).to(0) # 当前结果`c1`和`c2`位于GPU 0 # GPU 1的最大内存占用为7059 MB # GPU 0的最大内存占用为8777 MB,因结果存储于此而高于GPU 1 c = torch.concat([c1, c2], dim=0) # 因concat非原地操作导致OOM
疑问:
- 实现concat的原地操作能否解决该OOM问题?
- 或者应该先将
c1和c2移至CPU内存拼接,之后再移回GPU?
PyTorch 2.2张量并行尝试(进程重复生成张量问题)
使用PyTorch 2.2的张量并行功能,但启动双进程时,代码会执行两次,导致生成两个不同的big_tensor_1:
import torch import torch.distributed as distributed import os from torch.distributed._tensor import init_device_mesh, Shard, distribute_tensor from torch.distributed.tensor.parallel import parallelize_module, ColwiseParallel from visualize_sharding import visualize_sharding mesh = init_device_mesh("cuda", (2,)) rank = distributed.get_rank() big_tensor_1 = torch.randn(3, 2) big_tensor_2 = torch.randn(2, 6) print("big_tensor_1", big_tensor_1) my_dtensor_1 = distribute_tensor(big_tensor_1, mesh, [Shard(dim=0)]) my_dtensor_2 = distribute_tensor(big_tensor_2, mesh, [Shard(dim=1)]) # visualize_sharding(my_dtensor_1, header="my_dtensor_1") c = torch.mm(my_dtensor_1, my_dtensor_2) print("c: ", c)
运行命令:
python -m torch.distributed.launch --nproc_per_node=2 --nnodes=1 tmp.py
问题:如何修改代码,实现张量只生成一次,由双进程共同使用?
内容的提问来源于stack exchange,提问作者zenga
相关产品推荐
相关产品推荐

