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)