cuda算子--矩阵转置

复制代码
#include <cuda_runtime.h>

#define TILE 32
#define BLOCK_ROWS 8   // block 的 y 方向线程数,配合 TILE 减少 padding 浪费

// 用 padding 避免 shared bank conflict
__global__ void transpose_shared_kernel(const float* __restrict__ input,
                                        float* __restrict__ output,
                                        int rows, int cols) {
    __shared__ float tile[TILE][TILE + 1];

    const int tx = threadIdx.x;   // [0, TILE)
    const int ty = threadIdx.y;   // [0, BLOCK_ROWS)

    // ---- 读阶段 ----
    // 该 block 负责 input 的子块:
    //   行 [blockIdx.y*TILE, blockIdx.y*TILE + TILE)
    //   列 [blockIdx.x*TILE, blockIdx.x*TILE + TILE)
    const int col = blockIdx.x * TILE + tx;   // input 列
    const int row = blockIdx.y * TILE + ty;   // input 行起点

    #pragma unroll
    for (int j = 0; j < TILE; j += BLOCK_ROWS) {
        if (col < cols && (row + j) < rows)
            tile[ty + j][tx] = input[(size_t)(row + j) * cols + col];
    }

    __syncthreads();

    // ---- 写阶段 ----
    // 转置后,output 的子块:
    //   行 [blockIdx.x*TILE, blockIdx.x*TILE + TILE)  ← 原来是 input 的列
    //   列 [blockIdx.y*TILE, blockIdx.y*TILE + TILE)  ← 原来是 input 的行
    const int outCol     = blockIdx.y * TILE + tx;   // 输出列 = 原输入行
    const int outRowBase = blockIdx.x * TILE + ty;   // 输出行起点 = 原输入列

    #pragma unroll
    for (int j = 0; j < TILE; j += BLOCK_ROWS) {
        const int outRow = outRowBase + j;           // 关键:j 加到行上
        if (outCol < rows && outRow < cols)
            output[(size_t)outRow * rows + outCol] = tile[tx][ty + j];
    }
}

extern "C" void solve(const float* input, float* output, int rows, int cols) {
    if (rows <= 0 || cols <= 0) return;

    dim3 threadsPerBlock(TILE, BLOCK_ROWS);
    dim3 blocksPerGrid((cols + TILE - 1) / TILE,
                       (rows + TILE - 1) / TILE);

    matrix_transpose_kernel<<<blocksPerGrid, threadsPerBlock>>>(input, output, rows, cols);
    cudaDeviceSynchronize();
}

关键是

for (int j = 0; j < TILE; j += BLOCK_ROWS) {

const int outRow = outRowBase + j; // 关键:j 加到行上

if (outCol < rows && outRow < cols)

output(size_t)outRow \* rows + outCol = tiletxty + j;

}

如果让block都等于0的话,那么base row和col就是ty,tx,那么最后就是outputtytx=tiletxty]很简单的转置。接下来考虑block x增加,此时base row下移,但他不应该影响与tile的对应关系,因为tile只有一块,所以应该还是和block x等于0时一样,去tile的同样的位置去找。需要注意的是,为什么base row对应的是block x,而不是y,因为这样out才对应input矩阵的block级转置,而thread是对应tile级的转置。

相关推荐
xhy_070726 分钟前
AI 在同一步上反复打转怎么办?WES Code 循环检测怎么用
人工智能·大模型·ai编程·wes code
鱼宵1 小时前
Spring AI 提示词模板:{变量} 参数化 + few-shot,一条提示词反复用
java·人工智能·spring·few-shot·提示词工程·springai
蜗牛互联网1 小时前
Python消费Responses SSE事件:增量文本、超时与取消
java·开发语言·人工智能·后端·python
架构师那点事儿1 小时前
大模型如何私有化部署到生产环境
人工智能·架构·llm
HeyAI人工智能1 小时前
官网内容优化 vs 企业级 RAG:AI 搜索时代,企业到底该先做什么?
人工智能·aigc
wanderist.1 小时前
线性筛法详解:从筛质数到欧拉函数、Möbius 函数与约数函数
java·数据结构·算法
All for pursuit.1 小时前
【设计-1】208.实现Trie (前缀树)
数据结构·c++·算法·leetcode
黑妹天下第一乖2 小时前
第09讲 · 多媒体与音频 SDK:硬件编解码与端侧语音
人工智能·嵌入式硬件·深度学习·机器人·音视频·iot
weixin_307779132 小时前
C++代码实现MATLAB中的crossval函数功能
开发语言·c++·算法·matlab
青少儿编程课堂2 小时前
背包问题(0/1 背包与完全背包)解题精讲——动态规划入门
c++·python·算法·bfs·信息学竞赛