`flash_mla_with_kvcache` 算子输入规格

发布时间:2026/7/30 5:21:33
`flash_mla_with_kvcache` 算子输入规格
flash_mla_with_kvcache算子输入规格1. 基础配置参数值说明head(num_heads)64Q 的 head 数head_dim512Q 的每个 head 维度head_dim_v512V 的 head 维度softmax_scale1.0QK^T 后的缩放系数is_fp8_kvcacheTrueKV cache 走 FP8 packed 布局slide_window128主 KV cache 的滑动窗口大小2. Query 输入张量ShapeDtype说明q[token_num, 1, 64, 512]bfloat16最后两维分别为head、head_dim3. 主 KV Cache滑动窗口该算子接受两份 KV cache本节是第一份 —— 用滑动窗口的方式访问。张量ShapeDtype说明k_cache[2628, 256, 1, 584]uint8FP8 packed 布局stride(0)需按 TMA 对齐到 576 的倍数indices[token_num, 1, 128]int32每个 q token 在k_cache里关注的 KV 位置索引128 slide_windowtopk_length[token_num]int32每个 q token 实际关注的 KV token 数本用例全部填slide_window 128attn_sink[64]float32每个 head 一个 sink 值取值任意indices/topk_length只针对k_cache生效。4. Extra KV Cache扩展稀疏窗口第二份 KV cache用可变长度的 top-k 稀疏方式访问。张量ShapeDtype说明extra_k_cache[26279, 64, 1, 584]uint8同样 FP8 packedstride(0)亦需 576 对齐extra_indices_in_kvcache[token_num, 1, 512]int32每个 q token 在extra_k_cache里最多关注512个历史 KV token 的索引extra_topk_length[token_num]int32每个 q token 在extra_k_cache里实际关注的 KV token 数量≤ 512逐 token 可变5. 张量含义对照token_num一次 decode 的 query token 数batch × next_n。k_cache、extra_k_cache最后一维584FP8 数据 per-token scale/其他元数据的字节总长度。indices/extra_indices_in_kvcache的取值范围分别为[0, 2628*256)和[0, 26279*64)越界或-1视为 padding。topk_length、extra_topk_length描述有效长度用于在 topk 维度上做前缀掩码。