目的:感觉很多教程对于线程或者线程块和数据的对应关系讲的不是很清楚,所以记录一下分块矩阵乘的数据和线程块的对应关系便于理解。利用share_memory实现分块矩阵乘。

输入输出:输入为矩阵A(M,K),矩阵B(K,N),输出为矩阵C(M,N);

线程块和数据的对应关系:

  • 每一个thread负责计算C中的一个元素值,所以为了计算C(M,N)需要M行N列的thread。这些thread被分为很多个block,每个block为(block_size, block_size)。接下来把视角放到第一个block上,第一个block负责获取矩阵A的BLOCK*K个数据,同时负责获取矩阵B的K*BLOCK个数。
  • 但是可能由于block_size*K个数据过于庞大(K是远大于block_size的),所以第一个block的线程需要分批次读取数据。第一个批次读取矩阵A的[1:block_size,1:(1*block_size)]个数据,读取B的[1:(1*block_size),1:block_size]个数据;第二个批次读取A的[1:block_size,block_size+1:(2*block_size)],读取矩阵B的第[(block_size+1):block_size*2,1:block_size]个数据...以此类推。
  • 在第一个block读取了第一批次的数据并分别放入shareMemory_A和shareMemory_B中以后,第一个block线程的第一行第一列的thread负责计算shareMemory_A中第一行数据和shareMemory_B中第一列数据的乘累加工作,并将结果存放在第一个线程中的寄存器变量result_cache中;第一个block线程的第一行第二列的thread负责计算shareMemory_A中第一行数据和shareMemory_B中第二列数据的乘累加工作,并将结果存放在第二个线程中的寄存器变量result_cache中,以此类推。
  • 之后第一个block读取第二批数据并进行计算
  • 等到第一个block把所有的批次走完,结果也就累加到了第一个block的每一个thread的寄存器中。也就获取了C中的(block_size,block_size)个结果。

代码如下:

#define BLOCK_SIZE 16  // 每个线程块的大小

__global__ void matrixMulSharedMemory(int *A, int *B, int *C, int M, int N, int K) {
    // 线程索引
    int tx = threadIdx.x;
    int ty = threadIdx.y;

    // 线程块的起始位置
    int row = blockIdx.y * BLOCK_SIZE + ty;
    int col = blockIdx.x * BLOCK_SIZE + tx;

    // 临时变量存储 C 矩阵的计算结果
    int Cvalue = 0;

    // 每个线程块内的共享内存
    __shared__ int shared_A[BLOCK_SIZE][BLOCK_SIZE];
    __shared__ int shared_B[BLOCK_SIZE][BLOCK_SIZE];

    // 分批次加载矩阵 A 和 B 到共享内存并计算
    for (int m = 0; m < (K + BLOCK_SIZE - 1) / BLOCK_SIZE; ++m) {
        // 加载 A 和 B 到共享内存
        if (row < M && m * BLOCK_SIZE + tx < K) {
            shared_A[ty][tx] = A[row * K + m * BLOCK_SIZE + tx];
        } else {
            shared_A[ty][tx] = 0;
        }
        
        if (col < N && m * BLOCK_SIZE + ty < K) {
            shared_B[ty][tx] = B[(m * BLOCK_SIZE + ty) * N + col];
        } else {
            shared_B[ty][tx] = 0;
        }

        // 同步,确保每个线程块中的所有线程都加载完数据
        __syncthreads();

        // 执行矩阵乘法的核心计算
        for (int k = 0; k < BLOCK_SIZE; ++k) {
            Cvalue += shared_A[ty][k] * shared_B[k][tx];
        }

        // 同步,确保线程在加载下一批数据之前完成计算
        __syncthreads();
    }

    // 将计算结果存储到 C 矩阵中
    if (row < M && col < N) {
        C[row * N + col] = Cvalue;
    }
}

如果有动画版本将会一目了然,但是不知道怎么制作简易动画。

Logo

更多推荐