如何在Lightning的training_step中实现GPU间张量迁移?
如何在pl.LightningModule的training_step中实现跨GPU张量双向交换?
问题背景
已有用于PyTorch多GPU训练的双向邻接张量交换函数,但不清楚如何在Lightning框架中集成使用,核心困惑在于如何获取设备/进程ID,以及缺乏torch.distributed.P2POp与Lightning结合的实践示例。交换函数如下:
def neighbour_exchange_bidir(left_rank, right_rank, tensor_to_left, tensor_to_right, group=None): tensor_from_left = torch.zeros_like(tensor_to_right) tensor_from_right = torch.zeros_like(tensor_to_left) send_op_left = torch.distributed.P2POp( torch.distributed.isend, tensor_to_left, left_rank, group=group, ) send_op_right = torch.distributed.P2POp( torch.distributed.isend, tensor_to_right, right_rank, group=group, ) recv_op_left = torch.distributed.P2POp( torch.distributed.irecv, tensor_from_left, left_rank, group=group, ) recv_op_right = torch.distributed.P2POp( torch.distributed.irecv, tensor_from_right, right_rank, group=group, ) reqs = torch.distributed.batch_isend_irecv([send_op_right, send_op_left, recv_op_right, recv_op_left]) for req in reqs: req.wait() return tensor_from_right, tensor_from_left
解决方案
1. 获取进程Rank与分布式组
在LightningModule中,可直接通过内置属性获取当前进程的rank信息:
self.trainer.global_rank:全局进程ID(适用于多节点多GPU场景)self.local_rank:单节点内的GPU进程ID- 若需自定义通信组,可在
setup方法中创建:
def setup(self, stage: str): if stage == "fit": # 创建包含所有进程的自定义组(示例) self.comm_group = torch.distributed.new_group(list(range(self.trainer.world_size)))
2. 在training_step中调用交换函数
直接将Lightning获取的rank传入你的函数即可,无需额外处理设备——Lightning会自动将模型和张量分配到当前进程对应的GPU设备上。以下是完整的LightningModule示例:
import pytorch_lightning as pl import torch class MyModel(pl.LightningModule): def __init__(self, world_size): super().__init__() self.world_size = world_size self.comm_group = None def setup(self, stage: str): if stage == "fit": # 初始化自定义通信组(可选,默认用全局组) self.comm_group = torch.distributed.new_group(list(range(self.world_size))) def training_step(self, batch, batch_idx): # 假设当前进程需要发送给左右邻接进程的张量 tensor_to_left = torch.randn(32, 128, device=self.device) tensor_to_right = torch.randn(32, 128, device=self.device) # 计算左右邻接的rank(环形拓扑示例) current_rank = self.trainer.global_rank left_rank = (current_rank - 1) % self.world_size right_rank = (current_rank + 1) % self.world_size # 调用双向交换函数 tensor_from_right, tensor_from_left = neighbour_exchange_bidir( left_rank=left_rank, right_rank=right_rank, tensor_to_left=tensor_to_left, tensor_to_right=tensor_to_right, group=self.comm_group # 若用默认全局组,可传None ) # 后续逻辑:使用交换后的张量进行训练计算 loss = torch.mean(tensor_from_right + tensor_from_left) # 示例损失计算 self.log("train_loss", loss) return loss def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr=1e-3)
3. 关键注意事项
- 分布式启动方式:需用支持分布式的启动命令,例如:
或在Trainer中配置:torchrun --nproc_per_node=4 train.pytrainer = pl.Trainer( accelerator="gpu", devices=4, strategy="ddp", max_epochs=10 ) - 避免死锁:确保所有进程的邻接rank逻辑一致(比如环形拓扑中每个进程的左右rank计算规则统一),且发送/接收的张量形状、设备匹配。
- 张量设备:无需手动迁移张量,Lightning会自动将
training_step中的张量放到当前进程对应的GPU上,torch.distributed操作会自动处理跨设备通信。
内容的提问来源于stack exchange,提问作者The Hidden Reverse
相关产品推荐
相关产品推荐

