flashinfer.comm.pack_strided_memory

flashinfer.comm.pack_strided_memory(ptr: int, segment_size: int, segment_stride: int, num_segments: int, dtype: dtype, dev_id)

将 GPU 内存打包成具有指定步长的 PyTorch 张量。

参数:
  • ptr – 从 cudaMalloc 获得的 GPU 内存地址

  • segment_size – 每个段的内存大小,以字节为单位

  • segment_stride – 段之间的内存步长大小,以字节为单位

  • num_segments – 段的数量

  • dtype – 结果张量的 PyTorch 数据类型

  • dev_id – CUDA 设备 ID

返回值:

引用所提供内存的 PyTorch 张量

注意

即使指针相同,此函数每次调用都会创建一个新的 DLPack capsule。每个 capsule 仅被消耗一次。