flashinfer.comm.mnnvl.McastGPUBuffer

class flashinfer.comm.mnnvl.McastGPUBuffer(buf_size: int, group_size: int, group_rank: int, device: device, comm_backend_for_handle_transfer: CommBackend | None = None)

用于促进 PyTorch 张量创建的 SymmDeviceMemory 包装类。它管理一个可通过单播或多播进行多节点通信的缓冲区。

TensorRT-LLM 的 McastGPUBuffer 的 Python 移植

__init__(buf_size: int, group_size: int, group_rank: int, device: device, comm_backend_for_handle_transfer: CommBackend | None = None)

McastGpuBuffer 的构造函数。

参数:
  • buf_size – 缓冲区请求的大小,以字节为单位。实际可用大小可能因对齐要求而异。

  • group_size – 通信组中的进程数量

  • group_rank – 组内本地进程的排名

  • device – 缓冲区分配的 CUDA 设备

  • mn_nvlink – 标志,指示是否使用多节点 NVLink

  • comm_backend_for_handle_transfer – 用于句柄传输的通信后端

方法

__init__(buf_size, group_size, group_rank, ...)

McastGpuBuffer 的构造函数。

get_buffer_ptrs_dev()

获取缓冲区指针设备数组

get_multicast_buffer(sizes, dtype[, ...])

返回多播缓冲区部分的 PyTorch 张量视图。

get_multicast_ptr()

获取原始多播指针

get_unicast_buffer(sizes, dtype[, ...])

返回单播缓冲区部分的 PyTorch 张量视图。

get_unicast_ptr(rank)

获取指向给定排名的原始单播指针

lamport_initialize(rank, dtype)