如何在cuFFT函数中正确使用thrust::device_vector复杂类型?
问题:Thrust迭代器转换为cuFFT可用指针失败
我需要用cuFFT执行复数到复数的FFT运算,为简化代码采用Thrust库的thrust::complex类型。为贴近底层数学逻辑,我在host_vector中存储指向device_vector不同位置的迭代器,但尝试将这些迭代器转换为cufftDoubleComplex*用于cufftExecZ2Z调用时,出现编译错误。希望找到可行的转换方式,避免创建n个独立向量导致代码混乱。
精简后的报错代码
#include <stdlib.h> #include <cuda.h> #include <cufft.h> #include <thrust/host_vector.h> #include <thrust/device_vector.h> #include <thrust/complex.h> #define cuFFTFORWARD -1 #define cuFFTINVERSE 1 using namespace std; int main () { int m[3], M, r, n; cufftHandle pr_to_pk; n = 90; m[0] = 16; m[1] = m[0]; m[2] = m[0]; M = m[0]*m[1]*m[2]; // allocate memory for the propagators thrust::device_vector<thrust::complex<double>> pr(2*(n+1)*M); // fill pr with a number sequence (dummy data) thrust::sequence(pr.begin(), pr.end()); // allocate memory to store iterators pointing to device vector elements thrust::host_vector<thrust::device_vector<thrust::complex<double>>::iterator> p1(n+1); thrust::host_vector<thrust::device_vector<thrust::complex<double>>::iterator> p2(n+1); // save interators pointing to start of memory for p1 and p2 for (r=0; r<=n; r++) { p1[r] = pr.begin()+2*r*M; p2[(n+1-r)%(n+1)] = pr.begin()+(2*r+1)*M; } // set up the cufft plan cufftPlanMany(&pr_to_pk,3,m,NULL,1,0,NULL,1,0, CUFFT_Z2Z,2); // allocate memory for q(k) thrust::device_vector<thrust::complex<double>> pk(2*M); // attempt to cast to cufftDoubleComplex cufftDoubleComplex* _V1 = (cufftDoubleComplex*)thrust::raw_pointer_cast(p1[1]); cufftDoubleComplex* _V2 = (cufftDoubleComplex*)thrust::raw_pointer_cast(pk.data()); // attempt the first cufft cufftExecZ2Z(pr_to_pk, _V1, _V2, cuFFTFORWARD); cout << "complete" << endl; }
编译错误信息
error: class "thrust::detail::pointer_raw_pointer<thrust::detail::normal_iterator<thrust::device_ptr<thrust::complex
>>>" has no member "type"
detected during instantiation of class "thrust::detail::pointer_traits [with Ptr=thrust::detail::normal_iterator<thrust::device_ptr<thrust::complex>>]"
解决方案
核心问题原因
thrust::raw_pointer_cast只能直接处理thrust::device_ptr类型,无法直接作用于Thrust迭代器(thrust::detail::normal_iterator)。需要先将迭代器转换为device_ptr,再进行指针转换。
修改关键代码
将原来的指针转换代码:
cufftDoubleComplex* _V1 = (cufftDoubleComplex*)thrust::raw_pointer_cast(p1[1]);
替换为:
// 先将迭代器转为device_ptr,再转成cuFFT可用的原始指针 cufftDoubleComplex* _V1 = reinterpret_cast<cufftDoubleComplex*>( thrust::raw_pointer_cast(&*p1[1]) );
完整修正后的代码
#include <stdlib.h> #include <cuda.h> #include <cufft.h> #include <thrust/host_vector.h> #include <thrust/device_vector.h> #include <thrust/complex.h> #include <iostream> #define cuFFTFORWARD -1 #define cuFFTINVERSE 1 using namespace std; int main () { int m[3], M, r, n; cufftHandle pr_to_pk; n = 90; m[0] = 16; m[1] = m[0]; m[2] = m[0]; M = m[0]*m[1]*m[2]; // allocate memory for the propagators thrust::device_vector<thrust::complex<double>> pr(2*(n+1)*M); // fill pr with a number sequence (dummy data) thrust::sequence(pr.begin(), pr.end()); // allocate memory to store iterators pointing to device vector elements thrust::host_vector<thrust::device_vector<thrust::complex<double>>::iterator> p1(n+1); thrust::host_vector<thrust::device_vector<thrust::complex<double>>::iterator> p2(n+1); // save iterators pointing to start of memory for p1 and p2 for (r=0; r<=n; r++) { p1[r] = pr.begin()+2*r*M; p2[(n+1-r)%(n+1)] = pr.begin()+(2*r+1)*M; } // set up the cufft plan cufftPlanMany(&pr_to_pk,3,m,NULL,1,0,NULL,1,0, CUFFT_Z2Z,2); // allocate memory for q(k) thrust::device_vector<thrust::complex<double>> pk(2*M); // 正确转换迭代器为cuFFT可用的原始指针 cufftDoubleComplex* _V1 = reinterpret_cast<cufftDoubleComplex*>( thrust::raw_pointer_cast(&*p1[1]) ); cufftDoubleComplex* _V2 = reinterpret_cast<cufftDoubleComplex*>( thrust::raw_pointer_cast(pk.data()) ); // 执行FFT cufftExecZ2Z(pr_to_pk, _V1, _V2, cuFFTFORWARD); cout << "complete" << endl; // 释放cuFFT计划资源 cufftDestroy(pr_to_pk); return 0; }
关键说明
&*p1[1]:对迭代器解引用后取地址,得到thrust::device_ptr<thrust::complex<double>>,这是raw_pointer_cast支持的类型。reinterpret_cast:thrust::complex<double>与cufftDoubleComplex内存布局完全兼容(都是连续的双精度实部+虚部),因此可以安全转换指针类型。- 补充
#include <iostream>:原代码使用cout但未包含对应头文件,修正后避免编译错误。 - 添加
cufftDestroy(pr_to_pk):规范释放cuFFT计划资源,防止内存泄漏。
内容的提问来源于stack exchange,提问作者codephys
相关产品推荐
相关产品推荐

