Thrust CUDA中Lambda使用Tuple无法修改值,结果异常求助
问题原因及解决方法
核心问题:传值参数导致修改无效
你的Lambda函数参数t是传值传递,函数内操作的是tuple的副本,修改副本不会影响原始设备向量里的元素。另外还有个小问题:printf用了整数格式符%d打印浮点数myval,会导致输出错乱。
修复后的代码
#include <thrust/for_each.h> #include <thrust/device_vector.h> #include <thrust/iterator/zip_iterator.h> #include <iostream> #include <thrust/transform.h> #include <thrust/functional.h> int main(void) { // allocate storage thrust::device_vector<float> A(5); thrust::device_vector<float> B(5); thrust::device_vector<float> C(5); thrust::device_vector<float> D(5); // initialize input vectors A[0] = 3; B[0] = 6; C[0] = 2; A[1] = 4; B[1] = 7; C[1] = 5; A[2] = 0; B[2] = 2; C[2] = 7; A[3] = 8; B[3] = 1; C[3] = 4; A[4] = 2; B[4] = 8; C[4] = 3; auto start_zip = thrust::make_zip_iterator(thrust::make_tuple(A.begin(), B.begin(), C.begin(), D.begin())); auto end_zip =thrust::make_zip_iterator(thrust::make_tuple(A.end(), B.end(), C.end(), D.end())); thrust::for_each(thrust::device, start_zip, end_zip, [=] __device__ (thrust::tuple<float&, float&, float&, float&> t) { float myval = thrust::get<0>(t) + thrust::get<1>(t) * thrust::get<2>(t); thrust::get<3>(t) = myval; printf("Call for value : %f\n", myval); } ); // print the output for(int i = 0; i < 5; i++) std::cout << A[i] << " + " << B[i] << " * " << C[i] << " = " << D[i] << std::endl; }
关键修改点
- 将Lambda参数类型改为
thrust::tuple<float&, float&, float&, float&>,用引用传递,这样操作的就是原始设备向量的元素,修改会直接生效。 - 把
printf里的%d改成%f,匹配浮点数打印格式,同时补上换行符\n保证输出正常换行。
补充说明
Thrust的zip_iterator返回的tuple包含迭代器的引用类型,只有当Lambda参数是引用类型时,才能直接修改底层设备向量的数据。传值方式只会修改局部副本,无法影响原始数据。
内容的提问来源于stack exchange,提问作者Pablo Ninguno
相关产品推荐
相关产品推荐

