Count 2D Array Element
2026/6/6小于 1 分钟
Count 2D Array Element
题目描述
统计 的 32 位整数二维数组中值为 的元素数量。
实现要求
- 不允许使用外部库。
solve函数签名必须保持不变。- 最终结果必须存储在
output变量中。
示例
Input: input = [[1,2,3],[4,5,1]], k = 1
Output: 2约束条件
- ,。
解题思路
与一维版本相同:逐元素比较 + 规约求和。将二维索引 线性化为一维索引 ,每个线程比较一个元素,分块做 warp shuffle 规约求和。属于内存带宽受限内核。
代码实现
CUDA
#include <cuda_runtime.h>
__global__ void count2d(const int* input, int* output, int total, int k) {
__shared__ int bc; if(threadIdx.x==0)bc=0; __syncthreads();
for(int i=blockIdx.x*blockDim.x+threadIdx.x; i<total; i+=gridDim.x*blockDim.x)
if(input[i]==k) atomicAdd(&bc,1);
__syncthreads(); if(threadIdx.x==0)atomicAdd(output,bc);
}
extern "C" void solve(const int* input, int* output, int N, int M, int k) {
int t=N*M; count2d<<<min((t+255)/256,1024),256>>>(input,output,t,k);
cudaDeviceSynchronize();
}Triton
import triton, triton.language as tl
@triton.jit
def count2d(input_ptr,output_ptr,total, k, BLOCK:tl.constexpr):
idx=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK); mask=idx<total
x=tl.load(input_ptr+idx,mask=mask)
tl.atomic_add(output_ptr, tl.sum(tl.where(x==k,1,0),axis=0))