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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 02:31:32