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 仅被消耗一次。