You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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. 关键注意事项

  • 分布式启动方式:需用支持分布式的启动命令,例如:
    torchrun --nproc_per_node=4 train.py
    
    或在Trainer中配置:
    trainer = pl.Trainer(
        accelerator="gpu",
        devices=4,
        strategy="ddp",
        max_epochs=10
    )
    
  • 避免死锁:确保所有进程的邻接rank逻辑一致(比如环形拓扑中每个进程的左右rank计算规则统一),且发送/接收的张量形状、设备匹配。
  • 张量设备:无需手动迁移张量,Lightning会自动将training_step中的张量放到当前进程对应的GPU上,torch.distributed操作会自动处理跨设备通信。

内容的提问来源于stack exchange,提问作者The Hidden Reverse

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 22:15:06