
简介
该用户还未填写简介
擅长的技术栈
可提供的服务
暂无可提供的服务
本文介绍了Qwen3-0.6B模型中KVCache的显存开销问题及其优化方案。传统KVCache管理方式存在显存浪费严重的问题,分页注意力通过预分配连续空间并分割为等长块来提高空间利用率。文章详细讲解了分页注意力的核心原理,包括内存不连续处理、Packed模式等,并展示了flash-attn库的使用方法以及如何用triton实现分页注意力。该方案能有效减少显存占用,适用于长序列和大规模请求场景。
中构建了一个GPT2,于是就想着照着这个拓展到更新的模型,比如Qwen3(虽然现在Qwen3.5都出了)这个系列的目标首先是自己用Torch搭建一个Qwen3模型,实现推理,然后未来会实现KVCache、手写CUDA算子等项目链接:https://gitee.com/a1483795887/qwen3_from_scratch。
前文用CUDA实现了RMSNorm这个算子,这个算子还算简单的,但还需要自己手动管理各种内容,还需要额外编译,使用pytorch调用时还需要先加载so,然后用动态链接库调用函数本文使用 triton 这个框架来用 python 来实现同等功能,在大大简化开发的情况下性能额不输 CUDA 的版本。
前文用CUDA实现了RMSNorm这个算子,这个算子还算简单的,但还需要自己手动管理各种内容,还需要额外编译,使用pytorch调用时还需要先加载so,然后用动态链接库调用函数本文使用 triton 这个框架来用 python 来实现同等功能,在大大简化开发的情况下性能额不输 CUDA 的版本。
在本章的目标就是快速搭建一个模型,能够推理即可,之后再自己手写组件。
在上一章中,我们搭建了一个Qwen3模型并且进行推理,但推理速度较慢,而且随着输出变长越来越慢,在GPU上还好,较短的输出还感受不出来,CPU上超过20个token就能明显感受到越来越慢推理速度慢的速度后面后手写算子解决,现在先解决这个越来越慢的问题,按现在的速度完全无法生成长文。
每个特征向量除以方均根进行归一化,再乘以一个 gamma 进行尺度缩放RMSNormxijxij∑jxij2nϵ⋅γjRMSNormxijn∑jxij2ϵxij⋅γj融合操作减少内存访问次数Warp shuffle 比共享内存归约更高效减少 Python 层调用开销。
自注意力是大模型的核心组件,也是计算最密集的两个部分之一(另一个是注意力后的MLP,这也是参数最多的部分)第五章将依次用Triton实现基础版本和FlashAttn,本节先来实现基础版本。
前文中写了自注意力的基础版本,成功跑起了推理,但这个实现有个问题,在实际使用中不得不面对:它申请了一块M×NM\times NM×N的注意力权重矩阵当MN变大时,比如prefix阶段计算一个长度为S的提示词,就需要申请S2S^2S2大小的注意力权重,这至少会产生三个开销额外的显存开销开辟、释放显存的开销读写全局内存的开销如果能避免这个中间矩阵申请就能同时降低显存开销和计算开销。
前文用triton完成了FlashAttentionV2,成功将显存开销削减到ONO(N)ON,但16位下耗时显著长于官方实现。本文将分步提升性能,最后耗时从官方的237%下降到112%








