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

如何基于pybind11返回的设备内存指针创建PyTorch Tensor并接管内存?

解决方案

要让PyTorch直接接管Kokkos分配的CUDA设备内存,同时避免C库依赖PyTorch C API,核心是让PyTorch绑定自定义内存释放逻辑,确保Tensor销毁时自动调用Kokkos的内存释放函数。具体实现分两步:

1. 暴露Kokkos内存释放函数到Python

通过pybind11把Kokkos的kokkos_free封装成Python可调用的函数,用于后续给PyTorch Tensor指定释放逻辑。

C++端添加绑定代码

在你的pybind11模块定义中新增一个释放函数:

// 封装Kokkos的内存释放逻辑
void kokkos_free_device_ptr(uintptr_t data_ptr) {
    float* ptr = reinterpret_cast<float*>(data_ptr);
    Kokkos::kokkos_free<Kokkos::DefaultExecutionSpace>(ptr);
}

// 模块注册时添加该函数
PYBIND11_MODULE(tenex, m) {
    m.def("process_tensor", &process_tensor, "Process tensor via Kokkos GPU kernels");
    m.def("kokkos_free_device_ptr", &kokkos_free_device_ptr, "Free memory allocated by Kokkos on device");
}

2. 在Python层创建带自定义释放逻辑的Tensor

使用PyTorch的torch.Tensor._from_data_ptr方法,直接从设备指针创建Tensor,并绑定自定义释放函数,让PyTorch接管内存所有权。

Python实现代码

import torch
import tenex

def create_tensor_from_kokkos_ptr(data_ptr, size, dtype=torch.float32, device='cuda'):
    # 定义Tensor销毁时调用的释放函数
    def deleter(ptr):
        tenex.kokkos_free_device_ptr(ptr)
    
    # 从设备指针创建Tensor,绑定释放逻辑
    tensor = torch.Tensor._from_data_ptr(
        dtype=dtype,
        sizes=(size,),
        strides=(1,),  # 连续内存,步长设为1
        data_ptr=data_ptr,
        device=device,
        deleter=deleter
    )
    return tensor

# 使用示例
input_tensor = torch.randn(10, device='cuda')
result = tenex.process_tensor(input_tensor.data_ptr(), input_tensor.size(0))
output_tensor = create_tensor_from_kokkos_ptr(result.data_ptr, result.size)

# 验证计算结果
assert torch.allclose(output_tensor, input_tensor * 2.0)

关键细节说明

  • 零拷贝性能保障:_from_data_ptr直接复用Kokkos分配的设备内存,无主机/设备数据传输,完全满足性能要求。
  • 自动内存管理:通过deleter参数,PyTorch Tensor会在自身被垃圾回收时自动调用Kokkos的内存释放函数,避免内存泄漏。
  • 类型与设备一致性:确保C++中分配的内存类型(此处为float)与PyTorch的dtype(torch.float32)匹配,且Tensor指定的设备与Kokkos执行空间一致。
  • 计算同步:C++端的Kokkos::fence()必须执行完毕再返回指针,保证设备计算完成后PyTorch才能读取到有效数据。

内容的提问来源于stack exchange,提问作者Taylor Childers

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 07:32:06