如何在CUDA中使用泛型类型?CUDA 9.0替代GPU原始指针
CUDA泛型类型使用与设备端数组替代方案
1. CUDA中使用泛型类型的方法
CUDA完全兼容C++模板机制,实现泛型核心靠模板参数化:
- 给核函数添加模板参数,让它支持任意符合运算要求的类型(比如支持
+运算符的类型) - 模板核函数的定义要放在头文件中,或者在调用代码前显式实例化,避免编译时的链接错误
- 自定义类型要保证能在CUDA设备端正常使用:无复杂主机端依赖、支持默认构造/拷贝构造、运算符重载符合设备端要求
2. CUDA 9.0中用泛型封装替代设备端原始指针
不能直接用std::vector作为设备端数组——std::vector的内存分配在主机端,设备端线程无法直接访问。但我们可以实现一个泛型的设备端内存管理类,封装cudaMalloc/cudaMemcpy/cudaFree等操作,模拟类似vector的泛型使用体验,替代代码中的float* d_a、d_b、d_c。
修改后的完整代码
#include <vector> #include <iostream> // 泛型设备端内存管理类,模拟vector的泛型行为 template<typename T> class DeviceVector { public: DeviceVector(size_t size) : m_size(size), m_data(nullptr) { cudaMalloc(&m_data, size * sizeof(T)); } ~DeviceVector() { if (m_data != nullptr) { cudaFree(m_data); } } // 禁止拷贝构造和赋值,避免内存重复释放 DeviceVector(const DeviceVector&) = delete; DeviceVector& operator=(const DeviceVector&) = delete; // 移动构造和赋值(可选) DeviceVector(DeviceVector&& other) noexcept : m_size(other.m_size), m_data(other.m_data) { other.m_data = nullptr; other.m_size = 0; } DeviceVector& operator=(DeviceVector&& other) noexcept { if (this != &other) { if (m_data != nullptr) { cudaFree(m_data); } m_size = other.m_size; m_data = other.m_data; other.m_data = nullptr; other.m_size = 0; } return *this; } // 获取设备端指针 T* data() { return m_data; } const T* data() const { return m_data; } size_t size() const { return m_size; } // 从主机vector拷贝数据到设备 void copyFromHost(const std::vector<T>& host_vec) { if (host_vec.size() != m_size) { std::cerr << "Size mismatch when copying to device!" << std::endl; return; } cudaMemcpy(m_data, host_vec.data(), m_size * sizeof(T), cudaMemcpyHostToDevice); } // 从设备拷贝数据到主机vector void copyToHost(std::vector<T>& host_vec) const { if (host_vec.size() != m_size) { std::cerr << "Size mismatch when copying to host!" << std::endl; return; } cudaMemcpy(host_vec.data(), m_data, m_size * sizeof(T), cudaMemcpyDeviceToHost); } private: size_t m_size; T* m_data; }; // 泛型核函数,支持任意可相加的类型 template<typename T> __global__ void addVectors(T *a, T *b, T *c, int n) { int tid = blockIdx.x * blockDim.x + threadIdx.x; if (tid < n) { c[tid] = a[tid] + b[tid]; } } int main() { int n = 1024; // 主机端依然用std::vector std::vector<float> a(n), b(n), c(n); // 初始化主机向量 for (int i = 0; i < n; i++) { a[i] = i; b[i] = i * i; } // 使用泛型DeviceVector替代原始设备指针 DeviceVector<float> d_a(n); DeviceVector<float> d_b(n); DeviceVector<float> d_c(n); // 拷贝数据到设备 d_a.copyFromHost(a); d_b.copyFromHost(b); // 启动泛型核函数,显式指定模板类型(也可让编译器推导) addVectors<float><<<1, 1024>>>(d_a.data(), d_b.data(), d_c.data(), n); // 拷贝结果回主机 d_c.copyToHost(c); // 打印前10个结果 for (int i = 0; i < 10; i++) { std::cout << c[i] << '\n'; } return 0; }
关键说明
DeviceVector模板类封装了设备端内存的分配、释放和数据拷贝,使用起来和std::vector类似,同时保证内存安全(析构时自动释放设备内存)- 核函数
addVectors被修改为模板函数,支持任意实现了+运算符的类型,比如int、double都可以直接替换使用 - CUDA 9.0完全支持C++11的模板特性,上述代码可以正常编译运行
内容的提问来源于stack exchange,提问作者user366312
相关产品推荐
相关产品推荐

