From f3b4cf5db6f5c6c2514c1a3bed0a7d80189953ba Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Mon, 3 Aug 2026 14:55:28 +0000 Subject: [PATCH 1/9] feat: CUDA, RMSNorm, CPU Version --- src/kernels.cu | 84 ++++++++++++++++++++++++++++++++++--------------- src/kernels.o | Bin 0 -> 16536 bytes 2 files changed, 59 insertions(+), 25 deletions(-) create mode 100644 src/kernels.o diff --git a/src/kernels.cu b/src/kernels.cu index 2cc53e7e..60e5f3eb 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -1,5 +1,6 @@ -#include +#include #include +#include #include "../tester/utils.h" @@ -22,33 +23,60 @@ * @param[in] eps Numerical stability epsilon. */ template -void rmsNorm(const std::vector& h_input, const std::vector& h_weight, - std::vector& h_output, size_t rows, size_t hidden_dim, - float eps) { +void rmsNorm(const std::vector &h_input, const std::vector &h_weight, + std::vector &h_output, size_t rows, size_t hidden_dim, + float eps) { // TODO: Implement the rmsNorm function + + // 遍历 rows(batch_size*seq_len) + for (size_t i = 0; i < rows; i++) { + float square_sum = 0.0f; + // 1. 在第 i 行内,先计算所有元素的平方和 (还原你的 mean([i, :]^2) 逻辑) + for (size_t j = 0; j < hidden_dim; j++) { + float val = static_cast(h_input[i * hidden_dim + j]); + square_sum += val * val; // 自乘代替 ^2 + } + // 2. 计算均方根的倒数 (还原你的 rsqrt(mean + eps) 逻辑) + float mean_square = square_sum / hidden_dim; + float rsqrt_val = 1.0f / std::sqrt(mean_square + eps); + // 内层循环:更新当前行的每一个元素,应用缩放 + for (size_t j = 0; j < hidden_dim; j++) { + size_t idx = i * hidden_dim + j; + // 注意这里是 h_weight[j];统一转 float 计算,避免 half 重载歧义 + h_output[idx] = + static_cast(static_cast(h_input[idx]) * rsqrt_val * + static_cast(h_weight[j])); + } + } } /** * @brief Computes flash attention for given query, key, and value tensors. - * + * * @tparam T Data type (float) for input/output tensors - * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] - * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] + * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, + * head_dim] + * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, + * query_heads, head_dim] * @param[in] batch_size Batch dimension size * @param[in] target_seq_len Target sequence length - * @param[in] src_seq_len Source sequence length + * @param[in] src_seq_len Source sequence length * @param[in] query_heads Number of query attention heads - * @param[in] kv_heads Number of key/value heads (supports grouped query attention) + * @param[in] kv_heads Number of key/value heads (supports grouped query + * attention) * @param[in] head_dim Dimension size of each attention head * @param[in] is_causal Whether to apply causal masking */ template -void flashAttention(const std::vector& h_q, const std::vector& h_k, - const std::vector& h_v, std::vector& h_o, - int batch_size, int target_seq_len, int src_seq_len, - int query_heads, int kv_heads, int head_dim, bool is_causal) { +void flashAttention(const std::vector &h_q, const std::vector &h_k, + const std::vector &h_v, std::vector &h_o, + int batch_size, int target_seq_len, int src_seq_len, + int query_heads, int kv_heads, int head_dim, + bool is_causal) { // TODO: Implement the flash attention function } @@ -56,13 +84,19 @@ void flashAttention(const std::vector& h_q, const std::vector& h_k, // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION // ********************************************************************* -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, + int, int, bool); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, int, + int, bool); diff --git a/src/kernels.o b/src/kernels.o new file mode 100644 index 0000000000000000000000000000000000000000..a9eed28015ed1e866419eb3c00a429609a3725a4 GIT binary patch literal 16536 zcmcIrdvILUdB3Z*v9W=b2??ecyfrq67#6)BU_ivaW<5xbVP#k$o0*Q?Oq!9M8LYHvY97;gCjLW{l0jevMD1Zr>!SX? zbHDH2vqyKmm}$?<)pySCd!F|>_pbKYZ5t~>Axl!oy3g`Uf?C#r8b9BvvaQxS>mDV0 zEy`cQLwaAv^A$WR@Z5ywW;`qL5bt`X%Rz78{2zhdh$o8YkMXR+a|<4lUyWxC=WhkQ zjq_DZZ)aM~ltLuETBdcN=OQ=%S7*HN`*>mRh4{FAAzoMr)Y`w}7anJ9CR(_VF4!|Y z@$m&~$auV9U#OXg9JOZve**eRe7y5QPkillitQ|%CL6!>O%xl~Y`?6$#S15=PtDEE z#mBuO87K~-(hu>ZX@6{HvCnmFsJfh-UIDZFwMfQhD&xLW^1?L+`XdSfRin!Iq0GJI z=apYbu2Yh37_bg2T5DCFi!8G*j@uX0k?aIVs@k2rRnQ-#1un4`Gp%GAW!lPgl1i0gXsFddzCyWw}F-mJ(qQ(qm>$%*(bRE-sncl~ABhxt3hnQ|*`UulUnQmkH zIMa}tbEEg`?MhutZF_#)KJT?{;RFrabfJ>^!IP!|R~ViJYc_JUk|l^}V-MpR7N?C- z4Wa4CQR{OKY#L5HxS@n$OyU3X7ivb0&9qXqXCQD2!Au&C#_cl&durUCT7$u~uV7Cj zdwjZbjWudns#2#4l`rg_25;ZUT0}K@QZ)@A%NpG`ycE!>aclHf&!1g*Y;GTP?04V) zSBQ|#NBbTexyjQRbw#h9d(M^mTb7#J*E~{%>K>o^l5#i+hXYiDW>;anbNcuPS6_Jw zmd}nTzgFF@>&PrP1!`3UK+J7qWcqq65y!p>^|kg)!JY+%n(c$Y`4`&HBWCMZR554A zJI@!6zx_Kh#8xWNPvgUcepr}Q4qr&UJ2z)pSx+J}*&6Jt}MZ%;T&sWdhbB4L=p zW6#73|4vI5S*Jzt(h;h)^NbN9)Xt=jPeK#K8wP=*3mv0o#weh8AbrxFLCaZZo1tb6$CS(+0G|PfcE*furA~P> zH9Iyi6B^rl-g2g2NX;rE3&!l3A?nSkm4y@I_RR6A>#sbu)Sg8zqbo38&qXTjDfAqz z_U>ljt@hj{(R5+&hembLXdg<8p&|bh^}{p4ZJ3UH6RmKbOrH-n?ak}4nVpzUlbB8u z@j~aZc%l7BTfsg|qu}-UNqY)awBGN*5|@sAG(qSopsJ(JJ_yrO@yL(u!*Mm*r{f`7 zI$+D|OhkYs1XGqm!Jlm~QQ3a#T;#FiuB|b9Is~2R_&xUF+>env9RK#j?xo6EJo43v zvp10KH$TA8iC=j_;1e()mY$r1rS^Bo0scH@zfP6+Mr5Q4>`JwPor^rFoSsGn)Nnq# z*Qmg3JcM<{=Itt=BHQfK6!syEPg)7lVnsELRDr#1$05k34(W1JHVlI!5*|Svb zS*mvA#bMg~#%xrP8uT9fNaV#o{S3Qo@isezeuA~(Bj_izc>5vr(*X+S6?1Hyv`?#U z`xC9(CZ#jaw!_P5>b3)Jw@ta-rU&YcEH#d#ZQag8d%B%BsfTy;>ro4jS`8?k0jc^zM-?tPHfs>SzG?l`O}b9 zX*KhKslp0{S|QgL#vd&$!U0LnCIm0jczSl9d!Dab#9)Q>bJ}pdbb$|OyWL}wRll$v z%VjgpZLQW7bO?*(`|?iCNp>M8bkBfWRy2|C%l7QHV%d>cCYv5qneJqNe{U{1)Tc^A zxj|>p*)x<hd#6%gVza{4isXZw!)T$624_f@{twfO zpm|?<*&1tMqhI`DIqD9h%ax&4Hn)SR5bi!HT1$}&(*GVyl+&mC8xIZ7rpmCRIHM#( zein0??sAm)&Nbwg{J(M5k1`}N@^JU;P(Pp!^Yr^zztT6Ua^hd;6+o#MWqcdro?&qy zdSw2K>4idDpxerhS~$N5H7c5?Q5Il=DX&tTsgQL`?z$^%&8}J!8%gEzm@LueSksc2 zlg#y|oM?U^(NMc2mh9=trSkbiKJ#oU+E~8?N%a4Ieb+{ADuIgVyu}PlUVI!&HrD3j zzpKDft=$}C9lK1K&|At3ju_gV+r6cDzC`B1_3Q7BR&{m_XPx0_O#Nv9vH8YobW5YrwkB( z9rHgf;?H>e?aY6oi0?KK=|9i>!$c*%@T;ExgUtU!5&vC}|1;*(`y5Yyr(eG_W%w7F ze}k^S#n)d!ixyHT`|Frb@2Wiioxc8d=6_Y^Cw>0&%zsknSGWeKeGW3eWdE1ppJDzj zy8Xv}{}(U8CzYT?O}4=LSnY2nKWtrNLaH1NICWHm6Vh<*3agEAS!adjfLCbt)tpq; zX)a&OIiDOtmgGfF%|T{|Et&Pc{7;b$X}I|PE#p!adXaHjTPO*B6DmS>WPcOJ+~YKLh9zm4%(1HYH?3kKf7IE<;Z znE4sT7aRCK#w!i{yNoY4@EN90-%&K)ynK9dLoyGex!`52P`Q#$x@n+Qrv|QM_Lo2!2 zML6(>G+g3Zq3~s1&q#gSfYUzmPi#qQ=rH~s<7yA*eAueyTrmk-HN}7@Pq(1?DQwY~ zsN(b&D$eVyYMw7ltT!01G4M|ruQl*H)I}eQtC8*jzDnCq)!7y@t-xh0=qeja;EyUh z^i|JI>WqtQx&%I~aQf2cCK=g(rf}&cKI=l{M+%qS-Lj&r_jBMuasF1xi&e`~XHMuW zK}WCB`b#8ME8Ood!S7Z0t?bRR)Y%lJoWlKnyoDuRQMliag8zfUt2{e0{?92~N__T& zRMVBg_4-B$ysHHMYzh3kC2-3b7}}k764glE$%bl`p~9VwiKkP!Y^pyW>mKe&rc*Vs z?x7(o@pxNpZ6c8!No2AaC)JY}NM@70pgD(lbuKTRbhv!*2hXAXi%_9wcNc_%SknqF{cSCH;gskb%sY9z zgaEG0bq$HGjB_c0rZNI(r%PLCs4XW@S4N<*oQuYqau&+P&{WPvQ&|h;TC%yEi{{2M zwP-Ha7tJkYEVNuIhMF20F+Hh~On1s0Gi^%RO?yK z@W8;H#85Jq97yC+=`AC6LxM8O0$c|R5w`|{1A{%oSS>R>3A|rdD%FVuExj9)&iY|2 zr*XBMDh@&?PG>fk>doY_G=43UP3HD=I+=d8Wb!ghOSHE1CU#|b!;n_Z&m*pCP@h0GGt?$zT8uu> zuj1d_;WQ^O!)QsbLo3kg-)A-)p%(IbU3v&nZbObE16Fv`ZHZQ3CR%B+|B>*=-LH$JZ@pcpf<9AV0O_rxAYY z1N^Wa#ZwKr`Tl|hYU7;gN9Cdc2?=O_-B#&1F-j^!*{m=qnl&5b`l!lO} zB=T3eMe7Bog3BAS{Uz`NjFbH|9l|X!hwn9VNn_!Xx#PdoJPh8!f?aLY@-gbXySoJ_3;l*V|uiEQHf< z#YBF)=2w?*g$yOCZxbGo{}zErWan`_f*&Lh$@F{Hh3DaNA0c^q*CBS~cWmVMZVm55 zo|0+5+l3XLKwj*;tKp<8_(zPBeOx`f`|KTBJLq=WOW@CF_)gdn`}Dq;62;ktNAPC} zL?V1U9>G5)5Q%U-&IN4WjB|tLSGV7$;ktc2PF;Ra3Hd+Q@NI}m;`|#8r#2V-Jq_2# z53Z9c5ql%&bwW7R?L7vgObMe$QY&W&!&TZ#J2@xvGQLMOyqfbw4so93{o67=%4xWN zTehqz4X@|CB3xgxqklw~_ltx#Yx1p{e2a$f)bM*Wd{o2NY4{-xzgNR2G+gFBv8Ocr zeobB)SK3Y9r$q@s5?tPowHo-Vu9&scz~w#FsDVrW9@6l3wn1bKFvW6%u87X{&Pa4* zQih-+;QW(Pk!9^*WE`4GiP8C~h1O4=wN#=IjI?j7Exh<-A}?~1sEvLo(Usw34qfI}do z>dKSM<&t~+5me$>E1{M++56{B)* zb~uLLM-8QN4kA~xBGs2j<0^??OZDakhvi_W&&4m0cQCJ)1?aUhCg(+3LusF+{gRT< zQu_2=m(mKCSfu|j9?=t8N}tYvlw?06eLDAup3qYI(mZ zp3qYIa$eT`(>DXr6Ix3Db#jiR>nBkbJ)x!ak0|Zp_UpL}{bK?8smsux2+*hVoy0A) zRQ!_x`a3T}|5Si}*JbFR4$!Cbm&7eJ3XHb=EAZ%PiVHTM2ZU!7)&De+tyU3W`qFJ~ zrjHe9{L`;3WIilsxm!U=gtuo#@&78kk9z=TO#gQPC!3~z73)_S2E@O>bQU9j4D$KN3*?Xn_CqCH${lgrvU+hZK^1m0mrTP93mhgXw{VxWe68WdK-t>Q_ME!4E zN?QK+WVbZme-smvT$t=06zhW&r6Jar--QTM;+OKC68fX8AE(aIlBr){eRKbs zWXH13iV0KyThLeiF9J*tuXQt`Pwg*zC2<6Q75cOXh4rfq zKf73;`bKO literal 0 HcmV?d00001 From 77f19afb92910e64a060e7f84b417f994acfaf14 Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Mon, 3 Aug 2026 15:15:25 +0000 Subject: [PATCH 2/9] feat: CUDA, RMSNrom, CUDAVersion --- src/kernels.cu | 128 ++++++++++++++++++++++++++++++++++++++++--------- 1 file changed, 106 insertions(+), 22 deletions(-) diff --git a/src/kernels.cu b/src/kernels.cu index 60e5f3eb..4cdf124b 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -4,6 +4,78 @@ #include "../tester/utils.h" +// ===================================================================== +// CUDA 核函数辅助类型转换工具 (确保同时完美兼容 float 和 half) +// ===================================================================== +template __device__ __forceinline__ float to_float(T val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ float to_float(half val) { + return __half2float(val); +} + +template __device__ __forceinline__ T from_float(float val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ half from_float(float val) { + return __float2half(val); +} + +// ===================================================================== +// RMSNorm CUDA Kernel 实现 +// ===================================================================== +template +__global__ void rmsNormKernel(const T *input, const T *weight, T *output, + size_t rows, size_t hidden_dim, float eps) { + // 每个 Block 负责处理矩阵中的一个 Token (一行) + size_t i = blockIdx.x; + if (i >= rows) + return; + + // 定位当前行的起始指针 + const T *row_input = input + i * hidden_dim; + T *row_output = output + i * hidden_dim; + + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float sdata[]; + size_t tid = threadIdx.x; + + // 1. 每个线程并行计算自己分到的那一批元素的平方和 + float thread_sum = 0.0f; + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + thread_sum += val * val; + } + sdata[tid] = thread_sum; + __syncthreads(); // 等待全块线程完成局部平方和写入 + + // 2. 块内折半规约 (Block Reduction):将所有线程的和累加到 sdata[0] + // 保证 blockDim.x 是 2 的幂次(这里固定为 256),此逻辑绝对安全 + for (size_t s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + + // 3. 由 0 号线程算出这一行的 rsqrt 值,并共享给全块 + __shared__ float rsqrt_val; + if (tid == 0) { + float mean_square = sdata[0] / hidden_dim; + rsqrt_val = rsqrtf(mean_square + eps); // 使用 CUDA 硬件加速的 rsqrtf 指令 + } + __syncthreads(); // 等待 rsqrt_val 计算并同步完毕 + + // 4. 所有线程再次并行,计算当前行每个元素的最终缩放值并写回 + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + float w = to_float(weight[j]); + row_output[j] = from_float(val * rsqrt_val * w); + } +} + /** * @brief Computes RMSNorm over the last dimension of a 2D tensor. * @@ -26,28 +98,40 @@ template void rmsNorm(const std::vector &h_input, const std::vector &h_weight, std::vector &h_output, size_t rows, size_t hidden_dim, float eps) { - // TODO: Implement the rmsNorm function - - // 遍历 rows(batch_size*seq_len) - for (size_t i = 0; i < rows; i++) { - float square_sum = 0.0f; - // 1. 在第 i 行内,先计算所有元素的平方和 (还原你的 mean([i, :]^2) 逻辑) - for (size_t j = 0; j < hidden_dim; j++) { - float val = static_cast(h_input[i * hidden_dim + j]); - square_sum += val * val; // 自乘代替 ^2 - } - // 2. 计算均方根的倒数 (还原你的 rsqrt(mean + eps) 逻辑) - float mean_square = square_sum / hidden_dim; - float rsqrt_val = 1.0f / std::sqrt(mean_square + eps); - // 内层循环:更新当前行的每一个元素,应用缩放 - for (size_t j = 0; j < hidden_dim; j++) { - size_t idx = i * hidden_dim + j; - // 注意这里是 h_weight[j];统一转 float 计算,避免 half 重载歧义 - h_output[idx] = - static_cast(static_cast(h_input[idx]) * rsqrt_val * - static_cast(h_weight[j])); - } - } + // 1. 定义 Device 端的裸指针 + T *d_input = nullptr; + T *d_weight = nullptr; + T *d_output = nullptr; + + size_t input_size = rows * hidden_dim * sizeof(T); + size_t weight_size = hidden_dim * sizeof(T); + + // 2. 分配 GPU 显存 + cudaMalloc(&d_input, input_size); + cudaMalloc(&d_weight, weight_size); + cudaMalloc(&d_output, input_size); + + // 3. 将数据从 Host (CPU) 拷贝到 Device (GPU) + cudaMemcpy(d_input, h_input.data(), input_size, cudaMemcpyHostToDevice); + cudaMemcpy(d_weight, h_weight.data(), weight_size, cudaMemcpyHostToDevice); + + // 4. 配置配置网格和线程块尺寸 + // 固定使用 256 线程,它是 2 的幂次,能完美支持 Kernel 内部的折半规约 + unsigned int threads_per_block = 256; + unsigned int blocks_per_grid = rows; // 有多少行就启动多少个 Block + size_t shared_mem_size = threads_per_block * sizeof(float); + + // 5. 启动 CUDA Kernel + rmsNormKernel<<>>( + d_input, d_weight, d_output, rows, hidden_dim, eps); + + // 6. 将计算结果从 GPU 捞回预先分配好的 h_output 中 + cudaMemcpy(h_output.data(), d_output, input_size, cudaMemcpyDeviceToHost); + + // 7. 善后处理:释放显存防止内存泄漏 + cudaFree(d_input); + cudaFree(d_weight); + cudaFree(d_output); } /** From 0b4b18a5322c79beb1073502e36b31b883b3a12a Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Tue, 4 Aug 2026 16:15:48 +0000 Subject: [PATCH 3/9] feat: CUDA, FlashAtten, First Try --- src/kernels.cu | 124 ++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 123 insertions(+), 1 deletion(-) diff --git a/src/kernels.cu b/src/kernels.cu index 4cdf124b..3e56efe0 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -134,6 +134,98 @@ void rmsNorm(const std::vector &h_input, const std::vector &h_weight, cudaFree(d_output); } +// ===================================================================== +// Falsh Attention CUDA Kernel 实现 +// ===================================================================== +template +__global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, + int tgt_len, int src_len, int q_heads, + int kv_heads, int d, bool is_causal) { + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float smem[]; // size = src_len + blockDim.x + float *s_score = smem; + float *red = smem + src_len; // Reduction 缓冲区 + + int b = blockIdx.x; + int t = blockIdx.y; + int h = blockIdx.z; + int hkv = h / (q_heads / kv_heads); // ← GQA 分组查询注意力 + float scale = 1.0f / sqrtf((float)d); // ← 1/sqrt(d) + int tid = threadIdx.x, nthreads = blockDim.x; + + const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; + T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; + // K/V的第J行起始 = k + (((size_t)b * src_len+j)*kv_heads + hkv)*d; // j 变化 + + // ---- 阶段 A: 每个线程负责若干 j, 算 s_j = dot(q, k_j) * scale ---- + for (int j = tid; j < src_len; j += nthreads) { + // 1) causal 时若 j 被 mask, s_shared[j] = -INFINITY, continue + // 2) 否则定位 k_row(用 hkv!), 循环 d 维做点积(转 float 累加) + // 写进 s_shared[j] + if (is_causal && j > t) { // causal mask: 只能看 j <= t + s_score[j] = -INFINITY; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + // 在向量特征维度上循环迭代索引 + for (int dd = 0; dd < d; dd++) + dot += to_float(q_row[dd]) * to_float(k_row[dd]); + s_score[j] = dot * scale; + } + __syncthreads(); + + // ---- 阶段 B: 求 max —— 就是 rmsNorm 的折半归约, 把 + 换成 fmaxf ---- + // 每个线程先对自己的 j 集合求局部 max → 写入归约数组 → 折半归约 + // 结果 m = 全局最大 s_j + + // 找到最大的值 + float local_max = -INFINITY; + for (int j = tid; j < src_len; j += nthreads) + local_max = fmaxf(local_max, s_score[j]); + red[tid] = local_max; + __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] = fmaxf(red[tid], red[tid + s]); + __syncthreads(); + } + float m = red[0]; // 全Block都能读 + __syncthreads(); + + // ---- 阶段 C: s_j = exp(s_j - m), 再归约求和得 l ---- + // 原地改写 s_shared[j] = __expf(s_shared[j] - m) + // 再做一次加法归约(和 rmsNorm 一模一样)得到 l + float local_sum = 0.f; + for (int j = tid; j < src_len; j += nthreads) { + float e = __expf(s_score[j] - m); + s_score[j] = e; + local_sum += e; + } + red[tid] = local_sum; + __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] += red[tid + s]; + __syncthreads(); + } + float l = red[0]; + __syncthreads(); + + // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- + for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { + float acc = 0.f; + // 循环所有 j: acc += s_score[j] * v_row[dd](p_j 最后再除 l) + for (int j = 0; j < src_len; j++) { + const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + acc += s_score[j] * to_float(v_row[dd]); + } + o_row[dd] = from_float(acc / l); + } +} + +// Hidden_dim == num_heads * head_dim. +// query_heads和kv_heads是否相同,则决定了head_dim的大小 /** * @brief Computes flash attention for given query, key, and value tensors. * @@ -161,7 +253,37 @@ void flashAttention(const std::vector &h_q, const std::vector &h_k, int batch_size, int target_seq_len, int src_seq_len, int query_heads, int kv_heads, int head_dim, bool is_causal) { - // TODO: Implement the flash attention function + // 和 rmsNorm 一模一样的套路,只是换成 4 个张量: + T *d_q, *d_k, *d_v, *d_o; + size_t q_size = + (size_t)batch_size * target_seq_len * query_heads * head_dim * sizeof(T); + size_t kv_size = + (size_t)batch_size * src_seq_len * kv_heads * head_dim * sizeof(T); + // cudaMalloc x4 → cudaMemcpy q/k/v → launch → memcpy 回 h_o → cudaFree x4 + cudaMalloc(&d_q, q_size); + cudaMalloc(&d_k, kv_size); + cudaMalloc(&d_v, kv_size); + cudaMalloc(&d_o, q_size); + + cudaMemcpy(d_q, h_q.data(), q_size, cudaMemcpyHostToDevice); + cudaMemcpy(d_k, h_k.data(), kv_size, cudaMemcpyHostToDevice); + cudaMemcpy(d_v, h_v.data(), kv_size, cudaMemcpyHostToDevice); + + // 启动配置:一个 block 负责一个输出行 (b, t, h) + // 三维网格,天然映射, 剩下的一个就是head dim + dim3 grid(batch_size, target_seq_len, query_heads); + int threads = 128; + // s_score[src_len] + red[threads] 两块区域,缺一不可! + size_t shmem = ((size_t)src_seq_len + threads) * sizeof(float); + flashAttentionKernel<<>>( + d_q, d_k, d_v, d_o, target_seq_len, src_seq_len, query_heads, kv_heads, + head_dim, is_causal); + + cudaMemcpy(h_o.data(), d_o, q_size, cudaMemcpyDeviceToHost); + cudaFree(d_q); + cudaFree(d_k); + cudaFree(d_v); + cudaFree(d_o); } // ********************************************************************* From 94b8acab9f7a088acabbb7308c11c6e6bc07a176 Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Tue, 4 Aug 2026 16:38:17 +0000 Subject: [PATCH 4/9] feat: CUDA, FlashAtten --- src/kernels.cu | 156 ++++++++++++++++++++++++++++--------------------- 1 file changed, 88 insertions(+), 68 deletions(-) diff --git a/src/kernels.cu b/src/kernels.cu index 3e56efe0..228c9c5a 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -135,93 +135,113 @@ void rmsNorm(const std::vector &h_input, const std::vector &h_weight, } // ===================================================================== -// Falsh Attention CUDA Kernel 实现 +// Flash Attention CUDA Kernel 实现 (online-softmax / 单遍分块流式) // ===================================================================== +// 每个 block 负责一个输出行 (b, t, h);对 K/V 按 tile(大小 = blockDim) +// 流式扫描,维护 running max(m)/running sum(l)/running 输出累加(acc), +// 用校正因子 alpha = exp(m_old - m_new) 把旧累加量搬到新基准, +// 从而单遍完成且数值稳定(任意时刻 exp 参数 <= 0,不溢出)。 +// 线程分工:线程 tid 负责本 tile 的 key j0+tid 的打分;同时若 tid __global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, int tgt_len, int src_len, int q_heads, int kv_heads, int d, bool is_causal) { - // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) - extern __shared__ float smem[]; // size = src_len + blockDim.x - float *s_score = smem; - float *red = smem + src_len; // Reduction 缓冲区 - int b = blockIdx.x; int t = blockIdx.y; int h = blockIdx.z; - int hkv = h / (q_heads / kv_heads); // ← GQA 分组查询注意力 - float scale = 1.0f / sqrtf((float)d); // ← 1/sqrt(d) + int hkv = h / (q_heads / kv_heads); // GQA:多个 query head 共享一个 kv head + float scale = 1.0f / sqrtf((float)d); // 1/sqrt(d)(标准正确舍入版) int tid = threadIdx.x, nthreads = blockDim.x; const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; - // K/V的第J行起始 = k + (((size_t)b * src_len+j)*kv_heads + hkv)*d; // j 变化 - - // ---- 阶段 A: 每个线程负责若干 j, 算 s_j = dot(q, k_j) * scale ---- - for (int j = tid; j < src_len; j += nthreads) { - // 1) causal 时若 j 被 mask, s_shared[j] = -INFINITY, continue - // 2) 否则定位 k_row(用 hkv!), 循环 d 维做点积(转 float 累加) - // 写进 s_shared[j] - if (is_causal && j > t) { // causal mask: 只能看 j <= t - s_score[j] = -INFINITY; - continue; - } - const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; - float dot = 0.f; - // 在向量特征维度上循环迭代索引 - for (int dd = 0; dd < d; dd++) - dot += to_float(q_row[dd]) * to_float(k_row[dd]); - s_score[j] = dot * scale; - } - __syncthreads(); - // ---- 阶段 B: 求 max —— 就是 rmsNorm 的折半归约, 把 + 换成 fmaxf ---- - // 每个线程先对自己的 j 集合求局部 max → 写入归约数组 → 折半归约 - // 结果 m = 全局最大 s_j + // 共享内存布局: sh_q[d] | sh_p[nthreads] | sh_red[nthreads] + extern __shared__ float smem[]; + float *sh_q = smem; + float *sh_p = sh_q + d; + float *sh_red = sh_p + nthreads; - // 找到最大的值 - float local_max = -INFINITY; - for (int j = tid; j < src_len; j += nthreads) - local_max = fmaxf(local_max, s_score[j]); - red[tid] = local_max; + // 把 query 行缓存到共享内存(每个 tile 都要用,避免重复读全局) + for (int i = tid; i < d; i += nthreads) + sh_q[i] = to_float(q_row[i]); __syncthreads(); - for (int s = nthreads / 2; s > 0; s >>= 1) { - if (tid < s) - red[tid] = fmaxf(red[tid], red[tid + s]); + + // running 状态:m/l 每个线程各持一份相同副本;acc 每线程负责通道 dd=tid + float m = -INFINITY, l = 0.f, acc = 0.f; + + // causal: 只需扫到 j<=t;否则扫到 src_len + int jend = src_len; + if (is_causal && t + 1 < jend) + jend = t + 1; + + for (int j0 = 0; j0 < jend; j0 += nthreads) { + int j = j0 + tid; + + // (1) 本线程负责 key j 的打分 s = (q·k_j) * scale;越界/被 mask 记 -inf + float s = -INFINITY; + if (j < jend) { + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + for (int dd = 0; dd < d; dd++) + dot += sh_q[dd] * to_float(k_row[dd]); + s = dot * scale; + } + + // (2) tile 内最大值(折半归约) + sh_red[tid] = s; + __syncthreads(); + for (int r = nthreads / 2; r > 0; r >>= 1) { + if (tid < r) + sh_red[tid] = fmaxf(sh_red[tid], sh_red[tid + r]); + __syncthreads(); + } + float tile_max = sh_red[0]; __syncthreads(); - } - float m = red[0]; // 全Block都能读 - __syncthreads(); - // ---- 阶段 C: s_j = exp(s_j - m), 再归约求和得 l ---- - // 原地改写 s_shared[j] = __expf(s_shared[j] - m) - // 再做一次加法归约(和 rmsNorm 一模一样)得到 l - float local_sum = 0.f; - for (int j = tid; j < src_len; j += nthreads) { - float e = __expf(s_score[j] - m); - s_score[j] = e; - local_sum += e; - } - red[tid] = local_sum; - __syncthreads(); - for (int s = nthreads / 2; s > 0; s >>= 1) { - if (tid < s) - red[tid] += red[tid + s]; + // (3) 更新 running max,并算校正因子 alpha = exp(m_old - m_new) + float m_new = fmaxf(m, tile_max); + float alpha = __expf(m - m_new); + + // (4) 本 key 的 exp 权重(相对新基准 m_new) + float p = (j < jend) ? __expf(s - m_new) : 0.f; + sh_p[tid] = p; + + // (5) tile 内权重和(折半归约) + sh_red[tid] = p; + __syncthreads(); + for (int r = nthreads / 2; r > 0; r >>= 1) { + if (tid < r) + sh_red[tid] += sh_red[tid + r]; + __syncthreads(); + } + float tile_sum = sh_red[0]; __syncthreads(); - } - float l = red[0]; - __syncthreads(); - // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- - for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { - float acc = 0.f; - // 循环所有 j: acc += s_score[j] * v_row[dd](p_j 最后再除 l) - for (int j = 0; j < src_len; j++) { - const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; - acc += s_score[j] * to_float(v_row[dd]); + // (6) 更新归一化分母: l = alpha*l + 本 tile 权重和 + l = alpha * l + tile_sum; + + // (7) 更新输出累加:线程 tid 负责通道 dd=tid + if (tid < d) { + float delta = 0.f; + int cnt = jend - j0; + if (cnt > nthreads) + cnt = nthreads; + for (int jj = 0; jj < cnt; jj++) { + const T *v_row = + v + (((size_t)b * src_len + (j0 + jj)) * kv_heads + hkv) * d; + delta += sh_p[jj] * to_float(v_row[tid]); + } + acc = alpha * acc + delta; } - o_row[dd] = from_float(acc / l); + m = m_new; + __syncthreads(); // 复用 sh_p/sh_red 前同步 } + + // (8) 归一化写回 + if (tid < d) + o_row[tid] = from_float(acc / l); } // Hidden_dim == num_heads * head_dim. @@ -273,8 +293,8 @@ void flashAttention(const std::vector &h_q, const std::vector &h_k, // 三维网格,天然映射, 剩下的一个就是head dim dim3 grid(batch_size, target_seq_len, query_heads); int threads = 128; - // s_score[src_len] + red[threads] 两块区域,缺一不可! - size_t shmem = ((size_t)src_seq_len + threads) * sizeof(float); + // 共享内存: sh_q[head_dim] + sh_p[threads] + sh_red[threads] + size_t shmem = ((size_t)head_dim + 2 * threads) * sizeof(float); flashAttentionKernel<<>>( d_q, d_k, d_v, d_o, target_seq_len, src_seq_len, query_heads, kv_heads, head_dim, is_causal); From e1d89ac4bb51d28207ccd8dd2c11612749f18ce5 Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Wed, 5 Aug 2026 01:31:43 +0000 Subject: [PATCH 5/9] Revert "feat: CUDA, FlashAtten" This reverts commit 94b8acab9f7a088acabbb7308c11c6e6bc07a176. --- src/kernels.cu | 156 +++++++++++++++++++++---------------------------- 1 file changed, 68 insertions(+), 88 deletions(-) diff --git a/src/kernels.cu b/src/kernels.cu index 228c9c5a..3e56efe0 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -135,113 +135,93 @@ void rmsNorm(const std::vector &h_input, const std::vector &h_weight, } // ===================================================================== -// Flash Attention CUDA Kernel 实现 (online-softmax / 单遍分块流式) +// Falsh Attention CUDA Kernel 实现 // ===================================================================== -// 每个 block 负责一个输出行 (b, t, h);对 K/V 按 tile(大小 = blockDim) -// 流式扫描,维护 running max(m)/running sum(l)/running 输出累加(acc), -// 用校正因子 alpha = exp(m_old - m_new) 把旧累加量搬到新基准, -// 从而单遍完成且数值稳定(任意时刻 exp 参数 <= 0,不溢出)。 -// 线程分工:线程 tid 负责本 tile 的 key j0+tid 的打分;同时若 tid __global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, int tgt_len, int src_len, int q_heads, int kv_heads, int d, bool is_causal) { + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float smem[]; // size = src_len + blockDim.x + float *s_score = smem; + float *red = smem + src_len; // Reduction 缓冲区 + int b = blockIdx.x; int t = blockIdx.y; int h = blockIdx.z; - int hkv = h / (q_heads / kv_heads); // GQA:多个 query head 共享一个 kv head - float scale = 1.0f / sqrtf((float)d); // 1/sqrt(d)(标准正确舍入版) + int hkv = h / (q_heads / kv_heads); // ← GQA 分组查询注意力 + float scale = 1.0f / sqrtf((float)d); // ← 1/sqrt(d) int tid = threadIdx.x, nthreads = blockDim.x; const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; - - // 共享内存布局: sh_q[d] | sh_p[nthreads] | sh_red[nthreads] - extern __shared__ float smem[]; - float *sh_q = smem; - float *sh_p = sh_q + d; - float *sh_red = sh_p + nthreads; - - // 把 query 行缓存到共享内存(每个 tile 都要用,避免重复读全局) - for (int i = tid; i < d; i += nthreads) - sh_q[i] = to_float(q_row[i]); + // K/V的第J行起始 = k + (((size_t)b * src_len+j)*kv_heads + hkv)*d; // j 变化 + + // ---- 阶段 A: 每个线程负责若干 j, 算 s_j = dot(q, k_j) * scale ---- + for (int j = tid; j < src_len; j += nthreads) { + // 1) causal 时若 j 被 mask, s_shared[j] = -INFINITY, continue + // 2) 否则定位 k_row(用 hkv!), 循环 d 维做点积(转 float 累加) + // 写进 s_shared[j] + if (is_causal && j > t) { // causal mask: 只能看 j <= t + s_score[j] = -INFINITY; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + // 在向量特征维度上循环迭代索引 + for (int dd = 0; dd < d; dd++) + dot += to_float(q_row[dd]) * to_float(k_row[dd]); + s_score[j] = dot * scale; + } __syncthreads(); - // running 状态:m/l 每个线程各持一份相同副本;acc 每线程负责通道 dd=tid - float m = -INFINITY, l = 0.f, acc = 0.f; - - // causal: 只需扫到 j<=t;否则扫到 src_len - int jend = src_len; - if (is_causal && t + 1 < jend) - jend = t + 1; - - for (int j0 = 0; j0 < jend; j0 += nthreads) { - int j = j0 + tid; - - // (1) 本线程负责 key j 的打分 s = (q·k_j) * scale;越界/被 mask 记 -inf - float s = -INFINITY; - if (j < jend) { - const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; - float dot = 0.f; - for (int dd = 0; dd < d; dd++) - dot += sh_q[dd] * to_float(k_row[dd]); - s = dot * scale; - } + // ---- 阶段 B: 求 max —— 就是 rmsNorm 的折半归约, 把 + 换成 fmaxf ---- + // 每个线程先对自己的 j 集合求局部 max → 写入归约数组 → 折半归约 + // 结果 m = 全局最大 s_j - // (2) tile 内最大值(折半归约) - sh_red[tid] = s; - __syncthreads(); - for (int r = nthreads / 2; r > 0; r >>= 1) { - if (tid < r) - sh_red[tid] = fmaxf(sh_red[tid], sh_red[tid + r]); - __syncthreads(); - } - float tile_max = sh_red[0]; + // 找到最大的值 + float local_max = -INFINITY; + for (int j = tid; j < src_len; j += nthreads) + local_max = fmaxf(local_max, s_score[j]); + red[tid] = local_max; + __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] = fmaxf(red[tid], red[tid + s]); __syncthreads(); + } + float m = red[0]; // 全Block都能读 + __syncthreads(); - // (3) 更新 running max,并算校正因子 alpha = exp(m_old - m_new) - float m_new = fmaxf(m, tile_max); - float alpha = __expf(m - m_new); - - // (4) 本 key 的 exp 权重(相对新基准 m_new) - float p = (j < jend) ? __expf(s - m_new) : 0.f; - sh_p[tid] = p; - - // (5) tile 内权重和(折半归约) - sh_red[tid] = p; - __syncthreads(); - for (int r = nthreads / 2; r > 0; r >>= 1) { - if (tid < r) - sh_red[tid] += sh_red[tid + r]; - __syncthreads(); - } - float tile_sum = sh_red[0]; + // ---- 阶段 C: s_j = exp(s_j - m), 再归约求和得 l ---- + // 原地改写 s_shared[j] = __expf(s_shared[j] - m) + // 再做一次加法归约(和 rmsNorm 一模一样)得到 l + float local_sum = 0.f; + for (int j = tid; j < src_len; j += nthreads) { + float e = __expf(s_score[j] - m); + s_score[j] = e; + local_sum += e; + } + red[tid] = local_sum; + __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] += red[tid + s]; __syncthreads(); + } + float l = red[0]; + __syncthreads(); - // (6) 更新归一化分母: l = alpha*l + 本 tile 权重和 - l = alpha * l + tile_sum; - - // (7) 更新输出累加:线程 tid 负责通道 dd=tid - if (tid < d) { - float delta = 0.f; - int cnt = jend - j0; - if (cnt > nthreads) - cnt = nthreads; - for (int jj = 0; jj < cnt; jj++) { - const T *v_row = - v + (((size_t)b * src_len + (j0 + jj)) * kv_heads + hkv) * d; - delta += sh_p[jj] * to_float(v_row[tid]); - } - acc = alpha * acc + delta; + // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- + for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { + float acc = 0.f; + // 循环所有 j: acc += s_score[j] * v_row[dd](p_j 最后再除 l) + for (int j = 0; j < src_len; j++) { + const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + acc += s_score[j] * to_float(v_row[dd]); } - m = m_new; - __syncthreads(); // 复用 sh_p/sh_red 前同步 + o_row[dd] = from_float(acc / l); } - - // (8) 归一化写回 - if (tid < d) - o_row[tid] = from_float(acc / l); } // Hidden_dim == num_heads * head_dim. @@ -293,8 +273,8 @@ void flashAttention(const std::vector &h_q, const std::vector &h_k, // 三维网格,天然映射, 剩下的一个就是head dim dim3 grid(batch_size, target_seq_len, query_heads); int threads = 128; - // 共享内存: sh_q[head_dim] + sh_p[threads] + sh_red[threads] - size_t shmem = ((size_t)head_dim + 2 * threads) * sizeof(float); + // s_score[src_len] + red[threads] 两块区域,缺一不可! + size_t shmem = ((size_t)src_seq_len + threads) * sizeof(float); flashAttentionKernel<<>>( d_q, d_k, d_v, d_o, target_seq_len, src_seq_len, query_heads, kv_heads, head_dim, is_causal); From d6d1fa0ad282745f7b5fe8a6a4f208682ab2c64e Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Wed, 5 Aug 2026 01:32:32 +0000 Subject: [PATCH 6/9] remove *.o file --- src/kernels.o | Bin 16536 -> 0 bytes 1 file changed, 0 insertions(+), 0 deletions(-) delete mode 100644 src/kernels.o diff --git a/src/kernels.o b/src/kernels.o deleted file mode 100644 index a9eed28015ed1e866419eb3c00a429609a3725a4..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 16536 zcmcIrdvILUdB3Z*v9W=b2??ecyfrq67#6)BU_ivaW<5xbVP#k$o0*Q?Oq!9M8LYHvY97;gCjLW{l0jevMD1Zr>!SX? zbHDH2vqyKmm}$?<)pySCd!F|>_pbKYZ5t~>Axl!oy3g`Uf?C#r8b9BvvaQxS>mDV0 zEy`cQLwaAv^A$WR@Z5ywW;`qL5bt`X%Rz78{2zhdh$o8YkMXR+a|<4lUyWxC=WhkQ zjq_DZZ)aM~ltLuETBdcN=OQ=%S7*HN`*>mRh4{FAAzoMr)Y`w}7anJ9CR(_VF4!|Y z@$m&~$auV9U#OXg9JOZve**eRe7y5QPkillitQ|%CL6!>O%xl~Y`?6$#S15=PtDEE z#mBuO87K~-(hu>ZX@6{HvCnmFsJfh-UIDZFwMfQhD&xLW^1?L+`XdSfRin!Iq0GJI z=apYbu2Yh37_bg2T5DCFi!8G*j@uX0k?aIVs@k2rRnQ-#1un4`Gp%GAW!lPgl1i0gXsFddzCyWw}F-mJ(qQ(qm>$%*(bRE-sncl~ABhxt3hnQ|*`UulUnQmkH zIMa}tbEEg`?MhutZF_#)KJT?{;RFrabfJ>^!IP!|R~ViJYc_JUk|l^}V-MpR7N?C- z4Wa4CQR{OKY#L5HxS@n$OyU3X7ivb0&9qXqXCQD2!Au&C#_cl&durUCT7$u~uV7Cj zdwjZbjWudns#2#4l`rg_25;ZUT0}K@QZ)@A%NpG`ycE!>aclHf&!1g*Y;GTP?04V) zSBQ|#NBbTexyjQRbw#h9d(M^mTb7#J*E~{%>K>o^l5#i+hXYiDW>;anbNcuPS6_Jw zmd}nTzgFF@>&PrP1!`3UK+J7qWcqq65y!p>^|kg)!JY+%n(c$Y`4`&HBWCMZR554A zJI@!6zx_Kh#8xWNPvgUcepr}Q4qr&UJ2z)pSx+J}*&6Jt}MZ%;T&sWdhbB4L=p zW6#73|4vI5S*Jzt(h;h)^NbN9)Xt=jPeK#K8wP=*3mv0o#weh8AbrxFLCaZZo1tb6$CS(+0G|PfcE*furA~P> zH9Iyi6B^rl-g2g2NX;rE3&!l3A?nSkm4y@I_RR6A>#sbu)Sg8zqbo38&qXTjDfAqz z_U>ljt@hj{(R5+&hembLXdg<8p&|bh^}{p4ZJ3UH6RmKbOrH-n?ak}4nVpzUlbB8u z@j~aZc%l7BTfsg|qu}-UNqY)awBGN*5|@sAG(qSopsJ(JJ_yrO@yL(u!*Mm*r{f`7 zI$+D|OhkYs1XGqm!Jlm~QQ3a#T;#FiuB|b9Is~2R_&xUF+>env9RK#j?xo6EJo43v zvp10KH$TA8iC=j_;1e()mY$r1rS^Bo0scH@zfP6+Mr5Q4>`JwPor^rFoSsGn)Nnq# z*Qmg3JcM<{=Itt=BHQfK6!syEPg)7lVnsELRDr#1$05k34(W1JHVlI!5*|Svb zS*mvA#bMg~#%xrP8uT9fNaV#o{S3Qo@isezeuA~(Bj_izc>5vr(*X+S6?1Hyv`?#U z`xC9(CZ#jaw!_P5>b3)Jw@ta-rU&YcEH#d#ZQag8d%B%BsfTy;>ro4jS`8?k0jc^zM-?tPHfs>SzG?l`O}b9 zX*KhKslp0{S|QgL#vd&$!U0LnCIm0jczSl9d!Dab#9)Q>bJ}pdbb$|OyWL}wRll$v z%VjgpZLQW7bO?*(`|?iCNp>M8bkBfWRy2|C%l7QHV%d>cCYv5qneJqNe{U{1)Tc^A zxj|>p*)x<hd#6%gVza{4isXZw!)T$624_f@{twfO zpm|?<*&1tMqhI`DIqD9h%ax&4Hn)SR5bi!HT1$}&(*GVyl+&mC8xIZ7rpmCRIHM#( zein0??sAm)&Nbwg{J(M5k1`}N@^JU;P(Pp!^Yr^zztT6Ua^hd;6+o#MWqcdro?&qy zdSw2K>4idDpxerhS~$N5H7c5?Q5Il=DX&tTsgQL`?z$^%&8}J!8%gEzm@LueSksc2 zlg#y|oM?U^(NMc2mh9=trSkbiKJ#oU+E~8?N%a4Ieb+{ADuIgVyu}PlUVI!&HrD3j zzpKDft=$}C9lK1K&|At3ju_gV+r6cDzC`B1_3Q7BR&{m_XPx0_O#Nv9vH8YobW5YrwkB( z9rHgf;?H>e?aY6oi0?KK=|9i>!$c*%@T;ExgUtU!5&vC}|1;*(`y5Yyr(eG_W%w7F ze}k^S#n)d!ixyHT`|Frb@2Wiioxc8d=6_Y^Cw>0&%zsknSGWeKeGW3eWdE1ppJDzj zy8Xv}{}(U8CzYT?O}4=LSnY2nKWtrNLaH1NICWHm6Vh<*3agEAS!adjfLCbt)tpq; zX)a&OIiDOtmgGfF%|T{|Et&Pc{7;b$X}I|PE#p!adXaHjTPO*B6DmS>WPcOJ+~YKLh9zm4%(1HYH?3kKf7IE<;Z znE4sT7aRCK#w!i{yNoY4@EN90-%&K)ynK9dLoyGex!`52P`Q#$x@n+Qrv|QM_Lo2!2 zML6(>G+g3Zq3~s1&q#gSfYUzmPi#qQ=rH~s<7yA*eAueyTrmk-HN}7@Pq(1?DQwY~ zsN(b&D$eVyYMw7ltT!01G4M|ruQl*H)I}eQtC8*jzDnCq)!7y@t-xh0=qeja;EyUh z^i|JI>WqtQx&%I~aQf2cCK=g(rf}&cKI=l{M+%qS-Lj&r_jBMuasF1xi&e`~XHMuW zK}WCB`b#8ME8Ood!S7Z0t?bRR)Y%lJoWlKnyoDuRQMliag8zfUt2{e0{?92~N__T& zRMVBg_4-B$ysHHMYzh3kC2-3b7}}k764glE$%bl`p~9VwiKkP!Y^pyW>mKe&rc*Vs z?x7(o@pxNpZ6c8!No2AaC)JY}NM@70pgD(lbuKTRbhv!*2hXAXi%_9wcNc_%SknqF{cSCH;gskb%sY9z zgaEG0bq$HGjB_c0rZNI(r%PLCs4XW@S4N<*oQuYqau&+P&{WPvQ&|h;TC%yEi{{2M zwP-Ha7tJkYEVNuIhMF20F+Hh~On1s0Gi^%RO?yK z@W8;H#85Jq97yC+=`AC6LxM8O0$c|R5w`|{1A{%oSS>R>3A|rdD%FVuExj9)&iY|2 zr*XBMDh@&?PG>fk>doY_G=43UP3HD=I+=d8Wb!ghOSHE1CU#|b!;n_Z&m*pCP@h0GGt?$zT8uu> zuj1d_;WQ^O!)QsbLo3kg-)A-)p%(IbU3v&nZbObE16Fv`ZHZQ3CR%B+|B>*=-LH$JZ@pcpf<9AV0O_rxAYY z1N^Wa#ZwKr`Tl|hYU7;gN9Cdc2?=O_-B#&1F-j^!*{m=qnl&5b`l!lO} zB=T3eMe7Bog3BAS{Uz`NjFbH|9l|X!hwn9VNn_!Xx#PdoJPh8!f?aLY@-gbXySoJ_3;l*V|uiEQHf< z#YBF)=2w?*g$yOCZxbGo{}zErWan`_f*&Lh$@F{Hh3DaNA0c^q*CBS~cWmVMZVm55 zo|0+5+l3XLKwj*;tKp<8_(zPBeOx`f`|KTBJLq=WOW@CF_)gdn`}Dq;62;ktNAPC} zL?V1U9>G5)5Q%U-&IN4WjB|tLSGV7$;ktc2PF;Ra3Hd+Q@NI}m;`|#8r#2V-Jq_2# z53Z9c5ql%&bwW7R?L7vgObMe$QY&W&!&TZ#J2@xvGQLMOyqfbw4so93{o67=%4xWN zTehqz4X@|CB3xgxqklw~_ltx#Yx1p{e2a$f)bM*Wd{o2NY4{-xzgNR2G+gFBv8Ocr zeobB)SK3Y9r$q@s5?tPowHo-Vu9&scz~w#FsDVrW9@6l3wn1bKFvW6%u87X{&Pa4* zQih-+;QW(Pk!9^*WE`4GiP8C~h1O4=wN#=IjI?j7Exh<-A}?~1sEvLo(Usw34qfI}do z>dKSM<&t~+5me$>E1{M++56{B)* zb~uLLM-8QN4kA~xBGs2j<0^??OZDakhvi_W&&4m0cQCJ)1?aUhCg(+3LusF+{gRT< zQu_2=m(mKCSfu|j9?=t8N}tYvlw?06eLDAup3qYI(mZ zp3qYIa$eT`(>DXr6Ix3Db#jiR>nBkbJ)x!ak0|Zp_UpL}{bK?8smsux2+*hVoy0A) zRQ!_x`a3T}|5Si}*JbFR4$!Cbm&7eJ3XHb=EAZ%PiVHTM2ZU!7)&De+tyU3W`qFJ~ zrjHe9{L`;3WIilsxm!U=gtuo#@&78kk9z=TO#gQPC!3~z73)_S2E@O>bQU9j4D$KN3*?Xn_CqCH${lgrvU+hZK^1m0mrTP93mhgXw{VxWe68WdK-t>Q_ME!4E zN?QK+WVbZme-smvT$t=06zhW&r6Jar--QTM;+OKC68fX8AE(aIlBr){eRKbs zWXH13iV0KyThLeiF9J*tuXQt`Pwg*zC2<6Q75cOXh4rfq zKf73;`bKO From 65d9580c83bff3eadf223081f389cc46fba20d1c Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Wed, 5 Aug 2026 05:00:38 +0000 Subject: [PATCH 7/9] feat: CUDA, FalshAtten-fix1 --- src/kernels.cu | 58 ++++++++++++++++++++++++-------------------------- 1 file changed, 28 insertions(+), 30 deletions(-) diff --git a/src/kernels.cu b/src/kernels.cu index 3e56efe0..e29a19f5 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -141,81 +141,79 @@ template __global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, int tgt_len, int src_len, int q_heads, int kv_heads, int d, bool is_causal) { - // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) extern __shared__ float smem[]; // size = src_len + blockDim.x float *s_score = smem; - float *red = smem + src_len; // Reduction 缓冲区 + float *red = smem + src_len; int b = blockIdx.x; int t = blockIdx.y; int h = blockIdx.z; - int hkv = h / (q_heads / kv_heads); // ← GQA 分组查询注意力 - float scale = 1.0f / sqrtf((float)d); // ← 1/sqrt(d) + int hkv = h / (q_heads / kv_heads); // GQA 分组查询 + float scale = 1.0f / sqrtf((float)d); int tid = threadIdx.x, nthreads = blockDim.x; const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; - // K/V的第J行起始 = k + (((size_t)b * src_len+j)*kv_heads + hkv)*d; // j 变化 - // ---- 阶段 A: 每个线程负责若干 j, 算 s_j = dot(q, k_j) * scale ---- + // ---- 阶段 A: 计算 s_j = dot(q, k_j) * scale 并存 Shared Memory ---- for (int j = tid; j < src_len; j += nthreads) { - // 1) causal 时若 j 被 mask, s_shared[j] = -INFINITY, continue - // 2) 否则定位 k_row(用 hkv!), 循环 d 维做点积(转 float 累加) - // 写进 s_shared[j] - if (is_causal && j > t) { // causal mask: 只能看 j <= t + if (is_causal && j > t) { s_score[j] = -INFINITY; continue; } const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; float dot = 0.f; - // 在向量特征维度上循环迭代索引 for (int dd = 0; dd < d; dd++) dot += to_float(q_row[dd]) * to_float(k_row[dd]); s_score[j] = dot * scale; } __syncthreads(); - // ---- 阶段 B: 求 max —— 就是 rmsNorm 的折半归约, 把 + 换成 fmaxf ---- - // 每个线程先对自己的 j 集合求局部 max → 写入归约数组 → 折半归约 - // 结果 m = 全局最大 s_j - - // 找到最大的值 + // ---- 阶段 B: 求 max(s_j) ---- float local_max = -INFINITY; for (int j = tid; j < src_len; j += nthreads) local_max = fmaxf(local_max, s_score[j]); red[tid] = local_max; __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { if (tid < s) red[tid] = fmaxf(red[tid], red[tid + s]); __syncthreads(); } - float m = red[0]; // 全Block都能读 + float m = red[0]; __syncthreads(); - // ---- 阶段 C: s_j = exp(s_j - m), 再归约求和得 l ---- - // 原地改写 s_shared[j] = __expf(s_shared[j] - m) - // 再做一次加法归约(和 rmsNorm 一模一样)得到 l - float local_sum = 0.f; + // ---- 阶段 C: 重新计算 dot,保证 expf(dot * scale - m) 的 FMA 指令融合精度 + // ---- for (int j = tid; j < src_len; j += nthreads) { - float e = __expf(s_score[j] - m); - s_score[j] = e; - local_sum += e; + if (is_causal && j > t) { + s_score[j] = 0.f; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + for (int dd = 0; dd < d; dd++) + dot += to_float(q_row[dd]) * to_float(k_row[dd]); + s_score[j] = expf(dot * scale - m); } - red[tid] = local_sum; __syncthreads(); - for (int s = nthreads / 2; s > 0; s >>= 1) { - if (tid < s) - red[tid] += red[tid + s]; - __syncthreads(); + + // ---- 阶段 C.2: 由 0 号线程按 j 升序单线程串行求和 l (严格与 CPU + // 参考实现的结合树一致) ---- + if (tid == 0) { + float l_seq = 0.f; + for (int j = 0; j < src_len; j++) + l_seq += s_score[j]; + red[0] = l_seq; } + __syncthreads(); float l = red[0]; __syncthreads(); // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { float acc = 0.f; - // 循环所有 j: acc += s_score[j] * v_row[dd](p_j 最后再除 l) for (int j = 0; j < src_len; j++) { const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; acc += s_score[j] * to_float(v_row[dd]); From 4d1aec7c3c54b66a1f2d4bdae2356f5f0670fe1b Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Wed, 5 Aug 2026 05:15:31 +0000 Subject: [PATCH 8/9] feat: MetaX, RMSNorm+FlashAtten --- src/kernels.maca | 305 ++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 278 insertions(+), 27 deletions(-) diff --git a/src/kernels.maca b/src/kernels.maca index 4c320f21..2e59f005 100644 --- a/src/kernels.maca +++ b/src/kernels.maca @@ -1,8 +1,83 @@ -#include +#include +#include +#include #include +#include #include "../tester/utils.h" +// ===================================================================== +// CUDA 核函数辅助类型转换工具 (确保同时完美兼容 float 和 half) +// ===================================================================== +template __device__ __forceinline__ float to_float(T val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ float to_float(half val) { + return __half2float(val); +} + +template __device__ __forceinline__ T from_float(float val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ half from_float(float val) { + return __float2half(val); +} + +// ===================================================================== +// RMSNorm CUDA Kernel 实现 +// ===================================================================== +template +__global__ void rmsNormKernel(const T *input, const T *weight, T *output, + size_t rows, size_t hidden_dim, float eps) { + // 每个 Block 负责处理矩阵中的一个 Token (一行) + size_t i = blockIdx.x; + if (i >= rows) + return; + + // 定位当前行的起始指针 + const T *row_input = input + i * hidden_dim; + T *row_output = output + i * hidden_dim; + + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float sdata[]; + size_t tid = threadIdx.x; + + // 1. 每个线程并行计算自己分到的那一批元素的平方和 + float thread_sum = 0.0f; + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + thread_sum += val * val; + } + sdata[tid] = thread_sum; + __syncthreads(); // 等待全块线程完成局部平方和写入 + + // 2. 块内折半规约 (Block Reduction):将所有线程的和累加到 sdata[0] + // 保证 blockDim.x 是 2 的幂次(这里固定为 256),此逻辑绝对安全 + for (size_t s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + + // 3. 由 0 号线程算出这一行的 rsqrt 值,并共享给全块 + __shared__ float rsqrt_val; + if (tid == 0) { + float mean_square = sdata[0] / hidden_dim; + rsqrt_val = rsqrtf(mean_square + eps); // 使用 CUDA 硬件加速的 rsqrtf 指令 + } + __syncthreads(); // 等待 rsqrt_val 计算并同步完毕 + + // 4. 所有线程再次并行,计算当前行每个元素的最终缩放值并写回 + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + float w = to_float(weight[j]); + row_output[j] = from_float(val * rsqrt_val * w); + } +} + /** * @brief Computes RMSNorm over the last dimension of a 2D tensor. * @@ -22,47 +97,223 @@ * @param[in] eps Numerical stability epsilon. */ template -void rmsNorm(const std::vector& h_input, const std::vector& h_weight, - std::vector& h_output, size_t rows, size_t hidden_dim, - float eps) { - // TODO: Implement the rmsNorm function +void rmsNorm(const std::vector &h_input, const std::vector &h_weight, + std::vector &h_output, size_t rows, size_t hidden_dim, + float eps) { + // 1. 定义 Device 端的裸指针 + T *d_input = nullptr; + T *d_weight = nullptr; + T *d_output = nullptr; + + size_t input_size = rows * hidden_dim * sizeof(T); + size_t weight_size = hidden_dim * sizeof(T); + + // 2. 分配 GPU 显存 + mcMalloc(&d_input, input_size); + mcMalloc(&d_weight, weight_size); + mcMalloc(&d_output, input_size); + + // 3. 将数据从 Host (CPU) 拷贝到 Device (GPU) + mcMemcpy(d_input, h_input.data(), input_size, mcMemcpyHostToDevice); + mcMemcpy(d_weight, h_weight.data(), weight_size, mcMemcpyHostToDevice); + + // 4. 配置配置网格和线程块尺寸 + // 固定使用 256 线程,它是 2 的幂次,能完美支持 Kernel 内部的折半规约 + unsigned int threads_per_block = 256; + unsigned int blocks_per_grid = rows; // 有多少行就启动多少个 Block + size_t shared_mem_size = threads_per_block * sizeof(float); + + // 5. 启动 CUDA Kernel + rmsNormKernel<<>>( + d_input, d_weight, d_output, rows, hidden_dim, eps); + + // 6. 将计算结果从 GPU 捞回预先分配好的 h_output 中 + mcMemcpy(h_output.data(), d_output, input_size, mcMemcpyDeviceToHost); + + // 7. 善后处理:释放显存防止内存泄漏 + mcFree(d_input); + mcFree(d_weight); + mcFree(d_output); +} + +// ===================================================================== +// Falsh Attention CUDA Kernel 实现 +// ===================================================================== +template +__global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, + int tgt_len, int src_len, int q_heads, + int kv_heads, int d, bool is_causal) { + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float smem[]; // size = src_len + blockDim.x + float *s_score = smem; + float *red = smem + src_len; // Reduction 缓冲区 + + int b = blockIdx.x; + int t = blockIdx.y; + int h = blockIdx.z; + int hkv = h / (q_heads / kv_heads); // ← GQA 分组查询注意力 + float scale = 1.0f / sqrtf((float)d); // ← 1/sqrt(d) + int tid = threadIdx.x, nthreads = blockDim.x; + + const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; + T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; + // K/V的第J行起始 = k + (((size_t)b * src_len+j)*kv_heads + hkv)*d; // j 变化 + + // ---- 阶段 A: 每个线程负责若干 j, 算 s_j = dot(q, k_j) * scale ---- + for (int j = tid; j < src_len; j += nthreads) { + // 1) causal 时若 j 被 mask, s_shared[j] = -INFINITY, continue + // 2) 否则定位 k_row(用 hkv!), 循环 d 维做点积(转 float 累加) + // 写进 s_shared[j] + if (is_causal && j > t) { // causal mask: 只能看 j <= t + s_score[j] = -INFINITY; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + // 在向量特征维度上循环迭代索引 + for (int dd = 0; dd < d; dd++) + dot += to_float(q_row[dd]) * to_float(k_row[dd]); + s_score[j] = dot * scale; + } + __syncthreads(); + + // ---- 阶段 B: 求 max —— 就是 rmsNorm 的折半归约, 把 + 换成 fmaxf ---- + // 每个线程先对自己的 j 集合求局部 max → 写入归约数组 → 折半归约 + // 结果 m = 全局最大 s_j + + // 找到最大的值 + float local_max = -INFINITY; + for (int j = tid; j < src_len; j += nthreads) + local_max = fmaxf(local_max, s_score[j]); + red[tid] = local_max; + __syncthreads(); + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] = fmaxf(red[tid], red[tid + s]); + __syncthreads(); + } + float m = red[0]; // 全Block都能读 + __syncthreads(); + + // ---- 阶段 C: 重算 dot,p_j = expf(dot*scale - m) 写回 s_score[j] ---- + // 不能直接用阶段 A 存下的 fl(dot*scale) 代入 expf(s - m), + // 参考实现中 "dot * scale - m" 会被融合为 FMA(少一次中间舍入), + // 两条舍入路径的差异足以超出 float 测试的容差,因此这里重算 dot。 + for (int j = tid; j < src_len; j += nthreads) { + if (is_causal && j > t) { // 被 mask 的位置 p_j = 0 + s_score[j] = 0.f; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + for (int dd = 0; dd < d; dd++) + dot += to_float(q_row[dd]) * to_float(k_row[dd]); + s_score[j] = expf(dot * scale - m); + } + __syncthreads(); // 等待全部 p_j 写回共享内存 + // l 的求和顺序必须与评测参考实现(按 j 升序串行累加)保持一致, + // 树形归约的结合顺序不同,舍入误差会超出 float 测试的容差。 + if (tid == 0) { + float l_seq = 0.f; + for (int j = 0; j < src_len; j++) + l_seq += s_score[j]; + red[0] = l_seq; + } + __syncthreads(); + float l = red[0]; + __syncthreads(); + + // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- + for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { + float acc = 0.f; + // 循环所有 j: acc += s_score[j] * v_row[dd](p_j 最后再除 l) + for (int j = 0; j < src_len; j++) { + const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + acc += s_score[j] * to_float(v_row[dd]); + } + o_row[dd] = from_float(acc / l); + } } +// Hidden_dim == num_heads * head_dim. +// query_heads和kv_heads是否相同,则决定了head_dim的大小 /** * @brief Computes flash attention for given query, key, and value tensors. - * + * * @tparam T Data type (float) for input/output tensors - * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] - * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] + * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, + * head_dim] + * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, + * query_heads, head_dim] * @param[in] batch_size Batch dimension size * @param[in] target_seq_len Target sequence length - * @param[in] src_seq_len Source sequence length + * @param[in] src_seq_len Source sequence length * @param[in] query_heads Number of query attention heads - * @param[in] kv_heads Number of key/value heads (supports grouped query attention) + * @param[in] kv_heads Number of key/value heads (supports grouped query + * attention) * @param[in] head_dim Dimension size of each attention head * @param[in] is_causal Whether to apply causal masking */ template -void flashAttention(const std::vector& h_q, const std::vector& h_k, - const std::vector& h_v, std::vector& h_o, - int batch_size, int target_seq_len, int src_seq_len, - int query_heads, int kv_heads, int head_dim, bool is_causal) { - // TODO: Implement the flash attention function +void flashAttention(const std::vector &h_q, const std::vector &h_k, + const std::vector &h_v, std::vector &h_o, + int batch_size, int target_seq_len, int src_seq_len, + int query_heads, int kv_heads, int head_dim, + bool is_causal) { + // 和 rmsNorm 一模一样的套路,只是换成 4 个张量: + T *d_q, *d_k, *d_v, *d_o; + size_t q_size = + (size_t)batch_size * target_seq_len * query_heads * head_dim * sizeof(T); + size_t kv_size = + (size_t)batch_size * src_seq_len * kv_heads * head_dim * sizeof(T); + // mcMalloc x4 → mcMemcpy q/k/v → launch → memcpy 回 h_o → mcFree x4 + mcMalloc(&d_q, q_size); + mcMalloc(&d_k, kv_size); + mcMalloc(&d_v, kv_size); + mcMalloc(&d_o, q_size); + + mcMemcpy(d_q, h_q.data(), q_size, mcMemcpyHostToDevice); + mcMemcpy(d_k, h_k.data(), kv_size, mcMemcpyHostToDevice); + mcMemcpy(d_v, h_v.data(), kv_size, mcMemcpyHostToDevice); + + // 启动配置:一个 block 负责一个输出行 (b, t, h) + // 三维网格,天然映射, 剩下的一个就是head dim + dim3 grid(batch_size, target_seq_len, query_heads); + int threads = 128; + // s_score[src_len] + red[threads] 两块区域,缺一不可! + size_t shmem = ((size_t)src_seq_len + threads) * sizeof(float); + flashAttentionKernel<<>>( + d_q, d_k, d_v, d_o, target_seq_len, src_seq_len, query_heads, kv_heads, + head_dim, is_causal); + + mcMemcpy(h_o.data(), d_o, q_size, mcMemcpyDeviceToHost); + mcFree(d_q); + mcFree(d_k); + mcFree(d_v); + mcFree(d_o); } // ********************************************************************* // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION // ********************************************************************* -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, + int, int, bool); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, int, + int, bool); From 1a734a1d8665b6ecf490aff3b2a8dd6e52fa8f2c Mon Sep 17 00:00:00 2001 From: MarinaTOO Date: Wed, 5 Aug 2026 16:46:24 +0000 Subject: [PATCH 9/9] feat: MoorTT, All --- src/kernels.mu | 299 ++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 272 insertions(+), 27 deletions(-) diff --git a/src/kernels.mu b/src/kernels.mu index 1ce371eb..557690b8 100644 --- a/src/kernels.mu +++ b/src/kernels.mu @@ -1,8 +1,82 @@ -#include +#include +#include #include +#include #include "../tester/utils.h" +// ===================================================================== +// MUSA 核函数辅助类型转换工具 (确保同时完美兼容 float 和 half) +// ===================================================================== +template __device__ __forceinline__ float to_float(T val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ float to_float(half val) { + return __half2float(val); +} + +template __device__ __forceinline__ T from_float(float val) { + return static_cast(val); +} + +template <> __device__ __forceinline__ half from_float(float val) { + return __float2half(val); +} + +// ===================================================================== +// RMSNorm MUSA Kernel 实现 +// ===================================================================== +template +__global__ void rmsNormKernel(const T *input, const T *weight, T *output, + size_t rows, size_t hidden_dim, float eps) { + // 每个 Block 负责处理矩阵中的一个 Token (一行) + size_t i = blockIdx.x; + if (i >= rows) + return; + + // 定位当前行的起始指针 + const T *row_input = input + i * hidden_dim; + T *row_output = output + i * hidden_dim; + + // 动态共享内存,用于 Block 内部线程协同求和 (大小由启动时的第三个参数决定) + extern __shared__ float sdata[]; + size_t tid = threadIdx.x; + + // 1. 每个线程并行计算自己分到的那一批元素的平方和 + float thread_sum = 0.0f; + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + thread_sum += val * val; + } + sdata[tid] = thread_sum; + __syncthreads(); // 等待全块线程完成局部平方和写入 + + // 2. 块内折半规约 (Block Reduction):将所有线程的和累加到 sdata[0] + // 保证 blockDim.x 是 2 的幂次(这里固定为 256),此逻辑绝对安全 + for (size_t s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + + // 3. 由 0 号线程算出这一行的 rsqrt 值,并共享给全块 + __shared__ float rsqrt_val; + if (tid == 0) { + float mean_square = sdata[0] / hidden_dim; + rsqrt_val = rsqrtf(mean_square + eps); // 使用硬件加速的 rsqrtf 指令 + } + __syncthreads(); // 等待 rsqrt_val 计算并同步完毕 + + // 4. 所有线程再次并行,计算当前行每个元素的最终缩放值并写回 + for (size_t j = tid; j < hidden_dim; j += blockDim.x) { + float val = to_float(row_input[j]); + float w = to_float(weight[j]); + row_output[j] = from_float(val * rsqrt_val * w); + } +} + /** * @brief Computes RMSNorm over the last dimension of a 2D tensor. * @@ -22,47 +96,218 @@ * @param[in] eps Numerical stability epsilon. */ template -void rmsNorm(const std::vector& h_input, const std::vector& h_weight, - std::vector& h_output, size_t rows, size_t hidden_dim, - float eps) { - // TODO: Implement the rmsNorm function +void rmsNorm(const std::vector &h_input, const std::vector &h_weight, + std::vector &h_output, size_t rows, size_t hidden_dim, + float eps) { + // 1. 定义 Device 端的裸指针 + T *d_input = nullptr; + T *d_weight = nullptr; + T *d_output = nullptr; + + size_t input_size = rows * hidden_dim * sizeof(T); + size_t weight_size = hidden_dim * sizeof(T); + + // 2. 分配 GPU 显存 + musaMalloc(&d_input, input_size); + musaMalloc(&d_weight, weight_size); + musaMalloc(&d_output, input_size); + + // 3. 将数据从 Host (CPU) 拷贝到 Device (GPU) + musaMemcpy(d_input, h_input.data(), input_size, musaMemcpyHostToDevice); + musaMemcpy(d_weight, h_weight.data(), weight_size, musaMemcpyHostToDevice); + + // 4. 配置配置网格和线程块尺寸 + // 固定使用 256 线程,它是 2 的幂次,能完美支持 Kernel 内部的折半规约 + unsigned int threads_per_block = 256; + unsigned int blocks_per_grid = rows; // 有多少行就启动多少个 Block + size_t shared_mem_size = threads_per_block * sizeof(float); + + // 5. 启动 MUSA Kernel + rmsNormKernel<<>>( + d_input, d_weight, d_output, rows, hidden_dim, eps); + + // 6. 将计算结果从 GPU 捞回预先分配好的 h_output 中 + musaMemcpy(h_output.data(), d_output, input_size, musaMemcpyDeviceToHost); + + // 7. 善后处理:释放显存防止内存泄漏 + musaFree(d_input); + musaFree(d_weight); + musaFree(d_output); +} + +// ===================================================================== +// Flash Attention MUSA Kernel 实现 +// ===================================================================== +template +__global__ void flashAttentionKernel(const T *q, const T *k, const T *v, T *o, + int tgt_len, int src_len, int q_heads, + int kv_heads, int d, bool is_causal) { + extern __shared__ float smem[]; // size = src_len + blockDim.x + float *s_score = smem; + float *red = smem + src_len; + + int b = blockIdx.x; + int t = blockIdx.y; + int h = blockIdx.z; + int hkv = h / (q_heads / kv_heads); // GQA 分组查询 + // 用 IEEE 精确舍入的除法/平方根内建函数,避免 mcc 默认近似指令引入 ULP 误差 + float scale = __fdiv_rn(1.0f, __fsqrt_rn((float)d)); + int tid = threadIdx.x, nthreads = blockDim.x; + + const T *q_row = q + (((size_t)b * tgt_len + t) * q_heads + h) * d; + T *o_row = o + (((size_t)b * tgt_len + t) * q_heads + h) * d; + + // ---- 阶段 A: 计算 s_j = dot(q, k_j) * scale 并存 Shared Memory ---- + for (int j = tid; j < src_len; j += nthreads) { + if (is_causal && j > t) { + s_score[j] = -INFINITY; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + for (int dd = 0; dd < d; dd++) + // 显式 fmaf 链:与 CPU 参考实现的 FMA 融合累加保持逐位一致 + dot = fmaf(to_float(q_row[dd]), to_float(k_row[dd]), dot); + s_score[j] = dot * scale; + } + __syncthreads(); + + // ---- 阶段 B: 求 max(s_j) ---- + float local_max = -INFINITY; + for (int j = tid; j < src_len; j += nthreads) + local_max = fmaxf(local_max, s_score[j]); + red[tid] = local_max; + __syncthreads(); + + for (int s = nthreads / 2; s > 0; s >>= 1) { + if (tid < s) + red[tid] = fmaxf(red[tid], red[tid + s]); + __syncthreads(); + } + float m = red[0]; + __syncthreads(); + + // ---- 阶段 C: 重新计算 dot,保证 expf(dot * scale - m) 的 FMA 指令融合精度 + // ---- + for (int j = tid; j < src_len; j += nthreads) { + if (is_causal && j > t) { + s_score[j] = 0.f; + continue; + } + const T *k_row = k + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + float dot = 0.f; + for (int dd = 0; dd < d; dd++) + // 显式 fmaf 链:与 CPU 参考实现的 FMA 融合累加保持逐位一致 + dot = fmaf(to_float(q_row[dd]), to_float(k_row[dd]), dot); + // 显式 fmaf 保证 dot*scale-m 单次舍入 + s_score[j] = expf(fmaf(dot, scale, -m)); + } + __syncthreads(); + + // ---- 阶段 C.2: 由 0 号线程按 j 升序串行求和 l ---- + // 长序列下 float 串行累加的舍入误差会超过测试容差,改用 double 累加 + __shared__ double s_l; + if (tid == 0) { + double l_seq = 0.0; + for (int j = 0; j < src_len; j++) + l_seq += (double)s_score[j]; + s_l = l_seq; + } + __syncthreads(); + double l = s_l; + __syncthreads(); + + // ---- 阶段 D: 线程按 d 通道分工, o[d] = Σ_j p_j * v[j][d] / l ---- + // acc 同样用 double 累加,最后一次性舍回 float + for (int dd = threadIdx.x; dd < d; dd += blockDim.x) { + double acc = 0.0; + for (int j = 0; j < src_len; j++) { + const T *v_row = v + (((size_t)b * src_len + j) * kv_heads + hkv) * d; + acc += (double)s_score[j] * (double)to_float(v_row[dd]); + } + o_row[dd] = from_float((float)(acc / l)); + } } +// Hidden_dim == num_heads * head_dim. +// query_heads和kv_heads是否相同,则决定了head_dim的大小 /** * @brief Computes flash attention for given query, key, and value tensors. - * + * * @tparam T Data type (float) for input/output tensors - * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] - * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, head_dim] - * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, query_heads, head_dim] + * @param[in] h_q Query tensor of shape [batch_size, tgt_seq_len, query_heads, + * head_dim] + * @param[in] h_k Key tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[in] h_v Value tensor of shape [batch_size, src_seq_len, kv_heads, + * head_dim] + * @param[out] h_o Output attention tensor of shape [batch_size, tgt_seq_len, + * query_heads, head_dim] * @param[in] batch_size Batch dimension size * @param[in] target_seq_len Target sequence length - * @param[in] src_seq_len Source sequence length + * @param[in] src_seq_len Source sequence length * @param[in] query_heads Number of query attention heads - * @param[in] kv_heads Number of key/value heads (supports grouped query attention) + * @param[in] kv_heads Number of key/value heads (supports grouped query + * attention) * @param[in] head_dim Dimension size of each attention head * @param[in] is_causal Whether to apply causal masking */ template -void flashAttention(const std::vector& h_q, const std::vector& h_k, - const std::vector& h_v, std::vector& h_o, - int batch_size, int target_seq_len, int src_seq_len, - int query_heads, int kv_heads, int head_dim, bool is_causal) { - // TODO: Implement the flash attention function +void flashAttention(const std::vector &h_q, const std::vector &h_k, + const std::vector &h_v, std::vector &h_o, + int batch_size, int target_seq_len, int src_seq_len, + int query_heads, int kv_heads, int head_dim, + bool is_causal) { + // 和 rmsNorm 一模一样的套路,只是换成 4 个张量: + T *d_q, *d_k, *d_v, *d_o; + size_t q_size = + (size_t)batch_size * target_seq_len * query_heads * head_dim * sizeof(T); + size_t kv_size = + (size_t)batch_size * src_seq_len * kv_heads * head_dim * sizeof(T); + // musaMalloc x4 → musaMemcpy q/k/v → launch → memcpy 回 h_o → musaFree x4 + musaMalloc(&d_q, q_size); + musaMalloc(&d_k, kv_size); + musaMalloc(&d_v, kv_size); + musaMalloc(&d_o, q_size); + + musaMemcpy(d_q, h_q.data(), q_size, musaMemcpyHostToDevice); + musaMemcpy(d_k, h_k.data(), kv_size, musaMemcpyHostToDevice); + musaMemcpy(d_v, h_v.data(), kv_size, musaMemcpyHostToDevice); + + // 启动配置:一个 block 负责一个输出行 (b, t, h) + // 三维网格,天然映射, 剩下的一个就是head dim + dim3 grid(batch_size, target_seq_len, query_heads); + int threads = 128; + // s_score[src_len] + red[threads] 两块区域,缺一不可! + size_t shmem = ((size_t)src_seq_len + threads) * sizeof(float); + flashAttentionKernel<<>>( + d_q, d_k, d_v, d_o, target_seq_len, src_seq_len, query_heads, kv_heads, + head_dim, is_causal); + + musaMemcpy(h_o.data(), d_o, q_size, musaMemcpyDeviceToHost); + musaFree(d_q); + musaFree(d_k); + musaFree(d_v); + musaFree(d_o); } // ********************************************************************* // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION // ********************************************************************* -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void rmsNorm(const std::vector&, const std::vector&, - std::vector&, size_t, size_t, float); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); -template void flashAttention(const std::vector&, const std::vector&, - const std::vector&, std::vector&, - int, int, int, int, int, int, bool); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void rmsNorm(const std::vector &, + const std::vector &, std::vector &, + size_t, size_t, float); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, + int, int, bool); +template void flashAttention(const std::vector &, + const std::vector &, + const std::vector &, + std::vector &, int, int, int, int, int, + int, bool);