在CUDA类中使用Thrust处理向量的实现问题求助
在C++类中使用Thrust实现设备端三角形计算的修正方案
原代码的核心问题
- 类成员初始化错误:C++类无法直接在成员声明时使用带参数的构造函数(如
thrust::device_vector<double> v1(3,0)),必须通过构造函数的初始化列表或内部赋值完成。 - device_vector的使用限制:
thrust::device_vector是主机端容器,用于管理设备内存,但它的大部分操作(如赋值、运算符重载)只能在主机端执行,无法在__device__函数中直接使用。 - 变量与函数问题:
- 法向量
normal被声明为double,但实际是三维向量,类型错误; - 未实现设备端可用的叉乘(
cross)和向量模长(norm)函数; - 构造函数中使用类名
triangle访问成员(如triangle.v1)是语法错误,应使用this->或直接访问成员; - 变量拼写/大小写不一致(如
Areavsarea,Normalvsnormal)。
- 法向量
修正后的实现方案
要实现全设备端计算,我们需要:
- 使用原始数组或
thrust::device_ptr存储设备端数据,避免直接在设备端操作device_vector; - 实现设备端可调用的向量运算函数;
- 主机端负责数据拷贝与设备端计算的触发。
完整代码
#include <thrust/host_vector.h> #include <thrust/device_vector.h> #include <cmath> // 设备端向量运算函数 __device__ void cross(const double* a, const double* b, double* result) { result[0] = a[1]*b[2] - a[2]*b[1]; result[1] = a[2]*b[0] - a[0]*b[2]; result[2] = a[0]*b[1] - a[1]*b[0]; } __device__ double norm(const double* vec) { return sqrt(vec[0]*vec[0] + vec[1]*vec[1] + vec[2]*vec[2]); } // 三角形类:主机端管理设备内存,设备端执行计算 class Triangle { private: double* d_v1; // 设备端顶点1指针 double* d_v2; // 设备端顶点2指针 double* d_v3; // 设备端顶点3指针 double* d_E1; // 边向量v2-v1 double* d_E2; // 边向量v3-v1 double* d_normal;// 法向量 double d_area; // 面积 // 设备端计算函数 __device__ void compute() { // 计算边向量 d_E1[0] = d_v2[0] - d_v1[0]; d_E1[1] = d_v2[1] - d_v1[1]; d_E1[2] = d_v2[2] - d_v1[2]; d_E2[0] = d_v3[0] - d_v1[0]; d_E2[1] = d_v3[1] - d_v1[1]; d_E2[2] = d_v3[2] - d_v1[2]; // 计算叉乘得到法向量 cross(d_E2, d_E1, d_normal); // 计算面积 double n_norm = norm(d_normal); d_area = n_norm / 2.0; // 归一化法向量 if (n_norm > 1e-8) { d_normal[0] /= n_norm; d_normal[1] /= n_norm; d_normal[2] /= n_norm; } } public: // 主机端构造函数:从host_vector初始化设备内存 Triangle(const thrust::host_vector<double>& ver1, const thrust::host_vector<double>& ver2, const thrust::host_vector<double>& ver3) { // 分配设备内存 cudaMalloc(&d_v1, 3 * sizeof(double)); cudaMalloc(&d_v2, 3 * sizeof(double)); cudaMalloc(&d_v3, 3 * sizeof(double)); cudaMalloc(&d_E1, 3 * sizeof(double)); cudaMalloc(&d_E2, 3 * sizeof(double)); cudaMalloc(&d_normal, 3 * sizeof(double)); // 拷贝主机数据到设备 cudaMemcpy(d_v1, ver1.data(), 3*sizeof(double), cudaMemcpyHostToDevice); cudaMemcpy(d_v2, ver2.data(), 3*sizeof(double), cudaMemcpyHostToDevice); cudaMemcpy(d_v3, ver3.data(), 3*sizeof(double), cudaMemcpyHostToDevice); // 启动设备端计算(使用单个线程执行,适合单个三角形) compute<<<1,1>>>(); cudaDeviceSynchronize(); } // 主机端析构函数:释放设备内存 ~Triangle() { cudaFree(d_v1); cudaFree(d_v2); cudaFree(d_v3); cudaFree(d_E1); cudaFree(d_E2); cudaFree(d_normal); } // 主机端获取计算结果的函数 double getArea() const { double area; cudaMemcpy(&area, &d_area, sizeof(double), cudaMemcpyDeviceToHost); return area; } thrust::host_vector<double> getNormal() const { thrust::host_vector<double> normal(3); cudaMemcpy(normal.data(), d_normal, 3*sizeof(double), cudaMemcpyDeviceToHost); return normal; } // 包装compute函数为可调用的核函数 static __global__ void compute(double* v1, double* v2, double* v3, double* E1, double* E2, double* normal, double* area) { Triangle tri; tri.d_v1 = v1; tri.d_v2 = v2; tri.d_v3 = v3; tri.d_E1 = E1; tri.d_E2 = E2; tri.d_normal = normal; tri.d_area = *area; tri.compute(); *area = tri.d_area; } }; // 主函数示例 int main() { // 模拟从文件读取的顶点数据 thrust::host_vector<double> dum(9); dum[0] = 0.0; dum[1] = 0.0; dum[2] = 0.0; dum[3] = 1.0; dum[4] = 0.0; dum[5] = 0.0; dum[6] = 0.0; dum[7] = 1.0; dum[8] = 0.0; thrust::host_vector<double> ver1(dum.begin(), dum.begin()+3); thrust::host_vector<double> ver2(dum.begin()+3, dum.begin()+6); thrust::host_vector<double> ver3(dum.begin()+6, dum.end()); // 创建三角形对象,自动完成设备端计算 Triangle tri(ver1, ver2, ver3); // 获取结果 double area = tri.getArea(); thrust::host_vector<double> normal = tri.getNormal(); // 输出结果 printf("Triangle Area: %.4f\n", area); printf("Normal Vector: (%.4f, %.4f, %.4f)\n", normal[0], normal[1], normal[2]); return 0; }
关键改动说明
- 设备端数据存储:改用原始指针
double*存储设备端数据,避免device_vector在设备端的使用限制。 - 独立的设备端运算函数:实现了
cross和norm的设备端版本,确保计算完全在GPU上执行。 - 核函数触发计算:通过
compute<<<1,1>>>启动单个线程执行设备端计算,适合单个三角形的场景;如果需要批量处理多个三角形,可以修改为多线程核函数。 - 主机端管理内存:构造函数负责分配设备内存并拷贝数据,析构函数释放内存,主机端通过
getArea和getNormal获取计算结果。 - 语法修正:修正了成员访问、变量拼写、类型匹配等基础语法错误。
内容的提问来源于stack exchange,提问作者Seb Seb
相关产品推荐
相关产品推荐

