Skip to content

Instantly share code, notes, and snippets.

@aont
Created August 8, 2014 13:03
Show Gist options
  • Select an option

  • Save aont/a680bc7922cafee60a5a to your computer and use it in GitHub Desktop.

Select an option

Save aont/a680bc7922cafee60a5a to your computer and use it in GitHub Desktop.
Generating kernel from function object is helpful?
#include <thrust/device_vector.h>
#include <thrust/host_vector.h>
struct init_functor
{
float* data_device;
init_functor(float* const data_device)
: data_device(data_device)
{}
__device__ float operator()(char* const) const {
const int idx = blockDim.x * blockIdx.x + threadIdx.x;
return data_device[idx] = idx;
};
};
struct multiply_functor
{
float number;
float* data_device;
multiply_functor(const float number, float* const data_device)
: number(number), data_device(data_device)
{}
__device__ void operator()(char* const) const {
const int idx = blockDim.x * blockIdx.x + threadIdx.x;
data_device[idx] *= number;
};
__device__ void operator()(char* const, const float data_idx) const {
const int idx = blockDim.x * blockIdx.x + threadIdx.x;
data_device[idx] = data_idx * number;
};
};
template<typename F>
static __global__ void kernel
(const F f)
{
extern __shared__ char smem[];
f(smem);
};
template<typename F1, typename F2>
static __global__ void kernel
(const F1 f1, const F2 f2)
{
extern __shared__ char smem[];
f1(smem); f2(smem);
};
template<typename F1, typename F2, typename F3>
static __global__ void kernel
(const F1 f1, const F2 f2, const F3 f3)
{
extern __shared__ char smem[];
f1(smem); f2(smem); f3(smem);
};
template<typename F1, typename F2, typename F3, typename F4>
static __global__ void kernel
(const F1 f1, const F2 f2, const F3 f3, const F4 f4)
{
extern __shared__ char smem[];
f1(smem); f2(smem); f3(smem); f4(smem);
};
template<typename F1, typename F2>
static __global__ void kernel_chain
(const F1 f1, const F2 f2)
{
extern __shared__ char smem[];
f2(smem, f1(smem));
};
template<typename F1, typename F2, typename F3>
static __global__ void kernel_chain
(const F1 f1, const F2 f2, const F3 f3)
{
extern __shared__ char smem[];
f3(smem, f2(smem, f1(smem)));
};
template<typename F1, typename F2, typename F3, typename F4>
static __global__ void kernel_chain
(const F1 f1, const F2 f2, const F3 f3, const F4 f4)
{
extern __shared__ char smem[];
f4(smem, f3(smem, f2(smem, f1(smem))));
};
int main()
{
const int num_blocks = 128;
const int block_size = 128;
const int num_data = num_blocks * block_size;
thrust::device_vector<float> data_device(num_data);
thrust::host_vector<float> data_host(num_data);
// kernel<<<num_blocks, block_size>>>
// (init_functor(data_device.data().get() ) );
// kernel<<<num_blocks, block_size>>>
// (multiply_functor(2.0, data_device.data().get() ) );
// kernel<<<num_blocks, block_size>>>
// (init_functor(data_device.data().get() ),
// multiply_functor(2.0, data_device.data().get() ) );
kernel_chain<<<num_blocks, block_size>>>
(init_functor(data_device.data().get() ),
multiply_functor(2.0, data_device.data().get() ) );
data_host = data_device;
for(int i=0; i<num_data; ++i) {
printf("%d %g\n", i, data_host[i]);
}
return 0;
}
@aont

aont commented Aug 8, 2014

Copy link
Copy Markdown
Author

good points

  • easy to compose and customize kernel
  • argument list for each functor

bad points

  • reuse of local variable is not implemented ( though somehow possible )
  • increased byte size of kernel arguments ( <= 256 bytes )

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment