十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-transformer 中的 FP8 KV Cache 更新算子 scatter_pa_kv_cache_with_k_scale 使用指南

CANN ops-transformer 中的 FP8 KV Cache 更新算子 scatter_pa_kv_cache_with_k_scale 使用指南 算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载导读scatter_pa_kv_cache_with_k_scale是 CANN ops-transformer 开源仓库cann/ops-transformer中基于torch_npu的cann_ops_transformer扩展接口用于在 PagedAttention 推理场景下将 FP8 格式的 key/value 张量以及对应的 per-token-head key_scale 反量化系数按slot_mapping指定的位置原地写入 KV Cache。本文从算子功能与计算公式、函数原型与参数约束、底层实现原理算子定义、Tiling、SIMT Kernel、torch 适配层以及单算子/图模式两种调用方式四个层面展开帮助开发者在 Ascend 950 系列硬件上正确、高效地使用该接口完成 FP8 KV Cache 的 scatter 更新。产品支持情况该接口由算子ScatterPaKvCacheWithKScale支撑其硬件适配范围由算子定义中注册的 AICore 配置决定参见 scatter_pa_kv_cache_with_k_scale_def.cpp产品支持情况如下产品系列支持情况Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持从源码看算子定义通过AICore().AddConfig(ascend950, ...)与AICore().AddConfig(ascend350, ...)注册了两类 AICore 配置与文档中仅 Ascend 950 系列支持的说明保持一致。功能说明与计算公式接口功能scatter_pa_kv_cache_with_k_scale是基于torch_npu的cann_ops_transformer扩展接口用于调用ScatterPaKvCacheWithKScale算子完成 PagedAttention 场景下 FP8 格式的 key/value 及其对应 key_scale 的 KV Cache 更新。该算子是一个原地inplace更新算子调用后key_cache、value_cache、key_scale_cache三个 cache 张量被直接改写无返回值。计算公式对于每个 tokeni ∈ [0, num_tokens)和每个头j ∈ [0, num_head)$$ block_idx \lfloor slot_mapping[i] / block_size \rfloor $$$$ block_offset slot_mapping[i] \bmod block_size $$$$ key_cache[block_idx][j][block_offset][:] key[i][j][:] $$$$ value_cache[block_idx][j][block_offset][:] value[i][j][:] $$$$ key_scale_cache[block_idx][j][block_offset][0] key_scale[i][j] $$其中num_tokens batch * seq_len即当前 step 需要写入 cache 的 token 总数block_idxslot_mapping 映射到的 block 索引block_offsetblock 内的偏移量约定 BBatch表示输入样本批量大小、SSeq-Length表示输入样本序列长度、NHead-Num表示多头数、DHead-Dim表示隐藏层最小的单元尺寸 head_dimnum_tokens B × Snum_blocks 表示 KV cache 分块总数block_size 表示每个分块包含的 token 数num_slots num_blocks × block_size 表示 cache 可容纳的总 token 数。与无 k_scale 变体的差异从接口命名与参数设计可以推断该算子是 scatter_pa_kv_cache如有的 FP8 扩展除 key/value 外额外携带 per-token-head 的key_scaleFP8 反量化系数及其 cachekey_scale_cache二者同为原地更新对象保证 FP8 量化的 K 在推理时能同步反量化。函数原型cann_ops_transformer.scatter_pa_kv_cache_with_k_scale( key, value, key_cache, value_cache, slot_mapping, key_scale, key_scale_cache, *, cache_layoutBNBD ) - None参数说明参数名参数类型可选/必选描述数据类型维度(shape)keyTensor必选待更新的 key 值当前 step 多个 token 的 key。不支持空 Tensor。float8_e5m2、float8_e4m3fn(num_tokens, num_head, k_head_size)valueTensor必选待更新的 value 值当前 step 多个 token 的 value。不支持空 Tensor。float8_e5m2、float8_e4m3fn(num_tokens, num_head, v_head_size)key_cacheTensor必选需要更新的 key cache当前 layer 的 key cache。不支持空 Tensor。与 key 保持一致(num_blocks, num_head, block_size, k_head_size)value_cacheTensor必选需要更新的 value cache当前 layer 的 value cache。不支持空 Tensor。与 value 保持一致(num_blocks, num_head, block_size, v_head_size)slot_mappingTensor必选每个 token key 或 value 在 cache 中的存储偏移。不支持空 Tensor。int32、int64(num_tokens,)key_scaleTensor必选待更新的 key scale 值当前 step 多个 token 的 key scale尾轴可以不连续。不支持空 Tensor。float(num_tokens, num_head)key_scale_cacheTensor必选需要更新的 key scale cache当前 layer 的 key scale cache最后一维为 1尾轴必须连续。不支持空 Tensor。float(num_blocks, num_head, block_size, 1)cache_layoutstr可选表示 key_cache 和 value_cache 的内存排布格式。当传 BNBD 时表示格式为 [num_blocks, num_head, block_size, head_size]。默认值为 BNBD。STRING-dtype 组合说明从算子定义源码scatter_pa_kv_cache_with_k_scale_def.cpp的注释与DataType注册顺序可见实际支持 4 种 dtype 组合float8_e5m2 int64、float8_e4m3fn int64、float8_e5m2 int32、float8_e4m3fn int32。即 FP8 权重类型两种与 slot_mapping 索引类型两种自由组合。返回值说明无返回值。该接口为原地更新接口调用后key_cache、value_cache、key_scale_cache会被原地更新无需接收返回值。约束说明声明参数 slot_mapping 属于 tensor。由于算子在 Tiling 阶段无法获取 tensor 的具体数值tiling 侧不对值进行校验正确性需要用户自行保证。若传入非法值会触发未定义行为精度问题、非法内存访问导致的程序崩溃等。该接口支持推理场景下使用。该接口支持单算子模式和图模式torchair调用。key、value、key_cache、value_cache 的数据类型必须一致且必须为 float8_e5m2 或 float8_e4m3fn。key_scale 和 key_scale_cache 的数据类型必须为 float。key 和 value 的前两维 shape 必须相同。slot_mapping 的取值范围为 [0, num_blocks*block_size-1]且 slot_mapping 内的元素值保证不重复重复时不保证正确性。key_scale 是两维 tensorshape 为 [num_tokens, num_head]尾轴可以不连续。key_scale_cache 是四维 tensorshape 为 [num_blocks, num_head, block_size, 1]最后一维必须为 1尾轴必须连续。num_tokens 表示当前需要更新到 cache 中的 token 数量num_tokens batch * seq_len。num_blocks 表示 KV cache 分块的总数block_size 表示每个分块包含的 token 数。num_head 表示注意力头数k_head_size 和 v_head_size 分别表示 key 和 value 的头维度大小。约束的源码级印证上述 dtype 与 shape 约束并非只在文档层面Tiling 阶段有完整的运行时校验逻辑。例如 scatter_pa_kv_cache_with_k_scale_tiling.cpp 中的ValidateDtype会校验 key/value/key_cache/value_cache 四者 dtype 一致且为 FP8 类型、slot_mapping 为 INT32/INT64、key_scale/key_scale_cache 为 FLOATValidateShape同文件 L216-L336会逐一校验各张量维度数与关键维度数值如 cache 的 shape[1] 必须等于 num_head、shape[3] 必须等于对应 head_size、key_scale_cache shape[3] 必须为 1 等校验失败会返回GRAPH_FAILED。此外torch 适配层的 C 绑定csrc/scatter_pa_kv_cache_with_k_scale.cpp也会在 Python 侧调用前通过TORCH_CHECK做维度数与 dtype 的快速检查。确定性计算默认支持确定性计算。底层实现原理算子定义与 InferShapeop_host 层算子注册在 scatter_pa_kv_cache_with_k_scale_def.cpp 中通过OP_ADD(ScatterPaKvCacheWithKScale)完成注册。值得注意的实现细节7 个输入key/value/key_cache/value_cache/slot_mapping/key_scale/key_scale_cache全部为必选格式统一为FORMAT_ND且均声明AutoContiguous()输出key_cache/value_cache/key_scale_cache与对应输入同名即算子以输出与输入重名的方式实现原地更新语义属性cache_layout为可选属性默认值为字符串BNBDAICore 配置开启了DynamicShapeSupportFlag(true)支持动态 shape、DynamicRankSupportFlag(true)支持动态 rank与PrecisionReduceFlag(true)。InferShape 实现在 scatter_pa_kv_cache_with_k_scale_infershape.cpp 中三个输出的 shape 直接拷贝对应输入的 shape*outputKeyCacheShape *inputKeyCacheShape等输出 dtype 与对应输入一致。Tiling 策略op_host 层Tiling 实现在 scatter_pa_kv_cache_with_k_scale_tiling.cpp 中核心流程为获取平台信息AIV 核数、UB 大小→ 计算 workspace 大小 →ValidateDtype→ExtractTensorParams提取各输入 shape 与 stride含 view 输入支持→ValidateShape→ 填充 Tiling 数据 → 设置 block 维度与 tiling key。Tiling 数据通过ScatterPaKvCacheWithKScaleTilingData结构体传递定义见 scatter_pa_kv_cache_with_k_scale_tiling_data.h包含 numTokens/numHead/kHeadSize/vHeadSize/numBlocks/blockSize/maxSlot 及各张量 stride 数组核数分配needCoreNum max(1, min(numTokens, coreNum))即以 token 数为粒度在 AIV 核间切分然后context-SetBlockDim(tiling-needCoreNum)两种场景模式SetTilingKeyByScene同文件 L518-L527根据kHeadSize vHeadSize选择SCATTER_KV_CACHE_SCENE_SPECIALIZED特化或SCATTER_KV_CACHE_SCENE_GENERALIZED泛化tiling key驱动 kernel 选择不同执行路径UB 内存按ubSize - DCACHE_SIZEDCACHE 默认 32KB设定本地内存大小。Kernel 实现op_kernel 层Kernel 入口在 scatter_pa_kv_cache_with_k_scale_apt.cpp通过模板NsScatterPaKvCacheWithKScale::ProcessDTYPE_KEY, DTYPE_SLOT_MAPPING, schMode(...)实例化其中DTYPE_KEYFP8 类型与DTYPE_SLOT_MAPPINGint32/int64由框架依据 def.cpp 自动生成schMode由 tiling key 决定。核心计算逻辑位于 SIMT 实现 scatter_pa_kv_cache_with_k_scale_simt.h其设计要点包括四个 SIMT Vector FunctionScatterKvKVVfkeyvalue 合并写入仅特化模式使用、ScatterKvKeyVf、ScatterKvValueVf、ScatterKvKeyScaleVfscale 写入。特化模式kHeadSize vHeadSize用 1 个 VF 同时搬运 key/value泛化模式拆分为 3 个 VF边界防御每个 VF 内部都会检查slot 0 || slot maxSlot并continue跳过避免越界写对应 Python 层注释超出范围的 slot 会被跳过不更新对应 cache索引类型自适应Process中根据total numTokens * numHead * max(kHeadSize, vHeadSize)是否超过 INT32_MAX 选择 32 位或 64 位索引线程数也随之调整——32 位索引用 2048 线程64 位索引用 1024 线程64 位索引占用更多寄存器需减少线程数以保持寄存器预算除法优化block_idx/block_offset 的除法通过GetUintDivMagicAndShift生成的 magic number 与移位完成避免整数除法开销计算结束后通过SetFlag/WaitFlag的 V→S 事件同步保证可见性。torch 适配层torch_extension 层OpBuilder 注册ScatterPaKvCacheWithKScaleOpBuilderscatter_pa_kv_cache_with_k_scale.py声明算子 schema 为(Tensor key, Tensor value, Tensor(a!) key_cache, Tensor(b!) value_cache, Tensor slot_mapping, Tensor key_scale, Tensor(c!) key_scale_cache, *, str cache_layoutBNBD) - ()其中(a!)/(b!)/(c!)标注了三个原地更新的别名参数Meta 实现shape/dtype 推导返回None与无返回值语义一致C 绑定csrc/scatter_pa_kv_cache_with_k_scale.cpp 通过ACLNN_CMD(aclnnScatterPaKvCacheWithKScale, ...)最终下发到 ACLNN 接口执行图模式转换graph_convert_scatter_pa_kv_cache_with_k_scale.py 通过register_fx_node_ge_converter将torch.ops.cann_ops_transformer.scatter_pa_kv_cache_with_k_scale.default转换为 GE 算子ge.ScatterPaKvCacheWithKScale(...)支撑 torchair 图模式执行。调用示例单算子模式调用import torch import torch_npu import cann_ops_transformer torch_npu.npu.set_device(0) # 形状定义 num_tokens 4 # 本次需要写入的token数量 num_head 8 # 注意力头数 k_head_size 128 # key头维度 v_head_size 128 # value头维度 num_blocks 2 # KV cache分块总数 block_size 16 # 每个分块包含的token数 # FP8 dtypefloat8_e5m2 与float8_e4m3fn均支持此处以e4m3fn为例 kv_dtype torch.float8_e4m3fn # 构造输入key/value为待写入的新数据 key torch.randn(num_tokens, num_head, k_head_size, dtypetorch.float32, devicenpu).to(kv_dtype) value torch.randn(num_tokens, num_head, v_head_size, dtypetorch.float32, devicenpu).to(kv_dtype) # 构造KV cache被inplace更新的目标初始置0便于校验 key_cache torch.zeros(num_blocks, num_head, block_size, k_head_size, dtypekv_dtype, devicenpu) value_cache torch.zeros(num_blocks, num_head, block_size, v_head_size, dtypekv_dtype, devicenpu) # 构造slot_mapping每个token在cache中的偏移取值范围 [0, num_blocks*block_size-1] slot_mapping torch.tensor([0, 1, 16, 17], dtypetorch.int32, devicenpu) # 构造key_scale及其cacheper-token-head的FP8反量化scale key_scale torch.randn(num_tokens, num_head, dtypetorch.float32, devicenpu) key_scale_cache torch.zeros(num_blocks, num_head, block_size, 1, dtypetorch.float32, devicenpu) # 调用算子将key/value/key_scale按slot_mapping写入cache原地更新无返回值 cann_ops_transformer.scatter_pa_kv_cache_with_k_scale( key, value, key_cache, value_cache, slot_mapping, key_scale, key_scale_cache, cache_layoutBNBD, ) torch_npu.npu.synchronize() print(key_cache.shape, key_cache.dtype) print(value_cache.shape, value_cache.dtype) print(key_scale_cache.shape, key_scale_cache.dtype)上例中slot_mapping [0, 1, 16, 17]结合block_size 16可验证token 0/1 落入 block 0 的 offset 0/1token 2/3 落入 block 1 的 offset 0/1与计算公式block_idx slot_mapping // block_size、block_offset slot_mapping % block_size完全对应。图模式torchair调用import torch import torch_npu import torch.nn as nn import torchair from torchair.configs.compiler_config import CompilerConfig import cann_ops_transformer torch_npu.npu.set_device(0) # 形状定义 num_tokens 4 # 本次需要写入的token数量 num_head 8 # 注意力头数 k_head_size 128 # key头维度 v_head_size 128 # value头维度 num_blocks 2 # KV cache分块总数 block_size 16 # 每个分块包含的token数 # FP8 dtypefloat8_e5m2 与float8_e4m3fn均支持此处以e4m3fn为例 kv_dtype torch.float8_e4m3fn # 构造输入key/value为待写入的新数据 key torch.randn(num_tokens, num_head, k_head_size, dtypetorch.float32, devicenpu).to(kv_dtype) value torch.randn(num_tokens, num_head, v_head_size, dtypetorch.float32, devicenpu).to(kv_dtype) # 构造KV cache被inplace更新的目标初始置0便于校验 key_cache torch.zeros(num_blocks, num_head, block_size, k_head_size, dtypekv_dtype, devicenpu) value_cache torch.zeros(num_blocks, num_head, block_size, v_head_size, dtypekv_dtype, devicenpu) # 构造slot_mapping每个token在cache中的偏移取值范围 [0, num_blocks*block_size-1] slot_mapping torch.tensor([0, 1, 16, 17], dtypetorch.int32, devicenpu) # 构造key_scale及其cacheper-token-head的FP8反量化scale key_scale torch.randn(num_tokens, num_head, dtypetorch.float32, devicenpu) key_scale_cache torch.zeros(num_blocks, num_head, block_size, 1, dtypetorch.float32, devicenpu) class ScatterPaKvCacheWithKScaleNetwork(nn.Module): def __init__(self): super(ScatterPaKvCacheWithKScaleNetwork, self).__init__() torch._dynamo.disable def forward(self, key, value, key_cache, value_cache, slot_mapping, key_scale, key_scale_cache, cache_layoutBNBD): torch.ops.cann_ops_transformer.scatter_pa_kv_cache_with_k_scale( key, value, key_cache, value_cache, slot_mapping, key_scale, key_scale_cache, cache_layoutcache_layout, ) return key_cache, value_cache, key_scale_cache config CompilerConfig() config.mode reduce-overhead npu_backend torchair.get_npu_backend(compiler_configconfig) torch._dynamo.reset() npu_mode torch.compile(ScatterPaKvCacheWithKScaleNetwork(), backendnpu_backend, dynamicFalse) key_cache, value_cache, key_scale_cache npu_mode( key, value, key_cache, value_cache, slot_mapping, key_scale, key_scale_cache, cache_layoutBNBD) print(key_cache.shape, key_cache.dtype) print(value_cache.shape, value_cache.dtype) print(key_scale_cache.shape, key_scale_cache.dtype)使用建议与注意事项运行环境前提该接口依赖torch_npu与cann_ops_transformer扩展包且仅在 Ascend 950 系列产品上受支持在 A2/A3 等其他系列产品上调用会因算子未注册而失败。确保 slot_mapping 合法且无重复Tiling 阶段无法读取 tensor 数值算子只对越界 slot 做跳过防御对重复 slot不做正确性保证重复写入同一位置时以最后一次写入为准、结果不确定。dtype 一致性key/value/key_cache/value_cache 四者必须同为float8_e5m2或同为float8_e4m3fnslot_mapping 用torch.int32或torch.int64key_scale/key_scale_cache 用torch.float32torch 层与 Tiling 层均会做校验。cache_layout 语义目前仅支持BNBD一种排布[num_blocks, num_head, block_size, head_size]传入其他取值时行为由算子实现决定建议始终显式使用默认值。推理场景该接口面向 PagedAttention 推理中 KV Cache 的增量写入每次调用仅写入当前 step 的 num_tokens 个 token配合 decode 阶段的 token 追加使用。相关资源接口文档torchapi_scatter_pa_kv_cache_with_k_scale.mdACLNN 文档aclnnScatterPaKvCacheWithKScale.md算子定义与 InferShapescatter_pa_kv_cache_with_k_scale_def.cpp、scatter_pa_kv_cache_with_k_scale_infershape.cppTiling 实现scatter_pa_kv_cache_with_k_scale_tiling.cppKernel 实现scatter_pa_kv_cache_with_k_scale_apt.cpp、scatter_pa_kv_cache_with_k_scale_simt.htorch 适配层scatter_pa_kv_cache_with_k_scale.py、csrc/scatter_pa_kv_cache_with_k_scale.cpp、graph_convert_scatter_pa_kv_cache_with_k_scale.py示例与测试examples/、tests/赞分享算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载相关推荐CANN ops-transformer 算子解析aclnnScatterPaKvCacheWithKScale 的 FP8 KV Cache 与 K-Scale 原地更新实现CANN ops transformer 算子解析aclnnScatterPaKvCacheWithKScale 的 FP8 KV Cache 与 K Sca算子库人工智能深度学习AscendScatterPaKvCacheWithKScale 算子全解析CANN ops-transformer 中 FP8 量化 KV Cache 的分块写入与 scale 更新ScatterPaKvCacheWithKScale 算子全解析CANN ops transformer 中 FP8 量化 KV Cache 的分块写入与 s算子库人工智能深度学习Ascendcurl 的 --dns-ipv6-addr指定 IPv6 源地址发起 DNS 请求的完整指南curl 的 dns ipv6 addr 指定 IPv6 源地址发起 DNS 请求的完整指南 dns ipv6 addr 是 curl 在构建了 c ares算子库人工智能深度学习Ascend上一篇惊艳全场nodeppt自定义过渡效果完全指南从基础到高级创意设计下一篇30秒搞定国家中小学智慧教育平台电子课本一键下载神器创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表