CUDA 算子笔记 3

§ 参考资料 共 3 条

Stream Compaction

给一个包含正负数和 0 的数组,把所有的正数按照原来的顺序放到数组开头并输出

例子:

Input:  A = [1.0, -2.0, 3.0, 0.0, -1.0, 4.0]
Output: out = [1.0, 3.0, 4.0, 0.0, 0.0, 0.0]

这是 Radix Sort 的一个前置题目,用计数排序的思想,包含三个步骤:

  • Predicate:判断每个数是否是正数,得到 predicate[]
  • Scan:对 predicate[] 计算 Exclusive 前缀和,得到元素的新位置
  • Scatter:计算出来的新位置,把元素分散写回输出数组

最直接的方法

直接三个 Kernel 即可:

#include <cuda_runtime.h>

#define CEIL(x, y) (((x) + (y) - 1) / (y))
#define FLOAT4(x) (*(reinterpret_cast<float4*>(&(x))))
#define CONST_FLOAT4(x) (*(reinterpret_cast<const float4*>(&(x))))
#define INT4(x) (*(reinterpret_cast<int4*>(&(x))))
#define CONST_INT4(x) (*(reinterpret_cast<const int4*>(&(x))))

__device__ __forceinline__ int warp_inclusive_scan(int val) {
  const int lane = threadIdx.x & 31;
#pragma unroll
  for (int offset = 1; offset < 32; offset <<= 1) {
    const int other = __shfl_up_sync(0xffffffff, val, offset);
    if (lane >= offset) {
      val += other;
    }
  }
  return val;
}

__device__ __forceinline__ int warp_exclusive_scan(int val) {
  return warp_inclusive_scan(val) - val;
}

__device__ __forceinline__ int block_inclusive_scan(int val) {
  __shared__ int warp_offsets[32];

  const int warp = threadIdx.x >> 5;
  const int lane = threadIdx.x & 31;
  const int num_warps = blockDim.x >> 5;

  const int prefix = warp_inclusive_scan(val);

  if (lane == 31) {
    warp_offsets[warp] = prefix;
  }

  __syncthreads();

  if (warp == 0) {
    const int warp_sum = lane < num_warps ? warp_offsets[lane] : 0;
    const int warp_offset = warp_exclusive_scan(warp_sum);
    if (lane < num_warps) {
      warp_offsets[lane] = warp_offset;
    }
  }

  __syncthreads();

  return prefix + warp_offsets[warp];
}

__device__ __forceinline__ int block_exclusive_scan(int val) {
  return block_inclusive_scan(val) - val;
}

__device__ int g_block_offsets[20000];

template <int k_items_per_thread>
__global__ void stream_compaction_block_part1_kernel(
    const float* __restrict__ A, int N) {
  const int tid = threadIdx.x + blockIdx.x * blockDim.x;
  const int thread_idx = tid * k_items_per_thread;

  int thread_sum = 0;

  if (thread_idx + k_items_per_thread <= N) {
#pragma unroll
    for (int i = 0; i < k_items_per_thread; i += 4) {
      const float4 val = CONST_FLOAT4(A[thread_idx + i]);
      thread_sum += (val.x > 0) + (val.y > 0) + (val.z > 0) + (val.w > 0);
    }
  } else {
#pragma unroll
    for (int i = 0; i < k_items_per_thread; ++i) {
      const int idx = thread_idx + i;
      if (idx < N) {
        thread_sum += A[idx] > 0;
      }
    }
  }

  const int block_prefix = block_inclusive_scan(thread_sum);

  if (threadIdx.x == blockDim.x - 1) {
    g_block_offsets[blockIdx.x] = block_prefix;
  }
}

__global__ void stream_compaction_grid_kernel(int M) {
  const int items_per_thread = CEIL(M, blockDim.x);
  const int thread_idx = threadIdx.x * items_per_thread;

  int thread_offsets[128];
  int thread_sum = 0;

  for (int i = 0; i < items_per_thread; ++i) {
    const int idx = thread_idx + i;
    const int val = idx < M ? g_block_offsets[idx] : 0;
    thread_offsets[i] = thread_sum;
    thread_sum += val;
  }

  const int offset = block_exclusive_scan(thread_sum);

  for (int i = 0; i < items_per_thread; ++i) {
    const int idx = thread_idx + i;
    if (idx < M) {
      g_block_offsets[idx] = thread_offsets[i] + offset;
    }
  }
}

template <int k_items_per_thread>
__global__ void stream_compaction_block_part2_kernel(
    const float* __restrict__ A, float* __restrict__ out, int N) {
  const int tid = threadIdx.x + blockIdx.x * blockDim.x;
  const int thread_idx = tid * k_items_per_thread;

  float thread_inputs[k_items_per_thread];

  int thread_sum = 0;

  if (thread_idx + k_items_per_thread <= N) {
#pragma unroll
    for (int i = 0; i < k_items_per_thread; i += 4) {
      const float4 val = CONST_FLOAT4(A[thread_idx + i]);
      FLOAT4(thread_inputs[i]) = val;
      thread_sum += (val.x > 0) + (val.y > 0) + (val.z > 0) + (val.w > 0);
    }
  } else {
#pragma unroll
    for (int i = 0; i < k_items_per_thread; ++i) {
      const int idx = thread_idx + i;
      const float val = idx < N ? A[idx] : 0.0f;
      thread_inputs[i] = val;
      if (idx < N) {
        thread_sum += val > 0;
      }
    }
  }

  const int thread_offset = block_exclusive_scan(thread_sum);
  int output_pos = g_block_offsets[blockIdx.x] + thread_offset;

#pragma unroll
  for (int i = 0; i < k_items_per_thread; ++i) {
    const int idx = thread_idx + i;
    const float val = thread_inputs[i];
    if (idx < N && val > 0) {
      out[output_pos++] = val;
    }
  }
}

// A, out are device pointers
extern "C" void solve(const float* A, int N, float* out) {
  // 512 / 32 = 16
  constexpr int block_dim = 512;
  constexpr int k_items_per_thread = 12;

  // 100,000,000 / (512 * 12) = 16276
  const int grid_dim = CEIL(N, block_dim * k_items_per_thread);

  stream_compaction_block_part1_kernel<k_items_per_thread>
      <<<grid_dim, block_dim>>>(A, N);
  stream_compaction_grid_kernel<<<1, block_dim>>>(grid_dim);
  stream_compaction_block_part2_kernel<k_items_per_thread>
      <<<grid_dim, block_dim>>>(A, out, N);
}

基于 Decoupled Look-back Scan(DLBS)的方法

Radix Sort

输入一个 32 位无符号整数数组,实现从小到大的基数排序并输出

例子:

Input:  [170, 45, 75, 90, 2, 802, 24, 66]
Output: [2, 24, 45, 66, 75, 90, 170, 802]

基数排序将待排序的元素拆分为 $k$k 个关键字,逐一对各个关键字排序后完成对所有元素的排序

分为从第一关键字到最后一个关键字(Most Significant Digit first,MSD)和最后一个关键字到第一个关键字(Least Significant Digit first,LSD)两种遍历顺序,对每一层内部排序的时候,一般使用一种叫做计数排序的稳定排序(并不改变不关心元素的顺序,LSD 的正确性来源于此)来完成

MSD 通常只能使用递归的形式完成,时间常数比较大,且不适用于 GPU,所以一般使用 LSD:

一个 LSD 基数排序全流程的例子

CPU 代码长这样:

#include <algorithm>
#include <iostream>
#include <utility>

void radix_sort(int n, int a[]) {
  int *b = new int[n];  // 临时空间
  int *cnt = new int[1 << 8];
  int mask = (1 << 8) - 1;
  int *x = a, *y = b;
  for (int i = 0; i < 32; i += 8) {
    for (int j = 0; j != (1 << 8); ++j) cnt[j] = 0;
    for (int j = 0; j != n; ++j) ++cnt[x[j] >> i & mask];
    for (int sum = 0, j = 0; j != (1 << 8); ++j) {
      // 等价于 std::exclusive_scan(cnt, cnt + (1 << 8), cnt, 0);
      sum += cnt[j], cnt[j] = sum - cnt[j];
    }
    for (int j = 0; j != n; ++j) y[cnt[x[j] >> i & mask]++] = x[j];
    std::swap(x, y);
  }
  delete[] cnt;
  delete[] b;
}

int main() {
  std::ios::sync_with_stdio(false);
  std::cin.tie(nullptr);
  int n;
  std::cin >> n;
  int *a = new int[n];
  for (int i = 0; i < n; ++i) std::cin >> a[i];
  radix_sort(n, a);
  for (int i = 0; i < n; ++i) std::cout << a[i] << ' ';
  delete[] a;
  return 0;
}

可以发现,它的 $\log$ 次循环内部包含这三个操作:

  • Predicate:判断每个数的当前二进制位置是 0 还是 1,得到 predicate[](0 的话 predicate[i]=1,否则设置为 0)
  • Scan:对 predicate[] 计算 Exclusive 前缀和,得到元素的新位置
  • Scatter:计算出来的新位置,把元素分散写回输出数组