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

如何在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;
}

关键说明

  1. DeviceVector模板类封装了设备端内存的分配、释放和数据拷贝,使用起来和std::vector类似,同时保证内存安全(析构时自动释放设备内存)
  2. 核函数addVectors被修改为模板函数,支持任意实现了+运算符的类型,比如int、double都可以直接替换使用
  3. CUDA 9.0完全支持C++11的模板特性,上述代码可以正常编译运行

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 18:24:52