由于不喜欢的工作导致不喜欢的答辩的折磨以及跟QQ的折磨导致无法正常阅读所以需要这个。

test_triton_unified_attn

从测试接口进去看细节或许确实会很好。

image-20260629222935206

这部分是参数声明和pytest的mark,以及设置随机数种子

这些参数我暂时还不知道什么意思,往后看,试试能不能把所有参数的含义理解。

输入[a, b]作为pytest,那么pytest分别会用a和b调用这个函数

我们接下来的分析就以第一个数值为例子,以及如果是空的话就单独开个分析分支。

image-20260629231314369

如图,首先是

seq_lens,它包含的是query_lens和kv_lens,可以发现,这个的batch_size=3

然后是多头,query_heads和kv_heads都设置成4了。

image-20260629231440944

然后是这里,初始化各种tensor,

query: (这个batch的总query长度也就是总的n,query头,头大小)

kv_cache一起讲: block attention需要将kv cache分block,所以这里是这样的。形状是(block数, block大小, kv头数,kv头大小)

cu_query_lens是一个对query_lens的inclusive scan,也就是前缀和。

然后对每个sequence,以最大的blocknum分配block_tables,注意这里,豆包欺骗了我,block_tables形状应该是(3,83),只是随机从0-32768中取。这个确实更符合pagedattention

也就是说对每个sequence都有一个页表。

image-20260629231855563

这边属于量化和3d的路径,暂时先不管他。因为关于这些的参数似乎可以取0所以管不了。后续分析量化路径再说

image-20260629233315716

然后就是调用unified_attention了

unified_attention

image-20260629233602942

这是这个函数,前面是一些assert,还有一些用不到的东西,931-944跳过。

946-951是从qkv里面获取一些attention形状参数

image-20260629234032465

以下这部分是Q的分块计算逻辑,首先上面通过query_head数/kv_head数计算好了每个kv要处理多少query。

image-20260629234328953

i[ceil(query_len[i]/BLOCK_Q)]

说是把query_len_sum / BLOCK_Q进行Q的分块吧,然后他进行了一次变换对cpu友好(似乎说是这个计算在cpu上的)。

image-20260630000143623

这里是tile_size,具体是个啥估计得看算子了。

260630 0:02:35

预计早上起来更新

kernel_unified_attention_2d

image-20260630093401703

运行2d的条件有这些,这里满足的只有max_seqlen_q > 1 或3D=0的时候num_seqs=3 >3D=0. 注意到batch_invariant是一个值得研究的参数。

image-20260630093722959

image-20260630094842566

在这里注意到源码kv_scale标注错了。别的都类似。

image-20260630101900143

其中,find_seq_idx为一个二分,也就是在前缀和数组里面找当前block落在哪段上。

image-20260630102045684

image-20260630102734031

以不分块的为例子,发现确实是这样。不过我认为要深究可能也可以去形式化一下的。以及加mid的原因是之前说过的那个变换形式。

image-20260630104045884

image-20260630104244327

125行的这个计算也是这个原因。因为单个seq至少要占到一个BLOCK。这里以len_ptr作为index offset数组,计算当前Program的idx和边界。

 

image-20260630110948931

image-20260630112734477

这部分是计算Q边界、block_table坐标以及加载这个Q

image-20260630112655770

对这个例子num_queries_per_kv=1是这样。(红色区域为加载的Q)

认为154行不太可能出现0的情况, 所以一般只限制了batch_query和head

证明:

 

现在是260630 11:47:51,大体看完了,其实就是个flash attention,先吃饭,下午再细说。

image-20260630131834599

image-20260630132107956

这部分负责初始化一些临时数组,M,L,acc,context_len,上下文长度指的是当前query之前的所有query.

max_seq_prefix_len是到目前这个BLOCK(包含当前block)的前缀query长度,然后进行一下限界

image-20260630132252332

将前面的prefix按tile_size分tile,滑动窗口暂时不看吧。

image-20260630141941137

这里是加载KV,因为分页机制,所以这里通过页表进行转换之后的块地址索引。

我们知道四个stride:

那么对v的索引应该是这样的:physical_block_idx与seq_offset%BLOCK_SIZE共同决定是哪个范围的平面,offset_d选择这一块,最后是靠加kv_head_idx。所以根据内存连续性需求,应该是竖着的这一块面是连续的。

image-20260630143007564

在这张图里面,选择的V的seq_offset覆盖,其实大小就是TILE_SIZE

假设physical_block_idx都是同一个。 多个的时候很可能存在访存不连续。

而K就直接反过来了。同样可以用上面这张图进行解释。只不过TILE_SIZE变成了第二维。这是先选择了head_size再选择tile_size.也就是成为转置。

所以说这两句只是一个idx与形状变化结合的offset构造,具体访存还是得具体分析。

image-20260630144856486

这里就主要是加载和量化了。

到目前为止,三个矩阵的形状为

 

往后其实就是flash attention了,没什么可写的

image-20260630150943299

image-20260630150951839

image-20260630150958530

image-20260630151011292

这个相关的参考

image-20260630151053153

From arxiv2205.14135

img

From CS149GPT assignment.